"""SQLite の接続・マイグレーション・ログ書き込み。

単一ライターであることを前提とし、WAL で読み取りとの競合を避ける。
接続は 1 本のみ保持し、``asyncio.Lock`` と ``asyncio.to_thread`` の組で
イベント ループを塞がずに直列化する。
"""

from __future__ import annotations

import asyncio
import json
import logging
import os
import sqlite3
from dataclasses import dataclass, field
from datetime import datetime, timedelta
from typing import Any, Callable, Final, Iterable, Sequence

from .config import JST, Settings

logger = logging.getLogger(__name__)

SCHEMA_VERSION = 1

#: ``PRAGMA user_version`` による適用済み管理。番号順に一度だけ適用する。
MIGRATIONS: list[tuple[int, tuple[str, ...]]] = [
    (
        1,
        (
            """
            CREATE TABLE IF NOT EXISTS api_keys (
                id          INTEGER PRIMARY KEY AUTOINCREMENT,
                name        TEXT    NOT NULL,
                key_hash    TEXT    NOT NULL,
                key_prefix  TEXT    NOT NULL,
                enabled     INTEGER NOT NULL DEFAULT 1,
                created_at  TEXT    NOT NULL,
                expires_at  TEXT,
                note        TEXT    NOT NULL DEFAULT ''
            )
            """,
            "CREATE UNIQUE INDEX IF NOT EXISTS ux_api_keys_key_hash ON api_keys(key_hash)",
            """
            CREATE TABLE IF NOT EXISTS api_key_limits (
                key_id      INTEGER PRIMARY KEY REFERENCES api_keys(id) ON DELETE CASCADE,
                rpm         INTEGER,
                rpd         INTEGER,
                concurrency INTEGER,
                tpm         INTEGER
            )
            """,
            """
            CREATE TABLE IF NOT EXISTS api_key_models (
                key_id INTEGER NOT NULL REFERENCES api_keys(id) ON DELETE CASCADE,
                model  TEXT    NOT NULL,
                PRIMARY KEY (key_id, model)
            )
            """,
            """
            CREATE TABLE IF NOT EXISTS request_logs (
                id                INTEGER PRIMARY KEY AUTOINCREMENT,
                key_id            INTEGER,
                ts                TEXT    NOT NULL,
                method            TEXT,
                path              TEXT,
                model             TEXT,
                status            INTEGER,
                duration_ms       INTEGER,
                prompt_tokens     INTEGER,
                completion_tokens INTEGER,
                client_ip         TEXT
            )
            """,
            "CREATE INDEX IF NOT EXISTS ix_request_logs_key_ts ON request_logs(key_id, ts)",
            "CREATE INDEX IF NOT EXISTS ix_request_logs_ts ON request_logs(ts)",
            """
            CREATE TABLE IF NOT EXISTS admin_audit_log (
                id        INTEGER PRIMARY KEY AUTOINCREMENT,
                ts        TEXT NOT NULL,
                actor     TEXT,
                action    TEXT NOT NULL,
                target_id TEXT,
                before    TEXT,
                after     TEXT,
                source_ip TEXT
            )
            """,
            "CREATE INDEX IF NOT EXISTS ix_admin_audit_log_ts ON admin_audit_log(ts)",
            """
            CREATE TABLE IF NOT EXISTS rate_snapshots (
                key_id     INTEGER NOT NULL,
                "window"   TEXT    NOT NULL,
                count      INTEGER NOT NULL,
                updated_at TEXT    NOT NULL,
                PRIMARY KEY (key_id, "window")
            )
            """,
            """
            CREATE TABLE IF NOT EXISTS settings (
                k TEXT PRIMARY KEY,
                v TEXT NOT NULL
            )
            """,
        ),
    ),
]


def now_jst() -> datetime:
    """現在時刻（JST）。"""

    return datetime.now(JST)


def iso_jst(moment: datetime | None = None) -> str:
    """JST の ISO8601 文字列。文字列比較で時系列順になる。"""

    moment = moment or now_jst()
    if moment.tzinfo is None:
        moment = moment.replace(tzinfo=JST)
    return moment.astimezone(JST).isoformat(timespec="seconds")


def parse_iso(value: str | None) -> datetime | None:
    """ISO8601 文字列を JST の datetime へ。解釈できなければ None。"""

    if not value:
        return None
    try:
        parsed = datetime.fromisoformat(value)
    except ValueError:
        return None
    if parsed.tzinfo is None:
        parsed = parsed.replace(tzinfo=JST)
    return parsed.astimezone(JST)


@dataclass(slots=True)
class RequestLogEntry:
    """1 リクエスト分の利用ログ。"""

    key_id: int | None
    ts: str
    method: str
    path: str
    model: str | None
    status: int
    duration_ms: int
    prompt_tokens: int | None
    completion_tokens: int | None
    client_ip: str | None

    def as_row(self) -> tuple[Any, ...]:
        return (
            self.key_id,
            self.ts,
            self.method,
            self.path,
            self.model,
            self.status,
            self.duration_ms,
            self.prompt_tokens,
            self.completion_tokens,
            self.client_ip,
        )


class Database:
    """SQLite への直列化されたアクセス。"""

    def __init__(self, path: str) -> None:
        self.path = path
        self._conn: sqlite3.Connection | None = None
        self._lock = asyncio.Lock()

    # ------------------------------------------------------------------ 接続

    async def connect(self) -> None:
        """接続を張り、PRAGMA とマイグレーションを適用する。"""

        await asyncio.to_thread(self._connect_sync)

    def _connect_sync(self) -> None:
        directory = os.path.dirname(os.path.abspath(self.path))
        if directory:
            os.makedirs(directory, exist_ok=True)
        conn = sqlite3.connect(self.path, check_same_thread=False, timeout=30.0)
        conn.row_factory = sqlite3.Row
        conn.execute("PRAGMA journal_mode=WAL")
        conn.execute("PRAGMA synchronous=NORMAL")
        conn.execute("PRAGMA foreign_keys=ON")
        conn.execute("PRAGMA busy_timeout=30000")
        self._conn = conn
        self._migrate(conn)

    @staticmethod
    def _migrate(conn: sqlite3.Connection) -> None:
        current = int(conn.execute("PRAGMA user_version").fetchone()[0])
        for version, statements in MIGRATIONS:
            if version <= current:
                continue
            logger.info("スキーマ移行を適用します: user_version %s -> %s", current, version)
            for statement in statements:
                conn.execute(statement)
            conn.execute(f"PRAGMA user_version={version}")
            conn.commit()
            current = version

    async def close(self) -> None:
        conn = self._conn
        if conn is None:
            return
        self._conn = None
        await asyncio.to_thread(conn.close)

    # ------------------------------------------------------------ 実行ヘルパ

    async def run(self, fn: Callable[[sqlite3.Connection], Any]) -> Any:
        """接続を引数に取る同期関数を、直列化して別スレッドで実行する。"""

        async with self._lock:
            conn = self._conn
            if conn is None:
                raise RuntimeError("データベースへ接続していません")
            return await asyncio.to_thread(self._run_sync, conn, fn)

    @staticmethod
    def _run_sync(conn: sqlite3.Connection, fn: Callable[[sqlite3.Connection], Any]) -> Any:
        try:
            result = fn(conn)
        except Exception:
            conn.rollback()
            raise
        conn.commit()
        return result

    async def fetchall(self, sql: str, params: Sequence[Any] = ()) -> list[sqlite3.Row]:
        return await self.run(lambda conn: conn.execute(sql, params).fetchall())

    async def fetchone(self, sql: str, params: Sequence[Any] = ()) -> sqlite3.Row | None:
        return await self.run(lambda conn: conn.execute(sql, params).fetchone())

    async def execute(self, sql: str, params: Sequence[Any] = ()) -> int:
        def _do(conn: sqlite3.Connection) -> int:
            cur = conn.execute(sql, params)
            return cur.rowcount

        return await self.run(_do)

    # -------------------------------------------------------------- 監査ログ

    async def write_audit(
        self,
        *,
        action: str,
        actor: str | None,
        target_id: str | None = None,
        before: Any = None,
        after: Any = None,
        source_ip: str | None = None,
    ) -> None:
        """監査ログを 1 件書く。"""

        await self.execute(
            "INSERT INTO admin_audit_log (ts, actor, action, target_id, before, after, source_ip)"
            " VALUES (?, ?, ?, ?, ?, ?, ?)",
            (
                iso_jst(),
                actor,
                action,
                target_id,
                json.dumps(before, ensure_ascii=False) if before is not None else None,
                json.dumps(after, ensure_ascii=False) if after is not None else None,
                source_ip,
            ),
        )

    # ------------------------------------------------------------ 保守タスク

    async def purge_old_logs(self, retention_days: int) -> int:
        """保持日数を超えた利用ログを削除する。"""

        threshold = iso_jst(now_jst() - timedelta(days=retention_days))
        return await self.execute("DELETE FROM request_logs WHERE ts < ?", (threshold,))

    async def checkpoint(self) -> None:
        """WAL のチェックポイントと領域回収。"""

        def _do(conn: sqlite3.Connection) -> None:
            conn.execute("PRAGMA wal_checkpoint(TRUNCATE)")

        await self.run(_do)

    async def vacuum(self) -> None:
        def _do(conn: sqlite3.Connection) -> None:
            conn.isolation_level = None
            try:
                conn.execute("VACUUM")
            finally:
                conn.isolation_level = ""

        await self.run(_do)

    async def get_setting(self, key: str) -> str | None:
        row = await self.fetchone("SELECT v FROM settings WHERE k = ?", (key,))
        return row["v"] if row else None

    async def set_setting(self, key: str, value: str) -> None:
        await self.execute(
            "INSERT INTO settings (k, v) VALUES (?, ?) ON CONFLICT(k) DO UPDATE SET v = excluded.v",
            (key, value),
        )


class LogWriter:
    """利用ログの非同期書き込み。

    ホット パスは ``enqueue`` のみを呼ぶ。実際の書き込みは背後のタスクが
    まとめて行い、待ち行列が溢れた場合はログを捨てて中継を優先する。
    """

    #: 待ち行列を空にしてから終わるための番兵。
    _SENTINEL: Final[object] = object()

    def __init__(self, db: Database, *, maxsize: int = 10000, batch_size: int = 200) -> None:
        self._db = db
        self._queue: asyncio.Queue[Any] = asyncio.Queue(maxsize=maxsize)
        self._batch_size = batch_size
        self._task: asyncio.Task[None] | None = None
        self._dropped = 0

    @property
    def dropped(self) -> int:
        """待ち行列が溢れて捨てた件数。"""

        return self._dropped

    def enqueue(self, entry: RequestLogEntry) -> None:
        """ブロックせずに 1 件積む。"""

        try:
            self._queue.put_nowait(entry)
        except asyncio.QueueFull:
            self._dropped += 1
            if self._dropped % 100 == 1:
                logger.warning("利用ログの待ち行列が溢れています（累計 %s 件を破棄）", self._dropped)

    async def start(self) -> None:
        if self._task is None:
            self._task = asyncio.create_task(self._run(), name="log-writer")

    async def stop(self, timeout: float = 5.0) -> None:
        """待ち行列を空にしてから止める。

        単純に ``cancel`` すると、書き込み中のバッチを丸ごと失う。
        番兵を積んで自然終了させ、待ち切れない場合にのみ打ち切る。
        """

        task = self._task
        if task is None:
            return
        self._task = None
        try:
            self._queue.put_nowait(self._SENTINEL)
        except asyncio.QueueFull:  # pragma: no cover - 溢れている場合は打ち切りに委ねる
            pass
        try:
            await asyncio.wait_for(task, timeout)
        except (asyncio.TimeoutError, TimeoutError):  # pragma: no cover
            logger.warning("利用ログの書き出しが時間内に終わりませんでした")
        except asyncio.CancelledError:  # pragma: no cover
            pass
        await self._flush()

    async def _run(self) -> None:
        while True:
            try:
                item = await self._queue.get()
                batch: list[RequestLogEntry] = []
                stopping = item is self._SENTINEL
                if not stopping:
                    batch.append(item)
                while not stopping and len(batch) < self._batch_size:
                    try:
                        nxt = self._queue.get_nowait()
                    except asyncio.QueueEmpty:
                        break
                    if nxt is self._SENTINEL:
                        stopping = True
                        break
                    batch.append(nxt)
                if batch:
                    await self._write(batch)
                if stopping:
                    return
            except asyncio.CancelledError:
                raise
            except Exception:  # pragma: no cover - 書き込み失敗で中継を止めない
                logger.exception("利用ログの書き込みに失敗しました")
                await asyncio.sleep(1.0)

    async def _flush(self) -> None:
        batch: list[RequestLogEntry] = []
        while True:
            try:
                item = self._queue.get_nowait()
            except asyncio.QueueEmpty:
                break
            if item is not self._SENTINEL:
                batch.append(item)
        if batch:
            try:
                await self._write(batch)
            except Exception:  # pragma: no cover
                logger.exception("終了時の利用ログ書き込みに失敗しました")

    async def _write(self, batch: Iterable[RequestLogEntry]) -> None:
        rows = [entry.as_row() for entry in batch]

        def _do(conn: sqlite3.Connection) -> None:
            conn.executemany(
                "INSERT INTO request_logs"
                " (key_id, ts, method, path, model, status, duration_ms, prompt_tokens, completion_tokens, client_ip)"
                " VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
                rows,
            )

        await self._db.run(_do)


@dataclass(slots=True)
class MaintenanceTask:
    """保持期間の削除と WAL 回収を周期実行する。"""

    db: Database
    settings: Settings
    interval: float = 3600.0
    _task: asyncio.Task[None] | None = field(default=None, init=False, repr=False)

    async def start(self) -> None:
        if self._task is None:
            self._task = asyncio.create_task(self._run(), name="maintenance")

    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 _run(self) -> None:
        while True:
            try:
                await asyncio.sleep(self.interval)
                await self.run_once()
            except asyncio.CancelledError:
                raise
            except Exception:  # pragma: no cover
                logger.exception("保守タスクに失敗しました")

    async def run_once(self) -> None:
        """1 巡分の保守。日次で古いログを削除し、週次で領域を回収する。"""

        last_purge = await self.db.get_setting("last_purge_at")
        now = now_jst()
        last = parse_iso(last_purge)
        if last is None or (now - last) >= timedelta(days=1):
            removed = await self.db.purge_old_logs(self.settings.log_retention_days)
            await self.db.set_setting("last_purge_at", iso_jst(now))
            if removed:
                logger.info("保持期間を超えた利用ログを %s 件削除しました", removed)

        last_vacuum = parse_iso(await self.db.get_setting("last_vacuum_at"))
        if last_vacuum is None or (now - last_vacuum) >= timedelta(days=7):
            await self.db.checkpoint()
            await self.db.vacuum()
            await self.db.set_setting("last_vacuum_at", iso_jst(now))
            logger.info("WAL のチェックポイントと領域回収を実施しました")
