"""パス方針・モデル検査・アドミッション コントロール。

いずれも「拒否するか通すか」だけを判断し、値の書き換えは一切しない。
モデル検査は受信したバイト列の**コピー**に対して行い、上流へは元のバイト列を送る。
"""

from __future__ import annotations

import json
import re
from typing import Any, Final
from urllib.parse import unquote_to_bytes

from . import errors
from .config import Settings

#: 拒否するプレフィックス。llama.cpp のスロット照会は処理中の他利用者の
#: プロンプトと生成中テキストを返すため、公開してはならない。
DENY_PREFIXES: Final[tuple[str, ...]] = ("/slots", "/props", "/metrics", "/lora-adapters")

#: ``model`` を持ちうるエンドポイント。旧来の別名を漏らすと迂回路になる。
MODEL_CHECK_PATHS: Final[frozenset[str]] = frozenset(
    {
        "/v1/chat/completions",
        "/v1/completions",
        "/v1/embeddings",
        "/v1/rerank",
        "/v1/reranking",
        "/completion",
        "/completions",
        "/infill",
        "/v1/infill",
        "/apply-template",
        "/v1/audio/transcriptions",
    }
)

#: 許可リスト方式で常に通すパス。
ALLOWLIST_EXTRA: Final[frozenset[str]] = frozenset({"/v1/models"})

_MODELS_PATH: Final[str] = "/v1/models"
_MULTIPART_MODEL_RE = re.compile(
    rb'(?is)content-disposition:[^\r\n]*\bname="model"[^\r\n]*\r?\n(?:[^\r\n]+\r?\n)*\r?\n(.*?)\r?\n--'
)
_CONTROL_CHARS = {0x00, 0x09, 0x0A, 0x0D}


def normalize_path(raw_path: str | bytes) -> str:
    """方針判定に使う正規化済みパスを返す。

    ``%2e%2e`` や ``%2f`` による多重符号化、``..``、連続スラッシュ、末尾スラッシュ、
    大文字小文字の差で判定を迂回されないようにする。
    転送には使わない。転送は受信した生のパスをそのまま用いる。
    """

    data = raw_path.encode("utf-8", "surrogateescape") if isinstance(raw_path, str) else bytes(raw_path)
    # 多重符号化に備えて変化しなくなるまで復号する（上限付き）。
    for _ in range(4):
        decoded = unquote_to_bytes(data)
        if decoded == data:
            break
        data = decoded
    text = data.decode("utf-8", "replace")
    text = "".join(ch for ch in text if ord(ch) >= 0x20 or ord(ch) not in _CONTROL_CHARS)
    text = text.replace("\\", "/")
    segments: list[str] = []
    for segment in text.split("/"):
        if segment in ("", "."):
            continue
        if segment == "..":
            if segments:
                segments.pop()
            continue
        segments.append(segment)
    return ("/" + "/".join(segments)).lower() if segments else "/"


def is_denied(normalized: str) -> bool:
    """拒否プレフィックスに該当するか。パス境界で判定する。

    ``/slots`` と ``/slots/0`` は該当し、``/slotsfoo`` は該当しない。
    """

    for prefix in DENY_PREFIXES:
        if normalized == prefix or normalized.startswith(prefix + "/"):
            return True
    return False


def is_allowed(normalized: str) -> bool:
    """許可リスト方式で通すか。"""

    if normalized in MODEL_CHECK_PATHS or normalized in ALLOWLIST_EXTRA:
        return True
    return normalized == "/v1" or normalized.startswith("/v1/")


def check_path(normalized: str, settings: Settings) -> None:
    """方針に反するパスなら 404 を送出する。

    存在の有無を漏らさないため、拒否も許可外も同じ 404 とする。
    """

    if is_denied(normalized):
        raise errors.not_found()
    if settings.allowlist_mode and not is_allowed(normalized):
        raise errors.not_found()


def requires_model_check(normalized: str) -> bool:
    """モデル検査の対象パスか。"""

    return normalized in MODEL_CHECK_PATHS


def is_models_path(normalized: str) -> bool:
    """``/v1/models`` か。"""

    return normalized == _MODELS_PATH


def parse_json_body(body: bytes) -> dict[str, Any] | None:
    """検査用に本文のコピーを JSON として読む。読めなければ ``None``。"""

    if not body:
        return None
    try:
        parsed = json.loads(body)
    except (ValueError, UnicodeDecodeError):
        return None
    return parsed if isinstance(parsed, dict) else None


def extract_model(body: bytes, content_type: str) -> str | None:
    """本文からモデル名を取り出す。取り出せなければ ``None``。

    ``None`` は「検査不能」を意味し、許可リストのあるキーでは拒否につながる。
    """

    lowered = (content_type or "").lower()
    if "multipart/form-data" in lowered:
        return _extract_model_multipart(body)
    payload = parse_json_body(body)
    if payload is None:
        return None
    model = payload.get("model")
    if isinstance(model, str) and model.strip():
        return model.strip()
    return None


def _extract_model_multipart(body: bytes) -> str | None:
    """``/v1/audio/transcriptions`` 等の multipart から ``model`` を取り出す。"""

    if not body:
        return None
    match = _MULTIPART_MODEL_RE.search(body)
    if not match:
        return None
    try:
        value = match.group(1).decode("utf-8").strip()
    except UnicodeDecodeError:
        return None
    return value or None


def enforce_model(body: bytes, content_type: str, allowed: frozenset[str]) -> str | None:
    """モデル許可を判定する。フェイル クローズ。

    許可リストが空なら無制限として検査を省略する。
    許可リストがある場合、モデル名を取り出せなければ 403 で拒否する。
    """

    if not allowed:
        # 無制限キー。記録用にモデル名だけ拾う（失敗しても構わない）。
        return extract_model(body, content_type)

    model = extract_model(body, content_type)
    if model is None:
        raise errors.forbidden_model(
            "The model could not be determined from the request body, and this API key is"
            " restricted to specific models.",
            code="model_check_failed",
        )
    if model not in allowed:
        raise errors.forbidden_model(f"The model `{model}` is not allowed for this API key.")
    return model


def check_admission(body: bytes, settings: Settings) -> None:
    """逸脱パラメータを拒否する。値の書き換えは行わない。"""

    payload = parse_json_body(body)
    if payload is None:
        # JSON として読めない本文は対象の項目を持ちえない。
        # モデル検査はこの経路でもフェイル クローズで別途拒否される。
        return

    grammar = payload.get("grammar")
    if isinstance(grammar, str):
        size = len(grammar.encode("utf-8", "ignore"))
        if size > settings.max_grammar_bytes:
            raise errors.bad_request(
                f"`grammar` exceeds the maximum allowed size of {settings.max_grammar_bytes} bytes.",
                code="grammar_too_large",
            )

    n_probs = payload.get("n_probs")
    if isinstance(n_probs, bool):
        n_probs = None
    if isinstance(n_probs, (int, float)) and n_probs > settings.max_n_probs:
        raise errors.bad_request(
            f"`n_probs` exceeds the maximum allowed value of {settings.max_n_probs}.",
            code="n_probs_too_large",
        )


def wants_stream(body: bytes) -> bool:
    """クライアントがストリーミングを要求しているか。"""

    payload = parse_json_body(body)
    return bool(payload and payload.get("stream") is True)


def filter_models_payload(raw: bytes, allowed: frozenset[str]) -> bytes | None:
    """``/v1/models`` の応答を許可モデルで絞る。

    形が想定と違う場合は ``None`` を返し、呼び出し側は原本をそのまま返す。
    """

    if not allowed:
        return None
    try:
        payload = json.loads(raw)
    except (ValueError, UnicodeDecodeError):
        return None
    if not isinstance(payload, dict):
        return None
    data = payload.get("data")
    if not isinstance(data, list):
        return None
    payload["data"] = [
        item for item in data if isinstance(item, dict) and str(item.get("id", "")) in allowed
    ]
    return json.dumps(payload, ensure_ascii=False).encode("utf-8")
