"""Public tests for the TokenBucket. Names are the page's check list — they assert BEHAVIOUR
(when starts happened on the virtual clock), so a red row says what to satisfy, not how."""

from __future__ import annotations

import asyncio

from conftest import FakeClock, run

from llmclient import TokenBucket


async def _stamped_acquires(bucket: TokenBucket, clock: FakeClock, count: int) -> list[float]:
    """`count` concurrent acquire() calls; the virtual time each one was admitted at."""
    times: list[float] = []

    async def one() -> None:
        await bucket.acquire()
        times.append(clock.now())

    await asyncio.gather(*(one() for _ in range(count)))
    return sorted(times)


def test_bucket_never_exceeds_rate_in_any_one_second_window(clock: FakeClock) -> None:
    # capacity 1, rate 10/s: once the single token is gone, at most 10 starts per second,
    # measured from every admission — no window of 1 s may hold more than `rate` starts.
    bucket = TokenBucket(rate=10, capacity=1, clock=clock.now, sleep=clock.sleep)
    times = run(_stamped_acquires(bucket, clock, 24))
    assert len(times) == 24
    for i, t0 in enumerate(times):
        in_window = sum(1 for t in times[i:] if t < t0 + 1.0 - 1e-9)
        assert in_window <= 10, f"{in_window} starts in [{t0:.2f}, {t0 + 1:.2f}) — rate is 10/s"
    assert times[-1] >= 2.3, f"24 starts at 10/s need ≥ 2.3 s of virtual time, took {times[-1]:.2f}"


def test_bucket_bursts_to_capacity_then_throttles(clock: FakeClock) -> None:
    # rate 2/s, capacity 3: three starts free at t=0, then one every 0.5 s. After 10 s idle
    # the bucket is FULL again (capacity 3) — not holding 20 tokens of "saved-up" burst.
    bucket = TokenBucket(rate=2, capacity=3, clock=clock.now, sleep=clock.sleep)

    async def scenario() -> tuple[list[float], list[float]]:
        first = await _stamped_acquires(bucket, clock, 5)
        await clock.sleep(10.0)
        second = await _stamped_acquires(bucket, clock, 5)
        return first, second

    first, second = run(scenario())
    assert [round(t, 9) for t in first] == [0.0, 0.0, 0.0, 0.5, 1.0], first
    base = second[0]
    assert abs(base - 11.0) < 1e-9, f"first start after the idle gap should be t=11.0, was {base}"
    assert [round(t - base, 9) for t in second] == [0.0, 0.0, 0.0, 0.5, 1.0], second
