Simplify rust ocr entrypoint

This commit is contained in:
Ishaan Jaffer 2026-06-22 17:53:18 -07:00
parent 6a135105eb
commit fea5204a2f
No known key found for this signature in database
11 changed files with 395 additions and 348 deletions

View file

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

View file

@ -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 {}),
}

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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