diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index a1753df5aa7..32cb7a0dd21 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -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, + optional_params: Option>, +) -> PyResult> { + 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)?)?; diff --git a/litellm-rust/crates/runtime/src/ocr/mod.rs b/litellm-rust/crates/runtime/src/ocr/mod.rs index b3b11d824e6..89d4906c84e 100644 --- a/litellm-rust/crates/runtime/src/ocr/mod.rs +++ b/litellm-rust/crates/runtime/src/ocr/mod.rs @@ -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; diff --git a/litellm-rust/crates/runtime/src/ocr/prepare.rs b/litellm-rust/crates/runtime/src/ocr/prepare.rs index c7ef3246538..163bb1ebaed 100644 --- a/litellm-rust/crates/runtime/src/ocr/prepare.rs +++ b/litellm-rust/crates/runtime/src/ocr/prepare.rs @@ -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 { + 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 { @@ -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, + ) -> 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) + ); + } +} diff --git a/litellm-rust/crates/runtime/src/ocr/types.rs b/litellm-rust/crates/runtime/src/ocr/types.rs index 46821c66fad..28ca995c8bc 100644 --- a/litellm-rust/crates/runtime/src/ocr/types.rs +++ b/litellm-rust/crates/runtime/src/ocr/types.rs @@ -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, diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index b918f013700..aa517b93e90 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -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, diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 82297d35170..a9fea41e003 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -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, diff --git a/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py b/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py index 0c8b1cc2836..b73718ba175 100644 --- a/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py +++ b/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py @@ -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.""" diff --git a/tests/test_litellm/ocr/test_ocr_native_format.py b/tests/test_litellm/ocr/test_ocr_native_format.py index 463213a2071..4d0efd8934b 100644 --- a/tests/test_litellm/ocr/test_ocr_native_format.py +++ b/tests/test_litellm/ocr/test_ocr_native_format.py @@ -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 diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index acad249a2bb..60cf00371f5 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -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,