"""Control-plane registry for sites, signing keys, and JWT replay records."""

from __future__ import annotations

from datetime import datetime, timezone
from typing import Any
from urllib.parse import urlparse

from cryptography.hazmat.primitives.serialization import load_pem_public_key
from sqlalchemy import delete, select
from sqlalchemy.exc import IntegrityError

from config.control_plane import get_control_plane_session
from bridge_platform.tenants.models import JWTReplayRecord, PlatformSite, SiteSigningKey, Tenant


def get_signing_key_record(key_id: str) -> dict[str, Any] | None:
    with get_control_plane_session() as session:
        stmt = (
            select(SiteSigningKey, PlatformSite)
            .join(PlatformSite, PlatformSite.site_id == SiteSigningKey.site_id)
            .where(SiteSigningKey.key_id == key_id)
            .limit(1)
        )
        row = session.execute(stmt).first()
        if row is None:
            return None
        key, site = row
        return {
            "key_id": key.key_id,
            "site_id": site.site_id,
            "tenant_id": site.tenant_id,
            "issuer": site.issuer.rstrip("/"),
            "audience": site.audience,
            "public_key_pem": key.public_key_pem,
            "key_status": key.status,
            "site_status": site.status,
        }


def get_active_signing_key(key_id: str) -> dict[str, Any] | None:
    record = get_signing_key_record(key_id)
    if record is None:
        return None
    if record["key_status"] != "active" or record["site_status"] != "active":
        return None
    return record


def register_bootstrap_site(
    *,
    site_id: str,
    tenant_id: str,
    issuer: str,
    audience: str,
    key_id: str,
    public_key_pem: str,
) -> None:
    _validate_public_key(public_key_pem)
    now = datetime.utcnow()
    with get_control_plane_session() as session:
        tenant = session.get(Tenant, tenant_id)
        if tenant is None or not tenant.is_active or tenant.status != "active":
            raise PermissionError(f"Tenant is not active: {tenant_id}")

        site = session.get(PlatformSite, site_id)
        if site is None:
            site = PlatformSite(
                site_id=site_id,
                tenant_id=tenant_id,
                issuer=issuer.rstrip("/"),
                audience=audience,
                domain=urlparse(issuer).hostname,
                status="active",
            )
            session.add(site)
        else:
            _assert_site_binding(site, tenant_id=tenant_id, issuer=issuer, audience=audience)
            site.status = "active"
        site.last_seen_at = now

        key = session.get(SiteSigningKey, key_id)
        if key is None:
            session.add(
                SiteSigningKey(
                    key_id=key_id,
                    site_id=site_id,
                    algorithm="RS256",
                    public_key_pem=public_key_pem,
                    status="active",
                    activated_at=now,
                )
            )
        elif key.site_id != site_id or key.public_key_pem.strip() != public_key_pem.strip():
            raise PermissionError("Signing key identifier is already bound differently")
        session.commit()


def rotate_signing_key(
    *,
    site_id: str,
    current_key_id: str,
    new_key_id: str,
    public_key_pem: str,
) -> None:
    _validate_public_key(public_key_pem)
    now = datetime.utcnow()
    with get_control_plane_session() as session:
        current = session.get(SiteSigningKey, current_key_id)
        site = session.get(PlatformSite, site_id)
        if site is None or site.status != "active":
            raise LookupError("Registered site not found")
        if current is None or current.site_id != site_id or current.status != "active":
            raise PermissionError("Current signing key is not active")
        if session.get(SiteSigningKey, new_key_id) is not None:
            raise ValueError("New key identifier already exists")

        session.add(
            SiteSigningKey(
                key_id=new_key_id,
                site_id=site_id,
                algorithm="RS256",
                public_key_pem=public_key_pem,
                status="active",
                activated_at=now,
            )
        )
        site.last_seen_at = now
        session.commit()


def revoke_signing_key(*, site_id: str, key_id: str, current_key_id: str) -> None:
    if key_id == current_key_id:
        raise ValueError("Cannot revoke the key authenticating this request")
    with get_control_plane_session() as session:
        key = session.get(SiteSigningKey, key_id)
        if key is None or key.site_id != site_id:
            raise LookupError("Signing key not found")
        key.status = "revoked"
        key.revoked_at = datetime.utcnow()
        session.commit()


def consume_jti(*, site_id: str, jti: str, expires_at: datetime) -> None:
    """Persist one token identifier; duplicates are rejected."""
    now = datetime.utcnow()
    naive_expiry = expires_at.astimezone(timezone.utc).replace(tzinfo=None)
    with get_control_plane_session() as session:
        session.execute(delete(JWTReplayRecord).where(JWTReplayRecord.expires_at < now))
        session.add(JWTReplayRecord(jti=jti, site_id=site_id, expires_at=naive_expiry))
        try:
            session.commit()
        except IntegrityError as exc:
            session.rollback()
            raise PermissionError("JWT replay detected") from exc


def describe_site(site_id: str) -> dict[str, Any] | None:
    with get_control_plane_session() as session:
        site = session.get(PlatformSite, site_id)
        if site is None:
            return None
        keys = session.execute(
            select(SiteSigningKey)
            .where(SiteSigningKey.site_id == site_id)
            .order_by(SiteSigningKey.created_at.desc())
        ).scalars().all()
        return {
            "site_id": site.site_id,
            "tenant_id": site.tenant_id,
            "issuer": site.issuer,
            "audience": site.audience,
            "status": site.status,
            "keys": [
                {
                    "kid": key.key_id,
                    "status": key.status,
                    "algorithm": key.algorithm,
                    "activated_at": key.activated_at.isoformat(),
                    "revoked_at": key.revoked_at.isoformat() if key.revoked_at else None,
                }
                for key in keys
            ],
        }


def _assert_site_binding(site: PlatformSite, *, tenant_id: str, issuer: str, audience: str) -> None:
    if (
        site.tenant_id != tenant_id
        or site.issuer.rstrip("/") != issuer.rstrip("/")
        or site.audience != audience
    ):
        raise PermissionError("Site identity does not match its registered binding")


def _validate_public_key(public_key_pem: str) -> None:
    try:
        load_pem_public_key(public_key_pem.encode("utf-8"))
    except Exception as exc:
        raise ValueError("Invalid PEM public key") from exc
