refactor(harness): simplify benchmark worker coordination

This commit is contained in:
Yujong Lee 2026-09-05 13:07:01 -07:00
parent af1c8eb5ed
commit de28a65e96
6 changed files with 111 additions and 69 deletions

View file

@ -51,7 +51,7 @@ Request variants add PDF comment padding before the final `startxref` marker, pr
Each backend, route, size, and repeat gets a fresh SDK process for timing and another for memory. Python and Rust execute sequentially, with their order reversed on alternating repeats. The local provider serves preloaded bytes without parsing or capturing request JSON. Its CPU and RSS are outside the SDK measurements
Workers warm up their clients and run an untimed response check before measuring. Python and Rust response digests must match. Every provider request also checks the existing parity harness's User-Agent convention: a Rust run using Python's HTTP path fails instead of reporting a comparison between two Python runs. Missing native extensions, SDK exceptions, timeouts, and incomplete samples fail the run
Workers warm up their clients and run an untimed response check before measuring. Python and Rust response digests must match. Every provider request also checks the existing parity harness's User-Agent convention: a Rust run using Python's HTTP path fails instead of reporting a comparison between two Python runs. Missing native extensions, SDK exceptions, timeouts, and incomplete samples fail the run. Workers publish readiness and results atomically in temporary JSON files; stdout and stderr go to a diagnostic log. The controller waits for timing workers to exit and samples RSS only for memory workers. A timeout or interruption terminates and reaps the SDK worker, with bounded shutdown waits for both SDK and provider processes
Latency starts immediately before the SDK call and ends when its result has been returned and discarded. Async calls are awaited on a persistent event loop. CPU is process CPU time during the timed batch, including Python and native threads. Fixture loading, process startup, warmup, preflight serialization, and report generation are excluded. Default SDK behavior is retained, so deferred background work can extend beyond a call's return; these metrics describe the measurement window, not the eventual cost of every callback

View file

@ -5,15 +5,14 @@ import subprocess
import sys
import tempfile
from collections.abc import Generator, Iterator
from concurrent.futures import Future, ThreadPoolExecutor
from contextlib import contextmanager
from contextlib import contextmanager, suppress
from pathlib import Path
from time import monotonic, sleep
from typing import TYPE_CHECKING, Final, TextIO, cast
from typing import TYPE_CHECKING, Final, TextIO
import psutil
from .models import PREFIX, Backend, BenchmarkModel, Invocation, Measurement, Memory, Options, Ready, Route, Timing
from .models import Backend, BenchmarkModel, Invocation, Measurement, Memory, Options, Ready, Route, Timing
from .provider import PYTHON_SENTINEL, provider_process
if TYPE_CHECKING:
@ -30,13 +29,6 @@ def rss_bytes(process: psutil.Process) -> int:
return ProcessMemory.model_validate(process.memory_info(), from_attributes=True).rss
def _read_message(stream: TextIO) -> str:
for line in stream:
if line.startswith(PREFIX):
return line.removeprefix(PREFIX)
raise RuntimeError("SDK worker exited without returning a measurement")
@contextmanager
def sdk_process(case_file: Path, backend: Backend, repo_root: Path, log: TextIO) -> Generator[subprocess.Popen[str]]:
process: Final = subprocess.Popen(
@ -52,31 +44,40 @@ def sdk_process(case_file: Path, backend: Backend, repo_root: Path, log: TextIO)
"PYTHONPATH": str(repo_root),
},
stdin=subprocess.PIPE,
stdout=subprocess.PIPE,
stdout=log,
stderr=log,
text=True,
)
try:
yield process
except BaseException:
process.terminate()
raise
finally:
if process.stdin is not None:
process.stdin.close()
with suppress(BrokenPipeError):
process.stdin.close()
try:
process.wait(timeout=5)
except subprocess.TimeoutExpired:
process.kill()
process.wait(timeout=5)
if process.stdout is not None:
process.stdout.close()
def sample_rss(process: psutil.Process, completed: Future[str], interval: float, timeout: float) -> Iterator[int]:
deadline: Final = monotonic() + timeout
while not completed.done():
def wait_for_output(
output: Path, child: subprocess.Popen[str], options: Options, *, sample_memory: bool = False
) -> Iterator[int]:
deadline: Final = monotonic() + options.timeout
process: Final = psutil.Process(child.pid) if sample_memory else None
interval: Final = options.sample_interval_ms / 1000 if sample_memory else 0.01
while not output.exists():
if child.poll() is not None:
raise RuntimeError(f"SDK worker exited with code {child.returncode} before writing {output.name}")
if monotonic() >= deadline:
raise TimeoutError("memory measurement timed out")
yield rss_bytes(process)
sleep(interval)
raise TimeoutError(f"SDK worker timed out waiting for {output.name}")
if process is not None:
yield rss_bytes(process)
sleep(min(interval, max(0, deadline - monotonic())))
def execute_phase(
@ -86,35 +87,40 @@ def execute_phase(
directory: Final = Path(raw_directory)
case_file: Final = directory / "invocation.json"
case_file.write_text(invocation.model_dump_json())
ready_file: Final = directory / "ready.json"
timing_file: Final = directory / "timing.json"
with (directory / "worker.log").open("w+") as log:
try:
with ThreadPoolExecutor(max_workers=1) as reader:
with sdk_process(case_file, backend, repo_root, log) as child:
assert child.stdout is not None and child.stdin is not None
stdout: Final = cast(TextIO, child.stdout)
ready: Final = Ready.model_validate_json(
reader.submit(_read_message, stdout).result(timeout=options.timeout)
with sdk_process(case_file, backend, repo_root, log) as child:
tuple(wait_for_output(ready_file, child, options))
ready: Final = Ready.model_validate_json(ready_file.read_bytes())
if invocation.phase == "timing":
child.communicate(input="go\n", timeout=options.timeout)
if child.returncode != 0:
raise RuntimeError(f"SDK worker exited with code {child.returncode}")
return (
ready,
Timing.model_validate_json(timing_file.read_bytes()),
Memory(baseline_rss_bytes=0, sampled_peak_rss_bytes=0, retained_rss_bytes=0, samples=0),
)
process: Final = psutil.Process(child.pid)
baseline: Final = rss_bytes(process) if invocation.phase == "memory" else 0
child.stdin.write("go\n")
child.stdin.flush()
result: Final = reader.submit(_read_message, stdout)
samples: Final = (
tuple(sample_rss(process, result, options.sample_interval_ms / 1000, options.timeout))
if invocation.phase == "memory"
else ()
)
timing: Final = Timing.model_validate_json(result.result(timeout=options.timeout))
retained: Final = rss_bytes(process) if invocation.phase == "memory" else 0
memory: Final = Memory(
process: Final = psutil.Process(child.pid)
baseline: Final = rss_bytes(process)
assert child.stdin is not None
child.stdin.write("go\n")
child.stdin.flush()
samples: Final = tuple(wait_for_output(timing_file, child, options, sample_memory=True))
retained: Final = rss_bytes(process)
return (
ready,
Timing.model_validate_json(timing_file.read_bytes()),
Memory(
baseline_rss_bytes=baseline,
sampled_peak_rss_bytes=max((baseline, retained, *samples)),
retained_rss_bytes=retained,
samples=len(samples),
)
return ready, timing, memory
except (RuntimeError, OSError, ValueError, TimeoutError) as error:
),
)
except (RuntimeError, OSError, ValueError, TimeoutError, subprocess.TimeoutExpired, psutil.Error) as error:
log.seek(0)
raise RuntimeError(f"{backend}/{invocation.phase}: {error}\n{log.read()[-6000:]}") from error

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from typing import Final, Literal
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field
@ -8,7 +8,6 @@ Backend = Literal["python", "rust"]
Route = Literal["ocr", "aocr"]
Phase = Literal["timing", "memory"]
Profile = Literal["small", "request_medium", "request_large", "response_medium", "response_large"]
PREFIX: Final = "LITELLM_BENCHMARK "
class BenchmarkModel(BaseModel):

View file

@ -67,5 +67,5 @@ def provider_process(response: bytes, backend: Backend) -> Generator[str]:
process.join(timeout=5)
if process.is_alive():
process.kill()
process.join()
process.join(timeout=5)
process.close()

View file

@ -2,9 +2,10 @@ from __future__ import annotations
import asyncio
import base64
from concurrent.futures import Future
import subprocess
import sys
from pathlib import Path
from time import sleep
from time import monotonic, sleep
from typing import Final
import httpx
@ -16,7 +17,7 @@ from litellm.llms.base_llm.ocr.transformation import OCRResponse
from ...cli.catalog import load_catalog
from ...shared.reporting.models import RunStatus
from .execution import execute_phase, sample_rss
from .execution import execute_phase, sdk_process, wait_for_output
from .models import Invocation, Options
from .provider import PYTHON_SENTINEL, provider_process
from .reporting import percentile, render_measurements
@ -114,10 +115,43 @@ def test_memory_pass_does_not_accumulate_latency_samples() -> None:
assert result.latency_ms == ()
def test_memory_monitor_has_a_deadline() -> None:
pending: Final[Future[str]] = Future()
with pytest.raises(TimeoutError, match="memory measurement timed out"):
tuple(sample_rss(psutil.Process(), pending, interval=0.001, timeout=0.01))
@pytest.mark.parametrize("sample_memory", (False, True))
def test_worker_deadline_terminates_and_reaps_the_process(tmp_path: Path, sample_memory: bool) -> None:
case_file: Final = tmp_path / "invocation.json"
case_file.write_text(invocation().model_dump_json())
existing_children: Final = frozenset(process.pid for process in psutil.Process().children())
start: Final = monotonic()
with (tmp_path / "worker.log").open("w+") as log:
with pytest.raises(TimeoutError, match="timed out waiting for missing.json"):
with sdk_process(case_file, "python", REPO_ROOT, log) as child:
tuple(
wait_for_output(
tmp_path / "missing.json",
child,
Options(timeout=0.02, sample_interval_ms=10000),
sample_memory=sample_memory,
)
)
assert frozenset(process.pid for process in psutil.Process().children()) <= existing_children
assert monotonic() - start < 5
def test_worker_cleanup_handles_buffered_input_after_early_exit(tmp_path: Path) -> None:
case_file: Final = tmp_path / "invocation.json"
case_file.write_text(invocation().model_dump_json())
with (tmp_path / "worker.log").open("w+") as log:
with sdk_process(case_file, "python", REPO_ROOT, log) as child:
child.terminate()
child.wait(timeout=5)
assert child.stdin is not None
child.stdin.write("go\n")
assert child.returncode is not None
def test_worker_exit_is_detected_without_waiting_for_the_deadline(tmp_path: Path) -> None:
with subprocess.Popen((sys.executable, "-c", "raise SystemExit(7)"), cwd=REPO_ROOT, text=True) as child:
with pytest.raises(RuntimeError, match="exited with code 7"):
tuple(wait_for_output(tmp_path / "missing.json", child, Options(timeout=30)))
def test_replay_rejects_python_fallback_during_rust_measurement() -> None:

View file

@ -13,7 +13,7 @@ from typing import Final, cast
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from .models import PREFIX, Invocation, Ready, Timing
from .models import Invocation, Ready, Timing
def _ready(response: OCRResponse) -> Ready:
@ -31,20 +31,22 @@ def _ready(response: OCRResponse) -> Ready:
)
def _emit(value: Ready | Timing) -> None:
print(PREFIX + value.model_dump_json(), flush=True)
def _publish(value: Ready | Timing, destination: Path) -> None:
temporary: Final = destination.with_suffix(".tmp")
temporary.write_text(value.model_dump_json())
temporary.replace(destination)
def _handshake(ready: Ready) -> None:
def _handshake(ready: Ready, directory: Path) -> None:
gc.collect()
_emit(ready)
_publish(ready, directory / "ready.json")
if sys.stdin.readline().strip() != "go":
raise RuntimeError("benchmark controller disconnected before measurement")
def _finish(timing: Timing) -> None:
def _finish(timing: Timing, directory: Path) -> None:
gc.collect()
_emit(timing)
_publish(timing, directory / "timing.json")
sys.stdin.readline()
@ -84,15 +86,15 @@ async def measure_async(call: Callable[[], Awaitable[OCRResponse]], invocation:
return Timing(latency_ms=samples, cpu_ms=(process_time_ns() - cpu_start) / 1e6, elapsed_ms=elapsed / 1e6)
async def _run_async(call: Callable[[], Awaitable[OCRResponse]], invocation: Invocation) -> None:
async def _run_async(call: Callable[[], Awaitable[OCRResponse]], invocation: Invocation, directory: Path) -> None:
for _ in range(invocation.warmup):
await call()
ready: Final = _ready(await call())
_handshake(ready)
_finish(await measure_async(call, invocation))
_handshake(ready, directory)
_finish(await measure_async(call, invocation), directory)
def run_worker(invocation: Invocation) -> None:
def run_worker(invocation: Invocation, directory: Path) -> None:
import litellm
kwargs: Final = {
@ -105,16 +107,17 @@ def run_worker(invocation: Invocation) -> None:
}
if invocation.route == "aocr":
async_route: Final = cast(Callable[..., Awaitable[OCRResponse]], litellm.aocr)
asyncio.run(_run_async(lambda: async_route(**kwargs), invocation))
asyncio.run(_run_async(lambda: async_route(**kwargs), invocation, directory))
return
sync_route: Final = cast(Callable[..., OCRResponse], litellm.ocr)
call: Final = lambda: sync_route(**kwargs)
for _ in range(invocation.warmup):
call()
ready: Final = _ready(call())
_handshake(ready)
_finish(measure_sync(call, invocation))
_handshake(ready, directory)
_finish(measure_sync(call, invocation), directory)
if __name__ == "__main__":
run_worker(Invocation.model_validate_json(Path(sys.argv[1]).read_bytes()))
case_file: Final = Path(sys.argv[1])
run_worker(Invocation.model_validate_json(case_file.read_bytes()), case_file.parent)