refactor(ai-gateway): share one reqwest exception mapper across endpoints

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-17 02:57:08 +00:00
parent 6661462d5a
commit 95a8312a1e
4 changed files with 172 additions and 28 deletions

View file

@ -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<Duration>) -> 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());
}
}

View file

@ -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<Value> {
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<Value> {
.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 {

View file

@ -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(),

View file

@ -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;