"""スキーマ移行・キー保管・ログ書き込み・起動時の配線。"""

from __future__ import annotations

import hashlib
import sqlite3
import string
from datetime import timedelta

import pytest

from app import keys as keys_module
from app.db import (
    Database,
    LogWriter,
    MaintenanceTask,
    RequestLogEntry,
    iso_jst,
    now_jst,
    parse_iso,
)
from app.main import build_servers

# ------------------------------------------------------------------ スキーマ


async def test_schema_version_and_wal(db: Database) -> None:
    version = await db.run(lambda conn: conn.execute("PRAGMA user_version").fetchone()[0])
    assert version == 1
    mode = await db.run(lambda conn: conn.execute("PRAGMA journal_mode").fetchone()[0])
    assert mode.lower() == "wal"


async def test_required_indexes_exist(db: Database) -> None:
    """索引が無いと集計と日次削除が全表走査になる。"""

    rows = await db.fetchall("SELECT name FROM sqlite_master WHERE type = 'index'")
    names = {row["name"] for row in rows}
    assert "ux_api_keys_key_hash" in names
    assert "ix_request_logs_key_ts" in names
    assert "ix_request_logs_ts" in names


async def test_migration_is_idempotent(settings) -> None:
    first = Database(settings.db_path)
    await first.connect()
    await first.close()
    second = Database(settings.db_path)
    await second.connect()
    version = await second.run(lambda conn: conn.execute("PRAGMA user_version").fetchone()[0])
    await second.close()
    assert version == 1


async def test_key_hash_is_unique(db: Database) -> None:
    def _insert(conn: sqlite3.Connection) -> None:
        keys_module.insert_key(
            conn,
            name="a",
            key_hash="dup",
            prefix="sk-aig-aaaa",
            note="",
            expires_at=None,
            limits=keys_module.KeyLimits(),
            models=[],
        )

    await db.run(_insert)
    with pytest.raises(sqlite3.IntegrityError):
        await db.run(_insert)


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


def test_secret_format() -> None:
    secret = keys_module.generate_secret()
    assert secret.startswith("sk-aig-")
    body = secret[len("sk-aig-") :]
    assert len(body) == 32
    assert all(character in string.ascii_letters + string.digits for character in body)
    assert keys_module.looks_like_secret(secret)


def test_secrets_are_distinct() -> None:
    generated = {keys_module.generate_secret() for _ in range(200)}
    assert len(generated) == 200


def test_hash_is_plain_sha256_without_pepper() -> None:
    secret = keys_module.generate_secret()
    assert keys_module.hash_secret(secret) == hashlib.sha256(secret.encode()).hexdigest()


def test_pepper_changes_the_hash() -> None:
    secret = keys_module.generate_secret()
    assert keys_module.hash_secret(secret, "pepper") != keys_module.hash_secret(secret)


def test_prefix_is_first_twelve_characters() -> None:
    secret = keys_module.generate_secret()
    assert keys_module.key_prefix(secret) == secret[:12]
    assert len(keys_module.key_prefix(secret)) == 12


@pytest.mark.parametrize(
    "candidate",
    ["", "sk-aig-", "sk-aig-short", "sk-oth-" + "a" * 32, "sk-aig-" + "a" * 33, "sk-aig-" + "!" * 32],
)
def test_malformed_secrets_are_rejected_early(candidate: str) -> None:
    assert not keys_module.looks_like_secret(candidate)


async def test_cache_resolves_only_active_keys(db, keys, settings) -> None:
    from .conftest import create_key

    secret = await create_key(db, keys, settings, expires_at="2000-01-01T00:00:00+09:00")
    record = keys.resolve(secret)
    assert record is not None
    assert record.is_active() is False


async def test_cache_reload_picks_up_changes(db, keys, settings) -> None:
    from .conftest import create_key

    secret = await create_key(db, keys, settings, name="before")
    assert keys.resolve(secret).name == "before"
    await db.execute("UPDATE api_keys SET name = 'after'")
    assert keys.resolve(secret).name == "before", "キャッシュが読み直し前に更新されています"
    await keys.reload()
    assert keys.resolve(secret).name == "after"


# ------------------------------------------------------------------ ログ書き込み


def _entry(status: int = 200) -> RequestLogEntry:
    return RequestLogEntry(
        key_id=None,
        ts=iso_jst(),
        method="POST",
        path="/v1/chat/completions",
        model="m",
        status=status,
        duration_ms=1,
        prompt_tokens=None,
        completion_tokens=None,
        client_ip="203.0.113.1",
    )


async def test_log_writer_flushes_pending_entries_on_stop(db: Database) -> None:
    """終了時に書き込み中のバッチを失わない。"""

    writer = LogWriter(db)
    await writer.start()
    for _ in range(50):
        writer.enqueue(_entry())
    await writer.stop()

    row = await db.fetchone("SELECT COUNT(*) AS c FROM request_logs")
    assert row["c"] == 50


async def test_log_writer_stop_without_start_is_safe(db: Database) -> None:
    await LogWriter(db).stop()


async def test_log_writer_drops_instead_of_blocking(db: Database) -> None:
    """待ち行列が溢れてもホット パスを止めない。"""

    writer = LogWriter(db, maxsize=2)
    for _ in range(10):
        writer.enqueue(_entry())
    assert writer.dropped == 8
    await writer.stop()


async def test_retention_purges_old_logs(db: Database, settings) -> None:
    old = iso_jst(now_jst() - timedelta(days=settings.log_retention_days + 1))
    await db.execute(
        "INSERT INTO request_logs (key_id, ts, method, path, model, status, duration_ms,"
        " prompt_tokens, completion_tokens, client_ip) VALUES (NULL, ?, 'GET', '/x', NULL, 200, 1, NULL, NULL, NULL)",
        (old,),
    )
    await db.execute(
        "INSERT INTO request_logs (key_id, ts, method, path, model, status, duration_ms,"
        " prompt_tokens, completion_tokens, client_ip) VALUES (NULL, ?, 'GET', '/x', NULL, 200, 1, NULL, NULL, NULL)",
        (iso_jst(),),
    )

    removed = await db.purge_old_logs(settings.log_retention_days)
    assert removed == 1
    row = await db.fetchone("SELECT COUNT(*) AS c FROM request_logs")
    assert row["c"] == 1


async def test_maintenance_runs_purge_and_vacuum(db: Database, settings) -> None:
    task = MaintenanceTask(db=db, settings=settings)
    await task.run_once()
    assert await db.get_setting("last_purge_at")
    assert await db.get_setting("last_vacuum_at")


def test_iso_roundtrip() -> None:
    moment = now_jst()
    assert parse_iso(iso_jst(moment)).utcoffset() == timedelta(hours=9)
    assert parse_iso(None) is None
    assert parse_iso("not a date") is None


def test_iso_strings_sort_chronologically() -> None:
    earlier = iso_jst(now_jst() - timedelta(days=1))
    later = iso_jst(now_jst())
    assert earlier < later, "文字列比較で時系列順にならないと範囲検索が壊れる"


# ------------------------------------------------------------------ 起動の配線


async def test_build_servers_uses_two_ports(settings, db, keys, limiter, make_upstream, logs) -> None:
    proxy, control = build_servers(
        settings, db=db, keys=keys, limiter=limiter, upstream=make_upstream(), logs=logs
    )
    assert proxy.config.port == settings.port
    assert control.config.port == settings.control_port
    assert proxy.config.port != control.config.port
    assert proxy.config.timeout_graceful_shutdown == int(settings.graceful_timeout)


async def test_managed_server_does_not_capture_signals(settings, db, keys, limiter, make_upstream, logs) -> None:
    """2 つ並べるため、シグナル処理は呼び出し側が持つ。"""

    import signal

    proxy, _ = build_servers(
        settings, db=db, keys=keys, limiter=limiter, upstream=make_upstream(), logs=logs
    )
    before = signal.getsignal(signal.SIGTERM)
    with proxy.capture_signals():
        assert signal.getsignal(signal.SIGTERM) is before
