merge litellm_ocr_rust_default into litellm_shared_reqwest_error_mapper

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-17 04:08:17 +00:00
commit ed5cb16589
11 changed files with 1560 additions and 447 deletions

View file

@ -0,0 +1,85 @@
use crate::constants::ENV_REFERENCE_PREFIX;
pub(crate) fn resolve_env_reference(
value: Option<&str>,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Option<String> {
let value = value?;
let Some(name) = value.strip_prefix(ENV_REFERENCE_PREFIX) else {
return Some(value.to_string());
};
if name.trim().is_empty() {
return None;
}
env_lookup(name).filter(|resolved| !resolved.trim().is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
fn env_lookup(name: &str) -> Option<String> {
match name {
"PRESENT" => Some("resolved".to_string()),
"BLANK" => Some(" ".to_string()),
_ => None,
}
}
#[test]
fn preserves_explicit_value() {
assert_eq!(
resolve_env_reference(Some("explicit"), &env_lookup),
Some("explicit".to_string())
);
}
#[test]
fn preserves_value_that_only_contains_reference_prefix() {
assert_eq!(
resolve_env_reference(Some("prefix-os.environ/PRESENT"), &env_lookup),
Some("prefix-os.environ/PRESENT".to_string())
);
}
#[test]
fn resolves_present_reference() {
assert_eq!(
resolve_env_reference(Some("os.environ/PRESENT"), &env_lookup),
Some("resolved".to_string())
);
}
#[test]
fn missing_reference_is_absent() {
assert_eq!(
resolve_env_reference(Some("os.environ/MISSING"), &env_lookup),
None
);
}
#[test]
fn blank_reference_value_is_absent() {
assert_eq!(
resolve_env_reference(Some("os.environ/BLANK"), &env_lookup),
None
);
}
#[test]
fn malformed_reference_is_absent() {
assert_eq!(
resolve_env_reference(Some("os.environ/"), &env_lookup),
None
);
assert_eq!(
resolve_env_reference(Some("os.environ/ "), &env_lookup),
None
);
}
#[test]
fn absent_input_stays_absent() {
assert_eq!(resolve_env_reference(None, &env_lookup), None);
}
}

View file

@ -7,23 +7,31 @@
/// Default LiteLLM control-plane base URL for request-log egress when
/// `LITELLM_PROXY_BASE_URL` is unset.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_PROXY_BASE_URL: &str = "http://localhost:4000";
/// The logs ingest path appended to the proxy base. Not a tunable; it is the
/// proxy's API contract (the rust-control-plane router on the Python proxy).
#[cfg(feature = "server")]
pub(crate) const RUST_CONTROL_PLANE_LOGS_PATH: &str = "/v1/rust_control_plane/logs";
/// Default bounded channel depth for the log-egress worker.
/// Override: `LITELLM_LOG_CHANNEL_CAPACITY`.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_CHANNEL_CAPACITY: usize = 4096;
/// Default max records POSTed per request to the control plane.
/// Override: `LITELLM_LOG_BATCH_SIZE`.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_MAX_BATCH_SIZE: usize = 256;
/// Default partial-batch flush cadence, in ms.
/// Override: `LITELLM_LOG_FLUSH_INTERVAL_MS`.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500;
/// Provider attributed to realtime sessions in the logging payload.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_PROVIDER: &str = "openai";
pub(crate) const ENV_REFERENCE_PREFIX: &str = "os.environ/";

View file

@ -14,6 +14,8 @@ use litellm_core::ocr::transformation::{
use litellm_core::CoreResult;
use serde_json::{Map, Value};
use crate::config::resolve_env_reference;
mod common_utils;
use crate::errors::map_reqwest_error;
@ -74,6 +76,14 @@ pub struct OcrRequest<'a> {
///
/// Async: intended to be awaited directly by the Python bridge's async entrypoint.
pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
let env_lookup = |key: &str| std::env::var(key).ok();
ocr_with_env(request, &env_lookup).await
}
async fn ocr_with_env(
request: OcrRequest<'_>,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> CoreResult<Value> {
let model = request.model;
let config = ocr_provider_config(request.custom_llm_provider, model).ok_or_else(|| {
CoreError::InvalidProvider(format!(
@ -81,18 +91,19 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
request.custom_llm_provider
))
})?;
let env_lookup = |key: &str| std::env::var(key).ok();
let api_key = resolve_env_reference(request.api_key, env_lookup);
let api_base = resolve_env_reference(request.api_base, env_lookup);
let headers = string_headers(request.extra_headers)?;
let auth_strategy = config.auth_strategy();
let api_key = (!has_header(&headers, auth_strategy.header_name()))
.then(|| config.resolve_api_key(request.api_key, &env_lookup))
.then(|| config.resolve_api_key(api_key.as_deref(), env_lookup))
.transpose()?;
let url = config.complete_url(
request.api_base,
api_base.as_deref(),
model,
&request.optional_params,
&env_lookup,
env_lookup,
)?;
let filtered_params = config.map_ocr_params(&request.optional_params);
let upstream_headers = upstream_headers(&headers, auth_strategy, api_key.as_deref());
@ -167,394 +178,4 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
async fn read_http_headers(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
if request.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
String::from_utf8(request).expect("request is utf8")
}
#[test]
fn truncate_error_body_passes_short_strings_through() {
let body = "Unauthorized";
assert_eq!(truncate_error_body(body), "Unauthorized");
}
#[test]
fn truncate_error_body_caps_long_payloads() {
let body = "x".repeat(306);
let truncated = truncate_error_body(&body);
assert!(truncated.ends_with("... (truncated)"));
let prefix_chars = truncated
.strip_suffix("... (truncated)")
.expect("truncated marker present")
.chars()
.count();
assert_eq!(prefix_chars, 256);
}
#[test]
fn truncate_error_body_does_not_split_multibyte_chars() {
let body = "é".repeat(266);
let truncated = truncate_error_body(&body);
assert!(truncated.is_char_boundary(truncated.len()));
}
#[test]
fn ocr_dispatch_supports_migrated_providers() {
assert!(ocr_provider_config("mistral", "mistral-ocr-latest").is_some());
assert!(ocr_provider_config("azure_ai", "pixtral-12b-2409")
.expect("azure ai config resolves")
.requires_data_uri_document());
assert_eq!(
ocr_provider_config("azure_ai", "doc-intelligence/prebuilt-read")
.expect("document intelligence config resolves")
.response_handling(),
OcrResponseHandling::AzureDocumentIntelligencePoll
);
assert!(ocr_provider_config("vertex_ai", "deepseek-ocr-maas")
.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());
}
#[test]
fn string_headers_accepts_string_values() {
let headers = json!({
"x-trace-id": "trace-1"
})
.as_object()
.unwrap()
.clone();
assert_eq!(
string_headers(Some(headers)).expect("string headers accepted"),
vec![("x-trace-id".to_string(), "trace-1".to_string())]
);
}
#[test]
fn auth_header_detection_is_case_insensitive() {
let headers = vec![
("x-trace-id".to_string(), "trace-1".to_string()),
("authorization".to_string(), "Bearer sk-test".to_string()),
];
assert!(has_header(&headers, "authorization"));
let headers = vec![("Authorization".to_string(), "Bearer sk-test".to_string())];
assert!(has_header(&headers, "authorization"));
let headers = vec![("x-trace-id".to_string(), "trace-1".to_string())];
assert!(!has_header(&headers, "authorization"));
}
#[tokio::test]
async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let request = read_http_headers(&mut socket).await;
let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
request
});
let mut headers = Map::new();
headers.insert(
"Authorization".to_string(),
Value::String("Bearer sk-from-python".to_string()),
);
headers.insert(
"x-trace-id".to_string(),
Value::String("trace-1".to_string()),
);
let response = ocr(OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-for-rust-fallback"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: "mistral",
extra_headers: Some(headers),
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
})
.await
.expect("ocr request succeeds");
assert_eq!(response["pages"][0]["markdown"], "ok");
let request = server.await.expect("server task completes");
let authorization_count = request
.lines()
.filter(|line| line.to_ascii_lowercase().starts_with("authorization:"))
.count();
assert_eq!(authorization_count, 1, "{request}");
assert!(
request.contains("authorization: Bearer sk-from-python")
|| request.contains("Authorization: Bearer sk-from-python"),
"{request}"
);
}
#[tokio::test]
async fn document_intelligence_poll_uses_resolved_subscription_key() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let operation_url = format!("http://{addr}/operations/1");
let server = tokio::spawn(async move {
let (mut post_socket, _) = listener.accept().await.expect("accepts post request");
let post_request = read_http_headers(&mut post_socket).await;
let post_response = format!(
"HTTP/1.1 202 Accepted\r\noperation-location: {operation_url}\r\ncontent-length: 0\r\nconnection: close\r\n\r\n"
);
post_socket
.write_all(post_response.as_bytes())
.await
.expect("writes post response");
let (mut poll_socket, _) = listener.accept().await.expect("accepts poll request");
let poll_request = read_http_headers(&mut poll_socket).await;
let response_body = r#"{"status":"succeeded","analyzeResult":{"pages":[{"pageNumber":1,"lines":[{"content":"ok"}]}]}}"#;
let poll_response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
poll_socket
.write_all(poll_response.as_bytes())
.await
.expect("writes poll response");
(post_request, poll_request)
});
let response = ocr(OcrRequest {
model: "prebuilt-read",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("di-key"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: "azure_ai/doc-intelligence",
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
})
.await
.expect("document intelligence request succeeds");
assert_eq!(response["pages"][0]["markdown"], "ok");
let (post_request, poll_request) = server.await.expect("server task completes");
assert!(
post_request
.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: di-key"),
"{post_request}"
);
assert!(
poll_request
.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: di-key"),
"{poll_request}"
);
}
#[test]
fn string_headers_rejects_non_string_values() {
let headers = json!({
"x-retry-count": 3
})
.as_object()
.unwrap()
.clone();
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
assert_eq!(
err,
CoreError::InvalidRequest(
"OCR extra_headers.x-retry-count must be a string, got number".to_string()
)
);
}
async fn respond_once(status_line: &'static str, body: String) -> String {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let _ = read_http_headers(&mut socket).await;
let response = format!(
"HTTP/1.1 {status_line}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
});
format!("http://{addr}")
}
async fn run_mistral_ocr(api_base: String, timeout: Duration) -> CoreResult<Value> {
ocr(OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: Some(&api_base),
custom_llm_provider: "mistral",
extra_headers: None,
optional_params: Map::new(),
timeout: Some(timeout),
})
.await
}
#[tokio::test]
async fn ocr_preserves_upstream_error_status() {
for status in [
"401 Unauthorized",
"404 Not Found",
"500 Internal Server Error",
] {
let base = respond_once(status, r#"{"error":"nope"}"#.to_string()).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("upstream error surfaces");
let expected = status[..3].parse::<u16>().expect("status prefix parses");
match err {
CoreError::Http { status: got, .. } => assert_eq!(got, expected),
other => panic!("expected Http error, got {other:?}"),
}
assert_eq!(err.public_status_code(), Some(expected));
}
}
#[tokio::test]
async fn ocr_bounds_oversized_error_body() {
let base = respond_once("500 Internal Server Error", "x".repeat(6000)).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("oversized error surfaces");
match err {
CoreError::Http { body, .. } => {
assert!(body.ends_with("... (truncated)"));
assert!(
body.chars().count() < 300,
"body not bounded: {} chars",
body.chars().count()
);
}
other => panic!("expected Http error, got {other:?}"),
}
}
#[tokio::test]
async fn ocr_rejects_invalid_json_success_body() {
let base = respond_once("200 OK", "not json".to_string()).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("invalid JSON surfaces");
assert!(matches!(err, CoreError::InvalidResponse(_)));
assert_eq!(err.public_status_code(), Some(500));
}
#[tokio::test]
async fn ocr_rejects_empty_success_body() {
let base = respond_once("200 OK", " ".to_string()).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("empty success surfaces");
match err {
CoreError::InvalidResponse(message) => assert!(message.contains("empty")),
other => panic!("expected InvalidResponse, got {other:?}"),
}
}
#[tokio::test]
async fn ocr_classifies_per_request_timeout() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let _ = read_http_headers(&mut socket).await;
tokio::time::sleep(Duration::from_secs(3)).await;
let _ = socket.write_all(b"HTTP/1.1 200 OK\r\n\r\n").await;
});
let err = run_mistral_ocr(format!("http://{addr}"), Duration::from_millis(200))
.await
.expect_err("timeout surfaces");
assert_eq!(err, CoreError::Timeout);
assert_eq!(err.public_status_code(), Some(408));
}
#[tokio::test]
async fn ocr_maps_unregistered_provider_to_invalid_provider() {
let err = ocr(OcrRequest {
model: "some-model",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: "definitely-not-a-provider",
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
})
.await
.expect_err("unregistered provider surfaces");
assert!(matches!(err, CoreError::InvalidProvider(_)));
assert_eq!(err.public_status_code(), Some(400));
assert_eq!(err.public_message(), "Invalid OCR request");
}
}
mod tests;

View file

@ -0,0 +1,573 @@
use super::*;
use serde_json::json;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
async fn read_http_headers(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
if request.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
String::from_utf8(request).expect("request is utf8")
}
async fn read_http_request(socket: &mut TcpStream) -> (String, Value) {
let mut raw = Vec::new();
let mut buffer = [0_u8; 1024];
loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break;
}
raw.extend_from_slice(&buffer[..n]);
let header_end = raw
.windows(4)
.position(|window| window == b"\r\n\r\n")
.map(|pos| pos + 4);
if let Some(body_start) = header_end {
let text = String::from_utf8(raw.clone()).expect("request is utf8");
let content_length = text
.lines()
.find_map(|line| {
line.to_ascii_lowercase()
.strip_prefix("content-length:")
.map(|value| value.trim().parse::<usize>().expect("content-length"))
})
.unwrap_or(0);
if raw.len() >= body_start + content_length {
let headers = text[..body_start].to_string();
let body: Value =
serde_json::from_slice(&raw[body_start..body_start + content_length])
.expect("request body is json");
return (headers, body);
}
}
}
panic!("did not receive a complete request");
}
#[test]
fn truncate_error_body_passes_short_strings_through() {
let body = "Unauthorized";
assert_eq!(truncate_error_body(body), "Unauthorized");
}
#[test]
fn truncate_error_body_caps_long_payloads() {
let body = "x".repeat(306);
let truncated = truncate_error_body(&body);
assert!(truncated.ends_with("... (truncated)"));
let prefix_chars = truncated
.strip_suffix("... (truncated)")
.expect("truncated marker present")
.chars()
.count();
assert_eq!(prefix_chars, 256);
}
#[test]
fn truncate_error_body_does_not_split_multibyte_chars() {
let body = "é".repeat(266);
let truncated = truncate_error_body(&body);
assert!(truncated.is_char_boundary(truncated.len()));
}
#[test]
fn ocr_dispatch_supports_migrated_providers() {
assert!(ocr_provider_config("mistral", "mistral-ocr-latest").is_some());
assert!(ocr_provider_config("azure_ai", "pixtral-12b-2409")
.expect("azure ai config resolves")
.requires_data_uri_document());
assert_eq!(
ocr_provider_config("azure_ai", "doc-intelligence/prebuilt-read")
.expect("document intelligence config resolves")
.response_handling(),
OcrResponseHandling::AzureDocumentIntelligencePoll
);
assert!(ocr_provider_config("vertex_ai", "deepseek-ocr-maas")
.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());
}
#[test]
fn string_headers_accepts_string_values() {
let headers = json!({
"x-trace-id": "trace-1"
})
.as_object()
.unwrap()
.clone();
assert_eq!(
string_headers(Some(headers)).expect("string headers accepted"),
vec![("x-trace-id".to_string(), "trace-1".to_string())]
);
}
#[test]
fn auth_header_detection_is_case_insensitive() {
let headers = vec![
("x-trace-id".to_string(), "trace-1".to_string()),
("authorization".to_string(), "Bearer sk-test".to_string()),
];
assert!(has_header(&headers, "authorization"));
let headers = vec![("Authorization".to_string(), "Bearer sk-test".to_string())];
assert!(has_header(&headers, "authorization"));
let headers = vec![("x-trace-id".to_string(), "trace-1".to_string())];
assert!(!has_header(&headers, "authorization"));
}
#[tokio::test]
async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let request = read_http_headers(&mut socket).await;
let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
request
});
let mut headers = Map::new();
headers.insert(
"Authorization".to_string(),
Value::String("Bearer sk-from-python".to_string()),
);
headers.insert(
"x-trace-id".to_string(),
Value::String("trace-1".to_string()),
);
let response = ocr(OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-for-rust-fallback"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: "mistral",
extra_headers: Some(headers),
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
})
.await
.expect("ocr request succeeds");
assert_eq!(response["pages"][0]["markdown"], "ok");
let request = server.await.expect("server task completes");
let authorization_count = request
.lines()
.filter(|line| line.to_ascii_lowercase().starts_with("authorization:"))
.count();
assert_eq!(authorization_count, 1, "{request}");
assert!(
request.contains("authorization: Bearer sk-from-python")
|| request.contains("Authorization: Bearer sk-from-python"),
"{request}"
);
}
#[tokio::test]
async fn ocr_resolves_api_key_and_base_references_in_rust() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let resolved_base = format!("http://{addr}");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let request = read_http_headers(&mut socket).await;
let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
request
});
let env_lookup = |name: &str| match name {
"OCR_TEST_API_KEY" => Some("sk-resolved".to_string()),
"OCR_TEST_API_BASE" => Some(resolved_base.clone()),
_ => None,
};
let response = ocr_with_env(
OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("os.environ/OCR_TEST_API_KEY"),
api_base: Some("os.environ/OCR_TEST_API_BASE"),
custom_llm_provider: "mistral",
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
},
&env_lookup,
)
.await
.expect("ocr request succeeds");
assert_eq!(response["pages"][0]["markdown"], "ok");
let request = server.await.expect("server task completes");
assert!(request.starts_with("POST /v1/ocr HTTP/1.1"), "{request}");
assert!(
request
.to_ascii_lowercase()
.contains("authorization: bearer sk-resolved"),
"{request}"
);
assert!(!request.contains("os.environ/"), "{request}");
}
#[tokio::test]
async fn ocr_forwards_full_mistral_contract_and_filters_internal_params() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let (_headers, body) = read_http_request(&mut socket).await;
let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
body
});
let document = json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
});
let optional_params = json!({
"pages": [0, 2, 5],
"include_image_base64": true,
"image_limit": 10,
"image_min_size": 64,
"bbox_annotation_format": {"type": "text"},
"document_annotation_format": {"type": "json_schema"},
"document_annotation_prompt": "extract title",
"extract_header": true,
"extract_footer": false,
"table_format": "html",
"confidence_scores_granularity": "word",
"include_blocks": true,
"id": "ocr-req-9",
"litellm_metadata": {"trace": "internal"},
"metadata": {"trace": "internal"},
"num_retries": 3
})
.as_object()
.unwrap()
.clone();
ocr(OcrRequest {
model: "mistral-ocr-latest",
document: document.clone(),
api_key: Some("sk-test"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: "mistral",
extra_headers: None,
optional_params,
timeout: Some(Duration::from_secs(5)),
})
.await
.expect("ocr request succeeds");
let body = server.await.expect("server task completes");
assert_eq!(
body,
json!({
"model": "mistral-ocr-latest",
"document": document,
"pages": [0, 2, 5],
"include_image_base64": true,
"image_limit": 10,
"image_min_size": 64,
"bbox_annotation_format": {"type": "text"},
"document_annotation_format": {"type": "json_schema"},
"document_annotation_prompt": "extract title",
"extract_header": true,
"extract_footer": false,
"table_format": "html",
"confidence_scores_granularity": "word",
"include_blocks": true,
"id": "ocr-req-9"
})
);
}
#[tokio::test]
async fn document_intelligence_poll_uses_resolved_subscription_key() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let operation_url = format!("http://{addr}/operations/1");
let server = tokio::spawn(async move {
let (mut post_socket, _) = listener.accept().await.expect("accepts post request");
let post_request = read_http_headers(&mut post_socket).await;
let post_response = format!(
"HTTP/1.1 202 Accepted\r\noperation-location: {operation_url}\r\ncontent-length: 0\r\nconnection: close\r\n\r\n"
);
post_socket
.write_all(post_response.as_bytes())
.await
.expect("writes post response");
let (mut poll_socket, _) = listener.accept().await.expect("accepts poll request");
let poll_request = read_http_headers(&mut poll_socket).await;
let response_body = r#"{"status":"succeeded","analyzeResult":{"pages":[{"pageNumber":1,"lines":[{"content":"ok"}]}]}}"#;
let poll_response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
poll_socket
.write_all(poll_response.as_bytes())
.await
.expect("writes poll response");
(post_request, poll_request)
});
let response = ocr(OcrRequest {
model: "prebuilt-read",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("di-key"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: "azure_ai/doc-intelligence",
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
})
.await
.expect("document intelligence request succeeds");
assert_eq!(response["pages"][0]["markdown"], "ok");
let (post_request, poll_request) = server.await.expect("server task completes");
assert!(
post_request
.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: di-key"),
"{post_request}"
);
assert!(
poll_request
.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: di-key"),
"{poll_request}"
);
}
#[test]
fn string_headers_rejects_non_string_values() {
let headers = json!({
"x-retry-count": 3
})
.as_object()
.unwrap()
.clone();
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
assert_eq!(
err,
CoreError::InvalidRequest(
"OCR extra_headers.x-retry-count must be a string, got number".to_string()
)
);
}
async fn respond_once(status_line: &'static str, body: String) -> String {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let _ = read_http_headers(&mut socket).await;
let response = format!(
"HTTP/1.1 {status_line}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
});
format!("http://{addr}")
}
async fn run_mistral_ocr(api_base: String, timeout: Duration) -> CoreResult<Value> {
ocr(OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: Some(&api_base),
custom_llm_provider: "mistral",
extra_headers: None,
optional_params: Map::new(),
timeout: Some(timeout),
})
.await
}
#[tokio::test]
async fn ocr_preserves_upstream_error_status() {
for status in [
"401 Unauthorized",
"404 Not Found",
"500 Internal Server Error",
] {
let base = respond_once(status, r#"{"error":"nope"}"#.to_string()).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("upstream error surfaces");
let expected = status[..3].parse::<u16>().expect("status prefix parses");
match err {
CoreError::Http { status: got, .. } => assert_eq!(got, expected),
other => panic!("expected Http error, got {other:?}"),
}
assert_eq!(err.public_status_code(), Some(expected));
}
}
#[tokio::test]
async fn ocr_bounds_oversized_error_body() {
let base = respond_once("500 Internal Server Error", "x".repeat(6000)).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("oversized error surfaces");
match err {
CoreError::Http { body, .. } => {
assert!(body.ends_with("... (truncated)"));
assert!(
body.chars().count() < 300,
"body not bounded: {} chars",
body.chars().count()
);
}
other => panic!("expected Http error, got {other:?}"),
}
}
#[tokio::test]
async fn ocr_rejects_invalid_json_success_body() {
let base = respond_once("200 OK", "not json".to_string()).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("invalid JSON surfaces");
assert!(matches!(err, CoreError::InvalidResponse(_)));
assert_eq!(err.public_status_code(), Some(500));
}
#[tokio::test]
async fn ocr_rejects_empty_success_body() {
let base = respond_once("200 OK", " ".to_string()).await;
let err = run_mistral_ocr(base, Duration::from_secs(5))
.await
.expect_err("empty success surfaces");
match err {
CoreError::InvalidResponse(message) => assert!(message.contains("empty")),
other => panic!("expected InvalidResponse, got {other:?}"),
}
}
#[tokio::test]
async fn ocr_classifies_per_request_timeout() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let _ = read_http_headers(&mut socket).await;
tokio::time::sleep(Duration::from_secs(3)).await;
let _ = socket.write_all(b"HTTP/1.1 200 OK\r\n\r\n").await;
});
let err = run_mistral_ocr(format!("http://{addr}"), Duration::from_millis(200))
.await
.expect_err("timeout surfaces");
assert_eq!(err, CoreError::Timeout);
assert_eq!(err.public_status_code(), Some(408));
}
#[tokio::test]
async fn ocr_maps_unregistered_provider_to_invalid_provider() {
let err = ocr(OcrRequest {
model: "some-model",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: "definitely-not-a-provider",
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
})
.await
.expect_err("unregistered provider surfaces");
assert!(matches!(err, CoreError::InvalidProvider(_)));
assert_eq!(err.public_status_code(), Some(400));
assert_eq!(err.public_message(), "Invalid OCR request");
}

View file

@ -13,6 +13,9 @@
pub mod io;
mod config;
mod constants;
/// Shared reqwest-failure classification into typed [`litellm_core::error::CoreError`]
/// contracts. Always available — every I/O endpoint maps transport failures here.
mod errors;
@ -32,8 +35,6 @@ pub mod state;
// `server`-gated; `io::realtime` exposes the generic `observe` hook while the
// collector and callback fan-out live here.
#[cfg(feature = "server")]
mod constants;
#[cfg(feature = "server")]
pub mod integrations;
#[cfg(feature = "server")]
mod realtime;

View file

@ -4,6 +4,8 @@ use crate::CoreResult;
use super::types::{OcrRequestData, OcrResponseData};
pub const OCR_PUBLIC_PARAMS_RESERVED_BY_LITELLM: &[&str] = &["id"];
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OcrAuthStrategy {
Bearer,
@ -36,13 +38,15 @@ pub trait OcrProviderConfig: Sync {
fn supported_ocr_params(&self) -> &'static [&'static str];
fn map_ocr_params(&self, non_default_params: &Map<String, Value>) -> Map<String, Value> {
let mut mapped_params = Map::new();
for (param, value) in non_default_params {
if self.supported_ocr_params().contains(&param.as_str()) {
mapped_params.insert(param.clone(), value.clone());
}
}
mapped_params
non_default_params
.iter()
.filter(|(param, value)| {
self.supported_ocr_params().contains(&param.as_str())
&& !(value.is_null()
&& OCR_PUBLIC_PARAMS_RESERVED_BY_LITELLM.contains(&param.as_str()))
})
.map(|(param, value)| (param.clone(), value.clone()))
.collect()
}
fn transform_ocr_request(

View file

@ -211,6 +211,137 @@ mod tests {
assert!(!mapped.contains_key("unsupported_param"));
}
#[test]
fn map_ocr_params_forwards_id_and_drops_litellm_internal_params() {
let params = json!({
"id": "ocr-req-9",
"pages": [0, 1],
"metadata": {"trace": "internal"},
"litellm_metadata": {"trace": "internal"},
"num_retries": 3,
"original_generic_function": "canary"
});
let mapped = map_ocr_params(params.as_object().unwrap());
assert_eq!(mapped.get("id"), Some(&json!("ocr-req-9")));
assert_eq!(mapped.get("pages"), Some(&json!([0, 1])));
for internal in [
"metadata",
"litellm_metadata",
"num_retries",
"original_generic_function",
] {
assert!(!mapped.contains_key(internal), "{internal} must be dropped");
}
}
#[test]
fn map_ocr_params_omits_id_when_absent() {
let params = json!({
"pages": [0, 1],
"include_image_base64": true
});
let mapped = map_ocr_params(params.as_object().unwrap());
assert!(!mapped.contains_key("id"));
assert_eq!(mapped.get("pages"), Some(&json!([0, 1])));
}
#[test]
fn map_ocr_params_omits_id_when_null() {
let params = json!({
"id": null,
"pages": [0, 1],
"include_image_base64": true
});
let mapped = map_ocr_params(params.as_object().unwrap());
assert!(
!mapped.contains_key("id"),
"null id must be dropped for the standalone Axum path"
);
assert_eq!(mapped.get("pages"), Some(&json!([0, 1])));
assert_eq!(mapped.get("include_image_base64"), Some(&json!(true)));
}
#[test]
fn transform_ocr_request_omits_null_id_from_serialized_body() {
let document = json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
});
let supplied = json!({
"id": null,
"pages": [0],
"include_image_base64": true
});
let filtered = map_ocr_params(supplied.as_object().unwrap());
let result = transform_ocr_request("mistral-ocr-latest", document.clone(), filtered)
.expect("request should transform");
assert_eq!(
result.data,
json!({
"model": "mistral-ocr-latest",
"document": document,
"pages": [0],
"include_image_base64": true
})
);
assert!(result.data.get("id").is_none());
}
#[test]
fn transform_ocr_request_serializes_full_supported_contract() {
let document = json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
});
let supplied = json!({
"pages": [0, 2, 5],
"include_image_base64": true,
"image_limit": 10,
"image_min_size": 64,
"bbox_annotation_format": {"type": "text"},
"document_annotation_format": {"type": "json_schema"},
"document_annotation_prompt": "extract title",
"extract_header": true,
"extract_footer": false,
"table_format": "html",
"confidence_scores_granularity": "word",
"include_blocks": true,
"id": "ocr-req-9",
"litellm_metadata": {"trace": "internal"},
"num_retries": 3
});
let filtered = map_ocr_params(supplied.as_object().unwrap());
let result = transform_ocr_request("mistral-ocr-latest", document.clone(), filtered)
.expect("request should transform");
assert_eq!(
result.data,
json!({
"model": "mistral-ocr-latest",
"document": document,
"pages": [0, 2, 5],
"include_image_base64": true,
"image_limit": 10,
"image_min_size": 64,
"bbox_annotation_format": {"type": "text"},
"document_annotation_format": {"type": "json_schema"},
"document_annotation_prompt": "extract title",
"extract_header": true,
"extract_footer": false,
"table_format": "html",
"confidence_scores_granularity": "word",
"include_blocks": true,
"id": "ocr-req-9"
})
);
}
#[test]
fn transform_ocr_request_builds_mistral_body() {
let document = json!({

View file

@ -201,9 +201,17 @@ def _resolve_ocr_call_context(
verbose_logger.debug(f"OCR call - model: {model}, provider: {custom_llm_provider}")
forwarded_kwargs = {
**filter_out_litellm_params(kwargs=kwargs),
**{
key: kwargs[key]
for key in _OCR_PUBLIC_PARAMS_RESERVED_BY_LITELLM
if kwargs.get(key) is not None
},
}
optional_params = {
key: value
for key, value in filter_out_litellm_params(kwargs=kwargs).items()
for key, value in forwarded_kwargs.items()
if key not in _RUST_BRIDGE_INTERNAL_PARAMS
}
@ -475,6 +483,8 @@ _MIME_PATTERN = re.compile(r"^[\w.+-]+/[\w.+-]+$")
_RUST_BRIDGE_INTERNAL_PARAMS = {"original_generic_function"}
_OCR_PUBLIC_PARAMS_RESERVED_BY_LITELLM: frozenset[str] = frozenset({"id"})
_MIME_TYPE_MAP = {
".pdf": "application/pdf",
".png": "image/png",

View file

@ -0,0 +1,217 @@
from __future__ import annotations
import json
import queue
import re
import socket
import subprocess
import sys
import threading
import time
from contextlib import closing
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, HTTPServer
from pathlib import Path
from typing import Iterator, TextIO, cast
import httpx
import pytest
import yaml
from pydantic import BaseModel, ConfigDict
from litellm.rust_bridge import native_bridge_available
REPO_ROOT = Path(__file__).resolve().parents[3]
_SECRET_PATTERN = re.compile(r"(sk-[A-Za-z0-9_\-]+|Bearer\s+\S+)")
_CAPTURE_RESPONSE_BODY: bytes = json.dumps(
{
"pages": [{"index": 0, "markdown": "captured"}],
"model": "mistral-ocr-latest",
"usage_info": {"pages_processed": 1},
}
).encode()
_LIVENESS_DEADLINE_SECONDS = 90.0
_LIVENESS_POLL_SECONDS = 0.5
_PROXY_TERMINATE_TIMEOUT_SECONDS = 15
_SERVER_JOIN_TIMEOUT_SECONDS = 5
class CaptureLitellmParams(BaseModel):
model_config = ConfigDict(frozen=True)
model: str
api_key: str
api_base: str
class CaptureModelEntry(BaseModel):
model_config = ConfigDict(frozen=True)
model_name: str
litellm_params: CaptureLitellmParams
class CaptureGeneralSettings(BaseModel):
model_config = ConfigDict(frozen=True)
master_key: str
class CaptureLitellmSettings(BaseModel):
model_config = ConfigDict(frozen=True)
drop_params: bool
class CaptureProxyConfig(BaseModel):
model_config = ConfigDict(frozen=True)
model_list: tuple[CaptureModelEntry, ...]
general_settings: CaptureGeneralSettings
litellm_settings: CaptureLitellmSettings
@dataclass(frozen=True)
class CaptureProxy:
proxy_url: str
master_key: str
captures: queue.Queue[bytes]
def _sanitize(text: str) -> str:
return _SECRET_PATTERN.sub("[redacted]", text)
def _make_capture_handler(
captures: queue.Queue[bytes],
) -> type[BaseHTTPRequestHandler]:
class _CaptureHandler(BaseHTTPRequestHandler):
def do_POST(self) -> None:
length = int(self.headers.get("content-length", "0"))
captures.put(self.rfile.read(length))
self.send_response(200)
self.send_header("content-type", "application/json")
self.send_header("content-length", str(len(_CAPTURE_RESPONSE_BODY)))
self.end_headers()
self.wfile.write(_CAPTURE_RESPONSE_BODY)
def log_message(self, format: str, *args: object) -> None:
return
return _CaptureHandler
def _free_port() -> int:
with closing(socket.socket(socket.AF_INET, socket.SOCK_STREAM)) as sock:
sock.bind(("127.0.0.1", 0))
_, port = cast(tuple[str, int], sock.getsockname())
return port
def _wait_for_liveness(base_url: str, deadline: float) -> bool:
while time.monotonic() < deadline:
try:
resp = httpx.get(f"{base_url}/health/liveliness", timeout=2)
if resp.status_code == 200:
return True
except httpx.HTTPError:
pass
time.sleep(_LIVENESS_POLL_SECONDS)
return False
def _capture_config(capture_port: int, master_key: str) -> CaptureProxyConfig:
return CaptureProxyConfig(
model_list=(
CaptureModelEntry(
model_name="rust-ocr-mistral-capture",
litellm_params=CaptureLitellmParams(
model="mistral/mistral-ocr-latest",
api_key="sk-capture-test",
api_base=f"http://127.0.0.1:{capture_port}",
),
),
),
general_settings=CaptureGeneralSettings(master_key=master_key),
litellm_settings=CaptureLitellmSettings(drop_params=False),
)
@pytest.fixture
def capture_proxy(tmp_path: Path) -> Iterator[CaptureProxy]:
if not native_bridge_available():
pytest.skip("compiled Rust OCR bridge is required for the capture E2E")
master_key = "sk-1234"
captures: queue.Queue[bytes] = queue.Queue()
capture_server: HTTPServer | None = None
server_thread: threading.Thread | None = None
proxy: subprocess.Popen[bytes] | None = None
proxy_log: TextIO | None = None
proxy_log_path = tmp_path / "capture-proxy.log"
try:
capture_port = _free_port()
capture_server = HTTPServer(
("127.0.0.1", capture_port), _make_capture_handler(captures)
)
server_thread = threading.Thread(
target=capture_server.serve_forever, daemon=True
)
server_thread.start()
proxy_port = _free_port()
config_path = tmp_path / "capture-config.yml"
config_path.write_text(
yaml.safe_dump(_capture_config(capture_port, master_key).model_dump())
)
proxy_log = proxy_log_path.open("w")
proxy = subprocess.Popen(
[
sys.executable,
str(REPO_ROOT / "litellm" / "proxy" / "proxy_cli.py"),
"--config",
str(config_path),
"--host",
"127.0.0.1",
"--port",
str(proxy_port),
"--num_workers",
"1",
],
cwd=str(REPO_ROOT),
stdout=proxy_log,
stderr=subprocess.STDOUT,
)
proxy_url = f"http://127.0.0.1:{proxy_port}"
if not _wait_for_liveness(
proxy_url, time.monotonic() + _LIVENESS_DEADLINE_SECONDS
):
proxy_log.flush()
tail = _sanitize(proxy_log_path.read_text()[-4000:])
pytest.fail(
f"capture proxy did not become live while the Rust bridge is available; "
f"sanitized proxy log at {proxy_log_path}\n{tail}"
)
yield CaptureProxy(
proxy_url=proxy_url, master_key=master_key, captures=captures
)
finally:
if proxy is not None:
proxy.terminate()
try:
proxy.wait(timeout=_PROXY_TERMINATE_TIMEOUT_SECONDS)
except subprocess.TimeoutExpired:
proxy.kill()
proxy.wait()
if proxy_log is not None:
proxy_log.close()
if server_thread is not None and server_thread.ident is not None:
if capture_server is not None:
capture_server.shutdown()
server_thread.join(timeout=_SERVER_JOIN_TIMEOUT_SECONDS)
if capture_server is not None:
capture_server.server_close()

View file

@ -8,14 +8,155 @@ litellm --config tests/e2e/gateway/litellm-config.yml --port 4000
from __future__ import annotations
import json
import os
import time
import uuid
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from typing import Literal, cast
import httpx
import pytest
import yaml
from pydantic import BaseModel, ConfigDict, Field
from ocr_capture_proxy import CaptureProxy, capture_proxy
__all__ = ["capture_proxy"]
class OcrDocument(BaseModel):
model_config = ConfigDict(frozen=True)
type: str
document_url: str | None = None
image_url: str | None = None
class JsonSchemaProperty(BaseModel):
model_config = ConfigDict(frozen=True)
type: str
class AnnotationJsonSchemaBody(BaseModel):
model_config = ConfigDict(frozen=True, populate_by_name=True)
type: str
properties: dict[str, JsonSchemaProperty]
required: tuple[str, ...]
additional_properties: bool = Field(alias="additionalProperties")
class AnnotationJsonSchema(BaseModel):
model_config = ConfigDict(frozen=True, populate_by_name=True)
name: str
body: AnnotationJsonSchemaBody = Field(alias="schema")
strict: bool
class MistralAnnotationFormat(BaseModel):
model_config = ConfigDict(frozen=True)
type: str
json_schema: AnnotationJsonSchema
class MistralOcrParams(BaseModel):
model_config = ConfigDict(frozen=True)
pages: tuple[int, ...]
include_image_base64: bool
include_blocks: bool
image_limit: int
image_min_size: int
bbox_annotation_format: MistralAnnotationFormat
document_annotation_format: MistralAnnotationFormat
document_annotation_prompt: str
extract_header: bool
extract_footer: bool
table_format: str
confidence_scores_granularity: str
id: str
class TraceMetadata(BaseModel):
model_config = ConfigDict(frozen=True)
trace: str
class LitellmInternalCanaries(BaseModel):
model_config = ConfigDict(frozen=True)
metadata: TraceMetadata
litellm_metadata: TraceMetadata
num_retries: int
tags: tuple[str, ...]
litellm_session_id: str
original_generic_function: str
class MistralOcrUpstreamRequest(MistralOcrParams):
model_config = ConfigDict(frozen=True, extra="forbid")
model: str
document: OcrDocument
class OcrResponsePage(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
index: int
markdown: str
class OcrUsageInfo(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
pages_processed: int
class OcrResponseEnvelope(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
object: Literal["ocr"]
model: str = Field(min_length=1)
pages: tuple[OcrResponsePage, ...] = Field(min_length=1)
usage_info: OcrUsageInfo | None = None
class ModelInfoDetail(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
id: str | None = None
class ModelInfoEntry(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
model_name: str
model_info: ModelInfoDetail = Field(default_factory=ModelInfoDetail)
class ModelInfoResponse(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
data: tuple[ModelInfoEntry, ...]
class GatewayConfigEntry(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
model_name: str
class GatewayConfig(BaseModel):
model_config = ConfigDict(frozen=True, extra="allow")
model_list: tuple[GatewayConfigEntry, ...]
TEST_PDF_URL = (
"https://cdn.jsdelivr.net/gh/BerriAI/litellm"
@ -28,33 +169,102 @@ TEST_IMAGE_URL = (
"/tests/image_gen_tests/test_image.png"
)
CAPTURE_DOCUMENT = OcrDocument(type="document_url", document_url=TEST_PDF_URL)
SUPPORTED_PARAMS = MistralOcrParams(
pages=(0,),
include_image_base64=True,
include_blocks=True,
image_limit=10,
image_min_size=64,
bbox_annotation_format=MistralAnnotationFormat(
type="json_schema",
json_schema=AnnotationJsonSchema(
name="bbox_annotation",
schema=AnnotationJsonSchemaBody(
type="object",
properties={"description": JsonSchemaProperty(type="string")},
required=("description",),
additionalProperties=False,
),
strict=True,
),
),
document_annotation_format=MistralAnnotationFormat(
type="json_schema",
json_schema=AnnotationJsonSchema(
name="document_annotation",
schema=AnnotationJsonSchemaBody(
type="object",
properties={"title": JsonSchemaProperty(type="string")},
required=("title",),
additionalProperties=False,
),
strict=True,
),
),
document_annotation_prompt="extract the title",
extract_header=True,
extract_footer=False,
table_format="markdown",
confidence_scores_granularity="word",
id="ocr-req-parity-9",
)
INTERNAL_CANARIES = LitellmInternalCanaries(
metadata=TraceMetadata(trace="internal"),
litellm_metadata=TraceMetadata(trace="internal"),
num_retries=3,
tags=("internal",),
litellm_session_id="sess-internal",
original_generic_function="litellm-internal-should-be-filtered",
)
EXPECTED_UPSTREAM = MistralOcrUpstreamRequest(
model="mistral-ocr-latest",
document=CAPTURE_DOCUMENT,
pages=SUPPORTED_PARAMS.pages,
include_image_base64=SUPPORTED_PARAMS.include_image_base64,
include_blocks=SUPPORTED_PARAMS.include_blocks,
image_limit=SUPPORTED_PARAMS.image_limit,
image_min_size=SUPPORTED_PARAMS.image_min_size,
bbox_annotation_format=SUPPORTED_PARAMS.bbox_annotation_format,
document_annotation_format=SUPPORTED_PARAMS.document_annotation_format,
document_annotation_prompt=SUPPORTED_PARAMS.document_annotation_prompt,
extract_header=SUPPORTED_PARAMS.extract_header,
extract_footer=SUPPORTED_PARAMS.extract_footer,
table_format=SUPPORTED_PARAMS.table_format,
confidence_scores_granularity=SUPPORTED_PARAMS.confidence_scores_granularity,
id=SUPPORTED_PARAMS.id,
)
RUST_OCR_GATEWAY_CASES = [
pytest.param(
"rust-ocr-mistral",
{"type": "document_url", "document_url": TEST_PDF_URL},
OcrDocument(type="document_url", document_url=TEST_PDF_URL),
id="mistral",
),
pytest.param(
"rust-ocr-azure-ai",
{"type": "document_url", "document_url": TEST_PDF_URL},
OcrDocument(type="document_url", document_url=TEST_PDF_URL),
id="azure_ai",
),
pytest.param(
"rust-ocr-azure-document-intelligence",
{"type": "document_url", "document_url": TEST_PDF_URL},
OcrDocument(type="document_url", document_url=TEST_PDF_URL),
id="azure_document_intelligence",
),
pytest.param(
"rust-ocr-vertex-mistral",
{"type": "document_url", "document_url": TEST_PDF_URL},
OcrDocument(type="document_url", document_url=TEST_PDF_URL),
id="vertex_mistral",
),
pytest.param(
"rust-ocr-vertex-deepseek",
{
"type": "image_url",
"image_url": os.getenv("RUST_OCR_IMAGE_URL", TEST_IMAGE_URL),
},
OcrDocument(
type="image_url",
image_url=os.getenv("RUST_OCR_IMAGE_URL", TEST_IMAGE_URL),
),
id="vertex_deepseek",
),
]
@ -62,36 +272,113 @@ RUST_OCR_GATEWAY_CASES = [
CONFIG_PATH = Path(__file__).with_name("litellm-config.yml")
def _wire_payload(
model: str,
document: OcrDocument,
params: MistralOcrParams | None,
canaries: LitellmInternalCanaries | None = None,
) -> str:
return json.dumps(
{
"model": model,
"document": document.model_dump(mode="json", exclude_none=True),
**(
params.model_dump(mode="json", by_alias=True)
if params is not None
else {}
),
**(
canaries.model_dump(mode="json", by_alias=True)
if canaries is not None
else {}
),
}
)
@dataclass(frozen=True)
class OcrGateway:
base_url: str
master_key: str
def model_names(self) -> set[str]:
with httpx.Client(
timeout=float(os.getenv("E2E_REQUEST_TIMEOUT", "120"))
) as client:
def _client(self) -> httpx.Client:
return httpx.Client(timeout=float(os.getenv("E2E_REQUEST_TIMEOUT", "120")))
def model_names(self) -> frozenset[str]:
with self._client() as client:
response = client.get(
f"{self.base_url.rstrip('/')}/model/info",
headers={"Authorization": f"Bearer {self.master_key}"},
)
assert response.status_code == 200, response.text
return {
model["model_name"]
for model in response.json().get("data", [])
if "model_name" in model
}
parsed = ModelInfoResponse.model_validate_json(response.content)
return frozenset(entry.model_name for entry in parsed.data)
def ocr(self, model: str, document: dict[str, str]) -> httpx.Response:
with httpx.Client(
timeout=float(os.getenv("E2E_REQUEST_TIMEOUT", "120"))
) as client:
def ocr(
self,
model: str,
document: OcrDocument,
params: MistralOcrParams | None = None,
) -> httpx.Response:
with self._client() as client:
return client.post(
f"{self.base_url.rstrip('/')}/v1/ocr",
headers={"Authorization": f"Bearer {self.master_key}"},
json={"model": model, "document": document},
headers={
"Authorization": f"Bearer {self.master_key}",
"content-type": "application/json",
},
content=_wire_payload(model, document, params),
)
def create_model(
self, model_name: str, litellm_params: dict[str, str]
) -> httpx.Response:
with self._client() as client:
return client.post(
f"{self.base_url.rstrip('/')}/model/new",
headers={"Authorization": f"Bearer {self.master_key}"},
json={"model_name": model_name, "litellm_params": litellm_params},
)
def delete_model(self, model_id: str) -> httpx.Response:
with self._client() as client:
return client.post(
f"{self.base_url.rstrip('/')}/model/delete",
headers={"Authorization": f"Bearer {self.master_key}"},
json={"id": model_id},
)
def model_id(self, model_name: str) -> str | None:
with self._client() as client:
response = client.get(
f"{self.base_url.rstrip('/')}/model/info",
headers={"Authorization": f"Bearer {self.master_key}"},
)
assert response.status_code == 200, response.text
parsed = ModelInfoResponse.model_validate_json(response.content)
for entry in parsed.data:
if entry.model_name == model_name:
return entry.model_info.id
return None
def wait_for_model(self, model_name: str, attempts: int = 20) -> None:
for _ in range(attempts):
if model_name in self.model_names():
return
time.sleep(1)
raise AssertionError(
f"{model_name} did not appear on /model/info within {attempts}s"
)
def wait_for_model_absent(self, model_name: str, attempts: int = 20) -> None:
for _ in range(attempts):
if model_name not in self.model_names():
return
time.sleep(1)
raise AssertionError(
f"{model_name} still present on /model/info after {attempts}s"
)
@dataclass(frozen=True)
class OcrResources:
@ -113,36 +400,107 @@ def resources() -> OcrResources:
)
def _assert_ocr_response_shape(response_json: dict[str, Any]) -> None:
assert response_json["object"] == "ocr"
assert response_json["model"]
assert isinstance(response_json["pages"], list)
assert len(response_json["pages"]) > 0
assert "index" in response_json["pages"][0]
assert "markdown" in response_json["pages"][0]
class TestRustOcrGateway:
def test_rust_ocr_models_are_on_gateway_config(self) -> None:
config = yaml.safe_load(CONFIG_PATH.read_text())
configured_models = {
model_config["model_name"] for model_config in config["model_list"]
}
config = GatewayConfig.model_validate(
cast(object, yaml.safe_load(CONFIG_PATH.read_text()))
)
configured_models = frozenset(entry.model_name for entry in config.model_list)
expected_models = {case.values[0] for case in RUST_OCR_GATEWAY_CASES}
expected_models = {str(case.values[0]) for case in RUST_OCR_GATEWAY_CASES}
assert expected_models.issubset(configured_models)
def test_running_gateway_loaded_rust_ocr_models(
self, resources: OcrResources
) -> None:
expected_models = {case.values[0] for case in RUST_OCR_GATEWAY_CASES}
expected_models = {str(case.values[0]) for case in RUST_OCR_GATEWAY_CASES}
assert expected_models.issubset(resources.gateway.model_names())
@pytest.mark.parametrize(("model", "document"), RUST_OCR_GATEWAY_CASES)
def test_rust_ocr_model_gateway_response(
self, resources: OcrResources, model: str, document: dict[str, str]
self, resources: OcrResources, model: str, document: OcrDocument
) -> None:
response = resources.gateway.ocr(model, document)
assert response.status_code == 200, response.text
_assert_ocr_response_shape(response.json())
OcrResponseEnvelope.model_validate_json(response.content)
@pytest.mark.e2e
def test_rust_ocr_mistral_live_forwards_supported_params(
self, resources: OcrResources
) -> None:
if not os.getenv("MISTRAL_API_KEY"):
pytest.skip("MISTRAL_API_KEY not set for live Mistral OCR call")
response = resources.gateway.ocr(
"rust-ocr-mistral", CAPTURE_DOCUMENT, SUPPORTED_PARAMS
)
assert response.status_code == 200, response.text
parsed = OcrResponseEnvelope.model_validate_json(response.content)
assert parsed.pages[0].markdown != ""
if parsed.usage_info is not None:
assert parsed.usage_info.pages_processed >= 1
def test_rust_ocr_proxy_forwards_full_contract_to_capture_endpoint(
capture_proxy: CaptureProxy,
) -> None:
response = httpx.post(
f"{capture_proxy.proxy_url}/v1/ocr",
headers={
"Authorization": f"Bearer {capture_proxy.master_key}",
"content-type": "application/json",
},
content=_wire_payload(
"rust-ocr-mistral-capture",
CAPTURE_DOCUMENT,
SUPPORTED_PARAMS,
INTERNAL_CANARIES,
),
timeout=60,
)
assert response.status_code == 200, response.text
OcrResponseEnvelope.model_validate_json(response.content)
captured = MistralOcrUpstreamRequest.model_validate_json(
capture_proxy.captures.get(timeout=10)
)
assert captured == EXPECTED_UPSTREAM
class TestRustOcrDynamicDeployment:
def test_os_environ_api_key_deployment_lifecycle(
self, resources: OcrResources
) -> None:
if not os.getenv("MISTRAL_API_KEY"):
pytest.skip("Set MISTRAL_API_KEY on the proxy for the live OCR lifecycle")
gateway = resources.gateway
model_name = f"rust-ocr-env-e2e-{uuid.uuid4().hex[:8]}"
create = gateway.create_model(
model_name=model_name,
litellm_params={
"model": "mistral/mistral-ocr-latest",
"api_key": "os.environ/MISTRAL_API_KEY",
},
)
assert create.status_code == 200, create.text
try:
gateway.wait_for_model(model_name)
response = gateway.ocr(
model_name,
OcrDocument(type="document_url", document_url=TEST_PDF_URL),
)
assert response.status_code == 200, response.text
OcrResponseEnvelope.model_validate_json(response.content)
assert "os.environ/MISTRAL_API_KEY" not in response.text
finally:
deployed_id = gateway.model_id(model_name)
if deployed_id is not None:
delete = gateway.delete_model(deployed_id)
assert delete.status_code == 200, delete.text
gateway.wait_for_model_absent(model_name)

View file

@ -449,13 +449,71 @@ def test_ocr_filters_internal_litellm_params_before_rust(fake_bridge):
document=DOCUMENT,
api_key="sk-test",
include_image_base64=True,
original_generic_function=lambda: None,
original_generic_function="litellm-internal-should-be-filtered",
litellm_metadata={"trace": "internal"},
)
assert fake_bridge.calls[0]["optional_params"] == {"include_image_base64": True}
def test_ocr_forwards_public_id_but_drops_internal_litellm_params(fake_bridge):
litellm.ocr(
model=MODEL,
document=DOCUMENT,
api_key="sk-test",
id="ocr-req-9",
pages=[0, 1],
include_image_base64=True,
table_format="html",
metadata={"trace": "internal"},
litellm_metadata={"trace": "internal"},
litellm_session_id="sess-internal",
tags=["internal"],
num_retries=3,
original_generic_function="litellm-internal-should-be-filtered",
)
optional_params = fake_bridge.calls[0]["optional_params"]
assert optional_params["id"] == "ocr-req-9"
assert optional_params["pages"] == [0, 1]
assert optional_params["include_image_base64"] is True
assert optional_params["table_format"] == "html"
for internal in (
"metadata",
"litellm_metadata",
"litellm_session_id",
"tags",
"num_retries",
"original_generic_function",
):
assert internal not in optional_params
def test_ocr_omits_reserved_id_when_none(fake_bridge):
litellm.ocr(
model=MODEL,
document=DOCUMENT,
api_key="sk-test",
id=None,
include_image_base64=True,
)
optional_params = fake_bridge.calls[0]["optional_params"]
assert "id" not in optional_params
assert optional_params["include_image_base64"] is True
def test_ocr_omits_reserved_id_when_absent(fake_bridge):
litellm.ocr(
model=MODEL,
document=DOCUMENT,
api_key="sk-test",
include_image_base64=True,
)
assert "id" not in fake_bridge.calls[0]["optional_params"]
def test_ocr_routes_azure_ai_to_rust_by_default(fake_bridge):
response = litellm.ocr(
model="azure_ai/pixtral-12b-2409",
@ -801,3 +859,50 @@ def test_raise_ocr_exception_keeps_validation_error_off_bad_request(
)
assert spy.calls[0].original_exception is validation_info.value
def test_ocr_forwards_os_environ_api_key_reference_to_rust(
fake_bridge: RecordingBridge,
) -> None:
litellm.ocr(
model=MODEL, document=DOCUMENT, api_key="os.environ/MISTRAL_OCR_TEST_KEY"
)
assert fake_bridge.calls[0]["api_key"] == "os.environ/MISTRAL_OCR_TEST_KEY"
def test_ocr_forwards_provider_derived_os_environ_references_to_rust(
fake_bridge: RecordingBridge, monkeypatch: pytest.MonkeyPatch
) -> None:
def fake_get_llm_provider(
*,
model: str,
custom_llm_provider: str | None,
api_base: str | None,
api_key: str | None,
) -> tuple[str, str, str, str]:
return (
"mistral-ocr-latest",
"mistral",
"os.environ/MISTRAL_PROVIDER_KEY",
"os.environ/MISTRAL_PROVIDER_BASE",
)
monkeypatch.setattr(ocr_main.litellm, "get_llm_provider", fake_get_llm_provider)
litellm.ocr(model=MODEL, document=DOCUMENT)
call = fake_bridge.calls[0]
assert call["api_key"] == "os.environ/MISTRAL_PROVIDER_KEY"
assert call["api_base"] == "os.environ/MISTRAL_PROVIDER_BASE"
@pytest.mark.asyncio
async def test_aocr_forwards_os_environ_api_key_reference_to_rust(
fake_async_bridge: RecordingAsyncBridge,
) -> None:
await litellm.aocr(
model=MODEL, document=DOCUMENT, api_key="os.environ/MISTRAL_OCR_TEST_KEY"
)
assert fake_async_bridge.calls[0]["api_key"] == "os.environ/MISTRAL_OCR_TEST_KEY"