mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
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:
commit
ed5cb16589
11 changed files with 1560 additions and 447 deletions
85
litellm-rust/crates/ai-gateway/src/config.rs
Normal file
85
litellm-rust/crates/ai-gateway/src/config.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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/";
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
573
litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs
Normal file
573
litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs
Normal 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");
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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(¶m.as_str()) {
|
||||
mapped_params.insert(param.clone(), value.clone());
|
||||
}
|
||||
}
|
||||
mapped_params
|
||||
non_default_params
|
||||
.iter()
|
||||
.filter(|(param, value)| {
|
||||
self.supported_ocr_params().contains(¶m.as_str())
|
||||
&& !(value.is_null()
|
||||
&& OCR_PUBLIC_PARAMS_RESERVED_BY_LITELLM.contains(¶m.as_str()))
|
||||
})
|
||||
.map(|(param, value)| (param.clone(), value.clone()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn transform_ocr_request(
|
||||
|
|
|
|||
|
|
@ -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!({
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
217
tests/e2e/gateway/ocr_capture_proxy.py
Normal file
217
tests/e2e/gateway/ocr_capture_proxy.py
Normal 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()
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue