"""レート制限。スライディング ウィンドウの境界と同時実行の解放。"""

from __future__ import annotations

import asyncio
import dataclasses

import pytest

from app import errors
from app import limits as limits_module
from app.keys import KeyLimits, KeyRecord
from app.limits import RateLimiter

from .conftest import auth, create_key


class FakeTime:
    """``app.limits`` から見える ``time`` を差し替えるための影。"""

    def __init__(self, start: float = 1000.0) -> None:
        self.value = start

    def monotonic(self) -> float:
        return self.value

    def advance(self, delta: float) -> None:
        self.value += delta


@pytest.fixture
def clock(monkeypatch) -> FakeTime:
    fake = FakeTime()
    monkeypatch.setattr(limits_module, "time", fake)
    return fake


def make_key(**limits) -> KeyRecord:
    return KeyRecord(
        id=1,
        name="k",
        key_hash="h",
        key_prefix="sk-aig-abcd",
        enabled=True,
        created_at="2026-09-09T10:00:00+09:00",
        expires_at=None,
        note="",
        limits=KeyLimits(**limits),
        models=frozenset(),
    )


# ------------------------------------------------------ スライディング ウィンドウ


def test_sliding_window_blocks_after_limit(settings, clock) -> None:
    limiter = RateLimiter(settings)
    key = make_key(rpm=3)

    for _ in range(3):
        limiter.check_rates(key)

    with pytest.raises(errors.GatewayError) as excinfo:
        limiter.check_rates(key)
    assert excinfo.value.status_code == 429
    assert excinfo.value.headers["Retry-After"] == "60"


def test_sliding_window_boundary(settings, clock) -> None:
    """最古の記録が 60 秒経過するまでは開放しない。"""

    limiter = RateLimiter(settings)
    key = make_key(rpm=1)
    limiter.check_rates(key)  # t = 1000.0

    clock.advance(59.999)
    with pytest.raises(errors.GatewayError):
        limiter.check_rates(key)

    clock.advance(0.001)  # t = 1060.0 ちょうど
    limiter.check_rates(key)


def test_sliding_window_is_not_a_fixed_window(settings, clock) -> None:
    """固定ウィンドウなら境界で 2 倍通るが、スライディングでは通らない。"""

    limiter = RateLimiter(settings)
    key = make_key(rpm=2)
    limiter.check_rates(key)
    clock.advance(59.0)
    limiter.check_rates(key)
    clock.advance(0.5)
    with pytest.raises(errors.GatewayError):
        limiter.check_rates(key)


def test_retry_after_shrinks_as_window_slides(settings, clock) -> None:
    limiter = RateLimiter(settings)
    key = make_key(rpm=1)
    limiter.check_rates(key)
    clock.advance(30.0)
    with pytest.raises(errors.GatewayError) as excinfo:
        limiter.check_rates(key)
    assert excinfo.value.headers["Retry-After"] == "30"


def test_rejected_request_is_not_counted(settings, clock) -> None:
    """弾いたリクエストを計上すると、待っても永久に開放されない。"""

    limiter = RateLimiter(settings)
    key = make_key(rpm=1)
    limiter.check_rates(key)
    for _ in range(5):
        with pytest.raises(errors.GatewayError):
            limiter.check_rates(key)
    clock.advance(60.0)
    limiter.check_rates(key)


def test_daily_limit(settings, clock) -> None:
    limiter = RateLimiter(settings)
    key = make_key(rpd=2)
    limiter.check_rates(key)
    limiter.check_rates(key)
    with pytest.raises(errors.GatewayError) as excinfo:
        limiter.check_rates(key)
    assert excinfo.value.code == "daily_limit_exceeded"
    assert 1 <= int(excinfo.value.headers["Retry-After"]) <= 86400


def test_daily_counter_resets_on_new_jst_day(settings, clock) -> None:
    limiter = RateLimiter(settings)
    key = make_key(rpd=1)
    limiter.check_rates(key)
    limiter.state(key.id).day_date = "2000-01-01"
    limiter.check_rates(key)


def test_token_rate_limit(settings, clock) -> None:
    limiter = RateLimiter(settings)
    key = make_key(tpm=100)
    limiter.check_rates(key)
    limiter.record_tokens(key.id, 150)
    with pytest.raises(errors.GatewayError) as excinfo:
        limiter.check_rates(key)
    assert excinfo.value.code == "token_rate_limit_exceeded"
    clock.advance(60.0)
    limiter.check_rates(key)


def test_no_limits_means_unlimited(settings, clock) -> None:
    limiter = RateLimiter(settings)
    key = make_key()
    for _ in range(500):
        limiter.check_rates(key)


# ------------------------------------------------------------------ 同時実行


async def test_key_concurrency_rejects_immediately(settings) -> None:
    limiter = RateLimiter(settings)
    key = make_key(concurrency=1)

    first = await limiter.acquire(key)
    with pytest.raises(errors.GatewayError) as excinfo:
        await limiter.acquire(key)
    assert excinfo.value.code == "concurrency_limit_exceeded"

    first.release()
    second = await limiter.acquire(key)
    second.release()


async def test_global_limit_is_acquired_before_key_limit(settings) -> None:
    """横断上限が先。1 本のキーで上流のスロットを占有させない。"""

    narrow = dataclasses.replace(settings, global_concurrency=1, global_queue_timeout=0.0)
    limiter = RateLimiter(narrow)
    key_a = make_key(concurrency=10)
    key_b = dataclasses.replace(key_a, id=2)

    held = await limiter.acquire(key_a)
    with pytest.raises(errors.GatewayError) as excinfo:
        await limiter.acquire(key_b)
    assert excinfo.value.code == "gateway_busy"
    held.release()
    revived = await limiter.acquire(key_b)
    revived.release()


async def test_global_limit_waits_when_queue_timeout_is_set(settings) -> None:
    waiting = dataclasses.replace(settings, global_concurrency=1, global_queue_timeout=5.0)
    limiter = RateLimiter(waiting)
    key = make_key()

    held = await limiter.acquire(key)
    task = asyncio.create_task(limiter.acquire(key))
    await asyncio.sleep(0)
    assert not task.done(), "横断上限に空きが無いのに待たずに通っています"

    held.release()
    reservation = await asyncio.wait_for(task, timeout=1.0)
    reservation.release()


async def test_global_limit_times_out_into_429(settings) -> None:
    waiting = dataclasses.replace(settings, global_concurrency=1, global_queue_timeout=0.05)
    limiter = RateLimiter(waiting)
    key = make_key()

    held = await limiter.acquire(key)
    with pytest.raises(errors.GatewayError) as excinfo:
        await limiter.acquire(key)
    assert excinfo.value.status_code == 429
    held.release()


async def test_slots_are_released_when_a_later_check_fails(settings) -> None:
    """毎分上限で弾かれた場合も、取得済みの同時実行スロットを必ず返す。"""

    limiter = RateLimiter(dataclasses.replace(settings, global_concurrency=2))
    key = make_key(rpm=1, concurrency=1)

    first = await limiter.acquire(key)
    first.release()
    for _ in range(3):
        with pytest.raises(errors.GatewayError):
            await limiter.acquire(key)
    # 解放漏れがあれば、ここで concurrency 側が枯渇して別のコードになる。
    with pytest.raises(errors.GatewayError) as excinfo:
        await limiter.acquire(key)
    assert excinfo.value.code == "rate_limit_exceeded"


def test_release_is_idempotent(settings) -> None:
    limiter = RateLimiter(settings)
    reservation = limits_module.Reservation(limiter)
    reservation.release()
    reservation.release()


async def test_concurrency_capacity_change_does_not_break_release(settings) -> None:
    """上限変更後も、取得時の Semaphore へ返すため解放が破綻しない。"""

    limiter = RateLimiter(settings)
    key = make_key(concurrency=1)
    held = await limiter.acquire(key)

    widened = dataclasses.replace(key, limits=KeyLimits(concurrency=3))
    other = await limiter.acquire(widened)

    held.release()
    other.release()
    assert limiter.state(key.id).semaphore_capacity == 3


# ------------------------------------------------------------ スナップショット


async def test_snapshot_and_restore(settings, db, clock) -> None:
    limiter = RateLimiter(settings)
    key = make_key(rpm=5, rpd=10)
    limiter.check_rates(key)
    limiter.check_rates(key)
    await limiter.snapshot(db)

    revived = RateLimiter(settings)
    await revived.restore(db)
    state = revived.state(key.id)
    assert len(state.minute_hits) == 2
    assert state.day_count == 2

    # 復元した分は上限に算入される。
    revived.check_rates(key)
    revived.check_rates(key)
    revived.check_rates(key)
    with pytest.raises(errors.GatewayError):
        revived.check_rates(key)


async def test_stale_minute_snapshot_is_discarded(settings, db, clock) -> None:
    from datetime import timedelta

    from app.db import iso_jst, now_jst

    await db.execute(
        'INSERT INTO rate_snapshots (key_id, "window", count, updated_at) VALUES (?, ?, ?, ?)',
        (7, "minute", 99, iso_jst(now_jst() - timedelta(minutes=5))),
    )
    limiter = RateLimiter(settings)
    await limiter.restore(db)
    assert len(limiter.state(7).minute_hits) == 0


async def test_stale_day_snapshot_is_discarded(settings, db) -> None:
    await db.execute(
        'INSERT INTO rate_snapshots (key_id, "window", count, updated_at) VALUES (?, ?, ?, ?)',
        (7, "day:2000-01-01", 99, "2000-01-01T00:00:00+09:00"),
    )
    limiter = RateLimiter(settings)
    await limiter.restore(db)
    assert limiter.state(7).day_count == 0


# ------------------------------------------------------------------ HTTP 経由


async def test_http_429_shape(db, keys, settings, build_proxy):
    secret = await create_key(db, keys, settings, rpm=1)
    client, _ = build_proxy()
    async with client:
        first = await client.post(
            "/v1/chat/completions", headers=auth(secret), content=b'{"model":"m"}'
        )
        second = await client.post(
            "/v1/chat/completions", headers=auth(secret), content=b'{"model":"m"}'
        )

    assert first.status_code == 200
    assert second.status_code == 429
    assert second.headers["retry-after"]
    payload = second.json()
    assert payload["error"]["type"] == "rate_limit_error"
    assert payload["error"]["code"] == "rate_limit_exceeded"


async def test_sequential_requests_do_not_leak_concurrency(db, keys, settings, build_proxy):
    """クライアント切断や失敗でスロットが漏れないこと。"""

    secret = await create_key(db, keys, settings, concurrency=1)
    client, _ = build_proxy()
    async with client:
        for _ in range(5):
            response = await client.post(
                "/v1/chat/completions", headers=auth(secret), content=b'{"model":"m"}'
            )
            assert response.status_code == 200
