fix(harness): handle benchmark CLI validation and output errors

This commit is contained in:
Yujong Lee 2026-09-05 13:14:53 -07:00
parent de28a65e96
commit e9dbe047a2
6 changed files with 105 additions and 37 deletions

View file

@ -0,0 +1,3 @@
from typing import Final
PYTHON_SENTINEL: Final = "litellm-benchmark-python"

View file

@ -12,8 +12,9 @@ from typing import TYPE_CHECKING, Final, TextIO
import psutil
from .constants import PYTHON_SENTINEL
from .models import Backend, BenchmarkModel, Invocation, Measurement, Memory, Options, Ready, Route, Timing
from .provider import PYTHON_SENTINEL, provider_process
from .provider import provider_process
if TYPE_CHECKING:
from .workloads import Workload

View file

@ -20,8 +20,8 @@ class Options(BenchmarkModel):
repeats: int = Field(default=3, ge=1)
profiles: tuple[Profile, ...] = ("small", "request_medium", "request_large", "response_medium", "response_large")
routes: tuple[Route, ...] = ("ocr", "aocr")
timeout: float = Field(default=120, gt=0)
sample_interval_ms: float = Field(default=5, ge=1)
timeout: float = Field(default=120, gt=0, allow_inf_nan=False)
sample_interval_ms: float = Field(default=5, ge=1, allow_inf_nan=False)
output: str | None = None

View file

@ -7,10 +7,9 @@ from multiprocessing.connection import Connection
from typing import ClassVar, Final
from ...shared.parity.local_server import LocalHttpHandler, LocalHttpServer
from .constants import PYTHON_SENTINEL
from .models import Backend
PYTHON_SENTINEL: Final = "litellm-benchmark-python"
class Provider(LocalHttpServer):
def __init__(self, response: bytes, backend: Backend) -> None:

View file

@ -1,6 +1,5 @@
from __future__ import annotations
import argparse
import platform
import subprocess
from collections.abc import Sequence
@ -8,6 +7,9 @@ from pathlib import Path
from time import monotonic
from typing import TYPE_CHECKING, Final
import click
from pydantic import ValidationError
from ...shared.reporting.models import CaseResult, HarnessCase, HarnessRun, ResultArtifact, RunStatus
from ...shared.reporting.strategy import ModuleCaseSpec, UpdateCallback
from .execution import benchmark
@ -29,19 +31,32 @@ class Report(BenchmarkModel):
def parse_options(arguments: Sequence[str]) -> Options:
parser: Final = argparse.ArgumentParser(prog="e2e_benchmark", exit_on_error=False)
parser.add_argument("--iterations", type=int, default=100)
parser.add_argument("--warmup", type=int, default=10)
parser.add_argument("--repeats", type=int, default=3)
parser.add_argument("--profile", dest="profiles", action="append", default=argparse.SUPPRESS)
parser.add_argument("--route", dest="routes", action="append", default=argparse.SUPPRESS)
parser.add_argument("--timeout", type=float, default=120)
parser.add_argument("--sample-interval-ms", type=float, default=5)
parser.add_argument("--output")
parsed, unknown = parser.parse_known_args(arguments)
if unknown:
raise ValueError(f"unknown benchmark arguments: {' '.join(unknown)}")
return Options.model_validate(vars(parsed))
defaults: Final = Options()
command: Final = click.Command(
"e2e_benchmark",
params=[
click.Option(("--iterations",), type=click.IntRange(min=1), default=defaults.iterations),
click.Option(("--warmup",), type=click.IntRange(min=1), default=defaults.warmup),
click.Option(("--repeats",), type=click.IntRange(min=1), default=defaults.repeats),
click.Option(
("--profile", "profiles"),
type=click.Choice(defaults.profiles),
multiple=True,
default=defaults.profiles,
),
click.Option(
("--route", "routes"), type=click.Choice(defaults.routes), multiple=True, default=defaults.routes
),
click.Option(("--timeout",), type=click.FloatRange(min=0, min_open=True), default=defaults.timeout),
click.Option(("--sample-interval-ms",), type=click.FloatRange(min=1), default=defaults.sample_interval_ms),
click.Option(("--output",)),
],
)
with command.make_context("e2e_benchmark", list(arguments)) as context:
try:
return Options.model_validate(context.params)
except ValidationError as error:
raise click.UsageError(str(error)) from error
def _run_pair(
@ -128,6 +143,11 @@ def run_benchmark_cases(
measurements=measurements(tuple(run.results.values())),
failures=tuple(run.failures),
)
Path(options.output).write_text(report.model_dump_json(indent=2) + "\n")
try:
Path(options.output).write_text(report.model_dump_json(indent=2) + "\n")
except OSError as error:
raise click.ClickException(
f"cannot write benchmark report to {options.output}: {error.strerror}"
) from error
on_update(run)
return int(bool(run.failures)), run

View file

@ -8,20 +8,20 @@ from pathlib import Path
from time import monotonic, sleep
from typing import Final
import click
import httpx
import psutil
import pytest
from pydantic import ValidationError
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from ...cli.catalog import load_catalog
from ...shared.reporting.models import RunStatus
from ...cli import main
from .execution import execute_phase, sdk_process, wait_for_output
from .constants import PYTHON_SENTINEL
from .models import Invocation, Options
from .provider import PYTHON_SENTINEL, provider_process
from .provider import provider_process
from .reporting import percentile, render_measurements
from .runner import Report, parse_options, run_benchmark_cases
from .runner import Report, parse_options
from .worker import measure_async, measure_sync
from .workloads import JSON_OBJECT, JSON_PAGES, ocr_workload, padded_pdf
@ -71,12 +71,12 @@ def test_pdf_padding_preserves_existing_offsets_and_exact_size() -> None:
@pytest.mark.parametrize("arguments", (("--iterations=0",), ("--warmup=0",), ("--route=chat",), ("--profile=unknown",)))
def test_invalid_benchmark_options_fail_before_running(arguments: tuple[str, ...]) -> None:
with pytest.raises(ValidationError):
with pytest.raises(click.BadParameter):
parse_options(arguments)
def test_unknown_options_are_not_silently_ignored() -> None:
with pytest.raises(ValueError, match="unknown benchmark arguments"):
with pytest.raises(click.NoSuchOption, match="No such option"):
parse_options(("--concurrency=8",))
@ -170,18 +170,27 @@ def test_worker_errors_are_reported_instead_of_counted_as_fast_calls() -> None:
execute_phase(request, "python", Options(iterations=3, warmup=1), REPO_ROOT)
def test_strategy_runs_both_backends_and_exports_measurements(tmp_path: Path) -> None:
strategy: Final = next(strategy for strategy in load_catalog() if strategy.id == "e2e_benchmark")
case: Final = next(case for case in strategy.cases if case.sdk_function == "ocr")
def test_cli_runs_both_backends_and_exports_measurements(tmp_path: Path, capsys: pytest.CaptureFixture[str]) -> None:
output: Final = tmp_path / "measurements.json"
exit_code, run = run_benchmark_cases(
(case,),
REPO_ROOT,
lambda _: None,
("--profile=small", "--route=aocr", "--iterations=3", "--warmup=1", "--repeats=1", f"--output={output}"),
exit_code: Final = main(
(
"run",
"e2e_benchmark",
"--surface",
"sdk",
"--function",
"ocr",
"--benchmark-arg=--profile=small",
"--benchmark-arg=--route=aocr",
"--benchmark-arg=--iterations=3",
"--benchmark-arg=--warmup=1",
"--benchmark-arg=--repeats=1",
f"--benchmark-arg=--output={output}",
)
)
assert exit_code == 0, run.failures
assert run.results[case.key].status is RunStatus.PASSED
captured: Final = capsys.readouterr()
assert exit_code == 0, captured.out + captured.err
assert "Result: PASSED" in captured.out
report: Final = Report.model_validate_json(output.read_bytes())
assert {value.backend for value in report.measurements} == {"python", "rust"}
assert len({value.ready.response_digest for value in report.measurements}) == 1
@ -197,3 +206,39 @@ def test_strategy_runs_both_backends_and_exports_measurements(tmp_path: Path) ->
assert "aocr/small | python" in table
assert "aocr/small | rust" in table
assert "CPU ms/call" in table
@pytest.mark.parametrize(
"argument", ("--iterations=0", "--warmup=invalid", "--route=chat", "--unknown=1", "--timeout=nan", "--timeout=inf")
)
def test_cli_rejects_invalid_benchmark_options_without_a_traceback(
argument: str, capsys: pytest.CaptureFixture[str]
) -> None:
assert main(("run", "e2e_benchmark", "--function", "ocr", f"--benchmark-arg={argument}")) == 2
captured: Final = capsys.readouterr()
assert "Error:" in captured.err
assert "Traceback" not in captured.err
assert "sdk/ocr: running" not in captured.out
@pytest.mark.parametrize("destination", ("missing/report.json", "."))
def test_cli_reports_output_errors_without_a_traceback(
tmp_path: Path, destination: str, capsys: pytest.CaptureFixture[str]
) -> None:
output: Final = tmp_path / destination
assert main(("run", "e2e_benchmark", "--function", "chat_completions", f"--benchmark-arg=--output={output}")) == 1
captured: Final = capsys.readouterr()
assert "Error: cannot write benchmark report" in captured.err
assert str(output) in captured.err
assert "Traceback" not in captured.err
def test_cli_reports_unsupported_functions_without_measurements(
tmp_path: Path, capsys: pytest.CaptureFixture[str]
) -> None:
output: Final = tmp_path / "unsupported.json"
assert main(("run", "e2e_benchmark", "--function", "chat_completions", f"--benchmark-arg=--output={output}")) == 0
captured: Final = capsys.readouterr()
assert "not implemented" in captured.out
report: Final = Report.model_validate_json(output.read_bytes())
assert report.measurements == ()