mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
make it faster
This commit is contained in:
parent
4965ba7961
commit
ea7113dc20
4 changed files with 270 additions and 103 deletions
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue