"""Licensing authority migration and PostgreSQL DDL regression tests."""

from __future__ import annotations

from datetime import datetime, timedelta, timezone
from io import StringIO
import os
from pathlib import Path
import tempfile
import unittest
from unittest.mock import patch

from alembic import command
from alembic.config import Config
from sqlalchemy import MetaData, Table, create_engine, inspect, select


REPOSITORY_ROOT = Path(__file__).resolve().parents[2]
ALEMBIC_CONFIG_PATH = REPOSITORY_ROOT / "licensing_server/alembic.ini"
EXPECTED_DOMAIN_TABLES = {
    "activations",
    "audit_events",
    "catalog_releases",
    "device_installations",
    "devices",
    "grants",
    "idempotency_records",
    "leases",
    "licenses",
    "memberships",
    "notification_deliveries",
    "notification_feedback_events",
    "notification_suppressions",
    "organizations",
    "privacy_requests",
    "product_surface_profiles",
    "rate_limit_buckets",
    "organization_invitations",
    "organization_seats",
    "offline_certificates",
    "offline_requests",
    "outbox_events",
    "serial_batches",
    "serial_redemptions",
    "serials",
    "signing_key_metadata",
    "subscription_items",
    "subscriptions",
    "trials",
    "users",
    "webhook_events",
}


def _config(database_url: str, *, output_buffer: StringIO | None = None) -> Config:
    if not isinstance(database_url, str) or not database_url:
        raise ValueError("database_url must be a non-empty string")
    config = Config(str(ALEMBIC_CONFIG_PATH), output_buffer=output_buffer)
    config.set_main_option("sqlalchemy.url", database_url)
    return config


class LicensingMigrationTests(unittest.TestCase):
    def test_environment_database_url_file_overrides_placeholder(self) -> None:
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            database_path = root / "file-configured-migrations.db"
            ignored_path = root / "ignored-placeholder.db"
            database_url = f"sqlite+pysqlite:///{database_path.as_posix()}"
            database_url_file = root / "licensing_database_url"
            database_url_file.write_text(database_url, encoding="utf-8")
            with patch.dict(
                os.environ,
                {
                    "LICENSING_DATABASE_URL_FILE": str(database_url_file),
                    "LICENSING_ENVIRONMENT": "development",
                },
                clear=True,
            ):
                command.upgrade(
                    _config(
                        f"sqlite+pysqlite:///{ignored_path.as_posix()}"
                    ),
                    "head",
                )
            engine = create_engine(database_url)
            try:
                self.assertTrue(
                    EXPECTED_DOMAIN_TABLES.issubset(
                        set(inspect(engine).get_table_names())
                    )
                )
            finally:
                engine.dispose()
            self.assertFalse(ignored_path.exists())

    def test_future_staged_compromise_survives_revision_downgrade(self) -> None:
        with tempfile.TemporaryDirectory() as directory:
            database_path = Path(directory) / "signing-compromise-migration.db"
            database_url = f"sqlite+pysqlite:///{database_path.as_posix()}"
            config = _config(database_url)
            command.upgrade(config, "head")
            compromised_at = datetime(2026, 8, 16, 22, 0, tzinfo=timezone.utc)
            not_before = compromised_at + timedelta(days=30)
            engine = create_engine(database_url)
            try:
                metadata = MetaData()
                signing_keys = Table(
                    "signing_key_metadata",
                    metadata,
                    autoload_with=engine,
                )
                with engine.begin() as connection:
                    connection.execute(
                        signing_keys.insert().values(
                            id="signing_key.test.future_compromise",
                            key_id="license.test.future_compromise",
                            purpose="license_snapshot",
                            public_key_bytes=b"k" * 32,
                            status="compromised",
                            not_before=not_before,
                            expires_at=compromised_at,
                            activated_at=None,
                            retired_at=compromised_at,
                        )
                    )
            finally:
                engine.dispose()

            command.downgrade(config, "f3a7c9d52b18")
            engine = create_engine(database_url)
            try:
                metadata = MetaData()
                signing_keys = Table(
                    "signing_key_metadata",
                    metadata,
                    autoload_with=engine,
                )
                with engine.connect() as connection:
                    row = connection.execute(
                        select(signing_keys).where(
                            signing_keys.c.key_id
                            == "license.test.future_compromise"
                        )
                    ).mappings().one()
                self.assertEqual(row["status"], "compromised")
                self.assertGreater(row["expires_at"], row["not_before"])
            finally:
                engine.dispose()

            command.upgrade(config, "head")
            engine = create_engine(database_url)
            try:
                metadata = MetaData()
                signing_keys = Table(
                    "signing_key_metadata",
                    metadata,
                    autoload_with=engine,
                )
                with engine.connect() as connection:
                    self.assertEqual(
                        connection.scalar(
                            select(signing_keys.c.status).where(
                                signing_keys.c.key_id
                                == "license.test.future_compromise"
                            )
                        ),
                        "compromised",
                    )
            finally:
                engine.dispose()

    def test_sqlite_upgrade_downgrade_and_reupgrade_round_trip(self) -> None:
        with tempfile.TemporaryDirectory() as directory:
            database_path = Path(directory) / "licensing-migrations.db"
            database_url = f"sqlite+pysqlite:///{database_path.as_posix()}"
            config = _config(database_url)
            command.upgrade(config, "head")
            engine = create_engine(database_url)
            try:
                tables = set(inspect(engine).get_table_names())
                self.assertTrue(EXPECTED_DOMAIN_TABLES.issubset(tables))
                subscription_columns = {
                    value["name"] for value in inspect(engine).get_columns("subscriptions")
                }
                self.assertIn("provider_event_created_at", subscription_columns)
                self.assertIn("provider_customer_id", subscription_columns)
                self.assertIn("billing_hold_status", subscription_columns)
                self.assertIn("billing_hold_provider_id", subscription_columns)
                self.assertIn("billing_hold_event_created_at", subscription_columns)
                offline_request_columns = {
                    value["name"]
                    for value in inspect(engine).get_columns("offline_requests")
                }
                self.assertIn("decision_generation", offline_request_columns)
                self.assertIn("request_digest", offline_request_columns)
                serial_batch_columns = {
                    value["name"]
                    for value in inspect(engine).get_columns("serial_batches")
                }
                self.assertIn("revoked_at", serial_batch_columns)
                self.assertIn("revoked_by_user_id", serial_batch_columns)
                self.assertIn("revocation_reason", serial_batch_columns)
                self.assertIn("surface_profile_id", serial_batch_columns)
                self.assertIn("trial_duration_hours", serial_batch_columns)
                serial_batch_constraints = {
                    value["name"]
                    for value in inspect(engine).get_check_constraints(
                        "serial_batches"
                    )
                }
                self.assertIn("ck_serial_batches_status", serial_batch_constraints)
                self.assertIn(
                    "ck_serial_batches_revocation_lifecycle",
                    serial_batch_constraints,
                )
                self.assertIn(
                    "ck_serial_batches_trial_duration_hours",
                    serial_batch_constraints,
                )
                license_columns = {
                    value["name"] for value in inspect(engine).get_columns("licenses")
                }
                self.assertIn("surface_profile_id", license_columns)
                surface_profile_constraints = {
                    value["name"]
                    for value in inspect(engine).get_check_constraints(
                        "product_surface_profiles"
                    )
                }
                self.assertTrue(
                    {
                        "ck_product_surface_profiles_status",
                        "ck_product_surface_profiles_inventory_revision",
                    }.issubset(surface_profile_constraints)
                )
                serial_constraints = {
                    value["name"]
                    for value in inspect(engine).get_check_constraints("serials")
                }
                self.assertIn("ck_serials_status", serial_constraints)
                signing_key_columns = {
                    value["name"]
                    for value in inspect(engine).get_columns(
                        "signing_key_metadata"
                    )
                }
                self.assertIn("activated_at", signing_key_columns)
                self.assertIn("retired_at", signing_key_columns)
                signing_key_constraints = {
                    value["name"]
                    for value in inspect(engine).get_check_constraints(
                        "signing_key_metadata"
                    )
                }
                self.assertTrue(
                    {
                        "ck_signing_key_metadata_status",
                        "ck_signing_key_metadata_validity",
                        "ck_signing_key_metadata_lifecycle",
                        "ck_signing_key_metadata_retirement",
                    }.issubset(signing_key_constraints)
                )
                signing_key_indexes = {
                    value["name"]
                    for value in inspect(engine).get_indexes(
                        "signing_key_metadata"
                    )
                }
                self.assertIn(
                    "uq_signing_key_metadata_active_purpose",
                    signing_key_indexes,
                )
                catalog_release_constraints = {
                    value["name"]
                    for value in inspect(engine).get_check_constraints(
                        "catalog_releases"
                    )
                }
                self.assertTrue(
                    {
                        "ck_catalog_releases_revision",
                        "ck_catalog_releases_status",
                        "ck_catalog_releases_lifecycle",
                    }.issubset(catalog_release_constraints)
                )
                catalog_release_indexes = {
                    value["name"]
                    for value in inspect(engine).get_indexes("catalog_releases")
                }
                self.assertIn(
                    "uq_catalog_releases_active_product",
                    catalog_release_indexes,
                )
                outbox_columns = {
                    value["name"] for value in inspect(engine).get_columns("outbox_events")
                }
                self.assertIn("last_error_code", outbox_columns)
                delivery_columns = {
                    value["name"]
                    for value in inspect(engine).get_columns("notification_deliveries")
                }
                self.assertIn("provider_message_id", delivery_columns)
                self.assertIn("recipient_address_digest", delivery_columns)
                self.assertIn("feedback_status", delivery_columns)
                self.assertIn("delivered_at", delivery_columns)
                delivery_indexes = {
                    value["name"]
                    for value in inspect(engine).get_indexes(
                        "notification_deliveries"
                    )
                }
                self.assertIn(
                    "ix_notification_deliveries_provider_message",
                    delivery_indexes,
                )
                feedback_columns = {
                    value["name"]
                    for value in inspect(engine).get_columns(
                        "notification_feedback_events"
                    )
                }
                self.assertIn("delivery_reference_id", feedback_columns)
                self.assertIn("provider_message_id", feedback_columns)
                rate_limit_indexes = {
                    value["name"]
                    for value in inspect(engine).get_indexes("rate_limit_buckets")
                }
                self.assertIn("ix_rate_limit_buckets_expiry", rate_limit_indexes)
                seat_columns = {
                    value["name"]
                    for value in inspect(engine).get_columns("organization_seats")
                }
                self.assertIn("assignment_generation", seat_columns)
                seat_indexes = {
                    value["name"]
                    for value in inspect(engine).get_indexes("organization_seats")
                }
                self.assertIn(
                    "uq_organization_seats_license_assigned_user",
                    seat_indexes,
                )
                privacy_request_columns = {
                    value["name"]
                    for value in inspect(engine).get_columns("privacy_requests")
                }
                self.assertTrue(
                    {
                        "request_digest",
                        "state_generation",
                        "completed_by_user_id",
                        "completion_request_digest",
                        "completion_summary_json",
                    }.issubset(privacy_request_columns)
                )
                privacy_request_constraints = {
                    value["name"]
                    for value in inspect(engine).get_check_constraints(
                        "privacy_requests"
                    )
                }
                self.assertTrue(
                    {
                        "ck_privacy_requests_type",
                        "ck_privacy_requests_status",
                        "ck_privacy_requests_generation",
                        "ck_privacy_requests_lifecycle",
                    }.issubset(privacy_request_constraints)
                )
                privacy_request_indexes = {
                    value["name"]
                    for value in inspect(engine).get_indexes("privacy_requests")
                }
                self.assertIn(
                    "uq_privacy_requests_pending_user",
                    privacy_request_indexes,
                )
            finally:
                engine.dispose()

            command.downgrade(config, "base")
            engine = create_engine(database_url)
            try:
                tables = set(inspect(engine).get_table_names())
                self.assertFalse(EXPECTED_DOMAIN_TABLES.intersection(tables))
            finally:
                engine.dispose()

            command.upgrade(config, "head")
            engine = create_engine(database_url)
            try:
                self.assertTrue(
                    EXPECTED_DOMAIN_TABLES.issubset(
                        set(inspect(engine).get_table_names())
                    )
                )
            finally:
                engine.dispose()

    def test_postgresql_offline_ddl_compiles_to_head(self) -> None:
        output = StringIO()
        command.upgrade(
            _config(
                "postgresql+psycopg://licensing:unused@database/licensing",
                output_buffer=output,
            ),
            "head",
            sql=True,
        )
        ddl = output.getvalue().lower()
        self.assertIn("create table licenses", ddl)
        self.assertIn("create table activations", ddl)
        self.assertIn("provider_event_created_at", ddl)
        self.assertIn("provider_customer_id", ddl)
        self.assertIn("billing_hold_status", ddl)
        self.assertIn("ck_subscriptions_billing_hold", ddl)
        self.assertIn("create table organization_invitations", ddl)
        self.assertIn("create table organization_seats", ddl)
        self.assertIn("create table offline_requests", ddl)
        self.assertIn("create table offline_certificates", ddl)
        self.assertIn("create table notification_deliveries", ddl)
        self.assertIn("create table notification_feedback_events", ddl)
        self.assertIn("create table notification_suppressions", ddl)
        self.assertIn("revoked_by_user_id", ddl)
        self.assertIn("ck_serial_batches_status", ddl)
        self.assertIn("ck_serial_batches_revocation_lifecycle", ddl)
        self.assertIn("ck_serials_status", ddl)
        self.assertIn("activated_at", ddl)
        self.assertIn("retired_at", ddl)
        self.assertIn("ck_signing_key_metadata_status", ddl)
        self.assertIn("ck_signing_key_metadata_lifecycle", ddl)
        self.assertIn("uq_signing_key_metadata_active_purpose", ddl)
        self.assertIn("create table catalog_releases", ddl)
        self.assertIn("ck_catalog_releases_lifecycle", ddl)
        self.assertIn("uq_catalog_releases_active_product", ddl)
        self.assertIn("where status = 'active'", ddl)
        self.assertIn("create table rate_limit_buckets", ddl)
        self.assertIn("ck_rate_limit_buckets_window", ddl)
        self.assertIn("ck_notification_deliveries_status", ddl)
        self.assertIn("ck_notification_feedback_type", ddl)
        self.assertIn("ck_notification_suppression_lifecycle", ddl)
        self.assertIn("uq_organization_seats_license_assigned_user", ddl)
        self.assertIn("where status = 'assigned'", ddl)
        self.assertIn("create table privacy_requests", ddl)
        self.assertIn("ck_privacy_requests_lifecycle", ddl)
        self.assertIn("uq_privacy_requests_pending_user", ddl)
        self.assertIn("where status = 'pending'", ddl)
        self.assertIn("create table product_surface_profiles", ddl)
        self.assertIn("ck_product_surface_profiles_status", ddl)
        self.assertIn("surface_profile_id", ddl)
        self.assertIn("trial_duration_hours", ddl)
        self.assertIn("ck_serial_batches_trial_duration_hours", ddl)
        self.assertIn("constraint uq_idempotency_scope_subject_key unique", ddl)


if __name__ == "__main__":
    unittest.main()
