From ea7113dc20fa116f39122a3a70f91ddf237b6144 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 29 Aug 2026 14:44:19 -0700 Subject: [PATCH] make it faster --- tests/test_litellm/ocr/test_sdk_parity.py | 122 ++++++++++----- tests/test_litellm/parity/models.py | 28 +++- tests/test_litellm/parity/replay.py | 46 ++++-- tests/test_litellm/parity/runner.py | 177 ++++++++++++++++------ 4 files changed, 270 insertions(+), 103 deletions(-) diff --git a/tests/test_litellm/ocr/test_sdk_parity.py b/tests/test_litellm/ocr/test_sdk_parity.py index 0a463232afb..9b855e08bef 100644 --- a/tests/test_litellm/ocr/test_sdk_parity.py +++ b/tests/test_litellm/ocr/test_sdk_parity.py @@ -2,7 +2,8 @@ from __future__ import annotations import asyncio import sys -from collections.abc import Callable, Coroutine +import traceback +from collections.abc import Callable, Coroutine, Generator from enum import Enum from pathlib import Path from typing import Final, cast @@ -12,9 +13,14 @@ import pytest from litellm.llms.base_llm.ocr.transformation import OCRResponse from tests.test_litellm.ocr.fixture_models import MistralOcrParityInput, OcrParityCase from tests.test_litellm.parity.compare import assert_parity -from tests.test_litellm.parity.models import SDKReport -from tests.test_litellm.parity.replay import replay_response -from tests.test_litellm.parity.runner import PythonScriptRunner, run_execution +from tests.test_litellm.parity.models import SDKCommand, SDKReport, WorkerFailure, WorkerResult, WorkerSuccess +from tests.test_litellm.parity.runner import ( + WORKER_RESULT_PREFIX, + PythonScriptRunner, + PythonScriptWorker, + execution_worker, + run_execution, +) API_KEY: Final = "test-key" PYTHON_HTTP_SENTINEL: Final = "python-ocr-parity-fallback" @@ -34,7 +40,12 @@ def _call_kwargs(sdk_input: MistralOcrParityInput, mock_url: str, route: SDKRout } -def _execute_sdk_case(sdk_input: MistralOcrParityInput, route: SDKRoute, mock_url: str) -> SDKReport: +def _execute_sdk_case( + sdk_input: MistralOcrParityInput, + route: SDKRoute, + mock_url: str, + event_loop: asyncio.AbstractEventLoop, +) -> SDKReport: import litellm call_kwargs: Final = _call_kwargs(sdk_input, mock_url, route) @@ -43,56 +54,83 @@ def _execute_sdk_case(sdk_input: MistralOcrParityInput, route: SDKRoute, mock_ur response: Final = sync_route(**call_kwargs) return SDKReport(response=response) async_route: Final = cast(Callable[..., Coroutine[object, object, OCRResponse]], litellm.aocr) - async_response: Final = asyncio.run(async_route(**call_kwargs)) + async_response: Final = event_loop.run_until_complete(async_route(**call_kwargs)) return SDKReport(response=async_response) -@pytest.mark.parametrize("route", tuple(SDKRoute), ids=tuple(route.value for route in SDKRoute)) -def test_recorded_ocr_sdk_parity(ocr_fixture: OcrParityCase, route: SDKRoute, tmp_path: Path) -> None: - case_file: Final = tmp_path / f"{route.value}-ocr-parity-case.json" - case_file.write_text(ocr_fixture.model_dump_json(indent=2, exclude_unset=True), encoding="utf-8") - response: Final = ocr_fixture.provider_response - response_body: Final = response.body_bytes() - response_headers: Final = tuple((header.name, header.value) for header in response.headers) +@pytest.fixture(scope="module") +def sdk_workers() -> Generator[tuple[PythonScriptWorker, PythonScriptWorker]]: runner: Final = PythonScriptRunner( entrypoint=Path(__file__), rust_env_var="LITELLM_USE_RUST_OCR", python_user_agent=PYTHON_HTTP_SENTINEL, ) + with execution_worker(runner, rust_enabled=False) as python_worker: + with execution_worker(runner, rust_enabled=True) as rust_worker: + yield python_worker, rust_worker - with replay_response(response.status_code, response_headers, response_body) as python_provider: - python: Final = run_execution( - runner, - case_file, - route.value, - tmp_path / f"{route.value}-python-report.json", - python_provider, - rust_enabled=False, - ) - with replay_response(response.status_code, response_headers, response_body) as rust_provider: - rust: Final = run_execution( - runner, - case_file, - route.value, - tmp_path / f"{route.value}-rust-report.json", - rust_provider, - rust_enabled=True, - ) + +@pytest.mark.parametrize("route", tuple(SDKRoute), ids=tuple(route.value for route in SDKRoute)) +def test_recorded_ocr_sdk_parity( + ocr_fixture: OcrParityCase, + route: SDKRoute, + tmp_path: Path, + sdk_workers: tuple[PythonScriptWorker, PythonScriptWorker], +) -> None: + case_file: Final = tmp_path / f"{route.value}-ocr-parity-case.json" + case_file.write_text(ocr_fixture.model_dump_json(indent=2, exclude_unset=True), encoding="utf-8") + response: Final = ocr_fixture.provider_response + response_body: Final = response.body_bytes() + response_headers: Final = tuple((header.name, header.value) for header in response.headers) + python_worker, rust_worker = sdk_workers + python: Final = run_execution( + python_worker, + case_file, + route.value, + response.status_code, + response_headers, + response_body, + ) + rust: Final = run_execution( + rust_worker, + case_file, + route.value, + response.status_code, + response_headers, + response_body, + ) assert_parity(python, rust, PYTHON_HTTP_SENTINEL) -def _child_main() -> None: - if len(sys.argv) != 5: - raise SystemExit("usage: test_sdk_parity.py CASE_FILE ROUTE MOCK_URL REPORT_FILE") - case_file: Final = Path(sys.argv[1]) - route: Final = SDKRoute(sys.argv[2]) - mock_url: Final = sys.argv[3] - report_file: Final = Path(sys.argv[4]) - case: Final = OcrParityCase.model_validate_json(case_file.read_text(encoding="utf-8")) - report: Final = _execute_sdk_case(case.litellm_input, route, mock_url) - report_file.write_text(report.model_dump_json(indent=2), encoding="utf-8") +def _execute_worker_command( + command_json: str, + mock_url: str, + event_loop: asyncio.AbstractEventLoop, +) -> WorkerResult: + try: + command: Final = SDKCommand.model_validate_json(command_json) + case_file: Final = Path(command.case_file) + route: Final = SDKRoute(command.route) + case: Final = OcrParityCase.model_validate_json(case_file.read_text(encoding="utf-8")) + return WorkerSuccess(report=_execute_sdk_case(case.litellm_input, route, mock_url, event_loop)) + except Exception: + return WorkerFailure(error=traceback.format_exc()) + + +def _worker_main(mock_url: str) -> None: + event_loop: Final = asyncio.new_event_loop() + try: + for line in sys.stdin: + sys.stdout.write( + f"{WORKER_RESULT_PREFIX}{_execute_worker_command(line, mock_url, event_loop).model_dump_json()}\n" + ) + sys.stdout.flush() + finally: + event_loop.close() if __name__ == "__main__": - _child_main() + if len(sys.argv) != 3 or sys.argv[1] != "--parity-worker": + raise SystemExit("usage: test_sdk_parity.py --parity-worker MOCK_URL") + _worker_main(sys.argv[2]) diff --git a/tests/test_litellm/parity/models.py b/tests/test_litellm/parity/models.py index 9cf26127765..5747110b7e0 100644 --- a/tests/test_litellm/parity/models.py +++ b/tests/test_litellm/parity/models.py @@ -1,6 +1,8 @@ from __future__ import annotations -from pydantic import BaseModel, ConfigDict, JsonValue +from typing import Annotated, Literal + +from pydantic import BaseModel, ConfigDict, Field, JsonValue from litellm.llms.base_llm.ocr.transformation import OCRResponse @@ -26,3 +28,27 @@ class Execution(BaseModel): request: CapturedRequest report: SDKReport + + +class SDKCommand(BaseModel): + model_config = ConfigDict(frozen=True) + + case_file: str + route: str + + +class WorkerSuccess(BaseModel): + model_config = ConfigDict(frozen=True) + + status: Literal["ok"] = "ok" + report: SDKReport + + +class WorkerFailure(BaseModel): + model_config = ConfigDict(frozen=True) + + status: Literal["error"] = "error" + error: str + + +WorkerResult = Annotated[WorkerSuccess | WorkerFailure, Field(discriminator="status")] diff --git a/tests/test_litellm/parity/replay.py b/tests/test_litellm/parity/replay.py index 1b067530199..017f9002303 100644 --- a/tests/test_litellm/parity/replay.py +++ b/tests/test_litellm/parity/replay.py @@ -4,6 +4,7 @@ import queue import threading from collections.abc import Generator from contextlib import contextmanager +from dataclasses import dataclass from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import Final @@ -28,23 +29,37 @@ EXCLUDED_RESPONSE_HEADERS: Final = frozenset({"content-length", "transfer-encodi class ReplayServer(ThreadingHTTPServer): daemon_threads = True - def __init__(self, status_code: int, headers: tuple[tuple[str, str], ...], body: bytes) -> None: + def __init__(self) -> None: super().__init__(("127.0.0.1", 0), _ReplayHandler) - self.status_code: Final = status_code - self.headers: Final = headers - self.body: Final = body + self.responses: queue.Queue[ReplayResponse] = queue.Queue() self.requests: queue.Queue[CapturedRequest] = queue.Queue() @property def url(self) -> str: return f"http://127.0.0.1:{self.server_address[1]}" + def enqueue_response(self, status_code: int, headers: tuple[tuple[str, str], ...], body: bytes) -> None: + self.responses.put(ReplayResponse(status_code=status_code, headers=headers, body=body)) + def take_request(self) -> CapturedRequest: request_count: Final = self.requests.qsize() if request_count != 1: raise AssertionError(f"expected exactly one provider request, received {request_count}") return self.requests.get_nowait() + def reset(self) -> None: + while not self.responses.empty(): + self.responses.get_nowait() + while not self.requests.empty(): + self.requests.get_nowait() + + +@dataclass(frozen=True, slots=True) +class ReplayResponse: + status_code: int + headers: tuple[tuple[str, str], ...] + body: bytes + class _ReplayHandler(BaseHTTPRequestHandler): protocol_version = "HTTP/1.1" @@ -70,26 +85,27 @@ class _ReplayHandler(BaseHTTPRequestHandler): user_agent=self.headers.get("user-agent"), ) ) - self.send_response_only(provider.status_code) - for name, value in provider.headers: + try: + response: Final = provider.responses.get(timeout=5) + except queue.Empty: + self.send_error(500, "no replay response queued") + return + self.send_response_only(response.status_code) + for name, value in response.headers: if name.lower() not in EXCLUDED_RESPONSE_HEADERS: self.send_header(name, value) - self.send_header("content-length", str(len(provider.body))) + self.send_header("content-length", str(len(response.body))) self.end_headers() - self.wfile.write(provider.body) + self.wfile.write(response.body) def log_message(self, format: str, *args: object) -> None: return @contextmanager -def replay_response( - status_code: int, - headers: tuple[tuple[str, str], ...], - body: bytes, -) -> Generator[ReplayServer]: - server: Final = ReplayServer(status_code=status_code, headers=headers, body=body) - thread: Final = threading.Thread(target=server.serve_forever, daemon=True) +def replay_server() -> Generator[ReplayServer]: + server: Final = ReplayServer() + thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.01}, daemon=True) thread.start() try: yield server diff --git a/tests/test_litellm/parity/runner.py b/tests/test_litellm/parity/runner.py index 2fe0984f5d1..83f652a582c 100644 --- a/tests/test_litellm/parity/runner.py +++ b/tests/test_litellm/parity/runner.py @@ -3,14 +3,27 @@ from __future__ import annotations import os import subprocess import sys +from collections import deque +from collections.abc import Generator +from concurrent.futures import ThreadPoolExecutor, TimeoutError +from contextlib import contextmanager from dataclasses import dataclass from pathlib import Path -from typing import Final +from typing import Final, TextIO, cast -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError -from tests.test_litellm.parity.models import Execution, SDKReport -from tests.test_litellm.parity.replay import ReplayServer +from tests.test_litellm.parity.models import ( + Execution, + SDKCommand, + WorkerFailure, + WorkerResult, + WorkerSuccess, +) +from tests.test_litellm.parity.replay import ReplayServer, replay_server + +WORKER_RESULT_PREFIX: Final = "LITELLM_PARITY_RESULT " +WORKER_RESULT_ADAPTER: Final[TypeAdapter[WorkerResult]] = TypeAdapter(WorkerResult) @dataclass(frozen=True, slots=True) @@ -19,55 +32,129 @@ class PythonScriptRunner: rust_env_var: str python_user_agent: str - def command(self, case_file: Path, route: str, provider_url: str, report_file: Path) -> tuple[str, ...]: + def command(self, provider_url: str) -> tuple[str, ...]: return ( sys.executable, str(self.entrypoint.resolve()), - str(case_file), - route, + "--parity-worker", provider_url, - str(report_file), ) +class PythonScriptWorker: + def __init__(self, runner: PythonScriptRunner, provider: ReplayServer, rust_enabled: bool) -> None: + project_root: Final = str(runner.entrypoint.resolve().parents[3]) + existing_pythonpath: Final = os.environ.get("PYTHONPATH") + env: Final = { + **os.environ, + runner.rust_env_var: "1" if rust_enabled else "0", + "LITELLM_USER_AGENT": runner.python_user_agent, + "PYTHONPATH": os.pathsep.join(path for path in (project_root, existing_pythonpath) if path), + } + self.mode: Final = "Rust" if rust_enabled else "Python" + self.provider: Final = provider + self.process: Final = subprocess.Popen( + runner.command(provider.url), + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + bufsize=1, + env=env, + ) + self.output_reader: Final = ThreadPoolExecutor(max_workers=1) + self.recent_output: Final[deque[str]] = deque(maxlen=100) + + def execute( + self, + case_file: Path, + route: str, + status_code: int, + headers: tuple[tuple[str, str], ...], + body: bytes, + ) -> Execution: + stdin: Final = self.process.stdin + if stdin is None or self.process.poll() is not None: + raise AssertionError(f"{self.mode} OCR worker exited before processing {case_file}") + self.provider.enqueue_response(status_code, headers, body) + command: Final = SDKCommand(case_file=str(case_file), route=route) + try: + stdin.write(f"{command.model_dump_json()}\n") + stdin.flush() + result: Final = self.output_reader.submit(self._read_result).result(timeout=60) + except TimeoutError as error: + self.provider.reset() + self.close() + raise AssertionError(f"{self.mode} OCR worker timed out after 60s while processing {case_file}") from error + except AssertionError: + self.provider.reset() + raise + except (BrokenPipeError, OSError) as error: + self.provider.reset() + raise AssertionError(self._failure_message(f"worker pipe failed while processing {case_file}")) from error + if isinstance(result, WorkerFailure): + self.provider.reset() + raise AssertionError(f"{self.mode} OCR worker failed while processing {case_file}:\n{result.error}") + assert isinstance(result, WorkerSuccess) + try: + return Execution(request=self.provider.take_request(), report=result.report) + except AssertionError: + self.provider.reset() + raise + + def _read_result(self) -> WorkerResult: + process_stdout: Final = self.process.stdout + if process_stdout is None: + raise AssertionError(self._failure_message("worker stdout is unavailable")) + stdout: Final = cast(TextIO, process_stdout) + line: Final = stdout.readline() + if not line: + raise AssertionError(self._failure_message("worker exited without returning a result")) + stripped: Final = line.rstrip() + if not stripped.startswith(WORKER_RESULT_PREFIX): + self.recent_output.append(stripped) + return self._read_result() + payload: Final = stripped.removeprefix(WORKER_RESULT_PREFIX) + try: + return WORKER_RESULT_ADAPTER.validate_json(payload) + except ValidationError as error: + raise AssertionError(self._failure_message("worker returned an invalid result")) from error + + def _failure_message(self, message: str) -> str: + output: Final = "\n".join(self.recent_output) + return f"{self.mode} OCR {message}" if not output else f"{self.mode} OCR {message}\noutput:\n{output}" + + def close(self) -> None: + stdin: Final = self.process.stdin + if stdin is not None and not stdin.closed: + stdin.close() + try: + self.process.wait(timeout=10) + except subprocess.TimeoutExpired: + self.process.terminate() + self.process.wait(timeout=10) + self.output_reader.shutdown(wait=True, cancel_futures=True) + + +@contextmanager +def execution_worker( + runner: PythonScriptRunner, + rust_enabled: bool, +) -> Generator[PythonScriptWorker]: + with replay_server() as provider: + worker: Final = PythonScriptWorker(runner, provider, rust_enabled) + try: + yield worker + finally: + worker.close() + + def run_execution( - runner: PythonScriptRunner, + worker: PythonScriptWorker, case_file: Path, route: str, - report_file: Path, - provider: ReplayServer, - rust_enabled: bool, + status_code: int, + headers: tuple[tuple[str, str], ...], + body: bytes, ) -> Execution: - project_root: Final = str(runner.entrypoint.resolve().parents[3]) - existing_pythonpath: Final = os.environ.get("PYTHONPATH") - env: Final = { - **os.environ, - runner.rust_env_var: "1" if rust_enabled else "0", - "LITELLM_USER_AGENT": runner.python_user_agent, - "PYTHONPATH": os.pathsep.join(path for path in (project_root, existing_pythonpath) if path), - } - mode: Final = "Rust" if rust_enabled else "Python" - command: Final = runner.command(case_file, route, provider.url, report_file) - try: - completed: Final = subprocess.run( - command, - capture_output=True, - text=True, - env=env, - timeout=60, - check=False, - ) - except subprocess.TimeoutExpired as error: - raise AssertionError(f"{mode} OCR subprocess timed out after {error.timeout}s: {' '.join(command)}") from error - if completed.returncode != 0: - raise AssertionError( - f"{mode} OCR subprocess failed with exit code {completed.returncode}\n" - f"command: {' '.join(command)}\nstdout:\n{completed.stdout}\nstderr:\n{completed.stderr}" - ) - if not report_file.is_file(): - raise AssertionError(f"{mode} OCR subprocess succeeded without writing report {report_file}") - try: - report: Final = SDKReport.model_validate_json(report_file.read_text(encoding="utf-8")) - except (OSError, ValidationError, ValueError) as error: - raise AssertionError(f"{mode} OCR subprocess wrote an invalid report at {report_file}: {error}") from error - return Execution(request=provider.take_request(), report=report) + return worker.execute(case_file, route, status_code, headers, body)