"""Opt-in PostgreSQL transaction qualification for the licensing authority.

This module is discovered safely by the portable lane, but it mutates a
database only when the disposable qualification runner supplies both explicit
guards.  It must never be pointed at a production database.
"""

from __future__ import annotations

from concurrent.futures import ThreadPoolExecutor
import os
from threading import Barrier
import unittest
import uuid

from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import ec
from sqlalchemy import create_engine, func, select
from sqlalchemy.engine import make_url
from sqlalchemy.orm import Session
from sqlalchemy.orm import sessionmaker

from licensing_shared.catalog import load_builtin_catalog
from licensing_shared.offline import device_public_key_thumbprint
from licensing_server.app.config import ServerSettings
from licensing_server.app.errors import ServerErrorCode, ServerLicensingError
from licensing_server.app.models import (
    Activation,
    Grant,
    License,
    RateLimitBucket,
    Serial,
    User,
)
from licensing_server.app.rate_limits import (
    DatabaseRateLimiter,
    RateLimitQuota,
    RequestRateLimitContext,
)
from licensing_server.app.security import Ed25519SnapshotSigner
from licensing_server.app.services import DeviceEnrollment, LicensingService
from licensing_server.qualification_contract import (
    QUALIFICATION_CONCURRENT_REDEEMERS as CONCURRENT_REDEEMERS,
    QUALIFICATION_DATABASE_PREFIX,
    QUALIFICATION_PROJECT_PATTERN,
    QUALIFICATION_RATE_LIMIT_ATTEMPTS as RATE_LIMIT_QUALIFICATION_ATTEMPTS,
    QUALIFICATION_RATE_LIMIT_MAXIMUM_ACCEPTED as RATE_LIMIT_QUALIFICATION_MAXIMUM_ACCEPTED,
    QUALIFICATION_RATE_LIMIT_REPLICAS as RATE_LIMIT_QUALIFICATION_REPLICAS,
)


QUALIFICATION_ENABLED_VARIABLE = "LICENSING_LIVE_POSTGRESQL_QUALIFICATION"
QUALIFICATION_PROJECT_VARIABLE = "LICENSING_QUALIFICATION_PROJECT"
def _qualification_suffix(settings: ServerSettings) -> str:
    if not isinstance(settings, ServerSettings):
        raise TypeError("settings must be ServerSettings")
    project_name = os.environ.get(QUALIFICATION_PROJECT_VARIABLE, "")
    project_match = QUALIFICATION_PROJECT_PATTERN.fullmatch(project_name)
    if project_match is None:
        raise RuntimeError("the disposable qualification project guard is invalid")
    database_name = make_url(settings.database_url).database or ""
    expected_database = QUALIFICATION_DATABASE_PREFIX + project_match.group(1)
    if database_name != expected_database:
        raise RuntimeError(
            "the live PostgreSQL test refuses a database outside its exact "
            "disposable qualification project"
        )
    if make_url(settings.database_url).get_backend_name() != "postgresql":
        raise RuntimeError("the live transaction qualification requires PostgreSQL")
    return project_match.group(1)


@unittest.skipUnless(
    os.environ.get(QUALIFICATION_ENABLED_VARIABLE) == "1",
    "explicit disposable PostgreSQL qualification was not requested",
)
class LicensingPostgreSQLLiveTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls) -> None:
        cls.settings = ServerSettings.from_environment()
        cls.qualification_suffix = _qualification_suffix(cls.settings)
        cls.engine = create_engine(
            cls.settings.database_url,
            pool_pre_ping=True,
            pool_size=CONCURRENT_REDEEMERS + 1,
            max_overflow=CONCURRENT_REDEEMERS,
        )
        cls.signer = Ed25519SnapshotSigner.from_private_key_file(
            cls.settings.signing_key_id,
            cls.settings.signing_private_key_path,
            require_owner_only=cls.settings.allow_production_file_signer,
        )
        cls.catalog = load_builtin_catalog()

    @classmethod
    def tearDownClass(cls) -> None:
        cls.engine.dispose()

    def _service(self, session: Session) -> LicensingService:
        if not isinstance(session, Session):
            raise TypeError("session must be a Session")
        return LicensingService(
            session,
            self.catalog,
            self.signer,
            self.settings.serial_pepper,
            self.settings.fingerprint_pepper,
        )

    @staticmethod
    def _enrollment(suffix: str) -> DeviceEnrollment:
        if not isinstance(suffix, str) or not suffix:
            raise ValueError("suffix must be a non-empty string")
        private_key = ec.generate_private_key(ec.SECP256R1())
        public_key_der = private_key.public_key().public_bytes(
            encoding=serialization.Encoding.DER,
            format=serialization.PublicFormat.SubjectPublicKeyInfo,
        )
        return DeviceEnrollment(
            installation_id=f"installation.qualification.{suffix}",
            public_key_der=public_key_der,
            key_thumbprint=device_public_key_thumbprint(public_key_der),
            key_provider="software.qualification",
            friendly_name=f"Qualification device {suffix}",
            evidence={
                "system_uuid": f"qualification-{suffix}",
                "baseboard_serial": f"qualification-board-{suffix}",
            },
        )

    def test_one_serial_can_be_redeemed_only_once_under_a_real_row_lock(self) -> None:
        run_suffix = uuid.uuid4().hex
        admin_user_id = f"user.qualification.{run_suffix}"
        with Session(self.engine) as session, session.begin():
            session.add(
                User(
                    id=admin_user_id,
                    external_issuer="https://identity.qualification.invalid",
                    external_subject=f"admin-{run_suffix}",
                    verified_email=f"admin-{run_suffix}@qualification.invalid",
                    status="active",
                    is_server_admin=True,
                )
            )
        with Session(self.engine) as session:
            serial_value = self._service(session).generate_serial_batch(
                "edition.professional",
                1,
                admin_user_id,
                "Disposable PostgreSQL transaction qualification",
                f"correlation.qualification.issue.{run_suffix}",
            ).serials[0]

        barrier = Barrier(CONCURRENT_REDEEMERS + 1)
        enrollments = tuple(
            self._enrollment(f"{run_suffix}.{index}")
            for index in range(CONCURRENT_REDEEMERS)
        )

        def redeem(index: int) -> str:
            if isinstance(index, bool) or not isinstance(index, int):
                raise TypeError("index must be an integer")
            barrier.wait(timeout=20)
            try:
                with Session(self.engine) as session:
                    self._service(session).redeem_serial(
                        serial_value,
                        enrollments[index],
                        1,
                        f"idempotency.qualification.redeem.{run_suffix}.{index}",
                        f"correlation.qualification.redeem.{run_suffix}.{index}",
                    )
            except ServerLicensingError as exc:
                return exc.code.value
            return "redeemed"

        with ThreadPoolExecutor(max_workers=CONCURRENT_REDEEMERS) as executor:
            futures = tuple(
                executor.submit(redeem, index)
                for index in range(CONCURRENT_REDEEMERS)
            )
            barrier.wait(timeout=20)
            outcomes = sorted(future.result(timeout=30) for future in futures)

        self.assertEqual(
            outcomes,
            sorted(("redeemed", ServerErrorCode.SERIAL_ALREADY_REDEEMED.value)),
        )
        with Session(self.engine) as session:
            serial_row = session.scalar(
                select(Serial).where(Serial.batch.has(created_by_user_id=admin_user_id))
            )
            self.assertIsNotNone(serial_row)
            self.assertEqual(serial_row.status, "redeemed")
            self.assertEqual(
                session.scalar(
                    select(func.count()).select_from(License).where(
                        License.id == serial_row.redeemed_license_id
                    )
                ),
                1,
            )
            self.assertEqual(
                session.scalar(
                    select(func.count()).select_from(Grant).where(
                        Grant.license_id == serial_row.redeemed_license_id
                    )
                ),
                1,
            )
            self.assertEqual(
                session.scalar(
                    select(func.count()).select_from(Activation).where(
                        Activation.license_id == serial_row.redeemed_license_id
                    )
                ),
                1,
            )

    def test_rate_limit_burst_is_atomic_across_shared_postgresql_replicas(self) -> None:
        run_suffix = uuid.uuid4().hex
        quota = RateLimitQuota(
            f"qualification.rate.{run_suffix}",
            "source",
            RATE_LIMIT_QUALIFICATION_MAXIMUM_ACCEPTED,
            60,
        )
        factory = sessionmaker(
            bind=self.engine,
            expire_on_commit=False,
            autoflush=False,
        )
        limiters = tuple(
            DatabaseRateLimiter(
                factory,
                self.settings.fingerprint_pepper,
                request_quotas={"api.other": (quota,)},
                principal_quotas={},
            )
            for _index in range(RATE_LIMIT_QUALIFICATION_REPLICAS)
        )
        source_address = "2001:db8:" + ":".join(
            run_suffix[index : index + 4] for index in range(0, 24, 4)
        )
        context = RequestRateLimitContext("api.other", source_address)
        barrier = Barrier(RATE_LIMIT_QUALIFICATION_ATTEMPTS + 1)

        def attempt(index: int) -> str:
            if isinstance(index, bool) or not isinstance(index, int):
                raise TypeError("index must be an integer")
            barrier.wait(timeout=20)
            try:
                limiters[index % len(limiters)].enforce_request(context)
            except ServerLicensingError as exc:
                return exc.code.value
            return "accepted"

        with ThreadPoolExecutor(
            max_workers=RATE_LIMIT_QUALIFICATION_ATTEMPTS
        ) as executor:
            futures = tuple(
                executor.submit(attempt, index)
                for index in range(RATE_LIMIT_QUALIFICATION_ATTEMPTS)
            )
            barrier.wait(timeout=20)
            outcomes = tuple(future.result(timeout=30) for future in futures)

        self.assertEqual(
            outcomes.count("accepted"),
            RATE_LIMIT_QUALIFICATION_MAXIMUM_ACCEPTED,
        )
        self.assertEqual(
            outcomes.count(ServerErrorCode.RATE_LIMITED.value),
            RATE_LIMIT_QUALIFICATION_ATTEMPTS
            - RATE_LIMIT_QUALIFICATION_MAXIMUM_ACCEPTED,
        )
        with Session(self.engine) as session:
            bucket = session.scalar(
                select(RateLimitBucket).where(RateLimitBucket.policy == quota.name)
            )
            self.assertIsNotNone(bucket)
            assert bucket is not None
            self.assertEqual(
                bucket.request_count,
                RATE_LIMIT_QUALIFICATION_ATTEMPTS,
            )
            self.assertNotIn(source_address, bucket.id)


if __name__ == "__main__":
    unittest.main()
