From c41b740d489f805e8299ddfa1f756c237d8f59d0 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Tue, 1 Sep 2026 14:41:44 -0700 Subject: [PATCH] test(ocr): expand fixture model coverage --- .../vertex_ai/ocr/deepseek_transformation.py | 3 +- tests/route_parity/README.md | 15 +- tests/route_parity/fixtures/store.py | 13 +- tests/test_litellm/ocr/conftest.py | 16 +- tests/test_litellm/ocr/fixtures/azure.py | 66 ++++---- tests/test_litellm/ocr/fixtures/mistral.py | 26 ++- tests/test_litellm/ocr/fixtures/reducto.py | 9 +- tests/test_litellm/ocr/fixtures/vertex.py | 34 +++- tests/test_litellm/ocr/test_fixture_models.py | 150 ++++++++++++++++-- tests/test_litellm/ocr/test_fixture_store.py | 57 +++++++ .../test_litellm/ocr/test_record_fixtures.py | 105 +++++++++--- tests/test_litellm/ocr/test_sdk_parity.py | 2 - 12 files changed, 405 insertions(+), 91 deletions(-) create mode 100644 tests/test_litellm/ocr/test_fixture_store.py diff --git a/litellm/llms/vertex_ai/ocr/deepseek_transformation.py b/litellm/llms/vertex_ai/ocr/deepseek_transformation.py index 2603552152d..b57a87c3325 100644 --- a/litellm/llms/vertex_ai/ocr/deepseek_transformation.py +++ b/litellm/llms/vertex_ai/ocr/deepseek_transformation.py @@ -177,8 +177,9 @@ class VertexAIDeepSeekOCRConfig(BaseOCRConfig): content_item = {"type": "image_url", "image_url": document_url} # Build DeepSeek OCR request + provider_model: Final = model if model.startswith("deepseek-ai/") else f"deepseek-ai/{model}" data: Final = { - "model": "deepseek-ai/" + model, + "model": provider_model, "messages": [{"role": "user", "content": [content_item]}], } diff --git a/tests/route_parity/README.md b/tests/route_parity/README.md index 2bba0be66cd..98a141a832d 100644 --- a/tests/route_parity/README.md +++ b/tests/route_parity/README.md @@ -61,13 +61,20 @@ OCR strategies generate only public `litellm.ocr()` and `litellm.aocr()` inputs. provider wire payloads. Each boundary has a required corpus containing a baseline and one case for every supported top-level LiteLLM OCR -parameter. `--examples` controls additional Hypothesis-generated cases; it does not replace the required corpus. +parameter for every active registered model that uses that transformation. Models whose registry deprecation date has +passed are excluded. `--examples` controls additional Hypothesis-generated cases; it does not replace the required +corpus. The explicit boundaries are Mistral, Azure-hosted Mistral, Vertex-hosted Mistral, Azure Document Intelligence, Vertex DeepSeek, Reducto v3, and Reducto legacy. Provider credentials and endpoints only control target discovery, so -a machine records the boundaries it has configured and skips the rest. Reducto fixtures record both upload and parse -responses, but Reducto remains outside Python/Rust parity until the Rust OCR bridge supports it. -Azure-hosted Mistral discovery also requires `AZURE_AI_OCR_MODEL`, the deployment name or full `azure_ai/...` model. +a machine records the boundaries it has configured and skips the rest. Azure-hosted Mistral enumerates its active +registered models rather than requiring a separately configured deployment model. Reducto fixtures record both upload +and parse responses. Their parity cases are non-strict expected failures until the Rust OCR bridge supports Reducto, so +both expected failures and unexpected passes keep CI green during the rollout. + +The committed corpus does not need to contain live recordings for every configured target. In particular, Azure and +Vertex generation paths are covered by unit tests without requiring their credentials in CI. Recordings can be added +later without changing the fixture schema or runner. Invalid OCR inputs do not use recorded provider responses. The parity suite checks unsupported providers and models, malformed documents, invalid request formats, invalid Azure Document Intelligence parameters, and invalid headers in diff --git a/tests/route_parity/fixtures/store.py b/tests/route_parity/fixtures/store.py index af6e7417c65..41f488d45d6 100644 --- a/tests/route_parity/fixtures/store.py +++ b/tests/route_parity/fixtures/store.py @@ -110,6 +110,7 @@ def parametrize_recorded_fixtures( default_directory: Path, regeneration_command: str, id_builder: Callable[[CaseT], str], + marks_builder: Callable[[CaseT], tuple[pytest.MarkDecorator, ...]] | None = None, ) -> None: if fixture_name not in metafunc.fixturenames: return @@ -127,7 +128,17 @@ def parametrize_recorded_fixtures( f"Validation details: {error}" ) from error if fixtures: - metafunc.parametrize(fixture_name, fixtures, ids=tuple(id_builder(fixture) for fixture in fixtures)) + metafunc.parametrize( + fixture_name, + tuple( + pytest.param( + fixture, + id=id_builder(fixture), + marks=marks_builder(fixture) if marks_builder is not None else (), + ) + for fixture in fixtures + ), + ) return if configured is not None: raise pytest.UsageError(f"no recorded fixtures in {directory}") diff --git a/tests/test_litellm/ocr/conftest.py b/tests/test_litellm/ocr/conftest.py index d99c0de7cf9..73a2c12e0c3 100644 --- a/tests/test_litellm/ocr/conftest.py +++ b/tests/test_litellm/ocr/conftest.py @@ -11,13 +11,24 @@ from tests.test_litellm.ocr.fixtures.models import OcrParityCase FIXTURE_DIR_ENV: Final = "LITELLM_OCR_FIXTURE_DIR" -def _fixture_id(fixture: OcrParityCase) -> str: +def ocr_fixture_id(fixture: OcrParityCase) -> str: case_input: Final = fixture.litellm_input provider: Final = case_input.custom_llm_provider prefix: Final = f"{provider}/{case_input.model}" if provider else case_input.model return fixture_id(case_input, prefix) +def ocr_fixture_marks(fixture: OcrParityCase) -> tuple[pytest.MarkDecorator, ...]: + if fixture.litellm_input.boundary not in {"reducto_v3", "reducto_legacy"}: + return () + return ( + pytest.mark.xfail( + reason="Reducto does not have a Rust OCR boundary", + strict=False, + ), + ) + + def pytest_generate_tests(metafunc: pytest.Metafunc) -> None: default_directory: Final = Path(__file__).with_name("fixtures") / "data" parametrize_recorded_fixtures( @@ -29,5 +40,6 @@ def pytest_generate_tests(metafunc: pytest.Metafunc) -> None: regeneration_command=( f"uv run python -m tests.test_litellm.ocr.fixtures.record --fixture-dir {default_directory}" ), - id_builder=_fixture_id, + id_builder=ocr_fixture_id, + marks_builder=ocr_fixture_marks, ) diff --git a/tests/test_litellm/ocr/fixtures/azure.py b/tests/test_litellm/ocr/fixtures/azure.py index e662d839fa2..b7ea1eca3a5 100644 --- a/tests/test_litellm/ocr/fixtures/azure.py +++ b/tests/test_litellm/ocr/fixtures/azure.py @@ -24,10 +24,24 @@ from tests.test_litellm.ocr.fixtures.mistral import ( required_mistral_inputs, ) +AzureMistralModel = Literal["azure_ai/mistral-document-ai-2512",] +AzureDocumentIntelligenceModel = Literal[ + "azure_ai/doc-intelligence/prebuilt-read", + "azure_ai/doc-intelligence/prebuilt-layout", + "azure_ai/doc-intelligence/prebuilt-document", +] + +AZURE_MISTRAL_MODELS: Final[tuple[AzureMistralModel, ...]] = ("azure_ai/mistral-document-ai-2512",) +AZURE_DOCUMENT_INTELLIGENCE_MODELS: Final[tuple[AzureDocumentIntelligenceModel, ...]] = ( + "azure_ai/doc-intelligence/prebuilt-read", + "azure_ai/doc-intelligence/prebuilt-layout", + "azure_ai/doc-intelligence/prebuilt-document", +) + class AzureMistralOcrSdkInput(MistralCompatibleOcrSdkInput): boundary: str = Field(default="azure_mistral", pattern=r"^azure_mistral$") - model: str + model: AzureMistralModel custom_llm_provider: Literal["azure_ai"] | None = None @field_validator("model") @@ -40,11 +54,7 @@ class AzureMistralOcrSdkInput(MistralCompatibleOcrSdkInput): class AzureDocumentIntelligenceOcrSdkInput(OcrSdkInputBase): boundary: str = Field(default="azure_document_intelligence", pattern=r"^azure_document_intelligence$") - model: Literal[ - "azure_ai/doc-intelligence/prebuilt-read", - "azure_ai/doc-intelligence/prebuilt-layout", - "azure_ai/doc-intelligence/prebuilt-document", - ] + model: AzureDocumentIntelligenceModel document: OcrDocument custom_llm_provider: Literal["azure_ai"] | None = None pages: str | list[int] | None = None @@ -52,19 +62,23 @@ class AzureDocumentIntelligenceOcrSdkInput(OcrSdkInputBase): req_format: Literal["litellm"] = "litellm" -def _as_azure_mistral(case_input: MistralOcrSdkInput, model: str) -> AzureMistralOcrSdkInput: +def _as_azure_mistral(case_input: MistralOcrSdkInput, model: AzureMistralModel) -> AzureMistralOcrSdkInput: values: Final = case_input.model_dump(mode="python", exclude={"boundary", "model", "custom_llm_provider"}) return AzureMistralOcrSdkInput.model_validate({**values, "model": model}) def _required_document_intelligence_inputs() -> tuple[AzureDocumentIntelligenceOcrSdkInput, ...]: document: Final = pdf_document() - model: Final = "azure_ai/doc-intelligence/prebuilt-layout" - return ( - AzureDocumentIntelligenceOcrSdkInput(model=model, document=document), - AzureDocumentIntelligenceOcrSdkInput(model=model, document=document, pages=[0, 1]), - AzureDocumentIntelligenceOcrSdkInput(model=model, document=document, features=["languages"]), - AzureDocumentIntelligenceOcrSdkInput(model=model, document=document, req_format="litellm"), + cases: Final[tuple[dict[str, object], ...]] = ( + {}, + {"pages": [0, 1]}, + {"features": ["languages"]}, + {"req_format": "litellm"}, + ) + return tuple( + AzureDocumentIntelligenceOcrSdkInput.model_validate({"model": model, "document": document, **case}) + for model in AZURE_DOCUMENT_INTELLIGENCE_MODELS + for case in cases ) @@ -81,15 +95,7 @@ def azure_document_intelligence_input_strategy(draw: DrawFn) -> AzureDocumentInt ) return AzureDocumentIntelligenceOcrSdkInput.model_validate( { - "model": draw( - st.sampled_from( - ( - "azure_ai/doc-intelligence/prebuilt-read", - "azure_ai/doc-intelligence/prebuilt-layout", - "azure_ai/doc-intelligence/prebuilt-document", - ) - ) - ), + "model": draw(st.sampled_from(AZURE_DOCUMENT_INTELLIGENCE_MODELS)), "document": draw(public_document_strategy()), **optional_params, } @@ -101,22 +107,28 @@ def azure_mistral_recording_targets( ) -> tuple[OcrRecordingTarget, ...]: api_key: Final = environ.get("AZURE_AI_API_KEY") upstream_base: Final = environ.get("AZURE_AI_API_BASE") - configured_model: Final = environ.get("AZURE_AI_OCR_MODEL") - if not api_key or not upstream_base or not configured_model: + if not api_key or not upstream_base: return () - model: Final = configured_model if configured_model.startswith("azure_ai/") else f"azure_ai/{configured_model}" return ( OcrRecordingTarget( name="azure-mistral", provider_spec=ProviderSpec(upstream_base=upstream_base.rstrip("/")), strategy=cast( SearchStrategy[OcrSdkInputBase], - mistral_input_strategy(MISTRAL_MODEL).map(lambda case_input: _as_azure_mistral(case_input, model)), + st.sampled_from(AZURE_MISTRAL_MODELS).flatmap( + lambda model: mistral_input_strategy(MISTRAL_MODEL).map( + lambda case_input: _as_azure_mistral(case_input, model) + ) + ), ), invocation=invoke_with_api_key(client, api_key), required_inputs=cast( tuple[OcrSdkInputBase, ...], - tuple(_as_azure_mistral(case_input, model) for case_input in required_mistral_inputs(MISTRAL_MODEL)), + tuple( + _as_azure_mistral(case_input, model) + for model in AZURE_MISTRAL_MODELS + for case_input in required_mistral_inputs(MISTRAL_MODEL) + ), ), ), ) diff --git a/tests/test_litellm/ocr/fixtures/mistral.py b/tests/test_litellm/ocr/fixtures/mistral.py index 3bf1e82c3e5..ca05baa34e2 100644 --- a/tests/test_litellm/ocr/fixtures/mistral.py +++ b/tests/test_litellm/ocr/fixtures/mistral.py @@ -24,11 +24,15 @@ from tests.test_litellm.ocr.fixtures.common import ( ) MistralModel = Literal[ + "mistral/mistral-ocr-3", + "mistral/mistral-ocr-3-0", "mistral/mistral-ocr-2512", "mistral/mistral-ocr-4-0", "mistral/mistral-ocr-4-1", "mistral/mistral-ocr-4", "mistral/mistral-ocr-latest", + "mistral-ocr-3", + "mistral-ocr-3-0", "mistral-ocr-2512", "mistral-ocr-4-0", "mistral-ocr-4-1", @@ -36,6 +40,16 @@ MistralModel = Literal[ "mistral-ocr-latest", ] +MISTRAL_MODELS: Final[tuple[MistralModel, ...]] = ( + "mistral/mistral-ocr-3", + "mistral/mistral-ocr-3-0", + "mistral/mistral-ocr-2512", + "mistral/mistral-ocr-4", + "mistral/mistral-ocr-4-0", + "mistral/mistral-ocr-4-1", + "mistral/mistral-ocr-latest", +) + class MistralCompatibleOcrSdkInput(OcrSdkInputBase): document: OcrDocument @@ -72,7 +86,7 @@ class MistralOcrSdkInput(MistralCompatibleOcrSdkInput): return self -MISTRAL_MODEL: Final = "mistral/mistral-ocr-latest" +MISTRAL_MODEL: Final[MistralModel] = "mistral/mistral-ocr-latest" _VALUE_TEXT: Final = st.just("case-1") @@ -145,8 +159,14 @@ def mistral_recording_targets(environ: Mapping[str, str], client: OcrFixtureClie OcrRecordingTarget( name="mistral-ocr", provider_spec=ProviderSpec(upstream_base=upstream_base), - strategy=cast(SearchStrategy[OcrSdkInputBase], mistral_input_strategy(MISTRAL_MODEL)), + strategy=cast( + SearchStrategy[OcrSdkInputBase], + st.sampled_from(MISTRAL_MODELS).flatmap(mistral_input_strategy), + ), invocation=invoke_with_api_key(client, api_key), - required_inputs=cast(tuple[OcrSdkInputBase, ...], required_mistral_inputs(MISTRAL_MODEL)), + required_inputs=cast( + tuple[OcrSdkInputBase, ...], + tuple(case_input for model in MISTRAL_MODELS for case_input in required_mistral_inputs(model)), + ), ), ) diff --git a/tests/test_litellm/ocr/fixtures/reducto.py b/tests/test_litellm/ocr/fixtures/reducto.py index b2b0e9b591f..66e1d2650bc 100644 --- a/tests/test_litellm/ocr/fixtures/reducto.py +++ b/tests/test_litellm/ocr/fixtures/reducto.py @@ -151,6 +151,11 @@ class ReductoHybridVpcSettings(FixtureModel): ReductoPageSelection = ReductoPageRange | list[ReductoPageRange] | list[int] | list[str] +ReductoV3Model = Literal["reducto/parse-v3", "parse-v3"] +ReductoLegacyModel = Literal["reducto/parse-legacy", "parse-legacy"] + +REDUCTO_V3_MODELS: Final[tuple[Literal["reducto/parse-v3"], ...]] = ("reducto/parse-v3",) +REDUCTO_LEGACY_MODELS: Final[tuple[Literal["reducto/parse-legacy"], ...]] = ("reducto/parse-legacy",) class ReductoSettings(FixtureModel): @@ -181,7 +186,7 @@ class ReductoSettings(FixtureModel): class ReductoParseV3SdkInput(OcrSdkInputBase): boundary: str = Field(default="reducto_v3", pattern=r"^reducto_v3$") - model: Literal["reducto/parse-v3", "parse-v3"] + model: ReductoV3Model document: ReductoDocument custom_llm_provider: Literal["reducto"] | None = None formatting: ReductoFormatting = Field(default_factory=ReductoFormatting) @@ -197,7 +202,7 @@ class ReductoParseV3SdkInput(OcrSdkInputBase): class ReductoParseLegacySdkInput(OcrSdkInputBase): boundary: str = Field(default="reducto_legacy", pattern=r"^reducto_legacy$") - model: Literal["reducto/parse-legacy", "parse-legacy"] + model: ReductoLegacyModel document: ReductoDocument custom_llm_provider: Literal["reducto"] | None = None enhance: JsonObject | None = None diff --git a/tests/test_litellm/ocr/fixtures/vertex.py b/tests/test_litellm/ocr/fixtures/vertex.py index 64e592f4f86..b3495c1a255 100644 --- a/tests/test_litellm/ocr/fixtures/vertex.py +++ b/tests/test_litellm/ocr/fixtures/vertex.py @@ -24,10 +24,16 @@ from tests.test_litellm.ocr.fixtures.mistral import ( required_mistral_inputs, ) +VertexMistralModel = Literal["vertex_ai/mistral-ocr-2505"] +VertexDeepSeekModel = Literal["vertex_ai/deepseek-ai/deepseek-ocr-maas"] + +VERTEX_MISTRAL_MODELS: Final[tuple[VertexMistralModel, ...]] = ("vertex_ai/mistral-ocr-2505",) +VERTEX_DEEPSEEK_MODELS: Final[tuple[VertexDeepSeekModel, ...]] = ("vertex_ai/deepseek-ai/deepseek-ocr-maas",) + class VertexMistralOcrSdkInput(MistralCompatibleOcrSdkInput): boundary: str = Field(default="vertex_mistral", pattern=r"^vertex_mistral$") - model: Literal["vertex_ai/mistral-ocr-2505"] = "vertex_ai/mistral-ocr-2505" + model: VertexMistralModel = "vertex_ai/mistral-ocr-2505" custom_llm_provider: Literal["vertex_ai"] | None = None vertex_project: str vertex_location: str = "us-central1" @@ -35,26 +41,36 @@ class VertexMistralOcrSdkInput(MistralCompatibleOcrSdkInput): class VertexDeepSeekOcrSdkInput(OcrSdkInputBase): boundary: str = Field(default="vertex_deepseek", pattern=r"^vertex_deepseek$") - model: Literal["vertex_ai/deepseek-ocr-maas"] = "vertex_ai/deepseek-ocr-maas" + model: VertexDeepSeekModel = "vertex_ai/deepseek-ai/deepseek-ocr-maas" document: OcrDocument custom_llm_provider: Literal["vertex_ai"] | None = None vertex_project: str vertex_location: str = "us-central1" -def _as_vertex_mistral(case_input: MistralOcrSdkInput, project: str, location: str) -> VertexMistralOcrSdkInput: +def _as_vertex_mistral( + case_input: MistralOcrSdkInput, + project: str, + location: str, + model: VertexMistralModel, +) -> VertexMistralOcrSdkInput: values: Final = case_input.model_dump(mode="python", exclude={"boundary", "model", "custom_llm_provider"}) - return VertexMistralOcrSdkInput.model_validate({**values, "vertex_project": project, "vertex_location": location}) + return VertexMistralOcrSdkInput.model_validate( + {**values, "model": model, "vertex_project": project, "vertex_location": location} + ) def _required_deepseek_inputs(project: str, location: str) -> tuple[VertexDeepSeekOcrSdkInput, ...]: - return ( + cases: Final = ( VertexDeepSeekOcrSdkInput( + model=model, document=image_document("invoice 123", 24), vertex_project=project, vertex_location=location, - ), + ) + for model in VERTEX_DEEPSEEK_MODELS ) + return tuple(cases) @st.composite @@ -84,16 +100,18 @@ def vertex_recording_targets(environ: Mapping[str, str], client: OcrFixtureClien SearchStrategy[OcrSdkInputBase], st.builds( _as_vertex_mistral, - case_input=mistral_input_strategy(MISTRAL_MODEL), project=st.just(project), location=st.just(location), + model=st.sampled_from(VERTEX_MISTRAL_MODELS), + case_input=mistral_input_strategy(MISTRAL_MODEL), ), ), invocation=invocation, required_inputs=cast( tuple[OcrSdkInputBase, ...], tuple( - _as_vertex_mistral(case_input, project, location) + _as_vertex_mistral(case_input, project, location, model) + for model in VERTEX_MISTRAL_MODELS for case_input in required_mistral_inputs(MISTRAL_MODEL) ), ), diff --git a/tests/test_litellm/ocr/test_fixture_models.py b/tests/test_litellm/ocr/test_fixture_models.py index 0292e714b36..145a8699b07 100644 --- a/tests/test_litellm/ocr/test_fixture_models.py +++ b/tests/test_litellm/ocr/test_fixture_models.py @@ -1,17 +1,25 @@ from __future__ import annotations from collections.abc import Callable +from datetime import date +from pathlib import Path from typing import Final, cast import pytest from hypothesis import given, settings -from pydantic import ValidationError +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError +from litellm.llms.azure_ai.ocr.document_intelligence.transformation import AzureDocumentIntelligenceOCRConfig +from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig from litellm.llms.mistral.ocr.transformation import MistralOCRConfig from litellm.llms.reducto.ocr.transformation import ReductoParseLegacyConfig, ReductoParseV3Config from litellm.llms.vertex_ai.ocr.deepseek_transformation import VertexAIDeepSeekOCRConfig +from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig +from tests.test_litellm.ocr.conftest import ocr_fixture_marks from tests.test_litellm.ocr.fixtures.azure import ( + AZURE_DOCUMENT_INTELLIGENCE_MODELS, + AZURE_MISTRAL_MODELS, AzureDocumentIntelligenceOcrSdkInput, AzureMistralOcrSdkInput, azure_document_intelligence_input_strategy, @@ -24,8 +32,11 @@ from tests.test_litellm.ocr.fixtures.base import ( JsonSchemaResponseFormat, OcrSdkInputBase, ) -from tests.test_litellm.ocr.fixtures.mistral import MistralOcrSdkInput, mistral_input_strategy +from tests.test_litellm.ocr.fixtures.mistral import MISTRAL_MODELS, MistralOcrSdkInput, mistral_input_strategy +from tests.test_litellm.ocr.fixtures.models import OcrParityCase from tests.test_litellm.ocr.fixtures.reducto import ( + REDUCTO_LEGACY_MODELS, + REDUCTO_V3_MODELS, ReductoChunking, ReductoDocumentUrlDocument, ReductoFormatting, @@ -38,6 +49,8 @@ from tests.test_litellm.ocr.fixtures.reducto import ( reducto_v3_input_strategy, ) from tests.test_litellm.ocr.fixtures.vertex import ( + VERTEX_DEEPSEEK_MODELS, + VERTEX_MISTRAL_MODELS, VertexDeepSeekOcrSdkInput, VertexMistralOcrSdkInput, vertex_deepseek_input_strategy, @@ -46,6 +59,29 @@ from tests.test_litellm.ocr.fixtures.vertex import ( COMMON_FIELDS: Final = frozenset( {"boundary", "model", "document", "custom_llm_provider", "vertex_project", "vertex_location"} ) +SUPPORTED_OCR_PROVIDERS: Final = frozenset({"mistral", "azure_ai", "reducto", "vertex_ai"}) +ACTIVE_OCR_MODELS: Final = frozenset( + ( + *MISTRAL_MODELS, + *AZURE_MISTRAL_MODELS, + *AZURE_DOCUMENT_INTELLIGENCE_MODELS, + *VERTEX_MISTRAL_MODELS, + *VERTEX_DEEPSEEK_MODELS, + *REDUCTO_V3_MODELS, + *REDUCTO_LEGACY_MODELS, + ) +) + + +class _ModelRegistryEntry(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + mode: str | None = None + litellm_provider: str | None = None + deprecation_date: date | None = None + + +MODEL_REGISTRY: Final = TypeAdapter(dict[str, dict[str, JsonValue]]) def _provider_fields(model: type[OcrSdkInputBase]) -> set[str]: @@ -74,26 +110,47 @@ def _reducto_document() -> ReductoDocumentUrlDocument: ) -def test_mistral_fixture_fields_match_provider_config() -> None: - assert _provider_fields(MistralOcrSdkInput) == _supported_params(MistralOCRConfig(), "mistral-ocr-latest") - - -def test_reducto_fixture_fields_match_provider_configs() -> None: - assert _provider_fields(ReductoParseV3SdkInput) == _supported_params(ReductoParseV3Config(), "parse-v3") - assert _provider_fields(ReductoParseLegacySdkInput) == _supported_params(ReductoParseLegacyConfig(), "parse-legacy") - - -def test_deepseek_fixture_fields_match_provider_config() -> None: - assert _provider_fields(VertexDeepSeekOcrSdkInput) == _supported_params( - VertexAIDeepSeekOCRConfig(), "deepseek-ocr-maas" +def test_fixture_catalogs_match_active_registered_ocr_models() -> None: + registry_path: Final = Path(__file__).resolve().parents[3] / "model_prices_and_context_window.json" + registry: Final = MODEL_REGISTRY.validate_json(registry_path.read_text(encoding="utf-8")) + active_registered: Final = frozenset( + model + for model, raw_metadata in registry.items() + if raw_metadata.get("mode") == "ocr" and raw_metadata.get("litellm_provider") in SUPPORTED_OCR_PROVIDERS + for metadata in (_ModelRegistryEntry.model_validate(raw_metadata),) + if metadata.deprecation_date is None or metadata.deprecation_date > date.today() ) + assert ACTIVE_OCR_MODELS == active_registered + + +@pytest.mark.parametrize( + ("fixture_model", "provider_config", "model"), + ( + (MistralOcrSdkInput, MistralOCRConfig(), "mistral-ocr-latest"), + (AzureMistralOcrSdkInput, AzureAIOCRConfig(), "mistral-document-ai-2512"), + ( + AzureDocumentIntelligenceOcrSdkInput, + AzureDocumentIntelligenceOCRConfig(), + "doc-intelligence/prebuilt-layout", + ), + (VertexMistralOcrSdkInput, VertexAIOCRConfig(), "mistral-ocr-2505"), + (VertexDeepSeekOcrSdkInput, VertexAIDeepSeekOCRConfig(), "deepseek-ai/deepseek-ocr-maas"), + (ReductoParseV3SdkInput, ReductoParseV3Config(), "parse-v3"), + (ReductoParseLegacySdkInput, ReductoParseLegacyConfig(), "parse-legacy"), + ), +) +def test_fixture_fields_match_provider_config( + fixture_model: type[OcrSdkInputBase], provider_config: BaseOCRConfig, model: str +) -> None: + assert _provider_fields(fixture_model) == _supported_params(provider_config, model) + @pytest.mark.parametrize( "sdk_input", ( AzureMistralOcrSdkInput( - model="azure_ai/mistral-ocr-deployment", + model="azure_ai/mistral-document-ai-2512", document=ImageUrlDocument(type="image_url", image_url="data:image/png;base64,AA=="), ), VertexMistralOcrSdkInput( @@ -198,6 +255,69 @@ def test_unqualified_models_require_explicit_provider() -> None: ReductoParseV3SdkInput(model="parse-v3", document=_reducto_document()) +@pytest.mark.parametrize("model", tuple(model.removeprefix("mistral/") for model in MISTRAL_MODELS)) +def test_unqualified_mistral_models_accept_explicit_provider(model: str) -> None: + sdk_input: Final = MistralOcrSdkInput.model_validate( + { + "model": model, + "custom_llm_provider": "mistral", + "document": ImageUrlDocument(type="image_url", image_url="https://example.com/image.png"), + } + ) + + assert sdk_input.model == model + + +@pytest.mark.parametrize( + ("model", "model_type"), + (("parse-v3", ReductoParseV3SdkInput), ("parse-legacy", ReductoParseLegacySdkInput)), +) +def test_unqualified_reducto_models_accept_explicit_provider( + model: str, model_type: type[ReductoParseV3SdkInput] | type[ReductoParseLegacySdkInput] +) -> None: + sdk_input: Final = model_type.model_validate( + {"model": model, "custom_llm_provider": "reducto", "document": _reducto_document()} + ) + + assert sdk_input.model == model + + +@pytest.mark.parametrize("model", ("deepseek-ocr-maas", "deepseek-ai/deepseek-ocr-maas")) +def test_vertex_deepseek_request_uses_single_provider_namespace(model: str) -> None: + request: Final = VertexAIDeepSeekOCRConfig().transform_ocr_request( # pyright: ignore[reportUnknownMemberType] + model=model, + document={"type": "image_url", "image_url": "data:image/png;base64,AA=="}, + optional_params={}, + headers={}, + ) + + data: Final = cast(dict[str, object], request.data) + assert data["model"] == "deepseek-ai/deepseek-ocr-maas" + + +@pytest.mark.parametrize( + "sdk_input", + ( + ReductoParseV3SdkInput(model="reducto/parse-v3", document=_reducto_document()), + ReductoParseLegacySdkInput(model="reducto/parse-legacy", document=_reducto_document()), + ), +) +def test_reducto_parity_cases_are_non_strict_xfails( + sdk_input: ReductoParseV3SdkInput | ReductoParseLegacySdkInput, +) -> None: + marks: Final = ocr_fixture_marks(OcrParityCase(litellm_input=sdk_input, provider_responses=())) + + assert len(marks) == 1 + assert marks[0].mark.name == "xfail" + assert marks[0].mark.kwargs["strict"] is False + + +def test_supported_parity_cases_have_no_marks() -> None: + sdk_input: Final = _mistral_input() + + assert ocr_fixture_marks(OcrParityCase(litellm_input=sdk_input, provider_responses=())) == () + + def test_reducto_v3_preserves_nested_provider_params() -> None: sdk_input: Final = ReductoParseV3SdkInput( model="reducto/parse-v3", diff --git a/tests/test_litellm/ocr/test_fixture_store.py b/tests/test_litellm/ocr/test_fixture_store.py new file mode 100644 index 00000000000..e12c6b3aa4b --- /dev/null +++ b/tests/test_litellm/ocr/test_fixture_store.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from queue import Queue +from typing import Final, Protocol, cast + +import pytest + +from tests.route_parity.fixtures.store import parametrize_recorded_fixtures +from tests.test_litellm.ocr.conftest import ocr_fixture_id, ocr_fixture_marks +from tests.test_litellm.ocr.fixtures.models import OcrParityCase + + +class _Parameter(Protocol): + values: tuple[OcrParityCase, ...] + marks: tuple[pytest.Mark, ...] + + +@dataclass(frozen=True, slots=True) +class _MetafuncSpy: + fixturenames: tuple[str, ...] + calls: Queue[tuple[object, ...]] + + def parametrize(self, *args: object, **_kwargs: object) -> None: + self.calls.put(args) + + +def test_recorded_fixture_parametrization_applies_case_specific_marks() -> None: + calls: Final[Queue[tuple[object, ...]]] = Queue() + metafunc: Final = _MetafuncSpy(fixturenames=("ocr_fixture",), calls=calls) + + parametrize_recorded_fixtures( + cast(pytest.Metafunc, metafunc), + fixture_name="ocr_fixture", + case_type=OcrParityCase, + env_var="UNCONFIGURED_OCR_FIXTURE_TEST_DIRECTORY", + default_directory=Path(__file__).with_name("fixtures") / "data", + regeneration_command="unused", + id_builder=ocr_fixture_id, + marks_builder=ocr_fixture_marks, + ) + + parameters: Final = cast(tuple[_Parameter, ...], calls.get_nowait()[1]) + reducto_parameters: Final = tuple( + parameter + for parameter in parameters + if parameter.values[0].litellm_input.boundary in {"reducto_v3", "reducto_legacy"} + ) + supported_parameters: Final = tuple(parameter for parameter in parameters if parameter not in reducto_parameters) + + assert reducto_parameters + assert supported_parameters + assert all(len(parameter.marks) == 1 for parameter in reducto_parameters) + assert all(parameter.marks[0].name == "xfail" for parameter in reducto_parameters) + assert all(parameter.marks[0].kwargs["strict"] is False for parameter in reducto_parameters) + assert all(parameter.marks == () for parameter in supported_parameters) diff --git a/tests/test_litellm/ocr/test_record_fixtures.py b/tests/test_litellm/ocr/test_record_fixtures.py index 683511adf74..ca39c545a5f 100644 --- a/tests/test_litellm/ocr/test_record_fixtures.py +++ b/tests/test_litellm/ocr/test_record_fixtures.py @@ -9,11 +9,18 @@ import pytest from tests.route_parity.fixtures.inputs import generate_case_inputs from tests.route_parity.fixtures.pipeline import parse_recording_args +from tests.test_litellm.ocr.fixtures.azure import ( + AZURE_DOCUMENT_INTELLIGENCE_MODELS, + AZURE_MISTRAL_MODELS, +) from tests.test_litellm.ocr.fixtures.base import OcrSdkInputBase +from tests.test_litellm.ocr.fixtures.mistral import MISTRAL_MODELS from tests.test_litellm.ocr.fixtures.record import ( discover_targets, require_targets, ) +from tests.test_litellm.ocr.fixtures.reducto import REDUCTO_LEGACY_MODELS, REDUCTO_V3_MODELS +from tests.test_litellm.ocr.fixtures.vertex import VERTEX_DEEPSEEK_MODELS, VERTEX_MISTRAL_MODELS class _UnusedOcrClient: @@ -30,6 +37,29 @@ class _RecordingOcrClient: _UNUSED_OCR_CLIENT: Final = _UnusedOcrClient() +_MISTRAL_PARAMS: Final = frozenset( + { + "pages", + "include_image_base64", + "image_limit", + "image_min_size", + "bbox_annotation_format", + "document_annotation_format", + "document_annotation_prompt", + "extract_header", + "extract_footer", + "table_format", + "confidence_scores_granularity", + "include_blocks", + "id", + } +) + + +def _model(case_input: OcrSdkInputBase) -> str: + model: Final = case_input.canonical_input().get("model") + assert isinstance(model, str) + return model def test_parse_args_has_no_model_selection() -> None: @@ -66,7 +96,6 @@ def test_discovery_is_explicit_per_available_provider_boundary() -> None: "REDUCTO_API_KEY": "reducto-secret", "AZURE_AI_API_KEY": "azure-secret", "AZURE_AI_API_BASE": "https://azure.example", - "AZURE_AI_OCR_MODEL": "mistral-ocr-deployment", "AZURE_DOCUMENT_INTELLIGENCE_API_KEY": "document-secret", "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT": "https://document.example", "VERTEX_AI_API_KEY": "vertex-secret", @@ -87,17 +116,14 @@ def test_discovery_is_explicit_per_available_provider_boundary() -> None: assert all("secret" not in repr(target) for target in targets) -def test_azure_mistral_discovery_requires_and_normalizes_deployment_model() -> None: - incomplete: Final = { +def test_azure_mistral_discovery_enumerates_registered_models() -> None: + environ: Final = { "AZURE_AI_API_KEY": "azure-secret", "AZURE_AI_API_BASE": "https://azure.example", } - assert discover_targets(incomplete, _UNUSED_OCR_CLIENT) == () + target: Final = discover_targets(environ, _UNUSED_OCR_CLIENT)[0] - target: Final = discover_targets( - {**incomplete, "AZURE_AI_OCR_MODEL": "mistral-ocr-deployment"}, _UNUSED_OCR_CLIENT - )[0] - assert target.required_inputs[0].canonical_input()["model"] == "azure_ai/mistral-ocr-deployment" + assert {_model(case_input) for case_input in target.required_inputs} == set(AZURE_MISTRAL_MODELS) @pytest.mark.parametrize( @@ -125,29 +151,15 @@ def test_mistral_target_uses_canonical_model_and_normalized_base( assert "mistral-secret" not in repr(target) case_inputs: Final = generate_case_inputs(target.strategy, examples=1) assert len(case_inputs) == 1 - assert case_inputs[0].canonical_input()["model"] == "mistral/mistral-ocr-latest" - assert len(target.required_inputs) == 14 + assert case_inputs[0].canonical_input()["model"] in MISTRAL_MODELS + assert len(target.required_inputs) == 14 * len(MISTRAL_MODELS) covered_params: Final = { key for case_input in target.required_inputs for key in case_input.as_sdk_kwargs() if key not in {"model", "document", "custom_llm_provider"} } - assert covered_params == { - "pages", - "include_image_base64", - "image_limit", - "image_min_size", - "bbox_annotation_format", - "document_annotation_format", - "document_annotation_prompt", - "extract_header", - "extract_footer", - "table_format", - "confidence_scores_granularity", - "include_blocks", - "id", - } + assert covered_params == _MISTRAL_PARAMS def test_mistral_target_invocation_forwards_discovered_credentials() -> None: @@ -162,4 +174,45 @@ def test_mistral_target_invocation_forwards_discovered_credentials() -> None: kwargs: Final = calls.get_nowait() assert kwargs["api_base"] == "http://127.0.0.1:1234" assert kwargs["api_key"] == "mistral-secret" - assert kwargs["model"] == "mistral/mistral-ocr-latest" + assert kwargs["model"] in MISTRAL_MODELS + + +def test_every_target_covers_every_supported_param_for_every_model() -> None: + targets: Final = discover_targets( + { + "MISTRAL_API_KEY": "mistral-secret", + "REDUCTO_API_KEY": "reducto-secret", + "AZURE_AI_API_KEY": "azure-secret", + "AZURE_AI_API_BASE": "https://azure.example", + "AZURE_DOCUMENT_INTELLIGENCE_API_KEY": "document-secret", + "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT": "https://document.example", + "VERTEX_AI_API_KEY": "vertex-secret", + "VERTEXAI_PROJECT": "project-1", + }, + _UNUSED_OCR_CLIENT, + ) + expected: Final[dict[str, tuple[tuple[str, ...], frozenset[str]]]] = { + "mistral-ocr": (MISTRAL_MODELS, _MISTRAL_PARAMS), + "azure-mistral": (AZURE_MISTRAL_MODELS, _MISTRAL_PARAMS), + "azure-document-intelligence": ( + AZURE_DOCUMENT_INTELLIGENCE_MODELS, + frozenset({"pages", "features", "req_format"}), + ), + "vertex-mistral": (VERTEX_MISTRAL_MODELS, _MISTRAL_PARAMS), + "vertex-deepseek": (VERTEX_DEEPSEEK_MODELS, frozenset[str]()), + "reducto-v3": (REDUCTO_V3_MODELS, frozenset({"formatting", "retrieval", "settings"})), + "reducto-legacy": (REDUCTO_LEGACY_MODELS, frozenset({"enhance"})), + } + + for target in targets: + expected_models, expected_params = expected[target.name] + assert {_model(case_input) for case_input in target.required_inputs} == set(expected_models) + for model in expected_models: + covered = frozenset( + key + for case_input in target.required_inputs + if _model(case_input) == model + for key in case_input.as_sdk_kwargs() + if key not in {"model", "document", "custom_llm_provider", "vertex_project", "vertex_location"} + ) + assert covered == expected_params diff --git a/tests/test_litellm/ocr/test_sdk_parity.py b/tests/test_litellm/ocr/test_sdk_parity.py index 4352a85fbe1..cd45204d7c4 100644 --- a/tests/test_litellm/ocr/test_sdk_parity.py +++ b/tests/test_litellm/ocr/test_sdk_parity.py @@ -332,8 +332,6 @@ def test_recorded_ocr_sdk_parity( ocr_fixture: OcrParityCase, route: SDKRoute, ) -> None: - if ocr_fixture.litellm_input.boundary in {"reducto_v3", "reducto_legacy"}: - pytest.skip("Reducto does not have a Rust OCR boundary") sync_spy, async_spy = _native_spies() with _restore_rust_ocr_state(), replay_server() as provider: rust_ocr_bridge.use_litellm_rust(False, ocr=sync_spy, aocr=async_spy)