fix(ocr): honor timeout, logging, and missing-bridge fallback on Rust OCR path

- Forward the caller's timeout into the Rust bridge so the fixed 600s client
  ceiling no longer overrides shorter deadlines or the library default.
- Run update_from_kwargs and pre_call before invoking the Rust shortcut so
  observability, callbacks, and spend tracking match the Python path.
- Fall back to the Python OCR path when litellm_python_bridge isn't importable
  instead of raising ImportError to callers.
- Truncate upstream Mistral OCR error bodies before they cross the host
  boundary to avoid leaking document or prompt contents in CoreError::Http.
This commit is contained in:
Cursor Agent 2026-06-23 05:05:45 +00:00
parent 39d132f2a6
commit b126fbb439
No known key found for this signature in database
5 changed files with 231 additions and 35 deletions

View file

@ -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<reqwest::blocking::Client> = 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<String, Value>,
timeout: Option<Duration>,
) -> CoreResult<Value> {
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()));
}
}

View file

@ -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<String>,
api_base: Option<String>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
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,
)
});

View file

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

View file

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

View file

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