diff --git a/__init__.py b/__init__.py index e7c8aaf..521176b 100644 --- a/__init__.py +++ b/__init__.py @@ -1,7 +1,9 @@ +import asyncio + from fastapi import APIRouter +from loguru import logger from .crud import db -from .views import payroll_generic_router from .views_api import payroll_api_router payroll_static_files = [ @@ -12,12 +14,33 @@ payroll_static_files = [ ] payroll_ext: APIRouter = APIRouter(prefix="/payroll", tags=["payroll"]) -payroll_ext.include_router(payroll_generic_router) payroll_ext.include_router(payroll_api_router) +scheduled_tasks: list[asyncio.Task] = [] + + +def payroll_stop(): + for task in scheduled_tasks: + try: + task.cancel() + except Exception as ex: + logger.warning(ex) + + +def payroll_start(): + from lnbits.tasks import create_permanent_unique_task + + from .tasks import scheduler_loop + + scheduled_tasks.append( + create_permanent_unique_task("ext_payroll_scheduler", scheduler_loop) + ) + __all__ = [ "db", "payroll_ext", + "payroll_start", "payroll_static_files", + "payroll_stop", ] diff --git a/docs/operations.md b/docs/operations.md new file mode 100644 index 0000000..93e7c74 --- /dev/null +++ b/docs/operations.md @@ -0,0 +1,81 @@ +# Running payroll + +## The scheduler + +One permanent task (`ext_payroll_scheduler`), started with the extension. +It waits ~20s after boot, then every 5 minutes asks each **active** +contract whether it owes a payday. + +Minutes, not seconds, on purpose: a payday is a calendar event, so paying +a few minutes into the day costs nothing, and a tight loop would only +multiply log noise while a contract cannot be funded. The post-boot pass +means an instance that was down over a payday catches up on restart rather +than waiting a full interval. + +Days are **UTC**. An instance whose operators think in a far-eastern or +far-western timezone will see paydays land on what is locally the previous +or next day. + +## How one period is paid + +``` +create_invoice(employee_wallet, amount, currency, internal=True) + │ ← this is also where a fiat contract is priced in sats + ▼ + invoice.amount (msat) ← canonical for this period, recorded as-is + │ + ▼ +pay_invoice(source_wallet, invoice.bolt11) +``` + +A plain internal LNbits wallet-to-wallet transfer. The invoice carries +`extra = {tag: "payroll", contract_id, period, payday}`, so a payroll +transfer is identifiable from the payments list on either wallet. + +The sat amount is derived exactly once, by LNbits' own invoice pricing, and +everything downstream reuses that number. Re-deriving it from +`amount × rate` at any later point would drift — FX moves between quote and +settlement, and rounding accumulates over a year of paydays. + +If a payout fails after the invoice exists, an unpaid internal invoice is +left on the employee's wallet. It was never settled and expires on its own; +that is the deliberate price of not pricing the period a second time just +to run a balance check. + +## What a failure does + +A failed period does **not** advance `periods_done`, so the next tick +retries the same payday. Later periods do not jump the queue: the backlog +halts at the first failure so paydays settle in order. + +The common failure is an underfunded source wallet, and the log line says +so explicitly with both figures: + +``` +payroll: contract a1b2c3d4 period 3 failed: insufficient balance in source +wallet: 12000 sat available, 80000 sat required +``` + +## Back-dated start dates + +Creating a contract with a start date in the past is normally an +*anchoring* choice — "we pay on the 1st" — not a request for back-pay. +Periods whose payday fell before the contract row existed are therefore +**skipped**, and skipping still consumes the period, so the schedule stays +aligned to the anchor. + +Tick `backfill` on the contract to pay them instead. + +Either way a single tick will not fire more than `MAX_CATCH_UP_PERIODS` +(12) periods for one contract, so a mistyped start date cannot turn into a +hundred transfers. + +## Double-payment guard + +Each contract has an in-process `asyncio.Lock`, and `run_due_periods` +re-reads the contract row under it. The scheduler is a single task, so the +lock exists for the off-cycle paths that can land mid-tick. + +This is per-process. Two LNbits processes sharing one database would not +be serialised by it — payroll assumes the single-writer deployment LNbits +itself assumes. diff --git a/services.py b/services.py new file mode 100644 index 0000000..72c3376 --- /dev/null +++ b/services.py @@ -0,0 +1,363 @@ +"""Payroll scheduling and payout. + +Two separable things live here, and the split is deliberate: + +* **Schedule math** (`occurrence_on`, `due_period_indices`) is pure — no + DB, no wallets, no clock of its own. Paydays are the part of payroll that + is easy to get subtly wrong and expensive to get wrong in production, so + it is kept testable in isolation. +* **Payout** (`run_due_periods`) is the effectful half: it turns a due + period into an actual wallet-to-wallet transfer. + +Everything is anchored on `start_date`. The *n*-th payday is a function of +the start date and `n` alone — never of the previous payday, and never of +when the scheduler happened to run. That is what stops a missed tick, a +restart, or a short month from shifting the whole remaining schedule. +""" + +import asyncio +from calendar import monthrange +from dataclasses import dataclass +from datetime import date, datetime, timedelta, timezone + +from lnbits.core.crud import get_wallet +from lnbits.core.services import create_invoice, pay_invoice +from loguru import logger + +from . import crud +from .models import Contract, ContractStatus, Frequency + +# Frequencies that are an exact number of days: no calendar involved, so no +# clamping is possible or needed. +_DAY_STEP = { + Frequency.daily: 1, + Frequency.weekly: 7, + Frequency.biweekly: 14, +} + +# Frequencies that step whole months and therefore clamp into short months. +_MONTH_STEP = { + Frequency.monthly: 1, + Frequency.quarterly: 3, + Frequency.yearly: 12, +} + +# A single tick will not fire more than this many periods for one contract. +# Protects against a mistyped start date turning into a hundred transfers. +MAX_CATCH_UP_PERIODS = 12 + + +# --------------------------------------------------------------------------- +# Schedule math (pure) +# --------------------------------------------------------------------------- + + +def add_months(anchor: date, months: int) -> date: + """`anchor` shifted by `months`, clamped to the target month's last day. + + 31 Jan + 1 month is 28 Feb (29 in a leap year), because 31 Feb does not + exist. Note this is applied to the *anchor*, not to a previous result: + 31 Jan + 2 months is 31 Mar, not 28 Mar. Payroll anchored on month-end + should keep paying on month-end. + """ + total = anchor.year * 12 + (anchor.month - 1) + months + year, month_index = divmod(total, 12) + month = month_index + 1 + return date(year, month, min(anchor.day, monthrange(year, month)[1])) + + +def occurrence_on(start: date, frequency: Frequency, index: int) -> date: + """The payday for period `index`, counting the start date as period 0.""" + if frequency in _DAY_STEP: + return start + timedelta(days=_DAY_STEP[frequency] * index) + return add_months(start, _MONTH_STEP[frequency] * index) + + +def parse_start_date(contract: Contract) -> date: + return datetime.strptime(contract.start_date, "%Y-%m-%d").date() + + +def next_payday(contract: Contract) -> date | None: + """When this contract pays next, or None if it has no periods left.""" + remaining = contract.periods_remaining + if remaining is not None and remaining <= 0: + return None + return occurrence_on( + parse_start_date(contract), contract.frequency, contract.periods_done + ) + + +def upcoming_paydays(contract: Contract, count: int) -> list[date]: + """The next `count` paydays, truncated by any period cap.""" + start = parse_start_date(contract) + last = contract.total_periods if contract.total_periods is not None else None + indices = range(contract.periods_done, contract.periods_done + count) + return [ + occurrence_on(start, contract.frequency, i) + for i in indices + if last is None or i < last + ] + + +def due_period_indices(contract: Contract, today: date) -> list[int]: + """Every period index that is payable as of `today`, in order. + + Normally this is zero or one entry. It is a list because a contract can + legitimately have a backlog: a back-dated start date, or an instance + that was down over a payday. Returning them all lets the caller decide + what to do with each — `run_due_periods` pays or skips them one at a + time rather than collapsing a backlog into a single surprise transfer. + """ + if contract.status != ContractStatus.active: + return [] + + start = parse_start_date(contract) + cap = contract.total_periods + indices: list[int] = [] + index = contract.periods_done + while cap is None or index < cap: + if occurrence_on(start, contract.frequency, index) > today: + break + indices.append(index) + index += 1 + # A daily contract left unpaid for years would otherwise build an + # unbounded list. Anything beyond this is an operator problem, not + # something to silently drain in one tick. + if len(indices) >= MAX_CATCH_UP_PERIODS: + break + return indices + + +# --------------------------------------------------------------------------- +# Payout (effectful) +# --------------------------------------------------------------------------- + + +@dataclass +class PeriodOutcome: + """What happened to one period. `amount_msat` is the canonical settled + figure, taken from the invoice LNbits actually priced — never + recomputed from amount x rate.""" + + index: int + payday: date + status: str # "paid" | "skipped" | "failed" + detail: str = "" + amount_msat: int | None = None + payment_hash: str | None = None + + @property + def consumed(self) -> bool: + """Whether this period should advance the contract's position. + + Paid and skipped both consume the period. Failed does not — that is + what makes the next tick retry the same payday rather than dropping + it. + """ + return self.status in ("paid", "skipped") + + +# One lock per contract id. The scheduler is a single task, but an operator +# can trigger an off-cycle payout at any moment; without this, a manual run +# landing mid-tick could pay the same period twice. Double-paying is the +# worst thing this extension can do, so the guard is cheap insurance. +_contract_locks: dict[str, asyncio.Lock] = {} + + +def _lock_for(contract_id: str) -> asyncio.Lock: + lock = _contract_locks.get(contract_id) + if lock is None: + lock = asyncio.Lock() + _contract_locks[contract_id] = lock + return lock + + +def payout_memo(contract: Contract, payday: date) -> str: + label = contract.memo or contract.label or "Payroll" + return f"{label} — {payday.isoformat()}" + + +async def pay_period(contract: Contract, index: int, payday: date) -> PeriodOutcome: + """Move one period's money from the source wallet to the employee wallet. + + An internal invoice on the destination wallet, paid from the source + wallet — the standard LNbits wallet-to-wallet transfer. Creating the + invoice is also what prices a fiat contract in sats, so the amount is + converted exactly once, here, and the result is what gets recorded. + + A failure after the invoice exists leaves an unpaid internal invoice on + the employee's wallet. That is harmless — it expires on its own and was + never settled — and it is the price of not pre-computing the sat amount + a second time just to run a balance check. + """ + memo = payout_memo(contract, payday) + extra = { + "tag": "payroll", + "contract_id": contract.id, + "period": index, + "payday": payday.isoformat(), + } + + try: + invoice = await create_invoice( + wallet_id=contract.employee_wallet, + amount=contract.amount, + currency=contract.currency, + memo=memo, + internal=True, + extra=extra, + ) + except Exception as exc: + return PeriodOutcome( + index=index, + payday=payday, + status="failed", + detail=f"could not raise invoice: {exc}", + ) + + # invoice.amount is msat and is now the canonical figure for this period. + amount_msat = invoice.amount + + source = await get_wallet(contract.source_wallet) + if not source: + return PeriodOutcome( + index=index, + payday=payday, + status="failed", + detail="source wallet not found", + amount_msat=amount_msat, + ) + if source.balance_msat < amount_msat: + return PeriodOutcome( + index=index, + payday=payday, + status="failed", + detail=( + f"insufficient balance in source wallet: " + f"{source.balance_msat // 1000} sat available, " + f"{amount_msat // 1000} sat required" + ), + amount_msat=amount_msat, + ) + + try: + await pay_invoice( + wallet_id=contract.source_wallet, + payment_request=invoice.bolt11, + description=memo, + tag="payroll", + extra=extra, + ) + except Exception as exc: + return PeriodOutcome( + index=index, + payday=payday, + status="failed", + detail=f"payment failed: {exc}", + amount_msat=amount_msat, + payment_hash=invoice.payment_hash, + ) + + return PeriodOutcome( + index=index, + payday=payday, + status="paid", + amount_msat=amount_msat, + payment_hash=invoice.payment_hash, + ) + + +async def run_due_periods( + contract: Contract, today: date | None = None +) -> list[PeriodOutcome]: + """Settle every period this contract owes as of `today`. + + Stops at the first failure so a backlog cannot pay periods out of order: + if period 3 could not be funded, period 4 waits for it. Both will be + retried on the next tick, because a failed period does not advance the + contract's position. + """ + today = today or datetime.now(timezone.utc).date() + + async with _lock_for(contract.id): + # Re-read under the lock: a manual run may have moved the position + # since the caller loaded this row. + fresh = await crud.get_contract(contract.id) + if not fresh: + return [] + contract = fresh + + outcomes: list[PeriodOutcome] = [] + for index in due_period_indices(contract, today): + payday = occurrence_on( + parse_start_date(contract), contract.frequency, index + ) + outcome = await _settle(contract, index, payday) + outcomes.append(outcome) + if not outcome.consumed: + break + contract.periods_done = index + 1 + + if outcomes: + _maybe_complete(contract) + await crud.update_contract(contract) + + return outcomes + + +async def _settle(contract: Contract, index: int, payday: date) -> PeriodOutcome: + """Pay a period, or skip it as back-dated. + + A contract created with a start date in the past is normally an + *anchoring* choice — "we pay on the 1st" — not a request for back-pay. + Periods whose payday fell before the contract existed are therefore + skipped unless the operator explicitly asked to backfill, which keeps a + new contract from firing months of transfers on its first tick. + """ + if not contract.backfill and payday < contract.created_at.date(): + logger.info( + f"payroll: contract {contract.id} period {index} ({payday}) " + "predates the contract; skipping (backfill is off)" + ) + return PeriodOutcome( + index=index, + payday=payday, + status="skipped", + detail="payday predates the contract and backfill is off", + ) + + outcome = await pay_period(contract, index, payday) + if outcome.status == "paid": + logger.success( + f"payroll: contract {contract.id} period {index} paid " + f"{(outcome.amount_msat or 0) // 1000} sat to " + f"{contract.employee_wallet}" + ) + else: + logger.warning( + f"payroll: contract {contract.id} period {index} " + f"{outcome.status}: {outcome.detail}" + ) + return outcome + + +def _maybe_complete(contract: Contract) -> None: + if ( + contract.total_periods is not None + and contract.periods_done >= contract.total_periods + ): + contract.status = ContractStatus.completed + logger.info( + f"payroll: contract {contract.id} completed " + f"({contract.periods_done} periods)" + ) + + +async def tick(today: date | None = None) -> None: + """One scheduler pass over every active contract.""" + contracts = await crud.get_contracts_by_status(ContractStatus.active) + for contract in contracts: + try: + await run_due_periods(contract, today) + except Exception as exc: + logger.error(f"payroll: contract {contract.id} tick failed: {exc}") diff --git a/tasks.py b/tasks.py new file mode 100644 index 0000000..ceb0d48 --- /dev/null +++ b/tasks.py @@ -0,0 +1,36 @@ +"""The payroll scheduler. + +One permanent task that wakes up, asks every active contract whether it +owes a payday, and settles the ones that do. Deliberately dumb: all the +decisions live in services.py, and this file only owns the clock. + +The tick interval is minutes rather than seconds because a payday is a +calendar event — the cost of paying an hour into the day is nil, and a +tight loop would only multiply log noise when a contract cannot be funded. +A pass also runs shortly after startup so an instance that was down over a +payday catches up without waiting a full interval. +""" + +import asyncio + +from loguru import logger + +from . import services + +# How often to look for due paydays. +TICK_SECONDS = 300 + +# Let the funding source, DB and extension registry settle before the first +# pass — the first tick can move money, so it should not race startup. +STARTUP_DELAY_SECONDS = 20 + + +async def scheduler_loop(): + await asyncio.sleep(STARTUP_DELAY_SECONDS) + logger.info("payroll: scheduler started") + while True: + try: + await services.tick() + except Exception as exc: + logger.error(f"payroll: scheduler tick failed: {exc}") + await asyncio.sleep(TICK_SECONDS) diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..3bd7a28 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,43 @@ +"""Test fixtures for payroll. + +The schedule math is pure, so most tests need nothing but a `Contract` +instance — no LNbits DB, no wallets, no event loop. That is the point of +keeping `services` split into a pure half and an effectful half. +""" + +from datetime import datetime, timezone + +from ..models import Contract, ContractStatus, Frequency + + +def make_contract( + *, + contract_id: str = "c1", + start_date: str = "2026-01-15", + frequency: Frequency = Frequency.monthly, + total_periods: int | None = 12, + periods_done: int = 0, + status: ContractStatus = ContractStatus.active, + amount: float = 1000, + currency: str = "sat", + backfill: bool = False, + created_at: str = "2026-01-01", +) -> Contract: + return Contract( + id=contract_id, + employee_id="employee-account", + employee_username="alice", + employee_wallet="wallet-employee", + source_wallet="wallet-treasury", + amount=amount, + currency=currency, + frequency=frequency, + start_date=start_date, + total_periods=total_periods, + periods_done=periods_done, + status=status, + backfill=backfill, + created_at=datetime.strptime(created_at, "%Y-%m-%d").replace( + tzinfo=timezone.utc + ), + ) diff --git a/tests/test_payout.py b/tests/test_payout.py new file mode 100644 index 0000000..8a1399f --- /dev/null +++ b/tests/test_payout.py @@ -0,0 +1,132 @@ +"""Payout sequencing. + +`pay_period` itself talks to LNbits wallets, so it is stubbed here — what +these tests pin down is the surrounding contract: which periods get +attempted, in what order, and what a failure does to the contract's +position. Those are the invariants that decide whether a payday can be +paid twice or silently dropped. +""" + +import asyncio +from datetime import date + +import pytest + +from ..models import ContractStatus +from ..services import PeriodOutcome, run_due_periods +from .conftest import make_contract + + +@pytest.fixture +def payroll_stub(monkeypatch): + """Replace the DB and the wallet transfer with in-memory doubles. + + Returns a recorder holding the contract row and the period indices that + `pay_period` was asked to settle. + """ + + class Stub: + def __init__(self): + self.contract = None + self.attempted: list[int] = [] + self.outcome_status = "paid" + + stub = Stub() + + async def fake_get_contract(_contract_id): + return stub.contract + + async def fake_update_contract(contract): + stub.contract = contract + return contract + + async def fake_pay_period(contract, index, payday): + stub.attempted.append(index) + return PeriodOutcome( + index=index, + payday=payday, + status=stub.outcome_status, + detail="" if stub.outcome_status == "paid" else "stubbed failure", + amount_msat=int(contract.amount) * 1000, + ) + + from .. import services + + monkeypatch.setattr(services.crud, "get_contract", fake_get_contract) + monkeypatch.setattr(services.crud, "update_contract", fake_update_contract) + monkeypatch.setattr(services, "pay_period", fake_pay_period) + # Locks are created inside whichever loop is running; each test drives its + # own asyncio.run, so a lock left over from a previous test would raise. + services._contract_locks.clear() + return stub + + +def test_backdated_periods_are_skipped_when_backfill_is_off(payroll_stub): + """A contract created in May with a January start is anchoring a payday + ("we pay on the 15th"), not asking for four months of back-pay. + + The boundary is the contract's own creation date: Jan-Apr predate it and + are skipped, the 15 May period does not and is paid normally. + """ + payroll_stub.contract = make_contract( + start_date="2026-01-15", created_at="2026-05-01", backfill=False + ) + + outcomes = asyncio.run(run_due_periods(payroll_stub.contract, date(2026, 5, 20))) + + assert [o.status for o in outcomes] == ["skipped"] * 4 + ["paid"] + assert payroll_stub.attempted == [4] # only the first in-life period paid + # Skipped periods still consume the schedule, so the next payday is June. + assert payroll_stub.contract.periods_done == 5 + + +def test_backfill_pays_the_backlog(payroll_stub): + payroll_stub.contract = make_contract( + start_date="2026-01-15", created_at="2026-05-01", backfill=True + ) + + outcomes = asyncio.run(run_due_periods(payroll_stub.contract, date(2026, 3, 20))) + + assert [o.status for o in outcomes] == ["paid", "paid", "paid"] + assert payroll_stub.attempted == [0, 1, 2] + assert payroll_stub.contract.periods_done == 3 + + +def test_a_failure_halts_the_backlog_and_holds_the_position(payroll_stub): + """The crux: a failed period must not advance the counter, or that + payday is gone. And later periods must not jump the queue.""" + payroll_stub.contract = make_contract( + start_date="2026-01-15", created_at="2026-01-01", backfill=True + ) + payroll_stub.outcome_status = "failed" + + outcomes = asyncio.run(run_due_periods(payroll_stub.contract, date(2026, 4, 20))) + + assert [o.status for o in outcomes] == ["failed"] + assert payroll_stub.attempted == [0] # period 1 did not jump ahead + assert payroll_stub.contract.periods_done == 0 # retried next tick + + +def test_contract_completes_when_its_periods_run_out(payroll_stub): + payroll_stub.contract = make_contract( + start_date="2026-01-15", + created_at="2026-01-01", + total_periods=2, + backfill=True, + ) + + asyncio.run(run_due_periods(payroll_stub.contract, date(2026, 6, 1))) + + assert payroll_stub.contract.periods_done == 2 + assert payroll_stub.contract.status == ContractStatus.completed + + +def test_nothing_runs_before_the_first_payday(payroll_stub): + payroll_stub.contract = make_contract( + start_date="2026-06-01", created_at="2026-05-01" + ) + + outcomes = asyncio.run(run_due_periods(payroll_stub.contract, date(2026, 5, 31))) + + assert outcomes == [] + assert payroll_stub.attempted == [] diff --git a/tests/test_schedule.py b/tests/test_schedule.py new file mode 100644 index 0000000..c795a2a --- /dev/null +++ b/tests/test_schedule.py @@ -0,0 +1,149 @@ +"""Schedule math. + +These are the cases that make month-anchored payroll wrong in practice: +month-end clamping, clamping that must not become permanent, leap days, and +a schedule that must not drift no matter how late the scheduler runs. +""" + +from datetime import date + +import pytest + +from ..models import ContractStatus, Frequency +from ..services import ( + MAX_CATCH_UP_PERIODS, + add_months, + due_period_indices, + next_payday, + occurrence_on, + upcoming_paydays, +) +from .conftest import make_contract + +# --- month arithmetic ------------------------------------------------------ + + +def test_add_months_clamps_into_a_short_month(): + assert add_months(date(2026, 1, 31), 1) == date(2026, 2, 28) + + +def test_clamping_is_not_permanent(): + """The whole reason paydays are computed from the anchor and not from the + previous payday: 31 Jan must pay 28 Feb and then *31* Mar, not 28 Mar.""" + anchor = date(2026, 1, 31) + assert add_months(anchor, 1) == date(2026, 2, 28) + assert add_months(anchor, 2) == date(2026, 3, 31) + assert add_months(anchor, 3) == date(2026, 4, 30) + assert add_months(anchor, 4) == date(2026, 5, 31) + + +def test_add_months_crosses_year_boundaries(): + assert add_months(date(2026, 11, 30), 3) == date(2027, 2, 28) + assert add_months(date(2026, 3, 15), -4) == date(2025, 11, 15) + + +def test_leap_day_anchor_clamps_in_common_years(): + leap = date(2028, 2, 29) + assert add_months(leap, 12) == date(2029, 2, 28) + assert add_months(leap, 48) == date(2032, 2, 29) + + +# --- occurrences ----------------------------------------------------------- + + +@pytest.mark.parametrize( + ("frequency", "index", "expected"), + [ + (Frequency.daily, 0, date(2026, 1, 15)), + (Frequency.daily, 20, date(2026, 2, 4)), + (Frequency.weekly, 3, date(2026, 2, 5)), + (Frequency.biweekly, 2, date(2026, 2, 12)), + (Frequency.monthly, 2, date(2026, 3, 15)), + (Frequency.quarterly, 2, date(2026, 7, 15)), + (Frequency.yearly, 2, date(2028, 1, 15)), + ], +) +def test_occurrence_on(frequency, index, expected): + assert occurrence_on(date(2026, 1, 15), frequency, index) == expected + + +def test_period_zero_is_the_start_date_for_every_frequency(): + start = date(2026, 6, 30) + for frequency in Frequency: + assert occurrence_on(start, frequency, 0) == start + + +def test_schedule_does_not_drift_over_a_year(): + """A monthly contract must land on the same day twelve months later — + incrementally advancing a stored date is what would break this.""" + start = date(2026, 1, 15) + assert occurrence_on(start, Frequency.monthly, 12) == date(2027, 1, 15) + + +# --- due periods ----------------------------------------------------------- + + +def test_nothing_is_due_before_the_start_date(): + contract = make_contract(start_date="2026-03-01") + assert due_period_indices(contract, date(2026, 2, 28)) == [] + + +def test_the_start_date_itself_is_due(): + contract = make_contract(start_date="2026-03-01") + assert due_period_indices(contract, date(2026, 3, 1)) == [0] + + +def test_a_backlog_returns_every_missed_period_in_order(): + contract = make_contract(start_date="2026-01-15", frequency=Frequency.monthly) + assert due_period_indices(contract, date(2026, 4, 20)) == [0, 1, 2, 3] + + +def test_periods_already_done_are_not_due_again(): + contract = make_contract(start_date="2026-01-15", periods_done=3) + assert due_period_indices(contract, date(2026, 4, 20)) == [3] + + +def test_due_periods_stop_at_the_period_cap(): + contract = make_contract(start_date="2026-01-15", total_periods=2) + assert due_period_indices(contract, date(2026, 12, 31)) == [0, 1] + + +def test_catch_up_is_capped_for_open_ended_contracts(): + """A daily contract with a badly back-dated start must not try to fire + hundreds of transfers in one tick.""" + contract = make_contract( + start_date="2020-01-01", frequency=Frequency.daily, total_periods=None + ) + assert len(due_period_indices(contract, date(2026, 1, 1))) == MAX_CATCH_UP_PERIODS + + +@pytest.mark.parametrize( + "status", + [ContractStatus.paused, ContractStatus.cancelled, ContractStatus.completed], +) +def test_only_active_contracts_are_due(status): + contract = make_contract(start_date="2026-01-15", status=status) + assert due_period_indices(contract, date(2026, 6, 1)) == [] + + +# --- previews -------------------------------------------------------------- + + +def test_next_payday_follows_the_position(): + contract = make_contract(start_date="2026-01-31", periods_done=1) + assert next_payday(contract) == date(2026, 2, 28) + + +def test_next_payday_is_none_once_exhausted(): + contract = make_contract(total_periods=3, periods_done=3) + assert next_payday(contract) is None + + +def test_upcoming_paydays_are_truncated_by_the_period_cap(): + contract = make_contract(start_date="2026-01-15", total_periods=3, periods_done=1) + assert upcoming_paydays(contract, 10) == [date(2026, 2, 15), date(2026, 3, 15)] + + +def test_upcoming_paydays_are_unbounded_for_open_ended_contracts(): + contract = make_contract(start_date="2026-01-15", total_periods=None) + assert len(upcoming_paydays(contract, 24)) == 24