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] 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",