mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
39d132f2a6
commit
b126fbb439
5 changed files with 231 additions and 35 deletions
|
|
@ -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()));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue