make it faster

This commit is contained in:
Yujong Lee 2026-08-29 14:44:19 -07:00 committed by GitHub
parent 4965ba7961
commit ea7113dc20
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 270 additions and 103 deletions

View file

@ -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])

View file

@ -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")]

View file

@ -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

View file

@ -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)