From 66582fb33f3484a93ab2560ab50afeaa379841f9 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 16 Jul 2026 22:22:19 +0000 Subject: [PATCH] fix(ocr-security): complete URL and polling boundary enforcement Harden the Rust OCR fetch and Azure Document Intelligence polling paths. Parse the document URL host explicitly so IPv4-mapped IPv6 literals and bracketed forms are classified before any connection, resolve domain names and reject when any resolved address falls in a blocked range, and revalidate every redirect hop with per-hop connection accounting. Restrict the Azure Operation-Location to the same scheme, host and port as the original request and surface unsafe polling targets as typed input errors. Enforce the streaming download cap even when Content-Length is absent. Map OCR input-rejection ValueErrors (SSRF, cross-origin polling, oversized download, malformed input) to litellm.BadRequestError so callers and the proxy receive a typed 400 instead of a generic 500, while provider runtime failures still flow through exception_type. Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm-rust/Cargo.lock | 1 + litellm-rust/Cargo.toml | 1 + litellm-rust/crates/ai-gateway/Cargo.toml | 1 + .../ai-gateway/src/io/ocr/common_utils.rs | 345 ++++++++++++++++-- litellm/ocr/main.py | 23 +- tests/e2e/gateway/test_ocr_rust_e2e.py | 25 ++ tests/test_litellm/ocr/test_rust_bridge.py | 91 +++++ 7 files changed, 457 insertions(+), 30 deletions(-) diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 422dfb20065..8717ca7a190 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -663,6 +663,7 @@ dependencies = [ "subtle", "tokio", "tokio-tungstenite", + "url", ] [[package]] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 00af23b4c00..801ce894f85 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -28,3 +28,4 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"] tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] } futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] } base64 = "0.22" +url = "2" diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index 2f414159158..a3864a78c95 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -18,6 +18,7 @@ litellm-core.workspace = true # reqwest (rustls + json) is used by io/ocr and ships realtime logs to the # Python proxy callbacks API. reqwest.workspace = true +url.workspace = true # `sync` powers the bounded mpsc channel the realtime logger drains. tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "time", "sync"] } tokio-tungstenite.workspace = true 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 36d3d7da2e5..1eb1499202c 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 @@ -8,6 +8,7 @@ use litellm_core::ocr::transformation::OcrProviderConfig; use litellm_core::CoreResult; 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, @@ -169,36 +170,38 @@ fn blocked_url_error(url: &Url) -> CoreError { )) } +fn reject_blocked_literal(ip: IpAddr, url: &Url) -> CoreResult<()> { + if is_blocked_ip(ip) { + return Err(blocked_url_error(url)); + } + Ok(()) +} + +async fn validate_resolved_host(domain: &str, url: &Url) -> CoreResult<()> { + let port = url + .port_or_known_default() + .ok_or_else(|| blocked_url_error(url))?; + let addresses: Vec<_> = tokio::net::lookup_host((domain, port)) + .await + .map_err(|_| blocked_url_error(url))? + .collect(); + if addresses.is_empty() || addresses.iter().any(|address| is_blocked_ip(address.ip())) { + return Err(blocked_url_error(url)); + } + Ok(()) +} + 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(()); + match url.host() { + Some(Host::Ipv4(ip)) => reject_blocked_literal(IpAddr::V4(ip), url), + Some(Host::Ipv6(ip)) => reject_blocked_literal(IpAddr::V6(ip), url), + Some(Host::Domain(domain)) => validate_resolved_host(domain, url).await, + None => Err(blocked_url_error(url)), } - - 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 { @@ -214,6 +217,20 @@ fn redirect_location(response: &reqwest::Response, url: &Url) -> CoreResult } async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response)> { + fetch_with_redirects(url, |candidate| async move { + validate_safe_fetch_url(&candidate).await + }) + .await +} + +async fn fetch_with_redirects( + url: &str, + validate: V, +) -> CoreResult<(Url, reqwest::Response)> +where + V: Fn(Url) -> Fut, + Fut: std::future::Future>, +{ let client = reqwest::Client::builder() .redirect(reqwest::redirect::Policy::none()) .build() @@ -222,7 +239,7 @@ async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response) .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?; + validate(current_url.clone()).await?; let response = client .get(current_url.clone()) .send() @@ -258,8 +275,8 @@ fn enforce_download_size(content_length: u64, max_bytes: u64, url: &Url) -> Core async fn read_response_with_limit( mut response: reqwest::Response, url: &Url, + max_bytes: u64, ) -> 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 { @@ -306,7 +323,8 @@ 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 bytes = + read_response_with_limit(response, &final_url, max_document_download_bytes()).await?; let data_uri = format!( "data:{content_type};base64,{}", BASE64_STANDARD.encode(bytes) @@ -504,8 +522,8 @@ pub(super) async fn poll_document_intelligence( timeout: Option, ) -> CoreResult { if !same_origin(operation_url, original_url) { - return Err(CoreError::InvalidResponse( - "Azure Document Intelligence: rejected cross-origin polling URL".to_string(), + return Err(CoreError::InvalidRequest( + "Azure Document Intelligence: rejected unsafe polling target".to_string(), )); } @@ -557,6 +575,275 @@ pub(super) async fn poll_document_intelligence( mod tests { 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", + "ftp://8.8.8.8/x", + ]; + for raw in blocked { + let url = Url::parse(raw).unwrap(); + let error = validate_safe_fetch_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!(validate_safe_fetch_url(&allowed).await.is_ok()); + } + + #[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_port = redirector.addr.port(); + let start_url = format!("http://127.0.0.1:{redirector_port}/doc.png"); + + let validate = move |candidate: Url| async move { + if candidate.port() == Some(redirector_port) { + return Ok(()); + } + validate_safe_fetch_url(&candidate).await + }; + let error = fetch_with_redirects(&start_url, validate) + .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 url = Url::parse(&url_str).unwrap(); + let error = read_response_with_limit(response, &url, 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 url = Url::parse(&url_str).unwrap(); + let bytes = read_response_with_limit(response, &url, 1024) + .await + .unwrap(); + + assert_eq!(bytes.len(), 512); + } #[test] fn blocks_private_and_metadata_ips() { diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 9225b6ff13c..7f0a0a0768e 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -7,7 +7,7 @@ import mimetypes import os import re from io import IOBase -from typing import Any, Coroutine, Union, cast +from typing import Any, Coroutine, NoReturn, Union, cast import httpx @@ -199,6 +199,19 @@ def _missing_rust_bridge_error() -> RuntimeError: ) +def _raise_ocr_input_error( + e: ValueError, + *, + model: str, + custom_llm_provider: str | None, +) -> NoReturn: + raise litellm.BadRequestError( + message=str(e), + model=model, + llm_provider=custom_llm_provider or "", + ) from e + + async def _run_rust_aocr( rust_aocr: RustAocr, model: str, @@ -357,6 +370,10 @@ async def aocr( litellm_logging_obj=litellm_logging_obj, ) except Exception as e: + if isinstance(e, ValueError): + _raise_ocr_input_error( + e, model=model, custom_llm_provider=custom_llm_provider + ) raise litellm.exception_type( model=model, custom_llm_provider=custom_llm_provider, @@ -621,6 +638,10 @@ def ocr( litellm_logging_obj=litellm_logging_obj, ) except Exception as e: + if isinstance(e, ValueError): + _raise_ocr_input_error( + e, model=model, custom_llm_provider=custom_llm_provider + ) raise litellm.exception_type( model=model, custom_llm_provider=custom_llm_provider, diff --git a/tests/e2e/gateway/test_ocr_rust_e2e.py b/tests/e2e/gateway/test_ocr_rust_e2e.py index dcf85898365..423f69ae3f3 100644 --- a/tests/e2e/gateway/test_ocr_rust_e2e.py +++ b/tests/e2e/gateway/test_ocr_rust_e2e.py @@ -59,6 +59,19 @@ RUST_OCR_GATEWAY_CASES = [ ), ] +SSRF_BLOCKED_DOCUMENT_URLS = [ + pytest.param("http://127.0.0.1/secret", id="loopback"), + pytest.param("http://10.0.0.5/internal", id="rfc1918"), + pytest.param( + "http://169.254.169.254/latest/meta-data/", id="link_local_metadata" + ), + pytest.param("http://100.64.0.1/internal", id="cgnat"), + pytest.param("http://198.18.0.1/internal", id="benchmark"), + pytest.param("http://[::1]/secret", id="ipv6_loopback"), + pytest.param("http://[::ffff:169.254.169.254]/x", id="mapped_ipv6_metadata"), + pytest.param("http://[::ffff:10.0.0.5]/x", id="mapped_ipv6_rfc1918"), +] + CONFIG_PATH = Path(__file__).with_name("litellm-config.yml") @@ -146,3 +159,15 @@ class TestRustOcrGateway: assert response.status_code == 200, response.text _assert_ocr_response_shape(response.json()) + + @pytest.mark.parametrize("document_url", SSRF_BLOCKED_DOCUMENT_URLS) + def test_rust_ocr_rejects_special_use_urls_with_typed_4xx( + self, resources: OcrResources, document_url: str + ) -> None: + response = resources.gateway.ocr( + "rust-ocr-azure-ai", + {"type": "document_url", "document_url": document_url}, + ) + + assert response.status_code == 400, response.text + assert "SSRF protection" in response.text diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 5647b72667e..d30d33bfec6 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -130,6 +130,42 @@ class RaisingAsyncBridge: raise RuntimeError("bridge failed") +class SsrfRejectingBridge: + def __call__( + self, + model: str, + document: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> dict[str, object]: + raise ValueError( + "invalid request: OCR document URL rejected by SSRF protection: " + "http://169.254.169.254/latest/meta-data/" + ) + + +class SsrfRejectingAsyncBridge: + async def __call__( + self, + model: str, + document: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> dict[str, object]: + raise ValueError( + "invalid request: OCR document URL rejected by SSRF protection: " + "http://169.254.169.254/latest/meta-data/" + ) + + class RecordingLogging: """A spy standing in for ``LiteLLMLoggingObj`` to capture ``pre_call``.""" @@ -455,6 +491,61 @@ async def test_aocr_exception_type_uses_resolved_provider_context( assert captured["custom_llm_provider"] == "mistral" +def test_ocr_input_rejection_maps_to_bad_request(): + rust_bridge._set_rust_ocr_bridge(ocr=SsrfRejectingBridge()) + + with pytest.raises(litellm.BadRequestError) as exc_info: + litellm.ocr( + model="azure_ai/pixtral-12b-2409", + document={ + "type": "document_url", + "document_url": "http://169.254.169.254/latest/meta-data/", + }, + api_key="sk-test", + api_base="https://example.services.ai.azure.com", + ) + + assert exc_info.value.status_code == 400 + assert "SSRF protection" in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_aocr_input_rejection_maps_to_bad_request(): + rust_bridge._set_rust_ocr_bridge(aocr=SsrfRejectingAsyncBridge()) + + with pytest.raises(litellm.BadRequestError) as exc_info: + await litellm.aocr( + model="azure_ai/pixtral-12b-2409", + document={ + "type": "document_url", + "document_url": "http://169.254.169.254/latest/meta-data/", + }, + api_key="sk-test", + api_base="https://example.services.ai.azure.com", + ) + + assert exc_info.value.status_code == 400 + assert "SSRF protection" in str(exc_info.value) + + +def test_ocr_provider_runtime_error_is_not_downgraded_to_bad_request( + monkeypatch: pytest.MonkeyPatch, +): + captured: dict[str, object] = {} + + def fake_exception_type(**kwargs: object) -> CapturedException: + captured.update(kwargs) + return CapturedException("wrapped") + + monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) + rust_bridge._set_rust_ocr_bridge(ocr=RaisingBridge()) + + with pytest.raises(CapturedException): + litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + + assert captured["original_exception"].__class__ is RuntimeError + + def test_ocr_forwards_timeout_to_rust(fake_bridge): """Caller-supplied timeout must flow into the Rust bridge so the fixed 600s client ceiling doesn't silently override shorter deadlines."""