"""Signing-key registry, guarded activation, and rotation regression tests."""

from __future__ import annotations

from datetime import datetime, timedelta, timezone
import json
import unittest

from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
from sqlalchemy import create_engine, func, select
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool

from licensing_server.app.database import Base
from licensing_server.app.errors import ServerErrorCode, ServerLicensingError
from licensing_server.app.models import AuditEvent, SigningKeyMetadata, User
from licensing_server.app.security import Ed25519SnapshotSigner
from licensing_server.app.signing_key_administration import (
    RegistryBoundSnapshotSigner,
    SIGNING_KEY_PURPOSE,
    SigningKeyAdministrationService,
    inspect_signing_key_registry,
)


TEST_NOW = datetime(2026, 8, 16, 22, 0, tzinfo=timezone.utc)


def _signer(key_id: str) -> Ed25519SnapshotSigner:
    return Ed25519SnapshotSigner(key_id, Ed25519PrivateKey.generate())


class SigningKeyAdministrationTests(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)
        with Session(self.engine) as session, session.begin():
            session.add_all(
                (
                    User(
                        id="user.test.signing_admin",
                        external_issuer="https://identity.test",
                        external_subject="signing-admin",
                        verified_email="signing-admin@example.test",
                        status="active",
                        is_server_admin=True,
                    ),
                    User(
                        id="user.test.signing_customer",
                        external_issuer="https://identity.test",
                        external_subject="signing-customer",
                        verified_email="customer@example.test",
                        status="active",
                        is_server_admin=False,
                    ),
                )
            )
        self.signer = _signer("license.test.signing.primary")

    def tearDown(self) -> None:
        self.engine.dispose()

    def _service(
        self,
        session: Session,
        signer: Ed25519SnapshotSigner | None = None,
        now: datetime = TEST_NOW,
    ) -> SigningKeyAdministrationService:
        return SigningKeyAdministrationService(
            session,
            self.signer if signer is None else signer,
            now_factory=lambda: now,
        )

    def _register(
        self,
        session: Session,
        signer: Ed25519SnapshotSigner | None = None,
        suffix: str = "primary",
        now: datetime = TEST_NOW,
    ) -> dict[str, object]:
        return self._service(session, signer, now).register_current_key(
            TEST_NOW - timedelta(days=1),
            "user.test.signing_admin",
            "initial_provisioning" if suffix == "primary" else "planned_rotation",
            f"idempotency.signing_register.{suffix}",
            f"correlation.signing_register.{suffix}",
            note=f"Register {suffix} public metadata",
        )

    def test_registration_requires_admin_and_rejects_key_id_substitution(self) -> None:
        with Session(self.engine) as session:
            with self.assertRaises(ServerLicensingError) as denied:
                self._service(session).register_current_key(
                    TEST_NOW - timedelta(days=1),
                    "user.test.signing_customer",
                    "initial_provisioning",
                    "idempotency.signing_register.denied",
                    "correlation.signing_register.denied",
                )
            self.assertEqual(denied.exception.code, ServerErrorCode.AUTHORIZATION_DENIED)

            registered = self._register(session)
            self.assertTrue(registered["registered"])
            self.assertEqual(registered["status"], "staged")
            replay = self._register(session)
            self.assertTrue(replay["idempotentReplay"])

            substituted = Ed25519SnapshotSigner(
                self.signer.key_id,
                Ed25519PrivateKey.generate(),
            )
            with self.assertRaises(ServerLicensingError) as conflict:
                self._service(session, substituted).register_current_key(
                    TEST_NOW - timedelta(days=1),
                    "user.test.signing_admin",
                    "support_correction",
                    "idempotency.signing_register.substitution",
                    "correlation.signing_register.substitution",
                )
            self.assertEqual(conflict.exception.code, ServerErrorCode.CONFLICT)

        with Session(self.engine) as session:
            self.assertEqual(
                session.scalar(select(func.count()).select_from(SigningKeyMetadata)),
                1,
            )
            self.assertEqual(
                session.scalar(select(func.count()).select_from(AuditEvent)),
                1,
            )
            audit_document = json.dumps(
                session.scalars(select(AuditEvent)).one().metadata_json,
                sort_keys=True,
            )
            self.assertNotIn("publicKey\"", audit_document)
            self.assertIn("publicKeySha256", audit_document)

    def test_initial_activation_is_state_bound_idempotent_and_ready(self) -> None:
        with Session(self.engine) as session:
            self._register(session)
            service = self._service(session)
            preview = service.preview_current_key_activation(
                "user.test.signing_admin",
                "initial_provisioning",
                "correlation.signing_activate.primary",
            )
            self.assertTrue(preview["canActivateConfiguredKey"])
            self.assertFalse(preview["issuanceReady"])
            activated = service.activate_current_key(
                "user.test.signing_admin",
                "initial_provisioning",
                str(preview["stateDigest"]),
                "idempotency.signing_activate.primary",
                "correlation.signing_activate.primary",
            )
            self.assertTrue(activated["issuanceReady"])
            self.assertEqual(activated["activeKeyId"], self.signer.key_id)
            self.assertIsNone(activated["previousActiveKeyId"])
            replay = service.activate_current_key(
                "user.test.signing_admin",
                "initial_provisioning",
                str(preview["stateDigest"]),
                "idempotency.signing_activate.primary",
                "correlation.signing_activate.primary",
            )
            self.assertTrue(replay["idempotentReplay"])

        with Session(self.engine) as session:
            status = inspect_signing_key_registry(session, self.signer, TEST_NOW)
            self.assertTrue(status["issuanceReady"])
            self.assertEqual(status["issues"], [])
            row = session.scalars(select(SigningKeyMetadata)).one()
            self.assertEqual(row.status, "active")
            self.assertEqual(
                session.scalar(select(func.count()).select_from(AuditEvent)),
                2,
            )

    def test_rotation_retires_old_issuance_key_but_preserves_public_metadata(self) -> None:
        with Session(self.engine) as session:
            self._register(session)
            primary_service = self._service(session)
            primary_preview = primary_service.preview_current_key_activation(
                "user.test.signing_admin",
                "initial_provisioning",
                "correlation.signing_activate.primary",
            )
            primary_service.activate_current_key(
                "user.test.signing_admin",
                "initial_provisioning",
                str(primary_preview["stateDigest"]),
                "idempotency.signing_activate.primary",
                "correlation.signing_activate.primary",
            )

            replacement = _signer("license.test.signing.replacement")
            rotation_time = TEST_NOW + timedelta(days=7)
            self._register(
                session,
                replacement,
                suffix="replacement",
                now=rotation_time,
            )
            replacement_service = self._service(
                session,
                replacement,
                rotation_time,
            )
            preview = replacement_service.preview_current_key_activation(
                "user.test.signing_admin",
                "planned_rotation",
                "correlation.signing_activate.replacement",
            )
            self.assertEqual(preview["activeKeyId"], self.signer.key_id)
            self.assertTrue(preview["canActivateConfiguredKey"])
            result = replacement_service.activate_current_key(
                "user.test.signing_admin",
                "planned_rotation",
                str(preview["stateDigest"]),
                "idempotency.signing_activate.replacement",
                "correlation.signing_activate.replacement",
                note="Desktop overlap confirmed",
            )
            self.assertEqual(result["activeKeyId"], replacement.key_id)
            self.assertEqual(result["previousActiveKeyId"], self.signer.key_id)

        with Session(self.engine) as session:
            rows = {
                row.key_id: row
                for row in session.scalars(
                    select(SigningKeyMetadata).order_by(SigningKeyMetadata.key_id)
                )
            }
            self.assertEqual(rows[self.signer.key_id].status, "retired")
            self.assertEqual(
                rows[self.signer.key_id].expires_at.replace(tzinfo=timezone.utc),
                rotation_time,
            )
            self.assertEqual(rows[replacement.key_id].status, "active")
            self.assertEqual(len(rows), 2)
            old_status = inspect_signing_key_registry(session, self.signer, rotation_time)
            self.assertFalse(old_status["issuanceReady"])
            self.assertIn("configured_key_not_active", old_status["issues"])

    def test_activation_rejects_stale_registry_digest_and_changed_retry(self) -> None:
        replacement = _signer("license.test.signing.second")
        third = _signer("license.test.signing.third")
        with Session(self.engine) as session:
            self._register(session, replacement, suffix="second")
            replacement_service = self._service(session, replacement)
            preview = replacement_service.preview_current_key_activation(
                "user.test.signing_admin",
                "planned_rotation",
                "correlation.signing_activate.second",
            )
            self._register(session, third, suffix="third")
            with self.assertRaises(ServerLicensingError) as stale:
                replacement_service.activate_current_key(
                    "user.test.signing_admin",
                    "planned_rotation",
                    str(preview["stateDigest"]),
                    "idempotency.signing_activate.second",
                    "correlation.signing_activate.second",
                )
            self.assertEqual(stale.exception.code, ServerErrorCode.CONFLICT)

            refreshed = replacement_service.preview_current_key_activation(
                "user.test.signing_admin",
                "planned_rotation",
                "correlation.signing_activate.second",
            )
            replacement_service.activate_current_key(
                "user.test.signing_admin",
                "planned_rotation",
                str(refreshed["stateDigest"]),
                "idempotency.signing_activate.second.refreshed",
                "correlation.signing_activate.second",
            )
            with self.assertRaises(ServerLicensingError) as changed_retry:
                replacement_service.activate_current_key(
                    "user.test.signing_admin",
                    "support_correction",
                    str(refreshed["stateDigest"]),
                    "idempotency.signing_activate.second.refreshed",
                    "correlation.signing_activate.second",
                )
            self.assertEqual(
                changed_retry.exception.code,
                ServerErrorCode.IDEMPOTENCY_CONFLICT,
            )

        with Session(self.engine) as session:
            rows = session.scalars(
                select(SigningKeyMetadata).where(
                    SigningKeyMetadata.purpose == SIGNING_KEY_PURPOSE
                )
            ).all()
            self.assertEqual(sum(row.status == "active" for row in rows), 1)

    def test_active_key_compromise_is_reviewed_idempotent_and_stops_issuance(self) -> None:
        with Session(self.engine) as session:
            self._register(session)
            service = self._service(session)
            activation_preview = service.preview_current_key_activation(
                "user.test.signing_admin",
                "initial_provisioning",
                "correlation.signing_activate.primary",
            )
            service.activate_current_key(
                "user.test.signing_admin",
                "initial_provisioning",
                str(activation_preview["stateDigest"]),
                "idempotency.signing_activate.primary",
                "correlation.signing_activate.primary",
            )
            with self.assertRaises(ServerLicensingError) as denied:
                service.preview_key_compromise(
                    self.signer.key_id,
                    "user.test.signing_customer",
                    "compromise_response",
                    "correlation.signing_compromise.denied",
                )
            self.assertEqual(denied.exception.code, ServerErrorCode.AUTHORIZATION_DENIED)
            preview = service.preview_key_compromise(
                self.signer.key_id,
                "user.test.signing_admin",
                "compromise_response",
                "correlation.signing_compromise.primary",
                note="Confirmed signing-key exposure",
            )
            self.assertTrue(preview["canCompromise"])
            self.assertTrue(preview["wouldStopIssuance"])
            with self.assertRaises(ServerLicensingError) as stale:
                service.compromise_key(
                    self.signer.key_id,
                    "user.test.signing_admin",
                    "compromise_response",
                    "00" * 32,
                    "idempotency.signing_compromise.stale",
                    "correlation.signing_compromise.primary",
                )
            self.assertEqual(stale.exception.code, ServerErrorCode.CONFLICT)
            result = service.compromise_key(
                self.signer.key_id,
                "user.test.signing_admin",
                "compromise_response",
                str(preview["stateDigest"]),
                "idempotency.signing_compromise.primary",
                "correlation.signing_compromise.primary",
                note="Confirmed signing-key exposure",
            )
            self.assertTrue(result["executed"])
            self.assertTrue(result["issuanceStopped"])
            self.assertFalse(result["issuanceReady"])
            self.assertIsNone(result["activeKeyId"])
            self.assertEqual(result["compromisedStatus"], "compromised")
            with self.assertRaises(ServerLicensingError) as signing_stopped:
                RegistryBoundSnapshotSigner(
                    session,
                    self.signer,
                    now_factory=lambda: TEST_NOW,
                ).sign_payload({"test": "must-not-sign"})
            self.assertEqual(
                signing_stopped.exception.code,
                ServerErrorCode.SIGNING_UNAVAILABLE,
            )
            session.rollback()
            replay = service.compromise_key(
                self.signer.key_id,
                "user.test.signing_admin",
                "compromise_response",
                str(preview["stateDigest"]),
                "idempotency.signing_compromise.primary",
                "correlation.signing_compromise.primary",
                note="Confirmed signing-key exposure",
            )
            self.assertTrue(replay["idempotentReplay"])
            with self.assertRaises(ServerLicensingError) as changed_retry:
                service.compromise_key(
                    self.signer.key_id,
                    "user.test.signing_admin",
                    "compromise_response",
                    str(preview["stateDigest"]),
                    "idempotency.signing_compromise.primary",
                    "correlation.signing_compromise.primary",
                    note="Changed incident narrative",
                )
            self.assertEqual(
                changed_retry.exception.code,
                ServerErrorCode.IDEMPOTENCY_CONFLICT,
            )

        with Session(self.engine) as session:
            row = session.scalars(select(SigningKeyMetadata)).one()
            self.assertEqual(row.status, "compromised")
            self.assertEqual(
                row.expires_at.replace(tzinfo=timezone.utc),
                TEST_NOW,
            )
            self.assertEqual(
                row.retired_at.replace(tzinfo=timezone.utc),
                TEST_NOW,
            )
            compromise_audits = session.scalars(
                select(AuditEvent).where(
                    AuditEvent.action == "signing_key.compromised"
                )
            ).all()
            self.assertEqual(len(compromise_audits), 1)
            audit_document = json.dumps(
                compromise_audits[0].metadata_json,
                sort_keys=True,
            )
            self.assertNotIn("publicKey", audit_document)
            self.assertIn("issuanceStopped", audit_document)

    def test_future_staged_key_can_be_terminally_compromised(self) -> None:
        future_not_before = TEST_NOW + timedelta(days=30)
        with Session(self.engine) as session:
            service = self._service(session)
            service.register_current_key(
                future_not_before,
                "user.test.signing_admin",
                "planned_rotation",
                "idempotency.signing_register.future",
                "correlation.signing_register.future",
            )
            preview = service.preview_key_compromise(
                self.signer.key_id,
                "user.test.signing_admin",
                "compromise_response",
                "correlation.signing_compromise.future",
            )
            self.assertEqual(preview["targetStatus"], "staged")
            self.assertFalse(preview["wouldStopIssuance"])
            result = service.compromise_key(
                self.signer.key_id,
                "user.test.signing_admin",
                "compromise_response",
                str(preview["stateDigest"]),
                "idempotency.signing_compromise.future",
                "correlation.signing_compromise.future",
            )
            self.assertFalse(result["issuanceStopped"])
            self.assertNotIn("invalid_key_validity_window", result["issues"])

        with Session(self.engine) as session:
            row = session.scalars(select(SigningKeyMetadata)).one()
            self.assertEqual(row.status, "compromised")
            self.assertGreater(
                row.not_before.replace(tzinfo=timezone.utc),
                row.expires_at.replace(tzinfo=timezone.utc),
            )


if __name__ == "__main__":
    unittest.main()
