diff --git a/litellm-rust/crates/providers/src/ocr.rs b/litellm-rust/crates/providers/src/ocr.rs index 48e3bbe8894..bfcdfa16f47 100644 --- a/litellm-rust/crates/providers/src/ocr.rs +++ b/litellm-rust/crates/providers/src/ocr.rs @@ -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 { - 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 = 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, +/// 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, ) -> CoreResult { - match operation { - "map_params" => { - let params = required_object(payload, "non_default_params")?; - Ok(Value::Object(provider_config.map_ocr_params(¶ms))) - } - "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> { - 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, 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, - key: &'static str, -) -> CoreResult> { - 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, key: &'static str) -> CoreResult { - 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()) }