mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
wip
This commit is contained in:
parent
14790ef4f1
commit
8e0c2c60c4
15 changed files with 644 additions and 1211 deletions
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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], ...]:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
(
|
||||
|
|
|
|||
91
tests/test_litellm/ocr/fixture_models.py
Normal file
91
tests/test_litellm/ocr/fixture_models.py
Normal 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
|
||||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")) == ()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue