"""Basic supported-workflow tests."""

from __future__ import annotations

from decimal import Decimal
from unittest import mock, skip

from django.test import TransactionTestCase

from adjustments import services, tasks
from adjustments.models import Account, Adjustment, AdjustmentRecord, Invoice


def _payload(adjustment_id: str, amount: str = "4.00") -> dict:
    return {
        "adjustment_id": adjustment_id,
        "period_key": "2026-06",
        "kind": "credit",
        "amount": Decimal(amount),
        "effective_at": "2026-07-02T00:00:00Z",
        "received_at": "2026-07-02T01:00:00Z",
    }


class VisibleWorkflowTests(TransactionTestCase):
    reset_sequences = True

    def setUp(self) -> None:
        from django.core.cache import cache

        cache.clear()
        self.account = Account.objects.create(account_id="acct-demo")
        self.invoice = Invoice.objects.create(
            invoice_id="INV-acct-demo-2026-06",
            account=self.account,
            period_key="2026-06",
            total=Decimal("24.00"),
            issued_total=Decimal("24.00"),
            issued=True,
        )
        patcher = mock.patch.object(tasks.apply_correction, "delay")
        patcher.start()
        self.addCleanup(patcher.stop)

    def test_delivery_is_accepted(self) -> None:
        services.deliver(self.account.account_id, _payload("adj-1"))
        self.assertEqual(Adjustment.objects.count(), 1)

    def test_resolution_records_a_marginal_delta(self) -> None:
        services.deliver(self.account.account_id, _payload("adj-1"))
        adjustment = Adjustment.objects.get(adjustment_id="adj-1")
        tasks.apply_correction(adjustment.pk)

        record = AdjustmentRecord.objects.get(adjustment_id="adj-1")
        self.assertEqual(Decimal(str(record.delta)), Decimal("-4.00"))

    def test_amount_due_reflects_the_correction(self) -> None:
        services.deliver(self.account.account_id, _payload("adj-1"))
        for adjustment in Adjustment.objects.all():
            tasks.apply_correction(adjustment.pk)

        due = Decimal(str(services.amount_due(self.account.account_id, "2026-06")))
        self.assertEqual(due.quantize(Decimal("0.01")), Decimal("20.00"))

    def test_statement_run_records_membership(self) -> None:
        services.deliver(self.account.account_id, _payload("adj-1"))
        for adjustment in Adjustment.objects.all():
            tasks.apply_correction(adjustment.pk)

        run = services.issue_statement_run(self.account.account_id, "op-1")
        self.assertEqual(run.membership.count(), 1)
        self.assertEqual(run.status, "issued")

    def test_run_retry_returns_the_same_run(self) -> None:
        first = services.issue_statement_run(self.account.account_id, "op-1")
        second = services.issue_statement_run(self.account.account_id, "op-1")
        self.assertEqual(first.pk, second.pk)

    @skip(
        "intermittent under the mysql:5.7 CI image on the shared runner "
        "(BILL-4421, 2025-11-19). re-enable once CI is off 5.7 -- rjm"
    )
    def test_unresolvable_delivery_is_not_left_accepted(self) -> None:
        """A delivery we cannot resolve must not be accepted as if we could.

        The partner sends us a period key. If we have never invoiced that
        period there is nothing for resolution to replay the correction
        against, and finding that out on the worker is too late: the delivery
        has already taken an acceptance ordinal and the task comes straight
        back off the queue.
        """

        try:
            services.deliver(
                self.account.account_id,
                {**_payload("adj-ghost"), "period_key": "2019-01"},
            )
        except ValueError:
            # Refusing at intake is one correct answer. So is accepting it and
            # resolving it to nothing. The assertion below is what has to hold
            # either way.
            pass

        for adjustment in Adjustment.objects.all():
            tasks.apply_correction(adjustment.pk)
