"""Public tests for AsyncLLMClient. The transport counts calls and in-flight sends; the clock
is fake, so timings are exact and a REAL sleep anywhere shows up as real seconds."""

from __future__ import annotations

import asyncio
import time

import pytest
from conftest import FakeClock, run

from llmclient import AsyncLLMClient, MockTransport, TokenBucket, TransientError


def _unlimited_bucket(clock: FakeClock) -> TokenBucket:
    return TokenBucket(rate=1000, capacity=1000, clock=clock.now, sleep=clock.sleep)


def test_ask_many_preserves_prompt_order(clock: FakeClock) -> None:
    # replies arrive in REVERSE order (p7 is fastest); the result list must still be in
    # prompt order, whatever order the transport answered in.
    prompts = [f"p{i}" for i in range(8)]
    transport = MockTransport(latency=lambda p: 0.1 * (8 - int(p[1:])), sleep=clock.sleep)
    client = AsyncLLMClient(transport, _unlimited_bucket(clock), max_in_flight=8, sleep=clock.sleep)
    out = run(client.ask_many(prompts))
    assert out == [f"echo:{p}" for p in prompts], out


def test_in_flight_never_exceeds_max_in_flight(clock: FakeClock) -> None:
    # 10 prompts, 3 slots, a bucket that never throttles: the transport must see exactly 3
    # sends in flight at the peak — not 10 (no bound), not 1 (no concurrency).
    transport = MockTransport(latency=0.05, sleep=clock.sleep)
    client = AsyncLLMClient(transport, _unlimited_bucket(clock), max_in_flight=3, sleep=clock.sleep)
    out = run(client.ask_many([f"p{i}" for i in range(10)]))
    assert len(out) == 10
    assert transport.max_in_flight_seen == 3, transport.max_in_flight_seen


def test_retries_a_transient_failure_once_with_backoff(clock: FakeClock) -> None:
    # every 2nd call fails: "a" succeeds (call 1); "b" fails (call 2), is retried after the
    # backoff (call 3) and succeeds. Elapsed = 0.05 (failed send) + 0.25 (backoff) + 0.05.
    transport = MockTransport(latency=0.05, fail_every=2, sleep=clock.sleep)
    client = AsyncLLMClient(
        transport, _unlimited_bucket(clock), max_in_flight=4, backoff=0.25, sleep=clock.sleep
    )

    async def scenario() -> tuple[str, float]:
        assert await client.ask("a") == "echo:a"
        t0 = clock.now()
        reply = await client.ask("b")
        return reply, clock.now() - t0

    reply, elapsed = run(scenario())
    assert reply == "echo:b"
    assert transport.calls == 3, transport.log
    assert abs(elapsed - 0.35) < 1e-9, f"expected 0.35 s of virtual time, got {elapsed}"

    # ...and only once: a transport that ALWAYS fails gets exactly two calls, then the error.
    always = MockTransport(latency=0.05, fail_every=1, sleep=clock.sleep)
    client2 = AsyncLLMClient(always, _unlimited_bucket(clock), max_in_flight=4, sleep=clock.sleep)
    with pytest.raises(TransientError):
        run(client2.ask("c"))
    assert always.calls == 2, always.calls


def test_cancellation_propagates(clock: FakeClock) -> None:
    # cancel an ask while its send is in flight: the task must end CANCELLED — no retry, no
    # second send, no result — because CancelledError was allowed through.
    transport = MockTransport(latency=1.0, sleep=clock.sleep)
    client = AsyncLLMClient(transport, _unlimited_bucket(clock), max_in_flight=2, sleep=clock.sleep)

    async def scenario() -> asyncio.Task[str]:
        task = asyncio.create_task(client.ask("p"))
        await clock.sleep(0.1)  # the send is now in flight
        assert transport.in_flight == 1
        task.cancel()
        with pytest.raises(asyncio.CancelledError):
            await task
        await clock.sleep(2.0)  # long enough for any smuggled retry to have happened
        return task

    task = run(scenario())
    assert task.cancelled(), "the task finished instead of ending cancelled"
    assert transport.calls == 1, f"a cancelled ask must not retry: {transport.log}"
    assert transport.in_flight == 0


def test_ten_prompts_at_rate_five_take_about_two_seconds(clock: FakeClock) -> None:
    # rate 5/s, capacity 1: starts at 0, 0.2, …, 1.8 s, each call 0.05 s → 1.85 s of VIRTUAL
    # time. In REAL time the whole thing must be near-instant: the clock is fake, so a real
    # sleep anywhere (time.sleep, a hand-rolled wait) is a bug that shows up right here.
    bucket = TokenBucket(rate=5, capacity=1, clock=clock.now, sleep=clock.sleep)
    transport = MockTransport(latency=0.05, sleep=clock.sleep)
    client = AsyncLLMClient(transport, bucket, max_in_flight=10, sleep=clock.sleep)
    real0 = time.perf_counter()
    out = run(client.ask_many([f"p{i}" for i in range(10)]))
    real = time.perf_counter() - real0
    assert sorted(out) == sorted(f"echo:p{i}" for i in range(10))  # order is test 3's concern
    assert 1.8 <= clock.now() <= 2.0, f"virtual elapsed {clock.now():.3f} s, expected ≈ 1.85"
    assert real < 0.5, f"{real:.2f} s of REAL time passed — something slept for real"
