From fea5204a2f4cb9c2a73397c10208ed1c1241a523 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 22 Jun 2026 17:53:18 -0700 Subject: [PATCH] Simplify rust ocr entrypoint --- litellm/llms/base_llm/ocr/transformation.py | 6 - litellm/llms/mistral/ocr/rust_provider.py | 43 ++- litellm/llms/mistral/ocr/transformation.py | 4 - litellm/ocr/__init__.py | 4 +- litellm/ocr/main.py | 267 ++++++++++++++++-- litellm/rust_bridge/__init__.py | 2 - litellm/rust_bridge/loader.py | 6 +- litellm/rust_bridge/ocr/__init__.py | 5 +- litellm/rust_bridge/ocr/config.py | 199 ------------- litellm/rust_bridge/ocr/providers.py | 12 +- .../test_mistral_ocr_rust_bridge.py | 195 ++++++------- 11 files changed, 395 insertions(+), 348 deletions(-) delete mode 100644 litellm/rust_bridge/ocr/config.py diff --git a/litellm/llms/base_llm/ocr/transformation.py b/litellm/llms/base_llm/ocr/transformation.py index 4783a4762d0..263e0c094ce 100644 --- a/litellm/llms/base_llm/ocr/transformation.py +++ b/litellm/llms/base_llm/ocr/transformation.py @@ -101,12 +101,6 @@ class BaseOCRConfig: """ return [] - def get_rust_ocr_provider(self, model: str) -> Optional[str]: - """ - Return the Rust OCR provider id for providers that explicitly opt in. - """ - return None - def map_ocr_params( self, non_default_params: dict, diff --git a/litellm/llms/mistral/ocr/rust_provider.py b/litellm/llms/mistral/ocr/rust_provider.py index 960f7037810..fedaad3c2ac 100644 --- a/litellm/llms/mistral/ocr/rust_provider.py +++ b/litellm/llms/mistral/ocr/rust_provider.py @@ -1,15 +1,46 @@ from enum import Enum -from typing import Optional +from typing import Any + +from litellm.secret_managers.main import get_secret_str class MistralRustOcrProvider(str, Enum): MISTRAL = "mistral" -def get_mistral_rust_ocr_provider(config_type: type) -> Optional[str]: - if ( - config_type.__module__ != "litellm.llms.mistral.ocr.transformation" - or config_type.__name__ != "MistralOCRConfig" - ): +def get_mistral_rust_ocr_provider( + custom_llm_provider: str | None, +) -> str | None: + provider_value = getattr(custom_llm_provider, "value", custom_llm_provider) + if provider_value != MistralRustOcrProvider.MISTRAL.value: return None return MistralRustOcrProvider.MISTRAL.value + + +def get_mistral_rust_ocr_url(api_base: str | None) -> str: + if api_base is None: + api_base = "https://api.mistral.ai/v1" + + api_base = api_base.rstrip("/") + if api_base.endswith("/v1"): + return f"{api_base}/ocr" + return f"{api_base}/v1/ocr" + + +def get_mistral_rust_ocr_headers( + headers: dict[str, Any] | None, + api_key: str | None, +) -> dict[str, Any]: + if api_key is None: + api_key = get_secret_str("MISTRAL_API_KEY") + + if api_key is None: + raise ValueError( + "Missing Mistral API Key - A call is being made to Mistral but no key " + "is set either in the environment variables or via params" + ) + + return { + "Authorization": f"Bearer {api_key}", + **(headers or {}), + } diff --git a/litellm/llms/mistral/ocr/transformation.py b/litellm/llms/mistral/ocr/transformation.py index 99bdb2da45d..21e0e27a314 100644 --- a/litellm/llms/mistral/ocr/transformation.py +++ b/litellm/llms/mistral/ocr/transformation.py @@ -13,7 +13,6 @@ from litellm.llms.base_llm.ocr.transformation import ( OCRRequestData, OCRResponse, ) -from litellm.llms.mistral.ocr.rust_provider import get_mistral_rust_ocr_provider from litellm.secret_managers.main import get_secret_str @@ -27,9 +26,6 @@ class MistralOCRConfig(BaseOCRConfig): def __init__(self) -> None: super().__init__() - def get_rust_ocr_provider(self, model: str) -> Optional[str]: - return get_mistral_rust_ocr_provider(type(self)) - def get_supported_ocr_params(self, model: str) -> list: """ Get supported OCR parameters for Mistral OCR. diff --git a/litellm/ocr/__init__.py b/litellm/ocr/__init__.py index e97497b2db7..8281e58046d 100644 --- a/litellm/ocr/__init__.py +++ b/litellm/ocr/__init__.py @@ -1,5 +1,5 @@ """OCR module for LiteLLM.""" -from .main import aocr, ocr +from .main import aocr, ocr, rust_ocr -__all__ = ["ocr", "aocr"] +__all__ = ["ocr", "aocr", "rust_ocr"] diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index dcef3c79ffe..e6eba696c25 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -19,8 +19,16 @@ from litellm._logging import verbose_logger from litellm.constants import request_timeout from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse -from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler -from litellm.rust_bridge.ocr import get_rust_ocr_provider_config +from litellm.llms.custom_httpx.llm_http_handler import ( + BaseLLMHTTPHandler, + _get_httpx_client, +) +from litellm.llms.mistral.ocr.rust_provider import ( + get_mistral_rust_ocr_headers, + get_mistral_rust_ocr_provider, + get_mistral_rust_ocr_url, +) +from litellm.rust_bridge.ocr import call_ocr, rust_ocr_provider_enabled from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -29,6 +37,155 @@ base_llm_http_handler = BaseLLMHTTPHandler() ################################################# +def _prepare_ocr_document(document: dict[str, Any]) -> dict[str, str]: + if not isinstance(document, dict): + raise ValueError( + f"document must be a dict with 'type' and URL/file field, got {type(document)}" + ) + + doc_type = document.get("type") + if doc_type == "file": + document = convert_file_document_to_url_document(document) + doc_type = document.get("type") + + if doc_type not in ["document_url", "image_url"]: + raise ValueError( + f"Invalid document type: {doc_type}. " + "Must be 'document_url', 'image_url', or 'file'" + ) + + return document + + +def _get_rust_ocr_provider(custom_llm_provider: str | None) -> str | None: + return get_mistral_rust_ocr_provider(custom_llm_provider=custom_llm_provider) + + +def _should_route_to_rust_ocr(custom_llm_provider: str | None) -> bool: + rust_ocr_provider = _get_rust_ocr_provider(custom_llm_provider) + return rust_ocr_provider is not None and rust_ocr_provider_enabled( + rust_ocr_provider + ) + + +def _call_rust_ocr( + payload: dict[str, Any], + *, + require_enabled: bool, + fallback_on_unavailable: bool, +) -> dict[str, Any] | None: + result = call_ocr(payload, require_enabled=require_enabled) + if result is not None: + return result + if fallback_on_unavailable: + return None + raise ValueError("Rust OCR bridge is unavailable for this request") + + +def _rust_ocr_impl( + model: str, + document: dict[str, str], + api_key: str | None, + api_base: str | None, + timeout: float | httpx.Timeout | None, + custom_llm_provider: str | None, + extra_headers: dict[str, Any] | None, + litellm_logging_obj: LiteLLMLoggingObj, + litellm_call_id: str | None, + fallback_on_unavailable: bool, + require_enabled: bool, + kwargs: dict[str, Any], +) -> OCRResponse | None: + rust_ocr_provider = _get_rust_ocr_provider(custom_llm_provider) + if rust_ocr_provider is None: + if fallback_on_unavailable: + return None + raise ValueError( + f"Rust OCR is not supported for provider: {custom_llm_provider}" + ) + + if require_enabled and not rust_ocr_provider_enabled(rust_ocr_provider): + return None + + optional_params = _call_rust_ocr( + { + "provider": rust_ocr_provider, + "operation": "map_params", + "non_default_params": dict(kwargs), + }, + require_enabled=require_enabled, + fallback_on_unavailable=fallback_on_unavailable, + ) + if optional_params is None: + return None + + transformed_request = _call_rust_ocr( + { + "provider": rust_ocr_provider, + "operation": "transform_request", + "model": model, + "document": document, + "optional_params": optional_params, + }, + require_enabled=require_enabled, + fallback_on_unavailable=fallback_on_unavailable, + ) + if transformed_request is None: + return None + + request_data = transformed_request.get("data") + if not isinstance(request_data, dict): + raise ValueError(f"Rust OCR provider {rust_ocr_provider} returned invalid data") + + headers = get_mistral_rust_ocr_headers( + headers=extra_headers, + api_key=api_key, + ) + complete_url = get_mistral_rust_ocr_url(api_base=api_base) + + litellm_logging_obj.update_from_kwargs( + kwargs=kwargs, + model=model, + optional_params=optional_params, + litellm_params={ + "litellm_call_id": litellm_call_id, + "api_base": api_base, + }, + custom_llm_provider=custom_llm_provider, + ) + litellm_logging_obj.pre_call( + input="OCR document processing", + api_key=api_key, + additional_args={ + "complete_input_dict": request_data, + "api_base": complete_url, + "headers": headers, + }, + ) + + response = _get_httpx_client().post( + url=complete_url, + headers=headers, + json=request_data, + timeout=timeout or request_timeout, + ) + + transformed_response = _call_rust_ocr( + { + "provider": rust_ocr_provider, + "operation": "transform_response", + "model": model, + "response_json": response.json(), + }, + require_enabled=require_enabled, + fallback_on_unavailable=fallback_on_unavailable, + ) + if transformed_response is None: + return None + + return OCRResponse(**transformed_response) + + @client async def aocr( model: str, @@ -146,6 +303,70 @@ async def aocr( ) +@client +def rust_ocr( + model: str, + document: dict[str, Any], + api_key: str | None = None, + api_base: str | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, + extra_headers: dict[str, Any] | None = None, + **kwargs, +) -> OCRResponse: + """ + Direct Rust OCR entrypoint. + """ + local_vars = locals() + try: + litellm_logging_obj: LiteLLMLoggingObj = kwargs.pop("litellm_logging_obj") # type: ignore + litellm_call_id: str | None = kwargs.get("litellm_call_id", None) + document = _prepare_ocr_document(document) + + ( + model, + custom_llm_provider, + dynamic_api_key, + dynamic_api_base, + ) = litellm.get_llm_provider( + model=model, + custom_llm_provider=custom_llm_provider, + api_base=api_base, + api_key=api_key, + ) + + if dynamic_api_key: + api_key = dynamic_api_key + if dynamic_api_base: + api_base = dynamic_api_base + + response = _rust_ocr_impl( + model=model, + document=document, + api_key=api_key, + api_base=api_base, + timeout=timeout, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + litellm_logging_obj=litellm_logging_obj, + litellm_call_id=litellm_call_id, + fallback_on_unavailable=False, + require_enabled=False, + kwargs=kwargs, + ) + if response is None: + raise ValueError("Rust OCR bridge is unavailable for this request") + return response + except Exception as e: + raise litellm.exception_type( + model=model, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + @client def ocr( model: str, @@ -225,24 +446,7 @@ def ocr( litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aocr", False) is True - # Validate document parameter format - if not isinstance(document, dict): - raise ValueError( - f"document must be a dict with 'type' and URL/file field, got {type(document)}" - ) - - doc_type = document.get("type") - - # Handle file type: convert to document_url/image_url with base64 data URI - if doc_type == "file": - document = convert_file_document_to_url_document(document) - doc_type = document.get("type") - - if doc_type not in ["document_url", "image_url"]: - raise ValueError( - f"Invalid document type: {doc_type}. " - "Must be 'document_url', 'image_url', or 'file'" - ) + document = _prepare_ocr_document(document) ( model, @@ -262,6 +466,24 @@ def ocr( if dynamic_api_base: api_base = dynamic_api_base + if _should_route_to_rust_ocr(custom_llm_provider): + rust_response = _rust_ocr_impl( + model=model, + document=document, + api_key=api_key, + api_base=api_base, + timeout=timeout, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + litellm_logging_obj=litellm_logging_obj, + litellm_call_id=litellm_call_id, + fallback_on_unavailable=True, + require_enabled=True, + kwargs=kwargs, + ) + if rust_response is not None: + return rust_response + # Get provider config ocr_provider_config: Optional[BaseOCRConfig] = ( ProviderConfigManager.get_provider_ocr_config( @@ -279,11 +501,6 @@ def ocr( f"OCR call - model: {model}, provider: {custom_llm_provider}" ) - ocr_provider_config = get_rust_ocr_provider_config( - model=model, - fallback_config=ocr_provider_config, - ) - # Get litellm params using GenericLiteLLMParams (same as responses API) litellm_params = GenericLiteLLMParams(**kwargs) diff --git a/litellm/rust_bridge/__init__.py b/litellm/rust_bridge/__init__.py index df1467f3338..5e44170f89f 100644 --- a/litellm/rust_bridge/__init__.py +++ b/litellm/rust_bridge/__init__.py @@ -3,10 +3,8 @@ from litellm.rust_bridge.loader import ( set_rust_core_enabled, set_rust_core_strict, ) -from litellm.rust_bridge.ocr import get_rust_ocr_provider_config __all__ = [ - "get_rust_ocr_provider_config", "rust_core_available", "set_rust_core_enabled", "set_rust_core_strict", diff --git a/litellm/rust_bridge/loader.py b/litellm/rust_bridge/loader.py index 270ee1b2045..597e9e5d1e5 100644 --- a/litellm/rust_bridge/loader.py +++ b/litellm/rust_bridge/loader.py @@ -1,14 +1,14 @@ import importlib from functools import lru_cache from types import ModuleType -from typing import Any, Iterable, Optional, Union +from typing import Any, Iterable, Union _enabled_rust_core_scopes: set[str] = set() _rust_core_strict = False @lru_cache(maxsize=1) -def _load_rust_module() -> Optional[ModuleType]: +def _load_rust_module() -> ModuleType | None: try: return importlib.import_module("litellm_python_bridge") except Exception: @@ -48,7 +48,7 @@ def set_rust_core_strict(enabled: bool) -> None: _rust_core_strict = enabled -def call_rust_function(function_name: str, *args: Any) -> Optional[Any]: +def call_rust_function(function_name: str, *args: Any) -> Any | None: module = _load_rust_module() if module is None: return None diff --git a/litellm/rust_bridge/ocr/__init__.py b/litellm/rust_bridge/ocr/__init__.py index c4e958a0b50..c0cf35c642b 100644 --- a/litellm/rust_bridge/ocr/__init__.py +++ b/litellm/rust_bridge/ocr/__init__.py @@ -1,5 +1,6 @@ -from litellm.rust_bridge.ocr.config import get_rust_ocr_provider_config +from litellm.rust_bridge.ocr.providers import call_ocr, rust_ocr_provider_enabled __all__ = [ - "get_rust_ocr_provider_config", + "call_ocr", + "rust_ocr_provider_enabled", ] diff --git a/litellm/rust_bridge/ocr/config.py b/litellm/rust_bridge/ocr/config.py deleted file mode 100644 index fd38d0a104a..00000000000 --- a/litellm/rust_bridge/ocr/config.py +++ /dev/null @@ -1,199 +0,0 @@ -from typing import Any, Optional - -import httpx - -from litellm.llms.base_llm.ocr.transformation import ( - BaseOCRConfig, - DocumentType, - OCRRequestData, - OCRResponse, -) -from litellm.rust_bridge.ocr.providers import call_ocr - - -def get_rust_ocr_provider_config( - model: str, - fallback_config: BaseOCRConfig, -) -> BaseOCRConfig: - rust_ocr_provider = fallback_config.get_rust_ocr_provider(model=model) - if not rust_ocr_provider: - return fallback_config - - return RustOCRProviderConfig( - rust_ocr_provider=rust_ocr_provider, - fallback_config=fallback_config, - ) - - -class RustOCRProviderConfig(BaseOCRConfig): - def __init__( - self, - rust_ocr_provider: str, - fallback_config: BaseOCRConfig, - ) -> None: - super().__init__() - self.rust_ocr_provider = rust_ocr_provider - self.fallback_config = fallback_config - - def get_supported_ocr_params(self, model: str) -> list: - return self.fallback_config.get_supported_ocr_params(model=model) - - def validate_environment( - self, - headers: dict, - model: str, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - litellm_params: Optional[dict] = None, - **kwargs, - ) -> dict: - return self.fallback_config.validate_environment( - headers=headers, - model=model, - api_key=api_key, - api_base=api_base, - litellm_params=litellm_params, - **kwargs, - ) - - def get_complete_url( - self, - api_base: Optional[str], - model: str, - optional_params: dict, - litellm_params: Optional[dict] = None, - **kwargs, - ) -> str: - return self.fallback_config.get_complete_url( - api_base=api_base, - model=model, - optional_params=optional_params, - litellm_params=litellm_params, - **kwargs, - ) - - def map_ocr_params( - self, - non_default_params: dict, - optional_params: dict, - model: str, - ) -> dict: - mapped_params = call_ocr( - { - "provider": self.rust_ocr_provider, - "operation": "map_params", - "non_default_params": non_default_params, - } - ) - if mapped_params is not None: - return mapped_params - return self.fallback_config.map_ocr_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - ) - - def transform_ocr_request( - self, - model: str, - document: DocumentType, - optional_params: dict, - headers: dict, - **kwargs, - ) -> OCRRequestData: - if isinstance(document, dict): - transformed_request = call_ocr( - { - "provider": self.rust_ocr_provider, - "operation": "transform_request", - "model": model, - "document": document, - "optional_params": optional_params, - }, - ) - if transformed_request is not None: - request_data = transformed_request.get("data") - if not isinstance(request_data, dict): - raise ValueError( - f"Rust OCR provider {self.rust_ocr_provider} " - "returned invalid request data" - ) - return OCRRequestData( - data=request_data, - files=transformed_request.get("files"), - ) - - return self.fallback_config.transform_ocr_request( - model=model, - document=document, - optional_params=optional_params, - headers=headers, - **kwargs, - ) - - async def async_transform_ocr_request( - self, - model: str, - document: DocumentType, - optional_params: dict, - headers: dict, - **kwargs, - ) -> OCRRequestData: - return self.transform_ocr_request( - model=model, - document=document, - optional_params=optional_params, - headers=headers, - **kwargs, - ) - - def transform_ocr_response( - self, - model: str, - raw_response: httpx.Response, - logging_obj: Any, - **kwargs, - ) -> OCRResponse: - transformed_response = call_ocr( - { - "provider": self.rust_ocr_provider, - "operation": "transform_response", - "model": model, - "response_json": raw_response.json(), - }, - ) - if transformed_response is not None: - return OCRResponse(**transformed_response) - - return self.fallback_config.transform_ocr_response( - model=model, - raw_response=raw_response, - logging_obj=logging_obj, - **kwargs, - ) - - async def async_transform_ocr_response( - self, - model: str, - raw_response: httpx.Response, - logging_obj: Any, - **kwargs, - ) -> OCRResponse: - return self.transform_ocr_response( - model=model, - raw_response=raw_response, - logging_obj=logging_obj, - **kwargs, - ) - - def get_error_class( - self, - error_message: str, - status_code: int, - headers: dict, - ) -> Exception: - return self.fallback_config.get_error_class( - error_message=error_message, - status_code=status_code, - headers=headers, - ) diff --git a/litellm/rust_bridge/ocr/providers.py b/litellm/rust_bridge/ocr/providers.py index 160b9f8f4d9..d20edefc082 100644 --- a/litellm/rust_bridge/ocr/providers.py +++ b/litellm/rust_bridge/ocr/providers.py @@ -1,14 +1,18 @@ -from typing import Any, Optional +from typing import Any from litellm.rust_bridge.loader import call_rust_function, rust_core_enabled -def call_ocr(payload: dict[str, Any]) -> Optional[dict[str, Any]]: +def call_ocr( + payload: dict[str, Any], + *, + require_enabled: bool = True, +) -> dict[str, Any] | None: provider = payload.get("provider") if not isinstance(provider, str): return None - if not _rust_ocr_provider_enabled(provider): + if require_enabled and not rust_ocr_provider_enabled(provider): return None result = call_rust_function("ocr", payload) @@ -19,7 +23,7 @@ def call_ocr(payload: dict[str, Any]) -> Optional[dict[str, Any]]: return result -def _rust_ocr_provider_enabled(provider: str) -> bool: +def rust_ocr_provider_enabled(provider: str) -> bool: return ( rust_core_enabled("ocr") or rust_core_enabled(f"ocr:{provider}") diff --git a/tests/test_litellm/test_mistral_ocr_rust_bridge.py b/tests/test_litellm/test_mistral_ocr_rust_bridge.py index 27ab6744a3e..c59be3cf5d3 100644 --- a/tests/test_litellm/test_mistral_ocr_rust_bridge.py +++ b/tests/test_litellm/test_mistral_ocr_rust_bridge.py @@ -1,15 +1,34 @@ +import importlib + import httpx import pytest -from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig -from litellm.llms.mistral.ocr.rust_provider import MistralRustOcrProvider -from litellm.llms.mistral.ocr.transformation import MistralOCRConfig +import litellm +from litellm.llms.base_llm.ocr.transformation import OCRResponse +from litellm.llms.mistral.ocr.rust_provider import ( + MistralRustOcrProvider, + get_mistral_rust_ocr_provider, +) from litellm.rust_bridge import loader -from litellm.rust_bridge import ocr as rust_ocr from litellm.rust_bridge.ocr import providers -from litellm.rust_bridge.ocr.config import RustOCRProviderConfig + +ocr_main = importlib.import_module("litellm.ocr.main") MODEL = "mistral-ocr-latest" +SUPPORTED_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", + "id", +} DOCUMENT = { "type": "document_url", "document_url": "https://example.com/doc.pdf", @@ -36,7 +55,7 @@ class _FakeRustModule: return { key: value for key, value in payload["non_default_params"].items() - if key != "unsupported_param" + if key in SUPPORTED_PARAMS } if operation == "transform_request": return { @@ -59,31 +78,34 @@ class _FakeRustModule: raise AssertionError(f"Unexpected operation: {operation}") -class _FakeLoggingObj: - pass +class _FakeHTTPClient: + def __init__(self): + self.requests = [] - -class _InheritedMistralOCRConfig(MistralOCRConfig): - pass + def post(self, url, headers, json, timeout): + self.requests.append( + { + "url": url, + "headers": headers, + "json": json, + "timeout": timeout, + } + ) + return httpx.Response( + 200, + json={ + "pages": [{"index": 0, "markdown": "hello"}], + "model": "mistral-ocr-2505-completion", + "document_annotation": None, + "usage_info": {"pages_processed": 1}, + }, + ) def test_mistral_rust_ocr_provider_enum_is_owned_by_mistral(): assert MistralRustOcrProvider.MISTRAL.value == "mistral" - assert ( - MistralOCRConfig().get_rust_ocr_provider(model=MODEL) - == MistralRustOcrProvider.MISTRAL.value - ) - - -def test_inherited_mistral_ocr_provider_uses_python_fallback(): - fallback_config = _InheritedMistralOCRConfig() - - config = rust_ocr.get_rust_ocr_provider_config( - model=MODEL, - fallback_config=fallback_config, - ) - - assert config is fallback_config + assert get_mistral_rust_ocr_provider("mistral") == "mistral" + assert get_mistral_rust_ocr_provider("azure_ai") is None def test_rust_ocr_provider_returns_none_when_scope_disabled(monkeypatch): @@ -119,92 +141,75 @@ def test_mistral_ocr_map_params_uses_provider_gated_rust(monkeypatch): assert result == {"extract_header": True} -def test_mistral_ocr_provider_wrapper_uses_rust_when_enabled(monkeypatch): - loader.set_rust_core_enabled("ocr:mistral") +def test_litellm_rust_ocr_calls_rust_bridge(monkeypatch): + fake_client = _FakeHTTPClient() monkeypatch.setattr(loader, "_load_rust_module", lambda: _FakeRustModule) + monkeypatch.setattr(ocr_main, "_get_httpx_client", lambda: fake_client) - config = rust_ocr.get_rust_ocr_provider_config( - model=MODEL, - fallback_config=MistralOCRConfig(), - ) - - assert isinstance(config, BaseOCRConfig) - assert isinstance(config, RustOCRProviderConfig) - - request = config.transform_ocr_request( - model=MODEL, + response = litellm.rust_ocr( + model="mistral/mistral-ocr-latest", document=DOCUMENT, - optional_params={"include_image_base64": True}, - headers={}, - ) - assert request.data == { - "model": MODEL, - "document": DOCUMENT, - "include_image_base64": True, - } - assert request.files is None - - response = config.transform_ocr_response( - model=MODEL, - raw_response=httpx.Response( - 200, - json={ - "pages": [{"index": 0, "markdown": "hello"}], - "model": "mistral-ocr-2505-completion", - "document_annotation": None, - "usage_info": {"pages_processed": 1}, - }, - ), - logging_obj=_FakeLoggingObj(), + api_key="test-key", + pages=[0], + include_image_base64=False, + unsupported_param="drop", ) assert response.pages[0].index == 0 assert response.model == "mistral-ocr-2505-completion" assert response.usage_info.pages_processed == 1 - - -def test_mistral_ocr_provider_wrapper_falls_back_when_rust_module_missing( - monkeypatch, -): - loader.set_rust_core_enabled("ocr:mistral") - monkeypatch.setattr(loader, "_load_rust_module", lambda: None) - - config = rust_ocr.get_rust_ocr_provider_config( - model=MODEL, - fallback_config=MistralOCRConfig(), - ) - - result = config.transform_ocr_request( - model=MODEL, - document=DOCUMENT, - optional_params={"include_image_base64": True}, - headers={}, - ) - - assert result.data == { + assert len(fake_client.requests) == 1 + request = fake_client.requests[0] + assert request["url"] == "https://api.mistral.ai/v1/ocr" + assert request["headers"] == {"Authorization": "Bearer test-key"} + assert request["json"] == { "model": MODEL, "document": DOCUMENT, - "include_image_base64": True, + "pages": [0], + "include_image_base64": False, } -def test_inherited_mistral_ocr_config_does_not_use_mistral_rust(monkeypatch): +def test_litellm_ocr_routes_to_rust_when_mistral_scope_enabled(monkeypatch): + fake_client = _FakeHTTPClient() loader.set_rust_core_enabled("ocr:mistral") monkeypatch.setattr(loader, "_load_rust_module", lambda: _FakeRustModule) + monkeypatch.setattr(ocr_main, "_get_httpx_client", lambda: fake_client) - config = rust_ocr.get_rust_ocr_provider_config( - model=MODEL, - fallback_config=_InheritedMistralOCRConfig(), + response = litellm.ocr( + model="mistral/mistral-ocr-latest", + document=DOCUMENT, + api_key="test-key", + pages=[0], + include_image_base64=False, ) - assert isinstance(config, _InheritedMistralOCRConfig) - result = config.map_ocr_params( - non_default_params={ - "extract_header": True, - "unsupported_param": "value", - }, - optional_params={}, - model=MODEL, + assert response.pages[0].markdown == "hello" + assert len(fake_client.requests) == 1 + + +def test_litellm_ocr_uses_python_path_when_rust_scope_disabled(monkeypatch): + fake_client = _FakeHTTPClient() + monkeypatch.setattr(ocr_main, "_get_httpx_client", lambda: fake_client) + monkeypatch.setattr( + ocr_main.base_llm_http_handler, + "ocr", + lambda **kwargs: OCRResponse( + pages=[{"index": 0, "markdown": "python"}], + model="mistral-ocr-2505-completion", + document_annotation=None, + usage_info={"pages_processed": 1}, + object="ocr", + ), ) - assert result == {"extract_header": True} + response = litellm.ocr( + model="mistral/mistral-ocr-latest", + document=DOCUMENT, + api_key="test-key", + pages=[0], + include_image_base64=False, + ) + + assert response.pages[0].markdown == "python" + assert fake_client.requests == []