mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(ocr): route all OCR through Rust
This commit is contained in:
parent
d860cd9599
commit
1528c9b9d1
14 changed files with 618 additions and 330 deletions
17
litellm-rust/Cargo.lock
generated
17
litellm-rust/Cargo.lock
generated
|
|
@ -732,6 +732,16 @@ version = "0.3.17"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a"
|
||||
|
||||
[[package]]
|
||||
name = "mime_guess"
|
||||
version = "2.0.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f7c44f8e672c00fe5308fa235f821cb4198414e1c77935c1ab6948d3fd78550e"
|
||||
dependencies = [
|
||||
"mime",
|
||||
"unicase",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "mio"
|
||||
version = "1.2.1"
|
||||
|
|
@ -1025,6 +1035,7 @@ dependencies = [
|
|||
"hyper-util",
|
||||
"js-sys",
|
||||
"log",
|
||||
"mime_guess",
|
||||
"percent-encoding",
|
||||
"pin-project-lite",
|
||||
"quinn",
|
||||
|
|
@ -1552,6 +1563,12 @@ version = "1.20.1"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20"
|
||||
|
||||
[[package]]
|
||||
name = "unicase"
|
||||
version = "2.9.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "dbc4bc3a9f746d862c45cb89d705aa10f187bb96c76001afab07a0d35ce60142"
|
||||
|
||||
[[package]]
|
||||
name = "unicode-ident"
|
||||
version = "1.0.24"
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ axum = "0.7"
|
|||
pyo3 = "0.23.5"
|
||||
pyo3-async-runtimes = { version = "0.23.0", features = ["tokio-runtime"] }
|
||||
rand = "0.8"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "rustls-tls", "http2", "stream"] }
|
||||
reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "multipart", "rustls-tls", "http2", "stream"] }
|
||||
serde = { version = "1.0", features = ["derive"] }
|
||||
serde_json = "1.0"
|
||||
sha2 = "0.10"
|
||||
|
|
|
|||
|
|
@ -8,7 +8,9 @@ use std::sync::OnceLock;
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::ocr::transformation::{OcrAuthStrategy, OcrResponseHandling};
|
||||
use litellm_core::ocr::transformation::{
|
||||
OcrAuthStrategy, OcrDocumentPreparation, OcrResponseHandling,
|
||||
};
|
||||
use litellm_core::CoreResult;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
|
|
@ -16,7 +18,7 @@ mod common_utils;
|
|||
|
||||
use common_utils::{
|
||||
convert_document_url_to_data_uri, has_header, ocr_provider_config, poll_document_intelligence,
|
||||
string_headers, truncate_error_body,
|
||||
string_headers, truncate_error_body, upload_reducto_document,
|
||||
};
|
||||
|
||||
/// OCR over large documents can take a while; bound it generously rather than
|
||||
|
|
@ -88,15 +90,20 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
|
|||
&env_lookup,
|
||||
)?;
|
||||
let filtered_params = config.map_ocr_params(&request.optional_params);
|
||||
let document = if config.requires_data_uri_document() {
|
||||
convert_document_url_to_data_uri(request.document).await?
|
||||
} else {
|
||||
request.document
|
||||
let upstream_headers = upstream_headers(&headers, auth_strategy, api_key.as_deref());
|
||||
let document = match config.document_preparation() {
|
||||
OcrDocumentPreparation::None => request.document,
|
||||
OcrDocumentPreparation::DataUri => {
|
||||
convert_document_url_to_data_uri(request.document).await?
|
||||
}
|
||||
OcrDocumentPreparation::ReductoUpload => {
|
||||
upload_reducto_document(request.document, &url, &upstream_headers, request.timeout)
|
||||
.await?
|
||||
}
|
||||
};
|
||||
let body = config
|
||||
.transform_ocr_request(model, document, filtered_params)?
|
||||
.data;
|
||||
let upstream_headers = upstream_headers(&headers, auth_strategy, api_key.as_deref());
|
||||
|
||||
let mut request_builder = http_client().post(&url).json(&body);
|
||||
for (key, value) in &upstream_headers {
|
||||
|
|
@ -220,6 +227,8 @@ mod tests {
|
|||
.expect("vertex deepseek config resolves")
|
||||
.supported_ocr_params()
|
||||
.contains(&"temperature"));
|
||||
assert!(ocr_provider_config("reducto", "parse-v3").is_some());
|
||||
assert!(ocr_provider_config("reducto", "parse-legacy").is_some());
|
||||
assert!(ocr_provider_config("openai", "gpt-4o").is_none());
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,9 @@ use litellm_core::providers::azure_ai::ocr::transformation::{
|
|||
AZURE_AI_OCR_CONFIG, AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG,
|
||||
};
|
||||
use litellm_core::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG;
|
||||
use litellm_core::providers::reducto::ocr::transformation::{
|
||||
REDUCTO_PARSE_LEGACY_CONFIG, REDUCTO_PARSE_V3_CONFIG,
|
||||
};
|
||||
use litellm_core::providers::vertex_ai::ocr::transformation as vertex_ai;
|
||||
use litellm_core::providers::vertex_ai::ocr::transformation::{
|
||||
VERTEX_AI_DEEPSEEK_OCR_CONFIG, VERTEX_AI_OCR_CONFIG,
|
||||
|
|
@ -46,6 +49,8 @@ pub(super) fn ocr_provider_config(
|
|||
"azure_ai" => Some(&AZURE_AI_OCR_CONFIG),
|
||||
"vertex_ai" if vertex_ai::is_deepseek_model(model) => Some(&VERTEX_AI_DEEPSEEK_OCR_CONFIG),
|
||||
"vertex_ai" => Some(&VERTEX_AI_OCR_CONFIG),
|
||||
"reducto" if model == "parse-v3" => Some(&REDUCTO_PARSE_V3_CONFIG),
|
||||
"reducto" if model == "parse-legacy" => Some(&REDUCTO_PARSE_LEGACY_CONFIG),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
|
@ -296,6 +301,138 @@ pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreRes
|
|||
Ok(Value::Object(transformed))
|
||||
}
|
||||
|
||||
fn decode_reducto_data_uri(source_url: &str) -> CoreResult<(Vec<u8>, String)> {
|
||||
let (header, encoded) = source_url.split_once(',').ok_or_else(|| {
|
||||
CoreError::InvalidRequest("Invalid Reducto data URI provided.".to_string())
|
||||
})?;
|
||||
if !header.contains(";base64") {
|
||||
return Err(CoreError::InvalidRequest(
|
||||
"Reducto only supports base64-encoded data URIs.".to_string(),
|
||||
));
|
||||
}
|
||||
let mime = header
|
||||
.strip_prefix("data:")
|
||||
.and_then(|value| value.split(';').next())
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or("application/octet-stream")
|
||||
.to_string();
|
||||
let bytes = BASE64_STANDARD.decode(encoded).map_err(|_| {
|
||||
CoreError::InvalidRequest("Invalid Reducto base64 payload provided.".to_string())
|
||||
})?;
|
||||
Ok((bytes, mime))
|
||||
}
|
||||
|
||||
fn reducto_upload_url(parse_url: &str) -> CoreResult<Url> {
|
||||
let mut url = Url::parse(parse_url)
|
||||
.map_err(|err| CoreError::InvalidRequest(format!("invalid Reducto parse URL: {err}")))?;
|
||||
let path = url.path().trim_end_matches('/');
|
||||
let base_path = path
|
||||
.strip_suffix("/parse")
|
||||
.unwrap_or(path)
|
||||
.trim_end_matches('/');
|
||||
url.set_path(&format!("{base_path}/upload"));
|
||||
url.set_query(None);
|
||||
Ok(url)
|
||||
}
|
||||
|
||||
fn reducto_auth_headers(headers: &[(String, String)]) -> Vec<(String, String)> {
|
||||
headers
|
||||
.iter()
|
||||
.filter(|(key, _)| key.eq_ignore_ascii_case("authorization"))
|
||||
.cloned()
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn upload_reducto_bytes(
|
||||
bytes: Vec<u8>,
|
||||
mime: String,
|
||||
parse_url: &str,
|
||||
headers: &[(String, String)],
|
||||
timeout: Option<Duration>,
|
||||
) -> CoreResult<String> {
|
||||
let upload_url = reducto_upload_url(parse_url)?;
|
||||
let part = reqwest::multipart::Part::bytes(bytes)
|
||||
.file_name("document")
|
||||
.mime_str(&mime)
|
||||
.map_err(|err| {
|
||||
CoreError::InvalidRequest(format!("invalid Reducto upload MIME type: {err}"))
|
||||
})?;
|
||||
let form = reqwest::multipart::Form::new().part("file", part);
|
||||
let mut request_builder = http_client().post(upload_url).multipart(form);
|
||||
for (key, value) in reducto_auth_headers(headers) {
|
||||
request_builder = request_builder.header(key, value);
|
||||
}
|
||||
if let Some(duration) = timeout {
|
||||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
let status = response.status();
|
||||
let text = response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
if !status.is_success() {
|
||||
return Err(CoreError::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
});
|
||||
}
|
||||
let response_json: Value = serde_json::from_str(&text).map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("invalid Reducto upload response JSON: {err}"))
|
||||
})?;
|
||||
response_json
|
||||
.get("file_id")
|
||||
.and_then(Value::as_str)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_string)
|
||||
.ok_or_else(|| {
|
||||
CoreError::InvalidResponse(format!(
|
||||
"Reducto /upload returned 200 without a file_id; got payload={response_json}"
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn upload_reducto_document(
|
||||
document: Value,
|
||||
parse_url: &str,
|
||||
headers: &[(String, String)],
|
||||
timeout: Option<Duration>,
|
||||
) -> CoreResult<Value> {
|
||||
let Some((field, source_url)) = document_url_field(&document)? else {
|
||||
return Err(CoreError::InvalidRequest(
|
||||
"Reducto expected OCR preprocessing to produce document_url or image_url".to_string(),
|
||||
));
|
||||
};
|
||||
if source_url.starts_with("reducto://") {
|
||||
return Ok(document);
|
||||
}
|
||||
if source_url.starts_with("http://") || source_url.starts_with("https://") {
|
||||
return Err(CoreError::InvalidRequest(
|
||||
"Reducto requires type='file' (auto-uploaded) or a reducto:// id. Plain http(s) URLs are not supported; upload the file first."
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
if !source_url.starts_with("data:") {
|
||||
return Err(CoreError::InvalidRequest(
|
||||
"Reducto requires a reducto:// id or a base64 data URI after OCR preprocessing."
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
let (bytes, mime) = decode_reducto_data_uri(source_url)?;
|
||||
let file_id = upload_reducto_bytes(bytes, mime, parse_url, headers, timeout).await?;
|
||||
let mut transformed = document
|
||||
.as_object()
|
||||
.cloned()
|
||||
.ok_or_else(|| CoreError::InvalidRequest("OCR document must be an object".to_string()))?;
|
||||
transformed.insert(field.to_string(), Value::String(file_id));
|
||||
Ok(Value::Object(transformed))
|
||||
}
|
||||
|
||||
fn same_origin(left: &str, right: &str) -> bool {
|
||||
let Ok(left) = reqwest::Url::parse(left) else {
|
||||
return false;
|
||||
|
|
|
|||
|
|
@ -25,6 +25,13 @@ pub enum OcrResponseHandling {
|
|||
AzureDocumentIntelligencePoll,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum OcrDocumentPreparation {
|
||||
None,
|
||||
DataUri,
|
||||
ReductoUpload,
|
||||
}
|
||||
|
||||
pub trait OcrProviderConfig: Sync {
|
||||
fn supported_ocr_params(&self) -> &'static [&'static str];
|
||||
|
||||
|
|
@ -73,6 +80,14 @@ pub trait OcrProviderConfig: Sync {
|
|||
false
|
||||
}
|
||||
|
||||
fn document_preparation(&self) -> OcrDocumentPreparation {
|
||||
if self.requires_data_uri_document() {
|
||||
OcrDocumentPreparation::DataUri
|
||||
} else {
|
||||
OcrDocumentPreparation::None
|
||||
}
|
||||
}
|
||||
|
||||
fn response_handling(&self) -> OcrResponseHandling {
|
||||
OcrResponseHandling::Json
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
pub mod azure_ai;
|
||||
pub mod mistral;
|
||||
pub mod openai;
|
||||
pub mod reducto;
|
||||
pub mod vertex_ai;
|
||||
|
|
|
|||
1
litellm-rust/crates/core/src/providers/reducto/mod.rs
Normal file
1
litellm-rust/crates/core/src/providers/reducto/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub mod ocr;
|
||||
|
|
@ -0,0 +1 @@
|
|||
pub mod transformation;
|
||||
|
|
@ -0,0 +1,323 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use crate::error::{json_type_name, CoreError, CoreResult};
|
||||
use crate::ocr::transformation::{OcrDocumentPreparation, OcrProviderConfig};
|
||||
use crate::ocr::types::{OcrRequestData, OcrResponseData};
|
||||
use serde_json::{json, Map, Value};
|
||||
|
||||
const REDUCTO_API_BASE: &str = "https://platform.reducto.ai";
|
||||
const REDUCTO_API_KEY_ENV: &str = "REDUCTO_API_KEY";
|
||||
const REDUCTO_PARSE_V3_PARAMS: &[&str] = &["formatting", "retrieval", "settings"];
|
||||
const REDUCTO_PARSE_LEGACY_PARAMS: &[&str] = &["enhance"];
|
||||
|
||||
pub struct ReductoParseV3Config;
|
||||
pub struct ReductoParseLegacyConfig;
|
||||
|
||||
pub const REDUCTO_PARSE_V3_CONFIG: ReductoParseV3Config = ReductoParseV3Config;
|
||||
pub const REDUCTO_PARSE_LEGACY_CONFIG: ReductoParseLegacyConfig = ReductoParseLegacyConfig;
|
||||
|
||||
fn non_empty(value: Option<&str>) -> Option<&str> {
|
||||
value.map(str::trim).filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
pub fn resolve_api_key(
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
non_empty(api_key)
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(REDUCTO_API_KEY_ENV).filter(|value| !value.trim().is_empty()))
|
||||
.ok_or_else(|| {
|
||||
CoreError::Auth(
|
||||
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
|
||||
.to_string(),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
pub fn complete_url(api_base: Option<&str>) -> String {
|
||||
let base = non_empty(api_base).unwrap_or(REDUCTO_API_BASE);
|
||||
format!("{}/parse", base.trim_end_matches('/'))
|
||||
}
|
||||
|
||||
fn source_url<'a>(document: &'a Value, model: &str) -> CoreResult<&'a str> {
|
||||
let object = document.as_object().ok_or_else(|| CoreError::InvalidType {
|
||||
expected: "object",
|
||||
actual: json_type_name(document),
|
||||
})?;
|
||||
object
|
||||
.get("document_url")
|
||||
.or_else(|| object.get("image_url"))
|
||||
.and_then(Value::as_str)
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or_else(|| {
|
||||
CoreError::InvalidRequest(format!(
|
||||
"Reducto expected OCR preprocessing to produce document_url or image_url for model={model}"
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn page_no(block: &Value) -> Option<i64> {
|
||||
block
|
||||
.get("bbox")
|
||||
.and_then(|bbox| bbox.get("page"))
|
||||
.and_then(|page| match page {
|
||||
Value::Number(number) => number.as_i64(),
|
||||
Value::String(value) => value.parse::<i64>().ok(),
|
||||
_ => None,
|
||||
})
|
||||
}
|
||||
|
||||
fn block_content(block: &Value) -> Option<&str> {
|
||||
block.get("content").and_then(Value::as_str)
|
||||
}
|
||||
|
||||
fn build_pages_from_reducto(result: &Value) -> Vec<Value> {
|
||||
let chunks = result
|
||||
.get("chunks")
|
||||
.and_then(Value::as_array)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
let mut blocks_by_page: BTreeMap<i64, Vec<Value>> = BTreeMap::new();
|
||||
|
||||
for chunk in &chunks {
|
||||
let Some(blocks) = chunk.get("blocks").and_then(Value::as_array) else {
|
||||
continue;
|
||||
};
|
||||
for block in blocks {
|
||||
if let Some(page) = page_no(block) {
|
||||
blocks_by_page.entry(page).or_default().push(block.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if blocks_by_page.is_empty() {
|
||||
let markdown = chunks
|
||||
.iter()
|
||||
.filter_map(|chunk| chunk.get("content").and_then(Value::as_str))
|
||||
.filter(|content| !content.is_empty())
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n\n");
|
||||
if markdown.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
return vec![json!({"index": 0, "markdown": markdown})];
|
||||
}
|
||||
|
||||
blocks_by_page
|
||||
.into_iter()
|
||||
.map(|(page, blocks)| {
|
||||
let markdown = blocks
|
||||
.iter()
|
||||
.filter_map(block_content)
|
||||
.filter(|content| !content.is_empty())
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n\n");
|
||||
json!({
|
||||
"index": (page - 1).max(0),
|
||||
"markdown": markdown,
|
||||
"blocks": blocks,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
impl OcrProviderConfig for ReductoParseV3Config {
|
||||
fn supported_ocr_params(&self) -> &'static [&'static str] {
|
||||
REDUCTO_PARSE_V3_PARAMS
|
||||
}
|
||||
|
||||
fn transform_ocr_request(
|
||||
&self,
|
||||
model: &str,
|
||||
document: Value,
|
||||
optional_params: Map<String, Value>,
|
||||
) -> CoreResult<OcrRequestData> {
|
||||
let mut data = Map::new();
|
||||
data.insert(
|
||||
"input".to_string(),
|
||||
Value::String(source_url(&document, model)?.to_string()),
|
||||
);
|
||||
for (key, value) in optional_params {
|
||||
data.insert(key, value);
|
||||
}
|
||||
Ok(OcrRequestData {
|
||||
data: Value::Object(data),
|
||||
files: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> CoreResult<OcrResponseData> {
|
||||
transform_reducto_response(model, response_json)
|
||||
}
|
||||
|
||||
fn complete_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_optional_params: &Map<String, Value>,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
Ok(complete_url(api_base))
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
resolve_api_key(api_key, env_lookup)
|
||||
}
|
||||
|
||||
fn document_preparation(&self) -> OcrDocumentPreparation {
|
||||
OcrDocumentPreparation::ReductoUpload
|
||||
}
|
||||
}
|
||||
|
||||
impl OcrProviderConfig for ReductoParseLegacyConfig {
|
||||
fn supported_ocr_params(&self) -> &'static [&'static str] {
|
||||
REDUCTO_PARSE_LEGACY_PARAMS
|
||||
}
|
||||
|
||||
fn transform_ocr_request(
|
||||
&self,
|
||||
model: &str,
|
||||
document: Value,
|
||||
optional_params: Map<String, Value>,
|
||||
) -> CoreResult<OcrRequestData> {
|
||||
let mut data = Map::new();
|
||||
data.insert(
|
||||
"document_url".to_string(),
|
||||
Value::String(source_url(&document, model)?.to_string()),
|
||||
);
|
||||
if let Some(enhance) = optional_params.get("enhance") {
|
||||
data.insert("options".to_string(), json!({"enhance": enhance}));
|
||||
}
|
||||
Ok(OcrRequestData {
|
||||
data: Value::Object(data),
|
||||
files: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> CoreResult<OcrResponseData> {
|
||||
transform_reducto_response(model, response_json)
|
||||
}
|
||||
|
||||
fn complete_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_optional_params: &Map<String, Value>,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
Ok(complete_url(api_base))
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
resolve_api_key(api_key, env_lookup)
|
||||
}
|
||||
|
||||
fn document_preparation(&self) -> OcrDocumentPreparation {
|
||||
OcrDocumentPreparation::ReductoUpload
|
||||
}
|
||||
}
|
||||
|
||||
fn transform_reducto_response(model: &str, response_json: Value) -> CoreResult<OcrResponseData> {
|
||||
let response = response_json
|
||||
.as_object()
|
||||
.ok_or_else(|| CoreError::InvalidType {
|
||||
expected: "object",
|
||||
actual: json_type_name(&response_json),
|
||||
})?;
|
||||
let result = response.get("result").unwrap_or(&response_json);
|
||||
let usage = response.get("usage").cloned().unwrap_or_else(|| json!({}));
|
||||
Ok(OcrResponseData {
|
||||
pages: build_pages_from_reducto(result),
|
||||
model: model.to_string(),
|
||||
document_annotation: None,
|
||||
usage_info: Some(json!({
|
||||
"pages_processed": usage.get("num_pages").cloned().unwrap_or(Value::Null),
|
||||
"credits": usage.get("credits").cloned().unwrap_or(Value::Null),
|
||||
})),
|
||||
object: "ocr".to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn parse_v3_request_uses_uploaded_file_id() {
|
||||
let body = REDUCTO_PARSE_V3_CONFIG
|
||||
.transform_ocr_request(
|
||||
"parse-v3",
|
||||
json!({"type": "document_url", "document_url": "reducto://file-1"}),
|
||||
Map::from_iter([("settings".to_string(), json!({"ocr": true}))]),
|
||||
)
|
||||
.expect("request transforms")
|
||||
.data;
|
||||
|
||||
assert_eq!(
|
||||
body,
|
||||
json!({"input": "reducto://file-1", "settings": {"ocr": true}})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_legacy_request_maps_enhance_into_options() {
|
||||
let body = REDUCTO_PARSE_LEGACY_CONFIG
|
||||
.transform_ocr_request(
|
||||
"parse-legacy",
|
||||
json!({"type": "document_url", "document_url": "reducto://file-1"}),
|
||||
Map::from_iter([("enhance".to_string(), json!(true))]),
|
||||
)
|
||||
.expect("request transforms")
|
||||
.data;
|
||||
|
||||
assert_eq!(
|
||||
body,
|
||||
json!({"document_url": "reducto://file-1", "options": {"enhance": true}})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reducto_response_groups_blocks_by_page() {
|
||||
let response = REDUCTO_PARSE_V3_CONFIG
|
||||
.transform_ocr_response(
|
||||
"parse-v3",
|
||||
json!({
|
||||
"result": {
|
||||
"chunks": [{
|
||||
"blocks": [
|
||||
{"content": "p1", "bbox": {"page": 1}},
|
||||
{"content": "p2", "bbox": {"page": 2}}
|
||||
]
|
||||
}]
|
||||
},
|
||||
"usage": {"num_pages": 2, "credits": 1}
|
||||
}),
|
||||
)
|
||||
.expect("response transforms");
|
||||
|
||||
assert_eq!(response.pages[0]["index"], 0);
|
||||
assert_eq!(response.pages[0]["markdown"], "p1");
|
||||
assert_eq!(response.pages[1]["index"], 1);
|
||||
assert_eq!(
|
||||
response.usage_info,
|
||||
Some(json!({"pages_processed": 2, "credits": 1}))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1394,7 +1394,6 @@ from .skills.main import (
|
|||
)
|
||||
from .containers.main import *
|
||||
from .ocr.main import *
|
||||
from .ocr.rust_bridge import use_litellm_rust
|
||||
from .rag.main import *
|
||||
from .sandbox.main import *
|
||||
from .search.main import *
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
Main OCR function for LiteLLM.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import mimetypes
|
||||
import os
|
||||
|
|
@ -17,21 +16,15 @@ import litellm
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import request_timeout
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.ocr.rust_bridge import (
|
||||
RustAocr,
|
||||
RustOcr,
|
||||
load_rust_aocr,
|
||||
load_rust_ocr,
|
||||
rust_ocr_enabled,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager, client
|
||||
|
||||
####### ENVIRONMENT VARIABLES ###################
|
||||
base_llm_http_handler = BaseLLMHTTPHandler()
|
||||
#################################################
|
||||
from litellm.utils import client
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
@ -42,7 +35,6 @@ class _PreparedOCRRequest:
|
|||
api_base: str | None
|
||||
custom_llm_provider: str
|
||||
extra_headers: dict[str, object] | None
|
||||
provider_config: BaseOCRConfig
|
||||
optional_params: dict[str, object]
|
||||
litellm_params: dict[str, object]
|
||||
effective_timeout: Union[float, httpx.Timeout]
|
||||
|
|
@ -57,14 +49,6 @@ class _PreparedRustOCRCall:
|
|||
optional_params: dict[str, object]
|
||||
|
||||
|
||||
_RUST_OCR_PROVIDERS = {
|
||||
"mistral",
|
||||
"azure_ai",
|
||||
"azure_ai/doc-intelligence",
|
||||
"vertex_ai",
|
||||
}
|
||||
|
||||
|
||||
def _timeout_to_seconds(
|
||||
timeout: Union[float, httpx.Timeout] | None,
|
||||
) -> float | None:
|
||||
|
|
@ -128,29 +112,10 @@ def _prepare_ocr_request(
|
|||
if dynamic_api_base:
|
||||
api_base = dynamic_api_base
|
||||
|
||||
ocr_provider_config = ProviderConfigManager.get_provider_ocr_config(
|
||||
model=model,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if ocr_provider_config is None:
|
||||
raise ValueError(f"OCR is not supported for provider: {custom_llm_provider}")
|
||||
|
||||
verbose_logger.debug(f"OCR call - model: {model}, provider: {custom_llm_provider}")
|
||||
|
||||
litellm_params = GenericLiteLLMParams(**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)
|
||||
|
||||
optional_params = ocr_provider_config.map_ocr_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params={},
|
||||
model=model,
|
||||
)
|
||||
optional_params = dict(kwargs)
|
||||
|
||||
verbose_logger.debug(f"OCR optional_params after mapping: {optional_params}")
|
||||
|
||||
|
|
@ -174,7 +139,6 @@ def _prepare_ocr_request(
|
|||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=cast(dict[str, object] | None, extra_headers),
|
||||
provider_config=ocr_provider_config,
|
||||
optional_params=cast(dict[str, object], optional_params),
|
||||
litellm_params=dict(litellm_params),
|
||||
effective_timeout=effective_timeout,
|
||||
|
|
@ -182,10 +146,6 @@ def _prepare_ocr_request(
|
|||
)
|
||||
|
||||
|
||||
def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool:
|
||||
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
|
||||
|
||||
|
||||
def _rust_bridge_optional_params(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
|
|
@ -212,50 +172,62 @@ def _rust_bridge_optional_params(
|
|||
return optional_params
|
||||
|
||||
|
||||
def _is_azure_document_intelligence(prepared_request: _PreparedOCRRequest) -> bool:
|
||||
return prepared_request.custom_llm_provider == "azure_ai/doc-intelligence" or (
|
||||
prepared_request.custom_llm_provider == "azure_ai"
|
||||
and (
|
||||
"doc-intelligence" in prepared_request.model
|
||||
or "documentintelligence" in prepared_request.model
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _rust_bridge_api_base(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> str | None:
|
||||
if prepared_request.api_base is not None:
|
||||
return prepared_request.api_base
|
||||
if prepared_request.custom_llm_provider == "azure_ai/doc-intelligence":
|
||||
if _is_azure_document_intelligence(prepared_request):
|
||||
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
if prepared_request.custom_llm_provider == "azure_ai":
|
||||
if (
|
||||
"doc-intelligence" in prepared_request.model
|
||||
or "documentintelligence" in prepared_request.model
|
||||
):
|
||||
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
return resolve_secret("AZURE_AI_API_BASE")
|
||||
return None
|
||||
|
||||
|
||||
def _rust_bridge_api_key(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> str | None:
|
||||
if prepared_request.api_key is not None:
|
||||
return prepared_request.api_key
|
||||
|
||||
provider = prepared_request.custom_llm_provider
|
||||
if provider == "mistral":
|
||||
return resolve_api_key("MISTRAL_API_KEY")
|
||||
if provider == "azure_ai/doc-intelligence" or _is_azure_document_intelligence(
|
||||
prepared_request
|
||||
):
|
||||
return resolve_api_key("AZURE_DOCUMENT_INTELLIGENCE_API_KEY")
|
||||
if provider == "azure_ai":
|
||||
return resolve_api_key("AZURE_AI_API_KEY")
|
||||
if provider == "vertex_ai":
|
||||
return resolve_api_key("VERTEX_AI_API_KEY") or resolve_api_key("VERTEXAI_API_KEY")
|
||||
if provider == "reducto":
|
||||
return resolve_api_key("REDUCTO_API_KEY")
|
||||
return None
|
||||
|
||||
|
||||
def _prepare_rust_ocr_call(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> _PreparedRustOCRCall:
|
||||
provider_config = prepared_request.provider_config
|
||||
api_key_env_var = provider_config.get_api_key_env_var()
|
||||
resolved_api_key = prepared_request.api_key or (
|
||||
resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
|
||||
)
|
||||
resolved_headers = provider_config.validate_environment(
|
||||
headers=prepared_request.extra_headers or {},
|
||||
model=prepared_request.model,
|
||||
api_key=resolved_api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
resolved_complete_url = provider_config.get_complete_url(
|
||||
api_base=prepared_request.api_base,
|
||||
model=prepared_request.model,
|
||||
optional_params=prepared_request.optional_params,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
resolved_api_key = _rust_bridge_api_key(prepared_request, resolve_api_key)
|
||||
rust_api_base = _rust_bridge_api_base(prepared_request, resolve_api_key)
|
||||
rust_optional_params = _rust_bridge_optional_params(
|
||||
prepared_request, resolve_api_key
|
||||
)
|
||||
resolved_headers = prepared_request.extra_headers or {}
|
||||
prepared_request.litellm_logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=resolved_api_key,
|
||||
|
|
@ -265,7 +237,7 @@ def _prepare_rust_ocr_call(
|
|||
"document": prepared_request.document,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": resolved_complete_url,
|
||||
"api_base": rust_api_base or prepared_request.api_base,
|
||||
"headers": resolved_headers,
|
||||
},
|
||||
)
|
||||
|
|
@ -282,14 +254,6 @@ def _run_rust_ocr(
|
|||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> OCRResponse:
|
||||
"""Run the Mistral OCR call through the Rust bridge and wrap the result.
|
||||
|
||||
Resolves 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. The request that Rust actually sends (resolved URL and
|
||||
headers) is mirrored into pre_call so logs match the wire. Dependencies are
|
||||
injected so this stays unit-testable without patching module globals.
|
||||
"""
|
||||
prepared = _prepare_rust_ocr_call(
|
||||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
|
|
@ -308,6 +272,12 @@ def _run_rust_ocr(
|
|||
)
|
||||
|
||||
|
||||
def _missing_rust_bridge_error() -> RuntimeError:
|
||||
return RuntimeError(
|
||||
"Rust OCR bridge is required for litellm.ocr()/litellm.aocr(), but the native extension is unavailable"
|
||||
)
|
||||
|
||||
|
||||
async def _run_rust_aocr(
|
||||
rust_aocr: RustAocr,
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
|
|
@ -427,46 +397,17 @@ async def aocr(
|
|||
{"model": model, "custom_llm_provider": custom_llm_provider}
|
||||
)
|
||||
|
||||
if _rust_ocr_supported(prepared) and rust_ocr_enabled():
|
||||
rust_aocr = load_rust_aocr()
|
||||
if rust_aocr is None:
|
||||
verbose_logger.debug(
|
||||
"Async Rust OCR bridge unavailable; falling back to Python path"
|
||||
)
|
||||
else:
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
rust_aocr = load_rust_aocr()
|
||||
if rust_aocr is None:
|
||||
raise _missing_rust_bridge_error()
|
||||
|
||||
response = await _run_rust_aocr(
|
||||
rust_aocr=rust_aocr,
|
||||
prepared_request=prepared,
|
||||
resolve_api_key=get_secret_str,
|
||||
)
|
||||
return response
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
response = base_llm_http_handler.ocr(
|
||||
model=prepared.model,
|
||||
document=prepared.document,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared.effective_timeout,
|
||||
logging_obj=prepared.litellm_logging_obj,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared.custom_llm_provider,
|
||||
aocr=True,
|
||||
headers=prepared.extra_headers,
|
||||
provider_config=prepared.provider_config,
|
||||
litellm_params=prepared.litellm_params,
|
||||
return await _run_rust_aocr(
|
||||
rust_aocr=rust_aocr,
|
||||
prepared_request=prepared,
|
||||
resolve_api_key=get_secret_str,
|
||||
)
|
||||
|
||||
if asyncio.iscoroutine(response):
|
||||
response = await response
|
||||
|
||||
if response is None:
|
||||
raise ValueError(
|
||||
f"Got an unexpected None response from the OCR API: {response}"
|
||||
)
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
@ -704,38 +645,17 @@ def ocr(
|
|||
{"model": model, "custom_llm_provider": custom_llm_provider}
|
||||
)
|
||||
|
||||
# Optional Rust path: hand supported OCR calls to the Rust bridge.
|
||||
if _rust_ocr_supported(prepared) and rust_ocr_enabled():
|
||||
rust_ocr = load_rust_ocr()
|
||||
if rust_ocr is None:
|
||||
verbose_logger.debug(
|
||||
"Rust OCR bridge unavailable; falling back to Python path"
|
||||
)
|
||||
else:
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
rust_ocr = load_rust_ocr()
|
||||
if rust_ocr is None:
|
||||
raise _missing_rust_bridge_error()
|
||||
|
||||
return _run_rust_ocr(
|
||||
rust_ocr=rust_ocr,
|
||||
prepared_request=prepared,
|
||||
resolve_api_key=get_secret_str,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
response = base_llm_http_handler.ocr(
|
||||
model=prepared.model,
|
||||
document=prepared.document,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared.effective_timeout,
|
||||
logging_obj=prepared.litellm_logging_obj,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared.custom_llm_provider,
|
||||
aocr=_is_async,
|
||||
headers=prepared.extra_headers,
|
||||
provider_config=prepared.provider_config,
|
||||
litellm_params=prepared.litellm_params,
|
||||
return _run_rust_ocr(
|
||||
rust_ocr=rust_ocr,
|
||||
prepared_request=prepared,
|
||||
resolve_api_key=get_secret_str,
|
||||
)
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -1,9 +1,8 @@
|
|||
"""
|
||||
Optional Rust-backed OCR path.
|
||||
Rust-backed OCR path.
|
||||
|
||||
Enable with ``litellm.use_litellm_rust()``; the sync ``litellm.ocr()`` entrypoint
|
||||
then routes supported Mistral calls through the compiled ``litellm.rust_bridge._native``
|
||||
extension, which performs the whole OCR call (URL, headers, HTTP, parse) in Rust.
|
||||
Supported OCR providers route through the compiled ``litellm.rust_bridge._native``
|
||||
extension by default.
|
||||
|
||||
No module-level ``litellm`` imports keep this a leaf so ``litellm/ocr/main.py``
|
||||
can import it statically without forming an import cycle.
|
||||
|
|
@ -11,7 +10,6 @@ can import it statically without forming an import cycle.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Awaitable, Final, Protocol, cast
|
||||
|
||||
|
||||
|
|
@ -56,45 +54,27 @@ class _Unset:
|
|||
_UNSET: Final[_Unset] = _Unset()
|
||||
|
||||
|
||||
def _env_enables_rust_ocr() -> bool:
|
||||
return os.getenv("LITELLM_USE_RUST_OCR", "").strip().lower() in {
|
||||
"1",
|
||||
"true",
|
||||
"yes",
|
||||
"on",
|
||||
}
|
||||
|
||||
|
||||
_rust_ocr_enabled = _env_enables_rust_ocr()
|
||||
_rust_ocr_impl: RustOcr | None = None
|
||||
_rust_aocr_impl: RustAocr | None = None
|
||||
|
||||
|
||||
def use_litellm_rust(
|
||||
enabled: bool = True,
|
||||
*,
|
||||
def _set_rust_ocr_bridge(
|
||||
ocr: RustOcr | None | _Unset = _UNSET,
|
||||
aocr: RustAocr | None | _Unset = _UNSET,
|
||||
) -> None:
|
||||
"""Route supported OCR calls through the packaged Rust extension.
|
||||
"""Configure OCR bridge injection for tests.
|
||||
|
||||
``ocr`` and ``aocr`` inject bridge callables; when omitted the compiled
|
||||
extension is loaded on demand and any previously injected bridge is
|
||||
preserved. Pass ``None`` explicitly to clear a prior injection.
|
||||
``ocr`` and ``aocr`` inject bridge callables. When omitted, any previously
|
||||
injected bridge is preserved. Pass ``None`` explicitly to clear a prior
|
||||
injection.
|
||||
"""
|
||||
global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl
|
||||
_rust_ocr_enabled = enabled
|
||||
global _rust_ocr_impl, _rust_aocr_impl
|
||||
if not isinstance(ocr, _Unset):
|
||||
_rust_ocr_impl = ocr
|
||||
if not isinstance(aocr, _Unset):
|
||||
_rust_aocr_impl = aocr
|
||||
|
||||
|
||||
def rust_ocr_enabled() -> bool:
|
||||
"""Whether the Rust OCR path has been turned on via ``use_litellm_rust()``."""
|
||||
return _rust_ocr_enabled
|
||||
|
||||
|
||||
def load_rust_ocr() -> RustOcr | None:
|
||||
"""Return the Rust OCR callable, or ``None`` when no bridge is available.
|
||||
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Gateway E2E smoke for Rust-backed OCR.
|
|||
|
||||
Start the proxy with:
|
||||
|
||||
LITELLM_USE_RUST_OCR=1 litellm --config tests/e2e/gateway/litellm-config.yml --port 4000
|
||||
litellm --config tests/e2e/gateway/litellm-config.yml --port 4000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
"""Tests for the optional Rust-backed OCR path (``litellm/ocr/rust_bridge.py``)."""
|
||||
"""Tests for the Rust-backed OCR path (``litellm/ocr/rust_bridge.py``)."""
|
||||
|
||||
import importlib
|
||||
import builtins
|
||||
|
|
@ -151,41 +151,9 @@ class RecordingLogging:
|
|||
}
|
||||
|
||||
|
||||
class FakeOCRConfig:
|
||||
"""A stand-in ``BaseOCRConfig`` that echoes the request it would build."""
|
||||
|
||||
def __init__(self, api_key_env_var: str = "MISTRAL_API_KEY") -> None:
|
||||
self.api_key_env_var = api_key_env_var
|
||||
|
||||
def get_api_key_env_var(self) -> str:
|
||||
return self.api_key_env_var
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
*,
|
||||
headers: dict[str, object],
|
||||
model: str,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
litellm_params: dict[str, object],
|
||||
) -> dict[str, object]:
|
||||
return {"Authorization": f"Bearer {api_key}", **headers}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
*,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: dict[str, object],
|
||||
litellm_params: dict[str, object],
|
||||
) -> str:
|
||||
return f"{api_base or 'https://api.mistral.ai/v1'}/ocr"
|
||||
|
||||
|
||||
def build_prepared_request(
|
||||
*,
|
||||
logging_obj: RecordingLogging | None = None,
|
||||
provider_config: FakeOCRConfig | None = None,
|
||||
model: str = "mistral-ocr-latest",
|
||||
document: dict[str, object] = DOCUMENT,
|
||||
api_key: str | None = "sk-test",
|
||||
|
|
@ -203,7 +171,6 @@ def build_prepared_request(
|
|||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
provider_config=provider_config or FakeOCRConfig(),
|
||||
optional_params=optional_params or {},
|
||||
litellm_params=litellm_params or {},
|
||||
effective_timeout=timeout,
|
||||
|
|
@ -212,47 +179,34 @@ def build_prepared_request(
|
|||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_rust_flag():
|
||||
"""Keep the global toggle isolated between tests."""
|
||||
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None)
|
||||
def _reset_rust_bridge():
|
||||
"""Keep the global bridge state isolated between tests."""
|
||||
rust_bridge._set_rust_ocr_bridge(ocr=None, aocr=None)
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
yield
|
||||
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None)
|
||||
rust_bridge._set_rust_ocr_bridge(ocr=None, aocr=None)
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_bridge():
|
||||
"""Enable the Rust path with an injected recording bridge (no native wheel)."""
|
||||
"""Inject a recording bridge."""
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
rust_bridge._set_rust_ocr_bridge(ocr=bridge)
|
||||
return bridge
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_async_bridge():
|
||||
"""Enable the async Rust path with an injected recording bridge."""
|
||||
"""Inject an async recording bridge."""
|
||||
bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, aocr=bridge)
|
||||
rust_bridge._set_rust_ocr_bridge(aocr=bridge)
|
||||
return bridge
|
||||
|
||||
|
||||
def test_use_litellm_rust_toggles_flag():
|
||||
assert rust_bridge.rust_ocr_enabled() is False
|
||||
litellm.use_litellm_rust()
|
||||
assert rust_bridge.rust_ocr_enabled() is True
|
||||
litellm.use_litellm_rust(False)
|
||||
assert rust_bridge.rust_ocr_enabled() is False
|
||||
|
||||
|
||||
def test_env_var_enables_rust_ocr(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_USE_RUST_OCR", "1")
|
||||
assert rust_bridge._env_enables_rust_ocr() is True
|
||||
|
||||
|
||||
def test_load_rust_ocr_returns_injected_impl():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
rust_bridge._set_rust_ocr_bridge(ocr=bridge)
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
|
||||
|
||||
|
|
@ -296,25 +250,16 @@ def test_native_bridge_available_reflects_loader(monkeypatch):
|
|||
|
||||
def test_load_rust_aocr_returns_injected_impl():
|
||||
bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, aocr=bridge)
|
||||
rust_bridge._set_rust_ocr_bridge(aocr=bridge)
|
||||
assert rust_bridge.load_rust_aocr() is bridge
|
||||
|
||||
|
||||
def test_toggle_without_ocr_arg_preserves_injected_impl():
|
||||
"""Regression: routine enable/disable calls must not clobber a prior injection.
|
||||
|
||||
Earlier, ``use_litellm_rust()`` unconditionally assigned the keyword default
|
||||
of ``None`` to ``_rust_ocr_impl``, silently dropping a custom bridge whenever
|
||||
a caller toggled the flag without re-passing ``ocr=``.
|
||||
"""
|
||||
def test_bridge_injection_preserves_unspecified_impl():
|
||||
bridge = RecordingBridge()
|
||||
async_bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge, aocr=async_bridge)
|
||||
rust_bridge._set_rust_ocr_bridge(ocr=bridge, aocr=async_bridge)
|
||||
|
||||
litellm.use_litellm_rust(False)
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
assert rust_bridge.load_rust_aocr() is async_bridge
|
||||
litellm.use_litellm_rust(True)
|
||||
rust_bridge._set_rust_ocr_bridge()
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
assert rust_bridge.load_rust_aocr() is async_bridge
|
||||
|
||||
|
|
@ -327,9 +272,9 @@ def test_explicit_ocr_none_clears_injected_impl(monkeypatch):
|
|||
)
|
||||
bridge = RecordingBridge()
|
||||
async_bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge, aocr=async_bridge)
|
||||
rust_bridge._set_rust_ocr_bridge(ocr=bridge, aocr=async_bridge)
|
||||
|
||||
litellm.use_litellm_rust(True, ocr=None, aocr=None)
|
||||
rust_bridge._set_rust_ocr_bridge(ocr=None, aocr=None)
|
||||
assert rust_bridge.load_rust_ocr() is None
|
||||
assert rust_bridge.load_rust_aocr() is None
|
||||
|
||||
|
|
@ -342,7 +287,6 @@ def test_load_rust_ocr_none_when_extension_absent(monkeypatch):
|
|||
"get_native_bridge",
|
||||
lambda: None,
|
||||
)
|
||||
litellm.use_litellm_rust(True) # no impl injected; extension isn't built in CI
|
||||
assert rust_bridge.load_rust_ocr() is None
|
||||
assert rust_bridge.load_rust_aocr() is None
|
||||
|
||||
|
|
@ -360,7 +304,6 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch):
|
|||
lambda: fake_module,
|
||||
)
|
||||
|
||||
litellm.use_litellm_rust(True) # enabled, no impl injected -> import the extension
|
||||
assert rust_bridge.load_rust_ocr() is fake_module.ocr
|
||||
assert rust_bridge.load_rust_aocr() is fake_module.aocr
|
||||
|
||||
|
|
@ -396,10 +339,7 @@ def test_run_rust_ocr_forwards_args_and_wraps_response():
|
|||
"api_key": "sk-test",
|
||||
"api_base": "https://proxy.internal",
|
||||
"custom_llm_provider": "mistral",
|
||||
"extra_headers": {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
},
|
||||
"extra_headers": {"x-trace-id": "trace-1"},
|
||||
"optional_params": {"include_image_base64": True},
|
||||
"timeout_seconds": 12.5,
|
||||
}
|
||||
|
|
@ -421,7 +361,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
|
|||
assert bridge.calls[0]["api_key"] == "sk-from-vault"
|
||||
|
||||
|
||||
def test_run_rust_ocr_uses_provider_api_key_env_var():
|
||||
def test_run_rust_ocr_resolves_reducto_key():
|
||||
bridge = RecordingBridge()
|
||||
resolver_calls = []
|
||||
|
||||
|
|
@ -432,15 +372,15 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
|
|||
ocr_main._run_rust_ocr(
|
||||
rust_ocr=bridge,
|
||||
prepared_request=build_prepared_request(
|
||||
provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"),
|
||||
model="provider-ocr-model",
|
||||
custom_llm_provider="reducto",
|
||||
model="parse-v3",
|
||||
api_key=None,
|
||||
timeout=None,
|
||||
),
|
||||
resolve_api_key=_resolver,
|
||||
)
|
||||
|
||||
assert resolver_calls == ["PROVIDER_OCR_API_KEY"]
|
||||
assert resolver_calls == ["REDUCTO_API_KEY"]
|
||||
assert bridge.calls[0]["api_key"] == "sk-provider-env"
|
||||
|
||||
|
||||
|
|
@ -573,15 +513,11 @@ def test_run_rust_ocr_runs_pre_call_logging():
|
|||
complete_input = additional_args["complete_input_dict"]
|
||||
assert complete_input["document"] == DOCUMENT
|
||||
assert complete_input["include_image_base64"] is True
|
||||
# The logged request mirrors what Rust sends: resolved URL + headers.
|
||||
assert additional_args["api_base"] == "https://api.mistral.ai/v1/ocr"
|
||||
assert additional_args["headers"] == {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
}
|
||||
assert additional_args["api_base"] == "https://api.mistral.ai/v1"
|
||||
assert additional_args["headers"] == {"x-trace-id": "trace-1"}
|
||||
|
||||
|
||||
def test_ocr_routes_to_rust_when_enabled(fake_bridge):
|
||||
def test_ocr_routes_to_rust_by_default(fake_bridge):
|
||||
response = litellm.ocr(
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
|
|
@ -599,15 +535,12 @@ def test_ocr_routes_to_rust_when_enabled(fake_bridge):
|
|||
assert call["document"] == DOCUMENT
|
||||
assert call["api_key"] == "sk-test"
|
||||
assert call["custom_llm_provider"] == "mistral"
|
||||
assert call["extra_headers"] == {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
}
|
||||
assert call["extra_headers"] == {"x-trace-id": "trace-1"}
|
||||
# Raw OCR params ride along in optional_params; Rust filters to supported keys.
|
||||
assert call["optional_params"].get("include_image_base64") is True
|
||||
|
||||
|
||||
def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge):
|
||||
def test_ocr_routes_azure_ai_to_rust_by_default(fake_bridge):
|
||||
response = litellm.ocr(
|
||||
model="azure_ai/pixtral-12b-2409",
|
||||
document=DOCUMENT,
|
||||
|
|
@ -631,7 +564,7 @@ def test_ocr_exception_type_uses_resolved_provider_context(
|
|||
return CapturedException("wrapped")
|
||||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
litellm.use_litellm_rust(True, ocr=RaisingBridge())
|
||||
rust_bridge._set_rust_ocr_bridge(ocr=RaisingBridge())
|
||||
|
||||
with pytest.raises(CapturedException):
|
||||
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
|
@ -641,7 +574,7 @@ def test_ocr_exception_type_uses_resolved_provider_context(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_routes_to_async_rust_when_enabled(fake_async_bridge):
|
||||
async def test_aocr_routes_to_async_rust_by_default(fake_async_bridge):
|
||||
response = await litellm.aocr(
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
|
|
@ -658,10 +591,7 @@ async def test_aocr_routes_to_async_rust_when_enabled(fake_async_bridge):
|
|||
assert call["document"] == DOCUMENT
|
||||
assert call["api_key"] == "sk-test"
|
||||
assert call["custom_llm_provider"] == "mistral"
|
||||
assert call["extra_headers"] == {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
}
|
||||
assert call["extra_headers"] == {"x-trace-id": "trace-1"}
|
||||
assert call["optional_params"].get("include_image_base64") is True
|
||||
|
||||
|
||||
|
|
@ -676,7 +606,7 @@ async def test_aocr_exception_type_uses_resolved_provider_context(
|
|||
return CapturedException("wrapped")
|
||||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
litellm.use_litellm_rust(True, aocr=RaisingAsyncBridge())
|
||||
rust_bridge._set_rust_ocr_bridge(aocr=RaisingAsyncBridge())
|
||||
|
||||
with pytest.raises(CapturedException):
|
||||
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
|
@ -703,55 +633,10 @@ def test_ocr_passes_default_request_timeout_to_rust(fake_bridge):
|
|||
assert fake_bridge.calls[0]["timeout_seconds"] == float(request_timeout)
|
||||
|
||||
|
||||
def test_ocr_does_not_route_to_rust_when_disabled():
|
||||
"""With the flag off, the bridge must not be consulted even if an impl exists."""
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(False, ocr=bridge)
|
||||
|
||||
assert rust_bridge.rust_ocr_enabled() is False
|
||||
# The impl stays available for injection, but the disabled flag gates usage,
|
||||
# so ocr() never reaches the Rust path (asserted via the enabled-path test).
|
||||
assert bridge.calls == []
|
||||
|
||||
|
||||
def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch):
|
||||
"""Rust enabled but no bridge available (no injected impl, no compiled wheel):
|
||||
ocr() must degrade to the Python HTTP handler instead of raising."""
|
||||
def test_ocr_requires_rust_bridge_when_unavailable(monkeypatch):
|
||||
monkeypatch.setattr(ocr_main, "load_rust_ocr", lambda: None)
|
||||
litellm.use_litellm_rust(True) # enabled, but load_rust_ocr() returns None in CI
|
||||
|
||||
captured = {}
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
||||
def fake_handler_ocr(**kwargs):
|
||||
captured["called"] = True
|
||||
return OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr")
|
||||
|
||||
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fake_handler_ocr)
|
||||
|
||||
response = litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
||||
assert captured.get("called") is True # Python path was used
|
||||
assert isinstance(response, OCRResponse)
|
||||
|
||||
|
||||
def test_ocr_provider_configs_expose_api_key_env_vars():
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
)
|
||||
from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
|
||||
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
|
||||
from litellm.llms.vertex_ai.ocr.deepseek_transformation import (
|
||||
VertexAIDeepSeekOCRConfig,
|
||||
)
|
||||
from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig
|
||||
|
||||
assert BaseOCRConfig().get_api_key_env_var() is None
|
||||
assert MistralOCRConfig().get_api_key_env_var() == "MISTRAL_API_KEY"
|
||||
assert AzureAIOCRConfig().get_api_key_env_var() == "AZURE_AI_API_KEY"
|
||||
assert (
|
||||
AzureDocumentIntelligenceOCRConfig().get_api_key_env_var()
|
||||
== "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"
|
||||
)
|
||||
assert VertexAIOCRConfig().get_api_key_env_var() == "VERTEX_AI_API_KEY"
|
||||
assert VertexAIDeepSeekOCRConfig().get_api_key_env_var() == "VERTEX_AI_API_KEY"
|
||||
assert "Rust OCR bridge is required" in str(exc_info.value)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue