payroll/crud.py
Padreug 8a10c3cacc chore: annotate the aggregate row so mypy passes
`fetchone` without a model leaves its TModel unbound, so the COUNT(*) row
needs an explicit type. Completes the black + ruff + mypy pipeline the
Makefile declares.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_018jy52j9GRZ6XKa1Zt21LLj
2026-08-31 13:56:45 +02:00

162 lines
5.1 KiB
Python

"""Payroll CRUD.
Thin by design: the interesting logic (when a period falls due, whether a
payout may run) lives in services.py, and this module only owns row
lifecycle. The one rule enforced here is that `updated_at` is stamped on
every write, so "when did this contract last change" is answerable without
a separate audit trail.
"""
from datetime import datetime, timezone
from lnbits.db import Database
from lnbits.helpers import urlsafe_short_hash
from .models import Contract, ContractStatus, CreateContract, Payout, PayoutStatus
db = Database("ext_payroll")
# ---------------------------------------------------------------------------
# Contracts
# ---------------------------------------------------------------------------
async def create_contract(
data: CreateContract, employee_username: str = ""
) -> Contract:
contract = Contract(
**data.dict(),
id=urlsafe_short_hash()[:8],
employee_username=employee_username,
)
await db.insert("payroll.contracts", contract)
return contract
async def get_contract(contract_id: str) -> Contract | None:
return await db.fetchone(
"SELECT * FROM payroll.contracts WHERE id = :id",
{"id": contract_id},
Contract,
)
async def get_contracts() -> list[Contract]:
return await db.fetchall(
"SELECT * FROM payroll.contracts ORDER BY created_at DESC", model=Contract
)
async def get_contracts_by_status(status: ContractStatus) -> list[Contract]:
return await db.fetchall(
"SELECT * FROM payroll.contracts WHERE status = :status ORDER BY created_at",
{"status": status.value},
Contract,
)
async def get_contracts_for_employee(employee_id: str) -> list[Contract]:
return await db.fetchall(
"SELECT * FROM payroll.contracts WHERE employee_id = :eid "
"ORDER BY created_at DESC",
{"eid": employee_id},
Contract,
)
async def get_contracts_for_wallet(wallet_id: str) -> list[Contract]:
"""Contracts paying *into* a given wallet — the employee-side view, keyed
on the wallet whose key authenticated the request rather than on an
account id."""
return await db.fetchall(
"SELECT * FROM payroll.contracts WHERE employee_wallet = :wid "
"ORDER BY created_at DESC",
{"wid": wallet_id},
Contract,
)
async def update_contract(contract: Contract) -> Contract:
contract.updated_at = datetime.now(timezone.utc)
await db.update("payroll.contracts", contract)
return contract
async def delete_contract(contract_id: str) -> None:
await db.execute(
"DELETE FROM payroll.contracts WHERE id = :id", {"id": contract_id}
)
# ---------------------------------------------------------------------------
# Payout ledger
# ---------------------------------------------------------------------------
async def create_payout(payout: Payout) -> Payout:
await db.insert("payroll.payouts", payout)
return payout
async def get_payouts(
contract_id: str | None = None,
employee_wallet: str | None = None,
status: PayoutStatus | None = None,
since: str | None = None,
until: str | None = None,
limit: int = 200,
) -> list[Payout]:
"""Ledger rows, newest first.
`since`/`until` bound the **payday**, not the row's creation time: an
accounting period is about which paydays fall in it, and a payday that
was retried for three days would otherwise land in the wrong month.
Both are inclusive `YYYY-MM-DD`, which compares correctly as a string
because that is the only format `payday` is ever written in.
"""
where = []
values: dict = {}
if contract_id:
where.append("contract_id = :cid")
values["cid"] = contract_id
if employee_wallet:
where.append("employee_wallet = :wid")
values["wid"] = employee_wallet
if status:
where.append("status = :status")
values["status"] = status.value
if since:
where.append("payday >= :since")
values["since"] = since
if until:
where.append("payday <= :until")
values["until"] = until
clause = f"WHERE {' AND '.join(where)}" if where else ""
return await db.fetchall(
f"SELECT * FROM payroll.payouts {clause} "
"ORDER BY created_at DESC LIMIT :limit",
{**values, "limit": limit},
Payout,
)
async def count_period_failures(contract_id: str, period_index: int) -> int:
"""How many times this exact period has already failed.
Drives the bounded-retry cap. Counted from the ledger rather than from a
counter on the contract so the number survives a restart and stays
auditable — the rows that produced it are right there.
"""
# No model: fetchone's TModel is unbound for a bare aggregate, so the
# row comes back as a plain mapping and needs its own annotation.
row: dict | None = await db.fetchone(
"SELECT COUNT(*) AS n FROM payroll.payouts "
"WHERE contract_id = :cid AND period_index = :idx AND status = :status",
{
"cid": contract_id,
"idx": period_index,
"status": PayoutStatus.failed.value,
},
)
return int(row["n"]) if row else 0