diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr.rs b/litellm-rust/crates/ai-gateway/src/io/ocr.rs index 542880455a0..c103efbb941 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr.rs @@ -17,8 +17,8 @@ use serde_json::{Map, Value}; mod common_utils; use common_utils::{ - convert_document_url_to_data_uri, has_header, ocr_provider_config, poll_document_intelligence, - string_headers, truncate_error_body, upload_reducto_document, + classify_reqwest_error, convert_document_url_to_data_uri, has_header, ocr_provider_config, + poll_document_intelligence, string_headers, truncate_error_body, upload_reducto_document, }; /// OCR over large documents can take a while; bound it generously rather than @@ -74,8 +74,12 @@ 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 model = request.model; - let config = ocr_provider_config(request.custom_llm_provider, model) - .ok_or_else(|| CoreError::InvalidProvider(request.custom_llm_provider.to_string()))?; + let config = ocr_provider_config(request.custom_llm_provider, model).ok_or_else(|| { + CoreError::InvalidProvider(format!( + "no OCR provider '{}' registered for model '{model}'", + request.custom_llm_provider + )) + })?; let env_lookup = |key: &str| std::env::var(key).ok(); let headers = string_headers(request.extra_headers)?; @@ -116,7 +120,7 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult { let response = request_builder .send() .await - .map_err(|err| CoreError::Network(err.to_string()))?; + .map_err(classify_reqwest_error)?; let status = response.status(); if config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll @@ -141,10 +145,7 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult { .into_json()); } - let text = response - .text() - .await - .map_err(|err| CoreError::Network(err.to_string()))?; + let text = response.text().await.map_err(classify_reqwest_error)?; if !status.is_success() { return Err(CoreError::Http { @@ -153,6 +154,12 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult { }); } + if text.trim().is_empty() { + return Err(CoreError::InvalidResponse( + "OCR provider returned an empty success response".to_string(), + )); + } + let response_json: Value = serde_json::from_str(&text) .map_err(|err| CoreError::InvalidResponse(format!("invalid OCR response JSON: {err}")))?; @@ -412,4 +419,144 @@ mod tests { ) ); } + + async fn respond_once(status_line: &'static str, body: String) -> String { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("test listener binds"); + let addr = listener.local_addr().expect("listener has local addr"); + tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts one request"); + let _ = read_http_headers(&mut socket).await; + let response = format!( + "HTTP/1.1 {status_line}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + body.len(), + body + ); + socket + .write_all(response.as_bytes()) + .await + .expect("writes response"); + }); + format!("http://{addr}") + } + + async fn run_mistral_ocr(api_base: String, timeout: Duration) -> CoreResult { + ocr(OcrRequest { + model: "mistral-ocr-latest", + document: json!({ + "type": "document_url", + "document_url": "https://example.com/doc.pdf" + }), + api_key: Some("sk-test"), + api_base: Some(&api_base), + custom_llm_provider: "mistral", + extra_headers: None, + optional_params: Map::new(), + timeout: Some(timeout), + }) + .await + } + + #[tokio::test] + async fn ocr_preserves_upstream_error_status() { + for status in [ + "401 Unauthorized", + "404 Not Found", + "500 Internal Server Error", + ] { + let base = respond_once(status, r#"{"error":"nope"}"#.to_string()).await; + let err = run_mistral_ocr(base, Duration::from_secs(5)) + .await + .expect_err("upstream error surfaces"); + let expected = status[..3].parse::().expect("status prefix parses"); + match err { + CoreError::Http { status: got, .. } => assert_eq!(got, expected), + other => panic!("expected Http error, got {other:?}"), + } + assert_eq!(err.public_status_code(), Some(expected)); + } + } + + #[tokio::test] + async fn ocr_bounds_oversized_error_body() { + let base = respond_once("500 Internal Server Error", "x".repeat(6000)).await; + let err = run_mistral_ocr(base, Duration::from_secs(5)) + .await + .expect_err("oversized error surfaces"); + match err { + CoreError::Http { body, .. } => { + assert!(body.ends_with("... (truncated)")); + assert!( + body.chars().count() < 300, + "body not bounded: {} chars", + body.chars().count() + ); + } + other => panic!("expected Http error, got {other:?}"), + } + } + + #[tokio::test] + async fn ocr_rejects_invalid_json_success_body() { + let base = respond_once("200 OK", "not json".to_string()).await; + let err = run_mistral_ocr(base, Duration::from_secs(5)) + .await + .expect_err("invalid JSON surfaces"); + assert!(matches!(err, CoreError::InvalidResponse(_))); + assert_eq!(err.public_status_code(), Some(500)); + } + + #[tokio::test] + async fn ocr_rejects_empty_success_body() { + let base = respond_once("200 OK", " ".to_string()).await; + let err = run_mistral_ocr(base, Duration::from_secs(5)) + .await + .expect_err("empty success surfaces"); + match err { + CoreError::InvalidResponse(message) => assert!(message.contains("empty")), + other => panic!("expected InvalidResponse, got {other:?}"), + } + } + + #[tokio::test] + async fn ocr_classifies_per_request_timeout() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("test listener binds"); + let addr = listener.local_addr().expect("listener has local addr"); + tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts one request"); + let _ = read_http_headers(&mut socket).await; + tokio::time::sleep(Duration::from_secs(3)).await; + let _ = socket.write_all(b"HTTP/1.1 200 OK\r\n\r\n").await; + }); + let err = run_mistral_ocr(format!("http://{addr}"), Duration::from_millis(200)) + .await + .expect_err("timeout surfaces"); + assert_eq!(err, CoreError::Timeout); + assert_eq!(err.public_status_code(), Some(408)); + } + + #[tokio::test] + async fn ocr_maps_unregistered_provider_to_invalid_provider() { + let err = ocr(OcrRequest { + model: "some-model", + document: json!({ + "type": "document_url", + "document_url": "https://example.com/doc.pdf" + }), + api_key: Some("sk-test"), + api_base: None, + custom_llm_provider: "definitely-not-a-provider", + extra_headers: None, + optional_params: Map::new(), + timeout: Some(Duration::from_secs(5)), + }) + .await + .expect_err("unregistered provider surfaces"); + assert!(matches!(err, CoreError::InvalidProvider(_))); + assert_eq!(err.public_status_code(), Some(400)); + assert_eq!(err.public_message(), "Invalid OCR request"); + } } diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr/common_utils.rs b/litellm-rust/crates/ai-gateway/src/io/ocr/common_utils.rs index 36d3d7da2e5..e965209a0dd 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr/common_utils.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/common_utils.rs @@ -28,6 +28,14 @@ const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120; const DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB: f64 = 50.0; const MAX_SAFE_FETCH_REDIRECTS: usize = 10; +pub(super) fn classify_reqwest_error(err: reqwest::Error) -> CoreError { + if err.is_timeout() { + CoreError::Timeout + } else { + CoreError::Network(err.to_string()) + } +} + pub(super) fn truncate_error_body(body: &str) -> String { if body.chars().count() <= ERROR_BODY_MAX_CHARS { return body.to_string(); @@ -227,7 +235,7 @@ async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response) .get(current_url.clone()) .send() .await - .map_err(|err| CoreError::Network(err.to_string()))?; + .map_err(classify_reqwest_error)?; if !response.status().is_redirection() { return Ok((current_url, response)); } @@ -268,11 +276,7 @@ async fn read_response_with_limit( let mut bytes = Vec::new(); let mut bytes_downloaded: u64 = 0; - while let Some(chunk) = response - .chunk() - .await - .map_err(|err| CoreError::Network(err.to_string()))? - { + while let Some(chunk) = response.chunk().await.map_err(classify_reqwest_error)? { bytes_downloaded += chunk.len() as u64; enforce_download_size(bytes_downloaded, max_bytes, url)?; bytes.extend_from_slice(&chunk); @@ -388,12 +392,9 @@ async fn upload_reducto_bytes( let response = request_builder .send() .await - .map_err(|err| CoreError::Network(err.to_string()))?; + .map_err(classify_reqwest_error)?; let status = response.status(); - let text = response - .text() - .await - .map_err(|err| CoreError::Network(err.to_string()))?; + let text = response.text().await.map_err(classify_reqwest_error)?; if !status.is_success() { return Err(CoreError::Http { status: status.as_u16(), @@ -477,7 +478,7 @@ fn operation_status(response_json: &Value) -> CoreResult<&str> { let status = response_json .get("status") .and_then(Value::as_str) - .ok_or(CoreError::MissingField("status"))?; + .ok_or(CoreError::missing_response_field("status"))?; match status { "succeeded" => Ok("succeeded"), "running" | "notStarted" => Ok("running"), @@ -515,10 +516,7 @@ pub(super) async fn poll_document_intelligence( )); loop { if start.elapsed() > timeout { - return Err(CoreError::Network(format!( - "Azure Document Intelligence operation polling timed out after {} seconds", - timeout.as_secs() - ))); + return Err(CoreError::Timeout); } let mut request_builder = http_client().get(operation_url); @@ -530,13 +528,10 @@ pub(super) async fn poll_document_intelligence( let response = request_builder .send() .await - .map_err(|err| CoreError::Network(err.to_string()))?; + .map_err(classify_reqwest_error)?; let retry_after = retry_after_secs(&response); let status = response.status(); - let text = response - .text() - .await - .map_err(|err| CoreError::Network(err.to_string()))?; + let text = response.text().await.map_err(classify_reqwest_error)?; if !status.is_success() { return Err(CoreError::Http { status: status.as_u16(), diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index b3e0519b772..fa6fbe75e7f 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -21,12 +21,59 @@ pub enum CoreError { Auth(String), #[error("OCR request failed with status {status}: {body}")] Http { status: u16, body: String }, + #[error("OCR request timed out")] + Timeout, #[error("OCR network error: {0}")] Network(String), #[error("routing error: {0}")] Routing(String), } +impl CoreError { + pub fn unexpected_response_type(value: &serde_json::Value) -> Self { + CoreError::InvalidResponse(format!( + "expected object OCR response, got {}", + json_type_name(value) + )) + } + + pub fn missing_response_field(field: &'static str) -> Self { + CoreError::InvalidResponse(format!("OCR response missing required field: {field}")) + } + + pub fn public_status_code(&self) -> Option { + match self { + CoreError::Http { status, .. } => Some(*status), + CoreError::Auth(_) => Some(401), + CoreError::InvalidType { .. } + | CoreError::MissingField(_) + | CoreError::InvalidProvider(_) + | CoreError::InvalidRequest(_) => Some(400), + CoreError::Timeout => Some(408), + CoreError::InvalidResponse(_) | CoreError::Routing(_) => Some(500), + CoreError::Network(_) => None, + } + } + + pub fn public_message(&self) -> String { + match self { + CoreError::Http { status, .. } => format!("OCR request failed with status {status}"), + CoreError::Network(_) => "OCR request could not reach the provider".to_string(), + CoreError::InvalidResponse(_) => { + "OCR provider returned an invalid response".to_string() + } + CoreError::Routing(_) => "OCR request could not be routed".to_string(), + CoreError::Auth(_) => "OCR request failed provider authentication".to_string(), + CoreError::InvalidRequest(_) | CoreError::InvalidProvider(_) => { + "Invalid OCR request".to_string() + } + CoreError::Timeout | CoreError::InvalidType { .. } | CoreError::MissingField(_) => { + self.to_string() + } + } + } +} + pub fn json_type_name(value: &serde_json::Value) -> &'static str { match value { serde_json::Value::Null => "null", @@ -37,3 +84,172 @@ pub fn json_type_name(value: &serde_json::Value) -> &'static str { serde_json::Value::Object(_) => "object", } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn public_status_code_preserves_public_contracts() { + assert_eq!( + CoreError::Http { + status: 404, + body: "not found".to_string() + } + .public_status_code(), + Some(404) + ); + assert_eq!( + CoreError::Auth("bad key".to_string()).public_status_code(), + Some(401) + ); + assert_eq!(CoreError::Timeout.public_status_code(), Some(408)); + assert_eq!( + CoreError::InvalidRequest("bad".to_string()).public_status_code(), + Some(400) + ); + assert_eq!( + CoreError::MissingField("document.type").public_status_code(), + Some(400) + ); + assert_eq!( + CoreError::InvalidType { + expected: "object", + actual: "string" + } + .public_status_code(), + Some(400) + ); + assert_eq!( + CoreError::InvalidResponse("empty".to_string()).public_status_code(), + Some(500) + ); + assert_eq!( + CoreError::Routing("no deployment".to_string()).public_status_code(), + Some(500) + ); + assert_eq!( + CoreError::Network("dns".to_string()).public_status_code(), + None + ); + assert_eq!( + CoreError::InvalidProvider("mistral".to_string()).public_status_code(), + Some(400) + ); + } + + #[test] + fn timeout_message_is_data_minimized() { + assert_eq!(CoreError::Timeout.to_string(), "OCR request timed out"); + } + + #[test] + fn public_message_hides_upstream_body() { + let err = CoreError::Http { + status: 500, + body: "signed-url=https://secret.example/token=abc123 leaked".to_string(), + }; + let message = err.public_message(); + assert_eq!(message, "OCR request failed with status 500"); + assert!(!message.contains("secret")); + assert!(!message.contains("abc123")); + } + + #[test] + fn public_message_hides_network_and_response_detail() { + let network = CoreError::Network( + "error sending request for url (https://signed.example/token=xyz)".to_string(), + ); + assert_eq!( + network.public_message(), + "OCR request could not reach the provider" + ); + assert!(!network.public_message().contains("token=xyz")); + + let invalid = CoreError::InvalidResponse("expected value at line 1 column 2".to_string()); + assert_eq!( + invalid.public_message(), + "OCR provider returned an invalid response" + ); + } + + #[test] + fn public_message_hides_auth_detail() { + let err = CoreError::Auth( + "google auth: failed to load service account key /secrets/sa.json token=ya29.abc123" + .to_string(), + ); + let message = err.public_message(); + assert_eq!(message, "OCR request failed provider authentication"); + assert!(!message.contains("sa.json")); + assert!(!message.contains("ya29")); + assert!(!message.contains("token")); + } + + #[test] + fn public_message_hides_invalid_request_detail() { + let err = CoreError::InvalidRequest( + "document_url=https://signed.example/doc.pdf?token=SECRET123 \ + base64=QUJDREVG header=x-api-key page=42" + .to_string(), + ); + let message = err.public_message(); + assert_eq!(message, "Invalid OCR request"); + assert!(!message.contains("SECRET123")); + assert!(!message.contains("token")); + assert!(!message.contains("base64")); + assert!(!message.contains("QUJDREVG")); + assert!(!message.contains("x-api-key")); + assert!(!message.contains("signed.example")); + assert_eq!(err.public_status_code(), Some(400)); + } + + #[test] + fn public_message_keeps_only_static_detail() { + assert_eq!( + CoreError::MissingField("document.type").public_message(), + "missing required field: document.type" + ); + assert_eq!( + CoreError::InvalidType { + expected: "object", + actual: "string" + } + .public_message(), + "expected object, got string" + ); + assert_eq!( + CoreError::Routing("model_list parse failed".to_string()).public_message(), + "OCR request could not be routed" + ); + } + + #[test] + fn malformed_response_maps_to_sanitized_500() { + let wrong_type = CoreError::unexpected_response_type(&serde_json::json!("boom")); + assert!(matches!(wrong_type, CoreError::InvalidResponse(_))); + assert_eq!(wrong_type.public_status_code(), Some(500)); + assert_eq!( + wrong_type.public_message(), + "OCR provider returned an invalid response" + ); + + let missing = CoreError::missing_response_field("status"); + assert!(matches!(missing, CoreError::InvalidResponse(_))); + assert_eq!(missing.public_status_code(), Some(500)); + assert_eq!( + missing.public_message(), + "OCR provider returned an invalid response" + ); + } + + #[test] + fn public_message_hides_invalid_provider_detail() { + let err = CoreError::InvalidProvider( + "no OCR provider 'internal-secret-router' registered for model 'gpt-8'".to_string(), + ); + assert_eq!(err.public_message(), "Invalid OCR request"); + assert!(!err.public_message().contains("internal-secret-router")); + assert_eq!(err.public_status_code(), Some(400)); + } +} diff --git a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs index 060073acd47..57e14214581 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs @@ -362,14 +362,11 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig { ) -> CoreResult { let response = response_json .as_object() - .ok_or_else(|| CoreError::InvalidType { - expected: "object", - actual: json_type_name(&response_json), - })?; + .ok_or_else(|| CoreError::unexpected_response_type(&response_json))?; let status = response .get("status") .and_then(Value::as_str) - .ok_or(CoreError::MissingField("status"))?; + .ok_or(CoreError::missing_response_field("status"))?; if status != "succeeded" { return Err(CoreError::InvalidResponse(format!( "Azure Document Intelligence analysis failed with status: {status}" @@ -517,4 +514,24 @@ mod tests { Some(json!({"pages_processed": 1, "doc_size_bytes": null})) ); } + + #[test] + fn document_intelligence_response_missing_status_is_invalid_response() { + let err = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG + .transform_ocr_response("prebuilt-layout", json!({"analyzeResult": {"pages": []}})) + .expect_err("missing status should be rejected"); + + assert!(matches!(err, CoreError::InvalidResponse(_))); + assert_eq!(err.public_status_code(), Some(500)); + } + + #[test] + fn document_intelligence_response_non_object_is_invalid_response() { + let err = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG + .transform_ocr_response("prebuilt-layout", json!("boom")) + .expect_err("non-object provider response should be rejected"); + + assert!(matches!(err, CoreError::InvalidResponse(_))); + assert_eq!(err.public_status_code(), Some(500)); + } } 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 1a33bc1e951..4e6be1d8a04 100644 --- a/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs @@ -107,10 +107,7 @@ impl OcrProviderConfig for MistralOcrConfig { ) -> CoreResult { let response_object = response_json .as_object() - .ok_or_else(|| CoreError::InvalidType { - expected: "object", - actual: json_type_name(&response_json), - })?; + .ok_or_else(|| CoreError::unexpected_response_type(&response_json))?; let pages = response_object .get("pages") @@ -257,6 +254,15 @@ mod tests { ); } + #[test] + fn transform_ocr_response_rejects_non_object_as_invalid_response() { + let err = transform_ocr_response("mistral-ocr-latest", json!("boom")) + .expect_err("non-object provider response should be rejected"); + + assert!(matches!(err, CoreError::InvalidResponse(_))); + assert_eq!(err.public_status_code(), Some(500)); + } + #[test] fn transform_ocr_response_normalizes_mistral_json() { let response = json!({ diff --git a/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs index b1241a14fea..f3f358bbeec 100644 --- a/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs @@ -237,10 +237,7 @@ impl OcrProviderConfig for ReductoParseLegacyConfig { fn transform_reducto_response(model: &str, response_json: Value) -> CoreResult { let response = response_json .as_object() - .ok_or_else(|| CoreError::InvalidType { - expected: "object", - actual: json_type_name(&response_json), - })?; + .ok_or_else(|| CoreError::unexpected_response_type(&response_json))?; let result = response.get("result").unwrap_or(&response_json); let usage = response.get("usage").cloned().unwrap_or_else(|| json!({})); Ok(OcrResponseData { diff --git a/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs index 8639926c435..7a5b2c6094f 100644 --- a/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs @@ -292,10 +292,7 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig { ) -> CoreResult { let response = response_json .as_object() - .ok_or_else(|| CoreError::InvalidType { - expected: "object", - actual: json_type_name(&response_json), - })?; + .ok_or_else(|| CoreError::unexpected_response_type(&response_json))?; let usage = response.get("usage").cloned(); let content = first_choice_content(&response_json)?; let mut ocr_data = ocr_data_from_content(content.clone(), usage.clone(), model); @@ -314,10 +311,9 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig { }); } - let object = ocr_data.as_object().ok_or_else(|| CoreError::InvalidType { - expected: "object", - actual: json_type_name(&ocr_data), - })?; + let object = ocr_data + .as_object() + .ok_or_else(|| CoreError::unexpected_response_type(&ocr_data))?; let pages = object .get("pages") .and_then(Value::as_array) diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 82024e1bf47..0e27c04027d 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -2,7 +2,7 @@ use std::time::Duration; use litellm_ai_gateway::io::ocr::{ocr as run_ocr, OcrRequest}; use litellm_core::error::CoreError; -use pyo3::exceptions::{PyRuntimeError, PyValueError}; +use pyo3::exceptions::PyValueError; use pyo3::prelude::*; use pyo3::types::{PyAny, PyDict}; use serde_json::{Map, Value}; @@ -29,15 +29,22 @@ fn json_to_py(py: Python<'_>, value: Value) -> PyResult> { Ok(json.call_method1("loads", (encoded,))?.unbind()) } -fn core_error_to_pyerr(err: CoreError) -> PyErr { - match err { - CoreError::Auth(message) => PyValueError::new_err(message), - CoreError::InvalidProvider(_) - | CoreError::InvalidRequest(_) - | CoreError::InvalidType { .. } - | CoreError::MissingField(_) => PyValueError::new_err(err.to_string()), - other => PyRuntimeError::new_err(other.to_string()), - } +fn core_error_to_pyerr(py: Python<'_>, err: CoreError) -> PyErr { + let status_code = err.public_status_code(); + let message = err.public_message(); + build_rust_ocr_error(py, &message, status_code).unwrap_or_else(|import_err| import_err) +} + +fn build_rust_ocr_error( + py: Python<'_>, + message: &str, + status_code: Option, +) -> PyResult { + let exc_type = py + .import("litellm.ocr.rust_bridge")? + .getattr("RustOcrError")?; + let instance = exc_type.call1((message, status_code))?; + Ok(PyErr::from_value(instance)) } fn optional_object_to_map( @@ -120,7 +127,7 @@ fn ocr( match result { Ok(value) => json_to_py(py, value), - Err(err) => Err(core_error_to_pyerr(err)), + Err(err) => Err(core_error_to_pyerr(py, err)), } } @@ -159,7 +166,7 @@ fn aocr( timeout, }) .await - .map_err(core_error_to_pyerr)?; + .map_err(|err| Python::with_gil(|py| core_error_to_pyerr(py, err)))?; Python::with_gil(|py| json_to_py(py, value)) }) diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 9225b6ff13c..7cb5a462de9 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -10,6 +10,7 @@ from io import IOBase from typing import Any, Coroutine, Union, cast import httpx +from typing_extensions import Never import litellm from litellm._logging import verbose_logger @@ -19,12 +20,112 @@ from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.ocr.rust_bridge import ( RustAocr, RustOcr, + RustOcrError, load_rust_aocr, load_rust_ocr, ) from litellm.utils import client, filter_out_litellm_params +class _OCRInputError(ValueError): + pass + + +def _ocr_error_response(status_code: int) -> httpx.Response: + return httpx.Response( + status_code=status_code, + request=httpx.Request(method="POST", url="https://litellm.ai"), + ) + + +def _raise_rust_ocr_exception( + err: RustOcrError, model: str, custom_llm_provider: str | None +) -> Never: + provider = custom_llm_provider or "mistral" + status_code = err.status_code + message = err.message + match status_code: + case None: + raise litellm.APIConnectionError( + message=message, llm_provider=provider, model=model + ) + case 400: + raise litellm.BadRequestError( + message=message, model=model, llm_provider=provider + ) + case 401: + raise litellm.AuthenticationError( + message=message, llm_provider=provider, model=model + ) + case 403: + raise litellm.PermissionDeniedError( + message=message, + llm_provider=provider, + model=model, + response=_ocr_error_response(403), + ) + case 404: + raise litellm.NotFoundError( + message=message, model=model, llm_provider=provider + ) + case 408: + raise litellm.Timeout(message=message, model=model, llm_provider=provider) + case 422: + raise litellm.UnprocessableEntityError( + message=message, + model=model, + llm_provider=provider, + response=_ocr_error_response(422), + ) + case 429: + raise litellm.RateLimitError( + message=message, llm_provider=provider, model=model + ) + case 500: + raise litellm.InternalServerError( + message=message, llm_provider=provider, model=model + ) + case 502: + raise litellm.BadGatewayError( + message=message, llm_provider=provider, model=model + ) + case 503: + raise litellm.ServiceUnavailableError( + message=message, llm_provider=provider, model=model + ) + case _: + raise litellm.APIError( + status_code=status_code, + message=message, + llm_provider=provider, + model=model, + ) + + +def _raise_ocr_exception( + e: Exception, + model: str, + custom_llm_provider: str | None, + completion_kwargs: dict[str, object], + kwargs: dict[str, object], +) -> Never: + if isinstance(e, RustOcrError): + _raise_rust_ocr_exception(e, model, custom_llm_provider) + if isinstance(e, _OCRInputError): + raise litellm.BadRequestError( + message="Invalid OCR request", + model=model, + llm_provider=custom_llm_provider or "mistral", + ) from e + raise litellm.exception_type( + model=model, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=completion_kwargs, + extra_kwargs=kwargs, + ) + + def _timeout_to_seconds( timeout: Union[float, httpx.Timeout] | None, ) -> float | None: @@ -65,7 +166,7 @@ def _resolve_ocr_call_context( litellm_call_id = cast(str | None, kwargs.get("litellm_call_id", None)) if not isinstance(document, dict): - raise ValueError( + raise _OCRInputError( f"document must be a dict with 'type' and URL/file field, got {type(document)}" ) @@ -76,7 +177,7 @@ def _resolve_ocr_call_context( doc_type = document.get("type") if doc_type not in ["document_url", "image_url"]: - raise ValueError( + raise _OCRInputError( f"Invalid document type: {doc_type}. " "Must be 'document_url', 'image_url', or 'file'" ) @@ -357,12 +458,12 @@ async def aocr( litellm_logging_obj=litellm_logging_obj, ) except Exception as e: - raise litellm.exception_type( + _raise_ocr_exception( + e, model=model, custom_llm_provider=custom_llm_provider, - original_exception=e, completion_kwargs=completion_kwargs, - extra_kwargs=kwargs, + kwargs=kwargs, ) @@ -420,7 +521,7 @@ def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str, """ file_input = document.get("file") if file_input is None: - raise ValueError( + raise _OCRInputError( "document with type='file' must include a 'file' field containing " "a pathlib.Path, file-like object, or bytes" ) @@ -436,7 +537,7 @@ def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str, # Opening it as a path is an arbitrary local file read on the proxy # host, which is then base64-encoded and forwarded to the OCR # provider — an exfiltration primitive. - raise ValueError( + raise _OCRInputError( "OCR file input does not accept bare str values. Pass bytes, " "a pathlib.Path, or a file-like object. To OCR a local file " "from a path, call open(path, 'rb') yourself." @@ -446,7 +547,7 @@ def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str, # Python-level type that HTTP form values can't fabricate. file_path = str(file_input) if not os.path.isfile(file_path): - raise FileNotFoundError(f"File not found: {file_path}") + raise _OCRInputError("OCR file input path does not exist") mime_type = get_mime_type(file_path) file_name = os.path.basename(file_path) with open(file_path, "rb") as f: @@ -462,19 +563,19 @@ def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str, if isinstance(file_bytes, str): file_bytes = file_bytes.encode("utf-8") else: - raise ValueError( + raise _OCRInputError( f"Unsupported file input type: {type(file_input)}. " "Expected pathlib.Path, bytes, or a file-like object." ) if not file_bytes: - raise ValueError("File is empty or could not be read") + raise _OCRInputError("File is empty or could not be read") if "mime_type" in document: mime_type = document["mime_type"] if not _MIME_PATTERN.match(mime_type): - raise ValueError(f"Invalid MIME type: {mime_type}") + raise _OCRInputError(f"Invalid MIME type: {mime_type}") base64_data = base64.b64encode(file_bytes).decode("utf-8") data_uri = f"data:{mime_type};base64,{base64_data}" @@ -621,10 +722,10 @@ def ocr( litellm_logging_obj=litellm_logging_obj, ) except Exception as e: - raise litellm.exception_type( + _raise_ocr_exception( + e, model=model, custom_llm_provider=custom_llm_provider, - original_exception=e, completion_kwargs=completion_kwargs, - extra_kwargs=kwargs, + kwargs=kwargs, ) diff --git a/litellm/ocr/rust_bridge.py b/litellm/ocr/rust_bridge.py index c518c2cb874..21f4656f472 100644 --- a/litellm/ocr/rust_bridge.py +++ b/litellm/ocr/rust_bridge.py @@ -13,6 +13,13 @@ from __future__ import annotations from typing import Awaitable, Final, Protocol, cast +class RustOcrError(Exception): + def __init__(self, message: str, status_code: int | None = None) -> None: + super().__init__(message) + self.message = message + self.status_code = status_code + + class RustOcr(Protocol): """Signature of the compiled Rust OCR entrypoint.""" diff --git a/tests/test_litellm/ocr/test_ocr_file_input.py b/tests/test_litellm/ocr/test_ocr_file_input.py index e6216d7c580..935a73e7de6 100644 --- a/tests/test_litellm/ocr/test_ocr_file_input.py +++ b/tests/test_litellm/ocr/test_ocr_file_input.py @@ -200,12 +200,14 @@ class TestConvertFileDocumentToUrlDocument: convert_file_document_to_url_document({"type": "file"}) def test_should_raise_error_for_nonexistent_pathlib_path(self): - """Non-existent pathlib.Path should raise FileNotFoundError.""" - with pytest.raises(FileNotFoundError, match="File not found"): + """Non-existent pathlib.Path should raise a path-free input error.""" + with pytest.raises(ValueError, match="does not exist") as exc_info: convert_file_document_to_url_document( {"type": "file", "file": Path("/nonexistent/path/to/file.pdf")} ) + assert "/nonexistent/path/to/file.pdf" not in str(exc_info.value) + def test_should_raise_error_for_empty_file(self): """Empty file should raise ValueError.""" with tempfile.NamedTemporaryFile(suffix=".pdf", delete=False) as f: diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 5647b72667e..5a9a48ab81a 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -3,12 +3,15 @@ import importlib import builtins import types +from dataclasses import dataclass import httpx +import pydantic import pytest import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse +from litellm.ocr.rust_bridge import RustOcrError # `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr` # function onto `litellm.ocr` and shadows the submodule, so import the modules @@ -36,6 +39,40 @@ class CapturedException(Exception): pass +@dataclass(slots=True) +class ExceptionTypeCall: + model: str + custom_llm_provider: str | None + original_exception: Exception + completion_kwargs: dict[str, object] + extra_kwargs: dict[str, object] + + +class ExceptionTypeSpy: + def __init__(self) -> None: + self.calls: list[ExceptionTypeCall] = [] + + def __call__( + self, + *, + model: str, + custom_llm_provider: str | None, + original_exception: Exception, + completion_kwargs: dict[str, object], + extra_kwargs: dict[str, object], + ) -> CapturedException: + self.calls.append( + ExceptionTypeCall( + model=model, + custom_llm_provider=custom_llm_provider, + original_exception=original_exception, + completion_kwargs=completion_kwargs, + extra_kwargs=extra_kwargs, + ) + ) + return CapturedException("wrapped") + + class RecordingBridge: """A fake ``RustOcr`` callable that records the args it was handed.""" @@ -130,6 +167,44 @@ class RaisingAsyncBridge: raise RuntimeError("bridge failed") +class RustErrorBridge: + def __init__(self, message: str, status_code: int | None) -> None: + self.message = message + self.status_code = status_code + + def __call__( + self, + model: str, + document: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> dict[str, object]: + raise RustOcrError(self.message, self.status_code) + + +class RustErrorAsyncBridge: + def __init__(self, message: str, status_code: int | None) -> None: + self.message = message + self.status_code = status_code + + async def __call__( + self, + model: str, + document: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> dict[str, object]: + raise RustOcrError(self.message, self.status_code) + + class RecordingLogging: """A spy standing in for ``LiteLLMLoggingObj`` to capture ``pre_call``.""" @@ -317,6 +392,7 @@ def test_run_rust_ocr_forwards_args_and_wraps_response(): "timeout_seconds": 12.5, } + def test_run_rust_ocr_runs_pre_call_logging(): """The Rust shortcut must run pre_call so callbacks and spend tracking fire.""" logging_obj = RecordingLogging() @@ -396,21 +472,16 @@ def test_ocr_routes_azure_ai_to_rust_by_default(fake_bridge): def test_ocr_exception_type_uses_resolved_provider_context( monkeypatch: pytest.MonkeyPatch, -): - captured: dict[str, object] = {} - - def fake_exception_type(**kwargs: object) -> CapturedException: - captured.update(kwargs) - return CapturedException("wrapped") - - monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) +) -> None: + spy = ExceptionTypeSpy() + monkeypatch.setattr(ocr_main.litellm, "exception_type", spy) rust_bridge._set_rust_ocr_bridge(ocr=RaisingBridge()) with pytest.raises(CapturedException): litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") - assert captured["model"] == "mistral-ocr-latest" - assert captured["custom_llm_provider"] == "mistral" + assert spy.calls[0].model == "mistral-ocr-latest" + assert spy.calls[0].custom_llm_provider == "mistral" @pytest.mark.asyncio @@ -438,21 +509,16 @@ async def test_aocr_routes_to_async_rust_by_default(fake_async_bridge): @pytest.mark.asyncio async def test_aocr_exception_type_uses_resolved_provider_context( monkeypatch: pytest.MonkeyPatch, -): - captured: dict[str, object] = {} - - def fake_exception_type(**kwargs: object) -> CapturedException: - captured.update(kwargs) - return CapturedException("wrapped") - - monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) +) -> None: + spy = ExceptionTypeSpy() + monkeypatch.setattr(ocr_main.litellm, "exception_type", spy) rust_bridge._set_rust_ocr_bridge(aocr=RaisingAsyncBridge()) with pytest.raises(CapturedException): await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") - assert captured["model"] == "mistral-ocr-latest" - assert captured["custom_llm_provider"] == "mistral" + assert spy.calls[0].model == "mistral-ocr-latest" + assert spy.calls[0].custom_llm_provider == "mistral" def test_ocr_forwards_timeout_to_rust(fake_bridge): @@ -480,3 +546,258 @@ def test_ocr_requires_rust_bridge_when_unavailable(monkeypatch): litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") assert "Rust OCR bridge is required" in str(exc_info.value) + + +RUST_OCR_ERROR_CASES = [ + pytest.param(400, litellm.BadRequestError, 400, id="400_bad_request"), + pytest.param(401, litellm.AuthenticationError, 401, id="401_authentication"), + pytest.param(403, litellm.PermissionDeniedError, 403, id="403_permission_denied"), + pytest.param(404, litellm.NotFoundError, 404, id="404_not_found"), + pytest.param(408, litellm.Timeout, 408, id="408_timeout"), + pytest.param( + 422, litellm.UnprocessableEntityError, 422, id="422_unprocessable_entity" + ), + pytest.param(429, litellm.RateLimitError, 429, id="429_rate_limit"), + pytest.param(500, litellm.InternalServerError, 500, id="500_internal"), + pytest.param(502, litellm.BadGatewayError, 502, id="502_bad_gateway"), + pytest.param(503, litellm.ServiceUnavailableError, 503, id="503_unavailable"), + pytest.param(None, litellm.APIConnectionError, 500, id="none_connection"), +] + + +@pytest.mark.parametrize( + ("status_code", "expected_exception", "expected_status"), RUST_OCR_ERROR_CASES +) +def test_rust_ocr_error_maps_to_public_exception( + status_code: int | None, + expected_exception: type[Exception], + expected_status: int, +) -> None: + with pytest.raises(expected_exception) as exc_info: + ocr_main._raise_rust_ocr_exception( + RustOcrError("upstream boom", status_code), + model="mistral-ocr-latest", + custom_llm_provider="mistral", + ) + + exc = exc_info.value + assert exc.status_code == expected_status + assert exc.llm_provider == "mistral" + assert exc.model == "mistral-ocr-latest" + assert "upstream boom" in str(exc) + + +@pytest.mark.parametrize( + ("status_code", "expected_exception", "expected_status"), RUST_OCR_ERROR_CASES +) +def test_ocr_raises_typed_exception_from_rust_error( + status_code: int | None, + expected_exception: type[Exception], + expected_status: int, +) -> None: + rust_bridge._set_rust_ocr_bridge(ocr=RustErrorBridge("upstream boom", status_code)) + + with pytest.raises(expected_exception) as exc_info: + litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + + assert exc_info.value.status_code == expected_status + assert exc_info.value.llm_provider == "mistral" + + +@pytest.mark.parametrize( + ("status_code", "expected_exception", "expected_status"), RUST_OCR_ERROR_CASES +) +@pytest.mark.asyncio +async def test_aocr_raises_typed_exception_from_rust_error( + status_code: int | None, + expected_exception: type[Exception], + expected_status: int, +) -> None: + rust_bridge._set_rust_ocr_bridge( + aocr=RustErrorAsyncBridge("upstream boom", status_code) + ) + + with pytest.raises(expected_exception) as exc_info: + await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + + assert exc_info.value.status_code == expected_status + assert exc_info.value.llm_provider == "mistral" + + +UNKNOWN_STATUS_CASES = [ + pytest.param(409, id="409_conflict"), + pytest.param(451, id="451_legal_reasons"), + pytest.param(504, id="504_gateway_timeout"), + pytest.param(418, id="418_teapot"), +] + + +@pytest.mark.parametrize("status_code", UNKNOWN_STATUS_CASES) +def test_rust_ocr_error_unknown_status_preserves_exact_status( + status_code: int, +) -> None: + with pytest.raises(litellm.APIError) as exc_info: + ocr_main._raise_rust_ocr_exception( + RustOcrError("upstream boom", status_code), + model="mistral-ocr-latest", + custom_llm_provider="mistral", + ) + + exc = exc_info.value + assert type(exc) is litellm.APIError + assert exc.status_code == status_code + assert exc.llm_provider == "mistral" + assert exc.model == "mistral-ocr-latest" + + +def test_rust_ocr_error_message_is_preserved_bounded() -> None: + bounded = "x" * 256 + "... (truncated)" + with pytest.raises(litellm.InternalServerError) as exc_info: + ocr_main._raise_rust_ocr_exception( + RustOcrError(bounded, 500), + model="mistral-ocr-latest", + custom_llm_provider="mistral", + ) + + assert bounded in str(exc_info.value) + + +INVALID_OCR_INPUTS = [ + pytest.param({"type": "bogus", "document_url": "https://x/y.pdf"}, id="bad_type"), + pytest.param("not-a-dict", id="non_dict_document"), +] + + +@pytest.mark.parametrize("document", INVALID_OCR_INPUTS) +def test_ocr_invalid_input_raises_bad_request(document: object) -> None: + with pytest.raises(litellm.BadRequestError) as exc_info: + litellm.ocr(model=MODEL, document=document, api_key="sk-test") + + assert exc_info.value.status_code == 400 + + +@pytest.mark.parametrize("document", INVALID_OCR_INPUTS) +@pytest.mark.asyncio +async def test_aocr_invalid_input_raises_bad_request(document: object) -> None: + with pytest.raises(litellm.BadRequestError) as exc_info: + await litellm.aocr(model=MODEL, document=document, api_key="sk-test") + + assert exc_info.value.status_code == 400 + + +def test_raise_ocr_exception_maps_input_error_to_bad_request() -> None: + with pytest.raises(litellm.BadRequestError) as exc_info: + ocr_main._raise_ocr_exception( + ocr_main._OCRInputError("Invalid document type: bogus"), + model="mistral-ocr-latest", + custom_llm_provider="mistral", + completion_kwargs={}, + kwargs={}, + ) + + exc = exc_info.value + assert exc.status_code == 400 + assert "Invalid OCR request" in str(exc) + assert "bogus" not in str(exc) + + +OCR_INPUT_CANARIES = [ + "https://signed.example/doc.pdf", + "token=SECRET123", + "QUJDREVGYmFzZTY0", + "page=42", + "application/x-canary-mime", + "/var/secrets/service_account.json", + "sk-canary-secret", + "canary_document_type", +] + + +def test_raise_ocr_exception_input_error_publishes_generic_message() -> None: + canary = " ".join(OCR_INPUT_CANARIES) + with pytest.raises(litellm.BadRequestError) as exc_info: + ocr_main._raise_ocr_exception( + ocr_main._OCRInputError(canary), + model="mistral-ocr-latest", + custom_llm_provider="mistral", + completion_kwargs={}, + kwargs={}, + ) + + message = str(exc_info.value) + assert "Invalid OCR request" in message + for marker in OCR_INPUT_CANARIES: + assert marker not in message + + +def test_ocr_input_error_public_message_drops_canaries() -> None: + document = { + "type": " ".join(OCR_INPUT_CANARIES), + "document_url": "https://signed.example/doc.pdf?token=SECRET123&page=42", + } + with pytest.raises(litellm.BadRequestError) as exc_info: + litellm.ocr(model=MODEL, document=document, api_key="sk-test") + + message = str(exc_info.value) + assert "Invalid OCR request" in message + for marker in OCR_INPUT_CANARIES: + assert marker not in message + + +@pytest.mark.asyncio +async def test_aocr_input_error_public_message_drops_canaries() -> None: + document = { + "type": " ".join(OCR_INPUT_CANARIES), + "document_url": "https://signed.example/doc.pdf?token=SECRET123&page=42", + } + with pytest.raises(litellm.BadRequestError) as exc_info: + await litellm.aocr(model=MODEL, document=document, api_key="sk-test") + + message = str(exc_info.value) + assert "Invalid OCR request" in message + for marker in OCR_INPUT_CANARIES: + assert marker not in message + + +def test_raise_ocr_exception_keeps_plain_value_error_off_bad_request( + monkeypatch: pytest.MonkeyPatch, +) -> None: + spy = ExceptionTypeSpy() + monkeypatch.setattr(ocr_main.litellm, "exception_type", spy) + + internal = ValueError("internal invariant broke") + with pytest.raises(CapturedException): + ocr_main._raise_ocr_exception( + internal, + model="mistral-ocr-latest", + custom_llm_provider="mistral", + completion_kwargs={}, + kwargs={}, + ) + + assert spy.calls[0].original_exception is internal + + +def test_raise_ocr_exception_keeps_validation_error_off_bad_request( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class _Model(pydantic.BaseModel): + value: int + + adapter = pydantic.TypeAdapter(_Model) + with pytest.raises(pydantic.ValidationError) as validation_info: + adapter.validate_python({"value": "not-an-int"}) + + spy = ExceptionTypeSpy() + monkeypatch.setattr(ocr_main.litellm, "exception_type", spy) + + with pytest.raises(CapturedException): + ocr_main._raise_ocr_exception( + validation_info.value, + model="mistral-ocr-latest", + custom_llm_provider="mistral", + completion_kwargs={}, + kwargs={}, + ) + + assert spy.calls[0].original_exception is validation_info.value