diff --git a/litellm-rust/crates/ai-gateway/src/config.rs b/litellm-rust/crates/ai-gateway/src/config.rs new file mode 100644 index 00000000000..a98af8fb794 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/config.rs @@ -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 + Sync), +) -> Option { + 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 { + 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); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 3116a4c9932..10689943bff 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -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/"; diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr.rs b/litellm-rust/crates/ai-gateway/src/io/ocr.rs index 5d01b08bdae..66e00895216 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr.rs @@ -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 { + 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 + Sync), +) -> CoreResult { 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 { 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 { } #[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 { - 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::().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; diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs b/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs new file mode 100644 index 00000000000..901e20275bc --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs @@ -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::().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 { + 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::().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"); +} diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index 120c7e66365..713915717f1 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -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; diff --git a/litellm-rust/crates/core/src/ocr/transformation.rs b/litellm-rust/crates/core/src/ocr/transformation.rs index 39a17b44ca0..87fb658a3bd 100644 --- a/litellm-rust/crates/core/src/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/ocr/transformation.rs @@ -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) -> Map { - 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( diff --git a/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs index 4e6be1d8a04..73375229059 100644 --- a/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs @@ -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!({ diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 7cb5a462de9..11bb0f09733 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -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", diff --git a/tests/e2e/gateway/ocr_capture_proxy.py b/tests/e2e/gateway/ocr_capture_proxy.py new file mode 100644 index 00000000000..585f26fefc6 --- /dev/null +++ b/tests/e2e/gateway/ocr_capture_proxy.py @@ -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() diff --git a/tests/e2e/gateway/test_ocr_rust_e2e.py b/tests/e2e/gateway/test_ocr_rust_e2e.py index dcf85898365..6f799ed80fc 100644 --- a/tests/e2e/gateway/test_ocr_rust_e2e.py +++ b/tests/e2e/gateway/test_ocr_rust_e2e.py @@ -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) diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 5a9a48ab81a..8288b0196ac 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -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"