mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
c264758eb8
commit
66582fb33f
7 changed files with 457 additions and 30 deletions
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -663,6 +663,7 @@ dependencies = [
|
|||
"subtle",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
"url",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(¤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<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() {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue