"""control API。``control-api-contract.md`` との一致を確かめる。"""

from __future__ import annotations

import json

import httpx
import pytest

from .conftest import CONTROL_TOKEN, create_key, stream_response

TOKEN = {"x-control-token": CONTROL_TOKEN}
ACTOR = {"x-actor": "admin", "x-actor-ip": "203.0.113.1"}

KEY_FIELDS = {
    "id",
    "name",
    "key_prefix",
    "enabled",
    "created_at",
    "expires_at",
    "note",
    "limits",
    "models",
    "usage_24h",
}


def models_upstream(models: list[str]):
    payload = json.dumps({"object": "list", "data": [{"id": name} for name in models]}).encode()

    def respond(request: httpx.Request) -> httpx.Response:
        return stream_response(200, content=payload)

    return respond


# ------------------------------------------------------------------ 認証


async def test_missing_token_is_401(build_control, db):
    client = build_control()
    async with client:
        response = await client.get("/keys")

    assert response.status_code == 401
    payload = response.json()
    assert payload == {
        "error": {
            "message": payload["error"]["message"],
            "type": "invalid_request_error",
            "code": "unauthorized",
        }
    }

    row = await db.fetchone("SELECT action, source_ip FROM admin_audit_log ORDER BY id DESC LIMIT 1")
    assert row is not None
    assert row["action"] == "auth.fail"


async def test_wrong_token_is_401(build_control):
    client = build_control()
    async with client:
        response = await client.get("/keys", headers={"x-control-token": "wrong"})
    assert response.status_code == 401


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


def unreachable_upstream(request: httpx.Request) -> httpx.Response:
    raise httpx.ConnectError("refused", request=request)


async def test_healthz_needs_no_token(build_control):
    client = build_control(models_upstream(["qwen3-8b"]))
    async with client:
        response = await client.get("/healthz")

    assert response.status_code == 200
    payload = response.json()
    assert set(payload) == {"status", "upstream", "version", "checks"}
    assert payload["status"] == "ok"
    assert payload["checks"] == {"database": "ok", "listeners": "ok"}


async def test_healthz_stays_200_while_upstream_is_down(build_control):
    """上流の障害でコンテナを unhealthy にしない。

    unhealthy になると ``depends_on`` で待つ admin まで起動できず、
    上流障害時こそ必要な管理画面を失う。
    """

    client = build_control(unreachable_upstream)
    async with client:
        response = await client.get("/healthz")

    assert response.status_code == 200
    payload = response.json()
    assert payload["status"] == "ok"
    assert payload["checks"]["database"] == "ok"


async def test_healthz_does_not_call_the_upstream(build_control, recorder):
    """既定の経路は上流を叩かない。遅い上流に healthcheck を巻き込ませない。"""

    client = build_control(unreachable_upstream)
    async with client:
        await client.get("/healthz")

    assert recorder.requests == []


async def test_healthz_reports_last_known_upstream_state(build_control):
    """上流の状態は情報として返す。既定では最後に観測した値。"""

    client = build_control(models_upstream(["qwen3-8b"]))
    async with client:
        before = await client.get("/healthz")
        await client.get("/models", headers=TOKEN)
        after = await client.get("/healthz")

    assert before.json()["upstream"] == "unknown"
    assert after.json()["upstream"] == "ok"


async def test_healthz_strict_is_503_when_upstream_is_down(build_control):
    client = build_control(unreachable_upstream)
    async with client:
        response = await client.get("/healthz", params={"strict": "1"})

    assert response.status_code == 503
    payload = response.json()
    assert payload["upstream"] == "unreachable"
    assert payload["status"] == "ok", "ゲートウェイ自身は正常である"


async def test_healthz_strict_is_200_when_upstream_is_up(build_control, recorder):
    client = build_control(models_upstream(["qwen3-8b"]))
    async with client:
        response = await client.get("/healthz", params={"strict": "1"})

    assert response.status_code == 200
    assert response.json()["upstream"] == "ok"
    assert recorder.requests, "strict では能動的に上流を確かめる"


@pytest.mark.parametrize("value", ["1", "true", "yes", "on"])
async def test_healthz_strict_accepts_common_truthy_values(build_control, value):
    client = build_control(unreachable_upstream)
    async with client:
        response = await client.get("/healthz", params={"strict": value})
    assert response.status_code == 503


async def test_healthz_is_503_when_the_database_is_unusable(build_control, db):
    """ゲートウェイ自身が壊れている場合は従来どおり 503。"""

    client = build_control(models_upstream(["m"]))
    await db.close()
    async with client:
        response = await client.get("/healthz")

    assert response.status_code == 503
    payload = response.json()
    assert payload["status"] == "degraded"
    assert payload["checks"]["database"] == "error"


async def test_healthz_is_503_when_the_public_listener_is_down(build_control):
    client = build_control(models_upstream(["m"]), listeners_alive=lambda: False)
    async with client:
        response = await client.get("/healthz")

    assert response.status_code == 503
    payload = response.json()
    assert payload["status"] == "degraded"
    assert payload["checks"]["listeners"] == "down"


async def test_healthz_write_probe_touches_the_database(build_control, db):
    """読み取りだけではディスク満杯や読み取り専用マウントを見逃す。"""

    client = build_control(models_upstream(["m"]))
    async with client:
        await client.get("/healthz")

    assert await db.get_setting("healthz_at")


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


async def test_create_key_returns_secret_once(build_control, keys):
    client = build_control()
    async with client:
        response = await client.post(
            "/keys",
            headers={**TOKEN, **ACTOR},
            json={
                "name": "team-a",
                "note": "memo",
                "expires_at": None,
                "limits": {"rpm": 60, "rpd": 5000, "concurrency": 2, "tpm": None},
                "models": ["qwen3-8b"],
            },
        )

    assert response.status_code == 201
    payload = response.json()
    assert set(payload) == {"key", "secret"}
    assert payload["secret"].startswith("sk-aig-")
    assert len(payload["secret"]) == 7 + 32

    key = payload["key"]
    assert set(key) == KEY_FIELDS
    assert key["name"] == "team-a"
    assert key["key_prefix"] == payload["secret"][:12]
    assert key["enabled"] is True
    assert key["limits"] == {"rpm": 60, "rpd": 5000, "concurrency": 2, "tpm": None}
    assert key["models"] == ["qwen3-8b"]
    assert key["usage_24h"] == {"requests": 0, "prompt_tokens": 0, "completion_tokens": 0}
    assert "+09:00" in key["created_at"]


async def test_secret_is_never_returned_again(build_control, db):
    client = build_control()
    async with client:
        created = await client.post("/keys", headers=TOKEN, json={"name": "a"})
        listed = await client.get("/keys", headers=TOKEN)

    secret = created.json()["secret"]
    assert secret not in listed.text

    rows = await db.fetchall("SELECT key_hash FROM api_keys")
    assert all(secret not in row["key_hash"] for row in rows)


async def test_list_keys_shape(build_control, db, keys, settings):
    await create_key(db, keys, settings, name="existing", models=["m1"], rpm=10)
    client = build_control()
    async with client:
        response = await client.get("/keys", headers=TOKEN)

    assert response.status_code == 200
    payload = response.json()
    assert set(payload) == {"keys"}
    assert len(payload["keys"]) == 1
    assert set(payload["keys"][0]) == KEY_FIELDS
    assert payload["keys"][0]["limits"]["rpm"] == 10


async def test_list_keys_can_exclude_disabled(build_control, db, keys, settings):
    await create_key(db, keys, settings, name="on")
    await create_key(db, keys, settings, name="off", enabled=False)
    client = build_control()
    async with client:
        everything = await client.get("/keys", headers=TOKEN)
        enabled_only = await client.get("/keys", headers=TOKEN, params={"include_disabled": "false"})

    assert len(everything.json()["keys"]) == 2
    assert [key["name"] for key in enabled_only.json()["keys"]] == ["on"]


async def test_patch_updates_only_given_fields(build_control):
    client = build_control()
    async with client:
        created = await client.post(
            "/keys",
            headers=TOKEN,
            json={"name": "a", "limits": {"rpm": 60, "rpd": 100, "concurrency": 2}, "models": ["m1"]},
        )
        key_id = created.json()["key"]["id"]
        response = await client.patch(
            f"/keys/{key_id}", headers={**TOKEN, **ACTOR}, json={"limits": {"rpm": 5}}
        )

    assert response.status_code == 200
    limits = response.json()["key"]["limits"]
    assert limits == {"rpm": 5, "rpd": 100, "concurrency": 2, "tpm": None}
    assert response.json()["key"]["models"] == ["m1"]


async def test_patch_replaces_models_wholesale(build_control):
    client = build_control()
    async with client:
        created = await client.post("/keys", headers=TOKEN, json={"name": "a", "models": ["m1", "m2"]})
        key_id = created.json()["key"]["id"]
        response = await client.patch(f"/keys/{key_id}", headers=TOKEN, json={"models": ["m3"]})

    assert response.json()["key"]["models"] == ["m3"]


async def test_patch_can_disable_and_set_expiry(build_control, keys):
    client = build_control()
    async with client:
        created = await client.post("/keys", headers=TOKEN, json={"name": "a"})
        key_id = created.json()["key"]["id"]
        response = await client.patch(
            f"/keys/{key_id}",
            headers=TOKEN,
            json={"enabled": False, "expires_at": "2026-12-31T23:59:59+09:00"},
        )

    key = response.json()["key"]
    assert key["enabled"] is False
    assert key["expires_at"] == "2026-12-31T23:59:59+09:00"
    assert keys.get_by_id(key_id).enabled is False


async def test_patch_unknown_key_is_404(build_control):
    client = build_control()
    async with client:
        response = await client.patch("/keys/999", headers=TOKEN, json={"name": "x"})
    assert response.status_code == 404
    assert response.json()["error"]["code"] == "not_found"


async def test_delete_returns_204_and_keeps_audit(build_control, db, keys):
    client = build_control()
    async with client:
        created = await client.post("/keys", headers=TOKEN, json={"name": "doomed", "models": ["m"]})
        key_id = created.json()["key"]["id"]
        response = await client.delete(f"/keys/{key_id}", headers={**TOKEN, **ACTOR})
        listed = await client.get("/keys", headers=TOKEN)

    assert response.status_code == 204
    assert response.content == b""
    assert listed.json()["keys"] == []
    assert keys.get_by_id(key_id) is None

    row = await db.fetchone(
        "SELECT action, target_id, before FROM admin_audit_log WHERE action = 'key.delete'"
    )
    assert row is not None
    assert row["target_id"] == str(key_id)
    assert json.loads(row["before"])["name"] == "doomed"


async def test_deleted_key_is_rejected_by_the_proxy(build_control, db, keys, settings, build_proxy):
    from .conftest import auth

    control = build_control()
    async with control:
        created = await control.post("/keys", headers=TOKEN, json={"name": "temp"})
        secret = created.json()["secret"]
        key_id = created.json()["key"]["id"]

        proxy, _ = build_proxy()
        async with proxy:
            before = await proxy.post("/v1/chat/completions", headers=auth(secret), content=b'{"model":"m"}')
            await control.delete(f"/keys/{key_id}", headers=TOKEN)
            after = await proxy.post("/v1/chat/completions", headers=auth(secret), content=b'{"model":"m"}')

    assert before.status_code == 200
    assert after.status_code == 401


async def test_invalid_payload_uses_the_common_error_envelope(build_control):
    client = build_control()
    async with client:
        response = await client.post("/keys", headers=TOKEN, json={"note": "no name"})

    assert response.status_code == 400
    assert set(response.json()["error"]) == {"message", "type", "code"}


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


async def test_usage_shape(build_control, db, keys, settings, build_proxy, logs):
    from .conftest import auth

    secret = await create_key(db, keys, settings, name="team-a")
    proxy, _ = build_proxy()
    async with proxy:
        await proxy.post("/v1/chat/completions", headers=auth(secret), content=b'{"model":"qwen3-8b"}')
        await proxy.get("/slots", headers=auth(secret))
    await logs.stop()

    client = build_control()
    async with client:
        response = await client.get("/usage", headers=TOKEN, params={"days": 7})

    payload = response.json()
    assert set(payload) == {"totals", "by_day", "by_key", "by_model", "recent_errors"}
    assert set(payload["totals"]) == {"requests", "prompt_tokens", "completion_tokens"}
    assert payload["totals"]["requests"] == 2
    assert payload["by_day"][0]["requests"] == 2
    assert payload["by_key"][0]["name"] == "team-a"
    assert payload["by_model"][0] == {"model": "qwen3-8b", "requests": 1}
    assert payload["recent_errors"][0]["status"] == 404


async def test_usage_days_is_capped_at_90(build_control):
    client = build_control()
    async with client:
        ok = await client.get("/usage", headers=TOKEN, params={"days": 90})
        ng = await client.get("/usage", headers=TOKEN, params={"days": 91})
    assert ok.status_code == 200
    assert ng.status_code == 400


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


async def test_models_returns_upstream_list(build_control):
    client = build_control(models_upstream(["qwen3-8b", "gemma3-12b"]))
    async with client:
        response = await client.get("/models", headers=TOKEN)

    payload = response.json()
    assert set(payload) == {"models", "fetched_at", "stale"}
    assert payload["models"] == ["qwen3-8b", "gemma3-12b"]
    assert payload["stale"] is False
    assert payload["fetched_at"]


async def test_models_returns_stale_marker_when_upstream_is_down(build_control):
    def boom(request: httpx.Request) -> httpx.Response:
        raise httpx.ConnectError("refused", request=request)

    client = build_control(boom)
    async with client:
        response = await client.get("/models", headers=TOKEN)

    payload = response.json()
    assert payload["models"] == []
    assert payload["stale"] is True
    assert payload["fetched_at"] is None


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


async def test_audit_shape(build_control):
    client = build_control()
    async with client:
        await client.post("/keys", headers={**TOKEN, **ACTOR}, json={"name": "a"})
        response = await client.get("/audit", headers=TOKEN, params={"limit": 10})

    entries = response.json()["entries"]
    assert entries
    entry = entries[0]
    assert set(entry) == {"id", "ts", "actor", "action", "target_id", "before", "after", "source_ip"}
    assert entry["action"] == "key.create"
    assert entry["actor"] == "admin"
    assert entry["source_ip"] == "203.0.113.1"
    assert entry["before"] is None
    assert entry["after"]["name"] == "a"


async def test_audit_records_every_mutation(build_control):
    client = build_control()
    async with client:
        created = await client.post("/keys", headers=TOKEN, json={"name": "a"})
        key_id = created.json()["key"]["id"]
        await client.patch(f"/keys/{key_id}", headers=TOKEN, json={"name": "b"})
        await client.delete(f"/keys/{key_id}", headers=TOKEN)
        response = await client.get("/audit", headers=TOKEN)

    actions = [entry["action"] for entry in response.json()["entries"]]
    assert actions[:3] == ["key.delete", "key.update", "key.create"]


async def test_audit_never_contains_secrets(build_control):
    client = build_control()
    async with client:
        created = await client.post("/keys", headers=TOKEN, json={"name": "a"})
        response = await client.get("/audit", headers=TOKEN)

    assert created.json()["secret"] not in response.text


async def test_audit_limit_is_capped(build_control):
    client = build_control()
    async with client:
        ok = await client.get("/audit", headers=TOKEN, params={"limit": 1000})
        ng = await client.get("/audit", headers=TOKEN, params={"limit": 1001})
    assert ok.status_code == 200
    assert ng.status_code == 400
