"""Mapping between stored rows and the pricing domain's value objects.

The pricing engine and the correction-resolution algorithm are pure functions
over an in-memory account: base states plus the accepted adjustment history.
Nothing about them needs to know where that data came from.

So resolution is not reimplemented against the ORM. Rows are loaded into the
value objects the engine already expects, the existing algorithm runs, and the
results are written back. The arithmetic that decides what a customer owes
stays in one place, exercised by one implementation.
"""

from __future__ import annotations

from decimal import Decimal

from pricing.models import (
    Account as DomainAccount,
    Adjustment as DomainAdjustment,
    AdjustmentRecord as DomainAdjustmentRecord,
    BillingEntry as DomainBillingEntry,
    Cap,
    Credit,
    Invoice as DomainInvoice,
    LineItem as DomainLineItem,
    Period,
    Promotion,
    RuleSet,
    RunMembership as DomainRunMembership,
    RunSummaryEntry as DomainRunSummaryEntry,
    StatementRun as DomainStatementRun,
    ReceivableEntry as DomainReceivableEntry,
    Tax,
    UsageLine,
    UsageTier,
)

from adjustments.models import (
    Account,
    Adjustment,
    AdjustmentRecord,
    BillingEntry,
    Invoice,
    ReceivableEntry,
    StatementRun,
)


def _decimal(value: object) -> Decimal:
    return Decimal(str(value))


def _rule(raw: dict) -> object:
    kind = raw["type"]
    if kind == "promotion":
        return Promotion(
            raw["key"],
            _decimal(raw["percent"]),
            raw["kind"],
            tuple(raw["applies_to"]) if raw.get("applies_to") else None,
        )
    if kind == "credit":
        return Credit(raw["key"], _decimal(raw["amount"]))
    if kind == "cap":
        return Cap(raw["key"], _decimal(raw["amount"]))
    if kind == "tax":
        return Tax(raw["key"], _decimal(raw["rate"]))
    raise ValueError(f"unknown pricing rule type: {kind!r}")


def base_state(invoice: Invoice) -> tuple[Period, tuple[UsageLine, ...], RuleSet]:
    """Rebuild the immutable snapshot an invoice was priced from."""

    raw = invoice.base_state or {}
    period = Period(
        invoice.period_key, invoice.period_starts_on, invoice.period_ends_on
    )
    usage = tuple(
        UsageLine(line["key"], line["description"], _decimal(line["quantity"]))
        for line in raw.get("usage", [])
    )
    if not usage:
        # Compatibility for invoices created by the original visible workflow,
        # before immutable pricing snapshots were introduced.
        usage = (UsageLine("meter", "Metered service", Decimal("1")),)
        raw = {
            **raw,
            "usage_tiers": [
                {"up_to": None, "unit_price": str(invoice.total)}
            ],
        }
    rule_set = RuleSet(
        invoice.rule_set_key,
        raw.get("effective_from", ""),
        tuple(
            UsageTier(
                _decimal(tier["up_to"]) if tier.get("up_to") is not None else None,
                _decimal(tier["unit_price"]),
            )
            for tier in raw.get("usage_tiers", [])
        ),
        tuple(_rule(item) for item in raw.get("rules", [])),
    )
    return period, usage, rule_set


def domain_adjustment(row: Adjustment) -> DomainAdjustment:
    promotion = None
    if row.promotion_key:
        promotion = Promotion(
            row.promotion_key,
            _decimal(row.promotion_percent or 0),
            row.promotion_kind or "stackable",
            tuple(row.promotion_applies_to) if row.promotion_applies_to else None,
        )
    return DomainAdjustment(
        adjustment_id=row.adjustment_id,
        account_id=row.account.account_id,
        period_key=row.period_key,
        kind=row.kind,
        effective_at=row.effective_at,
        received_at=row.received_at,
        amount=_decimal(row.amount) if row.amount is not None else None,
        usage_key=row.usage_key,
        quantity=_decimal(row.quantity) if row.quantity is not None else None,
        promotion=promotion,
        target_period_key=row.target_period_key,
        subject_key=row.subject_key,
        replaces_adjustment_id=row.replaces_adjustment_id,
        withdraws_adjustment_id=row.withdraws_adjustment_id,
    )


def payload_adjustment(account_id: str, payload: dict) -> DomainAdjustment:
    promotion = payload.get("promotion")
    if isinstance(promotion, dict):
        promotion = Promotion(
            promotion["key"],
            _decimal(promotion["percent"]),
            promotion["kind"],
            (
                tuple(promotion["applies_to"])
                if promotion.get("applies_to")
                else None
            ),
        )
    return DomainAdjustment(
        adjustment_id=payload["adjustment_id"],
        account_id=account_id,
        period_key=payload["period_key"],
        kind=payload["kind"],
        effective_at=payload["effective_at"],
        received_at=payload["received_at"],
        amount=(
            _decimal(payload["amount"])
            if payload.get("amount") is not None
            else None
        ),
        usage_key=payload.get("usage_key"),
        quantity=(
            _decimal(payload["quantity"])
            if payload.get("quantity") is not None
            else None
        ),
        promotion=promotion,
        target_period_key=payload.get("target_period_key"),
        subject_key=payload.get("subject_key"),
        replaces_adjustment_id=payload.get("replaces_adjustment_id"),
        withdraws_adjustment_id=payload.get("withdraws_adjustment_id"),
    )


def load(
    account_id: str,
    *,
    before_ordinal: int | None = None,
    exclude_adjustment_pk: int | None = None,
) -> DomainAccount:
    """Load the whole account as the in-memory aggregate the engine expects."""

    row = Account.objects.get(account_id=account_id)
    domain = DomainAccount(
        account_id=row.account_id, settlement_policy=row.settlement_policy
    )
    domain.summaries_adopted_ordinal = row.summaries_adopted_ordinal
    domain._summary_adoption_id = row.summary_adoption_id
    for invoice in (
        Invoice.objects.filter(account=row, revision_of__isnull=True)
        .prefetch_related("line_items")
        .order_by("period_key")
    ):
        period, usage, rule_set = base_state(invoice)
        domain._base_states[invoice.period_key] = (period, usage, rule_set)
        line_items = tuple(
            DomainLineItem(
                item.usage_key,
                item.description,
                _decimal(item.quantity),
                _decimal(item.base),
                _decimal(item.net),
                _decimal(item.tax),
                _decimal(item.total),
            )
            for item in invoice.line_items.all()
        )
        domain.invoices[invoice.period_key] = DomainInvoice(
            invoice.invoice_id,
            row.account_id,
            period,
            usage,
            rule_set,
            line_items,
            _decimal(invoice.total),
            invoice.selected_exclusive_promotion,
            tuple(invoice.non_selected_exclusive_promotions),
            invoice.closed or invoice.issued,
            (
                _decimal(invoice.issued_total)
                if invoice.issued_total is not None
                else None
            ),
            invoice.revision_of,
        )
    for invoice in (
        Invoice.objects.filter(account=row, revision_of__isnull=False)
        .prefetch_related("line_items")
        .order_by("pk")
    ):
        period, usage, rule_set = base_state(invoice)
        domain.revisions.append(
            DomainInvoice(
                invoice.invoice_id,
                row.account_id,
                period,
                usage,
                rule_set,
                tuple(
                    DomainLineItem(
                        item.usage_key,
                        item.description,
                        _decimal(item.quantity),
                        _decimal(item.base),
                        _decimal(item.net),
                        _decimal(item.tax),
                        _decimal(item.total),
                    )
                    for item in invoice.line_items.all()
                ),
                _decimal(invoice.total),
                invoice.selected_exclusive_promotion,
                tuple(invoice.non_selected_exclusive_promotions),
                True,
                (
                    _decimal(invoice.issued_total)
                    if invoice.issued_total is not None
                    else None
                ),
                invoice.revision_of,
            )
        )
    adjustment_rows = Adjustment.objects.filter(account=row).select_related(
        "account"
    )
    if before_ordinal is not None:
        adjustment_rows = adjustment_rows.filter(
            acceptance_ordinal__lt=before_ordinal
        )
    if exclude_adjustment_pk is not None:
        adjustment_rows = adjustment_rows.exclude(pk=exclude_adjustment_pk)
    adjustment_rows = adjustment_rows.order_by("acceptance_ordinal", "pk")
    for item in adjustment_rows:
        value = domain_adjustment(item)
        domain.received_adjustments.append(value)
        if item.acceptance_ordinal is not None:
            domain.acceptance_ordinals[item.adjustment_id] = (
                item.acceptance_ordinal
            )
    for record in AdjustmentRecord.objects.filter(account=row).order_by(
        "acceptance_ordinal", "pk"
    ):
        domain.adjustment_records.append(
            DomainAdjustmentRecord(
                record.record_id,
                record.adjustment_id,
                record.acceptance_ordinal,
                record.original_invoice.invoice_id,
                record.target_period,
                _decimal(record.delta),
                _decimal(record.prior_recorded_total),
                _decimal(record.resulting_total),
                record.policy,
                record.effective_at,
                record.applied_period,
                record.replacement_invoice_id,
                (
                    _decimal(record.void_amount)
                    if record.void_amount is not None
                    else None
                ),
                (
                    _decimal(record.reissue_amount)
                    if record.reissue_amount is not None
                    else None
                ),
                record.replaces_adjustment_id,
                record.withdraws_adjustment_id,
            )
        )
    for entry in BillingEntry.objects.filter(account=row).order_by("pk"):
        domain.billing_entries.append(
            DomainBillingEntry(
                entry.entry_id,
                entry.source_type,
                entry.source_id,
                _decimal(entry.amount),
                entry.period_key,
                entry.invoice_id,
                entry.acceptance_ordinal,
            )
        )
    for run in (
        StatementRun.objects.filter(account=row)
        .prefetch_related("membership", "reconciliation_summary")
        .order_by("run_number")
    ):
        summary = None
        if (
            row.summaries_adopted_ordinal is not None
            and run.cutoff_ordinal >= row.summaries_adopted_ordinal
            and run.status == "issued"
        ):
            summary = tuple(
                DomainRunSummaryEntry(
                    item.invoice_id,
                    item.member_count,
                    _decimal(item.member_delta_total),
                )
                for item in run.reconciliation_summary.all()
            )
        domain.statement_runs.append(
            DomainStatementRun(
                run.run_id,
                row.account_id,
                run.run_number,
                run.operation_id,
                run.predecessor_run_id,
                run.previous_cutoff_ordinal,
                run.cutoff_ordinal,
                run.status,
                tuple(
                    DomainRunMembership(
                        item.adjustment_id,
                        item.acceptance_ordinal,
                        item.original_invoice_id,
                        item.record.record_id,
                    )
                    for item in run.membership.all()
                ),
                (
                    _decimal(run.recorded_demand)
                    if run.recorded_demand is not None
                    else None
                ),
                (
                    _decimal(run.incremental_demand)
                    if run.incremental_demand is not None
                    else None
                ),
                summary,
            )
        )
    for entry in ReceivableEntry.objects.filter(account=row).order_by("pk"):
        domain.receivable_entries.append(
            DomainReceivableEntry(
                entry.run_id,
                entry.predecessor_run_id,
                _decimal(entry.prior_demand),
                _decimal(entry.demand),
                _decimal(entry.incremental_demand),
            )
        )
    domain.balance = _decimal(row.balance)
    domain.amount_due = _decimal(row.amount_due)
    ordinals = [
        item.acceptance_ordinal
        for item in Adjustment.objects.filter(account=row)
        if item.acceptance_ordinal
    ]
    next_ordinal = (max(ordinals) + 1) if ordinals else 1
    if row.summaries_adopted_ordinal is not None:
        next_ordinal = max(next_ordinal, row.summaries_adopted_ordinal + 1)
    domain._next_acceptance_ordinal = next_ordinal
    return domain
