From 2cf209f265e1c1efde25ed3d196bcfb054fbc070 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 16 Jul 2026 20:25:22 -0700 Subject: [PATCH 1/2] fix(mistral-ocr): forward the complete public parameter contract (#33600) * fix(mistral-ocr): forward the complete public parameter contract The Rust OCR path already lists id in the supported Mistral parameters, but litellm/ocr/main.py ran filter_out_litellm_params before handing optional_params to the bridge. That generic filter drops every key in all_litellm_params, which includes id, so the supported public Mistral OCR field id was silently stripped and never reached the provider while the rest of the contract went through. Restore id after the generic filter via an OCR-specific reserved set, keeping genuine internal params filtered. Add Rust core and ai-gateway contract tests that pin the exact serialized request body for the full supported contract and prove internal canaries are dropped, a Python regression test that id survives while internal params do not, and a deterministic capture E2E plus a live Mistral param test. Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * test(mistral-ocr): drop unnecessary comments from parity tests Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * fix(mistral-ocr): only forward reserved id when set; harden capture e2e Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * refactor(mistral-ocr): typed queue-injected capture harness for e2e Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * test(mistral-ocr): fully type E2E capture models and move gateway ocr tests Replace the remaining Any/coarse object wire types in the OCR E2E module with concrete frozen Pydantic models for the supported-param contract, internal canaries, upstream capture request, and normalized response, and build every request payload immutably per call. The capture fixture now cleans up on every failure path (terminate then always wait, kill then always wait, shutdown plus server_close, join the server thread) and requires an HTTP 200 liveness probe. The live Mistral test now exercises the complete public parameter contract with valid json_schema annotation formats. Move the ai-gateway OCR test module out of io/ocr.rs into a focused io/ocr/tests.rs. Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * fix(mistral-ocr): drop reserved id upstream when null and type the E2E harness Extend map_ocr_params so a null-valued reserved public id is omitted from the serialized request, matching the Python OCR boundary for the standalone Axum path where a caller can submit JSON id:null; add core map and serialized-body tests for both the absent and null id cases. Extract the deterministic capture harness into a typed ocr_capture_proxy helper module with an exception-safe fixture that cleans up on every failure path from capture-server creation onward (terminate then always wait, kill then always wait, shutdown plus server_close, join the server thread, close logs after the process exits) and requires an HTTP 200 liveness probe. Serialize wire payloads with by_alias so the annotation schema/additionalProperties keys match the provider contract. Replace the lambda internal canary in the Python bridge regression with a serializable value and cover the litellm_session_id and tags internal fields. Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm-rust/crates/ai-gateway/src/io/ocr.rs | 392 +------------- .../crates/ai-gateway/src/io/ocr/tests.rs | 511 ++++++++++++++++++ .../crates/core/src/ocr/transformation.rs | 18 +- .../providers/mistral/ocr/transformation.rs | 131 +++++ litellm/ocr/main.py | 12 +- tests/e2e/gateway/ocr_capture_proxy.py | 217 ++++++++ tests/e2e/gateway/test_ocr_rust_e2e.py | 346 ++++++++++-- tests/test_litellm/ocr/test_rust_bridge.py | 60 +- 8 files changed, 1246 insertions(+), 441 deletions(-) create mode 100644 litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs create mode 100644 tests/e2e/gateway/ocr_capture_proxy.py diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr.rs b/litellm-rust/crates/ai-gateway/src/io/ocr.rs index c103efbb941..d233babb748 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr.rs @@ -169,394 +169,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..e2e9c65ab35 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs @@ -0,0 +1,511 @@ +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_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/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..f5d77afd3be 100644 --- a/tests/e2e/gateway/test_ocr_rust_e2e.py +++ b/tests/e2e/gateway/test_ocr_rust_e2e.py @@ -8,14 +8,147 @@ litellm --config tests/e2e/gateway/litellm-config.yml --port 4000 from __future__ import annotations +import json import os 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 ModelInfoEntry(BaseModel): + model_config = ConfigDict(frozen=True, extra="allow") + + model_name: str + + +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 +161,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,34 +264,62 @@ 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), ) @@ -113,36 +343,70 @@ 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 diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 5a9a48ab81a..c0897a642a6 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", From 4da89566eb1e9e55cd7e1a558095aa03ce526cb2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 16 Jul 2026 21:05:52 -0700 Subject: [PATCH 2/2] feat(ocr): resolve environment references in Rust (#33598) Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm-rust/crates/ai-gateway/src/config.rs | 85 ++++++++++++++++ .../crates/ai-gateway/src/constants.rs | 8 ++ litellm-rust/crates/ai-gateway/src/io/ocr.rs | 19 +++- .../crates/ai-gateway/src/io/ocr/tests.rs | 62 ++++++++++++ litellm-rust/crates/ai-gateway/src/lib.rs | 5 +- tests/e2e/gateway/test_ocr_rust_e2e.py | 96 ++++++++++++++++++- tests/test_litellm/ocr/test_rust_bridge.py | 47 +++++++++ 7 files changed, 315 insertions(+), 7 deletions(-) create mode 100644 litellm-rust/crates/ai-gateway/src/config.rs 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 d233babb748..0d42125bdbd 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 common_utils::{ @@ -73,6 +75,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!( @@ -80,18 +90,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()); diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs b/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs index e2e9c65ab35..901e20275bc 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/tests.rs @@ -198,6 +198,68 @@ async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() { ); } +#[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") diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index 6c04fbb7626..636a6d9252b 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; + /// GIL-activity tracking. Pure (atomics only); shared by the `server` routes and /// the `python-config` reader, so it is available without either feature. pub mod gil; @@ -28,8 +31,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/tests/e2e/gateway/test_ocr_rust_e2e.py b/tests/e2e/gateway/test_ocr_rust_e2e.py index f5d77afd3be..6f799ed80fc 100644 --- a/tests/e2e/gateway/test_ocr_rust_e2e.py +++ b/tests/e2e/gateway/test_ocr_rust_e2e.py @@ -10,6 +10,8 @@ from __future__ import annotations import json import os +import time +import uuid from dataclasses import dataclass from pathlib import Path from typing import Literal, cast @@ -126,10 +128,17 @@ class OcrResponseEnvelope(BaseModel): 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): @@ -149,7 +158,6 @@ class GatewayConfig(BaseModel): model_list: tuple[GatewayConfigEntry, ...] - TEST_PDF_URL = ( "https://cdn.jsdelivr.net/gh/BerriAI/litellm" "@d769e81c90d453240c61fc572cdb27fae06a89d0" @@ -322,6 +330,55 @@ class OcrGateway: 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: @@ -410,3 +467,40 @@ def test_rust_ocr_proxy_forwards_full_contract_to_capture_endpoint( 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 c0897a642a6..8288b0196ac 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -859,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"