"""Production desktop client to live TLS licensing-server journey."""

from __future__ import annotations

from dataclasses import replace
from datetime import datetime, timedelta, timezone
import ipaddress
import os
from pathlib import Path
import secrets
import socket
import ssl
import tempfile
import threading
import time
import unittest
from unittest.mock import patch
from urllib.request import HTTPSHandler, ProxyHandler, build_opener

from cryptography import x509
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import ec, rsa
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
from cryptography.x509.oid import NameOID
from fastapi import Response
import jwt
from sqlalchemy import create_engine, select
from sqlalchemy.orm import Session
import uvicorn

from desktop_licensing.account_auth import AccountSession
from desktop_licensing.api_client import LicensingApiClient, LicensingApiFailure
from desktop_licensing.config import DesktopLicensingConfig
from desktop_licensing.device_credential import DpapiDeviceCredentialStore
from desktop_licensing.manager import DesktopLicenseManager
from desktop_licensing.state import ProtectedStateStore, PublicStateStore
from licensing_shared.catalog import load_builtin_catalog
from licensing_shared.models import LicenseState
from licensing_shared.paths import LicensingPaths
from licensing_shared.policy import EntitlementRequirement
from licensing_shared.verifier import PublicKeyRecord, PublicKeyRing
from licensing_server.app.api import ADMIN_SERIAL_SCOPE, create_app
from licensing_server.app.auth import AuthenticatedPrincipal, OIDCTokenValidator
from licensing_server.app.config import ServerSettings
from licensing_server.app.constants import (
    MAX_BEARER_TOKEN_CHARACTERS,
    MAX_OIDC_JWKS_KEYS,
    MAX_OIDC_JWKS_RESPONSE_BYTES,
    MAX_OIDC_SCOPE_CHARACTERS,
    MAX_OIDC_SCOPE_COUNT,
    MAX_VERIFIED_EMAIL_CHARACTERS,
    OIDC_JWKS_HTTP_TIMEOUT_SECONDS,
)
from licensing_server.app.database import Base
from licensing_server.app.models import (
    Activation,
    CatalogRelease,
    License,
    Serial,
    SigningKeyMetadata,
    User,
)
from licensing_server.app.security import Ed25519SnapshotSigner


TEST_LOOPBACK_HOST = "127.0.0.1"
TEST_SERVER_START_TIMEOUT_SECONDS = 10.0
TEST_SERVER_STOP_TIMEOUT_SECONDS = 10.0
TEST_HTTP_TIMEOUT_SECONDS = 5.0
TEST_TLS_VALIDITY_MINUTES = 30
TEST_TLS_CLOCK_TOLERANCE_MINUTES = 5
TEST_AES_KEY_BITS = 256
TEST_AES_NONCE_BYTES = 12
TEST_PROTECTOR_DOMAIN = b"APOLON-LICENSING-LOOPBACK-TEST-V1\x00"
TEST_SERIAL_PEPPER = b"loopback-serial-pepper-material-32-bytes"
TEST_FINGERPRINT_PEPPER = b"loopback-fingerprint-pepper-material-32"
TEST_SIGNING_KEY_ID = "license.test.loopback"
TEST_ADMIN_TOKEN = "loopback-admin-token"
TEST_ADMIN_USER_ID = "user.test.loopback.admin"
TEST_OIDC_ISSUER = "https://identity.loopback.test"
TEST_OIDC_AUDIENCE = "licensing-api"
TEST_OIDC_KEY_ID = "oidc.test.loopback"
TEST_OIDC_TOKEN_VALIDITY_MINUTES = 10
TEST_OIDC_EXPIRED_TOKEN_VALIDITY_MINUTES = -1
TEST_OIDC_MAXIMUM_ABSOLUTE_VALIDITY_MINUTES = 60
TEST_OIDC_RSA_PUBLIC_EXPONENT = 65537
TEST_OIDC_RSA_KEY_BITS = 2048
TEST_CUSTOMER_SUBJECT = "loopback-customer"
TEST_CUSTOMER_EMAIL = "customer@loopback.test"
TEST_UNKNOWN_OIDC_KEY_ID = "oidc.test.unknown"
TEST_MALFORMED_JWKS = b"{"
TEST_DUPLICATE_KEY_JWKS = b'{"keys":[],"keys":[]}'
TEST_OVERSIZED_JWKS = b" " * (MAX_OIDC_JWKS_RESPONSE_BYTES + 1)
TEST_FUTURE_AUTHENTICATION_MINUTES = 1
TEST_OVERSIZED_VERIFIED_EMAIL = "e" * (MAX_VERIFIED_EMAIL_CHARACTERS + 1)
TEST_CONTROL_CHARACTER_EMAIL = "customer@loopback.test\r\ninvalid-header"
TEST_EXCESSIVE_OIDC_SCOPES = tuple(
    f"test.scope.{index}" for index in range(MAX_OIDC_SCOPE_COUNT + 1)
)
TEST_OVERSIZED_OIDC_SCOPE = "s" * (MAX_OIDC_SCOPE_CHARACTERS + 1)
TEST_OVERSIZED_BEARER_TOKEN = "x" * (MAX_BEARER_TOKEN_CHARACTERS + 1)


class _MemoryProtector:
    """Temporary authenticated storage envelope for a non-Windows test host."""

    def __init__(self) -> None:
        self._cipher = AESGCM(AESGCM.generate_key(bit_length=TEST_AES_KEY_BITS))

    def protect(self, plaintext: bytes) -> bytes:
        if not isinstance(plaintext, bytes) or not plaintext:
            raise ValueError("plaintext must be non-empty bytes")
        nonce = secrets.token_bytes(TEST_AES_NONCE_BYTES)
        return nonce + self._cipher.encrypt(
            nonce,
            plaintext,
            TEST_PROTECTOR_DOMAIN,
        )

    def unprotect(self, protected: bytes) -> bytes:
        if not isinstance(protected, bytes):
            raise TypeError("protected must be bytes")
        if len(protected) <= TEST_AES_NONCE_BYTES:
            raise ValueError("protected value is too short")
        nonce = protected[:TEST_AES_NONCE_BYTES]
        return self._cipher.decrypt(
            nonce,
            protected[TEST_AES_NONCE_BYTES:],
            TEST_PROTECTOR_DOMAIN,
        )


class _AdminTokenValidator:
    def validate(self, token: str) -> AuthenticatedPrincipal:
        if not isinstance(token, str):
            raise TypeError("token must be a string")
        if token != TEST_ADMIN_TOKEN:
            raise PermissionError("test token is not authorized")
        return AuthenticatedPrincipal(
            issuer=TEST_OIDC_ISSUER,
            subject="loopback-admin",
            scopes=frozenset((ADMIN_SERIAL_SCOPE,)),
            verified_email="admin@loopback.test",
            authentication_time=datetime.now(timezone.utc),
        )


class _ExistingAccountAuthorizer:
    """Existing in-memory account session for the production manager boundary."""

    def __init__(self, access_token: str) -> None:
        if not isinstance(access_token, str) or not access_token.strip():
            raise ValueError("access_token must be a non-empty string")
        self._access_token: str | None = access_token.strip()

    @property
    def is_signed_in(self) -> bool:
        return self._access_token is not None

    def sign_in(self) -> AccountSession:
        if self._access_token is None:
            raise LicensingApiFailure(
                "account_sign_in_required",
                "The test account session has been signed out.",
                status_code=401,
            )
        return AccountSession(
            self._access_token,
            datetime.now(timezone.utc)
            + timedelta(minutes=TEST_OIDC_TOKEN_VALIDITY_MINUTES),
        )

    def access_token(self) -> str:
        if self._access_token is None:
            raise LicensingApiFailure(
                "account_sign_in_required",
                "The test account session has been signed out.",
                status_code=401,
            )
        return self._access_token

    def current_access_token(self) -> str | None:
        return self._access_token

    def sign_out(self) -> None:
        self._access_token = None


class DesktopServerLoopbackTests(unittest.TestCase):
    def setUp(self) -> None:
        self.temporary_directory = tempfile.TemporaryDirectory()
        self.addCleanup(self.temporary_directory.cleanup)
        self.root = Path(self.temporary_directory.name)
        self.database_path = self.root / "licensing.sqlite3"
        self.database_url = f"sqlite+pysqlite:///{self.database_path.as_posix()}"
        self.engine = create_engine(
            self.database_url,
            connect_args={"check_same_thread": False},
        )
        self.addCleanup(self.engine.dispose)
        Base.metadata.create_all(self.engine)
        self.signer = Ed25519SnapshotSigner(
            TEST_SIGNING_KEY_ID,
            Ed25519PrivateKey.generate(),
        )
        self.oidc_private_key = rsa.generate_private_key(
            public_exponent=TEST_OIDC_RSA_PUBLIC_EXPONENT,
            key_size=TEST_OIDC_RSA_KEY_BITS,
        )
        self.oidc_jwk = jwt.algorithms.RSAAlgorithm.to_jwk(
            self.oidc_private_key.public_key(),
            as_dict=True,
        )
        self.oidc_jwk.update(
            {
                "alg": "RS256",
                "kid": TEST_OIDC_KEY_ID,
                "use": "sig",
            }
        )
        self.settings = ServerSettings(
            database_url=self.database_url,
            environment_name="test",
            signing_key_id=self.signer.key_id,
            signing_private_key_path=self.root / "unused-signing-key",
            serial_pepper=TEST_SERIAL_PEPPER,
            fingerprint_pepper=TEST_FINGERPRINT_PEPPER,
            oidc_issuer=TEST_OIDC_ISSUER,
            oidc_audience=TEST_OIDC_AUDIENCE,
            oidc_jwks_url=f"{TEST_OIDC_ISSUER}/.well-known/jwks.json",
            allow_development_signer=False,
        )
        self.app = create_app(
            self.settings,
            engine=self.engine,
            signer=self.signer,
            token_validator=_AdminTokenValidator(),
        )
        self.oidc_jwks_request_count = 0

        def oidc_jwks() -> dict[str, object]:
            self.oidc_jwks_request_count += 1
            return {"keys": [dict(self.oidc_jwk)]}

        def malformed_oidc_jwks() -> Response:
            return Response(
                content=TEST_MALFORMED_JWKS,
                media_type="application/json",
            )

        def non_object_oidc_jwks() -> list[object]:
            return []

        def empty_oidc_jwks() -> dict[str, object]:
            return {"keys": []}

        def duplicate_key_oidc_jwks() -> Response:
            return Response(
                content=TEST_DUPLICATE_KEY_JWKS,
                media_type="application/json",
            )

        def oversized_oidc_jwks() -> Response:
            return Response(
                content=TEST_OVERSIZED_JWKS,
                media_type="application/json",
            )

        def excessive_key_oidc_jwks() -> dict[str, object]:
            return {
                "keys": [
                    dict(self.oidc_jwk)
                    for _index in range(MAX_OIDC_JWKS_KEYS + 1)
                ]
            }

        self.app.add_api_route(
            "/identity/jwks.json",
            oidc_jwks,
            methods=["GET"],
        )
        self.app.add_api_route(
            "/identity/malformed-jwks.json",
            malformed_oidc_jwks,
            methods=["GET"],
        )
        self.app.add_api_route(
            "/identity/non-object-jwks.json",
            non_object_oidc_jwks,
            methods=["GET"],
        )
        self.app.add_api_route(
            "/identity/empty-jwks.json",
            empty_oidc_jwks,
            methods=["GET"],
        )
        self.app.add_api_route(
            "/identity/duplicate-key-jwks.json",
            duplicate_key_oidc_jwks,
            methods=["GET"],
        )
        self.app.add_api_route(
            "/identity/oversized-jwks.json",
            oversized_oidc_jwks,
            methods=["GET"],
        )
        self.app.add_api_route(
            "/identity/excessive-key-jwks.json",
            excessive_key_oidc_jwks,
            methods=["GET"],
        )
        self.catalog = load_builtin_catalog()
        now = datetime.now(timezone.utc)
        with Session(self.engine) as session, session.begin():
            session.add_all(
                (
                    User(
                        id=TEST_ADMIN_USER_ID,
                        external_issuer=TEST_OIDC_ISSUER,
                        external_subject="loopback-admin",
                        verified_email="admin@loopback.test",
                        status="active",
                        is_server_admin=True,
                    ),
                    CatalogRelease(
                        id="catalog_release.test.loopback",
                        product_id=self.catalog.product_id,
                        revision=self.catalog.revision,
                        document_json=self.catalog.to_mapping(),
                        document_digest=bytes.fromhex(self.catalog.sha256()),
                        status="active",
                        staged_by_user_id=TEST_ADMIN_USER_ID,
                        staged_at=now - timedelta(minutes=1),
                        published_by_user_id=TEST_ADMIN_USER_ID,
                        published_at=now - timedelta(minutes=1),
                        retired_at=None,
                    ),
                    SigningKeyMetadata(
                        id="signing_key.test.loopback",
                        key_id=self.signer.key_id,
                        purpose="license_snapshot",
                        public_key_bytes=self.signer.public_key_bytes(),
                        status="active",
                        not_before=now - timedelta(minutes=1),
                        expires_at=None,
                        activated_at=now - timedelta(minutes=1),
                        retired_at=None,
                    ),
                )
            )
        self.certificate_path, self.private_key_path = self._create_tls_identity()
        self.server, self.server_thread, port = self._start_server()
        self.addCleanup(self._stop_server)
        self.api_url = f"https://{TEST_LOOPBACK_HOST}:{port}"
        self.oidc_settings = replace(
            self.settings,
            oidc_jwks_url=f"{self.api_url}/identity/jwks.json",
        )
        context = ssl.create_default_context(cafile=str(self.certificate_path))
        self.opener = build_opener(
            ProxyHandler({}),
            HTTPSHandler(context=context),
        )

    def _create_tls_identity(self) -> tuple[Path, Path]:
        private_key = ec.generate_private_key(ec.SECP256R1())
        now = datetime.now(timezone.utc)
        subject = x509.Name(
            (x509.NameAttribute(NameOID.COMMON_NAME, "Apolon loopback test"),)
        )
        certificate = (
            x509.CertificateBuilder()
            .subject_name(subject)
            .issuer_name(subject)
            .public_key(private_key.public_key())
            .serial_number(x509.random_serial_number())
            .not_valid_before(
                now - timedelta(minutes=TEST_TLS_CLOCK_TOLERANCE_MINUTES)
            )
            .not_valid_after(now + timedelta(minutes=TEST_TLS_VALIDITY_MINUTES))
            .add_extension(
                x509.SubjectAlternativeName(
                    (x509.IPAddress(ipaddress.ip_address(TEST_LOOPBACK_HOST)),)
                ),
                critical=False,
            )
            .sign(private_key, hashes.SHA256())
        )
        certificate_path = self.root / "loopback-certificate.pem"
        private_key_path = self.root / "loopback-private-key.pem"
        certificate_path.write_bytes(certificate.public_bytes(serialization.Encoding.PEM))
        private_key_path.write_bytes(
            private_key.private_bytes(
                encoding=serialization.Encoding.PEM,
                format=serialization.PrivateFormat.PKCS8,
                encryption_algorithm=serialization.NoEncryption(),
            )
        )
        private_key_path.chmod(0o600)
        return certificate_path, private_key_path

    def _start_server(self) -> tuple[uvicorn.Server, threading.Thread, int]:
        listener = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
        listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
        listener.bind((TEST_LOOPBACK_HOST, 0))
        listener.listen()
        port = listener.getsockname()[1]
        config = uvicorn.Config(
            self.app,
            log_level="critical",
            access_log=False,
            lifespan="off",
            ssl_certfile=str(self.certificate_path),
            ssl_keyfile=str(self.private_key_path),
            timeout_graceful_shutdown=TEST_SERVER_STOP_TIMEOUT_SECONDS,
        )
        server = uvicorn.Server(config)
        thread = threading.Thread(
            target=server.run,
            kwargs={"sockets": (listener,)},
            name="licensing-loopback-test-server",
            daemon=True,
        )
        thread.start()
        deadline = time.monotonic() + TEST_SERVER_START_TIMEOUT_SECONDS
        while not server.started and time.monotonic() < deadline:
            if not thread.is_alive():
                self.fail("loopback licensing server stopped during startup")
            time.sleep(0.01)
        if not server.started:
            self.fail("loopback licensing server did not start in time")
        return server, thread, port

    def _stop_server(self) -> None:
        self.server.should_exit = True
        self.server_thread.join(TEST_SERVER_STOP_TIMEOUT_SECONDS)
        if self.server_thread.is_alive():
            self.fail("loopback licensing server did not stop in time")

    def _licensing_paths(self) -> LicensingPaths:
        root = self.root / "desktop-state"
        return LicensingPaths(
            root=root,
            entitlement_snapshot=root / "entitlement.json",
            offline_license=root / "offline-license.json",
            public_state=root / "state.json",
            protected_state=root / "protected-state.bin",
            protected_update_state=root / "protected-update-state.bin",
            update_downloads=root / "updates",
        )

    def _oidc_token(
        self,
        subject: str,
        verified_email: str,
        scopes: tuple[str, ...] = (),
        *,
        audience: str = TEST_OIDC_AUDIENCE,
        key_id: str | None = TEST_OIDC_KEY_ID,
        private_key: rsa.RSAPrivateKey | None = None,
        validity_minutes: int = TEST_OIDC_TOKEN_VALIDITY_MINUTES,
        authentication_time_offset_minutes: int = 0,
    ) -> str:
        for value, field_name in (
            (subject, "subject"),
            (verified_email, "verified_email"),
            (audience, "audience"),
        ):
            if not isinstance(value, str) or not value.strip():
                raise ValueError(f"{field_name} must be a non-empty string")
        if key_id is not None and (
            not isinstance(key_id, str) or not key_id.strip()
        ):
            raise ValueError("key_id must be a non-empty string or None")
        if private_key is not None and not isinstance(
            private_key,
            rsa.RSAPrivateKey,
        ):
            raise TypeError("private_key must be an RSA private key or None")
        if (
            isinstance(validity_minutes, bool)
            or not isinstance(validity_minutes, int)
            or validity_minutes == 0
            or abs(validity_minutes)
            > TEST_OIDC_MAXIMUM_ABSOLUTE_VALIDITY_MINUTES
        ):
            raise ValueError("validity_minutes is outside the supported test range")
        if (
            isinstance(authentication_time_offset_minutes, bool)
            or not isinstance(authentication_time_offset_minutes, int)
            or abs(authentication_time_offset_minutes)
            > TEST_OIDC_MAXIMUM_ABSOLUTE_VALIDITY_MINUTES
        ):
            raise ValueError(
                "authentication_time_offset_minutes is outside the supported test range"
            )
        if not isinstance(scopes, tuple) or any(
            not isinstance(scope, str) or not scope.strip() for scope in scopes
        ):
            raise ValueError("scopes must be a tuple of non-empty strings")
        now = datetime.now(timezone.utc)
        return jwt.encode(
            {
                "aud": audience.strip(),
                "auth_time": int(
                    (
                        now
                        + timedelta(minutes=authentication_time_offset_minutes)
                    ).timestamp()
                ),
                "email": verified_email.strip(),
                "email_verified": True,
                "exp": int(
                    (
                        now + timedelta(minutes=validity_minutes)
                    ).timestamp()
                ),
                "iat": int(now.timestamp()),
                "iss": TEST_OIDC_ISSUER,
                "scope": " ".join(scope.strip() for scope in scopes),
                "sub": subject.strip(),
            },
            self.oidc_private_key if private_key is None else private_key,
            algorithm="RS256",
            headers={} if key_id is None else {"kid": key_id.strip()},
        )

    def _new_manager(
        self,
        config: DesktopLicensingConfig,
        paths: LicensingPaths,
        protector: _MemoryProtector,
        api_client: LicensingApiClient,
        *,
        account_authorizer: _ExistingAccountAuthorizer | None = None,
        evidence_suffix: str = "serial",
    ) -> DesktopLicenseManager:
        if not isinstance(config, DesktopLicensingConfig):
            raise TypeError("config must be DesktopLicensingConfig")
        if not isinstance(paths, LicensingPaths):
            raise TypeError("paths must be LicensingPaths")
        if not isinstance(protector, _MemoryProtector):
            raise TypeError("protector must be _MemoryProtector")
        if not isinstance(api_client, LicensingApiClient):
            raise TypeError("api_client must be LicensingApiClient")
        if account_authorizer is not None and not isinstance(
            account_authorizer,
            _ExistingAccountAuthorizer,
        ):
            raise TypeError("account_authorizer must be _ExistingAccountAuthorizer or None")
        if not isinstance(evidence_suffix, str) or not evidence_suffix.strip():
            raise ValueError("evidence_suffix must be a non-empty string")
        key_ring = PublicKeyRing(
            (PublicKeyRecord(self.signer.key_id, self.signer.public_key_bytes()),)
        )
        return DesktopLicenseManager(
            config=config,
            paths=paths,
            key_ring=key_ring,
            api_client=api_client,
            account_authorizer=account_authorizer,
            credential_store=DpapiDeviceCredentialStore(
                paths.root / "device-key.bin",
                protector,
            ),
            public_state_store=PublicStateStore(paths.public_state),
            protected_state_store=ProtectedStateStore(
                paths.protected_state,
                protector,
            ),
            evidence_collector=lambda: {
                "system_uuid": f"loopback-{evidence_suffix.strip()}-system-uuid",
                "baseboard_serial": f"loopback-{evidence_suffix.strip()}-baseboard",
            },
        )

    def test_serial_activation_restart_refresh_authorization_and_release_over_tls(
        self,
    ) -> None:
        config = DesktopLicensingConfig(
            api_base_url=self.api_url,
            source_developer_access=False,
            api_timeout_seconds=TEST_HTTP_TIMEOUT_SECONDS,
        )
        api_client = LicensingApiClient(config)
        protector = _MemoryProtector()
        paths = self._licensing_paths()

        with patch("desktop_licensing.api_client.urlopen", new=self.opener.open):
            readiness = api_client.get("/v1/readiness")
            self.assertEqual(readiness["status"], "ready")
            self.assertEqual(readiness["catalogSha256"], self.catalog.sha256())
            public_catalog = api_client.get("/v1/catalog/public")
            self.assertEqual(public_catalog, self.catalog.to_mapping())
            generated = api_client.post(
                "/v1/admin/serial-batches",
                {
                    "skuId": "edition.professional",
                    "quantity": 1,
                    "reason": "Desktop/server TLS loopback qualification",
                    "correlationId": "correlation.test.loopback.serial",
                },
                access_token=TEST_ADMIN_TOKEN,
            )
            serial = generated["serials"][0]

            manager = self._new_manager(config, paths, protector, api_client)
            self.assertEqual(manager.initialize().state, LicenseState.MISSING)
            activated = manager.activate_serial(serial)
            self.assertEqual(activated.state, LicenseState.LICENSED_ACTIVE)
            self.assertTrue(paths.entitlement_snapshot.is_file())
            decision = manager.authorize(
                EntitlementRequirement(
                    access_id="analysis.advanced",
                    all_of=("feature.analysis.advanced",),
                )
            )
            self.assertTrue(decision.allowed)

            restarted = self._new_manager(config, paths, protector, api_client)
            self.assertEqual(
                restarted.initialize().state,
                LicenseState.LICENSED_ACTIVE,
            )
            self.assertEqual(
                restarted.refresh().state,
                LicenseState.LICENSED_ACTIVE,
            )
            self.assertEqual(restarted.deactivate().state, LicenseState.MISSING)
            self.assertFalse(paths.entitlement_snapshot.exists())

        with Session(self.engine) as session:
            stored_serial = session.scalar(select(Serial))
            activation = session.scalar(select(Activation))
            self.assertIsNotNone(stored_serial)
            self.assertEqual(stored_serial.status, "redeemed")
            self.assertIsNotNone(activation)
            self.assertEqual(activation.status, "deactivated")
        serial_bytes = serial.encode("ascii")
        for path in self.root.rglob("*"):
            if path.is_file():
                self.assertNotIn(serial_bytes, path.read_bytes())

    def test_production_oidc_account_recovery_and_scoped_administration_over_tls(
        self,
    ) -> None:
        self.app.state.settings = self.oidc_settings
        self.app.state.token_validator = OIDCTokenValidator(self.oidc_settings)
        admin_token = self._oidc_token(
            "loopback-admin",
            "admin@loopback.test",
            (ADMIN_SERIAL_SCOPE,),
        )
        customer_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_CUSTOMER_EMAIL,
        )
        wrong_audience_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_CUSTOMER_EMAIL,
            audience="wrong-licensing-audience",
        )
        unknown_key_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_CUSTOMER_EMAIL,
            key_id=TEST_UNKNOWN_OIDC_KEY_ID,
        )
        missing_key_id_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_CUSTOMER_EMAIL,
            key_id=None,
        )
        wrong_signature_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_CUSTOMER_EMAIL,
            private_key=rsa.generate_private_key(
                public_exponent=TEST_OIDC_RSA_PUBLIC_EXPONENT,
                key_size=TEST_OIDC_RSA_KEY_BITS,
            ),
        )
        expired_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_CUSTOMER_EMAIL,
            validity_minutes=TEST_OIDC_EXPIRED_TOKEN_VALIDITY_MINUTES,
        )
        oversized_email_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_OVERSIZED_VERIFIED_EMAIL,
        )
        control_character_email_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_CONTROL_CHARACTER_EMAIL,
        )
        excessive_scopes_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_CUSTOMER_EMAIL,
            TEST_EXCESSIVE_OIDC_SCOPES,
        )
        oversized_scope_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_CUSTOMER_EMAIL,
            (TEST_OVERSIZED_OIDC_SCOPE,),
        )
        future_authentication_token = self._oidc_token(
            TEST_CUSTOMER_SUBJECT,
            TEST_CUSTOMER_EMAIL,
            authentication_time_offset_minutes=TEST_FUTURE_AUTHENTICATION_MINUTES,
        )
        config = DesktopLicensingConfig(
            api_base_url=self.api_url,
            source_developer_access=False,
            api_timeout_seconds=TEST_HTTP_TIMEOUT_SECONDS,
        )
        api_client = LicensingApiClient(config)
        protector = _MemoryProtector()
        paths = self._licensing_paths()
        account_authorizer = _ExistingAccountAuthorizer(customer_token)

        with patch.dict(
            os.environ,
            {
                "NO_PROXY": TEST_LOOPBACK_HOST,
                "SSL_CERT_FILE": str(self.certificate_path),
                "no_proxy": TEST_LOOPBACK_HOST,
            },
        ), patch("desktop_licensing.api_client.urlopen", new=self.opener.open):
            jwks_requests_before_oversized_token = self.oidc_jwks_request_count
            with self.assertRaises(LicensingApiFailure) as captured:
                api_client.get(
                    "/v1/me/licenses",
                    access_token=TEST_OVERSIZED_BEARER_TOKEN,
                )
            self.assertEqual(captured.exception.status_code, 401)
            self.assertEqual(captured.exception.code, "authentication_required")
            self.assertFalse(captured.exception.retryable)
            self.assertEqual(
                self.oidc_jwks_request_count,
                jwks_requests_before_oversized_token,
            )

            for invalid_token_kind, invalid_token in (
                ("wrong_audience", wrong_audience_token),
                ("missing_key_id", missing_key_id_token),
                ("wrong_signature", wrong_signature_token),
                ("expired", expired_token),
                ("oversized_email", oversized_email_token),
                ("control_character_email", control_character_email_token),
                ("excessive_scopes", excessive_scopes_token),
                ("oversized_scope", oversized_scope_token),
                ("future_authentication", future_authentication_token),
            ):
                with self.subTest(invalid_token_kind=invalid_token_kind):
                    with self.assertRaises(LicensingApiFailure) as captured:
                        api_client.get(
                            "/v1/me/licenses",
                            access_token=invalid_token,
                        )
                    self.assertEqual(captured.exception.status_code, 401)
                    self.assertEqual(
                        captured.exception.code,
                        "authentication_required",
                    )
                    self.assertFalse(captured.exception.retryable)

            with self.assertRaises(LicensingApiFailure) as captured:
                api_client.post(
                    "/v1/admin/serial-batches",
                    {
                        "skuId": "edition.professional",
                        "quantity": 1,
                        "reason": "Rejected unscoped OIDC TLS request",
                        "correlationId": "correlation.test.loopback.unscoped",
                    },
                    access_token=customer_token,
                )
            self.assertEqual(captured.exception.status_code, 403)
            self.assertEqual(captured.exception.code, "authorization_denied")

            generated = api_client.post(
                "/v1/admin/serial-batches",
                {
                    "skuId": "edition.professional",
                    "quantity": 1,
                    "reason": "Production OIDC TLS loopback qualification",
                    "correlationId": "correlation.test.loopback.oidc",
                },
                access_token=admin_token,
            )
            serial = generated["serials"][0]

            manager = self._new_manager(
                config,
                paths,
                protector,
                api_client,
                account_authorizer=account_authorizer,
                evidence_suffix="oidc",
            )
            self.assertEqual(manager.initialize().state, LicenseState.MISSING)
            activated = manager.activate_serial(serial)
            self.assertEqual(activated.state, LicenseState.LICENSED_ACTIVE)
            self.assertTrue(manager.account_signed_in)
            self.assertIsNotNone(manager.verified_snapshot)
            license_id = manager.verified_snapshot.snapshot.license_id
            account_licenses = manager.list_account_licenses()
            self.assertEqual(len(account_licenses), 1)
            self.assertEqual(account_licenses[0].license_id, license_id)
            self.assertEqual(account_licenses[0].sku_ids, ("edition.professional",))

            self.assertEqual(manager.deactivate().state, LicenseState.MISSING)
            restarted = self._new_manager(
                config,
                paths,
                protector,
                api_client,
                account_authorizer=account_authorizer,
                evidence_suffix="oidc",
            )
            self.assertEqual(restarted.initialize().state, LicenseState.MISSING)
            recovered_licenses = restarted.list_account_licenses()
            self.assertEqual(recovered_licenses[0].license_id, license_id)
            recovered = restarted.activate_account_license(license_id)
            self.assertEqual(recovered.state, LicenseState.LICENSED_ACTIVE)
            self.assertTrue(
                restarted.authorize(
                    EntitlementRequirement(
                        access_id="analysis.advanced",
                        all_of=("feature.analysis.advanced",),
                    )
                ).allowed
            )
            self.assertEqual(restarted.deactivate().state, LicenseState.MISSING)

            jwks_requests_before_unknown_key = self.oidc_jwks_request_count
            with self.assertRaises(LicensingApiFailure) as captured:
                api_client.get(
                    "/v1/me/licenses",
                    access_token=unknown_key_token,
                )
            self.assertEqual(captured.exception.status_code, 401)
            self.assertEqual(captured.exception.code, "authentication_required")
            self.assertFalse(captured.exception.retryable)
            self.assertEqual(
                self.oidc_jwks_request_count,
                jwks_requests_before_unknown_key + 1,
            )

            timeout_identity_settings = replace(
                self.oidc_settings,
                oidc_jwks_url=f"{self.api_url}/identity/jwks.json?timeout-test=1",
            )
            self.app.state.settings = timeout_identity_settings
            self.app.state.token_validator = OIDCTokenValidator(
                timeout_identity_settings
            )
            with patch(
                "licensing_server.app.auth.urlopen",
                side_effect=TimeoutError("simulated identity-provider timeout"),
            ) as timeout_urlopen:
                with self.assertRaises(LicensingApiFailure) as captured:
                    api_client.get(
                        "/v1/me/licenses",
                        access_token=customer_token,
                    )
            self.assertEqual(captured.exception.status_code, 503)
            self.assertEqual(captured.exception.code, "identity_unavailable")
            self.assertTrue(captured.exception.retryable)
            self.assertEqual(
                timeout_urlopen.call_args.kwargs["timeout"],
                OIDC_JWKS_HTTP_TIMEOUT_SECONDS,
            )

            for jwks_path in (
                "/identity/malformed-jwks.json",
                "/identity/non-object-jwks.json",
                "/identity/empty-jwks.json",
                "/identity/duplicate-key-jwks.json",
                "/identity/oversized-jwks.json",
                "/identity/excessive-key-jwks.json",
                "/identity/unavailable-jwks.json",
            ):
                with self.subTest(jwks_path=jwks_path):
                    unavailable_identity_settings = replace(
                        self.oidc_settings,
                        oidc_jwks_url=f"{self.api_url}{jwks_path}",
                    )
                    self.app.state.settings = unavailable_identity_settings
                    self.app.state.token_validator = OIDCTokenValidator(
                        unavailable_identity_settings
                    )
                    with self.assertRaises(LicensingApiFailure) as captured:
                        api_client.get(
                            "/v1/me/licenses",
                            access_token=customer_token,
                        )
                    self.assertEqual(captured.exception.status_code, 503)
                    self.assertEqual(
                        captured.exception.code,
                        "identity_unavailable",
                    )
                    self.assertTrue(captured.exception.retryable)

        with Session(self.engine) as session:
            customer = session.scalar(
                select(User).where(
                    User.external_issuer == TEST_OIDC_ISSUER,
                    User.external_subject == TEST_CUSTOMER_SUBJECT,
                )
            )
            self.assertIsNotNone(customer)
            self.assertEqual(customer.verified_email, TEST_CUSTOMER_EMAIL)
            owned_license = session.scalar(
                select(License).where(License.owner_user_id == customer.id)
            )
            self.assertIsNotNone(owned_license)
            self.assertEqual(owned_license.id, license_id)
            self.assertEqual(owned_license.status, "active")
            activations = tuple(
                session.scalars(
                    select(Activation).where(Activation.license_id == license_id)
                )
            )
            self.assertGreaterEqual(len(activations), 1)
            self.assertTrue(all(value.status == "deactivated" for value in activations))
        serial_bytes = serial.encode("ascii")
        for path in self.root.rglob("*"):
            if path.is_file():
                self.assertNotIn(serial_bytes, path.read_bytes())


if __name__ == "__main__":
    unittest.main()
