"""Backend-only Kit verification; never authenticate from unsigned profile data."""
import hashlib
from contextlib import closing
import json
import secrets
import sqlite3
import time
import urllib.request
from urllib.parse import urlsplit

import jwt
from cryptography.hazmat.primitives.asymmetric import ec


def reject():
    raise ValueError("Invalid Kit launch token")


class NonceStore:
    """Use one shared durable DB outside web root, not one DB per app instance."""
    def __init__(self, filename):
        self.filename = filename
        with closing(sqlite3.connect(filename)) as db, db:
            db.execute("""CREATE TABLE IF NOT EXISTS kit_login_nonce (
                nonce TEXT PRIMARY KEY, pre_session TEXT NOT NULL, expires_ms INTEGER NOT NULL)""")

    def issue(self, pre_session_id, now_ms=None):
        if not isinstance(pre_session_id, str) or len(pre_session_id) < 32:
            reject()
        now_ms = time.time_ns() // 1_000_000 if now_ms is None else now_ms
        nonce = secrets.token_urlsafe(32)
        with closing(sqlite3.connect(self.filename, timeout=5)) as db, db:
            db.execute("DELETE FROM kit_login_nonce WHERE expires_ms <= ?", (now_ms,))
            db.execute("INSERT INTO kit_login_nonce VALUES (?, ?, ?)",
                       (nonce, hashlib.sha256(pre_session_id.encode()).hexdigest(), now_ms + 300_000))
        return nonce

    def consume(self, nonce, pre_session_id, now_ms):
        if not isinstance(pre_session_id, str) or len(pre_session_id) < 32:
            return False
        with closing(sqlite3.connect(self.filename, timeout=5)) as db, db:
            return db.execute("""DELETE FROM kit_login_nonce
                WHERE nonce = ? AND pre_session = ? AND expires_ms > ? RETURNING nonce""",
                (nonce, hashlib.sha256(pre_session_id.encode()).hexdigest(), now_ms)).fetchone() is not None


class NoRedirect(urllib.request.HTTPRedirectHandler):
    def redirect_request(self, req, fp, code, msg, headers, newurl):
        raise ValueError("JWKS redirects are not allowed")


class FixedJWKS(jwt.PyJWKClient):
    # PyJWKClient's per-key LRU has no TTL: keep it disabled. Its whole-set
    # cache expires after 300s and unknown kid causes at most one fresh lookup.
    def __init__(self, url):
        super().__init__(url, cache_keys=False, lifespan=300, timeout=5)

    def fetch_data(self):
        try:
            opener = urllib.request.build_opener(NoRedirect())
            with opener.open(urllib.request.Request(self.uri, headers={"Accept": "application/json"}), timeout=5) as response:
                raw = response.read(1_048_577)
                if len(raw) > 1_048_576:
                    reject()
                data = json.loads(raw)
            self.jwk_set_cache.put(data)
            return data
        except Exception:
            self.jwk_set_cache.put(None)
            raise


class KitVerifier:
    def __init__(self, issuer, kit_id, allowed_origins, nonce_store):
        import re
        if issuer not in ("https://ims.buko.app", "https://ims.koee.app") or not re.fullmatch(r"[a-z0-9.-]{5,100}", kit_id):
            reject()
        if not allowed_origins:
            reject()
        for origin in allowed_origins:
            url = urlsplit(origin)
            if url.scheme != "https" or url.username or url.password or url.path or url.query or url.fragment or not url.hostname:
                reject()
        self.issuer, self.kit_id = issuer, kit_id
        self.origins, self.nonce_store = set(allowed_origins), nonce_store
        self.jwks = FixedJWKS(f"{issuer}/kit-keys/{kit_id}/jwks.json")

    def refresh_keys(self):
        self.jwks.get_jwk_set(refresh=True)

    def verify(self, token, pre_session_id):
        header = checked_header(token)
        key = self.jwks.get_signing_key(header["kid"])
        if (key.key_type != "EC" or key.algorithm_name != "ES256"
                or not isinstance(key.key, ec.EllipticCurvePublicKey)
                or not isinstance(key.key.curve, ec.SECP256R1)
                or key.public_key_use not in (None, "sig")):
            reject()
        return verify_with_key(token, key.key, self.issuer, self.kit_id,
                               self.origins, self.nonce_store, pre_session_id)


def checked_header(token):
    if not isinstance(token, str) or len(token) > 16_384:
        reject()
    header = jwt.get_unverified_header(token)
    if (header.get("alg") != "ES256" or header.get("typ") != "kit-launch+jwt"
            or not isinstance(header.get("kid"), str) or not 1 <= len(header["kid"]) <= 128
            or any(name in header for name in ("jku", "jwk", "x5u", "x5c", "crit", "b64"))):
        reject()
    return header


def verify_with_key(token, key, issuer, kit_id, origins, nonce_store, pre_session_id, now_ms=None):
    """Testable core; application handlers use KitVerifier with configured JWKS."""
    import re
    checked_header(token)
    # Validate signature/issuer/audience with PyJWT. Time checks below deliberately
    # avoid global leeway: only iat may be at most 60 seconds in the future.
    claims = jwt.decode(token, key, algorithms=["ES256"], issuer=issuer, audience=kit_id,
                        options={"verify_exp": False, "verify_iat": False, "verify_nbf": False,
                                 "strict_aud": True,
                                 "require": ["iss", "aud", "sub", "iat", "exp", "jti", "origin", "scope", "nonce"]})
    fixed_time = now_ms is not None
    now_ms = time.time_ns() // 1_000_000 if now_ms is None else now_ms
    iat, exp = claims["iat"], claims["exp"]
    if (type(iat) is not int or type(exp) is not int or iat < 0 or exp <= iat or exp - iat > 300
            or exp > 9_007_199_254_740 or now_ms >= exp * 1000 or iat * 1000 > now_ms + 60_000
            or ("nbf" in claims and (type(claims["nbf"]) is not int or now_ms < claims["nbf"] * 1000))
            or not isinstance(claims["sub"], str) or not re.fullmatch(r"[A-Za-z0-9_-]{43}", claims["sub"])
            or not isinstance(claims["jti"], str) or not 1 <= len(claims["jti"]) <= 128
            or not isinstance(claims["origin"], str) or claims["origin"] not in origins
            or claims["scope"] not in ("profile", "profile handle")
            or not isinstance(claims.get("name"), str) or not isinstance(claims.get("locale"), str)
            or (claims["scope"] == "profile" and "handle" in claims)
            or (claims["scope"] == "profile handle" and "handle" not in claims)
            or (claims.get("handle") is not None and not isinstance(claims["handle"], str))
            or not isinstance(claims["nonce"], str) or not re.fullmatch(r"[A-Za-z0-9_-]{43}", claims["nonce"])):
        reject()
    if not nonce_store.consume(claims["nonce"], pre_session_id, now_ms):
        reject()
    if not fixed_time and time.time_ns() // 1_000_000 >= exp * 1000:
        reject()  # Database lock waits cannot extend token expiry.
    return claims
