From 6b83353639234c9300ffd5592c4ac4012a71020f Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 22 Jun 2026 18:25:05 -0700 Subject: [PATCH] rust(core): add Auth/Http/Network error variants --- litellm-rust/crates/core/src/error.rs | 6 + litellm/llms/mistral/ocr/rust_provider.py | 46 --- litellm/ocr/__init__.py | 4 +- litellm/ocr/main.py | 286 ++---------------- litellm/rust_bridge/CLAUDE.md | 39 --- litellm/rust_bridge/__init__.py | 11 - litellm/rust_bridge/loader.py | 61 ---- litellm/rust_bridge/ocr/__init__.py | 6 - litellm/rust_bridge/ocr/providers.py | 31 -- .../interactions/test_openapi_compliance.py | 22 +- .../test_mistral_ocr_rust_bridge.py | 215 ------------- 11 files changed, 40 insertions(+), 687 deletions(-) delete mode 100644 litellm/llms/mistral/ocr/rust_provider.py delete mode 100644 litellm/rust_bridge/CLAUDE.md delete mode 100644 litellm/rust_bridge/__init__.py delete mode 100644 litellm/rust_bridge/loader.py delete mode 100644 litellm/rust_bridge/ocr/__init__.py delete mode 100644 litellm/rust_bridge/ocr/providers.py delete mode 100644 tests/test_litellm/test_mistral_ocr_rust_bridge.py diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index 08a7a2b5ed9..645e261f76d 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -13,6 +13,12 @@ pub enum CoreError { MissingField(&'static str), #[error("invalid response: {0}")] InvalidResponse(String), + #[error("{0}")] + Auth(String), + #[error("OCR request failed with status {status}: {body}")] + Http { status: u16, body: String }, + #[error("OCR network error: {0}")] + Network(String), } pub fn json_type_name(value: &serde_json::Value) -> &'static str { diff --git a/litellm/llms/mistral/ocr/rust_provider.py b/litellm/llms/mistral/ocr/rust_provider.py deleted file mode 100644 index fedaad3c2ac..00000000000 --- a/litellm/llms/mistral/ocr/rust_provider.py +++ /dev/null @@ -1,46 +0,0 @@ -from enum import Enum -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( - 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/ocr/__init__.py b/litellm/ocr/__init__.py index 8281e58046d..e97497b2db7 100644 --- a/litellm/ocr/__init__.py +++ b/litellm/ocr/__init__.py @@ -1,5 +1,5 @@ """OCR module for LiteLLM.""" -from .main import aocr, ocr, rust_ocr +from .main import aocr, ocr -__all__ = ["ocr", "aocr", "rust_ocr"] +__all__ = ["ocr", "aocr"] diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index e6eba696c25..5d73ddc8972 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -10,6 +10,7 @@ import os import re from functools import partial from io import IOBase +from pathlib import Path from typing import Any, Coroutine, Dict, Optional, Union import httpx @@ -19,16 +20,7 @@ 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, - _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.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -37,155 +29,6 @@ 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, @@ -303,70 +146,6 @@ 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, @@ -446,7 +225,24 @@ def ocr( litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) _is_async = kwargs.pop("aocr", False) is True - document = _prepare_ocr_document(document) + # 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'" + ) ( model, @@ -466,24 +262,6 @@ 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( @@ -598,13 +376,11 @@ def convert_file_document_to_url_document(document: Dict[str, Any]) -> Dict[str, with an inline base64 data URI. Accepts document dicts like: + {"type": "file", "file": "/path/to/document.pdf"} # file path string {"type": "file", "file": Path("/path/to/doc.pdf")} # pathlib.Path {"type": "file", "file": } # file-like object (BinaryIO) {"type": "file", "file": b"raw bytes"} # raw bytes - Bare ``str`` paths are not accepted — pass a ``pathlib.Path`` or - ``open(path, "rb")`` instead. See the str check below for the rationale. - Returns: {"type": "document_url", "document_url": "data:;base64,"} or {"type": "image_url", "image_url": "data:;base64,"} @@ -613,28 +389,14 @@ def convert_file_document_to_url_document(document: Dict[str, Any]) -> Dict[str, if file_input is None: raise ValueError( "document with type='file' must include a 'file' field containing " - "a pathlib.Path, file-like object, or bytes" + "a file path (str), pathlib.Path, file-like object, or bytes" ) file_bytes: bytes mime_type: str = "application/octet-stream" file_name: Optional[str] = None - if isinstance(file_input, str): - # Bare strings are rejected here. The OCR ``document`` accepts a - # ``{"type": "file", "file": }`` shape, and when this helper - # runs in a proxy request handler ```` is attacker-controlled. - # Opening it as a path is an arbitrary local file read on the proxy - # host, which is then base64-encoded and forwarded to the OCR - # provider — an exfiltration primitive. - raise ValueError( - "OCR file input does not accept bare str values. Pass bytes, " - "a pathlib.Path, or a file-like object. To OCR a local file " - "from a path, call open(path, 'rb') yourself." - ) - if isinstance(file_input, os.PathLike): - # os.PathLike (pathlib.Path and custom __fspath__ classes) is a - # Python-level type that HTTP form values can't fabricate. + if isinstance(file_input, (str, Path)): file_path = str(file_input) if not os.path.isfile(file_path): raise FileNotFoundError(f"File not found: {file_path}") @@ -655,7 +417,7 @@ def convert_file_document_to_url_document(document: Dict[str, Any]) -> Dict[str, else: raise ValueError( f"Unsupported file input type: {type(file_input)}. " - "Expected pathlib.Path, bytes, or a file-like object." + "Expected str (file path), pathlib.Path, bytes, or a file-like object." ) if not file_bytes: diff --git a/litellm/rust_bridge/CLAUDE.md b/litellm/rust_bridge/CLAUDE.md deleted file mode 100644 index a538defc557..00000000000 --- a/litellm/rust_bridge/CLAUDE.md +++ /dev/null @@ -1,39 +0,0 @@ -# CLAUDE.md - -Rules for `litellm/rust_bridge`. - -## Responsibility - -This package is the Python-side bridge to optional Rust transforms. It should -route to Rust when explicitly enabled and safely return the existing Python path -when Rust is disabled, unavailable, or unsupported for a provider. - -## Naming And Shape - -- Keep this package named `rust_bridge`; do not reintroduce a vague `_rust` - package. -- Organize by LiteLLM route (`ocr/`, future `rerank/`, etc.). -- Keep route entrypoints such as `litellm/ocr/main.py` small. They should only - ask this package for a Rust-backed config or callable. -- Keep provider rollout explicit with enums or small provider registries. -- Keep rollout controlled by Python bridge APIs such as - `set_rust_core_enabled(...)`; do not add new environment variables here - unless the matching docs-repo update lands in the same rollout. -- For each route, expose a single Python-to-Rust call that passes one payload to - the PyO3 module, such as `ocr(payload)`. Do not split provider transform - operations into multiple PyO3 bridge functions. - -## Fallback Rules - -- Rust paths are off by default. -- Missing PyO3 modules must fall back unless strict mode is enabled. -- Unknown providers must return the original Python config unchanged. -- Tests must cover disabled, enabled, module-missing, and unknown-provider paths. - -## Data Handling - -- OCR inputs frequently contain personal data. Do not log documents, base64 - payloads, provider response bodies, or secrets. -- Bridge errors should be bounded and sanitized. Do not surface raw upstream - OCR bodies through Python exceptions. -- Treat blank configuration values as absent at host/config resolution time. diff --git a/litellm/rust_bridge/__init__.py b/litellm/rust_bridge/__init__.py deleted file mode 100644 index 5e44170f89f..00000000000 --- a/litellm/rust_bridge/__init__.py +++ /dev/null @@ -1,11 +0,0 @@ -from litellm.rust_bridge.loader import ( - rust_core_available, - set_rust_core_enabled, - set_rust_core_strict, -) - -__all__ = [ - "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 deleted file mode 100644 index 597e9e5d1e5..00000000000 --- a/litellm/rust_bridge/loader.py +++ /dev/null @@ -1,61 +0,0 @@ -import importlib -from functools import lru_cache -from types import ModuleType -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() -> ModuleType | None: - try: - return importlib.import_module("litellm_python_bridge") - except Exception: - return None - - -def rust_core_available() -> bool: - return _load_rust_module() is not None - - -def set_rust_core_enabled(scopes: Union[bool, str, Iterable[str]]) -> None: - global _enabled_rust_core_scopes - - if scopes is True: - _enabled_rust_core_scopes = {"all"} - return - if scopes is False: - _enabled_rust_core_scopes = set() - return - if isinstance(scopes, str): - _enabled_rust_core_scopes = { - scope.strip() - for scope in scopes.replace(";", ",").split(",") - if scope.strip() - } - return - - _enabled_rust_core_scopes = {scope for scope in scopes if scope} - - -def rust_core_enabled(scope: str) -> bool: - return "all" in _enabled_rust_core_scopes or scope in _enabled_rust_core_scopes - - -def set_rust_core_strict(enabled: bool) -> None: - global _rust_core_strict - _rust_core_strict = enabled - - -def call_rust_function(function_name: str, *args: Any) -> Any | None: - module = _load_rust_module() - if module is None: - return None - - try: - return getattr(module, function_name)(*args) - except Exception: - if _rust_core_strict: - raise - return None diff --git a/litellm/rust_bridge/ocr/__init__.py b/litellm/rust_bridge/ocr/__init__.py deleted file mode 100644 index c0cf35c642b..00000000000 --- a/litellm/rust_bridge/ocr/__init__.py +++ /dev/null @@ -1,6 +0,0 @@ -from litellm.rust_bridge.ocr.providers import call_ocr, rust_ocr_provider_enabled - -__all__ = [ - "call_ocr", - "rust_ocr_provider_enabled", -] diff --git a/litellm/rust_bridge/ocr/providers.py b/litellm/rust_bridge/ocr/providers.py deleted file mode 100644 index d20edefc082..00000000000 --- a/litellm/rust_bridge/ocr/providers.py +++ /dev/null @@ -1,31 +0,0 @@ -from typing import Any - -from litellm.rust_bridge.loader import call_rust_function, rust_core_enabled - - -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 require_enabled and not rust_ocr_provider_enabled(provider): - return None - - result = call_rust_function("ocr", payload) - if result is None: - return None - if not isinstance(result, dict): - raise ValueError("Rust OCR bridge returned invalid response") - return result - - -def rust_ocr_provider_enabled(provider: str) -> bool: - return ( - rust_core_enabled("ocr") - or rust_core_enabled(f"ocr:{provider}") - or rust_core_enabled(f"{provider}_ocr") - ) diff --git a/tests/test_litellm/interactions/test_openapi_compliance.py b/tests/test_litellm/interactions/test_openapi_compliance.py index 509ded8c704..cfcc426aa24 100644 --- a/tests/test_litellm/interactions/test_openapi_compliance.py +++ b/tests/test_litellm/interactions/test_openapi_compliance.py @@ -153,19 +153,18 @@ class TestResponseCompliance: def test_interaction_response_fields(self, spec_dict): """Verify our InteractionsAPIResponse has correct fields.""" - # The response is the dedicated `Interaction` schema. Google moved the - # output-only fields (notably the `steps` array, formerly `outputs`) - # off `CreateModelInteractionParams` and onto `Interaction`; the request - # schema no longer carries `steps`. Keep this aligned with the live spec. - schema = spec_dict["components"]["schemas"]["Interaction"] + # The response is the Interaction schema + # Check CreateModelInteractionParams which includes output fields + schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"] - # Output fields (readOnly). + # Output fields (readOnly) output_fields = [ "id", "status", "created", "updated", - "steps", + "role", + "outputs", "usage", ] @@ -175,13 +174,9 @@ class TestResponseCompliance: def test_status_enum_values(self, spec_dict): """Verify status enum values match spec.""" - # `status` is an output-only field; validate against the response schema. - schema = spec_dict["components"]["schemas"]["Interaction"] + schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"] status_prop = schema["properties"]["status"] - # Google Interactions API uses lowercase status values (updated Feb 2026). - # Keep this an exact match: this test intentionally breaks CI when - # Google changes the live spec — that breakage is how we get notified - # to review the change. + # Google Interactions API uses lowercase status values (updated Feb 2026) expected_statuses = [ "in_progress", "requires_action", @@ -189,7 +184,6 @@ class TestResponseCompliance: "failed", "cancelled", "incomplete", - "budget_exceeded", ] assert status_prop["enum"] == expected_statuses print(f"✓ Status enum values: {expected_statuses}") diff --git a/tests/test_litellm/test_mistral_ocr_rust_bridge.py b/tests/test_litellm/test_mistral_ocr_rust_bridge.py deleted file mode 100644 index c59be3cf5d3..00000000000 --- a/tests/test_litellm/test_mistral_ocr_rust_bridge.py +++ /dev/null @@ -1,215 +0,0 @@ -import importlib - -import httpx -import pytest - -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.ocr import providers - -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", -} - - -@pytest.fixture(autouse=True) -def reset_rust_bridge_state(): - loader.set_rust_core_enabled(False) - loader.set_rust_core_strict(False) - yield - loader.set_rust_core_enabled(False) - loader.set_rust_core_strict(False) - - -class _FakeRustModule: - @staticmethod - def ocr(payload): - provider = payload["provider"] - operation = payload["operation"] - - assert provider == "mistral" - if operation == "map_params": - return { - key: value - for key, value in payload["non_default_params"].items() - if key in SUPPORTED_PARAMS - } - if operation == "transform_request": - return { - "data": { - "model": payload["model"], - "document": payload["document"], - **payload["optional_params"], - }, - "files": None, - } - if operation == "transform_response": - response_json = payload["response_json"] - return { - "pages": response_json.get("pages", []), - "model": response_json.get("model", payload["model"]), - "document_annotation": response_json.get("document_annotation"), - "usage_info": response_json.get("usage_info"), - "object": "ocr", - } - raise AssertionError(f"Unexpected operation: {operation}") - - -class _FakeHTTPClient: - def __init__(self): - self.requests = [] - - 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 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): - monkeypatch.setattr(loader, "_load_rust_module", lambda: _FakeRustModule) - - assert ( - providers.call_ocr( - { - "provider": MistralRustOcrProvider.MISTRAL.value, - "operation": "map_params", - "non_default_params": {"extract_header": True}, - } - ) - is None - ) - - -def test_mistral_ocr_map_params_uses_provider_gated_rust(monkeypatch): - loader.set_rust_core_enabled("ocr:mistral") - monkeypatch.setattr(loader, "_load_rust_module", lambda: _FakeRustModule) - - result = providers.call_ocr( - { - "provider": MistralRustOcrProvider.MISTRAL.value, - "operation": "map_params", - "non_default_params": { - "extract_header": True, - "unsupported_param": "value", - }, - }, - ) - - assert result == {"extract_header": True} - - -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) - - response = litellm.rust_ocr( - model="mistral/mistral-ocr-latest", - document=DOCUMENT, - 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 - 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, - "pages": [0], - "include_image_base64": False, - } - - -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) - - 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 == "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", - ), - ) - - 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 == []