"""Stripe normalization, durable enqueue, and subscription projection tests."""

from __future__ import annotations

from dataclasses import replace
from datetime import datetime, timedelta, timezone
import json
import unittest

from sqlalchemy import create_engine, func, select
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool

from licensing_shared.catalog import LicensingCatalog, load_builtin_catalog
from licensing_shared.constants import PRODUCT_ID
from licensing_server.app.billing import ProviderBillingIncidentResolution
from licensing_server.app.database import Base
from licensing_server.app.models import (
    Grant,
    License,
    OutboxEvent,
    Subscription,
    SubscriptionItem,
    AuditEvent,
    User,
    WebhookEvent,
)
from licensing_server.app.notifications import CUSTOMER_NOTIFICATION_OUTBOX_EVENT_TYPE
from licensing_server.app.subscriptions import (
    enqueue_stripe_billing_incident_event,
    enqueue_stripe_subscription_event,
    normalize_stripe_billing_incident_event,
    normalize_stripe_subscription_event,
    process_subscription_outbox,
    requeue_failed_subscription_outbox,
)


TEST_NOW = datetime(2026, 8, 15, 16, 0, tzinfo=timezone.utc)
TEST_LICENSE_ID = "license.test.subscription"
TEST_PRICE_ID = "price_test_ai_assist"


class StaticBillingIncidentProvider:
    def __init__(self, hold_status: str | None, *, apply_hold: bool = True) -> None:
        self.hold_status = hold_status
        self.apply_hold = apply_hold

    def resolve_billing_incident(
        self,
        incident: dict[str, object],
    ) -> ProviderBillingIncidentResolution:
        return ProviderBillingIncidentResolution(
            provider_incident_id=str(incident["incidentId"]),
            subscription_id="sub_test_subscription",
            apply_hold=self.apply_hold,
            hold_status=self.hold_status,
        )


class SubscriptionProjectionTests(unittest.TestCase):
    def setUp(self) -> None:
        self.engine = create_engine(
            "sqlite+pysqlite:///:memory:",
            connect_args={"check_same_thread": False},
            poolclass=StaticPool,
        )
        Base.metadata.create_all(self.engine)
        self.catalog = load_builtin_catalog()
        with Session(self.engine) as session, session.begin():
            session.add(
                User(
                    id="user.test.subscription",
                    external_issuer="https://identity.test",
                    external_subject="subscription-subject",
                    verified_email="subscription@example.test",
                    status="active",
                    is_server_admin=False,
                )
            )
            session.add(
                License(
                    id=TEST_LICENSE_ID,
                    product_id=PRODUCT_ID,
                    owner_user_id="user.test.subscription",
                    owner_organization_id=None,
                    anonymous_subject_digest=None,
                    subject_type="user",
                    status="active",
                    device_policy_id="policy.named_user.standard",
                    revocation_generation=0,
                )
            )

    def tearDown(self) -> None:
        self.engine.dispose()

    @staticmethod
    def _event(
        event_id: str,
        status: str,
        created: datetime,
        *,
        period_start: datetime = TEST_NOW,
        period_end: datetime = TEST_NOW + timedelta(days=30),
    ) -> dict[str, object]:
        return {
            "id": event_id,
            "type": "customer.subscription.updated",
            "created": int(created.timestamp()),
            "data": {
                "object": {
                    "id": "sub_test_subscription",
                    "customer": "cus_test_subscription",
                    "status": status,
                    "cancel_at_period_end": False,
                    "metadata": {"license_id": TEST_LICENSE_ID},
                    "items": {
                        "data": [
                            {
                                "id": "si_test_ai_assist",
                                "quantity": 1,
                                "current_period_start": int(period_start.timestamp()),
                                "current_period_end": int(period_end.timestamp()),
                                "price": {"id": TEST_PRICE_ID},
                            }
                        ]
                    },
                }
            },
        }

    def _enqueue(self, event: dict[str, object]) -> bool:
        normalized = normalize_stripe_subscription_event(
            event,
            {TEST_PRICE_ID: "subscription.ai_assist"},
        )
        raw = json.dumps(event, sort_keys=True).encode("utf-8")
        with Session(self.engine) as session, session.begin():
            return enqueue_stripe_subscription_event(
                session,
                normalized,
                raw,
                now=TEST_NOW,
            )

    def _enqueue_incident(self, event: dict[str, object]) -> bool:
        normalized = normalize_stripe_billing_incident_event(event)
        raw = json.dumps(event, sort_keys=True).encode("utf-8")
        with Session(self.engine) as session, session.begin():
            return enqueue_stripe_billing_incident_event(
                session,
                normalized,
                raw,
                now=TEST_NOW,
            )

    def _process(self) -> int:
        with Session(self.engine) as session:
            return process_subscription_outbox(
                session,
                self.catalog,
                now=TEST_NOW,
            )

    def test_duplicate_events_queue_once_and_disabled_sku_grants_no_rights(self) -> None:
        event = self._event("evt_test_active", "active", TEST_NOW)
        self.assertTrue(self._enqueue(event))
        self.assertFalse(self._enqueue(event))
        self.assertEqual(self._process(), 1)
        with Session(self.engine) as session:
            self.assertEqual(session.scalar(select(func.count()).select_from(WebhookEvent)), 1)
            self.assertEqual(session.scalar(select(func.count()).select_from(OutboxEvent)), 1)
            subscription = session.scalar(select(Subscription))
            item = session.scalar(select(SubscriptionItem))
            grant = session.scalar(select(Grant))
            self.assertEqual(subscription.status, "active")
            self.assertEqual(
                subscription.provider_customer_id,
                "cus_test_subscription",
            )
            self.assertEqual(item.sku_id, "subscription.ai_assist")
            self.assertEqual(grant.status, "pending")
            self.assertEqual(grant.metadata_json["subscriptionState"], "active")

    def test_past_due_grace_and_out_of_order_event_do_not_rollback(self) -> None:
        active = self._event("evt_test_first", "active", TEST_NOW)
        past_due = self._event(
            "evt_test_past_due",
            "past_due",
            TEST_NOW + timedelta(hours=2),
        )
        older = self._event(
            "evt_test_older",
            "canceled",
            TEST_NOW + timedelta(hours=1),
        )
        self.assertTrue(self._enqueue(active))
        self.assertEqual(self._process(), 1)
        self.assertTrue(self._enqueue(past_due))
        self.assertEqual(self._process(), 1)
        self.assertTrue(self._enqueue(older))
        self.assertEqual(self._process(), 1)
        with Session(self.engine) as session:
            subscription = session.scalar(select(Subscription))
            grant = session.scalar(select(Grant))
            self.assertEqual(subscription.status, "past_due")
            self.assertEqual(
                subscription.grace_ends_at,
                (TEST_NOW + timedelta(days=37)).replace(tzinfo=None),
            )
            self.assertEqual(grant.metadata_json["subscriptionState"], "past_due")

    def test_unmapped_price_and_non_subscription_sku_fail_closed(self) -> None:
        event = self._event("evt_test_unmapped", "active", TEST_NOW)
        with self.assertRaises(ValueError):
            normalize_stripe_subscription_event(event, {})
        normalized = normalize_stripe_subscription_event(
            event,
            {TEST_PRICE_ID: "edition.personal"},
        )
        with Session(self.engine) as session, session.begin():
            enqueue_stripe_subscription_event(
                session,
                normalized,
                b"signed-test-payload",
                now=TEST_NOW,
            )
        with Session(self.engine) as session:
            self.assertEqual(
                process_subscription_outbox(
                    session,
                    self.catalog,
                    now=TEST_NOW,
                ),
                0,
            )
        with Session(self.engine) as session:
            self.assertEqual(session.scalar(select(OutboxEvent)).status, "failed")
            self.assertEqual(session.scalar(select(WebhookEvent)).status, "failed")
        with Session(self.engine) as session:
            self.assertEqual(
                requeue_failed_subscription_outbox(
                    session,
                    now=TEST_NOW + timedelta(minutes=1),
                ),
                1,
            )
        with Session(self.engine) as session:
            self.assertEqual(session.scalar(select(OutboxEvent)).status, "pending")
            self.assertEqual(session.scalar(select(WebhookEvent)).status, "queued")

    def test_dispute_and_full_refund_holds_survive_subscription_refresh(self) -> None:
        original_sku = self.catalog.skus["subscription.ai_assist"]
        active_sku = replace(original_sku, active=True)
        active_catalog = LicensingCatalog(
            product_id=self.catalog.product_id,
            revision=self.catalog.revision,
            entitlements=self.catalog.entitlements,
            device_policies=self.catalog.device_policies,
            skus={**self.catalog.skus, active_sku.sku_id: active_sku},
        )
        self.assertTrue(self._enqueue(self._event("evt_test_hold_base", "active", TEST_NOW)))
        with Session(self.engine) as session:
            self.assertEqual(
                process_subscription_outbox(session, active_catalog, now=TEST_NOW),
                1,
            )
        dispute = {
            "id": "evt_test_dispute_open",
            "type": "charge.dispute.created",
            "created": int((TEST_NOW + timedelta(minutes=1)).timestamp()),
            "data": {
                "object": {
                    "id": "dp_test_subscription",
                    "charge": "ch_test_subscription",
                    "status": "needs_response",
                }
            },
        }
        self.assertTrue(self._enqueue_incident(dispute))
        with Session(self.engine) as session:
            self.assertEqual(
                process_subscription_outbox(
                    session,
                    active_catalog,
                    now=TEST_NOW + timedelta(minutes=1),
                    billing_incident_provider=StaticBillingIncidentProvider(
                        "dispute_open"
                    ),
                ),
                1,
            )
        refreshed = self._event(
            "evt_test_hold_refresh",
            "active",
            TEST_NOW + timedelta(minutes=2),
        )
        self.assertTrue(self._enqueue(refreshed))
        with Session(self.engine) as session:
            self.assertEqual(
                process_subscription_outbox(
                    session,
                    active_catalog,
                    now=TEST_NOW + timedelta(minutes=2),
                ),
                1,
            )
            subscription = session.scalar(select(Subscription))
            grant = session.scalar(select(Grant))
            self.assertEqual(subscription.billing_hold_status, "dispute_open")
            self.assertEqual(grant.status, "suspended")
            self.assertEqual(grant.metadata_json["billingHold"], "dispute_open")

        won = {
            "id": "evt_test_dispute_won",
            "type": "charge.dispute.closed",
            "created": int((TEST_NOW + timedelta(minutes=3)).timestamp()),
            "data": {
                "object": {
                    "id": "dp_test_subscription",
                    "charge": "ch_test_subscription",
                    "status": "won",
                }
            },
        }
        self.assertTrue(self._enqueue_incident(won))
        with Session(self.engine) as session:
            self.assertEqual(
                process_subscription_outbox(
                    session,
                    active_catalog,
                    now=TEST_NOW + timedelta(minutes=3),
                    billing_incident_provider=StaticBillingIncidentProvider(None),
                ),
                1,
            )
            self.assertIsNone(session.scalar(select(Subscription)).billing_hold_status)
            self.assertEqual(session.scalar(select(Grant)).status, "active")

        refund = {
            "id": "evt_test_refund_full",
            "type": "refund.created",
            "created": int((TEST_NOW + timedelta(minutes=4)).timestamp()),
            "data": {
                "object": {
                    "id": "re_test_subscription",
                    "charge": "ch_test_subscription",
                    "status": "succeeded",
                }
            },
        }
        self.assertTrue(self._enqueue_incident(refund))
        with Session(self.engine) as session:
            self.assertEqual(
                process_subscription_outbox(
                    session,
                    active_catalog,
                    now=TEST_NOW + timedelta(minutes=4),
                    billing_incident_provider=StaticBillingIncidentProvider("refunded"),
                ),
                1,
            )
            self.assertEqual(
                session.scalar(select(Subscription)).billing_hold_status,
                "refunded",
            )
            self.assertEqual(session.scalar(select(Grant)).status, "suspended")
            self.assertEqual(
                session.scalar(
                    select(func.count()).select_from(AuditEvent).where(
                        AuditEvent.action == "subscription.billing_hold_changed"
                    )
                ),
                3,
            )
            notification_kinds = {
                value.payload_json["notificationKind"]
                for value in session.scalars(
                    select(OutboxEvent).where(
                        OutboxEvent.event_type == CUSTOMER_NOTIFICATION_OUTBOX_EVENT_TYPE
                    )
                )
            }
            self.assertEqual(
                notification_kinds,
                {
                    "subscription_started",
                    "dispute_open",
                    "dispute_resolved",
                    "refunded",
                },
            )


if __name__ == "__main__":
    unittest.main()
