This commit is contained in:
Yujong Lee 2026-08-29 14:06:59 -07:00 • committed by GitHub
parent 14790ef4f1
commit 8e0c2c60c4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 644 additions and 1211 deletions

View file

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

View file

@ -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], ...]:

View file

@ -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",
(

View file

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

View file

@ -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__":

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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")) == ()