diff --git a/litellm-rust/crates/providers/src/ocr.rs b/litellm-rust/crates/providers/src/ocr.rs index bfcdfa16f47..dcd56a5f0b4 100644 --- a/litellm-rust/crates/providers/src/ocr.rs +++ b/litellm-rust/crates/providers/src/ocr.rs @@ -16,9 +16,15 @@ use crate::mistral::ocr::transformation as mistral; use crate::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; /// OCR over large documents can take a while; bound it generously rather than -/// hanging forever on an unresponsive upstream. +/// hanging forever on an unresponsive upstream. The client-level limit is the +/// outer ceiling; callers can tighten it per request via ``run_ocr``'s ``timeout``. const OCR_TIMEOUT_SECS: u64 = 600; +/// Maximum upstream body characters retained in error messages. OCR responses +/// can echo document contents and prompts; keep enough for debugging without +/// forwarding sensitive payloads across the host boundary. +const ERROR_BODY_MAX_CHARS: usize = 256; + /// Process-wide blocking HTTP client (connection pool + TLS reused across calls). fn http_client() -> &'static reqwest::blocking::Client { static CLIENT: OnceLock = OnceLock::new(); @@ -30,6 +36,14 @@ fn http_client() -> &'static reqwest::blocking::Client { }) } +fn truncate_error_body(body: &str) -> String { + if body.chars().count() <= ERROR_BODY_MAX_CHARS { + return body.to_string(); + } + let truncated: String = body.chars().take(ERROR_BODY_MAX_CHARS).collect(); + format!("{truncated}... (truncated)") +} + /// Perform a Mistral OCR call end to end and return the normalized response as /// JSON (the shape the Python `OCRResponse` model expects). /// @@ -40,6 +54,7 @@ pub fn run_ocr( api_key: Option<&str>, api_base: Option<&str>, optional_params: Map, + timeout: Option, ) -> CoreResult { let config = &MISTRAL_OCR_CONFIG; @@ -50,10 +65,12 @@ pub fn run_ocr( .transform_ocr_request(model, document, filtered_params)? .data; - let response = http_client() - .post(&url) - .bearer_auth(&api_key) - .json(&body) + let mut request = http_client().post(&url).bearer_auth(&api_key).json(&body); + if let Some(duration) = timeout { + request = request.timeout(duration); + } + + let response = request .send() .map_err(|err| CoreError::Network(err.to_string()))?; @@ -65,7 +82,7 @@ pub fn run_ocr( if !status.is_success() { return Err(CoreError::Http { status: status.as_u16(), - body: text, + body: truncate_error_body(&text), }); } @@ -76,3 +93,35 @@ pub fn run_ocr( .transform_ocr_response(model, response_json)? .into_json()) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn truncate_error_body_passes_short_strings_through() { + let body = "Unauthorized"; + assert_eq!(truncate_error_body(body), "Unauthorized"); + } + + #[test] + fn truncate_error_body_caps_long_payloads() { + let body = "x".repeat(ERROR_BODY_MAX_CHARS + 50); + let truncated = truncate_error_body(&body); + + assert!(truncated.ends_with("... (truncated)")); + let prefix_chars = truncated + .strip_suffix("... (truncated)") + .expect("truncated marker present") + .chars() + .count(); + assert_eq!(prefix_chars, ERROR_BODY_MAX_CHARS); + } + + #[test] + fn truncate_error_body_does_not_split_multibyte_chars() { + let body = "é".repeat(ERROR_BODY_MAX_CHARS + 10); + let truncated = truncate_error_body(&body); + assert!(truncated.is_char_boundary(truncated.len())); + } +} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index a84aefc0e7e..15e93f7b00c 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -1,3 +1,5 @@ +use std::time::Duration; + use litellm_core::error::CoreError; use litellm_providers::ocr::run_ocr; use pyo3::exceptions::{PyRuntimeError, PyValueError}; @@ -35,7 +37,7 @@ fn core_error_to_pyerr(err: CoreError) -> PyErr { /// Perform a Mistral OCR call end to end and return the response as a dict. #[pyfunction] -#[pyo3(signature = (model, document, api_key=None, api_base=None, optional_params=None))] +#[pyo3(signature = (model, document, api_key=None, api_base=None, optional_params=None, timeout_seconds=None))] fn ocr( py: Python<'_>, model: String, @@ -43,6 +45,7 @@ fn ocr( api_key: Option, api_base: Option, optional_params: Option>, + timeout_seconds: Option, ) -> PyResult> { let document = py_to_json(py, document.bind(py))?; @@ -54,6 +57,14 @@ fn ocr( None => Map::new(), }; + let timeout = timeout_seconds.and_then(|secs| { + if secs.is_finite() && secs > 0.0 { + Some(Duration::from_secs_f64(secs)) + } else { + None + } + }); + // Release the GIL during the blocking HTTP call (counted for observability). let result = gil::release_gil(py, || { run_ocr( @@ -62,6 +73,7 @@ fn ocr( api_key.as_deref(), api_base.as_deref(), optional_params, + timeout, ) }); diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 1b4058c8f05..5b0288d5a02 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -29,6 +29,22 @@ base_llm_http_handler = BaseLLMHTTPHandler() ################################################# +def _timeout_to_seconds( + timeout: Optional[Union[float, httpx.Timeout]], +) -> Optional[float]: + """Convert the Python OCR timeout to a single seconds value for the Rust bridge. + + The Rust HTTP client takes one duration; ``httpx.Timeout`` carries separate + connect/read/write/pool values, so pick the read deadline as the closest + analog to a total-request timeout. + """ + if timeout is None: + return None + if isinstance(timeout, httpx.Timeout): + return timeout.read + return float(timeout) + + @client async def aocr( model: str, @@ -262,25 +278,6 @@ def ocr( if dynamic_api_base: api_base = dynamic_api_base - # Optional Rust path: hand the whole Mistral OCR call to the Rust bridge. - # Load via importlib to avoid a static import edge that can be flagged as - # part of a cyclic import graph during package initialization. - rust_bridge = importlib.import_module("litellm.ocr.rust_bridge") - - if custom_llm_provider == "mistral" and rust_bridge.rust_ocr_enabled(): - # Resolve the key the same way the Python path does so secret-manager - # backends (AWS/Azure/GCP/Vault) work — the Rust bridge's own fallback - # only reads the process environment. - from litellm.secret_managers.main import get_secret_str - - resolved_api_key = api_key or get_secret_str("MISTRAL_API_KEY") - return OCRResponse( - **rust_bridge.rust_ocr( - model, document, resolved_api_key, api_base, kwargs - ) - ) - - # Get provider config ocr_provider_config: Optional[BaseOCRConfig] = ( ProviderConfigManager.get_provider_ocr_config( model=model, @@ -297,17 +294,14 @@ def ocr( f"OCR call - model: {model}, provider: {custom_llm_provider}" ) - # Get litellm params using GenericLiteLLMParams (same as responses API) litellm_params = GenericLiteLLMParams(**kwargs) - # Extract OCR-specific parameters from kwargs supported_params = ocr_provider_config.get_supported_ocr_params(model=model) non_default_params = {} for param in supported_params: if param in kwargs: non_default_params[param] = kwargs.pop(param) - # Map parameters to provider-specific format optional_params = ocr_provider_config.map_ocr_params( non_default_params=non_default_params, optional_params={}, @@ -316,7 +310,8 @@ def ocr( verbose_logger.debug(f"OCR optional_params after mapping: {optional_params}") - # Pre Call logging + effective_timeout = timeout or request_timeout + litellm_logging_obj.update_from_kwargs( kwargs=kwargs, model=model, @@ -328,12 +323,56 @@ def ocr( custom_llm_provider=custom_llm_provider, ) - # Call the handler - pass document dict directly + # Optional Rust path: hand the whole Mistral OCR call to the Rust bridge. + # Load via importlib to avoid a static import edge that can be flagged as + # part of a cyclic import graph during package initialization. + rust_bridge = importlib.import_module("litellm.ocr.rust_bridge") + + if custom_llm_provider == "mistral" and rust_bridge.rust_ocr_enabled(): + try: + importlib.import_module("litellm_python_bridge") + except ImportError: + # Rust extension wheel isn't installed: degrade to the Python + # path instead of hard-failing the request. + verbose_logger.debug( + "Rust OCR bridge unavailable; falling back to Python path" + ) + else: + # Resolve the key the same way the Python path does so + # secret-manager backends (AWS/Azure/GCP/Vault) work — the Rust + # bridge's own fallback only reads the process environment. + from litellm.secret_managers.main import get_secret_str + + resolved_api_key = api_key or get_secret_str("MISTRAL_API_KEY") + litellm_logging_obj.pre_call( + input="OCR document processing", + api_key=resolved_api_key, + additional_args={ + "complete_input_dict": { + "model": model, + "document": document, + **optional_params, + }, + "api_base": api_base, + "headers": {}, + }, + ) + return OCRResponse( + **rust_bridge.rust_ocr( + model=model, + document=document, + api_key=resolved_api_key, + api_base=api_base, + optional_params=optional_params, + timeout_seconds=_timeout_to_seconds(effective_timeout), + ) + ) + response = base_llm_http_handler.ocr( model=model, - document=document, # Pass the entire document dict + document=document, optional_params=optional_params, - timeout=timeout or request_timeout, + timeout=effective_timeout, logging_obj=litellm_logging_obj, api_key=api_key, api_base=api_base, diff --git a/litellm/ocr/rust_bridge.py b/litellm/ocr/rust_bridge.py index c4104cd2b84..cfcb0ac4e62 100644 --- a/litellm/ocr/rust_bridge.py +++ b/litellm/ocr/rust_bridge.py @@ -28,6 +28,7 @@ def rust_ocr( api_key: str | None, api_base: str | None, optional_params: dict, + timeout_seconds: float | None = None, ) -> dict: """Call the Rust bridge and return the raw OCR response dict. @@ -38,5 +39,5 @@ def rust_ocr( import litellm_python_bridge return litellm_python_bridge.ocr( - model, document, api_key, api_base, optional_params + model, document, api_key, api_base, optional_params, timeout_seconds ) diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index eae5d3358de..1a627108cc4 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -40,7 +40,7 @@ def fake_bridge(monkeypatch): """Install a fake compiled ``litellm_python_bridge`` module and record calls.""" calls = [] - def _ocr(model, document, api_key, api_base, optional_params): + def _ocr(model, document, api_key, api_base, optional_params, timeout_seconds=None): calls.append( { "model": model, @@ -48,6 +48,7 @@ def fake_bridge(monkeypatch): "api_key": api_key, "api_base": api_base, "optional_params": optional_params, + "timeout_seconds": timeout_seconds, } ) return dict(FAKE_OCR_RESPONSE) @@ -135,3 +136,97 @@ def test_ocr_resolves_key_via_secret_manager(monkeypatch, fake_bridge): litellm.ocr(model=MODEL, document=DOCUMENT) # no api_key passed assert fake_bridge[0]["api_key"] == "sk-from-vault" + + +def test_ocr_forwards_timeout_to_rust(fake_bridge): + """Caller-supplied timeout must flow into the Rust bridge so the fixed 600s + client ceiling doesn't silently override shorter deadlines.""" + litellm.use_litellm_rust() + + litellm.ocr( + model=MODEL, + document=DOCUMENT, + api_key="sk-test", + timeout=12.5, + ) + + assert fake_bridge[0]["timeout_seconds"] == 12.5 + + +def test_ocr_passes_default_request_timeout_to_rust(fake_bridge): + """When no explicit timeout is given, the library default (request_timeout) + must still be forwarded so the Rust path matches the Python path's deadline.""" + from litellm.constants import request_timeout + + litellm.use_litellm_rust() + + litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + + assert fake_bridge[0]["timeout_seconds"] == float(request_timeout) + + +def test_ocr_runs_logging_on_rust_path(monkeypatch, fake_bridge): + """The Rust shortcut must run the same logging setup (update_from_kwargs + + pre_call) the Python path runs, otherwise callbacks and spend tracking break.""" + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + update_calls = [] + pre_calls = [] + + real_update = LiteLLMLoggingObj.update_from_kwargs + real_pre = LiteLLMLoggingObj.pre_call + + def _record_update(self, *args, **kwargs): + update_calls.append(kwargs) + return real_update(self, *args, **kwargs) + + def _record_pre(self, *args, **kwargs): + pre_calls.append(kwargs) + return real_pre(self, *args, **kwargs) + + monkeypatch.setattr(LiteLLMLoggingObj, "update_from_kwargs", _record_update) + monkeypatch.setattr(LiteLLMLoggingObj, "pre_call", _record_pre) + litellm.use_litellm_rust() + + litellm.ocr( + model=MODEL, + document=DOCUMENT, + api_key="sk-test", + include_image_base64=True, + ) + + assert update_calls, "update_from_kwargs must be invoked on the Rust path" + assert update_calls[0].get("custom_llm_provider") == "mistral" + assert update_calls[0].get("model") == "mistral-ocr-latest" + assert pre_calls, "pre_call must be invoked on the Rust path" + assert pre_calls[0].get("input") == "OCR document processing" + assert pre_calls[0]["additional_args"]["complete_input_dict"]["document"] == DOCUMENT + + +def test_ocr_falls_back_to_python_when_bridge_missing(monkeypatch): + """A missing ``litellm_python_bridge`` extension must degrade gracefully to + the Python provider path instead of bubbling up ImportError.""" + monkeypatch.delitem(sys.modules, "litellm_python_bridge", raising=False) + + real_import_module = importlib.import_module + + def _blocked_import_module(name, package=None): + if name == "litellm_python_bridge": + raise ImportError("litellm_python_bridge not built") + return real_import_module(name, package) + + monkeypatch.setattr(importlib, "import_module", _blocked_import_module) + + handler_calls = [] + + def _fake_handler(*_args, **kwargs): + handler_calls.append(kwargs) + return OCRResponse(pages=[], model="mistral-ocr-latest") + + monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", _fake_handler) + litellm.use_litellm_rust() + + response = litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + + assert isinstance(response, OCRResponse) + assert handler_calls, "Python handler must run when the Rust bridge is missing"