refactor(ocr-security): split OCR host modules and harden URL/IPv6/poll validation
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-17 00:41:05 +00:00
parent c7e5f71251
commit d321c60252
9 changed files with 1339 additions and 1062 deletions

View file

@ -7,23 +7,50 @@
/// Default LiteLLM control-plane base URL for request-log egress when
/// `LITELLM_PROXY_BASE_URL` is unset.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_PROXY_BASE_URL: &str = "http://localhost:4000";
/// The logs ingest path appended to the proxy base. Not a tunable; it is the
/// proxy's API contract (the rust-control-plane router on the Python proxy).
#[cfg(feature = "server")]
pub(crate) const RUST_CONTROL_PLANE_LOGS_PATH: &str = "/v1/rust_control_plane/logs";
/// Default bounded channel depth for the log-egress worker.
/// Override: `LITELLM_LOG_CHANNEL_CAPACITY`.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_CHANNEL_CAPACITY: usize = 4096;
/// Default max records POSTed per request to the control plane.
/// Override: `LITELLM_LOG_BATCH_SIZE`.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_MAX_BATCH_SIZE: usize = 256;
/// Default partial-batch flush cadence, in ms.
/// Override: `LITELLM_LOG_FLUSH_INTERVAL_MS`.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500;
/// Provider attributed to realtime sessions in the logging payload.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_PROVIDER: &str = "openai";
pub(crate) const OCR_ERROR_BODY_MAX_CHARS: usize = 256;
pub(crate) const OCR_ERROR_BODY_MAX_BYTES: usize = OCR_ERROR_BODY_MAX_CHARS * 4;
pub(crate) const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120;
pub(crate) const AZURE_POLL_DEFAULT_RETRY_AFTER_SECS: u64 = 2;
pub(crate) const DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB: f64 = 50.0;
pub(crate) const MAX_SAFE_FETCH_REDIRECTS: usize = 10;
pub(crate) const MAX_DOCUMENT_CLIENTS: usize = 128;
pub(crate) const DOCUMENT_FETCH_TIMEOUT_SECS: u64 = 600;
pub(crate) const OCR_DOCUMENT_FETCH_NETWORK_ERROR: &str = "OCR document fetch failed";
pub(crate) const AZURE_POLL_NETWORK_ERROR: &str =
"Azure Document Intelligence polling request failed";

View file

@ -14,12 +14,18 @@ use litellm_core::ocr::transformation::{
use litellm_core::CoreResult;
use serde_json::{Map, Value};
mod azure_poll;
mod common_utils;
mod document_fetch;
mod reducto;
mod ssrf;
#[cfg(test)]
mod test_support;
use common_utils::{
convert_document_url_to_data_uri, has_header, ocr_provider_config, poll_document_intelligence,
read_error_body, string_headers, upload_reducto_document,
};
use azure_poll::poll_document_intelligence;
use common_utils::{has_header, ocr_provider_config, read_error_body, string_headers};
use document_fetch::convert_document_url_to_data_uri;
use reducto::upload_reducto_document;
/// OCR over large documents can take a while; bound it generously rather than
/// hanging forever on an unresponsive upstream. The client-level limit is the
@ -330,7 +336,7 @@ mod tests {
}
#[tokio::test]
async fn document_intelligence_poll_uses_resolved_subscription_key() {
async fn document_intelligence_rejects_unsafe_operation_location_end_to_end() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
@ -347,23 +353,10 @@ mod tests {
.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)
post_request
});
let response = ocr(OcrRequest {
let error = ocr(OcrRequest {
model: "prebuilt-read",
document: json!({
"type": "document_url",
@ -377,23 +370,20 @@ mod tests {
timeout: Some(Duration::from_secs(5)),
})
.await
.expect("document intelligence request succeeds");
.expect_err("loopback operation-location must be rejected");
assert_eq!(response["pages"][0]["markdown"], "ok");
assert!(
matches!(&error, CoreError::InvalidRequest(message) if message.contains("SSRF protection")),
"got {error:?}"
);
let (post_request, poll_request) = server.await.expect("server task completes");
let post_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]

View file

@ -0,0 +1,336 @@
use std::net::SocketAddr;
use std::time::{Duration, Instant};
use litellm_core::error::CoreError;
use litellm_core::CoreResult;
use reqwest::Url;
use serde_json::Value;
use crate::constants::{
AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS, AZURE_POLL_DEFAULT_RETRY_AFTER_SECS,
AZURE_POLL_NETWORK_ERROR,
};
use super::common_utils::read_error_body;
use super::ssrf::{blocked_url_error, pinned_client, resolve_validated};
fn same_origin(left: &str, right: &str) -> bool {
let Ok(left) = Url::parse(left) else {
return false;
};
let Ok(right) = Url::parse(right) else {
return false;
};
left.scheme() == right.scheme()
&& left.host_str() == right.host_str()
&& left.port_or_known_default() == right.port_or_known_default()
}
fn retry_after_secs(response: &reqwest::Response) -> u64 {
response
.headers()
.get(reqwest::header::RETRY_AFTER)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<u64>().ok())
.unwrap_or(AZURE_POLL_DEFAULT_RETRY_AFTER_SECS)
}
fn operation_status(response_json: &Value) -> CoreResult<&str> {
let status = response_json
.get("status")
.and_then(Value::as_str)
.ok_or(CoreError::MissingField("status"))?;
match status {
"succeeded" => Ok("succeeded"),
"running" | "notStarted" => Ok("running"),
"failed" => {
let message = response_json
.get("error")
.and_then(|error| error.get("message"))
.and_then(Value::as_str)
.unwrap_or("Unknown error");
Err(CoreError::InvalidResponse(format!(
"Azure Document Intelligence analysis failed: {message}"
)))
}
other => Err(CoreError::InvalidResponse(format!(
"Unknown operation status: {other}"
))),
}
}
fn poll_timeout_error(overall: Duration) -> CoreError {
CoreError::Network(format!(
"Azure Document Intelligence operation polling timed out after {} seconds",
overall.as_secs()
))
}
pub(super) async fn poll_document_intelligence(
operation_url: &str,
original_url: &str,
headers: &[(String, String)],
timeout: Option<Duration>,
) -> CoreResult<Value> {
poll_document_intelligence_with(
operation_url,
original_url,
headers,
timeout,
|url| async move { resolve_validated(&url).await },
)
.await
}
pub(super) async fn poll_document_intelligence_with<P, Fut>(
operation_url: &str,
original_url: &str,
headers: &[(String, String)],
timeout: Option<Duration>,
resolve: P,
) -> CoreResult<Value>
where
P: Fn(Url) -> Fut,
Fut: std::future::Future<Output = CoreResult<Vec<SocketAddr>>>,
{
if !same_origin(operation_url, original_url) {
return Err(CoreError::InvalidRequest(
"Azure Document Intelligence: rejected unsafe polling target".to_string(),
));
}
let parsed = Url::parse(operation_url).map_err(|_| blocked_url_error())?;
let start = Instant::now();
let overall = timeout.unwrap_or(Duration::from_secs(
AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS,
));
loop {
let Some(remaining) = overall.checked_sub(start.elapsed()) else {
return Err(poll_timeout_error(overall));
};
if remaining.is_zero() {
return Err(poll_timeout_error(overall));
}
let addresses = resolve(parsed.clone()).await?;
let client = pinned_client(&parsed, &addresses)?;
let mut request_builder = client.get(parsed.clone()).timeout(remaining);
for (key, value) in headers {
if key.eq_ignore_ascii_case("ocp-apim-subscription-key") {
request_builder = request_builder.header(key, value);
}
}
let response = request_builder
.send()
.await
.map_err(|_| CoreError::Network(AZURE_POLL_NETWORK_ERROR.to_string()))?;
let retry_after = retry_after_secs(&response);
let status = response.status();
if !status.is_success() {
return Err(CoreError::Http {
status: status.as_u16(),
body: read_error_body(response).await,
});
}
let text = response
.text()
.await
.map_err(|_| CoreError::Network(AZURE_POLL_NETWORK_ERROR.to_string()))?;
let response_json: Value = serde_json::from_str(&text).map_err(|err| {
CoreError::InvalidResponse(format!("invalid Azure DI poll response JSON: {err}"))
})?;
if operation_status(&response_json)? == "succeeded" {
return Ok(response_json);
}
let remaining_after = overall
.checked_sub(start.elapsed())
.unwrap_or(Duration::ZERO);
if remaining_after.is_zero() {
return Err(poll_timeout_error(overall));
}
let sleep = Duration::from_secs(retry_after).min(remaining_after);
tokio::time::sleep(sleep).await;
}
}
#[cfg(test)]
mod tests {
use super::super::test_support::{http_response, spawn_counting_server};
use super::*;
#[tokio::test]
async fn foreign_operation_location_rejected_without_connecting() {
let foreign = spawn_counting_server(http_response(
"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n",
b"{}",
))
.await;
let operation_url = format!("http://127.0.0.1:{}/op", foreign.addr.port());
let original_url = format!("http://127.0.0.1:{}/analyze", foreign.addr.port() ^ 1);
let error = poll_document_intelligence(&operation_url, &original_url, &[], None)
.await
.unwrap_err();
assert!(
matches!(&error, CoreError::InvalidRequest(message) if message.contains("polling target")),
"got {error:?}"
);
assert_eq!(foreign.connection_count(), 0);
}
#[tokio::test]
async fn same_origin_operation_location_polls_pinned_target() {
let body = br#"{"status":"succeeded","analyzeResult":{}}"#;
let server = spawn_counting_server(http_response(
&format!("HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n", body.len()),
body,
))
.await;
let server_addr = server.addr;
let operation_url = format!("http://127.0.0.1:{}/op", server_addr.port());
let result = poll_document_intelligence_with(
&operation_url,
&operation_url,
&[],
None,
move |_url| async move { Ok(vec![server_addr]) },
)
.await;
assert!(result.is_ok(), "got {result:?}");
assert!(server.connection_count() >= 1);
}
#[tokio::test]
async fn poll_forwards_only_subscription_key_header() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let capture = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buffer = vec![0u8; 4096];
let n = socket.read(&mut buffer).await.unwrap();
let request = String::from_utf8_lossy(&buffer[..n]).to_string();
let body = br#"{"status":"succeeded","analyzeResult":{}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
socket.write_all(response.as_bytes()).await.unwrap();
socket.write_all(body).await.unwrap();
socket.flush().await.unwrap();
request
});
let operation_url = format!("http://127.0.0.1:{}/op", addr.port());
let headers = vec![
(
"ocp-apim-subscription-key".to_string(),
"di-secret".to_string(),
),
("authorization".to_string(), "Bearer leak".to_string()),
];
let result = poll_document_intelligence_with(
&operation_url,
&operation_url,
&headers,
None,
move |_url| async move { Ok(vec![addr]) },
)
.await
.expect("poll succeeds");
assert_eq!(result["status"], "succeeded");
let request = capture.await.unwrap().to_ascii_lowercase();
assert!(
request.contains("ocp-apim-subscription-key: di-secret"),
"{request}"
);
assert!(!request.contains("authorization:"), "{request}");
}
#[tokio::test]
async fn poll_rejects_operation_location_resolving_to_blocked_address() {
let operation_url = "http://127.0.0.1:9/op";
let error = poll_document_intelligence(operation_url, operation_url, &[], None)
.await
.unwrap_err();
assert!(
matches!(&error, CoreError::InvalidRequest(message) if message.contains("SSRF protection")),
"loopback poll target must be rejected, got {error:?}"
);
}
#[tokio::test]
async fn poll_timeout_override_is_enforced() {
let body = br#"{"status":"running"}"#;
let server = spawn_counting_server(http_response(
&format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\nRetry-After: 0\r\n\r\n",
body.len()
),
body,
))
.await;
let server_addr = server.addr;
let operation_url = format!("http://127.0.0.1:{}/op", server_addr.port());
let error = poll_document_intelligence_with(
&operation_url,
&operation_url,
&[],
Some(Duration::from_millis(150)),
move |_url| async move { Ok(vec![server_addr]) },
)
.await
.unwrap_err();
assert!(
matches!(&error, CoreError::Network(message) if message.contains("timed out")),
"got {error:?}"
);
}
#[tokio::test]
async fn poll_bounds_retry_after_by_remaining_timeout() {
let body = br#"{"status":"running"}"#;
let server = spawn_counting_server(http_response(
&format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\nRetry-After: 3600\r\n\r\n",
body.len()
),
body,
))
.await;
let server_addr = server.addr;
let operation_url = format!("http://127.0.0.1:{}/op", server_addr.port());
let started = Instant::now();
let error = poll_document_intelligence_with(
&operation_url,
&operation_url,
&[],
Some(Duration::from_millis(200)),
move |_url| async move { Ok(vec![server_addr]) },
)
.await
.unwrap_err();
assert!(
matches!(&error, CoreError::Network(message) if message.contains("timed out")),
"got {error:?}"
);
assert!(
started.elapsed() < Duration::from_secs(5),
"a huge Retry-After must be bounded by the remaining timeout"
);
}
}

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,401 @@
use std::net::SocketAddr;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use base64::Engine;
use litellm_core::error::CoreError;
use litellm_core::CoreResult;
use reqwest::Url;
use serde_json::{Map, Value};
use crate::constants::{
DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB, MAX_SAFE_FETCH_REDIRECTS,
OCR_DOCUMENT_FETCH_NETWORK_ERROR,
};
use super::common_utils::{document_url_field, read_error_body};
use super::ssrf::{parse_fetchable_url, pinned_client, resolve_validated};
fn document_fetch_network_error() -> CoreError {
CoreError::Network(OCR_DOCUMENT_FETCH_NETWORK_ERROR.to_string())
}
fn max_document_download_bytes() -> u64 {
let max_size_mb = std::env::var("MAX_IMAGE_URL_DOWNLOAD_SIZE_MB")
.ok()
.and_then(|value| value.parse::<f64>().ok())
.unwrap_or(DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB);
(max_size_mb.max(0.0) * 1024.0 * 1024.0) as u64
}
fn redirect_location(response: &reqwest::Response, url: &Url) -> CoreResult<Url> {
let location = response
.headers()
.get(reqwest::header::LOCATION)
.and_then(|value| value.to_str().ok())
.ok_or_else(|| {
CoreError::InvalidResponse("OCR document redirect missing Location header".to_string())
})?;
url.join(location)
.map_err(|_| CoreError::InvalidResponse("invalid OCR document redirect".to_string()))
}
async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response)> {
fetch_with_redirects(url, |candidate| async move {
resolve_validated(&candidate).await
})
.await
}
async fn fetch_with_redirects<P, Fut>(url: &str, resolve: P) -> CoreResult<(Url, reqwest::Response)>
where
P: Fn(Url) -> Fut,
Fut: std::future::Future<Output = CoreResult<Vec<SocketAddr>>>,
{
let mut current_url = Url::parse(url).map_err(|_| {
CoreError::InvalidRequest("OCR document URL rejected by SSRF protection".to_string())
})?;
for _ in 0..MAX_SAFE_FETCH_REDIRECTS {
let addresses = resolve(current_url.clone()).await?;
let client = pinned_client(&current_url, &addresses)?;
let response = client
.get(current_url.clone())
.send()
.await
.map_err(|_| document_fetch_network_error())?;
if !response.status().is_redirection() {
return Ok((current_url, response));
}
current_url = redirect_location(&response, &current_url)?;
}
Err(CoreError::InvalidRequest(
"Too many redirects while fetching OCR document URL".to_string(),
))
}
fn enforce_download_size(content_length: u64, max_bytes: u64) -> CoreResult<()> {
if max_bytes == 0 {
return Err(CoreError::InvalidRequest(
"OCR document URL download is disabled (MAX_IMAGE_URL_DOWNLOAD_SIZE_MB=0)".to_string(),
));
}
if content_length > max_bytes {
let size_mb = content_length as f64 / (1024.0 * 1024.0);
let max_size_mb = max_bytes as f64 / (1024.0 * 1024.0);
return Err(CoreError::InvalidRequest(format!(
"OCR document size ({size_mb:.2}MB) exceeds maximum allowed size ({max_size_mb:.2}MB)"
)));
}
Ok(())
}
async fn read_response_with_limit(
mut response: reqwest::Response,
max_bytes: u64,
) -> CoreResult<Vec<u8>> {
enforce_download_size(response.content_length().unwrap_or(0), max_bytes)?;
let mut bytes = Vec::new();
let mut bytes_downloaded: u64 = 0;
while let Some(chunk) = response
.chunk()
.await
.map_err(|_| document_fetch_network_error())?
{
bytes_downloaded += chunk.len() as u64;
enforce_download_size(bytes_downloaded, max_bytes)?;
bytes.extend_from_slice(&chunk);
}
Ok(bytes)
}
pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreResult<Value> {
let Some((field, url)) = document_url_field(&document)? else {
return Ok(document);
};
if url.starts_with("data:") {
return Ok(document);
}
parse_fetchable_url(url)?;
let (_final_url, response) = safe_get_document_url(url).await?;
let status = response.status();
if !status.is_success() {
return Err(CoreError::Http {
status: status.as_u16(),
body: read_error_body(response).await,
});
}
let content_type = response
.headers()
.get(reqwest::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.split(';').next())
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("application/octet-stream")
.to_string();
let bytes = read_response_with_limit(response, max_document_download_bytes()).await?;
let data_uri = format!(
"data:{content_type};base64,{}",
BASE64_STANDARD.encode(bytes)
);
let object = document
.as_object()
.ok_or_else(|| CoreError::InvalidRequest("OCR document must be an object".to_string()))?;
let transformed: Map<String, Value> = object
.iter()
.map(|(key, value)| {
if key == field {
(key.clone(), Value::String(data_uri.clone()))
} else {
(key.clone(), value.clone())
}
})
.collect();
Ok(Value::Object(transformed))
}
#[cfg(test)]
mod tests {
use super::super::ssrf::{blocked_url_error, resolve_validated};
use super::super::test_support::{http_response, spawn_counting_server};
use super::*;
use serde_json::json;
#[tokio::test]
async fn loopback_fetch_rejected_without_connecting() {
let server = spawn_counting_server(http_response(
"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n",
b"hi",
))
.await;
let url = format!("http://127.0.0.1:{}/doc.png", server.addr.port());
let error = convert_document_url_to_data_uri(json!({
"type": "image_url",
"image_url": url,
}))
.await
.unwrap_err();
assert!(
matches!(&error, CoreError::InvalidRequest(message) if message.contains("SSRF protection")),
"got {error:?}"
);
assert_eq!(server.connection_count(), 0);
}
#[tokio::test]
async fn mapped_ipv6_loopback_fetch_rejected_without_connecting() {
let server = spawn_counting_server(http_response(
"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n",
b"hi",
))
.await;
let url = format!("http://[::ffff:127.0.0.1]:{}/doc.png", server.addr.port());
let error = convert_document_url_to_data_uri(json!({
"type": "image_url",
"image_url": url,
}))
.await
.unwrap_err();
assert!(
matches!(&error, CoreError::InvalidRequest(message) if message.contains("SSRF protection")),
"got {error:?}"
);
assert_eq!(server.connection_count(), 0);
}
#[tokio::test]
async fn unsupported_scheme_rejected_without_fetch() {
for raw in ["file:///etc/passwd", "gopher://8.8.8.8/x"] {
let error = convert_document_url_to_data_uri(json!({
"type": "image_url",
"image_url": raw,
}))
.await
.unwrap_err();
assert!(
matches!(&error, CoreError::InvalidRequest(message) if message.contains("SSRF protection")),
"{raw} got {error:?}"
);
}
}
#[tokio::test]
async fn redirect_to_loopback_rejected_without_connecting_to_target() {
let target = spawn_counting_server(http_response(
"HTTP/1.1 200 OK\r\nContent-Length: 6\r\n\r\n",
b"secret",
))
.await;
let location = format!("http://127.0.0.1:{}/internal", target.addr.port());
let redirector = spawn_counting_server(
format!("HTTP/1.1 302 Found\r\nLocation: {location}\r\nContent-Length: 0\r\n\r\n")
.into_bytes(),
)
.await;
let redirector_addr = redirector.addr;
let redirector_port = redirector_addr.port();
let start_url = format!("http://127.0.0.1:{redirector_port}/doc.png");
let pin = move |candidate: Url| async move {
if candidate.port() == Some(redirector_port) {
return Ok(vec![redirector_addr]);
}
resolve_validated(&candidate).await
};
let error = fetch_with_redirects(&start_url, pin).await.unwrap_err();
assert!(
matches!(&error, CoreError::InvalidRequest(message) if message.contains("SSRF protection")),
"got {error:?}"
);
assert!(redirector.connection_count() >= 1);
assert_eq!(target.connection_count(), 0);
}
#[tokio::test]
async fn download_exceeding_cap_without_content_length_is_rejected() {
let server = spawn_counting_server(http_response(
"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n",
&vec![b'a'; 4096],
))
.await;
let url_str = format!("http://127.0.0.1:{}/big", server.addr.port());
let response = reqwest::Client::new().get(&url_str).send().await.unwrap();
assert!(response.content_length().is_none());
let error = read_response_with_limit(response, 1024).await.unwrap_err();
assert!(
matches!(&error, CoreError::InvalidRequest(message) if message.contains("exceeds maximum")),
"got {error:?}"
);
}
#[tokio::test]
async fn download_within_cap_without_content_length_succeeds() {
let server = spawn_counting_server(http_response(
"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n",
&vec![b'a'; 512],
))
.await;
let url_str = format!("http://127.0.0.1:{}/small", server.addr.port());
let response = reqwest::Client::new().get(&url_str).send().await.unwrap();
assert!(response.content_length().is_none());
let bytes = read_response_with_limit(response, 1024).await.unwrap();
assert_eq!(bytes.len(), 512);
}
#[tokio::test]
async fn domain_resolving_to_blocked_address_is_rejected_without_connecting() {
let server = spawn_counting_server(http_response(
"HTTP/1.1 200 OK\r\nContent-Length: 6\r\n\r\n",
b"secret",
))
.await;
let start_url = format!("http://localhost:{}/doc", server.addr.port());
let error = fetch_with_redirects(&start_url, |candidate| async move {
resolve_validated(&candidate).await
})
.await
.unwrap_err();
assert!(
matches!(&error, CoreError::InvalidRequest(message) if message.contains("SSRF protection")),
"a domain whose DNS answer is a blocked address must be rejected, got {error:?}"
);
assert_eq!(server.connection_count(), 0);
}
#[tokio::test]
async fn request_connects_only_to_pinned_address_without_a_second_dns_lookup() {
let server = spawn_counting_server(http_response(
"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n",
b"ok",
))
.await;
let pinned_addr = server.addr;
let url = format!("http://pinned.invalid:{}/doc", pinned_addr.port());
let (_final_url, response) =
fetch_with_redirects(&url, move |_candidate| async move { Ok(vec![pinned_addr]) })
.await
.expect("request must connect via the pinned address");
assert_eq!(response.status(), reqwest::StatusCode::OK);
assert_eq!(server.connection_count(), 1);
}
#[tokio::test]
async fn redirect_across_ports_repins_without_reusing_stale_client() {
let second = spawn_counting_server(http_response(
"HTTP/1.1 200 OK\r\nContent-Length: 6\r\n\r\n",
b"second",
))
.await;
let second_addr = second.addr;
let first = spawn_counting_server(http_response(
&format!(
"HTTP/1.1 302 Found\r\nLocation: http://pinned.invalid:{}/next\r\nContent-Length: 0\r\n\r\n",
second_addr.port()
),
b"",
))
.await;
let first_addr = first.addr;
let start_url = format!("http://pinned.invalid:{}/start", first_addr.port());
let (final_url, response) = fetch_with_redirects(&start_url, move |candidate| async move {
match candidate.port() {
Some(port) if port == first_addr.port() => Ok(vec![first_addr]),
Some(port) if port == second_addr.port() => Ok(vec![second_addr]),
_ => Err(blocked_url_error()),
}
})
.await
.expect("redirect across ports must repin to the second server");
assert_eq!(response.status(), reqwest::StatusCode::OK);
assert_eq!(final_url.port(), Some(second_addr.port()));
assert_eq!(first.connection_count(), 1);
assert_eq!(second.connection_count(), 1);
}
#[tokio::test]
async fn convert_document_url_rejects_loopback_fetch() {
let error = convert_document_url_to_data_uri(json!({
"type": "image_url",
"image_url": "http://127.0.0.1/image.png"
}))
.await
.unwrap_err();
assert!(matches!(
error,
CoreError::InvalidRequest(message)
if message.contains("SSRF protection")
));
}
#[tokio::test]
async fn convert_document_url_leaves_data_uri_untouched() {
let document = json!({
"type": "image_url",
"image_url": "data:image/png;base64,abcd"
});
let transformed = convert_document_url_to_data_uri(document.clone())
.await
.unwrap();
assert_eq!(transformed, document);
}
}

View file

@ -0,0 +1,149 @@
use std::time::Duration;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use base64::Engine;
use litellm_core::error::CoreError;
use litellm_core::CoreResult;
use reqwest::Url;
use serde_json::{Map, Value};
use super::common_utils::{document_url_field, read_error_body};
use super::http_client;
fn decode_reducto_data_uri(source_url: &str) -> CoreResult<(Vec<u8>, String)> {
let (header, encoded) = source_url.split_once(',').ok_or_else(|| {
CoreError::InvalidRequest("Invalid Reducto data URI provided.".to_string())
})?;
if !header.contains(";base64") {
return Err(CoreError::InvalidRequest(
"Reducto only supports base64-encoded data URIs.".to_string(),
));
}
let mime = header
.strip_prefix("data:")
.and_then(|value| value.split(';').next())
.filter(|value| !value.is_empty())
.unwrap_or("application/octet-stream")
.to_string();
let bytes = BASE64_STANDARD.decode(encoded).map_err(|_| {
CoreError::InvalidRequest("Invalid Reducto base64 payload provided.".to_string())
})?;
Ok((bytes, mime))
}
fn reducto_upload_url(parse_url: &str) -> CoreResult<Url> {
let mut url = Url::parse(parse_url)
.map_err(|err| CoreError::InvalidRequest(format!("invalid Reducto parse URL: {err}")))?;
let path = url.path().trim_end_matches('/');
let base_path = path
.strip_suffix("/parse")
.unwrap_or(path)
.trim_end_matches('/');
url.set_path(&format!("{base_path}/upload"));
url.set_query(None);
Ok(url)
}
fn reducto_auth_headers(headers: &[(String, String)]) -> Vec<(String, String)> {
headers
.iter()
.filter(|(key, _)| key.eq_ignore_ascii_case("authorization"))
.cloned()
.collect()
}
async fn upload_reducto_bytes(
bytes: Vec<u8>,
mime: String,
parse_url: &str,
headers: &[(String, String)],
timeout: Option<Duration>,
) -> CoreResult<String> {
let upload_url = reducto_upload_url(parse_url)?;
let part = reqwest::multipart::Part::bytes(bytes)
.file_name("document")
.mime_str(&mime)
.map_err(|err| {
CoreError::InvalidRequest(format!("invalid Reducto upload MIME type: {err}"))
})?;
let form = reqwest::multipart::Form::new().part("file", part);
let mut request_builder = http_client().post(upload_url).multipart(form);
for (key, value) in reducto_auth_headers(headers) {
request_builder = request_builder.header(key, value);
}
if let Some(duration) = timeout {
request_builder = request_builder.timeout(duration);
}
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
let status = response.status();
if !status.is_success() {
return Err(CoreError::Http {
status: status.as_u16(),
body: read_error_body(response).await,
});
}
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
let response_json: Value = serde_json::from_str(&text).map_err(|err| {
CoreError::InvalidResponse(format!("invalid Reducto upload response JSON: {err}"))
})?;
response_json
.get("file_id")
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.map(str::to_string)
.ok_or_else(|| {
CoreError::InvalidResponse("Reducto /upload returned 200 without a file_id".to_string())
})
}
pub(super) async fn upload_reducto_document(
document: Value,
parse_url: &str,
headers: &[(String, String)],
timeout: Option<Duration>,
) -> CoreResult<Value> {
let Some((field, source_url)) = document_url_field(&document)? else {
return Err(CoreError::InvalidRequest(
"Reducto expected OCR preprocessing to produce document_url or image_url".to_string(),
));
};
if source_url.starts_with("reducto://") {
return Ok(document);
}
if source_url.starts_with("http://") || source_url.starts_with("https://") {
return Err(CoreError::InvalidRequest(
"Reducto requires type='file' (auto-uploaded) or a reducto:// id. Plain http(s) URLs are not supported; upload the file first."
.to_string(),
));
}
if !source_url.starts_with("data:") {
return Err(CoreError::InvalidRequest(
"Reducto requires a reducto:// id or a base64 data URI after OCR preprocessing."
.to_string(),
));
}
let (bytes, mime) = decode_reducto_data_uri(source_url)?;
let file_id = upload_reducto_bytes(bytes, mime, parse_url, headers, timeout).await?;
let object = document
.as_object()
.ok_or_else(|| CoreError::InvalidRequest("OCR document must be an object".to_string()))?;
let transformed: Map<String, Value> = object
.iter()
.map(|(key, value)| {
if key == field {
(key.clone(), Value::String(file_id.clone()))
} else {
(key.clone(), value.clone())
}
})
.collect();
Ok(Value::Object(transformed))
}

View file

@ -0,0 +1,346 @@
use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::{Arc, Mutex, OnceLock, PoisonError};
use std::time::Duration;
use litellm_core::error::CoreError;
use litellm_core::CoreResult;
use reqwest::dns::{Addrs, Name, Resolve, Resolving};
use reqwest::Url;
use url::Host;
use crate::constants::{DOCUMENT_FETCH_TIMEOUT_SECS, MAX_DOCUMENT_CLIENTS};
const BLOCKED_IPV6_CIDRS: &[(Ipv6Addr, u32)] = &[
(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 0), 96),
(Ipv6Addr::new(0x64, 0xff9b, 0, 0, 0, 0, 0, 0), 96),
(Ipv6Addr::new(0x64, 0xff9b, 1, 0, 0, 0, 0, 0), 48),
(Ipv6Addr::new(0x100, 0, 0, 0, 0, 0, 0, 0), 64),
(Ipv6Addr::new(0x2001, 0, 0, 0, 0, 0, 0, 0), 23),
(Ipv6Addr::new(0x2001, 0xdb8, 0, 0, 0, 0, 0, 0), 32),
(Ipv6Addr::new(0x2002, 0, 0, 0, 0, 0, 0, 0), 16),
(Ipv6Addr::new(0x3fff, 0, 0, 0, 0, 0, 0, 0), 20),
(Ipv6Addr::new(0x5f00, 0, 0, 0, 0, 0, 0, 0), 16),
(Ipv6Addr::new(0xfc00, 0, 0, 0, 0, 0, 0, 0), 7),
(Ipv6Addr::new(0xfe80, 0, 0, 0, 0, 0, 0, 0), 10),
(Ipv6Addr::new(0xfec0, 0, 0, 0, 0, 0, 0, 0), 10),
(Ipv6Addr::new(0xff00, 0, 0, 0, 0, 0, 0, 0), 8),
];
const BLOCKED_IPV4_CIDRS: &[(Ipv4Addr, u32)] = &[
(Ipv4Addr::new(0, 0, 0, 0), 8),
(Ipv4Addr::new(100, 64, 0, 0), 10),
(Ipv4Addr::new(192, 0, 0, 0), 24),
(Ipv4Addr::new(192, 0, 2, 0), 24),
(Ipv4Addr::new(192, 88, 99, 0), 24),
(Ipv4Addr::new(198, 18, 0, 0), 15),
(Ipv4Addr::new(198, 51, 100, 0), 24),
(Ipv4Addr::new(203, 0, 113, 0), 24),
(Ipv4Addr::new(240, 0, 0, 0), 4),
];
pub(super) fn blocked_url_error() -> CoreError {
CoreError::InvalidRequest("OCR document URL rejected by SSRF protection".to_string())
}
fn ipv4_in_cidr(ip: Ipv4Addr, network: Ipv4Addr, prefix_length: u32) -> bool {
let mask = if prefix_length == 0 {
0
} else {
u32::MAX << (32 - prefix_length)
};
u32::from(ip) & mask == u32::from(network) & mask
}
fn ipv6_in_cidr(ip: Ipv6Addr, network: Ipv6Addr, prefix_length: u32) -> bool {
let mask = if prefix_length == 0 {
0
} else {
u128::MAX << (128 - prefix_length)
};
u128::from(ip) & mask == u128::from(network) & mask
}
pub(super) fn is_blocked_ip(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(ip) => {
ip.is_private()
|| ip.is_loopback()
|| ip.is_link_local()
|| ip.is_broadcast()
|| ip.is_multicast()
|| ip.is_unspecified()
|| BLOCKED_IPV4_CIDRS
.iter()
.any(|(network, prefix)| ipv4_in_cidr(ip, *network, *prefix))
}
IpAddr::V6(ip) => {
if let Some(v4) = ip.to_ipv4_mapped() {
return is_blocked_ip(IpAddr::V4(v4));
}
BLOCKED_IPV6_CIDRS
.iter()
.any(|(network, prefix)| ipv6_in_cidr(ip, *network, *prefix))
}
}
}
pub(super) fn parse_fetchable_url(raw: &str) -> CoreResult<Url> {
let url = Url::parse(raw).map_err(|_| blocked_url_error())?;
ensure_allowed_url(&url)?;
Ok(url)
}
fn ensure_allowed_url(url: &Url) -> CoreResult<()> {
if !matches!(url.scheme(), "http" | "https") {
return Err(blocked_url_error());
}
if !url.username().is_empty() || url.password().is_some() {
return Err(blocked_url_error());
}
Ok(())
}
pub(super) async fn resolve_validated(url: &Url) -> CoreResult<Vec<SocketAddr>> {
ensure_allowed_url(url)?;
let port = url.port_or_known_default().ok_or_else(blocked_url_error)?;
let addresses: Vec<SocketAddr> = match url.host() {
Some(Host::Ipv4(ip)) => vec![SocketAddr::from((ip, port))],
Some(Host::Ipv6(ip)) => vec![SocketAddr::from((ip, port))],
Some(Host::Domain(domain)) => tokio::net::lookup_host((domain, port))
.await
.map_err(|_| blocked_url_error())?
.collect(),
None => return Err(blocked_url_error()),
};
if addresses.is_empty() || addresses.iter().any(|address| is_blocked_ip(address.ip())) {
return Err(blocked_url_error());
}
Ok(addresses)
}
type PinnedAddrs = Arc<Mutex<HashMap<String, Vec<SocketAddr>>>>;
#[derive(Debug, Clone)]
struct PinnedResolver {
pins: PinnedAddrs,
}
impl Resolve for PinnedResolver {
fn resolve(&self, name: Name) -> Resolving {
let pins = self.pins.clone();
Box::pin(async move {
let pinned = pins
.lock()
.unwrap_or_else(PoisonError::into_inner)
.get(name.as_str())
.cloned();
match pinned {
Some(addresses) if !addresses.is_empty() => {
Ok(Box::new(addresses.into_iter()) as Addrs)
}
_ => Err(Box::<dyn std::error::Error + Send + Sync>::from(
"OCR document host was not pinned to a validated address",
)),
}
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct DocumentClientKey {
scheme: String,
host: String,
port: u16,
addresses: Vec<SocketAddr>,
}
fn document_client_cache() -> &'static Mutex<HashMap<DocumentClientKey, reqwest::Client>> {
static CACHE: OnceLock<Mutex<HashMap<DocumentClientKey, reqwest::Client>>> = OnceLock::new();
CACHE.get_or_init(|| Mutex::new(HashMap::new()))
}
pub(super) fn pinned_client(url: &Url, addresses: &[SocketAddr]) -> CoreResult<reqwest::Client> {
let host = url.host_str().ok_or_else(blocked_url_error)?.to_owned();
let port = url.port_or_known_default().ok_or_else(blocked_url_error)?;
let mut sorted = addresses.to_vec();
sorted.sort();
let key = DocumentClientKey {
scheme: url.scheme().to_owned(),
host: host.clone(),
port,
addresses: sorted.clone(),
};
let cache = document_client_cache();
if let Some(client) = cache
.lock()
.unwrap_or_else(PoisonError::into_inner)
.get(&key)
.cloned()
{
return Ok(client);
}
let mut pins = HashMap::new();
pins.insert(host, sorted);
let client = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.timeout(Duration::from_secs(DOCUMENT_FETCH_TIMEOUT_SECS))
.dns_resolver(Arc::new(PinnedResolver {
pins: Arc::new(Mutex::new(pins)),
}))
.build()
.map_err(|err| CoreError::Network(err.to_string()))?;
let mut guard = cache.lock().unwrap_or_else(PoisonError::into_inner);
if guard.len() >= MAX_DOCUMENT_CLIENTS {
if let Some(evicted) = guard.keys().next().cloned() {
guard.remove(&evicted);
}
}
guard.insert(key, client.clone());
Ok(client)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn validate_rejects_all_special_use_ranges() {
let blocked = [
"http://127.0.0.1/x",
"http://10.0.0.1/x",
"http://172.16.0.1/x",
"http://192.168.1.1/x",
"http://169.254.169.254/x",
"http://100.64.0.1/x",
"http://198.18.0.1/x",
"http://192.0.2.1/x",
"http://[::1]/x",
"http://[fd00::1]/x",
"http://[fe80::1]/x",
"http://[fec0::1]/x",
"http://[ff02::1]/x",
"http://[::ffff:169.254.169.254]/x",
"http://[::ffff:10.0.0.1]/x",
"http://[::ffff:100.64.0.1]/x",
"http://[::ffff:198.18.0.1]/x",
"http://[::1.2.3.4]/x",
"http://[2001:db8::1]/x",
"http://[100::1]/x",
"http://[2001::1]/x",
"http://[2001:2::1]/x",
"http://[2001:10::1]/x",
"http://[2001:20::1]/x",
"http://[2002:c0a8:101::1]/x",
"http://[2002:a9fe:a9fe::1]/x",
"http://[2002:808:808::1]/x",
"http://[3fff::1]/x",
"http://[5f00::1]/x",
"http://[64:ff9b:1::1]/x",
"http://[64:ff9b::192.168.1.1]/x",
"http://[64:ff9b::169.254.169.254]/x",
"http://[64:ff9b::8.8.8.8]/x",
"ftp://8.8.8.8/x",
"file:///etc/passwd",
"gopher://8.8.8.8/x",
"http://user:pass@8.8.8.8/x",
];
for raw in blocked {
let url = Url::parse(raw).unwrap();
let error = resolve_validated(&url).await.unwrap_err();
assert!(
matches!(&error, CoreError::InvalidRequest(message) if message.contains("SSRF protection")),
"{raw} should be rejected, got {error:?}"
);
}
let allowed = Url::parse("http://8.8.8.8/x").unwrap();
assert_eq!(
resolve_validated(&allowed).await.unwrap(),
vec![SocketAddr::from(([8, 8, 8, 8], 80))]
);
let allowed_v6 = [
"http://[2606:4700:4700::1111]/x",
"http://[::ffff:8.8.8.8]/x",
];
for raw in allowed_v6 {
let url = Url::parse(raw).unwrap();
assert!(
resolve_validated(&url).await.is_ok(),
"{raw} should be allowed"
);
}
}
#[test]
fn blocks_private_and_metadata_ips() {
assert!(is_blocked_ip("127.0.0.1".parse().unwrap()));
assert!(is_blocked_ip("10.0.0.1".parse().unwrap()));
assert!(is_blocked_ip("169.254.169.254".parse().unwrap()));
assert!(is_blocked_ip("100.64.0.1".parse().unwrap()));
assert!(is_blocked_ip("198.18.0.1".parse().unwrap()));
assert!(is_blocked_ip("192.0.2.1".parse().unwrap()));
assert!(is_blocked_ip("::1".parse().unwrap()));
assert!(is_blocked_ip("fd00::1".parse().unwrap()));
assert!(is_blocked_ip("fe80::1".parse().unwrap()));
assert!(is_blocked_ip("fec0::1".parse().unwrap()));
assert!(is_blocked_ip("2001:db8::1".parse().unwrap()));
assert!(is_blocked_ip("2002:c0a8:101::1".parse().unwrap()));
assert!(is_blocked_ip("64:ff9b::8.8.8.8".parse().unwrap()));
assert!(is_blocked_ip("3fff::1".parse().unwrap()));
assert!(is_blocked_ip("5f00::1".parse().unwrap()));
assert!(is_blocked_ip("::ffff:169.254.169.254".parse().unwrap()));
assert!(is_blocked_ip("::ffff:10.0.0.1".parse().unwrap()));
assert!(is_blocked_ip("::ffff:100.64.0.1".parse().unwrap()));
assert!(is_blocked_ip("::ffff:198.18.0.1".parse().unwrap()));
assert!(is_blocked_ip("::1.2.3.4".parse().unwrap()));
assert!(!is_blocked_ip("8.8.8.8".parse().unwrap()));
assert!(!is_blocked_ip("2606:4700:4700::1111".parse().unwrap()));
assert!(!is_blocked_ip("::ffff:8.8.8.8".parse().unwrap()));
}
#[test]
fn parse_fetchable_url_rejects_non_http_and_userinfo() {
assert!(parse_fetchable_url("file:///etc/passwd").is_err());
assert!(parse_fetchable_url("gopher://8.8.8.8/x").is_err());
assert!(parse_fetchable_url("http://user:pass@example.com/x").is_err());
assert_eq!(
parse_fetchable_url("HTTP://example.com/x")
.unwrap()
.scheme(),
"http"
);
}
#[test]
fn document_client_key_distinguishes_scheme_port_and_addresses() {
let addr_a = SocketAddr::from(([203, 0, 113, 1], 443));
let addr_b = SocketAddr::from(([203, 0, 113, 2], 443));
let base = DocumentClientKey {
scheme: "https".to_string(),
host: "docs.example".to_string(),
port: 443,
addresses: vec![addr_a],
};
let other_scheme = DocumentClientKey {
scheme: "http".to_string(),
..base.clone()
};
let other_port = DocumentClientKey {
port: 8443,
..base.clone()
};
let other_addrs = DocumentClientKey {
addresses: vec![addr_a, addr_b],
..base.clone()
};
assert_ne!(base, other_scheme);
assert_ne!(base, other_port);
assert_ne!(base, other_addrs);
assert_eq!(base, base.clone());
}
}

View file

@ -0,0 +1,53 @@
use std::net::SocketAddr;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use tokio::task::JoinHandle;
pub(super) struct CountingServer {
pub(super) addr: SocketAddr,
connections: Arc<AtomicUsize>,
handle: JoinHandle<()>,
}
impl CountingServer {
pub(super) fn connection_count(&self) -> usize {
self.connections.load(Ordering::SeqCst)
}
}
impl Drop for CountingServer {
fn drop(&mut self) {
self.handle.abort();
}
}
pub(super) async fn spawn_counting_server(response: Vec<u8>) -> CountingServer {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let connections = Arc::new(AtomicUsize::new(0));
let counter = connections.clone();
let handle = tokio::spawn(async move {
loop {
let Ok((mut socket, _)) = listener.accept().await else {
return;
};
counter.fetch_add(1, Ordering::SeqCst);
let mut discard = [0u8; 2048];
let _ = socket.read(&mut discard).await;
let _ = socket.write_all(&response).await;
let _ = socket.flush().await;
}
});
CountingServer {
addr,
connections,
handle,
}
}
pub(super) fn http_response(headers: &str, body: &[u8]) -> Vec<u8> {
[headers.as_bytes(), body].concat()
}

View file

@ -27,7 +27,6 @@ pub mod state;
// Realtime request logging. Only the server serves realtime, so these are
// `server`-gated; `io::realtime` exposes the generic `observe` hook while the
// collector and callback fan-out live here.
#[cfg(feature = "server")]
mod constants;
#[cfg(feature = "server")]
pub mod integrations;