From 7f526e0b08475e966a51b8d70e3a5432c6b4816e Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 17 Jul 2026 00:53:18 +0000 Subject: [PATCH] refactor(rust-gateway): harden OCR transport error contract and validation 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> --- .../crates/ai-gateway/src/constants.rs | 4 +- .../crates/ai-gateway/src/routes/ocr/mod.rs | 63 ++- .../crates/ai-gateway/src/routes/ocr/tests.rs | 134 +++-- .../ai-gateway/src/routes/ocr/transport.rs | 532 ------------------ .../src/routes/ocr/transport/mod.rs | 255 +++++++++ .../src/routes/ocr/transport/tests.rs | 441 +++++++++++++++ 6 files changed, 841 insertions(+), 588 deletions(-) delete mode 100644 litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs create mode 100644 litellm-rust/crates/ai-gateway/src/routes/ocr/transport/mod.rs create mode 100644 litellm-rust/crates/ai-gateway/src/routes/ocr/transport/tests.rs diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 33a1381599a..c8b50125bb0 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -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; diff --git a/litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs index 6233819f677..84bd06553ec 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs @@ -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> { 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 { @@ -123,22 +123,50 @@ async fn parse_multipart(request: Request, state: &AppState) -> CoreResult 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() } diff --git a/litellm-rust/crates/ai-gateway/src/routes/ocr/tests.rs b/litellm-rust/crates/ai-gateway/src/routes/ocr/tests.rs index 5a366672ea6..14deb8fda7e 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/ocr/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/ocr/tests.rs @@ -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) { 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) { (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"); +} diff --git a/litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs b/litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs deleted file mode 100644 index 5c3132cd492..00000000000 --- a/litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs +++ /dev/null @@ -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 { - 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, - #[serde(flatten)] - pub optional_params: Map, -} - -#[derive(Debug)] -pub struct OcrCall { - pub model: String, - pub document: Value, - pub optional_params: Map, - pub timeout: Option, -} - -fn parse_timeout(seconds: Option) -> CoreResult> { - 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) -> 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 { - 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, - filename: Option<&str>, - content_type: Option<&str>, -) -> CoreResult { - 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 { - 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::().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 = 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,")); - } -} diff --git a/litellm-rust/crates/ai-gateway/src/routes/ocr/transport/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/ocr/transport/mod.rs new file mode 100644 index 00000000000..6df6f676c28 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/ocr/transport/mod.rs @@ -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 { + 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, + #[serde(flatten)] + pub optional_params: Map, +} + +#[derive(Debug)] +pub struct OcrCall { + pub model: String, + pub document: Value, + pub optional_params: Map, + pub timeout: Option, +} + +fn parse_timeout(seconds: Option) -> CoreResult> { + 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) -> 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 { + 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, + filename: Option<&str>, + content_type: Option<&str>, +) -> CoreResult { + 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 { + 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::().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 = 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, + }) +} diff --git a/litellm-rust/crates/ai-gateway/src/routes/ocr/transport/tests.rs b/litellm-rust/crates/ai-gateway/src/routes/ocr/transport/tests.rs new file mode 100644 index 00000000000..dca33d2a951 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/ocr/transport/tests.rs @@ -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,")); +}