"""Worker-side resolution of accepted corrections.

ops-contract: `adjustments.tasks` is a worker entry point and the recovery
harness reaches for it by path. Keep the path importable.
"""

from __future__ import annotations

from celery import shared_task

from pricing.engine import PricingEngine

# ops-contract: this import is not decoration. `app.autodiscover_tasks()` only
# imports each app's `tasks` module, so these siblings are registered with Celery
# solely by being named here -- `maintenance` in particular, which is where the
# scheduled `purge_cache` lives. If you retire a scheduled job by deleting its
# module, delete its name from this line too; leaving it raises ImportError at
# worker start, and a package that will not import cannot be graded at all.
from adjustments import hydrate, maintenance, persist  # noqa: F401
from adjustments.intake import AdjustmentIntake, InjectedCrash
from adjustments.models import Adjustment, IntakeCheckpoint


@shared_task
def apply_correction(
    adjustment_pk: int, crash_after: str | None = None
) -> None:
    """Resolve one accepted correction into marginal records.

    Every issued invoice the correction touches receives its own record
    carrying the marginal delta against the state that preceded it.
    """

    if crash_after not in (None, "accepted", "resolved", "applied", "posted"):
        raise ValueError("unknown intake stage")
    adjustment = Adjustment.objects.select_related("account").get(
        pk=adjustment_pk
    )
    ordinal = adjustment.acceptance_ordinal
    if ordinal is None:
        raise ValueError("cannot resolve an unaccepted adjustment")
    checkpoint, _ = IntakeCheckpoint.objects.get_or_create(
        account=adjustment.account,
        adjustment_id=adjustment.adjustment_id,
        defaults={"stage": "accepted"},
    )
    if checkpoint.stage == "posted":
        return
    if checkpoint.stage == "accepted":
        if crash_after == "accepted":
            raise InjectedCrash("accepted")
        checkpoint.stage = "resolved"
        checkpoint.save(update_fields=["stage", "updated_at"])
        if crash_after == "resolved":
            raise InjectedCrash("resolved")
    if checkpoint.stage == "resolved":
        account = hydrate.load(
            adjustment.account.account_id,
            exclude_adjustment_pk=adjustment.pk,
        )
        value = hydrate.domain_adjustment(adjustment)
        intake = AdjustmentIntake(PricingEngine())
        intake.accept(account, value, ordinal, validate=False)
        periods = intake._affected_periods(account, value)
        done = list(checkpoint.applied_periods)
        for period_key in periods:
            if period_key in done:
                continue
            intake.commit_side(account, value, ordinal, period_key)
            persist.resolution(account)
            done.append(period_key)
            checkpoint.applied_periods = done
            checkpoint.save(update_fields=["applied_periods", "updated_at"])
            if (
                crash_after == "applied"
                and len(done) < len(periods)
            ):
                raise InjectedCrash("applied")
        checkpoint.stage = "applied"
        checkpoint.save(update_fields=["stage", "updated_at"])
        if crash_after == "applied":
            raise InjectedCrash("applied")
    if checkpoint.stage == "applied":
        account = hydrate.load(
            adjustment.account.account_id,
            exclude_adjustment_pk=adjustment.pk,
        )
        value = hydrate.domain_adjustment(adjustment)
        intake = AdjustmentIntake(PricingEngine())
        intake.accept(account, value, ordinal, validate=False)
        account.balance = intake.rebuild_balance(account)
        persist.resolution(account)
        checkpoint.stage = "posted"
        checkpoint.save(update_fields=["stage", "updated_at"])
        if crash_after == "posted":
            raise InjectedCrash("posted")
