diff --git a/litellm-rust/crates/ai-gateway/src/errors.rs b/litellm-rust/crates/ai-gateway/src/errors.rs new file mode 100644 index 00000000000..8452d74397c --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/errors.rs @@ -0,0 +1,154 @@ +//! Shared mapper from `reqwest` failures to typed [`CoreError`] contracts, +//! used by every endpoint in this crate that performs HTTP I/O. + +use litellm_core::error::CoreError; + +pub(crate) fn map_reqwest_error(err: reqwest::Error) -> CoreError { + if err.is_timeout() { + return CoreError::Timeout; + } + if let Some(status) = err.status() { + return CoreError::Http { + status: status.as_u16(), + body: String::new(), + }; + } + if err.is_decode() { + return CoreError::InvalidResponse(err.without_url().to_string()); + } + CoreError::Network(err.without_url().to_string()) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::time::Duration; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + const URL_CANARY: &str = "canary-host.invalid"; + const QUERY_CANARY: &str = "token=SECRET_CANARY_123"; + + async fn local_server(response: &'static str, delay: Option) -> 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 mut buffer = [0_u8; 1024]; + let _ = socket.read(&mut buffer).await; + if let Some(delay) = delay { + tokio::time::sleep(delay).await; + } + let _ = socket.write_all(response.as_bytes()).await; + }); + format!("http://{addr}") + } + + fn assert_no_canaries(text: &str) { + assert!(!text.contains(URL_CANARY), "leaked host: {text}"); + assert!(!text.contains("SECRET_CANARY_123"), "leaked query: {text}"); + } + + #[tokio::test] + async fn maps_request_timeout_to_timeout() { + let base = local_server( + "HTTP/1.1 200 OK\r\ncontent-length: 0\r\n\r\n", + Some(Duration::from_secs(3)), + ) + .await; + let err = reqwest::Client::new() + .get(format!("{base}/?{QUERY_CANARY}")) + .timeout(Duration::from_millis(50)) + .send() + .await + .expect_err("timeout surfaces"); + let mapped = map_reqwest_error(err); + assert_eq!(mapped, CoreError::Timeout); + assert_eq!(mapped.public_status_code(), Some(408)); + assert_no_canaries(&mapped.to_string()); + assert_no_canaries(&mapped.public_message()); + } + + #[tokio::test] + async fn maps_connect_failure_to_network() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("test listener binds"); + let addr = listener.local_addr().expect("listener has local addr"); + drop(listener); + let err = reqwest::Client::new() + .get(format!("http://{addr}/?{QUERY_CANARY}")) + .send() + .await + .expect_err("connect failure surfaces"); + assert!(err.is_connect()); + let mapped = map_reqwest_error(err); + assert!(matches!(mapped, CoreError::Network(_))); + assert_eq!(mapped.public_status_code(), None); + assert_no_canaries(&mapped.to_string()); + assert_no_canaries(&mapped.public_message()); + } + + #[tokio::test] + async fn maps_status_error_to_typed_http() { + let base = local_server( + "HTTP/1.1 503 Service Unavailable\r\ncontent-length: 0\r\nconnection: close\r\n\r\n", + None, + ) + .await; + let err = reqwest::Client::new() + .get(format!("{base}/?{QUERY_CANARY}")) + .send() + .await + .expect("response arrives") + .error_for_status() + .expect_err("status error surfaces"); + let mapped = map_reqwest_error(err); + assert_eq!( + mapped, + CoreError::Http { + status: 503, + body: String::new() + } + ); + assert_eq!(mapped.public_status_code(), Some(503)); + assert_no_canaries(&mapped.to_string()); + assert_no_canaries(&mapped.public_message()); + } + + #[tokio::test] + async fn maps_body_decode_failure_to_invalid_response() { + let base = local_server( + "HTTP/1.1 200 OK\r\ncontent-length: 100\r\nconnection: close\r\n\r\nshort", + None, + ) + .await; + let response = reqwest::Client::new() + .get(format!("{base}/?{QUERY_CANARY}")) + .send() + .await + .expect("response arrives"); + let err = response.text().await.expect_err("body read fails"); + assert!(err.is_decode()); + let mapped = map_reqwest_error(err); + assert!(matches!(mapped, CoreError::InvalidResponse(_))); + assert_eq!(mapped.public_status_code(), Some(500)); + assert_no_canaries(&mapped.to_string()); + assert_no_canaries(&mapped.public_message()); + } + + #[tokio::test] + async fn maps_unresolvable_host_to_network_without_url() { + let err = reqwest::Client::new() + .get(format!("http://{URL_CANARY}/?{QUERY_CANARY}")) + .send() + .await + .expect_err("dns failure surfaces"); + let mapped = map_reqwest_error(err); + assert!(matches!(mapped, CoreError::Network(_))); + assert_no_canaries(&mapped.to_string()); + assert_no_canaries(&mapped.public_message()); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr.rs b/litellm-rust/crates/ai-gateway/src/io/ocr.rs index c103efbb941..5d01b08bdae 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr.rs @@ -16,9 +16,10 @@ use serde_json::{Map, Value}; mod common_utils; +use crate::errors::map_reqwest_error; use common_utils::{ - 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, + 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 @@ -117,10 +118,7 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult { request_builder = request_builder.timeout(duration); } - let response = request_builder - .send() - .await - .map_err(classify_reqwest_error)?; + let response = request_builder.send().await.map_err(map_reqwest_error)?; let status = response.status(); if config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll @@ -145,7 +143,7 @@ pub async fn ocr(request: OcrRequest<'_>) -> CoreResult { .into_json()); } - let text = response.text().await.map_err(classify_reqwest_error)?; + let text = response.text().await.map_err(map_reqwest_error)?; if !status.is_success() { return Err(CoreError::Http { 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 e965209a0dd..7abeaecec94 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 @@ -9,6 +9,8 @@ use litellm_core::CoreResult; use reqwest::Url; use serde_json::{Map, Value}; +use crate::errors::map_reqwest_error; + use litellm_core::providers::azure_ai::ocr::transformation::{ AZURE_AI_OCR_CONFIG, AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG, }; @@ -28,14 +30,6 @@ 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(); @@ -225,7 +219,7 @@ async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response) let client = reqwest::Client::builder() .redirect(reqwest::redirect::Policy::none()) .build() - .map_err(|err| CoreError::Network(err.to_string()))?; + .map_err(map_reqwest_error)?; let mut current_url = Url::parse(url) .map_err(|err| CoreError::InvalidRequest(format!("invalid OCR document URL: {err}")))?; @@ -235,7 +229,7 @@ async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response) .get(current_url.clone()) .send() .await - .map_err(classify_reqwest_error)?; + .map_err(map_reqwest_error)?; if !response.status().is_redirection() { return Ok((current_url, response)); } @@ -276,7 +270,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(classify_reqwest_error)? { + while let Some(chunk) = response.chunk().await.map_err(map_reqwest_error)? { bytes_downloaded += chunk.len() as u64; enforce_download_size(bytes_downloaded, max_bytes, url)?; bytes.extend_from_slice(&chunk); @@ -389,12 +383,9 @@ async fn upload_reducto_bytes( request_builder = request_builder.timeout(duration); } - let response = request_builder - .send() - .await - .map_err(classify_reqwest_error)?; + let response = request_builder.send().await.map_err(map_reqwest_error)?; let status = response.status(); - let text = response.text().await.map_err(classify_reqwest_error)?; + let text = response.text().await.map_err(map_reqwest_error)?; if !status.is_success() { return Err(CoreError::Http { status: status.as_u16(), @@ -525,13 +516,10 @@ pub(super) async fn poll_document_intelligence( request_builder = request_builder.header(key, value); } } - let response = request_builder - .send() - .await - .map_err(classify_reqwest_error)?; + let response = request_builder.send().await.map_err(map_reqwest_error)?; let retry_after = retry_after_secs(&response); let status = response.status(); - let text = response.text().await.map_err(classify_reqwest_error)?; + let text = response.text().await.map_err(map_reqwest_error)?; if !status.is_success() { return Err(CoreError::Http { status: status.as_u16(), diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index 6c04fbb7626..120c7e66365 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -13,6 +13,10 @@ pub mod io; +/// Shared reqwest-failure classification into typed [`litellm_core::error::CoreError`] +/// contracts. Always available — every I/O endpoint maps transport failures here. +mod errors; + /// GIL-activity tracking. Pure (atomics only); shared by the `server` routes and /// the `python-config` reader, so it is available without either feature. pub mod gil;