mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
rust(providers): end-to-end run_ocr orchestrator with shared client + timeout
This commit is contained in:
parent
ba8c541804
commit
055bef8482
1 changed files with 59 additions and 137 deletions
|
|
@ -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(¶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<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())
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue