from __future__ import annotations

import argparse
import http.client
import json
import platform
import socket
import statistics
import threading
import time
from datetime import UTC, datetime
from pathlib import Path
from typing import Literal


FIRST_FRAME_DELAY_SECONDS = 0.02
TERMINAL_FRAME_HOLD_SECONDS = 0.18
READ_SIZE = 8192
TRIALS = 20

FIRST_FRAME = (
    b'data: {"id":"c1","choices":[{"index":0,"delta":{"content":"A"}}]}\n\n'
)
TERMINAL_FRAME = (
    b'data: {"id":"c1","choices":[{"index":0,"delta":{},'
    b'"finish_reason":"stop"}]}\n\ndata: [DONE]\n\n'
)


def _send_chunk(connection: socket.socket, payload: bytes) -> None:
    connection.sendall(f"{len(payload):X}\r\n".encode("ascii") + payload + b"\r\n")


class DelayedChunkServer:
    def __init__(self) -> None:
        self._listener = socket.create_server(("127.0.0.1", 0))
        self._listener.settimeout(5)
        self.port = int(self._listener.getsockname()[1])
        self.error: BaseException | None = None
        self._thread = threading.Thread(target=self._serve, daemon=True)
        self._thread.start()

    def _serve(self) -> None:
        try:
            connection, _ = self._listener.accept()
            with connection:
                connection.settimeout(5)
                request = bytearray()
                while b"\r\n\r\n" not in request:
                    block = connection.recv(4096)
                    if not block:
                        raise RuntimeError("client closed before sending request headers")
                    request.extend(block)
                connection.sendall(
                    b"HTTP/1.1 200 OK\r\n"
                    b"Content-Type: text/event-stream\r\n"
                    b"Transfer-Encoding: chunked\r\n"
                    b"Connection: close\r\n\r\n"
                )
                time.sleep(FIRST_FRAME_DELAY_SECONDS)
                _send_chunk(connection, FIRST_FRAME)
                time.sleep(TERMINAL_FRAME_HOLD_SECONDS)
                _send_chunk(connection, TERMINAL_FRAME)
                connection.sendall(b"0\r\n\r\n")
        except BaseException as exc:
            self.error = exc
        finally:
            self._listener.close()

    def finish(self) -> None:
        self._thread.join(timeout=5)
        if self._thread.is_alive():
            raise RuntimeError("delayed chunk server did not stop")
        if self.error is not None:
            raise RuntimeError("delayed chunk server failed") from self.error


def measure_first_return(mode: Literal["read1", "read"]) -> float:
    server = DelayedChunkServer()
    connection = http.client.HTTPConnection("127.0.0.1", server.port, timeout=5)
    try:
        connection.request("GET", "/stream")
        response = connection.getresponse()
        started = time.perf_counter()
        if mode == "read1":
            first = response.read1(READ_SIZE)
        else:
            first = response.read(READ_SIZE)
        elapsed_ms = (time.perf_counter() - started) * 1000
        if b'"content":"A"' not in first:
            raise RuntimeError(f"{mode} did not return the first visible-content frame")
        response.read()
        return elapsed_ms
    finally:
        connection.close()
        server.finish()


def percentile(values: list[float], proportion: float) -> float:
    ordered = sorted(values)
    position = (len(ordered) - 1) * proportion
    lower = int(position)
    upper = min(lower + 1, len(ordered) - 1)
    fraction = position - lower
    return ordered[lower] * (1 - fraction) + ordered[upper] * fraction


def summarize(values: list[float]) -> dict[str, object]:
    return {
        "samples_ms": [round(value, 3) for value in values],
        "median_ms": round(statistics.median(values), 3),
        "p25_ms": round(percentile(values, 0.25), 3),
        "p75_ms": round(percentile(values, 0.75), 3),
        "min_ms": round(min(values), 3),
        "max_ms": round(max(values), 3),
    }


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument(
        "--output",
        type=Path,
        default=Path(__file__).with_name("client-buffering-benchmark.json"),
    )
    args = parser.parse_args()

    samples: dict[str, list[float]] = {"read1": [], "read": []}
    for _ in range(TRIALS):
        for mode in ("read1", "read"):
            samples[mode].append(measure_first_return(mode))

    read1_summary = summarize(samples["read1"])
    read_summary = summarize(samples["read"])
    read1_median = float(read1_summary["median_ms"])
    read_median = float(read_summary["median_ms"])
    payload = {
        "schema_version": "1",
        "captured_at": datetime.now(UTC).isoformat(),
        "environment": {
            "python": platform.python_version(),
            "platform": platform.platform(),
            "transport": "HTTP/1.1 chunked response over 127.0.0.1",
        },
        "protocol": {
            "trials_per_mode": TRIALS,
            "read_size_bytes": READ_SIZE,
            "first_frame_delay_ms": FIRST_FRAME_DELAY_SECONDS * 1000,
            "terminal_frame_hold_ms": TERMINAL_FRAME_HOLD_SECONDS * 1000,
            "first_frame_bytes": len(FIRST_FRAME),
            "terminal_frame_bytes": len(TERMINAL_FRAME),
        },
        "read1": read1_summary,
        "read": read_summary,
        "derived": {
            "median_added_observation_delay_ms": round(read_median - read1_median, 3),
            "median_delay_ratio": round(read_median / read1_median, 3),
        },
        "scope": (
            "Controlled client-transport benchmark. These values do not measure "
            "provider, model, queue, prefill, decode, or GPU performance."
        ),
    }
    args.output.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
    print(json.dumps(payload["derived"], indent=2))


if __name__ == "__main__":
    main()
