rust(providers): end-to-end run_ocr orchestrator with shared client + timeout

This commit is contained in:
Ishaan Jaffer 2026-06-22 18:25:05 -07:00
parent ba8c541804
commit 055bef8482
No known key found for this signature in database

View file

@ -1,156 +1,78 @@
use litellm_core::error::{json_type_name, CoreError};
//! End-to-end OCR orchestration.
//!
//! Owns the whole Mistral OCR call so the Python side stays a thin bridge:
//! resolve the API key, build the URL + body via the pure transforms, POST it,
//! and normalize the response. The HTTP client is built once and reused.
use std::sync::OnceLock;
use std::time::Duration;
use litellm_core::error::CoreError;
use litellm_core::ocr::transformation::OcrProviderConfig;
use litellm_core::CoreResult;
use serde_json::{Map, Value};
use crate::mistral::ocr::transformation as mistral;
use crate::mistral::ocr::transformation::MISTRAL_OCR_CONFIG;
pub fn transform(payload: Value) -> CoreResult<Value> {
let payload = payload_object(&payload)?;
let provider = required_string(payload, "provider")?;
let operation = required_string(payload, "operation")?;
/// OCR over large documents can take a while; bound it generously rather than
/// hanging forever on an unresponsive upstream.
const OCR_TIMEOUT_SECS: u64 = 600;
match provider {
"mistral" => transform_with_provider(&MISTRAL_OCR_CONFIG, operation, payload),
_ => Err(CoreError::InvalidResponse(format!(
"unsupported OCR provider: {provider}"
))),
}
/// 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();
CLIENT.get_or_init(|| {
reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(OCR_TIMEOUT_SECS))
.build()
.expect("failed to build reqwest client")
})
}
fn transform_with_provider(
provider_config: &impl OcrProviderConfig,
operation: &str,
payload: &Map<String, Value>,
/// Perform a Mistral OCR call end to end and return the normalized response as
/// JSON (the shape the Python `OCRResponse` model expects).
///
/// Blocking: intended to be called with the GIL released from the Python bridge.
pub fn run_ocr(
model: &str,
document: Value,
api_key: Option<&str>,
api_base: Option<&str>,
optional_params: Map<String, Value>,
) -> CoreResult<Value> {
match operation {
"map_params" => {
let params = required_object(payload, "non_default_params")?;
Ok(Value::Object(provider_config.map_ocr_params(&params)))
}
"transform_request" => {
let model = required_string(payload, "model")?;
let document = required_value(payload, "document")?;
let optional_params = required_object(payload, "optional_params")?;
let transformed =
provider_config.transform_ocr_request(model, document, optional_params)?;
Ok(serde_json::json!({
"data": transformed.data,
"files": transformed.files,
}))
}
"transform_response" => {
let model = required_string(payload, "model")?;
let response_json = required_value(payload, "response_json")?;
let transformed = provider_config.transform_ocr_response(model, response_json)?;
Ok(transformed.into_json())
}
_ => Err(CoreError::InvalidResponse(format!(
"unsupported OCR operation: {operation}"
))),
}
}
let config = &MISTRAL_OCR_CONFIG;
fn payload_object(payload: &Value) -> CoreResult<&Map<String, Value>> {
payload.as_object().ok_or_else(|| CoreError::InvalidType {
expected: "object",
actual: json_type_name(payload),
})
}
let api_key = mistral::resolve_api_key(api_key, &|key| std::env::var(key).ok())?;
let url = mistral::complete_url(api_base);
let filtered_params = config.map_ocr_params(&optional_params);
let body = config
.transform_ocr_request(model, document, filtered_params)?
.data;
fn required_string<'a>(payload: &'a Map<String, Value>, key: &'static str) -> CoreResult<&'a str> {
let value = payload.get(key).ok_or(CoreError::MissingField(key))?;
value.as_str().ok_or_else(|| CoreError::InvalidType {
expected: "string",
actual: json_type_name(value),
})
}
let response = http_client()
.post(&url)
.bearer_auth(&api_key)
.json(&body)
.send()
.map_err(|err| CoreError::Network(err.to_string()))?;
fn required_object(
payload: &Map<String, Value>,
key: &'static str,
) -> CoreResult<Map<String, Value>> {
let value = payload.get(key).ok_or(CoreError::MissingField(key))?;
value
.as_object()
.cloned()
.ok_or_else(|| CoreError::InvalidType {
expected: "object",
actual: json_type_name(value),
})
}
let status = response.status();
let text = response
.text()
.map_err(|err| CoreError::Network(err.to_string()))?;
fn required_value(payload: &Map<String, Value>, key: &'static str) -> CoreResult<Value> {
payload
.get(key)
.cloned()
.ok_or(CoreError::MissingField(key))
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn transform_dispatches_mistral_map_params() {
let result = transform(json!({
"provider": "mistral",
"operation": "map_params",
"non_default_params": {
"extract_header": true,
"unsupported_param": "value"
}
}))
.expect("payload should transform");
assert_eq!(result, json!({"extract_header": true}));
}
#[test]
fn transform_dispatches_mistral_request() {
let document = json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
if !status.is_success() {
return Err(CoreError::Http {
status: status.as_u16(),
body: text,
});
let result = transform(json!({
"provider": "mistral",
"operation": "transform_request",
"model": "mistral-ocr-latest",
"document": document,
"optional_params": {"include_image_base64": true}
}))
.expect("payload should transform");
assert_eq!(
result,
json!({
"data": {
"model": "mistral-ocr-latest",
"document": {
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
},
"include_image_base64": true
},
"files": null
})
);
}
#[test]
fn transform_rejects_unknown_provider() {
let err = transform(json!({
"provider": "azure_ai",
"operation": "map_params",
"non_default_params": {}
}))
.expect_err("unsupported provider should fail");
let response_json: Value = serde_json::from_str(&text)
.map_err(|err| CoreError::InvalidResponse(format!("invalid OCR response JSON: {err}")))?;
assert_eq!(
err,
CoreError::InvalidResponse("unsupported OCR provider: azure_ai".to_string())
);
}
Ok(config
.transform_ocr_response(model, response_json)?
.into_json())
}