diff --git a/tests/test_litellm/_fixture_recorder.py b/tests/test_litellm/_fixture_recorder.py index fa9d73d47dd..a2d4ed06bbc 100644 --- a/tests/test_litellm/_fixture_recorder.py +++ b/tests/test_litellm/_fixture_recorder.py @@ -2,10 +2,9 @@ from __future__ import annotations import argparse import hashlib -import json import queue import threading -from collections.abc import Callable, Generator, Mapping +from collections.abc import Callable, Generator from contextlib import contextmanager from dataclasses import dataclass from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer @@ -13,95 +12,80 @@ from pathlib import Path from typing import Final, cast import httpx -from pydantic import BaseModel, ConfigDict, TypeAdapter from tests.test_litellm._json_fs_cache import JsonFileCache, canonical_json +from tests.test_litellm.ocr.fixture_models import ( + HttpHeader, + MistralOcrParityInput, + OcrParityCase, + RecordedHttpResponse, +) -JSON_OBJECT: Final = TypeAdapter(dict[str, object]) - - -class ProviderWireRequest(BaseModel): - model_config = ConfigDict(frozen=True) - - method: str - path: str - body: dict[str, object] - - -class FixtureRequest(BaseModel): - model_config = ConfigDict(frozen=True) - - provider: str - sdk_kwargs: dict[str, object] - provider_request: ProviderWireRequest - - -class FixtureResponse(BaseModel): - model_config = ConfigDict(frozen=True) - - status_code: int - headers: dict[str, str] - body: dict[str, object] - - -class Fixture(BaseModel): - model_config = ConfigDict(frozen=True) - - request: FixtureRequest - response: FixtureResponse +_HOP_BY_HOP_HEADERS: Final = frozenset( + { + "connection", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "te", + "trailer", + "transfer-encoding", + "upgrade", + } +) @dataclass(frozen=True, slots=True) class ProviderSpec: - name: str model: str upstream_base: str - api_key: str | None - upstream_model: str | None = None + api_key: str @dataclass(frozen=True, slots=True) class RecorderResult: - request: FixtureRequest - response: FixtureResponse | None + case: OcrParityCase cache_hit: bool @dataclass(frozen=True, slots=True) class GeneratorArgs: - providers: tuple[str, ...] examples: int fixture_dir: Path | None - requests_only: bool - responses_only: bool + model: str + + +def _excluded_headers(headers: tuple[tuple[str, str], ...]) -> frozenset[str]: + connection_values: Final = tuple(value for name, value in headers if name.lower() == "connection") + connection_headers: Final = frozenset( + token.strip().lower() for value in connection_values for token in value.split(",") if token.strip() + ) + return _HOP_BY_HOP_HEADERS | connection_headers + + +def _end_to_end_headers(headers: httpx.Headers) -> tuple[HttpHeader, ...]: + decoded: Final = tuple((name.decode("ascii"), value.decode("latin-1")) for name, value in headers.raw) + excluded: Final = _excluded_headers(decoded) | {"content-length"} + return tuple(HttpHeader(name=name, value=value) for name, value in decoded if name.lower() not in excluded) class _RecordingProvider(ThreadingHTTPServer): daemon_threads = True - def __init__( - self, - spec: ProviderSpec, - sdk_kwargs: dict[str, object], - cache: JsonFileCache, - requests_only: bool, - ) -> None: + def __init__(self, spec: ProviderSpec) -> None: super().__init__(("127.0.0.1", 0), _RecordingHandler) self.spec: Final = spec - self.sdk_kwargs: Final = sdk_kwargs - self.cache: Final = cache - self.requests_only: Final = requests_only - self.results: queue.Queue[RecorderResult] = queue.Queue() + self.responses: queue.Queue[RecordedHttpResponse] = queue.Queue() @property def url(self) -> str: return f"http://127.0.0.1:{self.server_address[1]}" - def take_result(self) -> RecorderResult: + def take_response(self) -> RecordedHttpResponse: try: - return self.results.get(timeout=5) + return self.responses.get(timeout=5) except queue.Empty as error: - raise RuntimeError("successful SDK call did not produce a recorder result") from error + raise RuntimeError("successful SDK call did not produce a recorded response") from error class _RecordingHandler(BaseHTTPRequestHandler): @@ -111,81 +95,41 @@ class _RecordingHandler(BaseHTTPRequestHandler): provider: Final = self.server assert isinstance(provider, _RecordingProvider) length: Final = int(self.headers.get("content-length") or "0") - body: Final = JSON_OBJECT.validate_json(self.rfile.read(length)) - fixture_request: Final = FixtureRequest( - provider=provider.spec.name, - sdk_kwargs=provider.sdk_kwargs, - provider_request=ProviderWireRequest(method=self.command, path=self.path, body=body), - ) - cache_key: Final = fixture_cache_key( - provider.spec.name, - fixture_request.sdk_kwargs, - fixture_request.provider_request, - ) - cached_value: Final = provider.cache.get(cache_key) - if cached_value is not None: - cached_request: Final = FixtureRequest.model_validate(cached_value["request"]) - raw_cached_response: Final = cached_value.get("response") - if raw_cached_response is not None: - cached_response: Final = FixtureResponse.model_validate(raw_cached_response) - provider.results.put(RecorderResult(request=cached_request, response=cached_response, cache_hit=True)) - self._send_fixture_response(cached_response) - return - if provider.requests_only: - provider.results.put(RecorderResult(request=cached_request, response=None, cache_hit=True)) - self._send_response(200, {"content-type": "application/json"}, b"{}") - return - - if provider.requests_only: - provider.results.put(RecorderResult(request=fixture_request, response=None, cache_hit=False)) - self._send_response(200, {"content-type": "application/json"}, b"{}") - return - + request_body: Final = self.rfile.read(length) + raw_headers: Final = tuple(self.headers.raw_items()) + excluded: Final = _excluded_headers(raw_headers) | {"host", "content-length"} + forwarded_headers: Final = tuple((name, value) for name, value in raw_headers if name.lower() not in excluded) upstream_url: Final = f"{provider.spec.upstream_base.rstrip('/')}{self.path}" - forwarded_headers: Final = { - name: value - for name, value in self.headers.items() - if name.lower() not in {"host", "content-length", "accept-encoding", "x-parity-case"} - } - upstream_body: Final = ( - {**body, "model": provider.spec.upstream_model} if provider.spec.upstream_model is not None else body - ) + try: - upstream_response: Final = httpx.post( + with httpx.stream( + self.command, upstream_url, headers=forwarded_headers, - content=json.dumps(upstream_body, separators=(",", ":")), + content=request_body, timeout=120, - ) + ) as upstream: + response_body: Final = b"".join(upstream.iter_raw()) + recorded_response: Final = RecordedHttpResponse.from_bytes( + status_code=upstream.status_code, + headers=_end_to_end_headers(upstream.headers), + body=response_body, + ) except httpx.HTTPError as error: - error_body: Final = json.dumps({"error": str(error)}).encode() - self._send_response(502, {"content-type": "application/json"}, error_body) + self._send_response(502, (), str(error).encode("utf-8")) return - raw_content_type: Final = cast(object, upstream_response.headers.get("content-type", "application/json")) - content_type: Final = raw_content_type if isinstance(raw_content_type, str) else "application/json" - response_headers: Final = {"content-type": content_type.split(";", 1)[0]} - if not upstream_response.is_success: - self._send_response(upstream_response.status_code, response_headers, upstream_response.content) - return + if 200 <= recorded_response.status_code < 300: + provider.responses.put(recorded_response) + self._send_recorded_response(recorded_response) - upstream_response_body: Final = JSON_OBJECT.validate_json(upstream_response.content) - fixture_response: Final = FixtureResponse( - status_code=upstream_response.status_code, - headers=response_headers, - body=upstream_response_body, - ) - provider.results.put(RecorderResult(request=fixture_request, response=fixture_response, cache_hit=False)) - self._send_fixture_response(fixture_response) + def _send_recorded_response(self, response: RecordedHttpResponse) -> None: + self._send_response(response.status_code, response.headers, response.body_bytes()) - def _send_fixture_response(self, response: FixtureResponse) -> None: - response_body: Final = json.dumps(response.body, separators=(",", ":")).encode() - self._send_response(response.status_code, response.headers, response_body) - - def _send_response(self, status_code: int, headers: Mapping[str, str], body: bytes) -> None: - self.send_response(status_code) - for name, value in headers.items(): - self.send_header(name, value) + def _send_response(self, status_code: int, headers: tuple[HttpHeader, ...], body: bytes) -> None: + self.send_response_only(status_code) + for header in headers: + self.send_header(header.name, header.value) self.send_header("content-length", str(len(body))) self.end_headers() self.wfile.write(body) @@ -195,18 +139,8 @@ class _RecordingHandler(BaseHTTPRequestHandler): @contextmanager -def _recording_provider( - spec: ProviderSpec, - sdk_kwargs: dict[str, object], - cache: JsonFileCache, - requests_only: bool, -) -> Generator[_RecordingProvider]: - server: Final = _RecordingProvider( - spec=spec, - sdk_kwargs=sdk_kwargs, - cache=cache, - requests_only=requests_only, - ) +def _recording_provider(spec: ProviderSpec) -> Generator[_RecordingProvider]: + server: Final = _RecordingProvider(spec) thread: Final = threading.Thread(target=server.serve_forever, daemon=True) thread.start() try: @@ -217,87 +151,41 @@ def _recording_provider( thread.join(timeout=5) -def fixture_cache_key( - provider: str, - sdk_kwargs: dict[str, object], - request: ProviderWireRequest, -) -> dict[str, object]: - return { - "provider": provider, - "sdk_kwargs": sdk_kwargs, - "request": request.model_dump(mode="json"), - } +def fixture_cache_key(case_input: MistralOcrParityInput) -> dict[str, object]: + return case_input.canonical_input() def record_case( spec: ProviderSpec, root: Path, - sdk_kwargs: dict[str, object], - requests_only: bool, + case_input: MistralOcrParityInput, sdk_call: Callable[..., object], ) -> RecorderResult: - cache: Final = JsonFileCache(root / spec.name) - with _recording_provider( - spec=spec, - sdk_kwargs=sdk_kwargs, - cache=cache, - requests_only=requests_only, - ) as recorder: - try: - sdk_call(api_base=recorder.url, api_key=spec.api_key, **sdk_kwargs) - except Exception: - if not requests_only: - raise - result: Final = recorder.take_result() + cache: Final = JsonFileCache(root) + cache_key: Final = fixture_cache_key(case_input) + cached: Final = cache.get(cache_key) + if cached is not None: + return RecorderResult(case=OcrParityCase.model_validate(cached), cache_hit=True) - if not result.cache_hit: - value: Final = ( - Fixture(request=result.request, response=result.response).model_dump(mode="json") - if result.response is not None - else {"request": result.request.model_dump(mode="json")} - ) - cache.put( - fixture_cache_key(spec.name, result.request.sdk_kwargs, result.request.provider_request), - value, - ) - return result + with _recording_provider(spec) as recorder: + sdk_call(api_base=recorder.url, api_key=spec.api_key, **case_input.as_sdk_kwargs()) + upstream_response: Final = recorder.take_response() + + case: Final = OcrParityCase(input=case_input, upstream_response=upstream_response) + cache.put(cache_key, cast(dict[str, object], case.model_dump(mode="json", exclude_unset=True))) + return RecorderResult(case=case, cache_hit=False) -def pending_requests(cache: JsonFileCache) -> tuple[FixtureRequest, ...]: - return tuple( - FixtureRequest.model_validate(value["request"]) for value in cache.values() if value.get("response") is None - ) - - -def fill_missing_responses( - specs: tuple[ProviderSpec, ...], - root: Path, - sdk_call: Callable[..., object], -) -> tuple[RecorderResult, ...]: - return tuple( - record_case(spec, root, request.sdk_kwargs, requests_only=False, sdk_call=sdk_call) - for spec in specs - for request in pending_requests(JsonFileCache(root / spec.name)) - if request.sdk_kwargs.get("model") == spec.model - ) - - -def parse_generator_args(provider_names: tuple[str, ...]) -> GeneratorArgs: +def parse_generator_args() -> GeneratorArgs: parser: Final = argparse.ArgumentParser() - parser.add_argument("--provider", action="append", choices=provider_names) parser.add_argument("--examples", type=int, default=4) parser.add_argument("--fixture-dir", type=Path) - mode: Final = parser.add_mutually_exclusive_group() - mode.add_argument("--requests-only", action="store_true", help="record deterministic requests without API calls") - mode.add_argument("--responses-only", action="store_true", help="fill responses for saved pending requests") + parser.add_argument("--model", default="mistral/mistral-ocr-latest") namespace: Final = parser.parse_args() - providers: Final = cast(list[str] | None, namespace.provider) return GeneratorArgs( - providers=tuple(providers) if providers else provider_names, examples=cast(int, namespace.examples), fixture_dir=cast(Path | None, namespace.fixture_dir), - requests_only=cast(bool, namespace.requests_only), - responses_only=cast(bool, namespace.responses_only), + model=cast(str, namespace.model), ) @@ -305,17 +193,11 @@ def fixture_directory(configured: Path | None, env_value: str | None, default: P return (configured or Path(env_value or default)).expanduser() -def recorded_fixtures(directory: Path) -> tuple[Fixture, ...]: - return tuple( - Fixture.model_validate(raw_fixture) - for raw_fixture in JsonFileCache(directory).values() - if raw_fixture.get("response") is not None - ) +def recorded_fixtures(directory: Path) -> tuple[OcrParityCase, ...]: + return tuple(OcrParityCase.model_validate(raw_fixture) for raw_fixture in JsonFileCache(directory).values()) -def fixture_id(fixture: Fixture) -> str: - raw_model: Final = fixture.request.sdk_kwargs.get("model") - model: Final = raw_model if isinstance(raw_model, str) else "unknown-model" - request_json: Final = canonical_json(fixture.request.provider_request.model_dump(mode="json")) - digest: Final = hashlib.sha256(request_json.encode("utf-8")).hexdigest()[:8] - return f"{fixture.request.provider}-{model.rsplit('/', 1)[-1]}-{digest}" +def fixture_id(fixture: OcrParityCase) -> str: + input_json: Final = canonical_json(fixture.input.canonical_input()) + digest: Final = hashlib.sha256(input_json.encode("utf-8")).hexdigest()[:8] + return f"mistral-{fixture.input.model.rsplit('/', 1)[-1]}-{digest}" diff --git a/tests/test_litellm/_json_fs_cache.py b/tests/test_litellm/_json_fs_cache.py index 61ae77b915a..9ce43931e36 100644 --- a/tests/test_litellm/_json_fs_cache.py +++ b/tests/test_litellm/_json_fs_cache.py @@ -33,7 +33,9 @@ class JsonFileCache: def put(self, key: Mapping[str, object], value: Mapping[str, object]) -> Path: self.root.mkdir(parents=True, exist_ok=True) path: Final = self.path_for(key) - path.write_text(json.dumps(value, indent=2, sort_keys=True) + "\n", encoding="utf-8") + temporary_path: Final = path.with_suffix(".tmp") + temporary_path.write_text(json.dumps(value, indent=2, sort_keys=True) + "\n", encoding="utf-8") + temporary_path.replace(path) return path def values(self) -> tuple[dict[str, object], ...]: diff --git a/tests/test_litellm/ocr/conftest.py b/tests/test_litellm/ocr/conftest.py index 6da67be1984..2822efe23ad 100644 --- a/tests/test_litellm/ocr/conftest.py +++ b/tests/test_litellm/ocr/conftest.py @@ -9,12 +9,15 @@ import pytest from tests.test_litellm._fixture_recorder import fixture_id, recorded_fixtures FIXTURE_DIR_ENV: Final = "LITELLM_OCR_FIXTURE_DIR" -pytest_plugins: Final = ("tests.test_litellm.parity.pytest_plugin",) def _fixture_directory() -> Path: - configured: Final = os.environ.get(FIXTURE_DIR_ENV) - return Path(configured).expanduser() if configured else Path(__file__).with_name(".fixtures") + if FIXTURE_DIR_ENV not in os.environ: + return Path(__file__).with_name(".fixtures") + configured: Final = os.environ[FIXTURE_DIR_ENV] + if not configured: + raise pytest.UsageError(f"{FIXTURE_DIR_ENV} is set but empty") + return Path(configured).expanduser() def pytest_generate_tests(metafunc: pytest.Metafunc) -> None: @@ -22,6 +25,8 @@ def pytest_generate_tests(metafunc: pytest.Metafunc) -> None: return fixtures: Final = recorded_fixtures(_fixture_directory()) if not fixtures: + if FIXTURE_DIR_ENV in os.environ: + raise pytest.UsageError(f"no recorded OCR fixtures in {_fixture_directory()}") metafunc.parametrize( "ocr_fixture", ( diff --git a/tests/test_litellm/ocr/fixture_models.py b/tests/test_litellm/ocr/fixture_models.py new file mode 100644 index 00000000000..2b06bf3392e --- /dev/null +++ b/tests/test_litellm/ocr/fixture_models.py @@ -0,0 +1,91 @@ +from __future__ import annotations + +import base64 +from typing import Annotated, Literal, cast + +from pydantic import BaseModel, ConfigDict, Field, JsonValue, NonNegativeInt, PositiveInt + + +class _FixtureModel(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid", populate_by_name=True, serialize_by_alias=True) + + +class ImageUrlDocument(_FixtureModel): + type: Literal["image_url"] + image_url: str + + +class DocumentUrlDocument(_FixtureModel): + type: Literal["document_url"] + document_url: str + + +MistralOcrDocument = Annotated[ImageUrlDocument | DocumentUrlDocument, Field(discriminator="type")] + + +class JsonSchemaDefinition(_FixtureModel): + name: str + schema_value: JsonValue = Field(alias="schema") + strict: bool | None = None + + +class AnnotationFormat(_FixtureModel): + type: Literal["json_schema"] + json_schema: JsonSchemaDefinition + + +class MistralOcrParityInput(_FixtureModel): + model: str + document: MistralOcrDocument + pages: list[NonNegativeInt] | None = None + include_image_base64: bool | None = None + image_limit: PositiveInt | None = None + image_min_size: NonNegativeInt | None = None + bbox_annotation_format: AnnotationFormat | None = None + document_annotation_format: AnnotationFormat | None = None + document_annotation_prompt: str | None = None + extract_header: bool | None = None + extract_footer: bool | None = None + table_format: Literal["markdown", "html"] | None = None + confidence_scores_granularity: Literal["word", "page", "block"] | None = None + include_blocks: bool | None = None + id: str | None = None + + def as_sdk_kwargs(self) -> dict[str, object]: + return cast(dict[str, object], self.model_dump(mode="python", exclude_unset=True)) + + def canonical_input(self) -> dict[str, object]: + return cast(dict[str, object], self.model_dump(mode="json", exclude_unset=True)) + + +class HttpHeader(_FixtureModel): + name: str + value: str + + +class RecordedHttpResponse(_FixtureModel): + kind: Literal["http"] = "http" + status_code: int + headers: tuple[HttpHeader, ...] + body_b64: str + + @classmethod + def from_bytes( + cls, + status_code: int, + headers: tuple[HttpHeader, ...], + body: bytes, + ) -> RecordedHttpResponse: + return cls( + status_code=status_code, + headers=headers, + body_b64=base64.b64encode(body).decode("ascii"), + ) + + def body_bytes(self) -> bytes: + return base64.b64decode(self.body_b64, validate=True) + + +class OcrParityCase(_FixtureModel): + input: MistralOcrParityInput + upstream_response: RecordedHttpResponse diff --git a/tests/test_litellm/ocr/generate_fixtures.py b/tests/test_litellm/ocr/generate_fixtures.py index efb73ddef0c..99f2ca05648 100644 --- a/tests/test_litellm/ocr/generate_fixtures.py +++ b/tests/test_litellm/ocr/generate_fixtures.py @@ -1,6 +1,5 @@ from __future__ import annotations -import base64 import logging import os from collections.abc import Callable @@ -8,345 +7,97 @@ from pathlib import Path from typing import Final, cast from urllib.parse import quote -import httpx from dotenv import load_dotenv from hypothesis import given, settings from hypothesis import strategies as st from hypothesis.strategies import SearchStrategy -from pydantic import TypeAdapter import litellm -from litellm.llms.reducto.common import ( - REDUCTO_API_BASE, - extract_file_id_or_bytes, - upload_bytes_sync, -) +from litellm.rust_bridge.ocr import use_litellm_rust from tests.test_litellm._fixture_recorder import ( ProviderSpec, - RecorderResult, - fill_missing_responses, fixture_directory, parse_generator_args, record_case, ) +from tests.test_litellm.ocr.fixture_models import ( + AnnotationFormat, + DocumentUrlDocument, + ImageUrlDocument, + JsonSchemaDefinition, + MistralOcrParityInput, +) FIXTURE_DIR_ENV: Final = "LITELLM_OCR_FIXTURE_DIR" -JSON_OBJECT: Final = TypeAdapter(dict[str, object]) -PROVIDER_NAMES: Final = ("mistral", "reducto") LOGGER: Final = logging.getLogger(__name__) _TEXT: Final = st.from_regex(r"[A-Za-z0-9 ]{1,24}", fullmatch=True) _VALUE_TEXT: Final = st.text(alphabet="abcdefghijklmnopqrstuvwxyz0123456789 -_", min_size=1, max_size=32) -_NULLABLE_TEXT: Final = st.one_of(st.none(), _VALUE_TEXT) -_POSITIVE_INTEGER: Final = st.integers(min_value=1, max_value=10_000) -_NON_NEGATIVE_INTEGER: Final = st.integers(min_value=0, max_value=10_000) -_SMALL_NUMBER: Final = st.floats(min_value=0.01, max_value=30.0, allow_nan=False, allow_infinity=False) -_REQUEST_ONLY_IMAGE: Final = ( - "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+A8AAQUBAScY42YAAAAASUVORK5CYII=" -) -_BLOCK_TYPES: Final = ( - "Header", - "Footer", - "Title", - "Section Header", - "Page Number", - "List Item", - "Figure", - "Table", - "Key Value", - "Text", - "Comment", - "Signature", -) - - -def _optional_object(optional: dict[str, SearchStrategy[object]]) -> SearchStrategy[dict[str, object]]: - return st.fixed_dictionaries({}, optional=optional) - - -def _merge_objects(left: dict[str, object], right: dict[str, object]) -> dict[str, object]: - return {**left, **right} - - -def _page_range(start: int, length: int) -> dict[str, object]: - return {"start": start, "end": start + length} - - -def _annotation_format(name: str, strict: bool) -> dict[str, object]: - return { - "type": "json_schema", - "json_schema": { - "name": name, - "schema": {"type": "object", "properties": {}, "additionalProperties": False}, - "strict": strict, - }, - } - - -def _annotation_format_only(value: dict[str, object]) -> dict[str, object]: - return {"document_annotation_format": value} - - -def _null_annotation_format() -> dict[str, object]: - return {"document_annotation_format": None} - - -def _annotation_format_and_prompt(value: dict[str, object], prompt: str | None) -> dict[str, object]: - return {"document_annotation_format": value, "document_annotation_prompt": prompt} - - -def _table_agentic(prompt: str | None, mode: str) -> dict[str, object]: - return {"scope": "table", "prompt": prompt, "mode": mode} - - -def _figure_agentic(prompt: str | None, advanced: bool, overlays: bool) -> dict[str, object]: - return { - "scope": "figure", - "prompt": prompt, - "advanced_chart_agent": advanced, - "return_overlays": overlays, - } - - -def _text_agentic(prompt: str | None) -> dict[str, object]: - return {"scope": "text", "prompt": prompt} - - -def _agentic_scope(value: dict[str, object]) -> object: - return value["scope"] - - -_PAGE_RANGE: Final = st.builds( - _page_range, st.integers(min_value=1, max_value=1000), st.integers(min_value=0, max_value=50) -) -_PAGE_RANGE_VALUE: Final = st.one_of( - _PAGE_RANGE, - st.lists(_PAGE_RANGE, min_size=1, max_size=3), - st.lists(st.integers(min_value=1, max_value=1000), min_size=1, max_size=5, unique=True), - st.lists(_VALUE_TEXT, min_size=1, max_size=3, unique=True), -) -_ANNOTATION_FORMAT: Final[SearchStrategy[dict[str, object]]] = st.builds(_annotation_format, _VALUE_TEXT, st.booleans()) -_DOCUMENT_ANNOTATION: Final[SearchStrategy[dict[str, object]]] = st.one_of( - st.just(dict[str, object]()), - st.just(_null_annotation_format()), - st.builds(_annotation_format_only, _ANNOTATION_FORMAT), - st.builds(_annotation_format_and_prompt, _ANNOTATION_FORMAT, _NULLABLE_TEXT), -) - - -def _mistral_options(model: str) -> SearchStrategy[dict[str, object]]: - model_name: Final = model.rsplit("/", 1)[-1] - confidence_values: Final = ("word", "page") if model_name == "mistral-ocr-4-0" else ("word", "page", "block") - confidence_option: Final[dict[str, SearchStrategy[object]]] = ( - {} - if model_name == "mistral-ocr-2512" - else {"confidence_scores_granularity": st.one_of(st.none(), st.sampled_from(confidence_values))} - ) - independent: Final = _optional_object( - { - "pages": st.one_of( - st.none(), - st.lists(_NON_NEGATIVE_INTEGER, min_size=1, max_size=5, unique=True), - st.sampled_from(("0", "0,1,2", "0-5", "0,2-4")), - ), - "include_image_base64": st.one_of(st.none(), st.booleans()), - "image_limit": st.one_of(st.none(), _POSITIVE_INTEGER), - "image_min_size": st.one_of(st.none(), _NON_NEGATIVE_INTEGER), - "bbox_annotation_format": st.one_of(st.none(), _ANNOTATION_FORMAT), - "extract_header": st.booleans(), - "extract_footer": st.booleans(), - "table_format": st.one_of(st.none(), st.sampled_from(("markdown", "html"))), - "include_blocks": st.booleans(), - **confidence_option, - } - ) - return st.builds(_merge_objects, independent, _DOCUMENT_ANNOTATION) - - -_AGENTIC_TABLE: Final[SearchStrategy[dict[str, object]]] = st.builds( - _table_agentic, - _NULLABLE_TEXT, - st.sampled_from(("default", "auto", "max")), -) -_AGENTIC_FIGURE: Final[SearchStrategy[dict[str, object]]] = st.builds( - _figure_agentic, - _NULLABLE_TEXT, - st.booleans(), - st.booleans(), -) -_AGENTIC_TEXT: Final[SearchStrategy[dict[str, object]]] = st.builds(_text_agentic, _NULLABLE_TEXT) -_REDUCTO_ENHANCE: Final = _optional_object( - { - "agentic": st.lists( - st.one_of(_AGENTIC_TABLE, _AGENTIC_FIGURE, _AGENTIC_TEXT), - max_size=3, - unique_by=_agentic_scope, - ), - "summarize_figures": st.booleans(), - "intelligent_ordering": st.booleans(), - } -) -_CHUNKING: Final = _optional_object( - { - "chunk_mode": st.sampled_from(("variable", "section", "page", "disabled", "block", "page_sections")), - "chunk_size": st.one_of(st.none(), _POSITIVE_INTEGER), - "chunk_overlap": _NON_NEGATIVE_INTEGER, - } -) -_LEGACY_CHUNKING: Final = _optional_object( - { - "chunk_mode": st.sampled_from(("variable", "section", "page", "disabled", "block", "page_sections")), - "chunk_size": _POSITIVE_INTEGER, - "chunk_overlap": _NON_NEGATIVE_INTEGER, - } -) -_REDUCTO_RETRIEVAL: Final = _optional_object( - { - "chunking": _CHUNKING, - "filter_blocks": st.lists(st.sampled_from(_BLOCK_TYPES), max_size=len(_BLOCK_TYPES), unique=True), - "embedding_optimized": st.booleans(), - } -) -_REDUCTO_FORMATTING: Final = _optional_object( - { - "add_page_markers": st.booleans(), - "table_output_format": st.sampled_from(("html", "json", "md", "jsonbbox", "dynamic", "csv")), - "merge_tables": st.booleans(), - "include": st.lists( - st.sampled_from( - ("change_tracking", "highlight", "comments", "hyperlinks", "signatures", "ignore_watermarks") - ), - max_size=6, - unique=True, - ), - } -) -_SPLIT_TABLE_SIZE: Final = st.one_of( - _POSITIVE_INTEGER, - _optional_object( - {"row": st.one_of(st.none(), _POSITIVE_INTEGER), "column": st.one_of(st.none(), _POSITIVE_INTEGER)} +_ANNOTATION_FORMAT: Final[SearchStrategy[AnnotationFormat]] = st.builds( + AnnotationFormat, + type=st.just("json_schema"), + json_schema=st.builds( + JsonSchemaDefinition, + name=_VALUE_TEXT, + schema_value=st.just({"type": "object", "properties": {}, "additionalProperties": False}), + strict=st.one_of(st.none(), st.booleans()), ), ) -_REDUCTO_SPREADSHEET: Final = _optional_object( - { - "split_large_tables": _optional_object({"enabled": st.booleans(), "size": _SPLIT_TABLE_SIZE}), - "include": st.lists(st.sampled_from(("cell_colors", "formula", "dropdowns")), max_size=3, unique=True), - "clustering": st.sampled_from(("accurate", "fast", "disabled")), - "exclude": st.lists( - st.sampled_from(("hidden_sheets", "hidden_rows", "hidden_cols", "styling", "spreadsheet_images")), - max_size=5, - unique=True, + + +def _image_document(text: str, font_size: int) -> ImageUrlDocument: + url: Final = f"https://dummyjson.com/image/800x300/ffffff/000000?text={quote(text)}&fontSize={font_size}" + return ImageUrlDocument(type="image_url", image_url=url) + + +def _input_strategy(model: str) -> SearchStrategy[MistralOcrParityInput]: + document_strategy: Final = st.one_of( + st.builds(_image_document, _TEXT, st.integers(min_value=12, max_value=36)), + st.just( + DocumentUrlDocument( + type="document_url", + document_url="https://arxiv.org/pdf/2201.04234", + ) ), - "max_cell_count": st.one_of(st.none(), _POSITIVE_INTEGER), - } -) -_TENANT_THROTTLING: Final = st.fixed_dictionaries( - {"tenant_id": st.text(alphabet="abcdefghijklmnopqrstuvwxyz0123456789-_", min_size=1, max_size=256)}, - optional={"max_share": st.floats(min_value=0.01, max_value=1.0, allow_nan=False, allow_infinity=False)}, -) -_REDUCTO_SETTINGS: Final = _optional_object( - { - "ocr_system": st.sampled_from(("standard", "legacy")), - "extraction_mode": st.sampled_from(("ocr", "hybrid")), - "force_url_result": st.booleans(), - "force_file_extension": _NULLABLE_TEXT, - "return_ocr_data": st.booleans(), - "return_images": st.lists(st.sampled_from(("figure", "table", "page")), max_size=3, unique=True), - "embed_pdf_metadata": st.booleans(), - "embed_pdf_metadata_dpi": st.integers(min_value=50, max_value=250), - "persist_results": st.booleans(), - "tenant_throttling": st.one_of(st.none(), _TENANT_THROTTLING), - "timeout": st.one_of(st.none(), _SMALL_NUMBER), - "page_range": st.one_of(st.none(), _PAGE_RANGE_VALUE), - "document_password": _NULLABLE_TEXT, - "hybrid_vpc": _optional_object({"environment": _NULLABLE_TEXT}), - } -) -_REDUCTO_V3_OPTIONS: Final = _optional_object( - { - "enhance": _REDUCTO_ENHANCE, - "retrieval": _REDUCTO_RETRIEVAL, - "formatting": _REDUCTO_FORMATTING, - "spreadsheet": _REDUCTO_SPREADSHEET, - "settings": _REDUCTO_SETTINGS, - } -) -_LEGACY_SUMMARY: Final = _optional_object( - {"enabled": st.booleans(), "prompt": _VALUE_TEXT, "override": st.booleans(), "advanced_chart_agent": st.booleans()} -) -_REDUCTO_LEGACY_OPTIONS: Final = _optional_object( - { - "ocr_mode": st.sampled_from(("standard", "agentic")), - "extraction_mode": st.sampled_from(("ocr", "metadata", "hybrid")), - "chunking": _LEGACY_CHUNKING, - "table_summary": _optional_object({"enabled": st.booleans(), "prompt": _VALUE_TEXT}), - "figure_summary": _LEGACY_SUMMARY, - "filter_blocks": st.lists(st.sampled_from(_BLOCK_TYPES), max_size=len(_BLOCK_TYPES), unique=True), - "force_url_result": st.booleans(), - } -) -_REDUCTO_LEGACY_ADVANCED: Final = _optional_object( - { - "ocr_system": st.sampled_from(("highres", "multilingual", "combined", "reducto", "legacy")), - "table_output_format": st.sampled_from(("html", "json", "md", "jsonbbox", "dynamic", "ai_json", "csv")), - "merge_tables": st.booleans(), - "include_formula_information": st.booleans(), - "include_color_information": st.booleans(), - "include_dropdown_information": st.booleans(), - "continue_hierarchy": st.booleans(), - "keep_line_breaks": st.booleans(), - "page_range": _PAGE_RANGE_VALUE, - "force_file_extension": _VALUE_TEXT, - "large_table_chunking": _optional_object({"enabled": st.booleans(), "size": _POSITIVE_INTEGER}), - "spreadsheet_table_clustering": st.sampled_from(("default", "disabled", "intelligent")), - "max_cell_count": st.one_of(st.none(), _POSITIVE_INTEGER), - "add_page_markers": st.booleans(), - "remove_text_formatting": st.booleans(), - "return_ocr_data": st.booleans(), - "document_password": _VALUE_TEXT, - "filter_line_numbers": st.booleans(), - "read_comments": st.booleans(), - "persist_results": st.booleans(), - "exclude_hidden_sheets": st.booleans(), - "exclude_hidden_rows_cols": st.booleans(), - "enable_change_tracking": st.booleans(), - "enable_highlight_detection": st.booleans(), - "ignore_watermarks": st.booleans(), - } -) -_REDUCTO_LEGACY_EXPERIMENTAL: Final = _optional_object( - { - "enrich": _optional_object( - { - "enabled": st.booleans(), - "mode": st.sampled_from(("standard", "page", "table", "table_auto")), - "prompt": _VALUE_TEXT, - } - ), - "layout_enrichment": st.booleans(), - "enable_checkboxes": st.booleans(), - "enable_equations": st.booleans(), - "rotate_pages": st.booleans(), - "rotate_figures": st.booleans(), - "enable_scripts": st.booleans(), - "return_figure_images": st.booleans(), - "return_table_images": st.booleans(), - "return_page_images": st.booleans(), - "layout_model": st.sampled_from(("default", "beta")), - "embed_text_metadata_pdf": st.booleans(), - "embed_pdf_metadata_dpi": st.integers(min_value=50, max_value=250), - "detect_signatures": st.booleans(), - "danger_filter_wide_boxes": st.booleans(), - "user_specified_timeout_seconds": st.one_of(st.none(), _SMALL_NUMBER), - } -) -_REDUCTO_LEGACY_ROOT: Final = _optional_object( - { - "options": _REDUCTO_LEGACY_OPTIONS, - "advanced_options": _REDUCTO_LEGACY_ADVANCED, - "experimental_options": _REDUCTO_LEGACY_EXPERIMENTAL, - "priority": st.booleans(), - } -) + ) + input_values: Final = st.fixed_dictionaries( + {"model": st.just(model), "document": document_strategy}, + optional={ + "pages": st.one_of( + st.none(), + st.lists(st.integers(min_value=0, max_value=20), min_size=1, max_size=5, unique=True), + ), + "include_image_base64": st.one_of(st.none(), st.booleans()), + "image_limit": st.one_of(st.none(), st.integers(min_value=1, max_value=100)), + "image_min_size": st.one_of(st.none(), st.integers(min_value=0, max_value=10_000)), + "bbox_annotation_format": st.one_of(st.none(), _ANNOTATION_FORMAT), + "document_annotation_format": st.one_of(st.none(), _ANNOTATION_FORMAT), + "document_annotation_prompt": st.one_of(st.none(), _VALUE_TEXT), + "extract_header": st.one_of(st.none(), st.booleans()), + "extract_footer": st.one_of(st.none(), st.booleans()), + "table_format": st.one_of(st.none(), st.sampled_from(("markdown", "html"))), + "confidence_scores_granularity": st.one_of( + st.none(), st.sampled_from(("word", "page", "block")) + ), + "include_blocks": st.one_of(st.none(), st.booleans()), + "id": st.one_of(st.none(), _VALUE_TEXT), + }, + ) + return input_values.map(MistralOcrParityInput.model_validate) + + +def _generate_examples( + spec: ProviderSpec, + root: Path, + examples: int, + sdk_call: Callable[..., object], +) -> None: + @settings(max_examples=examples, deadline=None, derandomize=True) + @given(case_input=_input_strategy(spec.model)) + def generate_case(case_input: MistralOcrParityInput) -> None: + result: Final = record_case(spec, root, case_input, sdk_call) + LOGGER.info("%s %s", "cached" if result.cache_hit else "recorded", result.case.input.model) + + generate_case() def _mistral_upstream_base() -> str: @@ -354,186 +105,21 @@ def _mistral_upstream_base() -> str: return configured.removesuffix("/v1") -def _provider_specs( - selected: tuple[str, ...], - requests_only: bool = False, - all_models: bool = False, -) -> tuple[ProviderSpec, ...]: - mistral_key: Final = os.environ.get("MISTRAL_API_KEY") or os.environ.get("LITELLM_API_KEY") - reducto_key: Final = os.environ.get("REDUCTO_API_KEY") - mistral_models: Final = ( - ("mistral/mistral-ocr-2512", "mistral/mistral-ocr-4-0", "mistral/mistral-ocr-4-1") - if requests_only or all_models - else ("mistral/mistral-ocr-latest",) - ) - reducto_models: Final = ("reducto/parse-v3", "reducto/parse-legacy") - mistral_specs: Final = ( - tuple( - ProviderSpec( - name="mistral", - model=model, - upstream_base=_mistral_upstream_base(), - api_key=mistral_key or "request-only-key", - upstream_model=os.environ.get("MISTRAL_OCR_UPSTREAM_MODEL"), - ) - for model in mistral_models - ) - if "mistral" in selected and (mistral_key is not None or requests_only) - else () - ) - reducto_specs: Final = ( - tuple( - ProviderSpec( - name="reducto", - model=model, - upstream_base=os.environ.get("REDUCTO_API_BASE", REDUCTO_API_BASE), - api_key=reducto_key or "request-only-key", - ) - for model in reducto_models - ) - if "reducto" in selected and (reducto_key is not None or requests_only) - else () - ) - specs: Final = (*mistral_specs, *reducto_specs) - present: Final = frozenset(spec.name for spec in specs) - missing: Final = tuple(name for name in selected if name not in present) - if missing: - LOGGER.warning("Skipping providers without credentials: %s", ", ".join(missing)) - return specs - - -def _image_data_uri(text: str, font_size: int) -> str: - url: Final = f"https://dummyjson.com/image/800x300/ffffff/000000?text={quote(text)}&fontSize={font_size}" - response: Final = httpx.get(url, timeout=30, follow_redirects=True) - response.raise_for_status() - raw_content_type: Final = cast(object, response.headers.get("content-type", "image/png")) - content_type: Final = raw_content_type if isinstance(raw_content_type, str) else "image/png" - media_type: Final = content_type.split(";", 1)[0] - encoded: Final = base64.b64encode(response.content).decode("ascii") - return f"data:{media_type};base64,{encoded}" - - -def _sdk_kwargs( - spec: ProviderSpec, - image_data_uri: str, - options: dict[str, object], -) -> dict[str, object]: - return { - "model": spec.model, - "document": {"type": "image_url", "image_url": image_data_uri}, - **options, - } - - -def _upload_reducto_document( - spec: ProviderSpec, - sdk_kwargs: dict[str, object], - requests_only: bool, -) -> dict[str, object]: - if spec.name != "reducto": - return sdk_kwargs - if requests_only: - return {**sdk_kwargs, "document": {"type": "image_url", "image_url": "reducto://fixture"}} - if spec.api_key is None: - raise ValueError("Reducto response fixture generation requires REDUCTO_API_KEY") - document: Final = JSON_OBJECT.validate_python(sdk_kwargs["document"]) - image_data_uri: Final = document.get("image_url") - if not isinstance(image_data_uri, str): - raise ValueError("Reducto fixture generation requires an image_url data URI") - _, raw_bytes, mime = extract_file_id_or_bytes(image_data_uri, model=spec.model) - file_id: Final = upload_bytes_sync( - raw_bytes=raw_bytes or b"", - mime=mime, - api_key=spec.api_key, - api_base=spec.upstream_base, - ) - return {**sdk_kwargs, "document": {"type": "image_url", "image_url": file_id}} - - -def _generate_provider_case( - spec: ProviderSpec, - root: Path, - image_data_uri: str, - options: dict[str, object], - requests_only: bool, - sdk_call: Callable[..., object], -) -> None: - sdk_kwargs: Final = _upload_reducto_document(spec, _sdk_kwargs(spec, image_data_uri, options), requests_only) - result: Final = record_case(spec, root, sdk_kwargs, requests_only, sdk_call) - state: Final = "cached" if result.cache_hit else "recorded request" if requests_only else "recorded" - LOGGER.info("%s %s %s", state, spec.name, result.request.provider_request.path) - - -def _generate_provider_examples( - spec: ProviderSpec, - root: Path, - examples: int, - requests_only: bool, - sdk_call: Callable[..., object], -) -> None: - options_strategy: Final = _options_strategy(spec) - image_strategy: Final = ( - st.just(_REQUEST_ONLY_IMAGE) - if requests_only - else st.builds(_image_data_uri, _TEXT, st.integers(min_value=12, max_value=36)) - ) - - @settings(max_examples=examples, deadline=None, derandomize=True) - @given(image_data_uri=image_strategy, options=options_strategy) - def generate_case(image_data_uri: str, options: dict[str, object]) -> None: - _generate_provider_case(spec, root, image_data_uri, options, requests_only, sdk_call) - - generate_case() - - -def _options_strategy(spec: ProviderSpec) -> SearchStrategy[dict[str, object]]: - model: Final = spec.model.rsplit("/", 1)[-1] - if spec.name == "mistral": - return _mistral_options(spec.model) - if model == "parse-v3": - return _REDUCTO_V3_OPTIONS - if model == "parse-legacy": - return _REDUCTO_LEGACY_ROOT - raise ValueError(f"Unsupported Reducto OCR fixture model: {spec.model}") - - -def _generate( - specs: tuple[ProviderSpec, ...], - root: Path, - examples: int, - requests_only: bool, - sdk_call: Callable[..., object], -) -> None: - for spec in specs: - _generate_provider_examples(spec, root, examples, requests_only, sdk_call) - - -def _log_filled_responses(results: tuple[RecorderResult, ...]) -> None: - for result in results: - LOGGER.info("filled response %s %s", result.request.provider, result.request.provider_request.path) - - def main() -> None: logging.basicConfig(level=logging.INFO, format="%(message)s") load_dotenv() - args: Final = parse_generator_args(PROVIDER_NAMES) + args: Final = parse_generator_args() + api_key: Final = os.environ.get("MISTRAL_API_KEY") or os.environ.get("LITELLM_API_KEY") + if api_key is None: + raise SystemExit("MISTRAL_API_KEY is required") root: Final = fixture_directory( args.fixture_dir, os.environ.get(FIXTURE_DIR_ENV), Path(__file__).with_name(".fixtures"), ) - specs: Final = _provider_specs( - args.providers, - requests_only=args.requests_only, - all_models=args.responses_only, - ) - if not specs: - raise SystemExit("No selected provider has the required credentials") - ocr_call: Final = cast(Callable[..., object], litellm.ocr) - if args.responses_only: - _log_filled_responses(fill_missing_responses(specs, root, ocr_call)) - return - _generate(specs, root, args.examples, args.requests_only, ocr_call) + spec: Final = ProviderSpec(model=args.model, upstream_base=_mistral_upstream_base(), api_key=api_key) + use_litellm_rust(False, ocr=None, aocr=None) + _generate_examples(spec, root, args.examples, cast(Callable[..., object], litellm.ocr)) if __name__ == "__main__": diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 5cd7fee0df7..d2f02dc9786 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -471,7 +471,9 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): ocr_main._run_rust_ocr( prepared_request=build_prepared_request(api_key=None, timeout=None), - resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None, + resolve_api_key=lambda name: ( + "sk-from-vault" if name == "MISTRAL_API_KEY" else None + ), ) assert bridge.calls[0]["api_key"] == "sk-from-vault" @@ -578,7 +580,9 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): api_base=None, timeout=None, ), - resolve_api_key=lambda name: "https://azure.example.com" if name == "AZURE_AI_API_BASE" else None, + resolve_api_key=lambda name: ( + "https://azure.example.com" if name == "AZURE_AI_API_BASE" else None + ), ) assert bridge.calls[0]["api_base"] == "https://azure.example.com" @@ -596,7 +600,9 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): timeout=None, ), resolve_api_key=lambda name: ( - "https://document-intelligence.example.com" if name == "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT" else None + "https://document-intelligence.example.com" + if name == "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT" + else None ), ) @@ -656,67 +662,34 @@ def test_ocr_routes_to_rust_when_enabled(fake_bridge): assert call["optional_params"].get("include_image_base64") is True -def test_ocr_routes_reducto_parse_v3_file_id_to_rust(fake_bridge: RecordingBridge) -> None: - document: dict[str, object] = { - "type": "document_url", - "document_url": "reducto://fixture", - } - +@pytest.mark.parametrize( + ("model", "provider"), + ( + ("azure_ai/pixtral-12b-2409", "azure_ai"), + ("vertex_ai/mistral-ocr-2505", "vertex_ai"), + ), +) +def test_ocr_routes_supported_provider_to_rust( + fake_bridge: RecordingBridge, + model: str, + provider: str, +) -> None: + provider_kwargs: dict[str, object] = ( + {"vertex_project": "project-1", "vertex_location": "us-central1"} if provider == "vertex_ai" else {} + ) response = litellm.ocr( - model="reducto/parse-v3", - document=document, + model=model, + document=DOCUMENT, api_key="sk-test", + api_base="https://example.com", + **provider_kwargs, ) assert isinstance(response, OCRResponse) assert len(fake_bridge.calls) == 1 call = fake_bridge.calls[0] - assert call["model"] == "parse-v3" - assert call["document"] == document - assert call["custom_llm_provider"] == "reducto" - - -@pytest.mark.parametrize( - ("model", "document"), - ( - ("azure_ai/pixtral-12b-2409", DOCUMENT), - ("vertex_ai/mistral-ocr-2505", DOCUMENT), - ( - "reducto/parse-v3", - { - "type": "document_url", - "document_url": "data:application/pdf;base64,AA==", - }, - ), - ( - "reducto/parse-legacy", - {"type": "document_url", "document_url": "reducto://fixture"}, - ), - ), -) -def test_ocr_falls_back_to_python_for_unsupported_rust_case( - monkeypatch: pytest.MonkeyPatch, - fake_bridge: RecordingBridge, - model: str, - document: dict[str, object], -) -> None: - captured: dict[str, object] = {} - - def fake_handler_ocr(**kwargs: object) -> OCRResponse: - captured.update(kwargs) - return OCRResponse(pages=[], model=model, object="ocr") - - monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fake_handler_ocr) - response = litellm.ocr( - model=model, - document=document, - api_key="sk-test", - api_base="https://example.com", - ) - - assert isinstance(response, OCRResponse) - assert fake_bridge.calls == [] - assert captured["model"] in {"pixtral-12b-2409", "mistral-ocr-2505"} + assert call["model"] == model.rsplit("/", 1)[-1] + assert call["custom_llm_provider"] == provider def test_ocr_rust_path_converts_file_document_before_bridge(fake_bridge): @@ -858,6 +831,9 @@ def test_ocr_provider_configs_expose_api_key_env_vars(): assert BaseOCRConfig().get_api_key_env_var() is None assert MistralOCRConfig().get_api_key_env_var() == "MISTRAL_API_KEY" assert AzureAIOCRConfig().get_api_key_env_var() == "AZURE_AI_API_KEY" - assert AzureDocumentIntelligenceOCRConfig().get_api_key_env_var() == "AZURE_DOCUMENT_INTELLIGENCE_API_KEY" + assert ( + AzureDocumentIntelligenceOCRConfig().get_api_key_env_var() + == "AZURE_DOCUMENT_INTELLIGENCE_API_KEY" + ) assert VertexAIOCRConfig().get_api_key_env_var() == "VERTEX_AI_API_KEY" assert VertexAIDeepSeekOCRConfig().get_api_key_env_var() == "VERTEX_AI_API_KEY" diff --git a/tests/test_litellm/ocr/test_sdk_parity.py b/tests/test_litellm/ocr/test_sdk_parity.py index a8b154c8de8..aa68cfd0b0c 100644 --- a/tests/test_litellm/ocr/test_sdk_parity.py +++ b/tests/test_litellm/ocr/test_sdk_parity.py @@ -5,29 +5,19 @@ import sys from collections.abc import Callable, Coroutine from enum import Enum from pathlib import Path -from typing import Final, Protocol, cast +from typing import Final, cast import pytest -from pydantic import BaseModel, ConfigDict, TypeAdapter -from tests.test_litellm._fixture_recorder import Fixture, FixtureResponse, ProviderWireRequest +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 ( - CapturedRequest, - ExceptionReport, - NativeEvidence, - ParityTrace, - ReplayResponse, - SDKOutput, - SDKReport, -) -from tests.test_litellm.parity.replay import replay_json_response +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 API_KEY: Final = "test-key" PYTHON_HTTP_SENTINEL: Final = "python-ocr-parity-fallback" -GIL_STATS: Final = TypeAdapter(dict[str, int]) -SUPPORTED_BY_RUST_PROVIDERS: Final = frozenset({"mistral", "reducto"}) class SDKRoute(str, Enum): @@ -35,180 +25,72 @@ class SDKRoute(str, Enum): AOCR = "aocr" -class SDKInput(BaseModel): - model_config = ConfigDict(frozen=True) - - route: SDKRoute - kwargs: dict[str, object] - native_expected: bool - - -class _NativeBridge(Protocol): - ocr: object - aocr: object - - def gil_stats(self) -> object: ... - - -def _qualified_name(value: object) -> str: - value_type: Final = type(value) - return f"{value_type.__module__}.{value_type.__qualname__}" - - -def _exception_report(error: Exception) -> ExceptionReport: - raw_status_code: Final = getattr(error, "status_code", None) - status_code: Final = raw_status_code if isinstance(raw_status_code, int) else None - return ExceptionReport(class_name=_qualified_name(error), status_code=status_code, message=str(error)) - - -def _gil_release_count(native_bridge: _NativeBridge) -> int: - stats: Final = GIL_STATS.validate_python(native_bridge.gil_stats()) - return stats["releases"] - - -def _response_trace(response: BaseModel) -> ParityTrace: - output: Final = SDKOutput( - response_type=_qualified_name(response), - response_json=response.model_dump(mode="json"), - ) - return ParityTrace(outputs=(output,), exception=None) - - -def _capture_sdk_call(sdk_input: SDKInput, mock_url: str) -> ParityTrace: - import litellm - from litellm.llms.base_llm.ocr.transformation import OCRResponse - - call_kwargs: Final[dict[str, object]] = { +def _call_kwargs(sdk_input: MistralOcrParityInput, mock_url: str, route: SDKRoute) -> dict[str, object]: + return { + **sdk_input.as_sdk_kwargs(), "api_base": mock_url, "api_key": API_KEY, - "extra_headers": {"x-parity-case": sdk_input.route.value}, - **sdk_input.kwargs, + "extra_headers": {"x-ocr-parity-route": route.value}, } - try: - if sdk_input.route is SDKRoute.OCR: - sync_route: Final = cast(Callable[..., OCRResponse], litellm.ocr) - return _response_trace(sync_route(**call_kwargs)) - async_route: Final = cast(Callable[..., Coroutine[object, object, OCRResponse]], litellm.aocr) - return _response_trace(asyncio.run(async_route(**call_kwargs))) - except Exception as error: - return ParityTrace(outputs=(), exception=_exception_report(error)) +def _execute_sdk_case(sdk_input: MistralOcrParityInput, route: SDKRoute, mock_url: str) -> SDKReport: + import litellm -def _native_sdk_report(trace: ParityTrace, native_handled_case: bool) -> SDKReport: - return SDKReport( - trace=trace, - native=NativeEvidence( - rust_enabled=True, - native_callable_loaded=True, - native_handled_case=native_handled_case, - ), - ) - - -def _execute_sdk_case(sdk_input: SDKInput, mock_url: str) -> SDKReport: - from litellm.rust_bridge.loader import get_native_bridge - from litellm.rust_bridge.ocr import load_rust_aocr, load_rust_ocr, rust_ocr_enabled - - rust_enabled: Final = rust_ocr_enabled() - if not rust_enabled: - return SDKReport( - trace=_capture_sdk_call(sdk_input, mock_url), - native=NativeEvidence( - rust_enabled=False, - native_callable_loaded=False, - native_handled_case=False, - ), - ) - - raw_native_bridge: Final = get_native_bridge() - if raw_native_bridge is None: - raise RuntimeError("LITELLM_USE_RUST_OCR=1 but the native bridge is unavailable") - native_bridge: Final = cast(_NativeBridge, raw_native_bridge) - - if not sdk_input.native_expected: - return _native_sdk_report(_capture_sdk_call(sdk_input, mock_url), native_handled_case=False) - - native_callable: Final = load_rust_ocr() if sdk_input.route is SDKRoute.OCR else load_rust_aocr() - expected_native_callable: Final = native_bridge.ocr if sdk_input.route is SDKRoute.OCR else native_bridge.aocr - if native_callable is None or native_callable is not expected_native_callable: - raise RuntimeError(f"native {sdk_input.route.value} callable is unavailable") - - if sdk_input.route is SDKRoute.AOCR: - async_trace: Final = _capture_sdk_call(sdk_input, mock_url) - return _native_sdk_report(async_trace, native_handled_case=bool(async_trace.outputs)) - - before_gil_releases: Final = _gil_release_count(native_bridge) - trace: Final = _capture_sdk_call(sdk_input, mock_url) - after_gil_releases: Final = _gil_release_count(native_bridge) - return _native_sdk_report(trace, native_handled_case=after_gil_releases == before_gil_releases + 1) - - -def _replay_response(response: FixtureResponse) -> ReplayResponse: - return ReplayResponse(status_code=response.status_code, headers=response.headers, body=response.body) - - -def _provider_wire_request(request: CapturedRequest) -> ProviderWireRequest: - return ProviderWireRequest(method=request.method, path=request.path, body=request.body) - - -def _native_expected(ocr_fixture: Fixture) -> bool: - provider: Final = ocr_fixture.request.provider - if provider != "reducto": - return provider in SUPPORTED_BY_RUST_PROVIDERS - model: Final = ocr_fixture.request.sdk_kwargs.get("model") - document: Final = ocr_fixture.request.sdk_kwargs.get("document") - if model != "reducto/parse-v3" or not isinstance(document, dict): - return False - source: Final = document.get("document_url") or document.get("image_url") - return isinstance(source, str) and source.startswith("reducto://") + call_kwargs: Final = _call_kwargs(sdk_input, mock_url, route) + if route is SDKRoute.OCR: + sync_route: Final = cast(Callable[..., OCRResponse], litellm.ocr) + 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)) + 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: Fixture, route: SDKRoute, tmp_path: Path) -> None: - native_expected: Final = _native_expected(ocr_fixture) - sdk_input: Final = SDKInput(route=route, kwargs=ocr_fixture.request.sdk_kwargs, native_expected=native_expected) - case_file: Final = tmp_path / f"{ocr_fixture.request.provider}-{route.value}-sdk-input.json" - case_file.write_text(sdk_input.model_dump_json(indent=2), encoding="utf-8") - expected_request: Final = ocr_fixture.request.provider_request +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.upstream_response + response_body: Final = response.body_bytes() + response_headers: Final = tuple((header.name, header.value) for header in response.headers) runner: Final = PythonScriptRunner( entrypoint=Path(__file__), rust_env_var="LITELLM_USE_RUST_OCR", python_user_agent=PYTHON_HTTP_SENTINEL, ) - with ( - replay_json_response(expected_request.path, _replay_response(ocr_fixture.response)) as python_provider, - replay_json_response(expected_request.path, _replay_response(ocr_fixture.response)) as rust_provider, - ): + 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, ) - assert_parity(python, rust, PYTHON_HTTP_SENTINEL, native_expected) - assert tuple(_provider_wire_request(request) for request in python.requests) == (expected_request,) - assert tuple(_provider_wire_request(request) for request in rust.requests) == (expected_request,) + assert_parity(python, rust, PYTHON_HTTP_SENTINEL) def _child_main() -> None: - if len(sys.argv) != 4: - raise SystemExit("usage: test_sdk_parity.py CASE_FILE MOCK_URL REPORT_FILE") + 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]) - mock_url: Final = sys.argv[2] - report_file: Final = Path(sys.argv[3]) - sdk_input: Final = SDKInput.model_validate_json(case_file.read_text(encoding="utf-8")) - report: Final = _execute_sdk_case(sdk_input, mock_url) + 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.input, route, mock_url) report_file.write_text(report.model_dump_json(indent=2), encoding="utf-8") diff --git a/tests/test_litellm/parity/compare.py b/tests/test_litellm/parity/compare.py index 4ffd4fef822..89e07ba717c 100644 --- a/tests/test_litellm/parity/compare.py +++ b/tests/test_litellm/parity/compare.py @@ -1,89 +1,27 @@ from __future__ import annotations -from collections.abc import Mapping, Sequence -from typing import Final, cast +from typing import Final -from pydantic import BaseModel - -from tests.test_litellm.parity.models import CapturedRequest, Execution, ParityTrace +from tests.test_litellm.parity.models import CapturedRequest, Execution -def _stable_request(request: CapturedRequest) -> dict[str, object]: - return request.model_dump(mode="json", exclude={"user_agent"}) - - -def assert_parity( - python: Execution, - rust: Execution, - python_user_agent: str, - native_expected: bool, -) -> None: - assert python.report.native.rust_enabled is False - assert rust.report.native.rust_enabled is True - assert tuple(request.user_agent for request in python.requests) == (python_user_agent,) * len(python.requests) - if native_expected: - assert rust.report.native.native_callable_loaded is True - assert rust.report.native.native_handled_case is True - assert all(request.user_agent != python_user_agent for request in rust.requests) - else: - assert rust.report.native.native_handled_case is False - assert tuple(request.user_agent for request in rust.requests) == (python_user_agent,) * len(rust.requests) - assert tuple(_stable_request(request) for request in rust.requests) == tuple( - _stable_request(request) for request in python.requests - ) - trace_differences: Final = _diff(rust.report.trace, python.report.trace) - assert not trace_differences, "\n".join(trace_differences) - assert python.report.trace.exception is None - assert python.report.trace.outputs - - -def _short(value: object) -> str: - rendered: Final = repr(value) - return rendered if len(rendered) <= 240 else f"{rendered[:237]}..." - - -def _diff(left: object, right: object, path: str = "$") -> tuple[str, ...]: - if isinstance(left, BaseModel) and isinstance(right, BaseModel): - return _diff(left.model_dump(mode="json"), right.model_dump(mode="json"), path) - if isinstance(left, Mapping) and isinstance(right, Mapping): - left_mapping: Final = cast(Mapping[str, object], left) - right_mapping: Final = cast(Mapping[str, object], right) - left_keys: Final = frozenset(left_mapping) - right_keys: Final = frozenset(right_mapping) - missing: Final = tuple(f"{path}.{key}: missing from left" for key in sorted(right_keys - left_keys)) - extra: Final = tuple(f"{path}.{key}: missing from right" for key in sorted(left_keys - right_keys)) - shared: Final = tuple( - line - for key in sorted(left_keys & right_keys) - for line in _diff(left_mapping[key], right_mapping[key], f"{path}.{key}") +def validate_harness(python: Execution, rust: Execution, python_user_agent: str) -> None: + if python.request.user_agent != python_user_agent: + raise AssertionError( + f"Python provider request did not carry fallback sentinel user-agent {python_user_agent!r}: " + f"{python.request.user_agent!r}" ) - return (*missing, *extra, *shared) - if ( - isinstance(left, Sequence) - and not isinstance(left, (str, bytes)) - and isinstance(right, Sequence) - and not isinstance(right, (str, bytes)) - ): - left_sequence: Final = cast(Sequence[object], left) - right_sequence: Final = cast(Sequence[object], right) - length_diff: Final = ( - (f"{path}: lengths differ ({len(left_sequence)} != {len(right_sequence)})",) - if len(left_sequence) != len(right_sequence) - else () - ) - item_diff: Final = tuple( - line - for index, (left_item, right_item) in enumerate(zip(left_sequence, right_sequence)) - for line in _diff(left_item, right_item, f"{path}[{index}]") - ) - return (*length_diff, *item_diff) - left_value: Final = cast(object, left) - right_value: Final = cast(object, right) - return () if left_value == right_value else (f"{path}: {_short(left_value)} != {_short(right_value)}",) + if rust.request.user_agent == python_user_agent: + raise AssertionError("Rust OCR fell back to the Python HTTP implementation") -def parity_comparison(left: object, right: object) -> list[str] | None: - if not isinstance(left, (ParityTrace, CapturedRequest)) or type(left) is not type(right): - return None - differences: Final = _diff(left, right) - return [f"Comparing {type(left).__name__} values:", *(f" {line}" for line in differences)] +def _request_after_transformation(request: CapturedRequest) -> CapturedRequest: + return request.model_copy(update={"user_agent": None}) + + +def assert_parity(python: Execution, rust: Execution, python_user_agent: str) -> None: + validate_harness(python, rust, python_user_agent) + python_request: Final = _request_after_transformation(python.request) + rust_request: Final = _request_after_transformation(rust.request) + assert python_request == rust_request + assert python.report.response == rust.report.response diff --git a/tests/test_litellm/parity/models.py b/tests/test_litellm/parity/models.py index f68cc690c71..9cf26127765 100644 --- a/tests/test_litellm/parity/models.py +++ b/tests/test_litellm/parity/models.py @@ -1,43 +1,8 @@ from __future__ import annotations -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, JsonValue - -class ExceptionReport(BaseModel): - model_config = ConfigDict(frozen=True) - - class_name: str - status_code: int | None - message: str - - -class SDKOutput(BaseModel): - model_config = ConfigDict(frozen=True) - - response_type: str - response_json: dict[str, object] - - -class ParityTrace(BaseModel): - model_config = ConfigDict(frozen=True) - - outputs: tuple[SDKOutput, ...] - exception: ExceptionReport | None - - -class NativeEvidence(BaseModel): - model_config = ConfigDict(frozen=True) - - rust_enabled: bool - native_callable_loaded: bool - native_handled_case: bool - - -class SDKReport(BaseModel): - model_config = ConfigDict(frozen=True) - - trace: ParityTrace - native: NativeEvidence +from litellm.llms.base_llm.ocr.transformation import OCRResponse class CapturedRequest(BaseModel): @@ -45,23 +10,19 @@ class CapturedRequest(BaseModel): method: str path: str - authorization: str | None - content_type: str | None - parity_case: str | None - body: dict[str, object] + headers: tuple[tuple[str, str], ...] + body: JsonValue user_agent: str | None +class SDKReport(BaseModel): + model_config = ConfigDict(frozen=True) + + response: OCRResponse + + class Execution(BaseModel): model_config = ConfigDict(frozen=True) + request: CapturedRequest report: SDKReport - requests: tuple[CapturedRequest, ...] - - -class ReplayResponse(BaseModel): - model_config = ConfigDict(frozen=True) - - status_code: int - headers: dict[str, str] - body: dict[str, object] diff --git a/tests/test_litellm/parity/pytest_plugin.py b/tests/test_litellm/parity/pytest_plugin.py deleted file mode 100644 index 1d8782c4255..00000000000 --- a/tests/test_litellm/parity/pytest_plugin.py +++ /dev/null @@ -1,11 +0,0 @@ -from __future__ import annotations - -import pytest - -from tests.test_litellm.parity.compare import parity_comparison - - -def pytest_assertrepr_compare(config: pytest.Config, op: str, left: object, right: object) -> list[str] | None: - if op != "==": - return None - return parity_comparison(left, right) diff --git a/tests/test_litellm/parity/replay.py b/tests/test_litellm/parity/replay.py index 3f2998593ac..1b067530199 100644 --- a/tests/test_litellm/parity/replay.py +++ b/tests/test_litellm/parity/replay.py @@ -1,6 +1,5 @@ from __future__ import annotations -import json import queue import threading from collections.abc import Generator @@ -8,77 +7,88 @@ from contextlib import contextmanager from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import Final -from pydantic import TypeAdapter +from pydantic import JsonValue, TypeAdapter -from tests.test_litellm.parity.models import CapturedRequest, ReplayResponse +from tests.test_litellm.parity.models import CapturedRequest -JSON_OBJECT: Final = TypeAdapter(dict[str, object]) +JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +EXCLUDED_REQUEST_HEADERS: Final = frozenset( + { + "host", + "content-length", + "connection", + "accept-encoding", + "user-agent", + "x-ocr-parity-route", + } +) +EXCLUDED_RESPONSE_HEADERS: Final = frozenset({"content-length", "transfer-encoding", "connection"}) -class JsonReplayServer(ThreadingHTTPServer): +class ReplayServer(ThreadingHTTPServer): daemon_threads = True - def __init__(self, expected_path: str, response: ReplayResponse) -> None: - super().__init__(("127.0.0.1", 0), _JsonReplayHandler) - self.expected_path: Final = expected_path - self.response: Final = response - self.response_body: Final = json.dumps(response.body, sort_keys=True, separators=(",", ":")).encode() + def __init__(self, status_code: int, headers: tuple[tuple[str, str], ...], body: bytes) -> None: + super().__init__(("127.0.0.1", 0), _ReplayHandler) + self.status_code: Final = status_code + self.headers: Final = headers + self.body: Final = body self.requests: queue.Queue[CapturedRequest] = queue.Queue() @property def url(self) -> str: return f"http://127.0.0.1:{self.server_address[1]}" - def take_requests(self) -> tuple[CapturedRequest, ...]: - try: - first: Final = self.requests.get(timeout=5) - except queue.Empty as error: - raise AssertionError("expected at least one provider request, received none") from error - remaining_count: Final = self.requests.qsize() - remaining: Final = tuple(self.requests.get_nowait() for _ in range(remaining_count)) - return (first, *remaining) + 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() -class _JsonReplayHandler(BaseHTTPRequestHandler): +class _ReplayHandler(BaseHTTPRequestHandler): protocol_version = "HTTP/1.1" def do_POST(self) -> None: provider: Final = self.server - assert isinstance(provider, JsonReplayServer) + assert isinstance(provider, ReplayServer) length: Final = int(self.headers.get("content-length") or "0") - body: Final = JSON_OBJECT.validate_json(self.rfile.read(length)) - content_type_header: Final = self.headers.get("content-type") - content_type: Final = content_type_header.split(";", 1)[0].lower() if content_type_header else None + body: Final = JSON_VALUE.validate_json(self.rfile.read(length)) + headers: Final = tuple( + sorted( + (name.lower(), value) + for name, value in self.headers.raw_items() + if name.lower() not in EXCLUDED_REQUEST_HEADERS + ) + ) provider.requests.put( CapturedRequest( method=self.command, path=self.path, - authorization=self.headers.get("authorization"), - content_type=content_type, - parity_case=self.headers.get("x-parity-case"), + headers=headers, body=body, user_agent=self.headers.get("user-agent"), ) ) - matched: Final = self.path == provider.expected_path - status_code: Final = provider.response.status_code if matched else 404 - response_body: Final = provider.response_body if matched else b'{"error":"unexpected path"}' - response_headers: Final = provider.response.headers if matched else {"content-type": "application/json"} - self.send_response(status_code) - for name, value in response_headers.items(): - if name.lower() not in {"content-length", "transfer-encoding", "content-encoding"}: + self.send_response_only(provider.status_code) + for name, value in provider.headers: + if name.lower() not in EXCLUDED_RESPONSE_HEADERS: self.send_header(name, value) - self.send_header("content-length", str(len(response_body))) + self.send_header("content-length", str(len(provider.body))) self.end_headers() - self.wfile.write(response_body) + self.wfile.write(provider.body) def log_message(self, format: str, *args: object) -> None: return @contextmanager -def replay_json_response(expected_path: str, response: ReplayResponse) -> Generator[JsonReplayServer]: - server: Final = JsonReplayServer(expected_path=expected_path, response=response) +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) thread.start() try: diff --git a/tests/test_litellm/parity/runner.py b/tests/test_litellm/parity/runner.py index b771408e707..2fe0984f5d1 100644 --- a/tests/test_litellm/parity/runner.py +++ b/tests/test_litellm/parity/runner.py @@ -7,8 +7,10 @@ from dataclasses import dataclass from pathlib import Path from typing import Final +from pydantic import ValidationError + from tests.test_litellm.parity.models import Execution, SDKReport -from tests.test_litellm.parity.replay import JsonReplayServer +from tests.test_litellm.parity.replay import ReplayServer @dataclass(frozen=True, slots=True) @@ -17,11 +19,12 @@ class PythonScriptRunner: rust_env_var: str python_user_agent: str - def command(self, case_file: Path, provider_url: str, report_file: Path) -> tuple[str, ...]: + def command(self, case_file: Path, route: str, provider_url: str, report_file: Path) -> tuple[str, ...]: return ( sys.executable, str(self.entrypoint.resolve()), str(case_file), + route, provider_url, str(report_file), ) @@ -30,8 +33,9 @@ class PythonScriptRunner: def run_execution( runner: PythonScriptRunner, case_file: Path, + route: str, report_file: Path, - provider: JsonReplayServer, + provider: ReplayServer, rust_enabled: bool, ) -> Execution: project_root: Final = str(runner.entrypoint.resolve().parents[3]) @@ -40,22 +44,30 @@ def run_execution( **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 - ), + "PYTHONPATH": os.pathsep.join(path for path in (project_root, existing_pythonpath) if path), } - completed: Final = subprocess.run( - runner.command(case_file, provider.url, report_file), - capture_output=True, - text=True, - env=env, - timeout=60, - check=False, - ) - assert completed.returncode == 0, ( - f"SDK subprocess failed with exit code {completed.returncode}\n" - f"stdout:\n{completed.stdout}\n" - f"stderr:\n{completed.stderr}" - ) - report: Final = SDKReport.model_validate_json(report_file.read_text(encoding="utf-8")) - return Execution(report=report, requests=provider.take_requests()) + 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) diff --git a/tests/test_litellm/parity/test_parity.py b/tests/test_litellm/parity/test_parity.py index af7e0f95038..c67938c366f 100644 --- a/tests/test_litellm/parity/test_parity.py +++ b/tests/test_litellm/parity/test_parity.py @@ -2,24 +2,54 @@ from __future__ import annotations from typing import Final -from tests.test_litellm.parity.compare import parity_comparison -from tests.test_litellm.parity.models import ParityTrace, SDKOutput +import pytest +from pydantic import JsonValue + +from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse +from tests.test_litellm.parity.compare import assert_parity +from tests.test_litellm.parity.models import CapturedRequest, Execution, SDKReport + +SENTINEL: Final = "python-ocr-parity-fallback" -def test_parity_comparison_reports_first_nested_output_difference() -> None: - python: Final = ParityTrace( - outputs=(SDKOutput(response_type="Chunk", response_json={"choices": [{"delta": {"content": "a"}}]}),), - exception=None, - ) - rust: Final = ParityTrace( - outputs=(SDKOutput(response_type="Chunk", response_json={"choices": [{"delta": {"content": "b"}}]}),), - exception=None, +def _execution(*, body: JsonValue = None, markdown: str = "same", user_agent: str | None = None) -> Execution: + return Execution( + request=CapturedRequest( + method="POST", + path="/v1/ocr?mode=test", + headers=(("authorization", "Bearer test-key"), ("content-type", "application/json")), + body={"model": "mistral-ocr-latest"} if body is None else body, + user_agent=user_agent, + ), + report=SDKReport( + response=OCRResponse( + pages=[OCRPage(index=0, markdown=markdown)], + model="mistral-ocr-latest", + object="ocr", + ) + ), ) - explanation: Final = parity_comparison(rust, python) - assert explanation is not None - assert explanation == [ - "Comparing ParityTrace values:", - " $.outputs[0].response_json.choices[0].delta.content: 'b' != 'a'", - ] +def test_parity_rejects_request_difference() -> None: + python: Final = _execution(user_agent=SENTINEL) + rust: Final = _execution(body={"model": "different"}, user_agent="litellm-rust") + + with pytest.raises(AssertionError): + assert_parity(python, rust, SENTINEL) + + +def test_parity_rejects_response_difference() -> None: + python: Final = _execution(user_agent=SENTINEL) + rust: Final = _execution(markdown="different", user_agent="litellm-rust") + + with pytest.raises(AssertionError): + assert_parity(python, rust, SENTINEL) + + +def test_parity_rejects_rust_fallback() -> None: + python: Final = _execution(user_agent=SENTINEL) + rust: Final = _execution(user_agent=SENTINEL) + + with pytest.raises(AssertionError, match="fell back"): + assert_parity(python, rust, SENTINEL) diff --git a/tests/test_litellm/test_fixture_recorder.py b/tests/test_litellm/test_fixture_recorder.py index 70f12809b35..8515ba811af 100644 --- a/tests/test_litellm/test_fixture_recorder.py +++ b/tests/test_litellm/test_fixture_recorder.py @@ -1,106 +1,174 @@ from __future__ import annotations import json -from collections.abc import Callable +import threading +from collections.abc import Callable, Generator +from contextlib import contextmanager +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from typing import Final, cast import httpx +import pytest +from pydantic import ValidationError -from tests.test_litellm._fixture_recorder import ( - ProviderSpec, - pending_requests, - record_case, -) +from litellm.llms.mistral.ocr.transformation import MistralOCRConfig +from tests.test_litellm._fixture_recorder import ProviderSpec, fixture_cache_key, record_case, recorded_fixtures from tests.test_litellm._json_fs_cache import JsonFileCache +from tests.test_litellm.ocr.fixture_models import ( + DocumentUrlDocument, + ImageUrlDocument, + MistralOcrParityInput, + OcrParityCase, +) + +_UPSTREAM_BODY: Final = b'{"pages":[],"model":"mistral-ocr-latest","usage_info":{"pages_processed":0}}\n' -def _request_value(model: str = "test/model") -> dict[str, object]: - return { - "request": { - "provider": "test-provider", - "sdk_kwargs": {"model": model}, - "provider_request": {"method": "POST", "path": "/v1/test", "body": {"model": model}}, - } - } +class _Upstream(ThreadingHTTPServer): + daemon_threads = True + + def __init__(self) -> None: + super().__init__(("127.0.0.1", 0), _UpstreamHandler) + self.requests: list[tuple[tuple[tuple[str, str], ...], bytes]] = [] + + @property + def url(self) -> str: + return f"http://127.0.0.1:{self.server_address[1]}" + + +class _UpstreamHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def do_POST(self) -> None: + upstream: Final = self.server + assert isinstance(upstream, _Upstream) + body: Final = self.rfile.read(int(self.headers.get("content-length") or "0")) + upstream.requests.append((tuple(self.headers.raw_items()), body)) + self.send_response_only(200) + self.send_header("content-type", "application/json") + self.send_header("set-cookie", "first=1") + self.send_header("set-cookie", "second=2") + self.send_header("connection", "keep-alive") + self.send_header("content-length", str(len(_UPSTREAM_BODY))) + self.end_headers() + self.wfile.write(_UPSTREAM_BODY) + + def log_message(self, format: str, *args: object) -> None: + return + + +@contextmanager +def _upstream() -> Generator[_Upstream]: + server: Final = _Upstream() + thread: Final = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + yield server + finally: + server.shutdown() + server.server_close() + thread.join(timeout=5) + + +def _case_input(model: str = "mistral/mistral-ocr-latest") -> MistralOcrParityInput: + return MistralOcrParityInput( + model=model, + document=ImageUrlDocument(type="image_url", image_url="https://example.test/image.png"), + ) def _sdk_call(**kwargs: object) -> None: api_base: Final = kwargs["api_base"] - model: Final = kwargs["model"] assert isinstance(api_base, str) - assert isinstance(model, str) + raw_body: Final = b'{ "document" : {"type":"image_url"}, "model" : "wire-model" }' response: Final = httpx.post( - f"{api_base}/v1/test", - content=json.dumps({"model": model}).encode(), - headers={"content-type": "application/json"}, + f"{api_base}/v1/ocr", + content=raw_body, + headers=[("content-type", "application/json"), ("x-repeat", "one"), ("x-repeat", "two")], ) + assert response.content == _UPSTREAM_BODY + assert response.headers.get_list("set-cookie") == ["first=1", "second=2"] response.raise_for_status() -def _same_wire_sdk_call(**kwargs: object) -> None: - api_base: Final = kwargs["api_base"] - assert isinstance(api_base, str) - response: Final = httpx.post( - f"{api_base}/v1/test", - content=b'{"constant":true}', - headers={"content-type": "application/json"}, +def test_case_schema_has_exact_fields_rejects_extras_and_preserves_null() -> None: + case_input: Final = MistralOcrParityInput( + model="mistral/mistral-ocr-latest", + document=DocumentUrlDocument(type="document_url", document_url="https://example.test/file.pdf"), + pages=None, ) - response.raise_for_status() - - -def test_request_only_recording_persists_pending_fixture_without_calling_upstream(tmp_path: Path) -> None: - spec: Final = ProviderSpec( - name="test-provider", - model="test/model", - upstream_base="http://127.0.0.1:1", - api_key="test-key", - ) - sdk_kwargs: Final[dict[str, object]] = {"model": spec.model} - - result: Final = record_case( - spec, - tmp_path, - sdk_kwargs, - requests_only=True, - sdk_call=cast(Callable[..., object], _sdk_call), - ) - values: Final = JsonFileCache(tmp_path / spec.name).values() - - assert result.response is None - assert len(values) == 1 - assert values[0] == {"request": result.request.model_dump(mode="json")} - - -def test_pending_requests_excludes_completed_fixtures(tmp_path: Path) -> None: - cache: Final = JsonFileCache(tmp_path) - pending: Final = _request_value() - response: Final[dict[str, object]] = {"status_code": 200, "headers": {}, "body": {}} - completed: Final[dict[str, object]] = { - **_request_value("test/completed"), - "response": response, + raw_case: Final[dict[str, object]] = { + "input": case_input.model_dump(mode="json", exclude_unset=True), + "upstream_response": {"kind": "http", "status_code": 200, "headers": [], "body_b64": "e30="}, } - cache.put({"case": "pending"}, pending) - cache.put({"case": "completed"}, completed) - requests: Final = pending_requests(cache) + case: Final = OcrParityCase.model_validate(raw_case) - assert len(requests) == 1 - assert requests[0].sdk_kwargs["model"] == "test/model" + assert set(case.model_dump(mode="json", exclude_unset=True)) == {"input", "upstream_response"} + assert case.input.as_sdk_kwargs()["pages"] is None + assert "include_blocks" not in case.input.as_sdk_kwargs() + with pytest.raises(ValidationError): + OcrParityCase.model_validate({**raw_case, "provider_request": {}}) + with pytest.raises(ValidationError): + MistralOcrParityInput.model_validate({**case_input.canonical_input(), "unknown": True}) -def test_request_cache_distinguishes_sdk_inputs_with_identical_wire_requests(tmp_path: Path) -> None: - spec: Final = ProviderSpec( - name="test-provider", - model="test/model-a", - upstream_base="http://127.0.0.1:1", - api_key="test-key", +def test_typed_input_fields_match_supported_mistral_params() -> None: + input_fields: Final = frozenset(MistralOcrParityInput.model_fields) - {"model", "document"} + + supported_params: Final = cast( + list[str], + MistralOCRConfig().get_supported_ocr_params( # pyright: ignore[reportUnknownMemberType] # legacy API returns an unparameterized list + model="mistral-ocr-latest" + ), ) - sdk_call: Final = cast(Callable[..., object], _same_wire_sdk_call) - first: Final = record_case(spec, tmp_path, {"model": "test/model-a"}, True, sdk_call) - second: Final = record_case(spec, tmp_path, {"model": "test/model-b"}, True, sdk_call) + assert input_fields == frozenset(supported_params) - assert not first.cache_hit - assert not second.cache_hit - assert len(JsonFileCache(tmp_path / spec.name).values()) == 2 + +def test_record_case_proxies_raw_request_and_roundtrips_response_bytes_and_headers(tmp_path: Path) -> None: + with _upstream() as upstream: + spec: Final = ProviderSpec( + model="mistral/mistral-ocr-latest", + upstream_base=upstream.url, + api_key="test-key", + ) + result: Final = record_case(spec, tmp_path, _case_input(), cast(Callable[..., object], _sdk_call)) + + assert not result.cache_hit + assert result.case.upstream_response.body_bytes() == _UPSTREAM_BODY + assert tuple((header.name, header.value) for header in result.case.upstream_response.headers) == ( + ("content-type", "application/json"), + ("set-cookie", "first=1"), + ("set-cookie", "second=2"), + ) + assert len(upstream.requests) == 1 + request_headers, request_body = upstream.requests[0] + assert request_body == b'{ "document" : {"type":"image_url"}, "model" : "wire-model" }' + assert tuple(value for name, value in request_headers if name.lower() == "x-repeat") == ("one", "two") + stored: Final = JsonFileCache(tmp_path).values() + assert len(stored) == 1 + assert set(stored[0]) == {"input", "upstream_response"} + assert "provider_request" not in json.dumps(stored[0]) + assert recorded_fixtures(tmp_path) == (result.case,) + + +def test_cache_identity_uses_only_canonical_unified_input(tmp_path: Path) -> None: + case_input: Final = _case_input() + key: Final = fixture_cache_key(case_input) + assert key == case_input.model_dump(mode="json", exclude_unset=True) + + cache: Final = JsonFileCache(tmp_path) + cache.put( + key, + { + "input": key, + "upstream_response": {"kind": "http", "status_code": 200, "headers": [], "body_b64": "e30="}, + }, + ) + unreachable: Final = ProviderSpec(model=case_input.model, upstream_base="http://127.0.0.1:1", api_key="key") + + result: Final = record_case(unreachable, tmp_path, case_input, cast(Callable[..., object], _sdk_call)) + + assert result.cache_hit diff --git a/tests/test_litellm/test_json_fs_cache.py b/tests/test_litellm/test_json_fs_cache.py index f2fc7c8d628..eb1232320f5 100644 --- a/tests/test_litellm/test_json_fs_cache.py +++ b/tests/test_litellm/test_json_fs_cache.py @@ -15,3 +15,4 @@ def test_json_file_cache_is_content_addressed_and_recursive(tmp_path: Path) -> N assert stored_path.name == cache.path_for(reordered_key).name assert cache.get(reordered_key) == value assert JsonFileCache(tmp_path).values() == (value,) + assert tuple((tmp_path / "provider").glob("*.tmp")) == ()