"""Pure correction resolution over pricing value objects.

Storage adapters hydrate an account, call this module, then persist the
result. Pricing arithmetic remains in ``pricing.engine``.

ops-contract: the recovery harness imports `adjustments.intake` and the
``InjectedCrash`` below, and passes ``crash_after`` stage names into
``services.deliver``, ``services.issue_statement_run`` and
``tasks.apply_correction``. It drives production code on purpose -- a seam that
only exists inside a test harness has never been proved against the path that
actually runs. It reads as test scaffolding left in production code because that
is what it is. Do not tidy it out.
"""

from __future__ import annotations

from decimal import Decimal

from pricing.engine import PricingEngine, money
from pricing.models import (
    Account,
    Adjustment,
    AdjustmentRecord,
    BillingEntry,
    ReceivableEntry,
    RunMembership,
    RunSummaryEntry,
    StatementRun,
)


class InjectedCrash(RuntimeError):
    def __init__(self, stage: str) -> None:
        super().__init__(f"processing stopped after {stage}")
        self.stage = stage


def _subject(adjustment: Adjustment) -> str:
    return adjustment.subject_key or f"adjustment:{adjustment.adjustment_id}"


class AdjustmentIntake:
    def __init__(self, engine: PricingEngine) -> None:
        self.engine = engine

    def _by_id(self, account: Account, adjustment_id: str) -> Adjustment | None:
        return next(
            (
                item
                for item in account.received_adjustments
                if item.adjustment_id == adjustment_id
            ),
            None,
        )

    def _validate_lineage(
        self, account: Account, adjustment: Adjustment
    ) -> None:
        if adjustment.kind == "withdrawal":
            if not adjustment.withdraws_adjustment_id:
                raise ValueError("withdrawal requires a target")
            if any(
                item.kind == "withdrawal"
                and item.withdraws_adjustment_id == adjustment.adjustment_id
                for item in account.received_adjustments
            ):
                raise ValueError("withdrawals cannot be withdrawal targets")
            target = self._by_id(account, adjustment.withdraws_adjustment_id)
            if target is not None and target.kind == "withdrawal":
                raise ValueError("withdrawals target adjustment versions")
            return
        subject = _subject(adjustment)
        if adjustment.replaces_adjustment_id is not None:
            parent = self._by_id(account, adjustment.replaces_adjustment_id)
            if parent is None or parent.kind == "withdrawal":
                raise ValueError("replacement parent is not accepted")
            if _subject(parent) != subject:
                raise ValueError("replacement subject does not match")
            if any(
                item.replaces_adjustment_id == parent.adjustment_id
                for item in account.received_adjustments
                if item.kind != "withdrawal"
            ):
                raise ValueError("replacement parent already has a child")
        elif any(
            item.kind != "withdrawal" and _subject(item) == subject
            for item in account.received_adjustments
        ):
            raise ValueError("correction subject already has a root")

    def validate(self, account: Account, adjustment: Adjustment) -> None:
        if adjustment.account_id != account.account_id:
            raise ValueError("adjustment account does not match")
        self._validate_lineage(account, adjustment)

    def _affected_periods(
        self, account: Account, adjustment: Adjustment
    ) -> list[str]:
        target = adjustment
        if adjustment.kind == "withdrawal" and adjustment.withdraws_adjustment_id:
            target = self._by_id(account, adjustment.withdraws_adjustment_id) or adjustment
        periods = [target.period_key]
        if target.target_period_key is not None:
            periods.append(target.target_period_key)
        return list(dict.fromkeys(periods))

    def _forward_period(self, account: Account, target: str) -> str:
        candidates = sorted(
            invoice.period.key
            for invoice in account.invoices.values()
            if not invoice.closed and invoice.period.key > target
        )
        return candidates[0] if candidates else "next-open-period"

    def resolve_accepted(
        self,
        account: Account,
        adjustment: Adjustment,
        acceptance_ordinal: int,
    ) -> AdjustmentRecord | None:
        """Resolve an adjustment whose durable intake row already exists."""

        self.accept(account, adjustment, acceptance_ordinal)
        first = None
        for period_key in self._affected_periods(account, adjustment):
            record = self.commit_side(
                account, adjustment, acceptance_ordinal, period_key
            )
            if first is None:
                first = record
        account.balance = self.rebuild_balance(account)
        return first

    def accept(
        self,
        account: Account,
        adjustment: Adjustment,
        acceptance_ordinal: int,
        *,
        validate: bool = True,
    ) -> None:
        if validate:
            self.validate(account, adjustment)
        account.received_adjustments.append(adjustment)
        account.acceptance_ordinals[adjustment.adjustment_id] = acceptance_ordinal
        account._next_acceptance_ordinal = max(
            account._next_acceptance_ordinal, acceptance_ordinal + 1
        )

    def commit_side(
        self,
        account: Account,
        adjustment: Adjustment,
        acceptance_ordinal: int,
        period_key: str,
    ) -> AdjustmentRecord | None:
        """Commit one affected invoice side; safe to repeat after a crash."""

        before = self.engine.resolved_invoice(
            account, period_key, acceptance_ordinal - 1
        )
        after = self.engine.resolved_invoice(
            account, period_key, acceptance_ordinal
        )
        invoice = account.invoices[period_key]
        if not invoice.closed:
            current = self.engine.resolved_invoice(
                account,
                period_key,
                account._next_acceptance_ordinal - 1,
            )
            invoice.usage = current.usage
            invoice.rule_set = current.rule_set
            invoice.line_items = current.line_items
            invoice.total = current.total
            invoice.selected_exclusive_promotion = (
                current.selected_exclusive_promotion
            )
            invoice.non_selected_exclusive_promotions = (
                current.non_selected_exclusive_promotions
            )
            return None

        record_id = f"AR-{adjustment.adjustment_id}-{invoice.invoice_id}"
        existing = next(
            (
                record
                for record in account.adjustment_records
                if record.record_id == record_id
            ),
            None,
        )
        if existing is not None:
            return existing
        delta = money(after.total - before.total)
        replacement_id = None
        applied_period = None
        void_amount = None
        reissue_amount = None
        if account.settlement_policy == "restate":
            number = 1 + sum(
                revision.revision_of == invoice.invoice_id
                for revision in account.revisions
            )
            replacement_id = f"{invoice.invoice_id}-R{number}"
            after.invoice_id = replacement_id
            after.revision_of = invoice.invoice_id
            after.closed = True
            after.issued_total = after.total
            account.revisions.append(after)
            void_amount = -before.total
            reissue_amount = after.total
        else:
            applied_period = self._forward_period(account, period_key)
        record = AdjustmentRecord(
            record_id,
            adjustment.adjustment_id,
            acceptance_ordinal,
            invoice.invoice_id,
            period_key,
            delta,
            before.total,
            after.total,
            account.settlement_policy,
            adjustment.effective_at,
            applied_period,
            replacement_id,
            void_amount,
            reissue_amount,
            adjustment.replaces_adjustment_id,
            adjustment.withdraws_adjustment_id,
        )
        account.adjustment_records.append(record)
        account.billing_entries.append(
            BillingEntry(
                f"BE-{record.record_id}",
                "adjustment",
                adjustment.adjustment_id,
                delta,
                period_key,
                invoice.invoice_id,
                acceptance_ordinal,
            )
        )
        return record

    def rebuild_balance(self, account: Account) -> Decimal:
        return money(
            sum(
                (entry.amount for entry in account.billing_entries),
                Decimal("0"),
            )
        )

    def start_statement_run(
        self, account: Account, operation_id: str
    ) -> StatementRun:
        existing = next(
            (
                run
                for run in account.statement_runs
                if run.operation_id == operation_id
            ),
            None,
        )
        if existing is not None:
            return existing
        if any(run.status == "started" for run in account.statement_runs):
            raise ValueError("another statement run is started")
        predecessor = next(
            (
                run
                for run in reversed(account.statement_runs)
                if run.status == "issued"
            ),
            None,
        )
        number = len(account.statement_runs) + 1
        run = StatementRun(
            f"ST-{account.account_id}-{number}",
            account.account_id,
            number,
            operation_id,
            predecessor.run_id if predecessor else None,
            predecessor.cutoff_ordinal if predecessor else 0,
            account._next_acceptance_ordinal - 1,
        )
        account.statement_runs.append(run)
        return run

    def issue_statement_run(
        self, account: Account, operation_id: str
    ) -> StatementRun:
        run = self.start_statement_run(account, operation_id)
        if run.status == "issued":
            return run
        run.membership = tuple(
            RunMembership(
                record.adjustment_id,
                record.acceptance_ordinal,
                record.original_invoice_id,
                record.record_id,
            )
            for record in sorted(
                account.adjustment_records,
                key=lambda item: (
                    item.acceptance_ordinal,
                    item.original_invoice_id,
                ),
            )
            if run.previous_cutoff_ordinal
            < record.acceptance_ordinal
            <= run.cutoff_ordinal
        )
        demand = money(
            sum(
                (
                    entry.amount
                    for entry in account.billing_entries
                    if entry.acceptance_ordinal is None
                    or entry.acceptance_ordinal <= run.cutoff_ordinal
                ),
                Decimal("0"),
            )
        )
        predecessor = next(
            (
                item
                for item in account.statement_runs
                if item.run_id == run.predecessor_run_id
            ),
            None,
        )
        prior = (
            predecessor.recorded_demand
            if predecessor and predecessor.recorded_demand is not None
            else Decimal("0.00")
        )
        run.recorded_demand = demand
        run.incremental_demand = money(demand - prior)
        if (
            account.summaries_adopted_ordinal is not None
            and run.cutoff_ordinal >= account.summaries_adopted_ordinal
        ):
            by_invoice: dict[str, list[Decimal]] = {}
            for member in run.membership:
                record = next(
                    item
                    for item in account.adjustment_records
                    if item.record_id == member.record_id
                )
                by_invoice.setdefault(member.original_invoice_id, []).append(
                    record.delta
                )
            run.reconciliation_summary = tuple(
                RunSummaryEntry(
                    invoice_id,
                    len(deltas),
                    money(sum(deltas, Decimal("0.00"))),
                )
                for invoice_id, deltas in sorted(by_invoice.items())
            )
        run.status = "issued"
        account.receivable_entries.append(
            ReceivableEntry(
                run.run_id,
                run.predecessor_run_id,
                prior,
                demand,
                run.incremental_demand,
            )
        )
        account.amount_due = demand
        return run

    def adopt_run_summaries(
        self, account: Account, adoption_id: str
    ) -> None:
        if account._summary_adoption_id == adoption_id:
            return
        if account._summary_adoption_id is not None:
            raise ValueError("run summaries already adopted")
        account._summary_adoption_id = adoption_id
        account.summaries_adopted_ordinal = account._next_acceptance_ordinal
        account._next_acceptance_ordinal += 1
