"""秘匿値のマスク。

内部キーと ``Authorization`` は、ログ・例外・エラー応答本文のいずれにも
そのまま現れてはならない。出力の直前に必ず本モジュールを通す。
"""

from __future__ import annotations

import logging
import re
from typing import Any, Iterable

REDACTED = "***REDACTED***"

#: 登録済みの秘匿文字列。長いものから順に置換するため降順で保持する。
_SECRETS: list[str] = []

#: ``Authorization: Bearer xxx`` 形式。ヘッダ名の直後の値のみを伏せる。
_AUTH_HEADER_RE = re.compile(
    r"(?i)\b(authorization|proxy-authorization|x-control-token|x-api-key)\b(\s*[:=]\s*)"
    r"(?:(bearer|basic|token)\s+)?([^\s,;\"'}\]]+)"
)

#: JSON / dict 表現に埋め込まれた秘密値。
_JSON_SECRET_RE = re.compile(
    r"(?i)([\"']?(?:api[_-]?key|secret|token|password|authorization)[\"']?\s*[:=]\s*[\"'])([^\"']{4,})([\"'])"
)

#: 発行キーの書式。ログに紛れ込んだ場合に備えた最後の網。
_ISSUED_KEY_RE = re.compile(r"sk-aig-[A-Za-z0-9]{8,}")


def register_secret(*values: str) -> None:
    """マスク対象の文字列を登録する。

    内部キー（``UPSTREAM_API_KEY``）と control トークンを起動時に登録する。
    短すぎる値は誤爆するため無視する。
    """

    changed = False
    for value in values:
        if not value or len(value) < 8:
            continue
        if value not in _SECRETS:
            _SECRETS.append(value)
            changed = True
    if changed:
        _SECRETS.sort(key=len, reverse=True)


def registered_secrets() -> tuple[str, ...]:
    """登録済みの秘匿文字列（試験用）。"""

    return tuple(_SECRETS)


def clear_secrets() -> None:
    """登録済みの秘匿文字列を消す（試験用）。"""

    _SECRETS.clear()


def mask_value(value: str, *, keep: int = 4) -> str:
    """値の中央を伏せる。8 文字以下は全て伏せる。"""

    if not value:
        return ""
    if len(value) <= keep * 2:
        return "****"
    return f"{value[:keep]}****{value[-keep:]}"


def scrub(text: str) -> str:
    """文字列から秘匿値を取り除く。"""

    if not text:
        return text
    for secret in _SECRETS:
        if secret in text:
            text = text.replace(secret, REDACTED)
    text = _AUTH_HEADER_RE.sub(lambda m: f"{m.group(1)}{m.group(2)}{(m.group(3) + ' ') if m.group(3) else ''}{REDACTED}", text)
    text = _JSON_SECRET_RE.sub(lambda m: f"{m.group(1)}{REDACTED}{m.group(3)}", text)
    text = _ISSUED_KEY_RE.sub(REDACTED, text)
    return text


def scrub_bytes(data: bytes) -> bytes:
    """バイト列から秘匿値を取り除く。復号できない部分はそのまま残す。"""

    if not data:
        return data
    try:
        text = data.decode("utf-8")
    except UnicodeDecodeError:
        out = data
        for secret in _SECRETS:
            out = out.replace(secret.encode("utf-8", "ignore"), REDACTED.encode())
        return out
    scrubbed = scrub(text)
    if scrubbed == text:
        return data
    return scrubbed.encode("utf-8")


def contains_secret(data: bytes) -> bool:
    """バイト列に登録済みの秘匿値が含まれるか。"""

    if not data or not _SECRETS:
        return False
    for secret in _SECRETS:
        if secret.encode("utf-8", "ignore") in data:
            return True
    return False


def safe_headers(headers: Iterable[tuple[str, str]] | dict[str, str]) -> dict[str, str]:
    """ログ出力用にヘッダ値をマスクする。"""

    items = headers.items() if isinstance(headers, dict) else headers
    out: dict[str, str] = {}
    for name, value in items:
        lowered = name.lower()
        if lowered in ("authorization", "proxy-authorization", "x-control-token", "x-api-key", "cookie"):
            out[name] = REDACTED
        else:
            out[name] = scrub(value)
    return out


class ScrubbingFilter(logging.Filter):
    """logging へ流れる全メッセージをマスクする。"""

    def filter(self, record: logging.LogRecord) -> bool:  # noqa: D102
        try:
            record.msg = scrub(str(record.msg))
            if record.args:
                if isinstance(record.args, dict):
                    record.args = {k: _scrub_any(v) for k, v in record.args.items()}
                else:
                    record.args = tuple(_scrub_any(a) for a in record.args)
        except Exception:  # pragma: no cover - ログ処理で例外を伝播させない
            return True
        return True


def _scrub_any(value: Any) -> Any:
    if isinstance(value, str):
        return scrub(value)
    if isinstance(value, bytes):
        return scrub_bytes(value)
    return value


def configure_logging(level: str = "INFO") -> None:
    """ルートロガーへマスク フィルタを取り付ける。"""

    logging.basicConfig(
        level=level.upper(),
        format="%(asctime)s %(levelname)s %(name)s %(message)s",
    )
    scrubber = ScrubbingFilter()
    root = logging.getLogger()
    for handler in root.handlers:
        handler.addFilter(scrubber)
    for name in ("uvicorn", "uvicorn.error", "uvicorn.access", "httpx", "httpcore"):
        logger = logging.getLogger(name)
        logger.addFilter(scrubber)
        for handler in logger.handlers:
            handler.addFilter(scrubber)
