refactor(rust-gateway): harden OCR transport error contract and validation
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled

Generic fixed messages for all JSON/body parse failures so serde field, tag,
and value text can never cross the Axum response, with unit and route marker
tests covering query, base64, credential-path, and unknown-document-type inputs.

Map every CoreError through a typed Serialize error envelope with an exhaustive
match instead of serde_json::Value, dropping error.to_string() for the wrapping
variants and preserving valid upstream 4xx/5xx status codes.

Normalize declared upload MIME values case-insensitively before image
classification and data-URI construction, and treat generic values with
parameters as generic so magic-byte sniffing still runs.

Reject a multipart document form field as reserved and reject blank JSON and
multipart model before routing, keeping duplicate file, model, and timeout
rejection.

Move the transport unit tests into transport/tests.rs, and join or abort the
gateway and mock upstream handles in route tests while draining request bodies
before mock responses.

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-17 00:53:18 +00:00
parent f09a33a9e1
commit 7f526e0b08
6 changed files with 841 additions and 588 deletions

View file

@ -51,7 +51,9 @@ pub(crate) const OCR_RESERVED_PARAM_KEYS: &[&str] = &[
"vertex_ai_location",
];
pub(crate) const OCR_MULTIPART_UNIQUE_FIELDS: &[&str] = &["model", "timeout", "document"];
pub(crate) const OCR_MULTIPART_UNIQUE_FIELDS: &[&str] = &["model", "timeout"];
pub(crate) const OCR_MULTIPART_RESERVED_TEXT_FIELDS: &[&str] = &["document"];
pub(crate) const MAX_OCR_REQUEST_BYTES: usize = 100 * 1024 * 1024;

View file

@ -13,7 +13,7 @@ use axum::routing::post;
use axum::{Json, Router};
use litellm_core::error::CoreError;
use litellm_core::CoreResult;
use serde_json::{json, Value};
use serde::Serialize;
use crate::auth::RequireMasterKey;
use crate::constants::MAX_OCR_REQUEST_BYTES;
@ -69,7 +69,7 @@ async fn read_body(body: Body) -> CoreResult<Vec<u8>> {
to_bytes(body, MAX_OCR_REQUEST_BYTES)
.await
.map(|bytes| bytes.to_vec())
.map_err(|err| CoreError::InvalidRequest(format!("could not read request body: {err}")))
.map_err(|_| CoreError::InvalidRequest("could not read the request body".to_string()))
}
async fn parse_multipart(request: Request, state: &AppState) -> CoreResult<OcrCall> {
@ -123,22 +123,50 @@ async fn parse_multipart(request: Request, state: &AppState) -> CoreResult<OcrCa
transport::assemble_multipart_call(document, &text_fields)
}
#[derive(Serialize)]
struct ErrorEnvelope {
error: ErrorDetail,
}
#[derive(Serialize)]
struct ErrorDetail {
message: String,
#[serde(rename = "type")]
error_type: &'static str,
}
fn error_response(error: &CoreError) -> Response {
let (status, error_type, message) = match error {
CoreError::InvalidRequest(_)
| CoreError::InvalidType { .. }
| CoreError::MissingField(_)
| CoreError::InvalidProvider(_) => (
let (status, error_type, message): (StatusCode, &'static str, String) = match error {
CoreError::InvalidRequest(message) => (
StatusCode::BAD_REQUEST,
"invalid_request_error",
error.to_string(),
message.clone(),
),
CoreError::MissingField(field) => (
StatusCode::BAD_REQUEST,
"invalid_request_error",
format!("missing required field: {field}"),
),
CoreError::InvalidType { .. } => (
StatusCode::BAD_REQUEST,
"invalid_request_error",
"a field in the request body has the wrong type".to_string(),
),
CoreError::InvalidProvider(_) => (
StatusCode::BAD_REQUEST,
"invalid_request_error",
"the requested model resolves to an unsupported provider".to_string(),
),
CoreError::Auth(_) => (
StatusCode::UNAUTHORIZED,
"authentication_error",
"authentication failed".to_string(),
),
CoreError::Routing(_) => (StatusCode::NOT_FOUND, "not_found_error", error.to_string()),
CoreError::Routing(_) => (
StatusCode::NOT_FOUND,
"not_found_error",
"no deployment was found for the requested model".to_string(),
),
CoreError::Http { status, .. } => {
let status = StatusCode::from_u16(*status).unwrap_or(StatusCode::BAD_GATEWAY);
(
@ -161,14 +189,11 @@ fn error_response(error: &CoreError) -> Response {
"the OCR provider returned an unexpected response".to_string(),
),
};
(status, Json(error_body(&message, error_type))).into_response()
}
fn error_body(message: &str, error_type: &str) -> Value {
json!({
"error": {
"message": message,
"type": error_type,
}
})
let envelope = ErrorEnvelope {
error: ErrorDetail {
message,
error_type,
},
};
(status, Json(envelope)).into_response()
}

View file

@ -11,20 +11,12 @@ use crate::state::AppState;
const MASTER_KEY: &str = "sk-master-test";
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 2048];
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;
}
struct ServerGuard(tokio::task::JoinHandle<()>);
impl Drop for ServerGuard {
fn drop(&mut self) {
self.0.abort();
}
String::from_utf8_lossy(&request).into_owned()
}
async fn read_full_http_request(socket: &mut TcpStream) -> String {
@ -89,7 +81,7 @@ async fn spawn_mock_upstream() -> (String, tokio::task::JoinHandle<String>) {
let addr = listener.local_addr().expect("upstream addr");
let handle = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let request = read_http_request(&mut socket).await;
let request = read_full_http_request(&mut socket).await;
let body = r#"{"pages":[{"index":0,"markdown":"hello ocr"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
@ -105,14 +97,14 @@ async fn spawn_mock_upstream() -> (String, tokio::task::JoinHandle<String>) {
(format!("http://{addr}"), handle)
}
async fn spawn_mock_upstream_error() -> String {
async fn spawn_mock_upstream_error() -> (String, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("binds upstream");
let addr = listener.local_addr().expect("upstream addr");
tokio::spawn(async move {
let handle = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let _ = read_http_request(&mut socket).await;
let _ = read_full_http_request(&mut socket).await;
let body = r#"{"error":"invalid_api_key: sk-leaked-secret-value"}"#;
let response = format!(
"HTTP/1.1 403 Forbidden\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
@ -124,7 +116,7 @@ async fn spawn_mock_upstream_error() -> String {
.await
.expect("writes response");
});
format!("http://{addr}")
(format!("http://{addr}"), handle)
}
fn app_with_deployment(api_base: &str) -> axum::Router {
@ -150,22 +142,22 @@ fn app_with_params(model: &str, custom_llm_provider: Option<&str>, api_base: &st
crate::routes::app(state)
}
async fn serve(app: axum::Router) -> SocketAddr {
async fn serve(app: axum::Router) -> (SocketAddr, ServerGuard) {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("binds gateway");
let addr = listener.local_addr().expect("gateway addr");
tokio::spawn(async move {
axum::serve(listener, app).await.expect("serves");
let handle = tokio::spawn(async move {
let _ = axum::serve(listener, app).await;
});
addr
(addr, ServerGuard(handle))
}
#[tokio::test]
async fn json_document_url_returns_normalized_ocr_on_both_paths() {
for path in ["/v1/ocr", "/ocr"] {
let (upstream, upstream_handle) = spawn_mock_upstream().await;
let addr = serve(app_with_deployment(&upstream)).await;
let (addr, _server) = serve(app_with_deployment(&upstream)).await;
let response = reqwest::Client::new()
.post(format!("http://{addr}{path}"))
@ -196,7 +188,7 @@ async fn json_document_url_returns_normalized_ocr_on_both_paths() {
#[tokio::test]
async fn explicit_custom_llm_provider_resolves_model_without_prefix() {
let (upstream, upstream_handle) = spawn_mock_upstream().await;
let addr = serve(app_with_params(
let (addr, _server) = serve(app_with_params(
"mistral-ocr-latest",
Some("mistral"),
&upstream,
@ -226,7 +218,7 @@ async fn explicit_custom_llm_provider_resolves_model_without_prefix() {
#[tokio::test]
async fn multipart_upload_returns_normalized_ocr() {
let (upstream, upstream_handle) = spawn_mock_upstream().await;
let addr = serve(app_with_deployment(&upstream)).await;
let (addr, _server) = serve(app_with_deployment(&upstream)).await;
let form = reqwest::multipart::Form::new()
.text("model", "rust-ocr-mistral")
@ -258,7 +250,7 @@ async fn multipart_upload_returns_normalized_ocr() {
#[tokio::test]
async fn multipart_unnamed_octet_stream_pdf_is_sniffed() {
let (upstream, upstream_handle) = spawn_mock_upstream_capture().await;
let addr = serve(app_with_deployment(&upstream)).await;
let (addr, _server) = serve(app_with_deployment(&upstream)).await;
let form = reqwest::multipart::Form::new()
.text("model", "rust-ocr-mistral")
@ -288,7 +280,7 @@ async fn multipart_unnamed_octet_stream_pdf_is_sniffed() {
#[tokio::test]
async fn multipart_unnamed_octet_stream_image_is_sniffed() {
let (upstream, upstream_handle) = spawn_mock_upstream_capture().await;
let addr = serve(app_with_deployment(&upstream)).await;
let (addr, _server) = serve(app_with_deployment(&upstream)).await;
let png = vec![0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x01];
let form = reqwest::multipart::Form::new()
@ -318,8 +310,8 @@ async fn multipart_unnamed_octet_stream_image_is_sniffed() {
#[tokio::test]
async fn upstream_error_status_is_propagated_without_leaking_provider_body() {
let upstream = spawn_mock_upstream_error().await;
let addr = serve(app_with_deployment(&upstream)).await;
let (upstream, upstream_handle) = spawn_mock_upstream_error().await;
let (addr, _server) = serve(app_with_deployment(&upstream)).await;
let response = reqwest::Client::new()
.post(format!("http://{addr}/v1/ocr"))
@ -340,11 +332,12 @@ async fn upstream_error_status_is_propagated_without_leaking_provider_body() {
!message.contains("sk-leaked-secret-value") && !message.contains("invalid_api_key"),
"provider body must not leak into the public error: {message}"
);
upstream_handle.await.expect("upstream served");
}
#[tokio::test]
async fn missing_master_key_is_unauthorized() {
let addr = serve(app_with_deployment("http://127.0.0.1:1")).await;
let (addr, _server) = serve(app_with_deployment("http://127.0.0.1:1")).await;
let response = reqwest::Client::new()
.post(format!("http://{addr}/v1/ocr"))
.json(&serde_json::json!({
@ -359,7 +352,7 @@ async fn missing_master_key_is_unauthorized() {
#[tokio::test]
async fn unknown_model_is_not_found() {
let addr = serve(app_with_deployment("http://127.0.0.1:1")).await;
let (addr, _server) = serve(app_with_deployment("http://127.0.0.1:1")).await;
let response = reqwest::Client::new()
.post(format!("http://{addr}/v1/ocr"))
.bearer_auth(MASTER_KEY)
@ -377,7 +370,7 @@ async fn unknown_model_is_not_found() {
#[tokio::test]
async fn file_document_over_json_is_rejected() {
let addr = serve(app_with_deployment("http://127.0.0.1:1")).await;
let (addr, _server) = serve(app_with_deployment("http://127.0.0.1:1")).await;
let response = reqwest::Client::new()
.post(format!("http://{addr}/v1/ocr"))
.bearer_auth(MASTER_KEY)
@ -395,7 +388,7 @@ async fn file_document_over_json_is_rejected() {
#[tokio::test]
async fn json_reserved_control_param_is_rejected() {
let addr = serve(app_with_deployment("http://127.0.0.1:1")).await;
let (addr, _server) = serve(app_with_deployment("http://127.0.0.1:1")).await;
let response = reqwest::Client::new()
.post(format!("http://{addr}/v1/ocr"))
.bearer_auth(MASTER_KEY)
@ -419,7 +412,7 @@ async fn json_reserved_control_param_is_rejected() {
#[tokio::test]
async fn multipart_reserved_control_param_is_rejected() {
let addr = serve(app_with_deployment("http://127.0.0.1:1")).await;
let (addr, _server) = serve(app_with_deployment("http://127.0.0.1:1")).await;
let form = reqwest::multipart::Form::new()
.text("model", "rust-ocr-mistral")
.text("vertex_credentials", "/etc/gcp/service-account.json")
@ -444,7 +437,7 @@ async fn multipart_reserved_control_param_is_rejected() {
#[tokio::test]
async fn duplicate_file_multipart_is_rejected() {
let addr = serve(app_with_deployment("http://127.0.0.1:1")).await;
let (addr, _server) = serve(app_with_deployment("http://127.0.0.1:1")).await;
let form = reqwest::multipart::Form::new()
.text("model", "rust-ocr-mistral")
.part(
@ -475,7 +468,7 @@ async fn duplicate_file_multipart_is_rejected() {
#[tokio::test]
async fn non_positive_timeout_is_rejected() {
let addr = serve(app_with_deployment("http://127.0.0.1:1")).await;
let (addr, _server) = serve(app_with_deployment("http://127.0.0.1:1")).await;
let response = reqwest::Client::new()
.post(format!("http://{addr}/v1/ocr"))
.bearer_auth(MASTER_KEY)
@ -491,3 +484,72 @@ async fn non_positive_timeout_is_rejected() {
let body: Value = response.json().await.expect("json body");
assert_eq!(body["error"]["type"], "invalid_request_error");
}
#[tokio::test]
async fn malformed_json_response_never_echoes_attacker_text() {
let (addr, _server) = serve(app_with_deployment("http://127.0.0.1:1")).await;
let response = reqwest::Client::new()
.post(format!("http://{addr}/v1/ocr"))
.bearer_auth(MASTER_KEY)
.header("content-type", "application/json")
.body(r#"{"model":"m","document":{"type":"/etc/creds/CREDMARKER.json?token=QUERYMARKER"}}"#)
.send()
.await
.expect("request sent");
assert_eq!(response.status(), reqwest::StatusCode::BAD_REQUEST);
let body: Value = response.json().await.expect("json body");
assert_eq!(body["error"]["type"], "invalid_request_error");
let message = body["error"]["message"].as_str().expect("message string");
for needle in ["CREDMARKER", "QUERYMARKER", "/etc/creds"] {
assert!(
!message.contains(needle),
"parse error must not echo attacker text ({needle}): {message}"
);
}
}
#[tokio::test]
async fn blank_model_is_rejected_before_routing() {
let (addr, _server) = serve(app_with_deployment("http://127.0.0.1:1")).await;
let response = reqwest::Client::new()
.post(format!("http://{addr}/v1/ocr"))
.bearer_auth(MASTER_KEY)
.json(&serde_json::json!({
"model": " ",
"document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
}))
.send()
.await
.expect("request sent");
assert_eq!(response.status(), reqwest::StatusCode::BAD_REQUEST);
let body: Value = response.json().await.expect("json body");
assert_eq!(body["error"]["type"], "invalid_request_error");
}
#[tokio::test]
async fn multipart_document_form_field_is_rejected() {
let (addr, _server) = serve(app_with_deployment("http://127.0.0.1:1")).await;
let form = reqwest::multipart::Form::new()
.text("model", "rust-ocr-mistral")
.text(
"document",
r#"{"type":"document_url","document_url":"reducto://smuggled"}"#,
)
.part(
"file",
reqwest::multipart::Part::bytes(b"%PDF-1.7 minimal".to_vec())
.file_name("doc.pdf")
.mime_str("application/pdf")
.expect("mime"),
);
let response = reqwest::Client::new()
.post(format!("http://{addr}/v1/ocr"))
.bearer_auth(MASTER_KEY)
.multipart(form)
.send()
.await
.expect("request sent");
assert_eq!(response.status(), reqwest::StatusCode::BAD_REQUEST);
let body: Value = response.json().await.expect("json body");
assert_eq!(body["error"]["type"], "invalid_request_error");
}

View file

@ -1,532 +0,0 @@
use std::time::Duration;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use base64::Engine;
use litellm_core::error::CoreError;
use litellm_core::ocr::mime::sniff_mime;
use litellm_core::CoreResult;
use serde::Deserialize;
use serde_json::{json, Map, Value};
use crate::constants::{
DEFAULT_UPLOAD_MIME_TYPE, GENERIC_UPLOAD_MIME_TYPES, OCR_MULTIPART_UNIQUE_FIELDS,
OCR_RESERVED_PARAM_KEYS, OCR_UPLOAD_MIME_BY_EXTENSION,
};
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum OcrDocument {
DocumentUrl { document_url: String },
ImageUrl { image_url: String },
File {},
}
impl OcrDocument {
fn into_value(self) -> CoreResult<Value> {
let (field, url) = match self {
OcrDocument::DocumentUrl { document_url } => ("document_url", document_url),
OcrDocument::ImageUrl { image_url } => ("image_url", image_url),
OcrDocument::File {} => {
return Err(CoreError::InvalidRequest(
"document type 'file' is not supported through the JSON API; upload the \
file via multipart/form-data with a 'file' field, or use a 'document_url' \
or 'image_url' document type"
.to_string(),
))
}
};
Ok(json!({ "type": field, field: url }))
}
}
#[derive(Debug, Deserialize)]
pub struct OcrJsonRequest {
pub model: String,
pub document: OcrDocument,
#[serde(default)]
pub timeout: Option<f64>,
#[serde(flatten)]
pub optional_params: Map<String, Value>,
}
#[derive(Debug)]
pub struct OcrCall {
pub model: String,
pub document: Value,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
}
fn parse_timeout(seconds: Option<f64>) -> CoreResult<Option<Duration>> {
let Some(value) = seconds else {
return Ok(None);
};
if !value.is_finite() || value <= 0.0 {
return Err(CoreError::InvalidRequest(
"'timeout' must be a positive, finite number of seconds".to_string(),
));
}
Duration::try_from_secs_f64(value).map(Some).map_err(|_| {
CoreError::InvalidRequest("'timeout' is larger than the supported range".to_string())
})
}
fn reject_reserved_params<'a>(names: impl IntoIterator<Item = &'a str>) -> CoreResult<()> {
let reserved = names.into_iter().find_map(|name| {
OCR_RESERVED_PARAM_KEYS
.iter()
.copied()
.find(|candidate| *candidate == name)
});
match reserved {
None => Ok(()),
Some(reserved) => Err(CoreError::InvalidRequest(format!(
"the '{reserved}' parameter is not accepted on an OCR request; deployment \
credentials, routing, and headers are server-controlled"
))),
}
}
pub fn parse_json_body(body: &[u8]) -> CoreResult<OcrCall> {
if body.is_empty() {
return Err(CoreError::InvalidRequest(
"empty request body; send a JSON body with 'model' and 'document', or use \
multipart/form-data for file uploads"
.to_string(),
));
}
let request: OcrJsonRequest = serde_json::from_slice(body)
.map_err(|err| CoreError::InvalidRequest(format!("invalid OCR request body: {err}")))?;
reject_reserved_params(request.optional_params.keys().map(String::as_str))?;
Ok(OcrCall {
model: request.model,
document: request.document.into_value()?,
optional_params: request.optional_params,
timeout: parse_timeout(request.timeout)?,
})
}
fn mime_from_filename(filename: &str) -> Option<&'static str> {
let extension = filename
.rsplit_once('.')
.map(|(_, ext)| ext.to_ascii_lowercase())?;
OCR_UPLOAD_MIME_BY_EXTENSION
.iter()
.find(|(candidate, _)| *candidate == extension)
.map(|(_, mime)| *mime)
}
fn is_generic_upload_mime(value: &str) -> bool {
GENERIC_UPLOAD_MIME_TYPES
.iter()
.any(|generic| value.eq_ignore_ascii_case(generic))
}
fn resolve_upload_mime(bytes: &[u8], content_type: Option<&str>, filename: Option<&str>) -> String {
let declared = content_type
.and_then(|value| value.split(';').next())
.map(str::trim)
.filter(|value| !value.is_empty() && !is_generic_upload_mime(value));
if let Some(declared) = declared {
return declared.to_string();
}
if let Some(sniffed) = sniff_mime(bytes) {
return sniffed.to_string();
}
filename
.and_then(mime_from_filename)
.unwrap_or(DEFAULT_UPLOAD_MIME_TYPE)
.to_string()
}
pub fn build_upload_document(
bytes: Vec<u8>,
filename: Option<&str>,
content_type: Option<&str>,
) -> CoreResult<Value> {
if bytes.is_empty() {
return Err(CoreError::InvalidRequest(
"uploaded file is empty".to_string(),
));
}
let mime = resolve_upload_mime(&bytes, content_type, filename);
let data_uri = format!("data:{mime};base64,{}", BASE64_STANDARD.encode(&bytes));
let field = if mime.starts_with("image/") {
"image_url"
} else {
"document_url"
};
Ok(json!({ "type": field, field: data_uri }))
}
fn coerce_form_field(value: &str) -> Value {
serde_json::from_str(value).unwrap_or_else(|_| Value::String(value.to_string()))
}
fn reject_duplicate_unique_fields(text_fields: &[(String, String)]) -> CoreResult<()> {
let duplicate = OCR_MULTIPART_UNIQUE_FIELDS
.iter()
.copied()
.find(|field| text_fields.iter().filter(|(name, _)| name == field).count() > 1);
match duplicate {
None => Ok(()),
Some(field) => Err(CoreError::InvalidRequest(format!(
"the '{field}' field must appear at most once in a multipart OCR request"
))),
}
}
pub fn assemble_multipart_call(
document: Value,
text_fields: &[(String, String)],
) -> CoreResult<OcrCall> {
reject_reserved_params(text_fields.iter().map(|(name, _)| name.as_str()))?;
reject_duplicate_unique_fields(text_fields)?;
let model = text_fields
.iter()
.find(|(name, _)| name == "model")
.map(|(_, value)| value.clone())
.ok_or_else(|| {
CoreError::InvalidRequest(
"multipart OCR request must include a 'model' form field".to_string(),
)
})?;
let timeout = match text_fields.iter().find(|(name, _)| name == "timeout") {
Some((_, raw)) => {
let seconds = raw.parse::<f64>().map_err(|_| {
CoreError::InvalidRequest(
"'timeout' form field must be a number of seconds".to_string(),
)
})?;
parse_timeout(Some(seconds))?
}
None => None,
};
let optional_params: Map<String, Value> = text_fields
.iter()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "timeout" | "file" | "document"))
.map(|(name, value)| (name.clone(), coerce_form_field(value)))
.collect();
Ok(OcrCall {
model,
document,
optional_params,
timeout,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_document_url_and_flattens_provider_params() {
let call = parse_json_body(
br#"{"model":"rust-ocr","document":{"type":"document_url","document_url":"https://x/doc.pdf"},"include_image_base64":true}"#,
)
.expect("valid body parses");
assert_eq!(call.model, "rust-ocr");
assert_eq!(call.document["type"], "document_url");
assert_eq!(call.document["document_url"], "https://x/doc.pdf");
assert_eq!(
call.optional_params["include_image_base64"],
Value::Bool(true)
);
assert!(call.timeout.is_none());
}
#[test]
fn parses_image_url_document() {
let call = parse_json_body(
br#"{"model":"m","document":{"type":"image_url","image_url":"https://x/i.png"}}"#,
)
.expect("valid body parses");
assert_eq!(call.document["type"], "image_url");
assert_eq!(call.document["image_url"], "https://x/i.png");
}
#[test]
fn extracts_timeout_and_keeps_it_out_of_provider_params() {
let call = parse_json_body(
br#"{"model":"m","document":{"type":"document_url","document_url":"https://x"},"timeout":12.5}"#,
)
.expect("valid body parses");
assert_eq!(call.timeout, Some(Duration::from_secs_f64(12.5)));
assert!(!call.optional_params.contains_key("timeout"));
}
#[test]
fn rejects_file_document_over_json() {
let err =
parse_json_body(br#"{"model":"m","document":{"type":"file","file":"/etc/passwd"}}"#)
.expect_err("file type rejected");
match err {
CoreError::InvalidRequest(message) => assert!(message.contains("multipart/form-data")),
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
#[test]
fn preserves_reducto_file_id_over_json() {
let call = parse_json_body(
br#"{"model":"m","document":{"type":"document_url","document_url":"reducto://abc123"}}"#,
)
.expect("reducto id preserved");
assert_eq!(call.document["type"], "document_url");
assert_eq!(call.document["document_url"], "reducto://abc123");
}
#[test]
fn empty_body_is_rejected() {
assert!(matches!(
parse_json_body(b""),
Err(CoreError::InvalidRequest(_))
));
}
#[test]
fn upload_prefers_declared_content_type() {
let document = build_upload_document(
b"%PDF-1.4".to_vec(),
Some("scan.bin"),
Some("application/pdf"),
)
.expect("builds document");
assert_eq!(document["type"], "document_url");
assert!(document["document_url"]
.as_str()
.expect("data uri")
.starts_with("data:application/pdf;base64,"));
}
#[test]
fn upload_sniffs_pdf_from_bytes_when_unnamed_octet_stream() {
let document = build_upload_document(
b"%PDF-1.7 minimal".to_vec(),
None,
Some("application/octet-stream"),
)
.expect("builds document");
assert_eq!(document["type"], "document_url");
assert!(document["document_url"]
.as_str()
.expect("data uri")
.starts_with("data:application/pdf;base64,"));
}
#[test]
fn upload_sniffs_png_from_bytes_when_unnamed_and_no_content_type() {
let document = build_upload_document(
vec![0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A, 0x00],
None,
None,
)
.expect("builds document");
assert_eq!(document["type"], "image_url");
assert!(document["image_url"]
.as_str()
.expect("data uri")
.starts_with("data:image/png;base64,"));
}
#[test]
fn upload_infers_mime_from_filename_when_octet_stream() {
let document = build_upload_document(
vec![0x89, b'P', b'N', b'G'],
Some("photo.PNG"),
Some("application/octet-stream"),
)
.expect("builds document");
assert_eq!(document["type"], "image_url");
assert!(document["image_url"]
.as_str()
.expect("data uri")
.starts_with("data:image/png;base64,"));
}
#[test]
fn upload_falls_back_to_octet_stream() {
let document = build_upload_document(vec![1, 2, 3], Some("data.unknown"), None)
.expect("builds document");
assert_eq!(document["type"], "document_url");
assert!(document["document_url"]
.as_str()
.expect("data uri")
.starts_with("data:application/octet-stream;base64,"));
}
#[test]
fn empty_upload_is_rejected() {
assert!(matches!(
build_upload_document(Vec::new(), Some("a.pdf"), Some("application/pdf")),
Err(CoreError::InvalidRequest(_))
));
}
#[test]
fn multipart_extracts_model_timeout_and_json_params() {
let document =
json!({"type": "document_url", "document_url": "data:application/pdf;base64,AA=="});
let fields = vec![
("model".to_string(), "rust-ocr".to_string()),
("timeout".to_string(), "30".to_string()),
("pages".to_string(), "[0,1,2]".to_string()),
("id".to_string(), "abc".to_string()),
];
let call = assemble_multipart_call(document, &fields).expect("assembles call");
assert_eq!(call.model, "rust-ocr");
assert_eq!(call.timeout, Some(Duration::from_secs(30)));
assert_eq!(call.optional_params["pages"], json!([0, 1, 2]));
assert_eq!(call.optional_params["id"], Value::String("abc".to_string()));
assert!(!call.optional_params.contains_key("timeout"));
assert!(!call.optional_params.contains_key("model"));
}
#[test]
fn multipart_requires_model() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let err = assemble_multipart_call(document, &[]).expect_err("model required");
assert!(matches!(err, CoreError::InvalidRequest(_)));
}
#[test]
fn json_rejects_reserved_control_params() {
for reserved in [
"api_key",
"api_base",
"custom_llm_provider",
"extra_headers",
"vertex_credentials",
"vertex_ai_credentials",
"vertex_project",
"vertex_ai_project",
"vertex_location",
"vertex_ai_location",
] {
let body = format!(
r#"{{"model":"m","document":{{"type":"document_url","document_url":"https://x"}},"{reserved}":"attacker"}}"#
);
let err = parse_json_body(body.as_bytes())
.expect_err("reserved control param must be rejected");
match err {
CoreError::InvalidRequest(message) => {
assert!(
message.contains(reserved),
"names the rejected key: {message}"
);
assert!(
!message.contains("attacker"),
"must not echo the attacker value: {message}"
);
}
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
}
#[test]
fn multipart_rejects_reserved_control_params() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let fields = vec![
("model".to_string(), "m".to_string()),
(
"vertex_credentials".to_string(),
"/etc/gcp/service-account.json".to_string(),
),
];
let err = assemble_multipart_call(document, &fields).expect_err("reserved param rejected");
match err {
CoreError::InvalidRequest(message) => {
assert!(message.contains("vertex_credentials"));
assert!(
!message.contains("service-account"),
"must not echo the attacker value: {message}"
);
}
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
#[test]
fn json_rejects_non_positive_timeout() {
for timeout in ["0", "-1", "-0.5", "1e300"] {
let body = format!(
r#"{{"model":"m","document":{{"type":"document_url","document_url":"https://x"}},"timeout":{timeout}}}"#
);
assert!(
matches!(
parse_json_body(body.as_bytes()),
Err(CoreError::InvalidRequest(_))
),
"timeout {timeout} must be rejected"
);
}
}
#[test]
fn multipart_rejects_non_positive_and_non_finite_timeout() {
let document = json!({"type": "document_url", "document_url": "data:x"});
for timeout in ["0", "-3", "inf", "-inf", "NaN", "1e300"] {
let fields = vec![
("model".to_string(), "m".to_string()),
("timeout".to_string(), timeout.to_string()),
];
assert!(
matches!(
assemble_multipart_call(document.clone(), &fields),
Err(CoreError::InvalidRequest(_))
),
"timeout {timeout} must be rejected"
);
}
}
#[test]
fn multipart_rejects_non_numeric_timeout_without_echoing_value() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let fields = vec![
("model".to_string(), "m".to_string()),
("timeout".to_string(), "not-a-number".to_string()),
];
let err =
assemble_multipart_call(document, &fields).expect_err("non-numeric timeout rejected");
match err {
CoreError::InvalidRequest(message) => assert!(
!message.contains("not-a-number"),
"must not echo the attacker value: {message}"
),
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
#[test]
fn multipart_rejects_duplicate_model_field() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let fields = vec![
("model".to_string(), "first".to_string()),
("model".to_string(), "second".to_string()),
];
let err = assemble_multipart_call(document, &fields).expect_err("duplicate model rejected");
assert!(matches!(err, CoreError::InvalidRequest(_)));
}
#[test]
fn upload_treats_binary_octet_stream_as_generic() {
let document = build_upload_document(
b"%PDF-1.7 minimal".to_vec(),
None,
Some("Binary/Octet-Stream"),
)
.expect("builds document");
assert!(document["document_url"]
.as_str()
.expect("data uri")
.starts_with("data:application/pdf;base64,"));
}
}

View file

@ -0,0 +1,255 @@
use std::time::Duration;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use base64::Engine;
use litellm_core::error::CoreError;
use litellm_core::ocr::mime::sniff_mime;
use litellm_core::CoreResult;
use serde::Deserialize;
use serde_json::{json, Map, Value};
use crate::constants::{
DEFAULT_UPLOAD_MIME_TYPE, GENERIC_UPLOAD_MIME_TYPES, OCR_MULTIPART_RESERVED_TEXT_FIELDS,
OCR_MULTIPART_UNIQUE_FIELDS, OCR_RESERVED_PARAM_KEYS, OCR_UPLOAD_MIME_BY_EXTENSION,
};
#[cfg(test)]
mod tests;
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum OcrDocument {
DocumentUrl { document_url: String },
ImageUrl { image_url: String },
File {},
}
impl OcrDocument {
fn into_value(self) -> CoreResult<Value> {
let (field, url) = match self {
OcrDocument::DocumentUrl { document_url } => ("document_url", document_url),
OcrDocument::ImageUrl { image_url } => ("image_url", image_url),
OcrDocument::File {} => {
return Err(CoreError::InvalidRequest(
"document type 'file' is not supported through the JSON API; upload the \
file via multipart/form-data with a 'file' field, or use a 'document_url' \
or 'image_url' document type"
.to_string(),
))
}
};
Ok(json!({ "type": field, field: url }))
}
}
#[derive(Debug, Deserialize)]
pub struct OcrJsonRequest {
pub model: String,
pub document: OcrDocument,
#[serde(default)]
pub timeout: Option<f64>,
#[serde(flatten)]
pub optional_params: Map<String, Value>,
}
#[derive(Debug)]
pub struct OcrCall {
pub model: String,
pub document: Value,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
}
fn parse_timeout(seconds: Option<f64>) -> CoreResult<Option<Duration>> {
let Some(value) = seconds else {
return Ok(None);
};
if !value.is_finite() || value <= 0.0 {
return Err(CoreError::InvalidRequest(
"'timeout' must be a positive, finite number of seconds".to_string(),
));
}
Duration::try_from_secs_f64(value).map(Some).map_err(|_| {
CoreError::InvalidRequest("'timeout' is larger than the supported range".to_string())
})
}
fn validate_model(model: &str) -> CoreResult<()> {
if model.trim().is_empty() {
return Err(CoreError::InvalidRequest(
"'model' must be a non-empty string".to_string(),
));
}
Ok(())
}
fn reject_reserved_params<'a>(names: impl IntoIterator<Item = &'a str>) -> CoreResult<()> {
let reserved = names.into_iter().find_map(|name| {
OCR_RESERVED_PARAM_KEYS
.iter()
.copied()
.find(|candidate| *candidate == name)
});
match reserved {
None => Ok(()),
Some(reserved) => Err(CoreError::InvalidRequest(format!(
"the '{reserved}' parameter is not accepted on an OCR request; deployment \
credentials, routing, and headers are server-controlled"
))),
}
}
pub fn parse_json_body(body: &[u8]) -> CoreResult<OcrCall> {
if body.is_empty() {
return Err(CoreError::InvalidRequest(
"empty request body; send a JSON body with 'model' and 'document', or use \
multipart/form-data for file uploads"
.to_string(),
));
}
let request: OcrJsonRequest = serde_json::from_slice(body).map_err(|_| {
CoreError::InvalidRequest(
"request body is not a valid OCR request; expected JSON with 'model' and 'document'"
.to_string(),
)
})?;
validate_model(&request.model)?;
reject_reserved_params(request.optional_params.keys().map(String::as_str))?;
Ok(OcrCall {
model: request.model,
document: request.document.into_value()?,
optional_params: request.optional_params,
timeout: parse_timeout(request.timeout)?,
})
}
fn mime_from_filename(filename: &str) -> Option<&'static str> {
let extension = filename
.rsplit_once('.')
.map(|(_, ext)| ext.to_ascii_lowercase())?;
OCR_UPLOAD_MIME_BY_EXTENSION
.iter()
.find(|(candidate, _)| *candidate == extension)
.map(|(_, mime)| *mime)
}
fn is_generic_upload_mime(value: &str) -> bool {
GENERIC_UPLOAD_MIME_TYPES
.iter()
.any(|generic| value.eq_ignore_ascii_case(generic))
}
fn resolve_upload_mime(bytes: &[u8], content_type: Option<&str>, filename: Option<&str>) -> String {
let declared = content_type
.and_then(|value| value.split(';').next())
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_ascii_lowercase)
.filter(|value| !is_generic_upload_mime(value));
if let Some(declared) = declared {
return declared;
}
if let Some(sniffed) = sniff_mime(bytes) {
return sniffed.to_string();
}
filename
.and_then(mime_from_filename)
.unwrap_or(DEFAULT_UPLOAD_MIME_TYPE)
.to_string()
}
pub fn build_upload_document(
bytes: Vec<u8>,
filename: Option<&str>,
content_type: Option<&str>,
) -> CoreResult<Value> {
if bytes.is_empty() {
return Err(CoreError::InvalidRequest(
"uploaded file is empty".to_string(),
));
}
let mime = resolve_upload_mime(&bytes, content_type, filename);
let data_uri = format!("data:{mime};base64,{}", BASE64_STANDARD.encode(&bytes));
let field = if mime.starts_with("image/") {
"image_url"
} else {
"document_url"
};
Ok(json!({ "type": field, field: data_uri }))
}
fn coerce_form_field(value: &str) -> Value {
serde_json::from_str(value).unwrap_or_else(|_| Value::String(value.to_string()))
}
fn reject_reserved_text_fields(text_fields: &[(String, String)]) -> CoreResult<()> {
let reserved = OCR_MULTIPART_RESERVED_TEXT_FIELDS
.iter()
.copied()
.find(|field| text_fields.iter().any(|(name, _)| name == field));
match reserved {
None => Ok(()),
Some(field) => Err(CoreError::InvalidRequest(format!(
"the '{field}' field is not accepted as a form field; upload the document through the \
'file' field"
))),
}
}
fn reject_duplicate_unique_fields(text_fields: &[(String, String)]) -> CoreResult<()> {
let duplicate = OCR_MULTIPART_UNIQUE_FIELDS
.iter()
.copied()
.find(|field| text_fields.iter().filter(|(name, _)| name == field).count() > 1);
match duplicate {
None => Ok(()),
Some(field) => Err(CoreError::InvalidRequest(format!(
"the '{field}' field must appear at most once in a multipart OCR request"
))),
}
}
pub fn assemble_multipart_call(
document: Value,
text_fields: &[(String, String)],
) -> CoreResult<OcrCall> {
reject_reserved_params(text_fields.iter().map(|(name, _)| name.as_str()))?;
reject_reserved_text_fields(text_fields)?;
reject_duplicate_unique_fields(text_fields)?;
let model = text_fields
.iter()
.find(|(name, _)| name == "model")
.map(|(_, value)| value.clone())
.ok_or_else(|| {
CoreError::InvalidRequest(
"multipart OCR request must include a 'model' form field".to_string(),
)
})?;
validate_model(&model)?;
let timeout = match text_fields.iter().find(|(name, _)| name == "timeout") {
Some((_, raw)) => {
let seconds = raw.parse::<f64>().map_err(|_| {
CoreError::InvalidRequest(
"'timeout' form field must be a number of seconds".to_string(),
)
})?;
parse_timeout(Some(seconds))?
}
None => None,
};
let optional_params: Map<String, Value> = text_fields
.iter()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "timeout" | "file" | "document"))
.map(|(name, value)| (name.clone(), coerce_form_field(value)))
.collect();
Ok(OcrCall {
model,
document,
optional_params,
timeout,
})
}

View file

@ -0,0 +1,441 @@
use super::*;
#[test]
fn parses_document_url_and_flattens_provider_params() {
let call = parse_json_body(
br#"{"model":"rust-ocr","document":{"type":"document_url","document_url":"https://x/doc.pdf"},"include_image_base64":true}"#,
)
.expect("valid body parses");
assert_eq!(call.model, "rust-ocr");
assert_eq!(call.document["type"], "document_url");
assert_eq!(call.document["document_url"], "https://x/doc.pdf");
assert_eq!(
call.optional_params["include_image_base64"],
Value::Bool(true)
);
assert!(call.timeout.is_none());
}
#[test]
fn parses_image_url_document() {
let call = parse_json_body(
br#"{"model":"m","document":{"type":"image_url","image_url":"https://x/i.png"}}"#,
)
.expect("valid body parses");
assert_eq!(call.document["type"], "image_url");
assert_eq!(call.document["image_url"], "https://x/i.png");
}
#[test]
fn extracts_timeout_and_keeps_it_out_of_provider_params() {
let call = parse_json_body(
br#"{"model":"m","document":{"type":"document_url","document_url":"https://x"},"timeout":12.5}"#,
)
.expect("valid body parses");
assert_eq!(call.timeout, Some(Duration::from_secs_f64(12.5)));
assert!(!call.optional_params.contains_key("timeout"));
}
#[test]
fn rejects_file_document_over_json() {
let err = parse_json_body(br#"{"model":"m","document":{"type":"file","file":"/etc/passwd"}}"#)
.expect_err("file type rejected");
match err {
CoreError::InvalidRequest(message) => assert!(message.contains("multipart/form-data")),
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
#[test]
fn preserves_reducto_file_id_over_json() {
let call = parse_json_body(
br#"{"model":"m","document":{"type":"document_url","document_url":"reducto://abc123"}}"#,
)
.expect("reducto id preserved");
assert_eq!(call.document["type"], "document_url");
assert_eq!(call.document["document_url"], "reducto://abc123");
}
#[test]
fn empty_body_is_rejected() {
assert!(matches!(
parse_json_body(b""),
Err(CoreError::InvalidRequest(_))
));
}
#[test]
fn malformed_json_error_is_generic_and_never_echoes_attacker_text() {
let markers = [
(
"query",
br#"{"model":"m","document":"https://x?secret=QUERYMARKERZZZ"}"#.to_vec(),
),
(
"base64",
br#"{"model":"m","document":{"type":"document_url","document_url":"https://x"},"timeout":"BASE64MARKERzz=="}"#.to_vec(),
),
(
"credential-path",
br#"{"model":"m","document":{"type":"/etc/creds/CREDMARKER.json"}}"#.to_vec(),
),
(
"unknown-type",
br#"{"model":"m","document":{"type":"UNKNOWNTYPEZZZ"}}"#.to_vec(),
),
];
for (label, body) in markers {
let err = parse_json_body(&body).expect_err("malformed body rejected");
let CoreError::InvalidRequest(message) = err else {
panic!("{label}: expected InvalidRequest, got {err:?}");
};
for needle in [
"QUERYMARKERZZZ",
"BASE64MARKERzz==",
"CREDMARKER",
"UNKNOWNTYPEZZZ",
"/etc/creds",
] {
assert!(
!message.contains(needle),
"{label}: parse error must not echo attacker text ({needle}): {message}"
);
}
assert_eq!(
message,
"request body is not a valid OCR request; expected JSON with 'model' and 'document'",
"{label}: parse errors must be a single fixed message"
);
}
}
#[test]
fn json_rejects_blank_model() {
for model in ["", " ", "\t\n"] {
let body = format!(
r#"{{"model":{model:?},"document":{{"type":"document_url","document_url":"https://x"}}}}"#
);
let err = parse_json_body(body.as_bytes()).expect_err("blank model rejected");
match err {
CoreError::InvalidRequest(message) => assert!(message.contains("model")),
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
}
#[test]
fn upload_prefers_declared_content_type() {
let document = build_upload_document(
b"%PDF-1.4".to_vec(),
Some("scan.bin"),
Some("application/pdf"),
)
.expect("builds document");
assert_eq!(document["type"], "document_url");
assert!(document["document_url"]
.as_str()
.expect("data uri")
.starts_with("data:application/pdf;base64,"));
}
#[test]
fn upload_normalizes_declared_mime_case_before_image_classification() {
let document = build_upload_document(
vec![0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A],
Some("scan.bin"),
Some("Image/PNG"),
)
.expect("builds document");
assert_eq!(document["type"], "image_url");
assert!(document["image_url"]
.as_str()
.expect("data uri")
.starts_with("data:image/png;base64,"));
}
#[test]
fn upload_treats_generic_with_parameters_as_generic() {
let document = build_upload_document(
b"%PDF-1.7 minimal".to_vec(),
None,
Some("application/octet-stream; charset=binary"),
)
.expect("builds document");
assert_eq!(document["type"], "document_url");
assert!(document["document_url"]
.as_str()
.expect("data uri")
.starts_with("data:application/pdf;base64,"));
}
#[test]
fn upload_sniffs_pdf_from_bytes_when_unnamed_octet_stream() {
let document = build_upload_document(
b"%PDF-1.7 minimal".to_vec(),
None,
Some("application/octet-stream"),
)
.expect("builds document");
assert_eq!(document["type"], "document_url");
assert!(document["document_url"]
.as_str()
.expect("data uri")
.starts_with("data:application/pdf;base64,"));
}
#[test]
fn upload_sniffs_png_from_bytes_when_unnamed_and_no_content_type() {
let document = build_upload_document(
vec![0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A, 0x00],
None,
None,
)
.expect("builds document");
assert_eq!(document["type"], "image_url");
assert!(document["image_url"]
.as_str()
.expect("data uri")
.starts_with("data:image/png;base64,"));
}
#[test]
fn upload_infers_mime_from_filename_when_octet_stream() {
let document = build_upload_document(
vec![0x89, b'P', b'N', b'G'],
Some("photo.PNG"),
Some("application/octet-stream"),
)
.expect("builds document");
assert_eq!(document["type"], "image_url");
assert!(document["image_url"]
.as_str()
.expect("data uri")
.starts_with("data:image/png;base64,"));
}
#[test]
fn upload_falls_back_to_octet_stream() {
let document =
build_upload_document(vec![1, 2, 3], Some("data.unknown"), None).expect("builds document");
assert_eq!(document["type"], "document_url");
assert!(document["document_url"]
.as_str()
.expect("data uri")
.starts_with("data:application/octet-stream;base64,"));
}
#[test]
fn empty_upload_is_rejected() {
assert!(matches!(
build_upload_document(Vec::new(), Some("a.pdf"), Some("application/pdf")),
Err(CoreError::InvalidRequest(_))
));
}
#[test]
fn multipart_extracts_model_timeout_and_json_params() {
let document =
json!({"type": "document_url", "document_url": "data:application/pdf;base64,AA=="});
let fields = vec![
("model".to_string(), "rust-ocr".to_string()),
("timeout".to_string(), "30".to_string()),
("pages".to_string(), "[0,1,2]".to_string()),
("id".to_string(), "abc".to_string()),
];
let call = assemble_multipart_call(document, &fields).expect("assembles call");
assert_eq!(call.model, "rust-ocr");
assert_eq!(call.timeout, Some(Duration::from_secs(30)));
assert_eq!(call.optional_params["pages"], json!([0, 1, 2]));
assert_eq!(call.optional_params["id"], Value::String("abc".to_string()));
assert!(!call.optional_params.contains_key("timeout"));
assert!(!call.optional_params.contains_key("model"));
}
#[test]
fn multipart_requires_model() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let err = assemble_multipart_call(document, &[]).expect_err("model required");
assert!(matches!(err, CoreError::InvalidRequest(_)));
}
#[test]
fn multipart_rejects_blank_model() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let fields = vec![("model".to_string(), " ".to_string())];
let err = assemble_multipart_call(document, &fields).expect_err("blank model rejected");
match err {
CoreError::InvalidRequest(message) => assert!(message.contains("model")),
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
#[test]
fn multipart_rejects_document_form_field() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let fields = vec![
("model".to_string(), "m".to_string()),
(
"document".to_string(),
r#"{"type":"document_url","document_url":"reducto://smuggled"}"#.to_string(),
),
];
let err = assemble_multipart_call(document, &fields).expect_err("document field rejected");
match err {
CoreError::InvalidRequest(message) => {
assert!(message.contains("document"));
assert!(
!message.contains("smuggled"),
"must not echo the attacker value: {message}"
);
}
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
#[test]
fn json_rejects_reserved_control_params() {
for reserved in [
"api_key",
"api_base",
"custom_llm_provider",
"extra_headers",
"vertex_credentials",
"vertex_ai_credentials",
"vertex_project",
"vertex_ai_project",
"vertex_location",
"vertex_ai_location",
] {
let body = format!(
r#"{{"model":"m","document":{{"type":"document_url","document_url":"https://x"}},"{reserved}":"attacker"}}"#
);
let err =
parse_json_body(body.as_bytes()).expect_err("reserved control param must be rejected");
match err {
CoreError::InvalidRequest(message) => {
assert!(
message.contains(reserved),
"names the rejected key: {message}"
);
assert!(
!message.contains("attacker"),
"must not echo the attacker value: {message}"
);
}
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
}
#[test]
fn multipart_rejects_reserved_control_params() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let fields = vec![
("model".to_string(), "m".to_string()),
(
"vertex_credentials".to_string(),
"/etc/gcp/service-account.json".to_string(),
),
];
let err = assemble_multipart_call(document, &fields).expect_err("reserved param rejected");
match err {
CoreError::InvalidRequest(message) => {
assert!(message.contains("vertex_credentials"));
assert!(
!message.contains("service-account"),
"must not echo the attacker value: {message}"
);
}
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
#[test]
fn json_rejects_non_positive_timeout() {
for timeout in ["0", "-1", "-0.5", "1e300"] {
let body = format!(
r#"{{"model":"m","document":{{"type":"document_url","document_url":"https://x"}},"timeout":{timeout}}}"#
);
assert!(
matches!(
parse_json_body(body.as_bytes()),
Err(CoreError::InvalidRequest(_))
),
"timeout {timeout} must be rejected"
);
}
}
#[test]
fn multipart_rejects_non_positive_and_non_finite_timeout() {
let document = json!({"type": "document_url", "document_url": "data:x"});
for timeout in ["0", "-3", "inf", "-inf", "NaN", "1e300"] {
let fields = vec![
("model".to_string(), "m".to_string()),
("timeout".to_string(), timeout.to_string()),
];
assert!(
matches!(
assemble_multipart_call(document.clone(), &fields),
Err(CoreError::InvalidRequest(_))
),
"timeout {timeout} must be rejected"
);
}
}
#[test]
fn multipart_rejects_non_numeric_timeout_without_echoing_value() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let fields = vec![
("model".to_string(), "m".to_string()),
("timeout".to_string(), "not-a-number".to_string()),
];
let err = assemble_multipart_call(document, &fields).expect_err("non-numeric timeout rejected");
match err {
CoreError::InvalidRequest(message) => assert!(
!message.contains("not-a-number"),
"must not echo the attacker value: {message}"
),
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
#[test]
fn multipart_rejects_duplicate_model_field() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let fields = vec![
("model".to_string(), "first".to_string()),
("model".to_string(), "second".to_string()),
];
let err = assemble_multipart_call(document, &fields).expect_err("duplicate model rejected");
assert!(matches!(err, CoreError::InvalidRequest(_)));
}
#[test]
fn multipart_rejects_duplicate_timeout_field() {
let document = json!({"type": "document_url", "document_url": "data:x"});
let fields = vec![
("model".to_string(), "m".to_string()),
("timeout".to_string(), "10".to_string()),
("timeout".to_string(), "20".to_string()),
];
let err = assemble_multipart_call(document, &fields).expect_err("duplicate timeout rejected");
assert!(matches!(err, CoreError::InvalidRequest(_)));
}
#[test]
fn upload_treats_binary_octet_stream_as_generic() {
let document = build_upload_document(
b"%PDF-1.7 minimal".to_vec(),
None,
Some("Binary/Octet-Stream"),
)
.expect("builds document");
assert!(document["document_url"]
.as_str()
.expect("data uri")
.starts_with("data:application/pdf;base64,"));
}