diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index ce86a0ee6ac..422dfb20065 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -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" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 5842ed5ba9b..00af23b4c00 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -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" diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr.rs b/litellm-rust/crates/ai-gateway/src/io/ocr.rs index 35e511fa982..542880455a0 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr.rs @@ -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 { &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()); } diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr/common_utils.rs b/litellm-rust/crates/ai-gateway/src/io/ocr/common_utils.rs index ee0d86c3000..35171de9a9f 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr/common_utils.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/common_utils.rs @@ -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, 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 { + 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, + mime: String, + parse_url: &str, + headers: &[(String, String)], + timeout: Option, +) -> CoreResult { + 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, +) -> CoreResult { + 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; diff --git a/litellm-rust/crates/core/src/ocr/transformation.rs b/litellm-rust/crates/core/src/ocr/transformation.rs index cb3e735e533..39a17b44ca0 100644 --- a/litellm-rust/crates/core/src/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/ocr/transformation.rs @@ -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 } diff --git a/litellm-rust/crates/core/src/providers/mod.rs b/litellm-rust/crates/core/src/providers/mod.rs index d75e750a0ba..cf12db66f0b 100644 --- a/litellm-rust/crates/core/src/providers/mod.rs +++ b/litellm-rust/crates/core/src/providers/mod.rs @@ -1,4 +1,5 @@ pub mod azure_ai; pub mod mistral; pub mod openai; +pub mod reducto; pub mod vertex_ai; diff --git a/litellm-rust/crates/core/src/providers/reducto/mod.rs b/litellm-rust/crates/core/src/providers/reducto/mod.rs new file mode 100644 index 00000000000..3621ff6a2fd --- /dev/null +++ b/litellm-rust/crates/core/src/providers/reducto/mod.rs @@ -0,0 +1 @@ +pub mod ocr; diff --git a/litellm-rust/crates/core/src/providers/reducto/ocr/mod.rs b/litellm-rust/crates/core/src/providers/reducto/ocr/mod.rs new file mode 100644 index 00000000000..f239b6921fa --- /dev/null +++ b/litellm-rust/crates/core/src/providers/reducto/ocr/mod.rs @@ -0,0 +1 @@ +pub mod transformation; diff --git a/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs new file mode 100644 index 00000000000..b1241a14fea --- /dev/null +++ b/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs @@ -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, +) -> CoreResult { + 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 { + 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::().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 { + let chunks = result + .get("chunks") + .and_then(Value::as_array) + .cloned() + .unwrap_or_default(); + let mut blocks_by_page: BTreeMap> = 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::>() + .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::>() + .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, + ) -> CoreResult { + 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 { + transform_reducto_response(model, response_json) + } + + fn complete_url( + &self, + api_base: Option<&str>, + _model: &str, + _optional_params: &Map, + _env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + Ok(complete_url(api_base)) + } + + fn resolve_api_key( + &self, + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + 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, + ) -> CoreResult { + 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 { + transform_reducto_response(model, response_json) + } + + fn complete_url( + &self, + api_base: Option<&str>, + _model: &str, + _optional_params: &Map, + _env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + Ok(complete_url(api_base)) + } + + fn resolve_api_key( + &self, + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + resolve_api_key(api_key, env_lookup) + } + + fn document_preparation(&self) -> OcrDocumentPreparation { + OcrDocumentPreparation::ReductoUpload + } +} + +fn transform_reducto_response(model: &str, response_json: Value) -> CoreResult { + 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})) + ); + } +} diff --git a/litellm/__init__.py b/litellm/__init__.py index 9650dc12c97..6085ebf9f96 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 * diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 6a196d41768..52244b40ab8 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -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, diff --git a/litellm/ocr/rust_bridge.py b/litellm/ocr/rust_bridge.py index 631fd4c63c5..c518c2cb874 100644 --- a/litellm/ocr/rust_bridge.py +++ b/litellm/ocr/rust_bridge.py @@ -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. diff --git a/tests/e2e/gateway/test_ocr_rust_e2e.py b/tests/e2e/gateway/test_ocr_rust_e2e.py index 6ce59b2b5ac..dcf85898365 100644 --- a/tests/e2e/gateway/test_ocr_rust_e2e.py +++ b/tests/e2e/gateway/test_ocr_rust_e2e.py @@ -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 diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 7e23e441f50..7136cac4362 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -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)