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 1eb1499202c..309231ea44d 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,4 +1,6 @@ -use std::net::{IpAddr, Ipv4Addr}; +use std::collections::HashMap; +use std::net::{IpAddr, Ipv4Addr, SocketAddr}; +use std::sync::{Arc, Mutex, PoisonError}; use std::time::{Duration, Instant}; use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; @@ -6,6 +8,7 @@ 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; @@ -164,43 +167,55 @@ fn is_blocked_ip(ip: IpAddr) -> bool { } } -fn blocked_url_error(url: &Url) -> CoreError { - CoreError::InvalidRequest(format!( - "OCR document URL rejected by SSRF protection: {url}" - )) +fn blocked_url_error() -> CoreError { + CoreError::InvalidRequest("OCR document URL rejected by SSRF protection".to_string()) } -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<()> { +async fn pin_validated_url(url: &Url) -> CoreResult> { if !matches!(url.scheme(), "http" | "https") { - return Err(blocked_url_error(url)); + 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) +} - 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)), +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", + )), + } + }) } } @@ -218,28 +233,32 @@ 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 + pin_validated_url(&candidate).await }) .await } -async fn fetch_with_redirects( - url: &str, - validate: V, -) -> CoreResult<(Url, reqwest::Response)> +async fn fetch_with_redirects(url: &str, pin: P) -> CoreResult<(Url, reqwest::Response)> where - V: Fn(Url) -> Fut, - Fut: std::future::Future>, + P: Fn(Url) -> Fut, + Fut: std::future::Future>>, { + let pins: PinnedAddrs = Arc::new(Mutex::new(HashMap::new())); let client = reqwest::Client::builder() .redirect(reqwest::redirect::Policy::none()) + .dns_resolver(Arc::new(PinnedResolver { pins: pins.clone() })) .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(current_url.clone()).await?; + let addresses = pin(current_url.clone()).await?; + if let Some(host) = current_url.host_str() { + pins.lock() + .unwrap_or_else(PoisonError::into_inner) + .insert(host.to_owned(), addresses); + } let response = client .get(current_url.clone()) .send() @@ -256,17 +275,17 @@ where )) } -fn enforce_download_size(content_length: u64, max_bytes: u64, url: &Url) -> CoreResult<()> { +fn enforce_download_size(content_length: u64, max_bytes: u64) -> 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}" - ))); + 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). url={url}" + "OCR document size ({size_mb:.2}MB) exceeds maximum allowed size ({max_size_mb:.2}MB)" ))); } Ok(()) @@ -274,13 +293,12 @@ 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> { if let Some(content_length) = response.content_length() { - enforce_download_size(content_length, max_bytes, url)?; + enforce_download_size(content_length, max_bytes)?; } else { - enforce_download_size(0, max_bytes, url)?; + enforce_download_size(0, max_bytes)?; } let mut bytes = Vec::new(); @@ -291,7 +309,7 @@ async fn read_response_with_limit( .map_err(|err| CoreError::Network(err.to_string()))? { bytes_downloaded += chunk.len() as u64; - enforce_download_size(bytes_downloaded, max_bytes, url)?; + enforce_download_size(bytes_downloaded, max_bytes)?; bytes.extend_from_slice(&chunk); } Ok(bytes) @@ -305,7 +323,7 @@ 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 (_final_url, response) = safe_get_document_url(url).await?; let status = response.status(); if !status.is_success() { let body = response.text().await.unwrap_or_default(); @@ -323,8 +341,7 @@ 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, max_document_download_bytes()).await?; + let bytes = read_response_with_limit(response, max_document_download_bytes()).await?; let data_uri = format!( "data:{content_type};base64,{}", BASE64_STANDARD.encode(bytes) @@ -650,7 +667,7 @@ mod tests { ]; for raw in blocked { let url = Url::parse(raw).unwrap(); - let error = validate_safe_fetch_url(&url).await.unwrap_err(); + 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:?}" @@ -658,7 +675,10 @@ mod tests { } let allowed = Url::parse("http://8.8.8.8/x").unwrap(); - assert!(validate_safe_fetch_url(&allowed).await.is_ok()); + assert_eq!( + pin_validated_url(&allowed).await.unwrap(), + vec![SocketAddr::from(([8, 8, 8, 8], 80))] + ); } #[tokio::test] @@ -718,18 +738,17 @@ mod tests { .into_bytes(), ) .await; - let redirector_port = redirector.addr.port(); + 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 validate = move |candidate: Url| async move { + let pin = move |candidate: Url| async move { if candidate.port() == Some(redirector_port) { - return Ok(()); + return Ok(vec![redirector_addr]); } - validate_safe_fetch_url(&candidate).await + pin_validated_url(&candidate).await }; - let error = fetch_with_redirects(&start_url, validate) - .await - .unwrap_err(); + let error = fetch_with_redirects(&start_url, pin).await.unwrap_err(); assert!( matches!(&error, CoreError::InvalidRequest(message) if message.contains("SSRF protection")), @@ -815,10 +834,7 @@ mod tests { 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(); + let error = read_response_with_limit(response, 1024).await.unwrap_err(); assert!( matches!(&error, CoreError::InvalidRequest(message) if message.contains("exceeds maximum")), @@ -837,10 +853,7 @@ mod tests { 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(); + let bytes = read_response_with_limit(response, 1024).await.unwrap(); assert_eq!(bytes.len(), 512); } @@ -866,6 +879,51 @@ mod tests { 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; + // `.invalid` is guaranteed non-resolvable (RFC 6761), so the connection can only + // land on the server if the request used the pinned SocketAddr instead of a second + // DNS lookup. This is the anti-rebinding contract: the connect uses exactly the + // validated address set, never an ambient resolver answer. + 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 convert_document_url_rejects_loopback_fetch() { let error = convert_document_url_to_data_uri(json!({ diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 82024e1bf47..9e0f697091b 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -2,6 +2,7 @@ use std::time::Duration; use litellm_ai_gateway::io::ocr::{ocr as run_ocr, OcrRequest}; use litellm_core::error::CoreError; +use pyo3::create_exception; use pyo3::exceptions::{PyRuntimeError, PyValueError}; use pyo3::prelude::*; use pyo3::types::{PyAny, PyDict}; @@ -9,6 +10,8 @@ use serde_json::{Map, Value}; mod gil; +create_exception!(_native, RustOcrInputError, PyValueError); + type MarshaledOcrInputs = ( Value, Option>, @@ -31,11 +34,11 @@ fn json_to_py(py: Python<'_>, value: Value) -> PyResult> { fn core_error_to_pyerr(err: CoreError) -> PyErr { match err { - CoreError::Auth(message) => PyValueError::new_err(message), - CoreError::InvalidProvider(_) + CoreError::Auth(_) + | CoreError::InvalidProvider(_) | CoreError::InvalidRequest(_) | CoreError::InvalidType { .. } - | CoreError::MissingField(_) => PyValueError::new_err(err.to_string()), + | CoreError::MissingField(_) => RustOcrInputError::new_err(err.to_string()), other => PyRuntimeError::new_err(other.to_string()), } } @@ -48,7 +51,7 @@ fn optional_object_to_map( match value { Some(value) => match py_to_json(py, value.bind(py))? { Value::Object(map) => Ok(map), - _ => Err(PyValueError::new_err(format!("{name} must be a dict"))), + _ => Err(RustOcrInputError::new_err(format!("{name} must be a dict"))), }, None => Ok(Map::new()), } @@ -177,5 +180,9 @@ fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> { module.add_function(wrap_pyfunction!(ocr, module)?)?; module.add_function(wrap_pyfunction!(aocr, module)?)?; module.add_function(wrap_pyfunction!(gil_stats, module)?)?; + module.add( + "RustOcrInputError", + module.py().get_type::(), + )?; Ok(()) } diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 7f0a0a0768e..33062a069ac 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -21,6 +21,7 @@ from litellm.ocr.rust_bridge import ( RustOcr, load_rust_aocr, load_rust_ocr, + rust_ocr_input_error_type, ) from litellm.utils import client, filter_out_litellm_params @@ -200,7 +201,7 @@ def _missing_rust_bridge_error() -> RuntimeError: def _raise_ocr_input_error( - e: ValueError, + e: BaseException, *, model: str, custom_llm_provider: str | None, @@ -212,6 +213,11 @@ def _raise_ocr_input_error( ) from e +def _is_rust_ocr_input_error(e: BaseException) -> bool: + input_error_type = rust_ocr_input_error_type() + return input_error_type is not None and isinstance(e, input_error_type) + + async def _run_rust_aocr( rust_aocr: RustAocr, model: str, @@ -370,7 +376,7 @@ async def aocr( litellm_logging_obj=litellm_logging_obj, ) except Exception as e: - if isinstance(e, ValueError): + if _is_rust_ocr_input_error(e): _raise_ocr_input_error( e, model=model, custom_llm_provider=custom_llm_provider ) @@ -638,7 +644,7 @@ def ocr( litellm_logging_obj=litellm_logging_obj, ) except Exception as e: - if isinstance(e, ValueError): + if _is_rust_ocr_input_error(e): _raise_ocr_input_error( e, model=model, custom_llm_provider=custom_llm_provider ) diff --git a/litellm/ocr/rust_bridge.py b/litellm/ocr/rust_bridge.py index c518c2cb874..e8437b07e2e 100644 --- a/litellm/ocr/rust_bridge.py +++ b/litellm/ocr/rust_bridge.py @@ -56,6 +56,21 @@ _UNSET: Final[_Unset] = _Unset() _rust_ocr_impl: RustOcr | None = None _rust_aocr_impl: RustAocr | None = None +_rust_ocr_input_error_override: type[BaseException] | None | _Unset = _UNSET + + +def _set_rust_ocr_input_error_type( + error_type: type[BaseException] | None | _Unset = _UNSET, +) -> None: + """Inject the exception type treated as a client-input rejection (tests only). + + Mirrors ``_set_rust_ocr_bridge`` so the native extension does not need to be + compiled to exercise the input-error mapping. Passing ``None`` clears a prior + override; omitting the argument preserves it. + """ + global _rust_ocr_input_error_override + if not isinstance(error_type, _Unset): + _rust_ocr_input_error_override = error_type def _set_rust_ocr_bridge( @@ -102,3 +117,24 @@ def load_rust_aocr() -> RustAocr | None: if native_bridge is None: return None return cast(RustAocr, getattr(native_bridge, "aocr", None)) + + +def rust_ocr_input_error_type() -> type[BaseException] | None: + """Return the native exception raised for client-input rejections, if available. + + The Rust bridge raises this dedicated type only for request-input problems + (SSRF-rejected URLs, malformed documents, unsafe polling targets); unrelated + internal failures surface as other exceptions and must not be downgraded to a + client error. + """ + if not isinstance(_rust_ocr_input_error_override, _Unset): + return _rust_ocr_input_error_override + from litellm.rust_bridge import get_native_bridge + + native_bridge = get_native_bridge() + if native_bridge is None: + return None + error_type = getattr(native_bridge, "RustOcrInputError", None) + if isinstance(error_type, type) and issubclass(error_type, BaseException): + return error_type + return None diff --git a/tests/e2e/gateway/test_ocr_rust_e2e.py b/tests/e2e/gateway/test_ocr_rust_e2e.py index 423f69ae3f3..f89023d1fdf 100644 --- a/tests/e2e/gateway/test_ocr_rust_e2e.py +++ b/tests/e2e/gateway/test_ocr_rust_e2e.py @@ -12,6 +12,7 @@ import os from dataclasses import dataclass from pathlib import Path from typing import Any +from urllib.parse import urlparse import httpx import pytest @@ -171,3 +172,10 @@ class TestRustOcrGateway: assert response.status_code == 400, response.text assert "SSRF protection" in response.text + # Public SSRF errors must be data-minimized: never echo the rejected + # URL, its host, or query back to the caller. + parsed = urlparse(document_url) + assert document_url not in response.text + assert parsed.netloc not in response.text + if parsed.query: + assert parsed.query not in response.text diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index d30d33bfec6..21d7a7522b1 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -36,6 +36,14 @@ class CapturedException(Exception): pass +class FakeRustOcrInputError(ValueError): + """Stands in for the native ``RustOcrInputError`` the compiled bridge raises. + + The native wheel isn't built in CI, so inject this type via + ``_set_rust_ocr_input_error_type`` to exercise the input-error mapping. + """ + + class RecordingBridge: """A fake ``RustOcr`` callable that records the args it was handed.""" @@ -142,9 +150,8 @@ class SsrfRejectingBridge: 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/" + raise FakeRustOcrInputError( + "invalid request: OCR document URL rejected by SSRF protection" ) @@ -160,12 +167,28 @@ class SsrfRejectingAsyncBridge: 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/" + raise FakeRustOcrInputError( + "invalid request: OCR document URL rejected by SSRF protection" ) +class UnrelatedValueErrorBridge: + """Raises a plain ``ValueError`` unrelated to request input (e.g. an internal bug).""" + + 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("some internal invariant broke") + + class RecordingLogging: """A spy standing in for ``LiteLLMLoggingObj`` to capture ``pre_call``.""" @@ -190,9 +213,11 @@ class RecordingLogging: def _reset_rust_bridge(): """Keep the global bridge state isolated between tests.""" rust_bridge._set_rust_ocr_bridge(ocr=None, aocr=None) + rust_bridge._set_rust_ocr_input_error_type(None) rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL yield rust_bridge._set_rust_ocr_bridge(ocr=None, aocr=None) + rust_bridge._set_rust_ocr_input_error_type(None) rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL @@ -492,6 +517,7 @@ async def test_aocr_exception_type_uses_resolved_provider_context( def test_ocr_input_rejection_maps_to_bad_request(): + rust_bridge._set_rust_ocr_input_error_type(FakeRustOcrInputError) rust_bridge._set_rust_ocr_bridge(ocr=SsrfRejectingBridge()) with pytest.raises(litellm.BadRequestError) as exc_info: @@ -511,6 +537,7 @@ def test_ocr_input_rejection_maps_to_bad_request(): @pytest.mark.asyncio async def test_aocr_input_rejection_maps_to_bad_request(): + rust_bridge._set_rust_ocr_input_error_type(FakeRustOcrInputError) rust_bridge._set_rust_ocr_bridge(aocr=SsrfRejectingAsyncBridge()) with pytest.raises(litellm.BadRequestError) as exc_info: @@ -546,6 +573,27 @@ def test_ocr_provider_runtime_error_is_not_downgraded_to_bad_request( assert captured["original_exception"].__class__ is RuntimeError +def test_ocr_unrelated_value_error_is_not_downgraded_to_bad_request( + monkeypatch: pytest.MonkeyPatch, +): + """Only the dedicated Rust input error becomes a 400; an unrelated internal + ValueError must stay a server error and flow through exception_type.""" + 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_input_error_type(FakeRustOcrInputError) + rust_bridge._set_rust_ocr_bridge(ocr=UnrelatedValueErrorBridge()) + + with pytest.raises(CapturedException): + litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + + assert captured["original_exception"].__class__ is ValueError + + 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."""