"""The transport is PROVIDED — do not edit it. It stands in for the HTTP layer of a real
LLM client: every `send` costs `latency` seconds of awaiting, every `fail_every`-th call raises
a `TransientError` (a 429, a dropped connection), and it keeps the counters the tests read:
how many calls it saw and how many were in flight at once.

`sleep` is injectable so the tests can run it on a fake clock; production code leaves the
default (`asyncio.sleep`).
"""

from __future__ import annotations

import asyncio
from collections.abc import Awaitable, Callable

SleepFn = Callable[[float], Awaitable[None]]


class TransientError(Exception):
    """A retryable failure — the kind a client is allowed to try once more."""


class MockTransport:
    def __init__(
        self,
        latency: float | Callable[[str], float] = 0.05,
        fail_every: int = 0,
        *,
        sleep: SleepFn = asyncio.sleep,
    ) -> None:
        self.latency = latency
        self.fail_every = fail_every
        self._sleep = sleep
        self.calls = 0  # every send, successful or not
        self.in_flight = 0  # sends currently awaiting their reply
        self.max_in_flight_seen = 0  # the high-water mark of in_flight
        self.log: list[str] = []  # prompts in the order they were SENT

    async def send(self, prompt: str) -> str:
        """Reply "echo:<prompt>" after `latency` seconds; raise TransientError on every
        `fail_every`-th call (counted across all prompts, retries included)."""
        self.calls += 1
        n = self.calls
        self.log.append(prompt)
        self.in_flight += 1
        self.max_in_flight_seen = max(self.max_in_flight_seen, self.in_flight)
        try:
            lat = self.latency(prompt) if callable(self.latency) else self.latency
            await self._sleep(lat)
            if self.fail_every and n % self.fail_every == 0:
                raise TransientError(f"429 on call {n}")
            return f"echo:{prompt}"
        finally:
            self.in_flight -= 1
