From 8170d9b0716b169897ceb02661b1a83bb205c150 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 29 Aug 2026 09:49:04 -0700 Subject: [PATCH] feat(rust): add bounded safe fetch foundation --- .../crates/ai-gateway/src/ocr/common_utils.rs | 177 +------- litellm-rust/crates/core/CLAUDE.md | 4 + litellm-rust/crates/core/Cargo.toml | 1 + litellm-rust/crates/core/src/constants.rs | 6 + litellm-rust/crates/core/src/http_utils.rs | 2 + .../crates/core/src/http_utils/safe_fetch.rs | 415 ++++++++++++++++++ 6 files changed, 435 insertions(+), 170 deletions(-) create mode 100644 litellm-rust/crates/core/src/http_utils/safe_fetch.rs diff --git a/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs b/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs index 9bc2818b6e7..7b6c1f833d2 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs @@ -1,12 +1,11 @@ -use std::net::IpAddr; use std::time::{Duration, Instant}; use base64::Engine; use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; use litellm_core::CoreResult; use litellm_core::error::CoreError; +use litellm_core::http_utils::safe_fetch::{SafeFetchOptions, safe_fetch}; use litellm_core::ocr::transformation::OcrProviderConfig; -use reqwest::Url; use serde_json::{Map, Value}; use litellm_core::providers::azure_ai::ocr::transformation::{ @@ -23,7 +22,6 @@ use crate::client::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; pub(super) fn truncate_error_body(body: &str) -> String { if body.chars().count() <= ERROR_BODY_MAX_CHARS { @@ -110,151 +108,6 @@ fn max_document_download_bytes() -> u64 { (max_size_mb.max(0.0) * 1024.0 * 1024.0) as u64 } -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() - } - IpAddr::V6(ip) => { - let first_segment = ip.segments()[0]; - let is_unique_local = (first_segment & 0xfe00) == 0xfc00; - let is_link_local = (first_segment & 0xffc0) == 0xfe80; - ip.is_loopback() - || ip.is_unspecified() - || ip.is_multicast() - || is_unique_local - || is_link_local - || ip - .to_ipv4_mapped() - .or_else(|| ip.to_ipv4()) - .map(|v4| is_blocked_ip(IpAddr::V4(v4))) - .unwrap_or(false) - } - } -} - -fn blocked_url_error(url: &Url) -> CoreError { - CoreError::InvalidRequest(format!( - "OCR document URL rejected by SSRF protection: {url}" - )) -} - -async fn validate_safe_fetch_url(url: &Url) -> CoreResult<()> { - if !matches!(url.scheme(), "http" | "https") { - return Err(blocked_url_error(url)); - } - - let host = url.host_str().ok_or_else(|| blocked_url_error(url))?; - if let Ok(ip) = host.parse::() { - if is_blocked_ip(ip) { - return Err(blocked_url_error(url)); - } - return Ok(()); - } - - let port = url - .port_or_known_default() - .ok_or_else(|| blocked_url_error(url))?; - let addresses = tokio::net::lookup_host((host, port)) - .await - .map_err(|err| CoreError::Network(err.to_string()))?; - let mut saw_address = false; - for address in addresses { - saw_address = true; - if is_blocked_ip(address.ip()) { - return Err(blocked_url_error(url)); - } - } - if !saw_address { - return Err(blocked_url_error(url)); - } - Ok(()) -} - -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)> { - let client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .map_err(|err| CoreError::Network(err.to_string()))?; - 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 { - validate_safe_fetch_url(¤t_url).await?; - 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, url: &Url) -> CoreResult<()> { - if max_bytes == 0 { - return Err(CoreError::InvalidRequest(format!( - "OCR document URL download is disabled (MAX_IMAGE_URL_DOWNLOAD_SIZE_MB=0). url={url}" - ))); - } - 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). url={url}" - ))); - } - Ok(()) -} - -async fn read_response_with_limit( - mut response: reqwest::Response, - url: &Url, -) -> CoreResult> { - let max_bytes = max_document_download_bytes(); - if let Some(content_length) = response.content_length() { - enforce_download_size(content_length, max_bytes, url)?; - } else { - enforce_download_size(0, max_bytes, url)?; - } - - 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, url)?; - 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); @@ -263,17 +116,16 @@ pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreRes return Ok(document); } - let (final_url, response) = safe_get_document_url(url).await?; - let status = response.status(); + let response = safe_fetch(url, SafeFetchOptions::new(max_document_download_bytes())).await?; + let status = response.status; if !status.is_success() { - let body = response.text().await.unwrap_or_default(); return Err(CoreError::Http { status: status.as_u16(), - body: truncate_error_body(&body), + body: truncate_error_body(&String::from_utf8_lossy(&response.body)), }); } let content_type = response - .headers() + .headers .get(reqwest::header::CONTENT_TYPE) .and_then(|value| value.to_str().ok()) .and_then(|value| value.split(';').next()) @@ -281,10 +133,9 @@ pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreRes .filter(|value| !value.is_empty()) .unwrap_or("application/octet-stream") .to_string(); - let bytes = read_response_with_limit(response, &final_url).await?; let data_uri = format!( "data:{content_type};base64,{}", - BASE64_STANDARD.encode(bytes) + BASE64_STANDARD.encode(response.body) ); let mut transformed = document @@ -401,20 +252,6 @@ mod tests { use super::*; use serde_json::json; - #[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("::1".parse().unwrap())); - assert!(is_blocked_ip("fd00::1".parse().unwrap())); - assert!(is_blocked_ip("fe80::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("8.8.8.8".parse().unwrap())); - assert!(!is_blocked_ip("::ffff:8.8.8.8".parse().unwrap())); - } - #[tokio::test] async fn convert_document_url_rejects_loopback_fetch() { let error = convert_document_url_to_data_uri(json!({ @@ -427,7 +264,7 @@ mod tests { assert!(matches!( error, CoreError::InvalidRequest(message) - if message.contains("SSRF protection") + if message.contains("not public") )); } diff --git a/litellm-rust/crates/core/CLAUDE.md b/litellm-rust/crates/core/CLAUDE.md index 5d36305ded5..2bd6b548923 100644 --- a/litellm-rust/crates/core/CLAUDE.md +++ b/litellm-rust/crates/core/CLAUDE.md @@ -30,6 +30,10 @@ Not allowed: - Provider-specific branching that belongs in `providers`. - Panics for user/provider-controlled input. +User-controlled remote media and file URLs must be downloaded through +`http_utils::safe_fetch`. It owns SSRF-safe DNS, redirect validation, timeouts, +and byte limits; route/provider code owns media interpretation and encoding. + ## Typed Contracts (core rule) Trait and function boundaries MUST be strongly typed. No stringly-typed JSON diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index ab8050734f2..abb19d21751 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -12,6 +12,7 @@ serde.workspace = true serde_json.workspace = true thiserror.workspace = true sha2.workspace = true +tokio = { workspace = true, features = ["net", "time"] } aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"], optional = true } aws-sdk-sts = { version = "1.108.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index e1ac0a4fc8f..feadf487a7d 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -14,6 +14,12 @@ pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10; /// before truncation, so provider bodies are bounded and data-minimized. pub(crate) const UPSTREAM_ERROR_BODY_MAX_CHARS: usize = 256; +/// Defaults for bounded fetches of user-controlled public URLs. Callers must +/// still supply a route-appropriate response byte limit. +pub(crate) const SAFE_FETCH_CONNECT_TIMEOUT_SECS: u64 = 10; +pub(crate) const SAFE_FETCH_TIMEOUT_SECS: u64 = 60; +pub(crate) const SAFE_FETCH_MAX_REDIRECTS: usize = 10; + /// Provider name used for Anthropic Messages when a deployment's provider model /// does not carry an explicit provider prefix. pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; diff --git a/litellm-rust/crates/core/src/http_utils.rs b/litellm-rust/crates/core/src/http_utils.rs index c541f50275b..02788b2b28b 100644 --- a/litellm-rust/crates/core/src/http_utils.rs +++ b/litellm-rust/crates/core/src/http_utils.rs @@ -5,6 +5,8 @@ use serde_json::{Map, Value}; use crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS; use crate::error::{CoreError, CoreResult, json_type_name}; +pub mod safe_fetch; + /// Bound an upstream error body before it crosses a host boundary, so provider /// bodies stay data-minimized. pub fn truncate_error_body(body: &str) -> String { diff --git a/litellm-rust/crates/core/src/http_utils/safe_fetch.rs b/litellm-rust/crates/core/src/http_utils/safe_fetch.rs new file mode 100644 index 00000000000..c8ef48bb0e7 --- /dev/null +++ b/litellm-rust/crates/core/src/http_utils/safe_fetch.rs @@ -0,0 +1,415 @@ +//! Bounded downloads of user-controlled public URLs. +//! +//! This module owns transport policy only: SSRF protection, redirect handling, +//! timeouts, connection reuse, and response-size limits. Route code decides +//! whether the bytes become a data URI, multipart part, provider upload, or +//! something else. + +use std::io; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; +use std::sync::{Arc, OnceLock}; +use std::time::Duration; + +use reqwest::dns::{Addrs, Name, Resolve, Resolving}; +use reqwest::{StatusCode, Url, header::HeaderMap}; + +use crate::constants::{ + SAFE_FETCH_CONNECT_TIMEOUT_SECS, SAFE_FETCH_MAX_REDIRECTS, SAFE_FETCH_TIMEOUT_SECS, +}; +use crate::error::{CoreError, CoreResult}; + +/// Required policy for one public-URL fetch. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct SafeFetchOptions { + max_response_bytes: u64, + timeout: Duration, + max_redirects: usize, +} + +impl SafeFetchOptions { + /// Construct a fetch policy with an explicit route-appropriate byte limit. + pub const fn new(max_response_bytes: u64) -> Self { + Self { + max_response_bytes, + timeout: Duration::from_secs(SAFE_FETCH_TIMEOUT_SECS), + max_redirects: SAFE_FETCH_MAX_REDIRECTS, + } + } + + pub const fn with_timeout(mut self, timeout: Duration) -> Self { + self.timeout = timeout; + self + } + + pub const fn with_max_redirects(mut self, max_redirects: usize) -> Self { + self.max_redirects = max_redirects; + self + } +} + +/// Fully buffered, bounded response from a validated public URL. +#[derive(Debug)] +pub struct SafeFetchResponse { + pub final_url: Url, + pub status: StatusCode, + pub headers: HeaderMap, + pub body: Vec, +} + +#[derive(Debug)] +struct PublicIpResolver; + +impl Resolve for PublicIpResolver { + fn resolve(&self, name: Name) -> Resolving { + let host = name.as_str().to_string(); + Box::pin(async move { + let addresses = tokio::net::lookup_host((host.as_str(), 0)) + .await + .map_err(boxed_io_error)? + .collect::>(); + if addresses.is_empty() { + return Err(boxed_io_error(io::Error::new( + io::ErrorKind::NotFound, + "host resolved to no addresses", + ))); + } + if addresses.iter().any(|address| !is_public_ip(address.ip())) { + return Err(boxed_io_error(io::Error::new( + io::ErrorKind::PermissionDenied, + "host resolved to a non-public address", + ))); + } + Ok(Box::new(addresses.into_iter()) as Addrs) + }) + } +} + +fn boxed_io_error(error: io::Error) -> Box { + Box::new(error) +} + +fn safe_fetch_client() -> CoreResult<&'static reqwest::Client> { + static CLIENT: OnceLock> = OnceLock::new(); + CLIENT + .get_or_init(|| { + reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .dns_resolver(Arc::new(PublicIpResolver)) + // A proxy could resolve the target itself and bypass our resolver. + .no_proxy() + .connect_timeout(Duration::from_secs(SAFE_FETCH_CONNECT_TIMEOUT_SECS)) + .build() + .map_err(|error| error.to_string()) + }) + .as_ref() + .map_err(|error| CoreError::Connect(error.clone())) +} + +fn is_public_ip(ip: IpAddr) -> bool { + match ip { + IpAddr::V4(ip) => is_public_ipv4(ip), + IpAddr::V6(ip) => ip + .to_ipv4_mapped() + .or_else(|| ip.to_ipv4()) + .map(is_public_ipv4) + .unwrap_or_else(|| is_public_ipv6(ip)), + } +} + +fn is_public_ipv4(ip: Ipv4Addr) -> bool { + let [a, b, c, d] = ip.octets(); + let is_shared = a == 100 && b & 0b1100_0000 == 0b0100_0000; + let is_protocol_assignment = a == 192 && b == 0 && c == 0 && d != 9 && d != 10; + let is_documentation = (a == 192 && b == 0 && c == 2) + || (a == 198 && b == 51 && c == 100) + || (a == 203 && b == 0 && c == 113); + let is_benchmarking = a == 198 && matches!(b, 18 | 19); + let is_deprecated_relay = a == 192 && b == 88 && c == 99; + let is_cloud_metadata = [a, b, c, d] == [168, 63, 129, 16]; + + !(a == 0 + || ip.is_private() + || is_shared + || ip.is_loopback() + || ip.is_link_local() + || is_protocol_assignment + || is_documentation + || is_benchmarking + || is_deprecated_relay + || a >= 224 + || is_cloud_metadata) +} + +fn is_public_ipv6(ip: Ipv6Addr) -> bool { + let segments = ip.segments(); + let value = u128::from_be_bytes(ip.octets()); + let is_ietf_assignment = segments[0] == 0x2001 + && segments[1] < 0x0200 + && !(value == 0x2001_0001_0000_0000_0000_0000_0000_0001 + || value == 0x2001_0001_0000_0000_0000_0000_0000_0002 + || segments[1] == 0x0003 + || (segments[1] == 0x0004 && segments[2] == 0x0112) + || (0x0020..=0x003f).contains(&segments[1])); + let is_documentation = (segments[0] == 0x2001 && segments[1] == 0x0db8) + || (segments[0] == 0x3fff && segments[1] <= 0x0fff); + + !(ip.is_unspecified() + || ip.is_loopback() + || ip.is_multicast() + || matches!(segments, [0x0064, 0xff9b, 0x0001, _, _, _, _, _]) + || matches!(segments, [0x0100, 0, 0, 0, _, _, _, _]) + || is_ietf_assignment + || segments[0] == 0x2002 + || is_documentation + || segments[0] == 0x5f00 + || (segments[0] & 0xfe00) == 0xfc00 + || (segments[0] & 0xffc0) == 0xfe80) +} + +fn invalid_url(message: impl Into) -> CoreError { + CoreError::InvalidRequest(format!("unsafe remote URL: {}", message.into())) +} + +fn validate_url(url: &Url) -> CoreResult<()> { + if !matches!(url.scheme(), "http" | "https") { + return Err(invalid_url("only http and https are allowed")); + } + if !url.username().is_empty() || url.password().is_some() { + return Err(invalid_url("embedded credentials are not allowed")); + } + let host = url + .host_str() + .ok_or_else(|| invalid_url("hostname is required"))?; + let ip_literal = host + .strip_prefix('[') + .and_then(|host| host.strip_suffix(']')) + .unwrap_or(host); + if ip_literal + .parse::() + .is_ok_and(|ip| !is_public_ip(ip)) + { + return Err(invalid_url("target address is not public")); + } + Ok(()) +} + +fn is_followable_redirect(status: StatusCode) -> bool { + matches!( + status, + StatusCode::MOVED_PERMANENTLY + | StatusCode::FOUND + | StatusCode::SEE_OTHER + | StatusCode::TEMPORARY_REDIRECT + | StatusCode::PERMANENT_REDIRECT + ) +} + +fn redirect_url(current_url: &Url, location: Option<&str>) -> CoreResult { + let location = location.ok_or_else(|| { + CoreError::InvalidResponse("safe fetch redirect missing Location header".to_string()) + })?; + current_url.join(location).map_err(|error| { + CoreError::InvalidResponse(format!("safe fetch returned invalid redirect: {error}")) + }) +} + +fn enforce_size(size: u64, max_bytes: u64) -> CoreResult<()> { + if max_bytes == 0 { + return Err(CoreError::InvalidRequest( + "remote URL fetching is disabled by a zero-byte limit".to_string(), + )); + } + if size > max_bytes { + return Err(CoreError::InvalidRequest(format!( + "remote response exceeds the {max_bytes}-byte limit" + ))); + } + Ok(()) +} + +fn append_bounded_chunk( + body: &mut Vec, + downloaded: &mut u64, + chunk: &[u8], + max_bytes: u64, +) -> CoreResult<()> { + let next_size = downloaded + .checked_add(chunk.len() as u64) + .ok_or_else(|| CoreError::InvalidRequest("remote response size overflowed".to_string()))?; + enforce_size(next_size, max_bytes)?; + body.extend_from_slice(chunk); + *downloaded = next_size; + Ok(()) +} + +async fn read_bounded_response( + mut response: reqwest::Response, + max_bytes: u64, +) -> CoreResult<(StatusCode, HeaderMap, Vec)> { + if let Some(content_length) = response.content_length() { + enforce_size(content_length, max_bytes)?; + } else { + enforce_size(0, max_bytes)?; + } + + let capacity = response + .content_length() + .and_then(|length| usize::try_from(length).ok()) + .unwrap_or_default(); + let status = response.status(); + let headers = response.headers().clone(); + let mut body = Vec::with_capacity(capacity); + let mut downloaded = 0_u64; + while let Some(chunk) = response + .chunk() + .await + .map_err(|error| CoreError::Network(error.to_string()))? + { + append_bounded_chunk(&mut body, &mut downloaded, &chunk, max_bytes)?; + } + Ok((status, headers, body)) +} + +fn map_send_error(error: reqwest::Error) -> CoreError { + if error.is_connect() || error.is_builder() { + CoreError::Connect(error.to_string()) + } else { + CoreError::Network(error.to_string()) + } +} + +async fn fetch_inner(url: &str, options: SafeFetchOptions) -> CoreResult { + enforce_size(0, options.max_response_bytes)?; + let client = safe_fetch_client()?; + let mut current_url = + Url::parse(url).map_err(|error| invalid_url(format!("could not be parsed: {error}")))?; + let mut redirects = 0_usize; + + loop { + validate_url(¤t_url)?; + let response = client + .get(current_url.clone()) + .send() + .await + .map_err(map_send_error)?; + if !is_followable_redirect(response.status()) { + let (status, headers, body) = + read_bounded_response(response, options.max_response_bytes).await?; + return Ok(SafeFetchResponse { + final_url: current_url, + status, + headers, + body, + }); + } + if redirects == options.max_redirects { + return Err(CoreError::InvalidRequest( + "remote URL exceeded the redirect limit".to_string(), + )); + } + let location = response + .headers() + .get(reqwest::header::LOCATION) + .and_then(|value| value.to_str().ok()); + current_url = redirect_url(¤t_url, location)?; + redirects += 1; + } +} + +/// Fetch a user-controlled public URL under an explicit bounded policy. +/// +/// DNS answers used by the connector are validated, proxies are disabled, and +/// every redirect target is checked before it is requested. The returned body +/// can never exceed `options.max_response_bytes`. +pub async fn safe_fetch(url: &str, options: SafeFetchOptions) -> CoreResult { + tokio::time::timeout(options.timeout, fetch_inner(url, options)) + .await + .map_err(|_| CoreError::Network("safe fetch timed out".to_string()))? +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn only_public_addresses_are_allowed() { + for blocked in [ + "0.0.0.0", + "10.0.0.1", + "100.64.0.1", + "127.0.0.1", + "168.63.129.16", + "169.254.169.254", + "192.0.2.1", + "198.18.0.1", + "224.0.0.1", + "::1", + "fc00::1", + "fe80::1", + "2001:db8::1", + "::ffff:169.254.169.254", + ] { + assert!(!is_public_ip(blocked.parse().unwrap()), "allowed {blocked}"); + } + for public in [ + "8.8.8.8", + "1.1.1.1", + "2001:4860:4860::8888", + "::ffff:8.8.8.8", + ] { + assert!(is_public_ip(public.parse().unwrap()), "blocked {public}"); + } + } + + #[test] + fn validates_scheme_credentials_and_ip_literals() { + assert!(validate_url(&Url::parse("ftp://example.com/file").unwrap()).is_err()); + assert!(validate_url(&Url::parse("https://user:pass@example.com/file").unwrap()).is_err()); + assert!(validate_url(&Url::parse("http://127.0.0.1/file").unwrap()).is_err()); + assert!(validate_url(&Url::parse("http://[::1]/file").unwrap()).is_err()); + assert!(validate_url(&Url::parse("https://example.com/file").unwrap()).is_ok()); + } + + #[test] + fn resolves_relative_redirects_and_rejects_missing_locations() { + let current = Url::parse("https://example.com/media/start").unwrap(); + assert_eq!( + redirect_url(¤t, Some("../next")).unwrap().as_str(), + "https://example.com/next" + ); + assert!(redirect_url(¤t, None).is_err()); + + let loopback = redirect_url(¤t, Some("http://127.0.0.1/private")).unwrap(); + assert!(validate_url(&loopback).is_err()); + } + + #[test] + fn enforces_declared_and_streamed_size_limits() { + assert!(enforce_size(10, 10).is_ok()); + assert!(enforce_size(11, 10).is_err()); + assert!(enforce_size(0, 0).is_err()); + + let mut body = Vec::new(); + let mut downloaded = 0; + append_bounded_chunk(&mut body, &mut downloaded, b"123456", 10).unwrap(); + assert!(append_bounded_chunk(&mut body, &mut downloaded, b"78901", 10).is_err()); + assert_eq!(body, b"123456"); + assert_eq!(downloaded, 6); + } + + #[test] + fn reuses_one_hardened_client() { + let first = safe_fetch_client().unwrap(); + let second = safe_fetch_client().unwrap(); + assert!(std::ptr::eq(first, second)); + } + + #[tokio::test] + async fn resolver_rejects_non_public_dns_answers() { + let name = "localhost".parse::().unwrap(); + let Err(error) = PublicIpResolver.resolve(name).await else { + panic!("localhost must not resolve through the public-IP resolver"); + }; + assert!(error.to_string().contains("non-public address")); + } +}