feat(ocr): route all OCR through Rust

This commit is contained in:
Ishaan Jaff 2026-06-25 14:16:42 -07:00
parent d860cd9599
commit 1528c9b9d1
No known key found for this signature in database
14 changed files with 618 additions and 330 deletions

View file

@ -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"

View file

@ -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"

View file

@ -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());
}

View file

@ -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;

View file

@ -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
}

View file

@ -1,4 +1,5 @@
pub mod azure_ai;
pub mod mistral;
pub mod openai;
pub mod reducto;
pub mod vertex_ai;

View file

@ -0,0 +1 @@
pub mod ocr;

View file

@ -0,0 +1 @@
pub mod transformation;

View file

@ -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}))
);
}
}

View file

@ -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 *

View file

@ -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,

View file

@ -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.

View file

@ -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

View file

@ -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)