From d321c60252cb8f615768f2fb68d72f4c1c2c55e0 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 17 Jul 2026 00:41:05 +0000 Subject: [PATCH] refactor(ocr-security): split OCR host modules and harden URL/IPv6/poll validation Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- .../crates/ai-gateway/src/constants.rs | 27 + litellm-rust/crates/ai-gateway/src/io/ocr.rs | 48 +- .../ai-gateway/src/io/ocr/azure_poll.rs | 336 ++++++ .../ai-gateway/src/io/ocr/common_utils.rs | 1040 +---------------- .../ai-gateway/src/io/ocr/document_fetch.rs | 401 +++++++ .../crates/ai-gateway/src/io/ocr/reducto.rs | 149 +++ .../crates/ai-gateway/src/io/ocr/ssrf.rs | 346 ++++++ .../ai-gateway/src/io/ocr/test_support.rs | 53 + litellm-rust/crates/ai-gateway/src/lib.rs | 1 - 9 files changed, 1339 insertions(+), 1062 deletions(-) create mode 100644 litellm-rust/crates/ai-gateway/src/io/ocr/azure_poll.rs create mode 100644 litellm-rust/crates/ai-gateway/src/io/ocr/document_fetch.rs create mode 100644 litellm-rust/crates/ai-gateway/src/io/ocr/reducto.rs create mode 100644 litellm-rust/crates/ai-gateway/src/io/ocr/ssrf.rs create mode 100644 litellm-rust/crates/ai-gateway/src/io/ocr/test_support.rs diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 3116a4c9932..9d0fbe81d0d 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -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"; diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr.rs b/litellm-rust/crates/ai-gateway/src/io/ocr.rs index 5dceb637808..44d1544485f 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr.rs @@ -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] diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr/azure_poll.rs b/litellm-rust/crates/ai-gateway/src/io/ocr/azure_poll.rs new file mode 100644 index 00000000000..4619508ec13 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/azure_poll.rs @@ -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::().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, +) -> CoreResult { + 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( + operation_url: &str, + original_url: &str, + headers: &[(String, String)], + timeout: Option, + resolve: P, +) -> CoreResult +where + P: Fn(Url) -> Fut, + Fut: std::future::Future>>, +{ + 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" + ); + } +} 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 19bee7c6eaa..65eec1b54db 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 @@ -1,17 +1,7 @@ -use std::collections::HashMap; -use std::net::{IpAddr, Ipv4Addr, SocketAddr}; -use std::sync::{Arc, Mutex, OnceLock, PoisonError}; -use std::time::{Duration, Instant}; - -use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; -use base64::Engine; use litellm_core::error::CoreError; use litellm_core::ocr::transformation::OcrProviderConfig; use litellm_core::CoreResult; -use reqwest::dns::{Addrs, Name, Resolve, Resolving}; -use reqwest::Url; use serde_json::{Map, Value}; -use url::Host; use litellm_core::providers::azure_ai::ocr::transformation::{ AZURE_AI_OCR_CONFIG, AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG, @@ -25,30 +15,22 @@ use litellm_core::providers::vertex_ai::ocr::transformation::{ VERTEX_AI_DEEPSEEK_OCR_CONFIG, VERTEX_AI_OCR_CONFIG, }; -use super::http_client; - -const ERROR_BODY_MAX_CHARS: usize = 256; -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; -const ERROR_BODY_MAX_BYTES: usize = ERROR_BODY_MAX_CHARS * 4; -const MAX_DOCUMENT_CLIENTS: usize = 128; -const DOCUMENT_FETCH_TIMEOUT_SECS: u64 = 600; +use crate::constants::{OCR_ERROR_BODY_MAX_BYTES, OCR_ERROR_BODY_MAX_CHARS}; pub(super) fn truncate_error_body(body: &str) -> String { - if body.chars().count() <= ERROR_BODY_MAX_CHARS { + if body.chars().count() <= OCR_ERROR_BODY_MAX_CHARS { return body.to_string(); } - let truncated: String = body.chars().take(ERROR_BODY_MAX_CHARS).collect(); + let truncated: String = body.chars().take(OCR_ERROR_BODY_MAX_CHARS).collect(); format!("{truncated}... (truncated)") } pub(super) async fn read_error_body(mut response: reqwest::Response) -> String { let mut collected: Vec = Vec::new(); - while collected.len() < ERROR_BODY_MAX_BYTES { + while collected.len() < OCR_ERROR_BODY_MAX_BYTES { match response.chunk().await { Ok(Some(chunk)) => { - let take = (ERROR_BODY_MAX_BYTES - collected.len()).min(chunk.len()); + let take = (OCR_ERROR_BODY_MAX_BYTES - collected.len()).min(chunk.len()); collected.extend_from_slice(&chunk[..take]); if take < chunk.len() { break; @@ -110,7 +92,7 @@ pub(super) fn has_header(headers: &[(String, String)], name: &str) -> bool { .any(|(key, _)| key.eq_ignore_ascii_case(name)) } -fn document_url_field(document: &Value) -> CoreResult> { +pub(super) fn document_url_field(document: &Value) -> CoreResult> { let Some(object) = document.as_object() else { return Ok(None); }; @@ -128,860 +110,10 @@ fn document_url_field(document: &Value) -> CoreResult> { Ok(Some((field, url))) } -fn is_url_requiring_fetch(url: &str) -> bool { - !url.starts_with("data:") && (url.starts_with("http://") || url.starts_with("https://")) -} - -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::().ok()) - .unwrap_or(DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB); - (max_size_mb.max(0.0) * 1024.0 * 1024.0) as u64 -} - -fn ipv4_in_cidr(ip: Ipv4Addr, network: Ipv4Addr, prefix_length: u32) -> bool { - let mask = u32::MAX << (32 - prefix_length); - u32::from(ip) & mask == u32::from(network) & mask -} - -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() - || ipv4_in_cidr(ip, Ipv4Addr::new(0, 0, 0, 0), 8) - || ipv4_in_cidr(ip, Ipv4Addr::new(100, 64, 0, 0), 10) - || ipv4_in_cidr(ip, Ipv4Addr::new(192, 0, 0, 0), 24) - || ipv4_in_cidr(ip, Ipv4Addr::new(192, 0, 2, 0), 24) - || ipv4_in_cidr(ip, Ipv4Addr::new(192, 88, 99, 0), 24) - || ipv4_in_cidr(ip, Ipv4Addr::new(198, 18, 0, 0), 15) - || ipv4_in_cidr(ip, Ipv4Addr::new(198, 51, 100, 0), 24) - || ipv4_in_cidr(ip, Ipv4Addr::new(203, 0, 113, 0), 24) - || ipv4_in_cidr(ip, Ipv4Addr::new(240, 0, 0, 0), 4) - } - IpAddr::V6(ip) => { - let segments = ip.segments(); - let first_segment = segments[0]; - let is_unique_local = (first_segment & 0xfe00) == 0xfc00; - let is_link_local = (first_segment & 0xffc0) == 0xfe80; - let is_site_local = (first_segment & 0xffc0) == 0xfec0; - let is_documentation = first_segment == 0x2001 && segments[1] == 0x0db8; - let is_discard = - first_segment == 0x0100 && segments[1] == 0 && segments[2] == 0 && segments[3] == 0; - let is_teredo = first_segment == 0x2001 && segments[1] == 0x0000; - let is_benchmarking = - first_segment == 0x2001 && segments[1] == 0x0002 && segments[2] == 0x0000; - let is_orchid = first_segment == 0x2001 && (segments[1] & 0xfff0) == 0x0010; - let is_orchidv2 = first_segment == 0x2001 && (segments[1] & 0xfff0) == 0x0020; - let is_nat64_wellknown = first_segment == 0x0064 - && segments[1] == 0xff9b - && segments[2] == 0 - && segments[3] == 0 - && segments[4] == 0 - && segments[5] == 0; - let is_nat64_local = - first_segment == 0x0064 && segments[1] == 0xff9b && segments[2] == 1; - let is_6to4_embedded_blocked = - first_segment == 0x2002 && embedded_ipv4_blocked(segments[1], segments[2]); - let is_nat64_embedded_blocked = - is_nat64_wellknown && embedded_ipv4_blocked(segments[6], segments[7]); - ip.is_loopback() - || ip.is_unspecified() - || ip.is_multicast() - || is_unique_local - || is_link_local - || is_site_local - || is_documentation - || is_discard - || is_teredo - || is_benchmarking - || is_orchid - || is_orchidv2 - || is_nat64_local - || is_nat64_embedded_blocked - || is_6to4_embedded_blocked - || ip - .to_ipv4_mapped() - .or_else(|| ip.to_ipv4()) - .map(|v4| is_blocked_ip(IpAddr::V4(v4))) - .unwrap_or(false) - } - } -} - -fn embedded_ipv4_blocked(hi: u16, lo: u16) -> bool { - let v4 = Ipv4Addr::new( - (hi >> 8) as u8, - (hi & 0xff) as u8, - (lo >> 8) as u8, - (lo & 0xff) as u8, - ); - is_blocked_ip(IpAddr::V4(v4)) -} - -fn blocked_url_error() -> CoreError { - CoreError::InvalidRequest("OCR document URL rejected by SSRF protection".to_string()) -} - -async fn pin_validated_url(url: &Url) -> CoreResult> { - if !matches!(url.scheme(), "http" | "https") { - return Err(blocked_url_error()); - } - let port = url.port_or_known_default().ok_or_else(blocked_url_error)?; - let addresses: Vec = 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>>>; - -#[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::::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, -} - -fn document_client_cache() -> &'static Mutex> { - static CACHE: OnceLock>> = OnceLock::new(); - CACHE.get_or_init(|| Mutex::new(HashMap::new())) -} - -fn document_fetch_client(url: &Url, addresses: &[SocketAddr]) -> CoreResult { - 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) -} - -fn redirect_location(response: &reqwest::Response, url: &Url) -> CoreResult { - 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(|err| CoreError::InvalidResponse(format!("invalid OCR document redirect: {err}"))) -} - -async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response)> { - fetch_with_redirects(url, |candidate| async move { - pin_validated_url(&candidate).await - }) - .await -} - -async fn fetch_with_redirects(url: &str, resolve: P) -> CoreResult<(Url, reqwest::Response)> -where - P: Fn(Url) -> Fut, - Fut: std::future::Future>>, -{ - let mut current_url = Url::parse(url) - .map_err(|err| CoreError::InvalidRequest(format!("invalid OCR document URL: {err}")))?; - - for _ in 0..MAX_SAFE_FETCH_REDIRECTS { - let addresses = resolve(current_url.clone()).await?; - let client = document_fetch_client(¤t_url, &addresses)?; - let response = client - .get(current_url.clone()) - .send() - .await - .map_err(|err| CoreError::Network(err.to_string()))?; - if !response.status().is_redirection() { - return Ok((current_url, response)); - } - current_url = redirect_location(&response, ¤t_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> { - if let Some(content_length) = response.content_length() { - enforce_download_size(content_length, max_bytes)?; - } else { - enforce_download_size(0, max_bytes)?; - } - - 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()))? - { - 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 { - let Some((field, url)) = document_url_field(&document)? else { - return Ok(document); - }; - if !is_url_requiring_fetch(url) { - return Ok(document); - } - - 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 mut transformed = document - .as_object() - .cloned() - .ok_or_else(|| CoreError::InvalidRequest("OCR document must be an object".to_string()))?; - transformed.insert(field.to_string(), Value::String(data_uri)); - Ok(Value::Object(transformed)) -} - -fn decode_reducto_data_uri(source_url: &str) -> CoreResult<(Vec, 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 { - 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, - mime: String, - parse_url: &str, - headers: &[(String, String)], - timeout: Option, -) -> CoreResult { - 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(format!( - "Reducto /upload returned 200 without a file_id; got payload={response_json}" - )) - }) -} - -pub(super) async fn upload_reducto_document( - document: Value, - parse_url: &str, - headers: &[(String, String)], - timeout: Option, -) -> CoreResult { - 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 mut transformed = document - .as_object() - .cloned() - .ok_or_else(|| CoreError::InvalidRequest("OCR document must be an object".to_string()))?; - transformed.insert(field.to_string(), Value::String(file_id)); - Ok(Value::Object(transformed)) -} - -fn same_origin(left: &str, right: &str) -> bool { - let Ok(left) = reqwest::Url::parse(left) else { - return false; - }; - let Ok(right) = reqwest::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::().ok()) - .unwrap_or(2) -} - -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}" - ))), - } -} - -pub(super) async fn poll_document_intelligence( - operation_url: &str, - original_url: &str, - headers: &[(String, String)], - timeout: Option, -) -> CoreResult { - if !same_origin(operation_url, original_url) { - return Err(CoreError::InvalidRequest( - "Azure Document Intelligence: rejected unsafe polling target".to_string(), - )); - } - - let start = Instant::now(); - let timeout = timeout.unwrap_or(Duration::from_secs( - AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS, - )); - loop { - if start.elapsed() > timeout { - return Err(CoreError::Network(format!( - "Azure Document Intelligence operation polling timed out after {} seconds", - timeout.as_secs() - ))); - } - - let mut request_builder = http_client().get(operation_url); - 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(|err| CoreError::Network(err.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(|err| CoreError::Network(err.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); - } - tokio::time::sleep(Duration::from_secs(retry_after)).await; - } -} - #[cfg(test)] mod tests { + use super::super::test_support::spawn_counting_server; use super::*; - use serde_json::json; - 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; - - struct CountingServer { - addr: SocketAddr, - connections: Arc, - handle: JoinHandle<()>, - } - - impl CountingServer { - fn connection_count(&self) -> usize { - self.connections.load(Ordering::SeqCst) - } - } - - impl Drop for CountingServer { - fn drop(&mut self) { - self.handle.abort(); - } - } - - async fn spawn_counting_server(response: Vec) -> 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, - } - } - - fn http_response(headers: &str, body: &[u8]) -> Vec { - [headers.as_bytes(), body].concat() - } - - #[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://[::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://[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://[64:ff9b:1::1]/x", - "http://[64:ff9b::192.168.1.1]/x", - "http://[64:ff9b::169.254.169.254]/x", - "ftp://8.8.8.8/x", - ]; - for raw in blocked { - let url = Url::parse(raw).unwrap(); - let error = pin_validated_url(&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!( - pin_validated_url(&allowed).await.unwrap(), - vec![SocketAddr::from(([8, 8, 8, 8], 80))] - ); - - let allowed_v6 = [ - "http://[2606:4700:4700::1111]/x", - "http://[64:ff9b::8.8.8.8]/x", - "http://[2002:808:808::1]/x", - ]; - for raw in allowed_v6 { - let url = Url::parse(raw).unwrap(); - assert!( - pin_validated_url(&url).await.is_ok(), - "{raw} should be allowed" - ); - } - } - - #[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 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]); - } - pin_validated_url(&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 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_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 operation_url = format!("http://127.0.0.1:{}/op", server.addr.port()); - - let result = poll_document_intelligence(&operation_url, &operation_url, &[], None).await; - - assert!(result.is_ok(), "got {result:?}"); - assert!(server.connection_count() >= 1); - } - - #[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 operation_url = format!("http://127.0.0.1:{}/op", server.addr.port()); - - let error = poll_document_intelligence( - &operation_url, - &operation_url, - &[], - Some(Duration::from_millis(150)), - ) - .await - .unwrap_err(); - - assert!( - matches!(&error, CoreError::Network(message) if message.contains("timed out")), - "got {error:?}" - ); - } - - #[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 oversized_error_body_without_content_length_is_bounded() { @@ -998,163 +130,7 @@ mod tests { let prefix = body .strip_suffix("... (truncated)") .expect("oversized error body must be truncated to the cap"); - assert_eq!(prefix.chars().count(), ERROR_BODY_MAX_CHARS); + assert_eq!(prefix.chars().count(), OCR_ERROR_BODY_MAX_CHARS); assert!(prefix.chars().all(|c| c == 'a')); } - - #[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("::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("8.8.8.8".parse().unwrap())); - assert!(!is_blocked_ip("::ffff:8.8.8.8".parse().unwrap())); - } - - #[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 { - pin_validated_url(&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); - } - - #[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()); - } - - #[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); - } } diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr/document_fetch.rs b/litellm-rust/crates/ai-gateway/src/io/ocr/document_fetch.rs new file mode 100644 index 00000000000..62fa5eaabdc --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/document_fetch.rs @@ -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::().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 { + 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(url: &str, resolve: P) -> CoreResult<(Url, reqwest::Response)> +where + P: Fn(Url) -> Fut, + Fut: std::future::Future>>, +{ + 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(¤t_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, ¤t_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> { + 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 { + 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 = 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); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr/reducto.rs b/litellm-rust/crates/ai-gateway/src/io/ocr/reducto.rs new file mode 100644 index 00000000000..eb364ad0d62 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/reducto.rs @@ -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, 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 { + 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, + mime: String, + parse_url: &str, + headers: &[(String, String)], + timeout: Option, +) -> CoreResult { + 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, +) -> CoreResult { + 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 = 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)) +} diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr/ssrf.rs b/litellm-rust/crates/ai-gateway/src/io/ocr/ssrf.rs new file mode 100644 index 00000000000..f87e90a916c --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/ssrf.rs @@ -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 { + 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> { + ensure_allowed_url(url)?; + let port = url.port_or_known_default().ok_or_else(blocked_url_error)?; + let addresses: Vec = 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>>>; + +#[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::::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, +} + +fn document_client_cache() -> &'static Mutex> { + static CACHE: OnceLock>> = OnceLock::new(); + CACHE.get_or_init(|| Mutex::new(HashMap::new())) +} + +pub(super) fn pinned_client(url: &Url, addresses: &[SocketAddr]) -> CoreResult { + 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()); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr/test_support.rs b/litellm-rust/crates/ai-gateway/src/io/ocr/test_support.rs new file mode 100644 index 00000000000..574dd461d84 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/io/ocr/test_support.rs @@ -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, + 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) -> 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 { + [headers.as_bytes(), body].concat() +} diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index 6c04fbb7626..e91964c0fe9 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -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;