"""試験共通のフィクスチャ。

上流は ``httpx.MockTransport`` で置き換え、ゲートウェイが実際に送出した
リクエストのバイト列とヘッダをそのまま検証できるようにする。
"""

from __future__ import annotations

import dataclasses
import sqlite3
from typing import Any, Callable, Iterable, Sequence

import httpx
import pytest

from app import logging_utils
from app.config import Settings, load_settings
from app.control import create_control_app
from app.db import Database, LogWriter
from app.keys import (
    KeyCache,
    KeyLimits,
    generate_secret,
    hash_secret,
    insert_key,
    key_prefix,
)
from app.limits import RateLimiter
from app.proxy import create_proxy_app
from app.upstream import UpstreamClient

UPSTREAM_BASE = "http://upstream.test"
INTERNAL_KEY = "internal-upstream-key-DO-NOT-LEAK-0123456789"
CONTROL_TOKEN = "control-token-for-tests-0123456789"


class AsyncChunks(httpx.AsyncByteStream):
    """任意のチャンク列を返す応答本体。"""

    def __init__(self, chunks: Sequence[bytes]) -> None:
        self._chunks = list(chunks)

    async def __aiter__(self):  # noqa: D105
        for chunk in self._chunks:
            yield chunk


class Recorder:
    """上流が受け取ったリクエストを記録する。"""

    def __init__(self) -> None:
        self.requests: list[httpx.Request] = []

    @property
    def last(self) -> httpx.Request:
        assert self.requests, "上流へのリクエストが記録されていません"
        return self.requests[-1]

    def header(self, name: str) -> str | None:
        return self.last.headers.get(name)

    def header_names(self) -> set[str]:
        return {name.decode("latin-1").lower() for name, _ in self.last.headers.raw}


@pytest.fixture(autouse=True)
def _reset_secrets():
    """秘匿値の登録は全体で共有されるため、試験ごとに初期化する。"""

    logging_utils.clear_secrets()
    yield
    logging_utils.clear_secrets()


@pytest.fixture
def settings(tmp_path) -> Settings:
    base = load_settings(require_secrets=False)
    return dataclasses.replace(
        base,
        upstream_base_url=UPSTREAM_BASE,
        upstream_api_key=INTERNAL_KEY,
        control_token=CONTROL_TOKEN,
        db_path=str(tmp_path / "gateway.db"),
        global_concurrency=8,
        global_queue_timeout=0.0,
        max_request_body_bytes=1024 * 1024,
        max_grammar_bytes=1024,
        max_n_probs=20,
        max_request_duration=30.0,
        idle_timeout=5.0,
    )


@pytest.fixture
async def db(settings: Settings):
    database = Database(settings.db_path)
    await database.connect()
    try:
        yield database
    finally:
        await database.close()


@pytest.fixture
async def keys(db: Database, settings: Settings) -> KeyCache:
    cache = KeyCache(db=db, pepper=settings.key_pepper, ttl=settings.key_cache_ttl)
    await cache.reload()
    return cache


@pytest.fixture
def limiter(settings: Settings) -> RateLimiter:
    return RateLimiter(settings)


@pytest.fixture
async def logs(db: Database):
    writer = LogWriter(db)
    await writer.start()
    try:
        yield writer
    finally:
        await writer.stop()


@pytest.fixture
def recorder() -> Recorder:
    return Recorder()


def stream_response(
    status_code: int = 200,
    *,
    content: bytes = b'{"ok":true}',
    headers: Sequence[tuple[bytes, bytes]] | None = None,
    content_type: str = "application/json",
    chunks: Sequence[bytes] | None = None,
) -> httpx.Response:
    """未読の応答を組み立てる。

    ``httpx.Response(json=...)`` は生成時に本文を読み込んでしまい、
    ``aiter_raw()`` が使えなくなる。実際の通信と同じ条件を保つため常に stream を渡す。
    """

    body = list(chunks) if chunks is not None else [content]
    total = sum(len(chunk) for chunk in body)
    base: list[tuple[bytes, bytes]] = []
    if content_type:
        base.append((b"content-type", content_type.encode("latin-1")))
    if headers is None:
        base.append((b"content-length", str(total).encode("latin-1")))
    else:
        base.extend(headers)
    return httpx.Response(status_code, headers=base, stream=AsyncChunks(body))


def echo_handler(recorder: Recorder, response_factory: Callable[[httpx.Request], httpx.Response] | None = None):
    """受け取ったリクエストを記録して定型の応答を返すハンドラ。"""

    def handler(request: httpx.Request) -> httpx.Response:
        recorder.requests.append(request)
        if response_factory is not None:
            return response_factory(request)
        return stream_response()

    return handler


@pytest.fixture
def make_upstream(settings: Settings, recorder: Recorder):
    """上流クライアントを差し替えるための組み立て関数。"""

    created: list[UpstreamClient] = []

    def _make(
        response_factory: Callable[[httpx.Request], httpx.Response] | None = None,
    ) -> UpstreamClient:
        client = httpx.AsyncClient(transport=httpx.MockTransport(echo_handler(recorder, response_factory)))
        upstream = UpstreamClient(settings, client=client)
        created.append(upstream)
        return upstream

    return _make


@pytest.fixture
def build_proxy(settings: Settings, keys: KeyCache, limiter: RateLimiter, make_upstream, logs: LogWriter):
    """公開用アプリと、それを叩くクライアントを組み立てる。"""

    def _build(
        response_factory: Callable[[httpx.Request], httpx.Response] | None = None,
        *,
        overrides: dict[str, Any] | None = None,
    ) -> tuple[httpx.AsyncClient, Settings]:
        effective = dataclasses.replace(settings, **(overrides or {}))
        upstream = make_upstream(response_factory)
        upstream._settings = effective  # type: ignore[attr-defined]
        app = create_proxy_app(
            settings=effective, keys=keys, limiter=limiter, upstream=upstream, logs=logs
        )
        client = httpx.AsyncClient(
            transport=httpx.ASGITransport(app=app), base_url="http://gateway.test"
        )
        return client, effective

    return _build


@pytest.fixture
def build_control(settings: Settings, db: Database, keys: KeyCache, limiter: RateLimiter, make_upstream):
    def _build(
        response_factory: Callable[[httpx.Request], httpx.Response] | None = None,
        *,
        listeners_alive: Callable[[], bool] | None = None,
    ) -> httpx.AsyncClient:
        upstream = make_upstream(response_factory)
        app = create_control_app(
            settings=settings,
            db=db,
            keys=keys,
            upstream=upstream,
            limiter=limiter,
            listeners_alive=listeners_alive,
        )
        return httpx.AsyncClient(
            transport=httpx.ASGITransport(app=app), base_url="http://control.test"
        )

    return _build


async def create_key(
    db: Database,
    keys: KeyCache,
    settings: Settings,
    *,
    name: str = "test",
    models: Iterable[str] = (),
    rpm: int | None = None,
    rpd: int | None = None,
    concurrency: int | None = None,
    tpm: int | None = None,
    expires_at: str | None = None,
    enabled: bool = True,
) -> str:
    """キーを 1 本発行し、平文を返す。"""

    secret = generate_secret()

    def _insert(conn: sqlite3.Connection) -> int:
        key_id = insert_key(
            conn,
            name=name,
            key_hash=hash_secret(secret, settings.key_pepper),
            prefix=key_prefix(secret),
            note="",
            expires_at=expires_at,
            limits=KeyLimits(rpm=rpm, rpd=rpd, concurrency=concurrency, tpm=tpm),
            models=list(models),
        )
        if not enabled:
            conn.execute("UPDATE api_keys SET enabled = 0 WHERE id = ?", (key_id,))
        return key_id

    await db.run(_insert)
    await keys.reload()
    return secret


async def send_raw(
    client: httpx.AsyncClient,
    method: str,
    url: str,
    *,
    headers: dict[str, str] | None = None,
    content: bytes | None = None,
    drop: Iterable[str] = (),
) -> httpx.Response:
    """httpx の既定ヘッダを取り除いた上で送出する。

    「クライアントが送っていないヘッダ」を再現するために使う。
    """

    request = client.build_request(method, url, headers=headers, content=content)
    for name in drop:
        if name in request.headers:
            del request.headers[name]
    return await client.send(request)


def auth(secret: str) -> dict[str, str]:
    return {"authorization": f"Bearer {secret}"}
