mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
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:
parent
6661462d5a
commit
95a8312a1e
4 changed files with 172 additions and 28 deletions
154
litellm-rust/crates/ai-gateway/src/errors.rs
Normal file
154
litellm-rust/crates/ai-gateway/src/errors.rs
Normal 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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue