"""control API。admin から呼ばれる内部 API。

``control-api-contract.md`` が唯一の結合点であり、本モジュールはその契約に一致する。
Compose ネットワーク内のみで到達可能とし、ホストへは公開しない。
共有シークレットは論理的な境界であり、実質的な境界はネットワークの分離である。
"""

from __future__ import annotations

import hmac
import logging
import sqlite3
from datetime import datetime, timedelta
from typing import Any, Callable, Iterable

from fastapi import APIRouter, Depends, FastAPI, Query, Request, status
from fastapi.exceptions import RequestValidationError
from pydantic import BaseModel, ConfigDict, Field
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.responses import JSONResponse, Response

from . import errors
from . import keys as keys_module
from .config import JST, Settings
from .db import Database, iso_jst, now_jst
from .keys import KeyCache, KeyLimits, KeyRecord
from .limits import RateLimiter
from .logging_utils import register_secret
from .upstream import UPSTREAM_OK, UpstreamClient

logger = logging.getLogger(__name__)

ACTION_CREATE = "key.create"
ACTION_UPDATE = "key.update"
ACTION_DELETE = "key.delete"
ACTION_AUTH_FAIL = "auth.fail"


# ------------------------------------------------------------------ 入力モデル


class LimitsIn(BaseModel):
    """キー単位の上限。``null`` は当該軸の制限なし。"""

    model_config = ConfigDict(extra="forbid")

    rpm: int | None = Field(default=None, ge=0)
    rpd: int | None = Field(default=None, ge=0)
    concurrency: int | None = Field(default=None, ge=0)
    tpm: int | None = Field(default=None, ge=0)


class KeyCreateIn(BaseModel):
    model_config = ConfigDict(extra="forbid")

    name: str = Field(min_length=1, max_length=200)
    note: str = ""
    expires_at: datetime | None = None
    limits: LimitsIn = Field(default_factory=LimitsIn)
    models: list[str] = Field(default_factory=list)


class KeyPatchIn(BaseModel):
    model_config = ConfigDict(extra="forbid")

    name: str | None = Field(default=None, min_length=1, max_length=200)
    note: str | None = None
    enabled: bool | None = None
    expires_at: datetime | None = None
    limits: LimitsIn | None = None
    models: list[str] | None = None


# ------------------------------------------------------------------ 補助関数


def _store_datetime(value: datetime | None) -> str | None:
    """タイムゾーンなしは JST とみなして ISO8601 で保存する。"""

    if value is None:
        return None
    if value.tzinfo is None:
        value = value.replace(tzinfo=JST)
    return iso_jst(value)


def _key_dict(record: KeyRecord, usage: dict[int, dict[str, int]]) -> dict[str, Any]:
    """契約どおりのキー表現を組み立てる。"""

    stats = usage.get(record.id, {"requests": 0, "prompt_tokens": 0, "completion_tokens": 0})
    return {
        "id": record.id,
        "name": record.name,
        "key_prefix": record.key_prefix,
        "enabled": record.enabled,
        "created_at": record.created_at,
        "expires_at": record.expires_at,
        "note": record.note,
        "limits": record.limits.as_dict(),
        "models": sorted(record.models),
        "usage_24h": {
            "requests": stats.get("requests", 0),
            "prompt_tokens": stats.get("prompt_tokens", 0),
            "completion_tokens": stats.get("completion_tokens", 0),
        },
    }


async def _usage_24h(db: Database) -> dict[int, dict[str, int]]:
    since = iso_jst(now_jst() - timedelta(hours=24))
    rows = await db.fetchall(
        "SELECT key_id,"
        " COUNT(*) AS requests,"
        " COALESCE(SUM(prompt_tokens), 0) AS prompt_tokens,"
        " COALESCE(SUM(completion_tokens), 0) AS completion_tokens"
        " FROM request_logs WHERE ts >= ? AND key_id IS NOT NULL GROUP BY key_id",
        (since,),
    )
    return {
        int(row["key_id"]): {
            "requests": int(row["requests"]),
            "prompt_tokens": int(row["prompt_tokens"]),
            "completion_tokens": int(row["completion_tokens"]),
        }
        for row in rows
    }


def _actor(request: Request) -> tuple[str | None, str | None]:
    """監査用のヘッダ。任意項目のため欠けていてもよい。"""

    return request.headers.get("x-actor"), request.headers.get("x-actor-ip")


def _limits_from(payload: LimitsIn) -> KeyLimits:
    return KeyLimits(rpm=payload.rpm, rpd=payload.rpd, concurrency=payload.concurrency, tpm=payload.tpm)


def _audit_snapshot(record: KeyRecord) -> dict[str, Any]:
    """監査ログへ残す表現。平文キーもハッシュも含めない。"""

    return {
        "id": record.id,
        "name": record.name,
        "key_prefix": record.key_prefix,
        "enabled": record.enabled,
        "created_at": record.created_at,
        "expires_at": record.expires_at,
        "note": record.note,
        "limits": record.limits.as_dict(),
        "models": sorted(record.models),
    }


# ------------------------------------------------------------------ アプリ本体


def create_control_app(
    *,
    settings: Settings,
    db: Database,
    keys: KeyCache,
    upstream: UpstreamClient,
    limiter: RateLimiter | None = None,
    listeners_alive: Callable[[], bool] | None = None,
) -> FastAPI:
    """control ポート用の FastAPI アプリを組み立てる。

    ``listeners_alive`` は公開リスナーの生死を返す。``None`` の場合は生きているとみなす。
    """

    register_secret(settings.upstream_api_key, settings.control_token)

    app = FastAPI(
        title="LLM gateway control API",
        version=settings.version,
        docs_url=None,
        redoc_url=None,
        openapi_url=None,
    )

    async def require_token(request: Request) -> None:
        """共有シークレットを照合する。失敗は監査ログへ残す。"""

        provided = request.headers.get("x-control-token", "")
        if not settings.control_token or not hmac.compare_digest(provided, settings.control_token):
            actor, actor_ip = _actor(request)
            source_ip = actor_ip or (request.client.host if request.client else None)
            try:
                await db.write_audit(
                    action=ACTION_AUTH_FAIL,
                    actor=actor,
                    target_id=request.url.path,
                    after={"method": request.method, "path": request.url.path},
                    source_ip=source_ip,
                )
            except Exception:  # pragma: no cover - 監査書き込み失敗で応答は変えない
                logger.exception("認証失敗の監査ログ書き込みに失敗しました")
            raise errors.control_unauthorized()

    router = APIRouter(dependencies=[Depends(require_token)])

    # ------------------------------------------------------------ ヘルスチェック

    @app.get("/healthz")
    async def healthz(strict: bool = Query(default=False)) -> JSONResponse:
        """ゲートウェイ自身の生死を返す。

        上流の状態は生死判定に含めない。含めると、上流が一時的に落ちただけで
        コンテナが unhealthy になり、``depends_on`` で待つ admin まで起動できなくなる。
        管理画面は上流障害時こそ必要であり、この依存は逆効果である。

        ``strict=1`` を付けた場合のみ、上流が ``ok`` でなければ 503 を返す。外形監視用。
        """

        database = await _probe_database(db)
        listeners = "ok" if (listeners_alive is None or listeners_alive()) else "down"
        self_ok = database == "ok" and listeners == "ok"

        if strict:
            # 能動的に上流を叩く。応答が遅い上流で Docker の healthcheck を
            # 巻き添えにしないため、これは strict のときだけ行う。
            upstream_state = await upstream.probe()
        else:
            upstream_state = upstream.upstream_status

        healthy = self_ok and (not strict or upstream_state == UPSTREAM_OK)
        payload = {
            "status": "ok" if self_ok else "degraded",
            "upstream": upstream_state,
            "version": settings.version,
            "checks": {"database": database, "listeners": listeners},
        }
        code = status.HTTP_200_OK if healthy else status.HTTP_503_SERVICE_UNAVAILABLE
        return JSONResponse(payload, status_code=code)

    # -------------------------------------------------------------------- キー

    @router.get("/keys")
    async def list_keys(include_disabled: bool = Query(default=True)) -> JSONResponse:
        await keys.reload()
        usage = await _usage_24h(db)
        records = sorted(keys.all_records(), key=lambda record: record.id)
        payload = [
            _key_dict(record, usage)
            for record in records
            if include_disabled or record.enabled
        ]
        return JSONResponse({"keys": payload})

    @router.post("/keys", status_code=status.HTTP_201_CREATED)
    async def create_key(body: KeyCreateIn, request: Request) -> JSONResponse:
        secret = keys_module.generate_secret()
        digest = keys_module.hash_secret(secret, settings.key_pepper)
        prefix = keys_module.key_prefix(secret)
        limits = _limits_from(body.limits)
        expires = _store_datetime(body.expires_at)

        def _insert(conn: sqlite3.Connection) -> int:
            return keys_module.insert_key(
                conn,
                name=body.name,
                key_hash=digest,
                prefix=prefix,
                note=body.note,
                expires_at=expires,
                limits=limits,
                models=body.models,
            )

        key_id = await db.run(_insert)
        await keys.invalidate()
        record = keys.get_by_id(key_id)
        if record is None:  # pragma: no cover - 直後に読み直せない場合のみ
            raise errors.GatewayError(500, "Failed to load the created key.", code="internal_error")

        actor, actor_ip = _actor(request)
        await db.write_audit(
            action=ACTION_CREATE,
            actor=actor,
            target_id=str(key_id),
            after=_audit_snapshot(record),
            source_ip=actor_ip or (request.client.host if request.client else None),
        )
        usage = await _usage_24h(db)
        return JSONResponse(
            {"key": _key_dict(record, usage), "secret": secret},
            status_code=status.HTTP_201_CREATED,
        )

    @router.patch("/keys/{key_id}")
    async def update_key(key_id: int, body: KeyPatchIn, request: Request) -> JSONResponse:
        before = keys.get_by_id(key_id)
        if before is None:
            await keys.reload()
            before = keys.get_by_id(key_id)
        if before is None:
            raise errors.not_found("The requested API key does not exist.")

        provided = body.model_fields_set
        column_updates: list[tuple[str, Any]] = []
        if "name" in provided and body.name is not None:
            column_updates.append(("name", body.name))
        if "note" in provided and body.note is not None:
            column_updates.append(("note", body.note))
        if "enabled" in provided and body.enabled is not None:
            column_updates.append(("enabled", 1 if body.enabled else 0))
        if "expires_at" in provided:
            column_updates.append(("expires_at", _store_datetime(body.expires_at)))

        limits_updates: list[tuple[str, Any]] = []
        if "limits" in provided and body.limits is not None:
            for field in body.limits.model_fields_set:
                limits_updates.append((field, getattr(body.limits, field)))

        models_update = body.models if "models" in provided and body.models is not None else None

        def _apply(conn: sqlite3.Connection) -> None:
            if column_updates:
                assignments = ", ".join(f"{name} = ?" for name, _ in column_updates)
                conn.execute(
                    f"UPDATE api_keys SET {assignments} WHERE id = ?",
                    [value for _, value in column_updates] + [key_id],
                )
            if limits_updates:
                conn.execute(
                    "INSERT OR IGNORE INTO api_key_limits (key_id) VALUES (?)",
                    (key_id,),
                )
                assignments = ", ".join(f"{name} = ?" for name, _ in limits_updates)
                conn.execute(
                    f"UPDATE api_key_limits SET {assignments} WHERE key_id = ?",
                    [value for _, value in limits_updates] + [key_id],
                )
            if models_update is not None:
                keys_module.replace_models(conn, key_id, models_update)

        await db.run(_apply)
        await keys.invalidate()
        after = keys.get_by_id(key_id)
        if after is None:  # pragma: no cover
            raise errors.not_found("The requested API key does not exist.")

        actor, actor_ip = _actor(request)
        await db.write_audit(
            action=ACTION_UPDATE,
            actor=actor,
            target_id=str(key_id),
            before=_audit_snapshot(before),
            after=_audit_snapshot(after),
            source_ip=actor_ip or (request.client.host if request.client else None),
        )
        usage = await _usage_24h(db)
        return JSONResponse({"key": _key_dict(after, usage)})

    @router.delete("/keys/{key_id}", status_code=status.HTTP_204_NO_CONTENT)
    async def delete_key(key_id: int, request: Request) -> Response:
        before = keys.get_by_id(key_id)
        if before is None:
            await keys.reload()
            before = keys.get_by_id(key_id)
        if before is None:
            raise errors.not_found("The requested API key does not exist.")

        def _delete(conn: sqlite3.Connection) -> None:
            conn.execute("DELETE FROM api_key_models WHERE key_id = ?", (key_id,))
            conn.execute("DELETE FROM api_key_limits WHERE key_id = ?", (key_id,))
            conn.execute('DELETE FROM rate_snapshots WHERE key_id = ?', (key_id,))
            conn.execute("DELETE FROM api_keys WHERE id = ?", (key_id,))

        await db.run(_delete)
        await keys.invalidate()
        if limiter is not None:
            limiter.forget(key_id)

        actor, actor_ip = _actor(request)
        await db.write_audit(
            action=ACTION_DELETE,
            actor=actor,
            target_id=str(key_id),
            before=_audit_snapshot(before),
            source_ip=actor_ip or (request.client.host if request.client else None),
        )
        return Response(status_code=status.HTTP_204_NO_CONTENT)

    # ------------------------------------------------------------------ 使用量

    @router.get("/usage")
    async def usage(days: int = Query(default=7, ge=1, le=90)) -> JSONResponse:
        since = iso_jst(now_jst() - timedelta(days=days))

        totals_row = await db.fetchone(
            "SELECT COUNT(*) AS requests,"
            " COALESCE(SUM(prompt_tokens), 0) AS prompt_tokens,"
            " COALESCE(SUM(completion_tokens), 0) AS completion_tokens"
            " FROM request_logs WHERE ts >= ?",
            (since,),
        )
        by_day_rows = await db.fetchall(
            "SELECT substr(ts, 1, 10) AS date, COUNT(*) AS requests,"
            " COALESCE(SUM(prompt_tokens), 0) AS prompt_tokens,"
            " COALESCE(SUM(completion_tokens), 0) AS completion_tokens"
            " FROM request_logs WHERE ts >= ? GROUP BY date ORDER BY date",
            (since,),
        )
        by_key_rows = await db.fetchall(
            "SELECT r.key_id AS key_id, COALESCE(k.name, '') AS name, COUNT(*) AS requests,"
            " COALESCE(SUM(r.prompt_tokens), 0) AS prompt_tokens,"
            " COALESCE(SUM(r.completion_tokens), 0) AS completion_tokens"
            " FROM request_logs r LEFT JOIN api_keys k ON k.id = r.key_id"
            " WHERE r.ts >= ? AND r.key_id IS NOT NULL"
            " GROUP BY r.key_id, name ORDER BY requests DESC",
            (since,),
        )
        by_model_rows = await db.fetchall(
            "SELECT model, COUNT(*) AS requests FROM request_logs"
            " WHERE ts >= ? AND model IS NOT NULL AND model <> ''"
            " GROUP BY model ORDER BY requests DESC",
            (since,),
        )
        error_rows = await db.fetchall(
            "SELECT ts, key_id, path, status FROM request_logs"
            " WHERE ts >= ? AND status >= 400 ORDER BY ts DESC LIMIT 50",
            (since,),
        )

        return JSONResponse(
            {
                "totals": {
                    "requests": int(totals_row["requests"]) if totals_row else 0,
                    "prompt_tokens": int(totals_row["prompt_tokens"]) if totals_row else 0,
                    "completion_tokens": int(totals_row["completion_tokens"]) if totals_row else 0,
                },
                "by_day": [
                    {
                        "date": row["date"],
                        "requests": int(row["requests"]),
                        "prompt_tokens": int(row["prompt_tokens"]),
                        "completion_tokens": int(row["completion_tokens"]),
                    }
                    for row in by_day_rows
                ],
                "by_key": [
                    {
                        "key_id": int(row["key_id"]),
                        "name": row["name"],
                        "requests": int(row["requests"]),
                        "prompt_tokens": int(row["prompt_tokens"]),
                        "completion_tokens": int(row["completion_tokens"]),
                    }
                    for row in by_key_rows
                ],
                "by_model": [
                    {"model": row["model"], "requests": int(row["requests"])} for row in by_model_rows
                ],
                "recent_errors": [
                    {
                        "ts": row["ts"],
                        "key_id": row["key_id"],
                        "path": row["path"],
                        "status": int(row["status"]),
                    }
                    for row in error_rows
                ],
            }
        )

    # ------------------------------------------------------------------ モデル

    @router.get("/models")
    async def models() -> JSONResponse:
        snapshot = await upstream.get_models()
        return JSONResponse(
            {
                "models": list(snapshot.models),
                "fetched_at": snapshot.fetched_at,
                "stale": snapshot.stale,
            }
        )

    # ------------------------------------------------------------------ 監査

    @router.get("/audit")
    async def audit(limit: int = Query(default=100, ge=1, le=1000)) -> JSONResponse:
        rows = await db.fetchall(
            "SELECT id, ts, actor, action, target_id, before, after, source_ip"
            " FROM admin_audit_log ORDER BY id DESC LIMIT ?",
            (limit,),
        )
        return JSONResponse({"entries": [_audit_row(row) for row in rows]})

    app.include_router(router)
    _install_error_handlers(app)
    return app


async def _probe_database(db: Database) -> str:
    """読み書きの両方を確かめる。

    読み取りだけでは、ディスク満杯や読み取り専用マウントを見逃す。
    いずれもキー発行と利用ログの記録が止まる致命的な状態である。
    """

    stamp = iso_jst()

    def _do(conn: sqlite3.Connection) -> str | None:
        conn.execute(
            "INSERT INTO settings (k, v) VALUES ('healthz_at', ?)"
            " ON CONFLICT(k) DO UPDATE SET v = excluded.v",
            (stamp,),
        )
        row = conn.execute("SELECT v FROM settings WHERE k = 'healthz_at'").fetchone()
        return row["v"] if row else None

    try:
        written = await db.run(_do)
    except Exception:
        logger.exception("ヘルスチェックでデータベースへ読み書きできませんでした")
        return "error"
    return "ok" if written == stamp else "error"


def _audit_row(row: Any) -> dict[str, Any]:
    import json

    def _load(value: str | None) -> Any:
        if value is None:
            return None
        try:
            return json.loads(value)
        except ValueError:
            return value

    return {
        "id": int(row["id"]),
        "ts": row["ts"],
        "actor": row["actor"],
        "action": row["action"],
        "target_id": row["target_id"],
        "before": _load(row["before"]),
        "after": _load(row["after"]),
        "source_ip": row["source_ip"],
    }


def _install_error_handlers(app: FastAPI) -> None:
    """全ての 4xx / 5xx を共通エラー形へ揃える。"""

    @app.exception_handler(errors.GatewayError)
    async def _gateway_error(_: Request, exc: errors.GatewayError) -> Response:
        return exc.to_response()

    @app.exception_handler(StarletteHTTPException)
    async def _http_error(_: Request, exc: StarletteHTTPException) -> Response:
        code = "not_found" if exc.status_code == 404 else "invalid_request"
        return errors.error_response(
            exc.status_code,
            str(exc.detail),
            type=errors.TYPE_INVALID_REQUEST,
            code=code,
            headers=getattr(exc, "headers", None),
        )

    @app.exception_handler(RequestValidationError)
    async def _validation_error(_: Request, exc: RequestValidationError) -> Response:
        return errors.error_response(
            400,
            _format_validation(exc.errors()),
            type=errors.TYPE_INVALID_REQUEST,
            code="invalid_request",
        )

    @app.exception_handler(Exception)
    async def _unhandled(_: Request, exc: Exception) -> Response:  # pragma: no cover
        logger.exception("control API で予期しない例外が発生しました")
        return errors.error_response(
            500,
            "The control API encountered an internal error.",
            type=errors.TYPE_API,
            code="internal_error",
        )


def _format_validation(items: Iterable[dict[str, Any]]) -> str:
    parts: list[str] = []
    for item in items:
        location = ".".join(str(part) for part in item.get("loc", ()) if part != "body")
        message = item.get("msg", "invalid value")
        parts.append(f"{location or 'body'}: {message}")
    return "; ".join(parts) or "Invalid request."
