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>
This commit is contained in:
Devin AI 2026-07-16 22:22:19 +00:00
parent c264758eb8
commit 66582fb33f
7 changed files with 457 additions and 30 deletions

View file

@ -663,6 +663,7 @@ dependencies = [
"subtle",
"tokio",
"tokio-tungstenite",
"url",
]
[[package]]

View file

@ -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"

View file

@ -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

View file

@ -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::<IpAddr>() {
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<Url> {
@ -214,6 +217,20 @@ fn redirect_location(response: &reqwest::Response, url: &Url) -> CoreResult<Url>
}
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<V, Fut>(
url: &str,
validate: V,
) -> CoreResult<(Url, reqwest::Response)>
where
V: Fn(Url) -> Fut,
Fut: std::future::Future<Output = CoreResult<()>>,
{
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(&current_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<Vec<u8>> {
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<Duration>,
) -> CoreResult<Value> {
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<AtomicUsize>,
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<u8>) -> CountingServer {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let connections = Arc::new(AtomicUsize::new(0));
let counter = connections.clone();
let handle = tokio::spawn(async move {
loop {
let Ok((mut socket, _)) = listener.accept().await else {
return;
};
counter.fetch_add(1, Ordering::SeqCst);
let mut discard = [0u8; 2048];
let _ = socket.read(&mut discard).await;
let _ = socket.write_all(&response).await;
let _ = socket.flush().await;
}
});
CountingServer {
addr,
connections,
handle,
}
}
fn http_response(headers: &str, body: &[u8]) -> Vec<u8> {
[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() {

View file

@ -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,

View file

@ -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

View file

@ -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."""