"""ex-vllm-load — the load generator. STARTER: signatures and docstrings are the contract;
fill in the bodies of run_at_concurrency, run_load_test and format_report.

Fires `requests_per_level` requests at an OpenAI-compatible endpoint (POST .../v1/completions
or .../v1/chat/completions), `concurrency` of them in flight at once, for each level in
`concurrency_levels`, and reports p50/p95 latency per level. `transport` is the same
httpx.AsyncClient(transport=...) parameter httpx always takes: left unset, it opens a real
socket at base_url (the 3060 run); in the tests it is an httpx.MockTransport wired to the
fake server in tests/conftest.py — the same code, the same call sequence, either way.
"""

from __future__ import annotations

import time
from collections.abc import Sequence
from dataclasses import dataclass, field

import httpx

# run_at_concurrency needs `import asyncio` (for asyncio.Semaphore / asyncio.gather) and
# `from .stats import percentile` once its body is filled in — add both then.

DEFAULT_PAYLOAD = {
    "model": "local-model",
    "prompt": "The quick brown fox jumps over the lazy dog.",
    "max_tokens": 16,
}


@dataclass
class ConcurrencyResult:
    """One concurrency level's outcome. p50_s/p95_s cover only the SUCCESSFUL requests'
    wall-clock latencies, in seconds; `errors` counts the rest."""

    concurrency: int
    n: int
    errors: int
    p50_s: float
    p95_s: float


@dataclass
class LoadReport:
    results: list[ConcurrencyResult] = field(default_factory=list)


async def _one_request(
    client: httpx.AsyncClient, url: str, payload: dict[str, object]
) -> float | None:
    """POST `payload` to `url`; return its wall-clock latency in seconds, or None if it
    raised or came back non-2xx."""
    t0 = time.perf_counter()
    try:
        resp = await client.post(url, json=payload)
        resp.raise_for_status()
    except httpx.HTTPError:
        return None
    return time.perf_counter() - t0


async def run_at_concurrency(
    client: httpx.AsyncClient,
    url: str,
    payload: dict[str, object],
    concurrency: int,
    n_requests: int,
) -> ConcurrencyResult:
    """Fire `n_requests` requests at `url` through `client`, at most `concurrency` of them in
    flight at once (an asyncio.Semaphore, not one flat gather — n_requests may exceed
    concurrency). Raise RuntimeError if every request fails (nothing to compute a percentile
    over)."""
    ...


async def run_load_test(
    base_url: str,
    *,
    concurrency_levels: Sequence[int],
    requests_per_level: int,
    path: str = "/v1/completions",
    payload: dict[str, object] | None = None,
    transport: httpx.AsyncBaseTransport | None = None,
) -> LoadReport:
    """Run every level of `concurrency_levels`, IN ORDER, against
    `base_url.rstrip("/") + path`, collecting one ConcurrencyResult per level into a
    LoadReport. `payload` defaults to DEFAULT_PAYLOAD; `transport` is passed straight to
    httpx.AsyncClient (see module docstring)."""
    ...


def format_report(report: LoadReport) -> str:
    """Render `report` as a markdown table: a header row naming concurrency / requests /
    errors / p50 (ms) / p95 (ms), a separator row, then one row per ConcurrencyResult in
    `report.results`, in order. Latencies are reported in MILLISECONDS (p50_s/p95_s are
    seconds)."""
    ...
