mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
fix(ocr): harden rust document fetch SSRF checks
This commit is contained in:
parent
62f93a3343
commit
378ab8f04d
2 changed files with 244 additions and 41 deletions
|
|
@ -35,6 +35,7 @@ fn http_client() -> &'static reqwest::Client {
|
|||
CLIENT.get_or_init(|| {
|
||||
reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(OCR_TIMEOUT_SECS))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.expect("failed to build reqwest client")
|
||||
})
|
||||
|
|
@ -386,6 +387,67 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn document_intelligence_post_does_not_follow_redirects_with_credentials() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("test listener binds");
|
||||
let addr = listener.local_addr().expect("listener has local addr");
|
||||
let redirect_url = format!("http://{addr}/leak");
|
||||
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut post_socket, _) = listener.accept().await.expect("accepts post request");
|
||||
let post_request = read_http_headers(&mut post_socket).await;
|
||||
let post_response = format!(
|
||||
"HTTP/1.1 307 Temporary Redirect\r\nlocation: {redirect_url}\r\ncontent-length: 0\r\nconnection: close\r\n\r\n"
|
||||
);
|
||||
post_socket
|
||||
.write_all(post_response.as_bytes())
|
||||
.await
|
||||
.expect("writes redirect response");
|
||||
|
||||
let second_connection =
|
||||
tokio::time::timeout(Duration::from_millis(200), listener.accept()).await;
|
||||
(post_request, second_connection.is_ok())
|
||||
});
|
||||
|
||||
let error = ocr(OcrRequest {
|
||||
model: "prebuilt-read",
|
||||
document: json!({
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/doc.pdf"
|
||||
}),
|
||||
api_key: Some("di-key"),
|
||||
api_base: Some(&format!("http://{addr}")),
|
||||
custom_llm_provider: "azure_ai/doc-intelligence",
|
||||
extra_headers: None,
|
||||
optional_params: Map::new(),
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
})
|
||||
.await
|
||||
.expect_err("redirect is returned to caller");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
CoreError::Http {
|
||||
status: 307,
|
||||
body: "".to_string()
|
||||
}
|
||||
);
|
||||
|
||||
let (post_request, followed_redirect) = server.await.expect("server task completes");
|
||||
assert!(
|
||||
post_request
|
||||
.to_ascii_lowercase()
|
||||
.contains("ocp-apim-subscription-key: di-key"),
|
||||
"{post_request}"
|
||||
);
|
||||
assert!(
|
||||
!followed_redirect,
|
||||
"credentialed DI POST followed redirect: {post_request}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn string_headers_rejects_non_string_values() {
|
||||
let headers = json!({
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use std::net::IpAddr;
|
||||
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
|
||||
|
|
@ -24,6 +24,13 @@ 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 AZURE_WIRE_SERVER_IP: Ipv4Addr = Ipv4Addr::new(168, 63, 129, 16);
|
||||
|
||||
struct ValidatedFetchDestination {
|
||||
url: Url,
|
||||
host: String,
|
||||
addresses: Vec<SocketAddr>,
|
||||
}
|
||||
|
||||
pub(super) fn truncate_error_body(body: &str) -> String {
|
||||
if body.chars().count() <= ERROR_BODY_MAX_CHARS {
|
||||
|
|
@ -111,31 +118,76 @@ 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 {
|
||||
fn is_private_ipv4(ip: Ipv4Addr, prefix: [u8; 4], prefix_len: u32) -> bool {
|
||||
let ip = u32::from(ip);
|
||||
let prefix = u32::from(Ipv4Addr::from(prefix));
|
||||
let mask = u32::MAX.checked_shl(32 - prefix_len).unwrap_or(0);
|
||||
(ip & mask) == (prefix & mask)
|
||||
}
|
||||
|
||||
fn is_global_ipv4(ip: Ipv4Addr) -> bool {
|
||||
!(is_private_ipv4(ip, [0, 0, 0, 0], 8)
|
||||
|| is_private_ipv4(ip, [10, 0, 0, 0], 8)
|
||||
|| is_private_ipv4(ip, [100, 64, 0, 0], 10)
|
||||
|| is_private_ipv4(ip, [127, 0, 0, 0], 8)
|
||||
|| is_private_ipv4(ip, [169, 254, 0, 0], 16)
|
||||
|| is_private_ipv4(ip, [172, 16, 0, 0], 12)
|
||||
|| is_private_ipv4(ip, [192, 0, 0, 0], 24)
|
||||
|| is_private_ipv4(ip, [192, 0, 2, 0], 24)
|
||||
|| is_private_ipv4(ip, [192, 88, 99, 0], 24)
|
||||
|| is_private_ipv4(ip, [192, 168, 0, 0], 16)
|
||||
|| is_private_ipv4(ip, [198, 18, 0, 0], 15)
|
||||
|| is_private_ipv4(ip, [198, 51, 100, 0], 24)
|
||||
|| is_private_ipv4(ip, [203, 0, 113, 0], 24)
|
||||
|| is_private_ipv4(ip, [224, 0, 0, 0], 4)
|
||||
|| is_private_ipv4(ip, [240, 0, 0, 0], 4))
|
||||
}
|
||||
|
||||
fn is_private_ipv6(ip: Ipv6Addr, prefix: u128, prefix_len: u32) -> bool {
|
||||
let ip = u128::from(ip);
|
||||
let mask = u128::MAX.checked_shl(128 - prefix_len).unwrap_or(0);
|
||||
(ip & mask) == (prefix & mask)
|
||||
}
|
||||
|
||||
fn is_global_ipv6(ip: Ipv6Addr) -> bool {
|
||||
if ip
|
||||
.to_ipv4_mapped()
|
||||
.or_else(|| ip.to_ipv4())
|
||||
.map(is_global_ipv4)
|
||||
.is_some_and(|is_global| !is_global)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
!(ip.is_unspecified()
|
||||
|| ip.is_loopback()
|
||||
|| ip.is_multicast()
|
||||
|| is_private_ipv6(ip, 0x0064_ff9b_0001_0000_0000_0000_0000_0000, 48)
|
||||
|| is_private_ipv6(ip, 0x0100_0000_0000_0000_0000_0000_0000_0000, 64)
|
||||
|| is_private_ipv6(ip, 0x2001_0000_0000_0000_0000_0000_0000_0000, 23)
|
||||
|| is_private_ipv6(ip, 0x2001_0db8_0000_0000_0000_0000_0000_0000, 32)
|
||||
|| is_private_ipv6(ip, 0x2002_0000_0000_0000_0000_0000_0000_0000, 16)
|
||||
|| is_private_ipv6(ip, 0xfc00_0000_0000_0000_0000_0000_0000_0000, 7)
|
||||
|| is_private_ipv6(ip, 0xfe80_0000_0000_0000_0000_0000_0000_0000, 10))
|
||||
}
|
||||
|
||||
fn is_cloud_metadata_exception(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)
|
||||
}
|
||||
IpAddr::V4(ip) => ip == AZURE_WIRE_SERVER_IP,
|
||||
IpAddr::V6(ip) => ip
|
||||
.to_ipv4_mapped()
|
||||
.or_else(|| ip.to_ipv4())
|
||||
.is_some_and(|ip| ip == AZURE_WIRE_SERVER_IP),
|
||||
}
|
||||
}
|
||||
|
||||
fn is_blocked_ip(ip: IpAddr) -> bool {
|
||||
if is_cloud_metadata_exception(ip) {
|
||||
return true;
|
||||
}
|
||||
match ip {
|
||||
IpAddr::V4(ip) => !is_global_ipv4(ip),
|
||||
IpAddr::V6(ip) => !is_global_ipv6(ip),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -145,36 +197,44 @@ fn blocked_url_error(url: &Url) -> CoreError {
|
|||
))
|
||||
}
|
||||
|
||||
async fn validate_safe_fetch_url(url: &Url) -> CoreResult<()> {
|
||||
async fn validate_safe_fetch_url(url: &Url) -> CoreResult<ValidatedFetchDestination> {
|
||||
if !matches!(url.scheme(), "http" | "https") {
|
||||
return Err(blocked_url_error(url));
|
||||
}
|
||||
|
||||
let host = url.host_str().ok_or_else(|| blocked_url_error(url))?;
|
||||
let port = url
|
||||
.port_or_known_default()
|
||||
.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(());
|
||||
return Ok(ValidatedFetchDestination {
|
||||
url: url.clone(),
|
||||
host: host.to_string(),
|
||||
addresses: vec![SocketAddr::new(ip, port)],
|
||||
});
|
||||
}
|
||||
|
||||
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;
|
||||
let mut safe_addresses = Vec::new();
|
||||
for address in addresses {
|
||||
saw_address = true;
|
||||
if is_blocked_ip(address.ip()) {
|
||||
return Err(blocked_url_error(url));
|
||||
}
|
||||
safe_addresses.push(address);
|
||||
}
|
||||
if !saw_address {
|
||||
if safe_addresses.is_empty() {
|
||||
return Err(blocked_url_error(url));
|
||||
}
|
||||
Ok(())
|
||||
Ok(ValidatedFetchDestination {
|
||||
url: url.clone(),
|
||||
host: host.to_string(),
|
||||
addresses: safe_addresses,
|
||||
})
|
||||
}
|
||||
|
||||
fn redirect_location(response: &reqwest::Response, url: &Url) -> CoreResult<Url> {
|
||||
|
|
@ -190,22 +250,23 @@ fn redirect_location(response: &reqwest::Response, url: &Url) -> CoreResult<Url>
|
|||
}
|
||||
|
||||
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 destination = validate_safe_fetch_url(¤t_url).await?;
|
||||
let client = reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.resolve_to_addrs(&destination.host, &destination.addresses)
|
||||
.build()
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
let response = client
|
||||
.get(current_url.clone())
|
||||
.get(destination.url.clone())
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
if !response.status().is_redirection() {
|
||||
return Ok((current_url, response));
|
||||
return Ok((destination.url, response));
|
||||
}
|
||||
current_url = redirect_location(&response, ¤t_url)?;
|
||||
}
|
||||
|
|
@ -401,16 +462,38 @@ pub(super) async fn poll_document_intelligence(
|
|||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::{TcpListener, TcpStream};
|
||||
|
||||
async fn read_http_headers(socket: &mut TcpStream) -> String {
|
||||
let mut request = Vec::new();
|
||||
let mut buffer = [0_u8; 1024];
|
||||
loop {
|
||||
let n = socket.read(&mut buffer).await.expect("reads request");
|
||||
if n == 0 {
|
||||
break;
|
||||
}
|
||||
request.extend_from_slice(&buffer[..n]);
|
||||
if request.windows(4).any(|window| window == b"\r\n\r\n") {
|
||||
break;
|
||||
}
|
||||
}
|
||||
String::from_utf8(request).expect("request is utf8")
|
||||
}
|
||||
|
||||
#[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("100.64.0.1".parse().unwrap()));
|
||||
assert!(is_blocked_ip("192.0.2.1".parse().unwrap()));
|
||||
assert!(is_blocked_ip("169.254.169.254".parse().unwrap()));
|
||||
assert!(is_blocked_ip("168.63.129.16".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:168.63.129.16".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()));
|
||||
|
|
@ -445,4 +528,62 @@ mod tests {
|
|||
|
||||
assert_eq!(transformed, document);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn document_intelligence_poll_does_not_follow_redirects_with_credentials() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("test listener binds");
|
||||
let addr = listener.local_addr().expect("listener has local addr");
|
||||
let operation_url = format!("http://{addr}/operations/1");
|
||||
let redirect_url = format!("http://{addr}/leak");
|
||||
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut poll_socket, _) = listener.accept().await.expect("accepts poll request");
|
||||
let poll_request = read_http_headers(&mut poll_socket).await;
|
||||
let poll_response = format!(
|
||||
"HTTP/1.1 302 Found\r\nlocation: {redirect_url}\r\ncontent-length: 0\r\nconnection: close\r\n\r\n"
|
||||
);
|
||||
poll_socket
|
||||
.write_all(poll_response.as_bytes())
|
||||
.await
|
||||
.expect("writes redirect response");
|
||||
|
||||
let second_connection =
|
||||
tokio::time::timeout(Duration::from_millis(200), listener.accept()).await;
|
||||
(poll_request, second_connection.is_ok())
|
||||
});
|
||||
|
||||
let error = poll_document_intelligence(
|
||||
&operation_url,
|
||||
&format!("http://{addr}/documentintelligence/documentModels/prebuilt-read:analyze"),
|
||||
&[(
|
||||
"Ocp-Apim-Subscription-Key".to_string(),
|
||||
"di-key".to_string(),
|
||||
)],
|
||||
Some(Duration::from_secs(5)),
|
||||
)
|
||||
.await
|
||||
.expect_err("poll redirect is returned to caller");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
CoreError::Http {
|
||||
status: 302,
|
||||
body: "".to_string()
|
||||
}
|
||||
);
|
||||
|
||||
let (poll_request, followed_redirect) = server.await.expect("server task completes");
|
||||
assert!(
|
||||
poll_request
|
||||
.to_ascii_lowercase()
|
||||
.contains("ocp-apim-subscription-key: di-key"),
|
||||
"{poll_request}"
|
||||
);
|
||||
assert!(
|
||||
!followed_redirect,
|
||||
"credentialed DI poll followed redirect: {poll_request}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue