"""Subscription reconciliation and provider-record retention tests."""

from __future__ import annotations

from datetime import datetime, timedelta, timezone
import hashlib
import unittest

from sqlalchemy import create_engine, func, select
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool

from licensing_shared.catalog import load_builtin_catalog
from licensing_shared.constants import PRODUCT_ID
from licensing_server.app.database import Base
from licensing_server.app.models import (
    AuditEvent,
    Grant,
    IdempotencyRecord,
    License,
    OutboxEvent,
    Subscription,
    SubscriptionItem,
    User,
    WebhookEvent,
)
from licensing_server.app.subscription_maintenance import (
    purge_subscription_provider_records,
    reconcile_stripe_subscriptions,
)
from licensing_server.app.subscriptions import process_subscription_outbox


TEST_NOW = datetime(2026, 8, 15, 18, 0, tzinfo=timezone.utc)
TEST_LICENSE_ID = "license.test.subscription_maintenance"
TEST_SUBSCRIPTION_ID = "subscription.test.maintenance"
TEST_PROVIDER_SUBSCRIPTION_ID = "sub_test_maintenance"
TEST_PROVIDER_CUSTOMER_ID = "cus_test_maintenance"
TEST_PRICE_ID = "price_test_ai_assist"


class FakeSubscriptionProvider:
    def __init__(self, provider_object: dict[str, object]) -> None:
        self.provider_object = provider_object
        self.requests: list[str] = []

    def retrieve_subscription(self, subscription_id: str) -> dict[str, object]:
        self.requests.append(subscription_id)
        return dict(self.provider_object)


class SubscriptionMaintenanceTests(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_maintenance",
                    external_issuer="https://identity.test",
                    external_subject="subscription-maintenance-subject",
                    verified_email="maintenance@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_maintenance",
                    owner_organization_id=None,
                    anonymous_subject_digest=None,
                    subject_type="user",
                    status="active",
                    device_policy_id="policy.named_user.standard",
                    revocation_generation=0,
                )
            )
            session.add(
                Subscription(
                    id=TEST_SUBSCRIPTION_ID,
                    license_id=TEST_LICENSE_ID,
                    provider="stripe",
                    provider_subscription_id=TEST_PROVIDER_SUBSCRIPTION_ID,
                    provider_customer_id=TEST_PROVIDER_CUSTOMER_ID,
                    status="active",
                    current_period_start=TEST_NOW,
                    current_period_end=TEST_NOW + timedelta(days=30),
                    cancel_at_period_end=False,
                    grace_ends_at=None,
                    provider_event_created_at=TEST_NOW,
                )
            )
            session.add(
                SubscriptionItem(
                    id="subscription_item.test.maintenance",
                    subscription_id=TEST_SUBSCRIPTION_ID,
                    sku_id="subscription.ai_assist",
                    quantity=1,
                    provider_item_id="si_test_maintenance",
                )
            )
            session.add(
                Grant(
                    id="grant.test.subscription_maintenance",
                    license_id=TEST_LICENSE_ID,
                    sku_id="subscription.ai_assist",
                    source_type="subscription",
                    source_reference=(
                        "stripe:sub_test_maintenance:subscription.ai_assist"
                    ),
                    catalog_revision=self.catalog.revision,
                    starts_at=TEST_NOW,
                    ends_at=TEST_NOW + timedelta(days=30),
                    status="pending",
                    metadata_json={"subscriptionState": "active"},
                )
            )

    def tearDown(self) -> None:
        self.engine.dispose()

    @staticmethod
    def _provider_object(status: str = "active") -> dict[str, object]:
        return {
            "id": TEST_PROVIDER_SUBSCRIPTION_ID,
            "customer": TEST_PROVIDER_CUSTOMER_ID,
            "status": status,
            "cancel_at_period_end": status == "canceled",
            "metadata": {"license_id": TEST_LICENSE_ID},
            "items": {
                "data": [
                    {
                        "id": "si_test_maintenance",
                        "quantity": 1,
                        "current_period_start": int(TEST_NOW.timestamp()),
                        "current_period_end": int(
                            (TEST_NOW + timedelta(days=30)).timestamp()
                        ),
                        "price": {"id": TEST_PRICE_ID},
                    }
                ]
            },
        }

    def test_reconciliation_skips_matching_state_and_queues_drift_once(self) -> None:
        provider = FakeSubscriptionProvider(self._provider_object())
        with Session(self.engine) as session:
            matching = reconcile_stripe_subscriptions(
                session,
                self.catalog,
                provider,
                {TEST_PRICE_ID: "subscription.ai_assist"},
                now=TEST_NOW + timedelta(hours=1),
                correlation_id="correlation.test.reconcile_matching",
            )
        self.assertEqual(matching.unchanged, 1)
        self.assertEqual(matching.queued, 0)

        provider.provider_object = self._provider_object("canceled")
        with Session(self.engine) as session:
            drift = reconcile_stripe_subscriptions(
                session,
                self.catalog,
                provider,
                {TEST_PRICE_ID: "subscription.ai_assist"},
                now=TEST_NOW + timedelta(hours=2),
                correlation_id="correlation.test.reconcile_drift",
            )
        self.assertEqual(drift.queued, 1)
        with Session(self.engine) as session:
            self.assertEqual(session.scalar(select(func.count()).select_from(OutboxEvent)), 1)
            self.assertEqual(session.scalar(select(Subscription)).status, "active")
            self.assertEqual(session.scalar(select(Grant)).status, "pending")

        with Session(self.engine) as session:
            duplicate = reconcile_stripe_subscriptions(
                session,
                self.catalog,
                provider,
                {TEST_PRICE_ID: "subscription.ai_assist"},
                now=TEST_NOW + timedelta(hours=3),
                correlation_id="correlation.test.reconcile_duplicate",
            )
        self.assertEqual(duplicate.already_queued, 1)
        with Session(self.engine) as session:
            self.assertEqual(
                process_subscription_outbox(
                    session,
                    self.catalog,
                    now=TEST_NOW + timedelta(hours=4),
                ),
                1,
            )
        with Session(self.engine) as session:
            self.assertEqual(session.scalar(select(Subscription)).status, "canceled")
            self.assertEqual(session.scalar(select(Grant)).status, "suspended")
            self.assertEqual(
                session.scalar(
                    select(func.count())
                    .select_from(AuditEvent)
                    .where(AuditEvent.action == "subscription.reconciliation_completed")
                ),
                3,
            )

    def test_retention_deletes_only_expired_terminal_provider_records(self) -> None:
        with Session(self.engine) as session, session.begin():
            for suffix, status, age_days in (
                ("processed", "processed", 31),
                ("failed", "failed", 366),
                ("pending", "pending", 500),
            ):
                event_id = f"evt_test_retention_{suffix}"
                processed_at = (
                    None if status == "pending" else TEST_NOW - timedelta(days=age_days)
                )
                session.add(
                    OutboxEvent(
                        id=f"outbox.test.retention.{suffix}",
                        event_type="stripe.subscription.project",
                        aggregate_type="subscription",
                        aggregate_id=TEST_PROVIDER_SUBSCRIPTION_ID,
                        payload_json={"providerEventId": event_id},
                        status=status,
                        available_at=TEST_NOW - timedelta(days=age_days),
                        attempts=1,
                        processed_at=processed_at,
                    )
                )
                session.add(
                    WebhookEvent(
                        id=f"webhook.test.retention.{suffix}",
                        provider="stripe",
                        provider_event_id=event_id,
                        provider_object_id=TEST_PROVIDER_SUBSCRIPTION_ID,
                        signature_valid=True,
                        status=status,
                        payload_digest=hashlib.sha256(event_id.encode("ascii")).digest(),
                        processed_at=processed_at,
                        error_code=None,
                    )
                )
            session.add(
                IdempotencyRecord(
                    id="idempotency.test.expired",
                    scope="account.billing.checkout",
                    subject_key="user.test.subscription_maintenance",
                    idempotency_key="idempotency-test-expired",
                    request_digest=hashlib.sha256(b"expired").digest(),
                    response_status=200,
                    response_json={"url": "https://checkout.example.test/expired"},
                    expires_at=TEST_NOW - timedelta(seconds=1),
                )
            )
            session.add(
                IdempotencyRecord(
                    id="idempotency.test.current",
                    scope="account.billing.checkout",
                    subject_key="user.test.subscription_maintenance",
                    idempotency_key="idempotency-test-current",
                    request_digest=hashlib.sha256(b"current").digest(),
                    response_status=200,
                    response_json={"url": "https://checkout.example.test/current"},
                    expires_at=TEST_NOW + timedelta(days=1),
                )
            )
        with Session(self.engine) as session:
            result = purge_subscription_provider_records(
                session,
                now=TEST_NOW,
                correlation_id="correlation.test.provider_retention",
            )
        self.assertEqual(result.processed_outbox_deleted, 1)
        self.assertEqual(result.failed_outbox_deleted, 1)
        self.assertEqual(result.webhook_records_deleted, 2)
        self.assertEqual(result.expired_idempotency_deleted, 1)
        with Session(self.engine) as session:
            self.assertEqual(session.scalar(select(func.count()).select_from(OutboxEvent)), 1)
            self.assertEqual(session.scalar(select(func.count()).select_from(WebhookEvent)), 1)
            retained = session.scalar(select(IdempotencyRecord))
            self.assertEqual(retained.id, "idempotency.test.current")
            self.assertIsNotNone(
                session.scalar(
                    select(AuditEvent).where(
                        AuditEvent.action == "subscription.provider_records_purged"
                    )
                )
            )


if __name__ == "__main__":
    unittest.main()
