refactor(ocr): dispatch through Rust before Python preparation

This commit is contained in:
Yujong Lee 2026-08-30 16:33:22 -07:00
parent 0d16842e67
commit c13f0aeedc
9 changed files with 491 additions and 358 deletions

View file

@ -12,7 +12,7 @@ use litellm_core::chat_completions::{
use litellm_core::error::CoreError;
use litellm_core::messages::messages as run_messages;
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
use litellm_runtime::ocr::{OcrRequest, ocr as run_ocr};
use litellm_runtime::ocr::{OcrRequest, ocr as run_ocr, ocr_decline_reason};
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::{PyAny, PyDict};
@ -208,6 +208,29 @@ fn marshal_inputs(
Ok((document, extra_headers, optional_params, timeout))
}
#[pyfunction]
#[pyo3(signature = (model, custom_llm_provider=None, optional_params=None))]
fn ocr_decline(
py: Python<'_>,
model: String,
custom_llm_provider: Option<String>,
optional_params: Option<Py<PyAny>>,
) -> PyResult<Option<String>> {
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
Ok(ocr_decline_reason(OcrRequest {
model: &model,
document: Value::Null,
api_key: None,
api_base: None,
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers: None,
optional_params,
timeout: None,
litellm_call_id: Some("ocr-decline"),
})
.map(|reason| reason.to_string()))
}
#[pyfunction]
#[pyo3(signature = (model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
@ -613,6 +636,7 @@ fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> {
let py = module.py();
module.add_function(wrap_pyfunction!(ocr, module)?)?;
module.add_function(wrap_pyfunction!(aocr, module)?)?;
module.add_function(wrap_pyfunction!(ocr_decline, module)?)?;
module.add_function(wrap_pyfunction!(transcription, module)?)?;
module.add_function(wrap_pyfunction!(atranscription, module)?)?;
module.add_function(wrap_pyfunction!(messages, module)?)?;

View file

@ -14,8 +14,8 @@ mod prepare;
mod provider;
mod types;
pub use prepare::{prepare_ocr_request, prepare_provider_request};
pub use types::{OcrRequest, PreparedOcrRequest, ProviderOcrRequest};
pub use prepare::{ocr_decline_reason, prepare_ocr_request, prepare_provider_request};
pub use types::{OcrDeclineReason, OcrRequest, PreparedOcrRequest, ProviderOcrRequest};
use handler::execute_ocr_provider_call;

View file

@ -4,21 +4,35 @@ use std::time::{SystemTime, UNIX_EPOCH};
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use serde_json::Value;
use super::common_utils::{convert_document_url_to_data_uri, has_header, string_headers};
use super::provider::ocr_provider_config;
use super::types::{OcrAuthStrategy, OcrRequest, PreparedOcrRequest, ProviderOcrRequest};
use super::types::{
OcrAuthStrategy, OcrDeclineReason, OcrRequest, PreparedOcrRequest, ProviderOcrRequest,
};
pub fn ocr_decline_reason(request: OcrRequest<'_>) -> Option<OcrDeclineReason> {
if request
.optional_params
.get("req_format")
.and_then(Value::as_str)
== Some("native")
{
return Some(OcrDeclineReason::NativeRequestFormat);
}
let provider = select_ocr_provider(request.model, request.custom_llm_provider);
ocr_provider_config(provider.custom_llm_provider, provider.model)
.is_none()
.then_some(OcrDeclineReason::UnsupportedProvider)
}
pub fn prepare_ocr_request(request: OcrRequest<'_>) -> PreparedOcrRequest {
let litellm_call_id = request
.litellm_call_id
.map(str::to_string)
.unwrap_or_else(new_ocr_call_id);
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.unwrap_or(CustomLlmProvider {
model: request.model,
custom_llm_provider: "mistral",
});
let provider_info = select_ocr_provider(request.model, request.custom_llm_provider);
PreparedOcrRequest {
model: provider_info.model.to_string(),
@ -33,6 +47,16 @@ pub fn prepare_ocr_request(request: OcrRequest<'_>) -> PreparedOcrRequest {
}
}
fn select_ocr_provider<'a>(
model: &'a str,
custom_llm_provider: Option<&'a str>,
) -> CustomLlmProvider<'a> {
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
model,
custom_llm_provider: "mistral",
})
}
pub async fn prepare_provider_request(
request: PreparedOcrRequest,
) -> CoreResult<ProviderOcrRequest> {
@ -102,3 +126,53 @@ fn new_ocr_call_id() -> String {
.unwrap_or(0);
format!("ocr-{timestamp}-{sequence}")
}
#[cfg(test)]
mod tests {
use serde_json::{Map, json};
use super::*;
fn request<'a>(
model: &'a str,
custom_llm_provider: Option<&'a str>,
optional_params: Map<String, Value>,
) -> OcrRequest<'a> {
OcrRequest {
model,
document: json!({}),
api_key: None,
api_base: None,
custom_llm_provider,
extra_headers: None,
optional_params,
timeout: None,
litellm_call_id: Some("decline-test"),
}
}
#[test]
fn decline_uses_runtime_provider_selection_without_credentials() {
assert_eq!(
ocr_decline_reason(request("mistral/mistral-ocr-latest", None, Map::new())),
None
);
assert_eq!(
ocr_decline_reason(request("gpt-4o", Some("openai"), Map::new())),
Some(OcrDeclineReason::UnsupportedProvider)
);
}
#[test]
fn decline_rejects_python_only_native_format() {
let optional_params = Map::from_iter([("req_format".to_string(), json!("native"))]);
assert_eq!(
ocr_decline_reason(request(
"azure_ai/doc-intelligence/prebuilt-layout",
None,
optional_params,
)),
Some(OcrDeclineReason::NativeRequestFormat)
);
}
}

View file

@ -4,6 +4,25 @@ use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest};
use litellm_core::ocr::transformation::OcrProviderTransformation;
use serde_json::{Map, Value};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OcrDeclineReason {
NativeRequestFormat,
UnsupportedProvider,
}
impl std::fmt::Display for OcrDeclineReason {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NativeRequestFormat => {
formatter.write_str("native OCR response format requires Python")
}
Self::UnsupportedProvider => {
formatter.write_str("OCR provider is not supported by the Rust runtime")
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OcrAuthStrategy {
Bearer,

View file

@ -7,7 +7,7 @@ import base64
import mimetypes
import os
import re
from collections.abc import Callable, Coroutine, Mapping
from collections.abc import Coroutine, Mapping
from dataclasses import dataclass
from io import IOBase
from typing import Any, Final, cast
@ -52,21 +52,6 @@ class _PreparedOCRRequest:
litellm_logging_obj: LiteLLMLoggingObj
@dataclass
class _PreparedRustOCRCall:
api_key: str | None
api_base: str | None
headers: dict[str, object]
optional_params: dict[str, object]
_RUST_OCR_PROVIDERS: Final = {
"mistral",
"azure_ai",
"vertex_ai",
}
def _prepare_ocr_request(
model: str,
document: Mapping[str, object],
@ -188,115 +173,59 @@ def _prepare_ocr_request(
)
def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool:
if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native":
return False
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
def _rust_bridge_optional_params(
prepared_request: _PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
) -> dict[str, object]:
optional_params: Final = dict(prepared_request.optional_params)
if prepared_request.custom_llm_provider == "vertex_ai":
vertex_project: Final = (
prepared_request.litellm_params.get("vertex_project")
or prepared_request.litellm_params.get("vertex_ai_project")
or litellm.vertex_project
or resolve_secret("VERTEXAI_PROJECT")
def _materialize_ocr_document(document: Mapping[str, object]) -> dict[str, object]:
if not isinstance(document, dict):
raise ValueError(f"document must be a dict with 'type' and URL/file field, got {type(document)}")
if document.get("type") == "file":
return cast(dict[str, object], convert_file_document_to_url_document(document))
if document.get("type") not in ("document_url", "image_url"):
raise ValueError(
f"Invalid document type: {document.get('type')}. Must be 'document_url', 'image_url', or 'file'"
)
vertex_location: Final = (
prepared_request.litellm_params.get("vertex_location")
or prepared_request.litellm_params.get("vertex_ai_location")
or litellm.vertex_location
or resolve_secret("VERTEXAI_LOCATION")
or resolve_secret("VERTEX_LOCATION")
)
if vertex_project is not None:
optional_params["vertex_project"] = vertex_project
if vertex_location is not None:
optional_params["vertex_location"] = vertex_location
return optional_params
return document
def _rust_bridge_api_base(
prepared_request: _PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
) -> str | None:
if prepared_request.api_base is not None:
return prepared_request.api_base
if prepared_request.custom_llm_provider == "azure_ai":
if is_azure_document_intelligence_model(prepared_request.model):
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
return resolve_secret("AZURE_AI_API_BASE")
return None
def _native_ocr_params(kwargs: Mapping[str, object]) -> dict[str, object]:
return {key: value for key, value in kwargs.items() if key != "litellm_logging_obj"}
def _prepare_rust_ocr_call(
prepared_request: _PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> _PreparedRustOCRCall:
provider_config: Final = prepared_request.provider_config
api_key_env_var: Final = provider_config.get_api_key_env_var()
resolved_api_key: Final = prepared_request.api_key or (
resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
)
resolved_headers: Final = provider_config.validate_environment(
headers=prepared_request.extra_headers or {},
model=prepared_request.model,
api_key=resolved_api_key,
api_base=prepared_request.api_base,
litellm_params=prepared_request.litellm_params,
)
resolved_complete_url: Final = provider_config.get_complete_url(
api_base=prepared_request.api_base,
model=prepared_request.model,
optional_params=prepared_request.optional_params,
litellm_params=prepared_request.litellm_params,
)
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
prepared_request.litellm_logging_obj.pre_call(
input="OCR document processing",
api_key=resolved_api_key,
additional_args={
"complete_input_dict": {
"model": prepared_request.model,
"document": prepared_request.document,
**rust_optional_params,
},
"api_base": resolved_complete_url,
"headers": resolved_headers,
},
)
return _PreparedRustOCRCall(
api_key=resolved_api_key,
api_base=rust_api_base,
headers=cast(dict[str, object], resolved_headers),
optional_params=rust_optional_params,
)
def _validate_ocr_request_format(
model: str,
custom_llm_provider: str | None,
kwargs: Mapping[str, object],
) -> None:
requested_format: Final = kwargs.get(OCR_REQUEST_FORMAT_PARAM)
if requested_format is not None:
try:
parse_ocr_request_format(requested_format)
except ValueError as error:
raise litellm.exceptions.UnsupportedParamsError(
message=str(error),
model=model,
llm_provider=custom_llm_provider or "",
) from error
def _run_rust_ocr(
prepared_request: _PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
*,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> OCRResponse | None:
if rust_ocr_bridge.load_rust_ocr() is None:
return None
prepared: Final = _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
)
rust_response: Final = rust_ocr_bridge.ocr(
model=prepared_request.model,
document=prepared_request.document,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
extra_headers=prepared.headers,
optional_params=prepared.optional_params,
timeout=prepared_request.effective_timeout,
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
)
if rust_response is None:
return None
@ -304,24 +233,25 @@ def _run_rust_ocr(
async def _run_rust_aocr(
prepared_request: _PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
*,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> OCRResponse | None:
if rust_ocr_bridge.load_rust_aocr() is None:
return None
prepared: Final = _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
)
rust_response: Final = await rust_ocr_bridge.aocr(
model=prepared_request.model,
document=prepared_request.document,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
extra_headers=prepared.headers,
optional_params=prepared.optional_params,
timeout=prepared_request.effective_timeout,
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
)
if rust_response is None:
return None
@ -408,9 +338,40 @@ async def aocr(
"kwargs": kwargs,
}
try:
_validate_ocr_request_format(model, custom_llm_provider, kwargs)
native_document: Final = _materialize_ocr_document(document)
native_params: Final = _native_ocr_params(kwargs)
native_available: Final = (
rust_ocr_bridge.rust_ocr_enabled()
and rust_ocr_bridge.load_rust_aocr() is not None
and rust_ocr_bridge.load_rust_ocr_decline() is not None
)
if native_available:
decline_reason: Final = rust_ocr_bridge.ocr_decline(
model=model,
custom_llm_provider=custom_llm_provider,
optional_params=native_params,
)
if decline_reason is None:
rust_response: Final = await _run_rust_aocr(
model=model,
document=native_document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=native_params,
timeout=timeout,
)
if rust_response is not None:
return rust_response
verbose_logger.debug("Async Rust OCR bridge unavailable; falling back to Python path")
else:
verbose_logger.debug("Async Rust OCR declined request: %s", decline_reason)
prepared: Final = _prepare_ocr_request(
model=model,
document=document,
document=native_document,
api_key=api_key,
api_base=api_base,
timeout=timeout,
@ -422,18 +383,6 @@ async def aocr(
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
if _rust_ocr_supported(prepared) and rust_ocr_bridge.rust_ocr_enabled():
from litellm.secret_managers.main import get_secret_str
rust_response: Final = await _run_rust_aocr(
prepared_request=prepared,
resolve_api_key=get_secret_str,
)
if rust_response is None:
verbose_logger.debug("Async Rust OCR bridge unavailable; falling back to Python path")
else:
return rust_response
response = base_llm_http_handler.ocr(
model=prepared.model,
document=prepared.document,
@ -680,9 +629,40 @@ def ocr(
try:
_is_async: Final = kwargs.pop("aocr", False) is True
completion_kwargs["aocr"] = _is_async
_validate_ocr_request_format(model, custom_llm_provider, kwargs)
native_document: Final = _materialize_ocr_document(document)
native_params: Final = _native_ocr_params(kwargs)
native_available: Final = (
rust_ocr_bridge.rust_ocr_enabled()
and rust_ocr_bridge.load_rust_ocr() is not None
and rust_ocr_bridge.load_rust_ocr_decline() is not None
)
if native_available:
decline_reason: Final = rust_ocr_bridge.ocr_decline(
model=model,
custom_llm_provider=custom_llm_provider,
optional_params=native_params,
)
if decline_reason is None:
rust_response: Final = _run_rust_ocr(
model=model,
document=native_document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=native_params,
timeout=timeout,
)
if rust_response is not None:
return rust_response
verbose_logger.debug("Rust OCR bridge unavailable; falling back to Python path")
else:
verbose_logger.debug("Rust OCR declined request: %s", decline_reason)
prepared: Final = _prepare_ocr_request(
model=model,
document=document,
document=native_document,
api_key=api_key,
api_base=api_base,
kwargs=kwargs,
@ -694,18 +674,6 @@ def ocr(
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
if _rust_ocr_supported(prepared) and rust_ocr_bridge.rust_ocr_enabled():
from litellm.secret_managers.main import get_secret_str
rust_response: Final = _run_rust_ocr(
prepared_request=prepared,
resolve_api_key=get_secret_str,
)
if rust_response is None:
verbose_logger.debug("Rust OCR bridge unavailable; falling back to Python path")
else:
return rust_response
response: Final = base_llm_http_handler.ocr(
model=prepared.model,
document=prepared.document,

View file

@ -44,6 +44,16 @@ class RustAocr(Protocol):
raise NotImplementedError
class RustOcrDecline(Protocol):
def __call__(
self,
model: str,
custom_llm_provider: str | None,
optional_params: dict[str, object],
) -> str | None:
raise NotImplementedError
class _Unset:
pass
@ -63,6 +73,7 @@ def _env_enables_rust_ocr() -> bool:
_rust_ocr_enabled = _env_enables_rust_ocr()
_rust_ocr_impl: RustOcr | None = None
_rust_aocr_impl: RustAocr | None = None
_rust_ocr_decline_impl: RustOcrDecline | None = None
def use_litellm_rust(
@ -70,14 +81,17 @@ def use_litellm_rust(
*,
ocr: RustOcr | None | _Unset = _UNSET,
aocr: RustAocr | None | _Unset = _UNSET,
ocr_decline: RustOcrDecline | None | _Unset = _UNSET,
messages: RustMessages | None | _Unset = _UNSET,
amessages: RustAmessages | None | _Unset = _UNSET,
responses_websocket: Any | None | _Unset = _UNSET,
transcription: Any | None | _Unset = _UNSET,
atranscription: Any | None | _Unset = _UNSET,
) -> None:
global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl
configuring_ocr: Final = not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset)
global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl, _rust_ocr_decline_impl
configuring_ocr: Final = (
not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset) or not isinstance(ocr_decline, _Unset)
)
configuring_messages: Final = not isinstance(messages, _Unset) or not isinstance(amessages, _Unset)
configuring_responses_websocket: Final = not isinstance(responses_websocket, _Unset)
configuring_transcription: Final = not isinstance(transcription, _Unset) or not isinstance(atranscription, _Unset)
@ -87,6 +101,8 @@ def use_litellm_rust(
_rust_ocr_impl = ocr
if not isinstance(aocr, _Unset):
_rust_aocr_impl = aocr
if not isinstance(ocr_decline, _Unset):
_rust_ocr_decline_impl = ocr_decline
if configuring_transcription:
from litellm.rust_bridge.transcription import configure_rust_transcription
@ -138,6 +154,33 @@ def load_rust_aocr() -> RustAocr | None:
return cast(RustAocr, getattr(native_bridge, "aocr", None))
def load_rust_ocr_decline() -> RustOcrDecline | None:
if _rust_ocr_decline_impl is not None:
return _rust_ocr_decline_impl
from litellm.rust_bridge import get_native_bridge
native_bridge: Final = get_native_bridge()
if native_bridge is None:
return None
return cast(RustOcrDecline, getattr(native_bridge, "ocr_decline", None))
def ocr_decline(
*,
model: str,
custom_llm_provider: str | None,
optional_params: dict[str, object],
) -> str | None:
decline: Final = load_rust_ocr_decline()
if decline is None:
return "Rust OCR decline bridge is unavailable"
return decline(
model=model,
custom_llm_provider=custom_llm_provider,
optional_params=optional_params,
)
def ocr(
*,
model: str,

View file

@ -11,10 +11,9 @@ supplied api_base is always honoured.
from litellm.llms.azure_ai.ocr.common_utils import (
is_azure_document_intelligence_model,
)
from litellm.ocr.main import _prepare_ocr_request, _rust_bridge_api_base
from litellm.ocr.main import _prepare_ocr_request
_DOC = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
_DOC_INTELLIGENCE_ENDPOINT = "https://di.cognitiveservices.azure.com"
_AZURE_AI_API_BASE = "https://generic-azure-ai.example.com"
@ -23,13 +22,6 @@ class _FakeLogging:
return None
def _resolve_secret(name: str) -> str | None:
return {
"AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT": _DOC_INTELLIGENCE_ENDPOINT,
"AZURE_AI_API_BASE": _AZURE_AI_API_BASE,
}.get(name)
def _prepare(model: str, api_base: str | None):
return _prepare_ocr_request(
model=model,
@ -64,7 +56,6 @@ class TestDocIntelligenceApiBaseResolution:
prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", None)
assert prepared.api_base is None
assert _rust_bridge_api_base(prepared, _resolve_secret) == _DOC_INTELLIGENCE_ENDPOINT
def test_explicit_api_base_is_honoured_for_doc_intelligence(self, monkeypatch):
"""A caller-supplied api_base must always win, even for doc-intelligence."""
@ -74,7 +65,6 @@ class TestDocIntelligenceApiBaseResolution:
prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", custom)
assert prepared.api_base == custom
assert _rust_bridge_api_base(prepared, _resolve_secret) == custom
def test_generic_azure_ai_base_still_applies_to_mistral_ocr(self, monkeypatch):
"""Non doc-intelligence azure_ai models keep using AZURE_AI_API_BASE."""

View file

@ -4,39 +4,32 @@ providers that don't support a native response must reject it, and the Rust
bridge (which only returns the normalized shape) must not serve native requests.
"""
from unittest.mock import MagicMock
import pytest
import litellm
from litellm.ocr.main import _PreparedOCRRequest, _rust_ocr_supported
from litellm.rust_bridge import ocr as rust_ocr_bridge
DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
def _prepared(optional_params: dict[str, object]) -> _PreparedOCRRequest:
return _PreparedOCRRequest(
model="doc-intelligence/prebuilt-layout",
document=dict(DOCUMENT),
api_key="fake-key",
api_base="https://example.cognitiveservices.azure.com",
custom_llm_provider="azure_ai",
extra_headers=None,
provider_config=MagicMock(),
optional_params=optional_params,
litellm_params={},
effective_timeout=60.0,
litellm_logging_obj=MagicMock(),
def test_native_decline_wrapper_uses_runtime_result():
rust_ocr_bridge.use_litellm_rust(
True,
ocr_decline=lambda model, custom_llm_provider, optional_params: (
"native OCR response format requires Python" if optional_params.get("req_format") == "native" else None
),
)
@pytest.mark.parametrize("optional_params", [{}, {"req_format": "litellm"}])
def test_rust_ocr_serves_default_format(optional_params):
assert _rust_ocr_supported(_prepared(optional_params)) is True
def test_rust_ocr_skipped_for_native_format():
assert _rust_ocr_supported(_prepared({"req_format": "native"})) is False
try:
assert (
rust_ocr_bridge.ocr_decline(
model="azure_ai/doc-intelligence/prebuilt-layout",
custom_llm_provider=None,
optional_params={"req_format": "native"},
)
== "native OCR response format requires Python"
)
finally:
rust_ocr_bridge.use_litellm_rust(False, ocr_decline=None)
@pytest.mark.asyncio

View file

@ -1,9 +1,10 @@
"""Tests for the optional Rust-backed OCR path."""
import importlib
import builtins
import importlib
import types
from typing import Any
from unittest.mock import MagicMock
import httpx
import pytest
@ -101,6 +102,27 @@ class RecordingAsyncBridge:
return dict(FAKE_OCR_RESPONSE)
class RecordingDecline:
def __init__(self, reason: str | None = None) -> None:
self.reason = reason
self.calls: list[dict[str, object]] = []
def __call__(
self,
model: str,
custom_llm_provider: str | None,
optional_params: dict[str, object],
) -> str | None:
self.calls.append(
{
"model": model,
"custom_llm_provider": custom_llm_provider,
"optional_params": optional_params,
}
)
return self.reason
class RaisingBridge:
def __call__(
self,
@ -214,10 +236,10 @@ def build_prepared_request(
@pytest.fixture(autouse=True)
def _reset_rust_flag():
"""Keep the global toggle isolated between tests."""
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None)
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None, ocr_decline=None)
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
yield
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None)
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None, ocr_decline=None)
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
@ -225,7 +247,7 @@ def _reset_rust_flag():
def fake_bridge():
"""Enable the Rust path with an injected recording bridge (no native wheel)."""
bridge = RecordingBridge()
litellm.use_litellm_rust(True, ocr=bridge)
litellm.use_litellm_rust(True, ocr=bridge, ocr_decline=RecordingDecline())
return bridge
@ -233,7 +255,7 @@ def fake_bridge():
def fake_async_bridge():
"""Enable the async Rust path with an injected recording bridge."""
bridge = RecordingAsyncBridge()
litellm.use_litellm_rust(True, aocr=bridge)
litellm.use_litellm_rust(True, aocr=bridge, ocr_decline=RecordingDecline())
return bridge
@ -434,90 +456,85 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
def test_run_rust_ocr_prepares_request_and_wraps_response():
bridge = RecordingBridge()
logging_obj = RecordingLogging()
litellm.use_litellm_rust(True, ocr=bridge)
response = ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
logging_obj=logging_obj,
api_base="https://proxy.internal",
extra_headers={"x-trace-id": "trace-1"},
optional_params={"include_image_base64": True},
timeout=12.5,
),
resolve_api_key=lambda _name: None,
model=MODEL,
document=DOCUMENT,
api_key="sk-test",
api_base="https://proxy.internal",
custom_llm_provider=None,
extra_headers={"x-trace-id": "trace-1"},
optional_params={"include_image_base64": True},
timeout=12.5,
)
assert isinstance(response, OCRResponse)
assert response.pages[0].markdown == "hello world"
assert bridge.calls[0] == {
"model": "mistral-ocr-latest",
"model": MODEL,
"document": DOCUMENT,
"api_key": "sk-test",
"api_base": "https://proxy.internal",
"custom_llm_provider": "mistral",
"extra_headers": {
"Authorization": "Bearer sk-test",
"x-trace-id": "trace-1",
},
"custom_llm_provider": None,
"extra_headers": {"x-trace-id": "trace-1"},
"optional_params": {"include_image_base64": True},
"timeout_seconds": 12.5,
}
def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
def test_run_rust_ocr_leaves_missing_key_for_runtime_resolution():
bridge = RecordingBridge()
litellm.use_litellm_rust(True, ocr=bridge)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(api_key=None, timeout=None),
resolve_api_key=lambda name: (
"sk-from-vault" if name == "MISTRAL_API_KEY" else None
),
model=MODEL,
document=DOCUMENT,
api_key=None,
api_base=None,
custom_llm_provider=None,
extra_headers=None,
optional_params={},
timeout=None,
)
assert bridge.calls[0]["api_key"] == "sk-from-vault"
assert bridge.calls[0]["api_key"] is None
def test_run_rust_ocr_prefers_explicit_key_over_resolver():
def test_run_rust_ocr_forwards_explicit_key():
bridge = RecordingBridge()
litellm.use_litellm_rust(True, ocr=bridge)
def _resolver(name: str) -> str | None:
raise AssertionError(f"resolver should not be called for {name}")
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
api_key="sk-explicit",
timeout=None,
),
resolve_api_key=_resolver,
model=MODEL,
document=DOCUMENT,
api_key="sk-explicit",
api_base=None,
custom_llm_provider=None,
extra_headers=None,
optional_params={},
timeout=None,
)
assert bridge.calls[0]["api_key"] == "sk-explicit"
def test_run_rust_ocr_uses_provider_api_key_env_var():
def test_run_rust_ocr_does_not_resolve_provider_api_key_env_var():
bridge = RecordingBridge()
resolver_calls = []
litellm.use_litellm_rust(True, ocr=bridge)
def _resolver(name):
resolver_calls.append(name)
return "sk-provider-env"
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"),
model="provider-ocr-model",
api_key=None,
timeout=None,
),
resolve_api_key=_resolver,
model="provider/provider-ocr-model",
document=DOCUMENT,
api_key=None,
api_base=None,
custom_llm_provider=None,
extra_headers=None,
optional_params={},
timeout=None,
)
assert resolver_calls == ["PROVIDER_OCR_API_KEY"]
assert bridge.calls[0]["api_key"] == "sk-provider-env"
assert bridge.calls[0]["api_key"] is None
def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
@ -525,18 +542,18 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
litellm.use_litellm_rust(True, ocr=bridge)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
custom_llm_provider="vertex_ai",
model="mistral-ocr-maas",
litellm_params={
"vertex_project": "project-1",
"vertex_location": "us-central1",
"vertex_credentials": "redacted",
},
optional_params={"include_image_base64": True},
timeout=None,
),
resolve_api_key=lambda _name: None,
model="vertex_ai/mistral-ocr-maas",
document=DOCUMENT,
api_key="sk-test",
api_base=None,
custom_llm_provider=None,
extra_headers=None,
optional_params={
"include_image_base64": True,
"vertex_project": "project-1",
"vertex_location": "us-central1",
},
timeout=None,
)
assert bridge.calls[0]["optional_params"] == {
@ -546,99 +563,85 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
}
def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager():
bridge = RecordingBridge()
litellm.use_litellm_rust(True, ocr=bridge)
def _resolver(name: str) -> str | None:
return {
"VERTEXAI_PROJECT": "project-from-secret",
"VERTEXAI_LOCATION": "us-east5",
}.get(name)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
custom_llm_provider="vertex_ai",
model="mistral-ocr-maas",
timeout=None,
),
resolve_api_key=_resolver,
)
assert bridge.calls[0]["optional_params"]["vertex_project"] == "project-from-secret"
assert bridge.calls[0]["optional_params"]["vertex_location"] == "us-east5"
def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
def test_prepare_rust_ocr_call_leaves_vertex_env_resolution_to_runtime():
bridge = RecordingBridge()
litellm.use_litellm_rust(True, ocr=bridge)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
custom_llm_provider="azure_ai",
model="pixtral-12b-2409",
api_base=None,
timeout=None,
),
resolve_api_key=lambda name: (
"https://azure.example.com" if name == "AZURE_AI_API_BASE" else None
),
model="vertex_ai/mistral-ocr-maas",
document=DOCUMENT,
api_key=None,
api_base=None,
custom_llm_provider=None,
extra_headers=None,
optional_params={},
timeout=None,
)
assert bridge.calls[0]["api_base"] == "https://azure.example.com"
assert bridge.calls[0]["optional_params"] == {}
def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
def test_prepare_rust_ocr_call_leaves_azure_ai_api_base_for_runtime_resolution():
bridge = RecordingBridge()
litellm.use_litellm_rust(True, ocr=bridge)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
custom_llm_provider="azure_ai",
model="doc-intelligence/prebuilt-layout",
api_base=None,
timeout=None,
),
resolve_api_key=lambda name: (
"https://document-intelligence.example.com"
if name == "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT"
else None
),
model="azure_ai/pixtral-12b-2409",
document=DOCUMENT,
api_key=None,
api_base=None,
custom_llm_provider=None,
extra_headers=None,
optional_params={},
timeout=None,
)
assert bridge.calls[0]["api_base"] == "https://document-intelligence.example.com"
assert bridge.calls[0]["api_base"] is None
def test_run_rust_ocr_runs_pre_call_logging():
def test_prepare_rust_ocr_call_leaves_document_intelligence_endpoint_for_runtime():
bridge = RecordingBridge()
litellm.use_litellm_rust(True, ocr=bridge)
ocr_main._run_rust_ocr(
model="azure_ai/doc-intelligence/prebuilt-layout",
document=DOCUMENT,
api_key=None,
api_base=None,
custom_llm_provider=None,
extra_headers=None,
optional_params={},
timeout=None,
)
assert bridge.calls[0]["api_base"] is None
def test_run_rust_ocr_does_not_run_python_provider_logging_preparation():
logging_obj = RecordingLogging()
bridge = RecordingBridge()
litellm.use_litellm_rust(True, ocr=bridge)
ocr_main._run_rust_ocr(
prepared_request=build_prepared_request(
logging_obj=logging_obj,
api_base="https://api.mistral.ai/v1",
extra_headers={"x-trace-id": "trace-1"},
optional_params={"include_image_base64": True},
timeout=None,
),
resolve_api_key=lambda _name: None,
model=MODEL,
document=DOCUMENT,
api_key="sk-test",
api_base=None,
custom_llm_provider=None,
extra_headers={"x-trace-id": "trace-1"},
optional_params={"include_image_base64": True},
timeout=None,
)
assert logging_obj.pre_call_kwargs is not None
assert logging_obj.pre_call_kwargs["input"] == "OCR document processing"
additional_args = logging_obj.pre_call_kwargs["additional_args"]
complete_input = additional_args["complete_input_dict"]
assert complete_input["document"] == DOCUMENT
assert complete_input["include_image_base64"] is True
assert additional_args["api_base"] == "https://api.mistral.ai/v1/ocr"
assert additional_args["headers"] == {
"Authorization": "Bearer sk-test",
"x-trace-id": "trace-1",
}
assert logging_obj.pre_call_kwargs is None
def test_ocr_routes_to_rust_when_enabled(fake_bridge):
def test_ocr_routes_to_rust_before_python_preparation(fake_bridge, monkeypatch):
monkeypatch.setattr(
ocr_main,
"_prepare_ocr_request",
lambda **kwargs: (_ for _ in ()).throw(AssertionError("Python OCR preparation must not run")),
)
response = litellm.ocr(
model=MODEL,
document=DOCUMENT,
@ -651,14 +654,11 @@ def test_ocr_routes_to_rust_when_enabled(fake_bridge):
assert response.pages[0].markdown == "hello world"
assert len(fake_bridge.calls) == 1
call = fake_bridge.calls[0]
assert call["model"] == "mistral-ocr-latest"
assert call["model"] == MODEL
assert call["document"] == DOCUMENT
assert call["api_key"] == "sk-test"
assert call["custom_llm_provider"] == "mistral"
assert call["extra_headers"] == {
"Authorization": "Bearer sk-test",
"x-trace-id": "trace-1",
}
assert call["custom_llm_provider"] is None
assert call["extra_headers"] == {"x-trace-id": "trace-1"}
assert call["optional_params"].get("include_image_base64") is True
@ -672,8 +672,8 @@ def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge):
assert isinstance(response, OCRResponse)
assert len(fake_bridge.calls) == 1
assert fake_bridge.calls[0]["model"] == "pixtral-12b-2409"
assert fake_bridge.calls[0]["custom_llm_provider"] == "azure_ai"
assert fake_bridge.calls[0]["model"] == "azure_ai/pixtral-12b-2409"
assert fake_bridge.calls[0]["custom_llm_provider"] is None
def test_ocr_rust_path_converts_file_document_before_bridge(fake_bridge):
@ -699,17 +699,22 @@ def test_ocr_exception_type_uses_resolved_provider_context(
return CapturedException("wrapped")
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
litellm.use_litellm_rust(True, ocr=RaisingBridge())
litellm.use_litellm_rust(True, ocr=RaisingBridge(), ocr_decline=RecordingDecline())
with pytest.raises(CapturedException):
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert captured["model"] == "mistral-ocr-latest"
assert captured["custom_llm_provider"] == "mistral"
assert captured["model"] == MODEL
assert captured["custom_llm_provider"] is None
@pytest.mark.asyncio
async def test_aocr_routes_to_async_rust_when_enabled(fake_async_bridge):
async def test_aocr_routes_to_async_rust_before_python_preparation(fake_async_bridge, monkeypatch):
monkeypatch.setattr(
ocr_main,
"_prepare_ocr_request",
lambda **kwargs: (_ for _ in ()).throw(AssertionError("Python OCR preparation must not run")),
)
response = await litellm.aocr(
model=MODEL,
document=DOCUMENT,
@ -722,14 +727,11 @@ async def test_aocr_routes_to_async_rust_when_enabled(fake_async_bridge):
assert response.pages[0].markdown == "hello world"
assert len(fake_async_bridge.calls) == 1
call = fake_async_bridge.calls[0]
assert call["model"] == "mistral-ocr-latest"
assert call["model"] == MODEL
assert call["document"] == DOCUMENT
assert call["api_key"] == "sk-test"
assert call["custom_llm_provider"] == "mistral"
assert call["extra_headers"] == {
"Authorization": "Bearer sk-test",
"x-trace-id": "trace-1",
}
assert call["custom_llm_provider"] is None
assert call["extra_headers"] == {"x-trace-id": "trace-1"}
assert call["optional_params"].get("include_image_base64") is True
@ -744,13 +746,13 @@ async def test_aocr_exception_type_uses_resolved_provider_context(
return CapturedException("wrapped")
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
litellm.use_litellm_rust(True, aocr=RaisingAsyncBridge())
litellm.use_litellm_rust(True, aocr=RaisingAsyncBridge(), ocr_decline=RecordingDecline())
with pytest.raises(CapturedException):
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert captured["model"] == "mistral-ocr-latest"
assert captured["custom_llm_provider"] == "mistral"
assert captured["model"] == MODEL
assert captured["custom_llm_provider"] is None
def test_ocr_forwards_timeout_to_rust(fake_bridge):
@ -764,9 +766,7 @@ def test_ocr_forwards_timeout_to_rust(fake_bridge):
def test_ocr_passes_default_request_timeout_to_rust(fake_bridge):
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
from litellm.constants import request_timeout
assert fake_bridge.calls[0]["timeout_seconds"] == float(request_timeout)
assert fake_bridge.calls[0]["timeout_seconds"] is None
def test_ocr_does_not_route_to_rust_when_disabled():
@ -800,6 +800,28 @@ def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch):
assert isinstance(response, OCRResponse)
def test_ocr_falls_back_to_python_when_runtime_declines(monkeypatch):
bridge = RecordingBridge()
decline = RecordingDecline("OCR provider is not supported by the Rust runtime")
litellm.use_litellm_rust(True, ocr=bridge, ocr_decline=decline)
prepared = build_prepared_request()
fake_prepare = MagicMock(return_value=prepared)
monkeypatch.setattr(ocr_main, "_prepare_ocr_request", fake_prepare)
monkeypatch.setattr(
ocr_main.base_llm_http_handler,
"ocr",
lambda **kwargs: OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr"),
)
response = litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert isinstance(response, OCRResponse)
assert len(decline.calls) == 1
assert bridge.calls == []
fake_prepare.assert_called_once()
def test_ocr_provider_configs_expose_api_key_env_vars():
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,