"""発行キーの生成・照合・インメモリ キャッシュ。

平文は保存しない。保管するのは SHA-256（``KEY_PEPPER`` 指定時は HMAC-SHA-256）のみ。
ホット パスで SQLite を引かないよう、全キーをメモリへ載せて周期的に読み直す。
"""

from __future__ import annotations

import asyncio
import hashlib
import hmac
import logging
import secrets
import sqlite3
import string
from dataclasses import dataclass, field
from datetime import datetime

from .db import Database, iso_jst, now_jst, parse_iso

logger = logging.getLogger(__name__)

SECRET_PREFIX = "sk-aig-"
SECRET_BODY_LENGTH = 32
PREFIX_LENGTH = 12
_ALPHABET = string.ascii_letters + string.digits


def generate_secret() -> str:
    """発行キーを 1 本生成する。``sk-aig-`` + 英数 32 文字。"""

    body = "".join(secrets.choice(_ALPHABET) for _ in range(SECRET_BODY_LENGTH))
    return SECRET_PREFIX + body


def hash_secret(secret: str, pepper: str = "") -> str:
    """発行キーのハッシュ。高エントロピーのためソルトとストレッチは行わない。"""

    raw = secret.encode("utf-8")
    if pepper:
        return hmac.new(pepper.encode("utf-8"), raw, hashlib.sha256).hexdigest()
    return hashlib.sha256(raw).hexdigest()


def key_prefix(secret: str) -> str:
    """画面表示用の先頭 12 文字。"""

    return secret[:PREFIX_LENGTH]


def looks_like_secret(candidate: str) -> bool:
    """発行キーの書式に一致するか。照合前の足切りに使う。"""

    if not candidate.startswith(SECRET_PREFIX):
        return False
    body = candidate[len(SECRET_PREFIX) :]
    return len(body) == SECRET_BODY_LENGTH and all(c in _ALPHABET for c in body)


@dataclass(frozen=True, slots=True)
class KeyLimits:
    """キー単位の上限。``None`` は当該軸の制限なしを意味する。"""

    rpm: int | None = None
    rpd: int | None = None
    concurrency: int | None = None
    tpm: int | None = None

    def as_dict(self) -> dict[str, int | None]:
        return {"rpm": self.rpm, "rpd": self.rpd, "concurrency": self.concurrency, "tpm": self.tpm}


@dataclass(frozen=True, slots=True)
class KeyRecord:
    """照合済みの発行キー。"""

    id: int
    name: str
    key_hash: str
    key_prefix: str
    enabled: bool
    created_at: str
    expires_at: str | None
    note: str
    limits: KeyLimits
    models: frozenset[str]

    @property
    def unrestricted_models(self) -> bool:
        """許可リストが空なら全モデル許可。"""

        return not self.models

    def allows_model(self, model: str) -> bool:
        return self.unrestricted_models or model in self.models

    def is_active(self, moment: datetime | None = None) -> bool:
        """無効化されておらず、有効期限内か。"""

        if not self.enabled:
            return False
        expires = parse_iso(self.expires_at)
        if expires is None:
            return True
        return (moment or now_jst()) < expires


@dataclass(slots=True)
class KeyCache:
    """全キーのインメモリ キャッシュ。"""

    db: Database
    pepper: str = ""
    ttl: float = 60.0
    _by_hash: dict[str, KeyRecord] = field(default_factory=dict, init=False, repr=False)
    _by_id: dict[int, KeyRecord] = field(default_factory=dict, init=False, repr=False)
    _task: asyncio.Task[None] | None = field(default=None, init=False, repr=False)
    _loaded: bool = field(default=False, init=False, repr=False)

    # ------------------------------------------------------------------ 参照

    def get_by_hash(self, digest: str) -> KeyRecord | None:
        return self._by_hash.get(digest)

    def get_by_id(self, key_id: int) -> KeyRecord | None:
        return self._by_id.get(key_id)

    def resolve(self, secret: str) -> KeyRecord | None:
        """平文の発行キーからレコードを引く。"""

        if not looks_like_secret(secret):
            return None
        return self._by_hash.get(hash_secret(secret, self.pepper))

    def all_records(self) -> tuple[KeyRecord, ...]:
        return tuple(self._by_id.values())

    @property
    def loaded(self) -> bool:
        return self._loaded

    # ------------------------------------------------------------------ 更新

    async def reload(self) -> None:
        """SQLite から全キーを読み直す。"""

        rows = await self.db.run(_load_all)
        by_hash: dict[str, KeyRecord] = {}
        by_id: dict[int, KeyRecord] = {}
        for record in rows:
            by_hash[record.key_hash] = record
            by_id[record.id] = record
        self._by_hash = by_hash
        self._by_id = by_id
        self._loaded = True

    async def invalidate(self) -> None:
        """control API の更新直後に呼ぶ。"""

        await self.reload()

    async def start(self) -> None:
        await self.reload()
        if self._task is None:
            self._task = asyncio.create_task(self._refresh_loop(), name="key-cache-refresh")

    async def stop(self) -> None:
        task = self._task
        if task is None:
            return
        self._task = None
        task.cancel()
        try:
            await task
        except asyncio.CancelledError:
            pass

    async def _refresh_loop(self) -> None:
        while True:
            try:
                await asyncio.sleep(self.ttl)
                await self.reload()
            except asyncio.CancelledError:
                raise
            except Exception:  # pragma: no cover - 読み直し失敗で中継を止めない
                logger.exception("キャッシュの読み直しに失敗しました")


def _load_all(conn: sqlite3.Connection) -> list[KeyRecord]:
    """全キーを 3 クエリで読み出す（キーごとの N+1 を避ける）。"""

    limits: dict[int, KeyLimits] = {}
    for row in conn.execute("SELECT key_id, rpm, rpd, concurrency, tpm FROM api_key_limits"):
        limits[row["key_id"]] = KeyLimits(
            rpm=row["rpm"], rpd=row["rpd"], concurrency=row["concurrency"], tpm=row["tpm"]
        )

    models: dict[int, set[str]] = {}
    for row in conn.execute("SELECT key_id, model FROM api_key_models"):
        models.setdefault(row["key_id"], set()).add(row["model"])

    records: list[KeyRecord] = []
    for row in conn.execute(
        "SELECT id, name, key_hash, key_prefix, enabled, created_at, expires_at, note FROM api_keys"
    ):
        records.append(
            KeyRecord(
                id=row["id"],
                name=row["name"],
                key_hash=row["key_hash"],
                key_prefix=row["key_prefix"],
                enabled=bool(row["enabled"]),
                created_at=row["created_at"],
                expires_at=row["expires_at"],
                note=row["note"] or "",
                limits=limits.get(row["id"], KeyLimits()),
                models=frozenset(models.get(row["id"], ())),
            )
        )
    return records


# --------------------------------------------------------------------- 書き込み


def insert_key(
    conn: sqlite3.Connection,
    *,
    name: str,
    key_hash: str,
    prefix: str,
    note: str,
    expires_at: str | None,
    limits: KeyLimits,
    models: list[str],
) -> int:
    """キーを 1 本追加し、その id を返す。"""

    cur = conn.execute(
        "INSERT INTO api_keys (name, key_hash, key_prefix, enabled, created_at, expires_at, note)"
        " VALUES (?, ?, ?, 1, ?, ?, ?)",
        (name, key_hash, prefix, iso_jst(), expires_at, note),
    )
    key_id = int(cur.lastrowid)
    conn.execute(
        "INSERT INTO api_key_limits (key_id, rpm, rpd, concurrency, tpm) VALUES (?, ?, ?, ?, ?)",
        (key_id, limits.rpm, limits.rpd, limits.concurrency, limits.tpm),
    )
    replace_models(conn, key_id, models)
    return key_id


def replace_models(conn: sqlite3.Connection, key_id: int, models: list[str]) -> None:
    """許可モデルを全置換する。"""

    conn.execute("DELETE FROM api_key_models WHERE key_id = ?", (key_id,))
    unique = sorted({model.strip() for model in models if model and model.strip()})
    if unique:
        conn.executemany(
            "INSERT INTO api_key_models (key_id, model) VALUES (?, ?)",
            [(key_id, model) for model in unique],
        )
