mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(ocr-security): pin validated addresses and add dedicated Rust input error
Close the DNS-rebinding TOCTOU by resolving and validating each candidate URL once, then pinning reqwest to exactly that validated SocketAddr set via a per-request resolver; every redirect hop is revalidated and repinned, so a connection can never use an unvalidated second DNS answer. Data-minimize public SSRF errors to a generic message with no URL, host, or query. Add a dedicated status-aware RustOcrInputError so only genuine client-input rejections map to BadRequestError; unrelated internal ValueError stays a server error. Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
parent
66582fb33f
commit
9aa73ad92d
6 changed files with 246 additions and 83 deletions
|
|
@ -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<Vec<SocketAddr>> {
|
||||
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<SocketAddr> = 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<Mutex<HashMap<String, Vec<SocketAddr>>>>;
|
||||
|
||||
#[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::<dyn std::error::Error + Send + Sync>::from(
|
||||
"OCR document host was not pinned to a validated address",
|
||||
)),
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -218,28 +233,32 @@ 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
|
||||
pin_validated_url(&candidate).await
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn fetch_with_redirects<V, Fut>(
|
||||
url: &str,
|
||||
validate: V,
|
||||
) -> CoreResult<(Url, reqwest::Response)>
|
||||
async fn fetch_with_redirects<P, Fut>(url: &str, pin: P) -> CoreResult<(Url, reqwest::Response)>
|
||||
where
|
||||
V: Fn(Url) -> Fut,
|
||||
Fut: std::future::Future<Output = CoreResult<()>>,
|
||||
P: Fn(Url) -> Fut,
|
||||
Fut: std::future::Future<Output = CoreResult<Vec<SocketAddr>>>,
|
||||
{
|
||||
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<Vec<u8>> {
|
||||
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!({
|
||||
|
|
|
|||
|
|
@ -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<Map<String, Value>>,
|
||||
|
|
@ -31,11 +34,11 @@ fn json_to_py(py: Python<'_>, value: Value) -> PyResult<Py<PyAny>> {
|
|||
|
||||
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::<RustOcrInputError>(),
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue