From 47d8165a522f0bb269b55568ab29fc4a710f29dd 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:28:06 +0000 Subject: [PATCH] feat(rust-gateway): harden OCR transport inputs and validation - reject host/config control params (api_key, api_base, custom_llm_provider, extra_headers, vertex_*) on both JSON and multipart OCR requests - reject duplicate file/model/timeout/document multipart fields - reject non-finite/zero/negative timeouts with a typed 400 instead of silently dropping them - treat generic upload MIME case-insensitively and include binary/octet-stream - tighten PDF magic-byte sniff to %PDF- - data-minimize multipart parse errors (no attacker-controlled field names/values) - extend the missing-master-key startup warning to cover OCR routes - move route tests into a dedicated tests module Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- .../crates/ai-gateway/src/constants.rs | 22 +- litellm-rust/crates/ai-gateway/src/main.rs | 2 +- .../crates/ai-gateway/src/routes/ocr/mod.rs | 425 +-------------- .../crates/ai-gateway/src/routes/ocr/tests.rs | 493 ++++++++++++++++++ .../ai-gateway/src/routes/ocr/transport.rs | 206 +++++++- litellm-rust/crates/core/src/constants.rs | 1 + litellm-rust/crates/core/src/ocr/mime.rs | 20 +- 7 files changed, 737 insertions(+), 432 deletions(-) create mode 100644 litellm-rust/crates/ai-gateway/src/routes/ocr/tests.rs diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 0abc7b04592..33a1381599a 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -6,8 +6,8 @@ //! read + fallback happens at the host/config layer. use litellm_core::constants::{ - MIME_APPLICATION_OCTET_STREAM, MIME_APPLICATION_PDF, MIME_IMAGE_BMP, MIME_IMAGE_GIF, - MIME_IMAGE_JPEG, MIME_IMAGE_PNG, MIME_IMAGE_TIFF, MIME_IMAGE_WEBP, + MIME_APPLICATION_OCTET_STREAM, MIME_APPLICATION_PDF, MIME_BINARY_OCTET_STREAM, MIME_IMAGE_BMP, + MIME_IMAGE_GIF, MIME_IMAGE_JPEG, MIME_IMAGE_PNG, MIME_IMAGE_TIFF, MIME_IMAGE_WEBP, }; /// Default LiteLLM control-plane base URL for request-log egress when @@ -35,6 +35,24 @@ pub(crate) const DEFAULT_PROVIDER: &str = "openai"; pub(crate) const DEFAULT_UPLOAD_MIME_TYPE: &str = MIME_APPLICATION_OCTET_STREAM; +pub(crate) const GENERIC_UPLOAD_MIME_TYPES: &[&str] = + &[MIME_APPLICATION_OCTET_STREAM, MIME_BINARY_OCTET_STREAM]; + +pub(crate) const OCR_RESERVED_PARAM_KEYS: &[&str] = &[ + "api_key", + "api_base", + "custom_llm_provider", + "extra_headers", + "vertex_credentials", + "vertex_ai_credentials", + "vertex_project", + "vertex_ai_project", + "vertex_location", + "vertex_ai_location", +]; + +pub(crate) const OCR_MULTIPART_UNIQUE_FIELDS: &[&str] = &["model", "timeout", "document"]; + pub(crate) const MAX_OCR_REQUEST_BYTES: usize = 100 * 1024 * 1024; pub(crate) const OCR_UPLOAD_MIME_BY_EXTENSION: &[(&str, &str)] = &[ diff --git a/litellm-rust/crates/ai-gateway/src/main.rs b/litellm-rust/crates/ai-gateway/src/main.rs index 3220c74ae4d..a7d775d2735 100644 --- a/litellm-rust/crates/ai-gateway/src/main.rs +++ b/litellm-rust/crates/ai-gateway/src/main.rs @@ -37,7 +37,7 @@ async fn main() { .map(Arc::from); if master_key.is_none() { eprintln!( - "warning: LITELLM_MASTER_KEY is not set; /v1/realtime will reject all requests (fail closed)" + "warning: LITELLM_MASTER_KEY is not set; /v1/realtime, /v1/ocr and /ocr will reject all requests (fail closed)" ); } 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 a3de414dd9b..6233819f677 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs @@ -1,6 +1,9 @@ mod service; mod transport; +#[cfg(test)] +mod tests; + use axum::body::{to_bytes, Body}; use axum::extract::{DefaultBodyLimit, FromRequest, Multipart, Request, State}; use axum::http::header::CONTENT_TYPE; @@ -72,9 +75,7 @@ async fn read_body(body: Body) -> CoreResult> { async fn parse_multipart(request: Request, state: &AppState) -> CoreResult { let mut multipart = Multipart::from_request(request, state) .await - .map_err(|err| { - CoreError::InvalidRequest(format!("could not read multipart form: {err}")) - })?; + .map_err(|_| CoreError::InvalidRequest("could not read the multipart form".to_string()))?; let mut file: Option<(Vec, Option, Option)> = None; let mut text_fields: Vec<(String, String)> = Vec::new(); @@ -82,22 +83,28 @@ async fn parse_multipart(request: Request, state: &AppState) -> CoreResult { + if file.is_some() { + return Err(CoreError::InvalidRequest( + "the 'file' field must appear at most once in a multipart OCR request" + .to_string(), + )); + } let filename = field.file_name().map(str::to_string); let content_type = field.content_type().map(str::to_string); - let bytes = field.bytes().await.map_err(|err| { - CoreError::InvalidRequest(format!("could not read uploaded file: {err}")) + let bytes = field.bytes().await.map_err(|_| { + CoreError::InvalidRequest("could not read the uploaded file".to_string()) })?; file = Some((bytes.to_vec(), filename, content_type)); } Some(name) => { let name = name.to_string(); - let text = field.text().await.map_err(|err| { - CoreError::InvalidRequest(format!("could not read form field '{name}': {err}")) + let text = field.text().await.map_err(|_| { + CoreError::InvalidRequest("could not read a multipart form field".to_string()) })?; text_fields.push((name, text)); } @@ -165,405 +172,3 @@ fn error_body(message: &str, error_type: &str) -> Value { } }) } - -#[cfg(test)] -mod tests { - use std::net::SocketAddr; - use std::sync::Arc; - - use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter}; - use serde_json::Value; - use tokio::io::{AsyncReadExt, AsyncWriteExt}; - use tokio::net::{TcpListener, TcpStream}; - - use crate::io::realtime_pool::RealtimePool; - 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; - } - } - String::from_utf8_lossy(&request).into_owned() - } - - async fn read_full_http_request(socket: &mut TcpStream) -> String { - let mut request = Vec::new(); - let mut buffer = [0_u8; 4096]; - let mut body_start: Option = None; - let mut content_length = 0_usize; - loop { - let n = socket.read(&mut buffer).await.expect("reads request"); - if n == 0 { - break; - } - request.extend_from_slice(&buffer[..n]); - if body_start.is_none() { - if let Some(pos) = request.windows(4).position(|window| window == b"\r\n\r\n") { - let start = pos + 4; - body_start = Some(start); - let headers = String::from_utf8_lossy(&request[..pos]).to_ascii_lowercase(); - content_length = headers - .lines() - .find_map(|line| line.strip_prefix("content-length:")) - .and_then(|value| value.trim().parse::().ok()) - .unwrap_or(0); - } - } - if let Some(start) = body_start { - if request.len() >= start + content_length { - break; - } - } - } - String::from_utf8_lossy(&request).into_owned() - } - - async fn spawn_mock_upstream_capture() -> (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"); - let handle = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts one request"); - 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{}", - body.len(), - body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - request - }); - (format!("http://{addr}"), handle) - } - - async fn spawn_mock_upstream() -> (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"); - 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 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{}", - body.len(), - body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - request - }); - (format!("http://{addr}"), handle) - } - - async fn spawn_mock_upstream_error() -> String { - 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 (mut socket, _) = listener.accept().await.expect("accepts one request"); - let _ = read_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{}", - body.len(), - body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - }); - format!("http://{addr}") - } - - fn app_with_deployment(api_base: &str) -> axum::Router { - app_with_params("mistral/mistral-ocr-latest", None, api_base) - } - - fn app_with_params( - model: &str, - custom_llm_provider: Option<&str>, - api_base: &str, - ) -> axum::Router { - let router = ModelRouter::new(vec![Deployment { - model_name: "rust-ocr-mistral".to_string(), - litellm_params: LiteLLMParams { - model: model.to_string(), - api_key: Some("sk-upstream".to_string()), - api_base: Some(api_base.to_string()), - custom_llm_provider: custom_llm_provider.map(str::to_string), - }, - }]); - let state = AppState { - router: Arc::new(router), - master_key: Some(Arc::from(MASTER_KEY)), - loggers: Arc::new(Vec::new()), - realtime_pool: RealtimePool::disabled(), - }; - crate::routes::app(state) - } - - async fn serve(app: axum::Router) -> SocketAddr { - 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"); - }); - addr - } - - #[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 response = reqwest::Client::new() - .post(format!("http://{addr}{path}")) - .bearer_auth(MASTER_KEY) - .json(&serde_json::json!({ - "model": "rust-ocr-mistral", - "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"} - })) - .send() - .await - .expect("request sent"); - - assert_eq!(response.status(), reqwest::StatusCode::OK, "path {path}"); - let body: Value = response.json().await.expect("json body"); - assert_eq!(body["object"], "ocr", "path {path}"); - assert_eq!(body["model"], "rust-ocr-mistral", "path {path}"); - assert_eq!(body["pages"][0]["markdown"], "hello ocr", "path {path}"); - - let upstream_request = upstream_handle.await.expect("upstream served"); - assert!( - upstream_request.contains("authorization: Bearer sk-upstream") - || upstream_request.contains("Authorization: Bearer sk-upstream"), - "upstream must receive the deployment credential: {upstream_request}" - ); - } - } - - #[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( - "mistral-ocr-latest", - Some("mistral"), - &upstream, - )) - .await; - - let response = reqwest::Client::new() - .post(format!("http://{addr}/v1/ocr")) - .bearer_auth(MASTER_KEY) - .json(&serde_json::json!({ - "model": "rust-ocr-mistral", - "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"} - })) - .send() - .await - .expect("request sent"); - - assert_eq!(response.status(), reqwest::StatusCode::OK); - let body: Value = response.json().await.expect("json body"); - assert_eq!(body["object"], "ocr"); - assert_eq!(body["model"], "rust-ocr-mistral"); - - let upstream_request = upstream_handle.await.expect("upstream served"); - assert!(upstream_request.starts_with("POST"), "{upstream_request}"); - } - - #[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 form = reqwest::multipart::Form::new() - .text("model", "rust-ocr-mistral") - .part( - "file", - reqwest::multipart::Part::bytes(b"%PDF-1.4 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::OK); - let body: Value = response.json().await.expect("json body"); - assert_eq!(body["object"], "ocr"); - assert_eq!(body["pages"][0]["markdown"], "hello ocr"); - - let upstream_request = upstream_handle.await.expect("upstream served"); - assert!(upstream_request.starts_with("POST"), "{upstream_request}"); - } - - #[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 form = reqwest::multipart::Form::new() - .text("model", "rust-ocr-mistral") - .part( - "file", - reqwest::multipart::Part::bytes(b"%PDF-1.7 minimal pdf bytes".to_vec()) - .mime_str("application/octet-stream") - .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::OK); - let upstream_request = upstream_handle.await.expect("upstream served"); - assert!( - upstream_request.contains("data:application/pdf;base64,"), - "unnamed octet-stream PDF must be sniffed to application/pdf: {upstream_request}" - ); - } - - #[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 png = vec![0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x01]; - let form = reqwest::multipart::Form::new() - .text("model", "rust-ocr-mistral") - .part( - "file", - reqwest::multipart::Part::bytes(png) - .mime_str("application/octet-stream") - .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::OK); - let upstream_request = upstream_handle.await.expect("upstream served"); - assert!( - upstream_request.contains("data:image/png;base64,"), - "unnamed octet-stream PNG must be sniffed to image/png: {upstream_request}" - ); - } - - #[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 response = reqwest::Client::new() - .post(format!("http://{addr}/v1/ocr")) - .bearer_auth(MASTER_KEY) - .json(&serde_json::json!({ - "model": "rust-ocr-mistral", - "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"} - })) - .send() - .await - .expect("request sent"); - - assert_eq!(response.status(), reqwest::StatusCode::FORBIDDEN); - let body: Value = response.json().await.expect("json body"); - assert_eq!(body["error"]["type"], "upstream_error"); - let message = body["error"]["message"].as_str().expect("message string"); - assert!( - !message.contains("sk-leaked-secret-value") && !message.contains("invalid_api_key"), - "provider body must not leak into the public error: {message}" - ); - } - - #[tokio::test] - async fn missing_master_key_is_unauthorized() { - let addr = 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!({ - "model": "rust-ocr-mistral", - "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"} - })) - .send() - .await - .expect("request sent"); - assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED); - } - - #[tokio::test] - async fn unknown_model_is_not_found() { - let addr = 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": "does-not-exist", - "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"} - })) - .send() - .await - .expect("request sent"); - assert_eq!(response.status(), reqwest::StatusCode::NOT_FOUND); - let body: Value = response.json().await.expect("json body"); - assert_eq!(body["error"]["type"], "not_found_error"); - } - - #[tokio::test] - async fn file_document_over_json_is_rejected() { - let addr = 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": "rust-ocr-mistral", - "document": {"type": "file", "file": "/etc/passwd"} - })) - .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/tests.rs b/litellm-rust/crates/ai-gateway/src/routes/ocr/tests.rs new file mode 100644 index 00000000000..5a366672ea6 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/ocr/tests.rs @@ -0,0 +1,493 @@ +use std::net::SocketAddr; +use std::sync::Arc; + +use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter}; +use serde_json::Value; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::{TcpListener, TcpStream}; + +use crate::io::realtime_pool::RealtimePool; +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; + } + } + String::from_utf8_lossy(&request).into_owned() +} + +async fn read_full_http_request(socket: &mut TcpStream) -> String { + let mut request = Vec::new(); + let mut buffer = [0_u8; 4096]; + let mut body_start: Option = None; + let mut content_length = 0_usize; + loop { + let n = socket.read(&mut buffer).await.expect("reads request"); + if n == 0 { + break; + } + request.extend_from_slice(&buffer[..n]); + if body_start.is_none() { + if let Some(pos) = request.windows(4).position(|window| window == b"\r\n\r\n") { + let start = pos + 4; + body_start = Some(start); + let headers = String::from_utf8_lossy(&request[..pos]).to_ascii_lowercase(); + content_length = headers + .lines() + .find_map(|line| line.strip_prefix("content-length:")) + .and_then(|value| value.trim().parse::().ok()) + .unwrap_or(0); + } + } + if let Some(start) = body_start { + if request.len() >= start + content_length { + break; + } + } + } + String::from_utf8_lossy(&request).into_owned() +} + +async fn spawn_mock_upstream_capture() -> (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"); + let handle = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts one request"); + 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{}", + body.len(), + body + ); + socket + .write_all(response.as_bytes()) + .await + .expect("writes response"); + request + }); + (format!("http://{addr}"), handle) +} + +async fn spawn_mock_upstream() -> (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"); + 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 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{}", + body.len(), + body + ); + socket + .write_all(response.as_bytes()) + .await + .expect("writes response"); + request + }); + (format!("http://{addr}"), handle) +} + +async fn spawn_mock_upstream_error() -> String { + 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 (mut socket, _) = listener.accept().await.expect("accepts one request"); + let _ = read_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{}", + body.len(), + body + ); + socket + .write_all(response.as_bytes()) + .await + .expect("writes response"); + }); + format!("http://{addr}") +} + +fn app_with_deployment(api_base: &str) -> axum::Router { + app_with_params("mistral/mistral-ocr-latest", None, api_base) +} + +fn app_with_params(model: &str, custom_llm_provider: Option<&str>, api_base: &str) -> axum::Router { + let router = ModelRouter::new(vec![Deployment { + model_name: "rust-ocr-mistral".to_string(), + litellm_params: LiteLLMParams { + model: model.to_string(), + api_key: Some("sk-upstream".to_string()), + api_base: Some(api_base.to_string()), + custom_llm_provider: custom_llm_provider.map(str::to_string), + }, + }]); + let state = AppState { + router: Arc::new(router), + master_key: Some(Arc::from(MASTER_KEY)), + loggers: Arc::new(Vec::new()), + realtime_pool: RealtimePool::disabled(), + }; + crate::routes::app(state) +} + +async fn serve(app: axum::Router) -> SocketAddr { + 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"); + }); + addr +} + +#[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 response = reqwest::Client::new() + .post(format!("http://{addr}{path}")) + .bearer_auth(MASTER_KEY) + .json(&serde_json::json!({ + "model": "rust-ocr-mistral", + "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"} + })) + .send() + .await + .expect("request sent"); + + assert_eq!(response.status(), reqwest::StatusCode::OK, "path {path}"); + let body: Value = response.json().await.expect("json body"); + assert_eq!(body["object"], "ocr", "path {path}"); + assert_eq!(body["model"], "rust-ocr-mistral", "path {path}"); + assert_eq!(body["pages"][0]["markdown"], "hello ocr", "path {path}"); + + let upstream_request = upstream_handle.await.expect("upstream served"); + assert!( + upstream_request.contains("authorization: Bearer sk-upstream") + || upstream_request.contains("Authorization: Bearer sk-upstream"), + "upstream must receive the deployment credential: {upstream_request}" + ); + } +} + +#[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( + "mistral-ocr-latest", + Some("mistral"), + &upstream, + )) + .await; + + let response = reqwest::Client::new() + .post(format!("http://{addr}/v1/ocr")) + .bearer_auth(MASTER_KEY) + .json(&serde_json::json!({ + "model": "rust-ocr-mistral", + "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"} + })) + .send() + .await + .expect("request sent"); + + assert_eq!(response.status(), reqwest::StatusCode::OK); + let body: Value = response.json().await.expect("json body"); + assert_eq!(body["object"], "ocr"); + assert_eq!(body["model"], "rust-ocr-mistral"); + + let upstream_request = upstream_handle.await.expect("upstream served"); + assert!(upstream_request.starts_with("POST"), "{upstream_request}"); +} + +#[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 form = reqwest::multipart::Form::new() + .text("model", "rust-ocr-mistral") + .part( + "file", + reqwest::multipart::Part::bytes(b"%PDF-1.4 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::OK); + let body: Value = response.json().await.expect("json body"); + assert_eq!(body["object"], "ocr"); + assert_eq!(body["pages"][0]["markdown"], "hello ocr"); + + let upstream_request = upstream_handle.await.expect("upstream served"); + assert!(upstream_request.starts_with("POST"), "{upstream_request}"); +} + +#[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 form = reqwest::multipart::Form::new() + .text("model", "rust-ocr-mistral") + .part( + "file", + reqwest::multipart::Part::bytes(b"%PDF-1.7 minimal pdf bytes".to_vec()) + .mime_str("application/octet-stream") + .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::OK); + let upstream_request = upstream_handle.await.expect("upstream served"); + assert!( + upstream_request.contains("data:application/pdf;base64,"), + "unnamed octet-stream PDF must be sniffed to application/pdf: {upstream_request}" + ); +} + +#[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 png = vec![0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x01]; + let form = reqwest::multipart::Form::new() + .text("model", "rust-ocr-mistral") + .part( + "file", + reqwest::multipart::Part::bytes(png) + .mime_str("application/octet-stream") + .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::OK); + let upstream_request = upstream_handle.await.expect("upstream served"); + assert!( + upstream_request.contains("data:image/png;base64,"), + "unnamed octet-stream PNG must be sniffed to image/png: {upstream_request}" + ); +} + +#[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 response = reqwest::Client::new() + .post(format!("http://{addr}/v1/ocr")) + .bearer_auth(MASTER_KEY) + .json(&serde_json::json!({ + "model": "rust-ocr-mistral", + "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"} + })) + .send() + .await + .expect("request sent"); + + assert_eq!(response.status(), reqwest::StatusCode::FORBIDDEN); + let body: Value = response.json().await.expect("json body"); + assert_eq!(body["error"]["type"], "upstream_error"); + let message = body["error"]["message"].as_str().expect("message string"); + assert!( + !message.contains("sk-leaked-secret-value") && !message.contains("invalid_api_key"), + "provider body must not leak into the public error: {message}" + ); +} + +#[tokio::test] +async fn missing_master_key_is_unauthorized() { + let addr = 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!({ + "model": "rust-ocr-mistral", + "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"} + })) + .send() + .await + .expect("request sent"); + assert_eq!(response.status(), reqwest::StatusCode::UNAUTHORIZED); +} + +#[tokio::test] +async fn unknown_model_is_not_found() { + let addr = 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": "does-not-exist", + "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"} + })) + .send() + .await + .expect("request sent"); + assert_eq!(response.status(), reqwest::StatusCode::NOT_FOUND); + let body: Value = response.json().await.expect("json body"); + assert_eq!(body["error"]["type"], "not_found_error"); +} + +#[tokio::test] +async fn file_document_over_json_is_rejected() { + let addr = 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": "rust-ocr-mistral", + "document": {"type": "file", "file": "/etc/passwd"} + })) + .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 json_reserved_control_param_is_rejected() { + let addr = 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": "rust-ocr-mistral", + "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"}, + "api_base": "http://attacker.example" + })) + .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"); + assert!( + !message.contains("attacker.example"), + "must not echo the attacker value: {message}" + ); +} + +#[tokio::test] +async fn multipart_reserved_control_param_is_rejected() { + let addr = 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") + .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"); +} + +#[tokio::test] +async fn duplicate_file_multipart_is_rejected() { + let addr = serve(app_with_deployment("http://127.0.0.1:1")).await; + let form = reqwest::multipart::Form::new() + .text("model", "rust-ocr-mistral") + .part( + "file", + reqwest::multipart::Part::bytes(b"%PDF-1.7 first".to_vec()) + .file_name("a.pdf") + .mime_str("application/pdf") + .expect("mime"), + ) + .part( + "file", + reqwest::multipart::Part::bytes(b"%PDF-1.7 second".to_vec()) + .file_name("b.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"); +} + +#[tokio::test] +async fn non_positive_timeout_is_rejected() { + let addr = 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": "rust-ocr-mistral", + "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"}, + "timeout": -1 + })) + .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 index c106c0cf5b6..66a9e7b5ef1 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs @@ -8,7 +8,10 @@ use litellm_core::CoreResult; use serde::Deserialize; use serde_json::{json, Map, Value}; -use crate::constants::{DEFAULT_UPLOAD_MIME_TYPE, OCR_UPLOAD_MIME_BY_EXTENSION}; +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")] @@ -54,10 +57,30 @@ pub struct OcrCall { pub timeout: Option, } -fn positive_duration(seconds: Option) -> Option { - seconds - .filter(|secs| secs.is_finite() && *secs > 0.0) - .map(Duration::from_secs_f64) +fn parse_timeout(seconds: Option) -> CoreResult> { + match seconds { + None => Ok(None), + Some(value) if value.is_finite() && value > 0.0 => Ok(Some(Duration::from_secs_f64(value))), + Some(_) => Err(CoreError::InvalidRequest( + "'timeout' must be a positive, finite number of seconds".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 { @@ -70,11 +93,12 @@ pub fn parse_json_body(body: &[u8]) -> CoreResult { } 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: positive_duration(request.timeout), + timeout: parse_timeout(request.timeout)?, }) } @@ -88,11 +112,17 @@ fn mime_from_filename(filename: &str) -> Option<&'static str> { .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() && *value != DEFAULT_UPLOAD_MIME_TYPE); + .filter(|value| !value.is_empty() && !is_generic_upload_mime(value)); if let Some(declared) = declared { return declared.to_string(); } @@ -129,10 +159,26 @@ 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") @@ -144,9 +190,14 @@ pub fn assemble_multipart_call( })?; let timeout = match text_fields.iter().find(|(name, _)| name == "timeout") { - Some((_, value)) => positive_duration(Some(value.parse::().map_err(|_| { - CoreError::InvalidRequest(format!("invalid 'timeout' form field: {value:?}")) - })?)), + 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, }; @@ -339,4 +390,139 @@ mod tests { 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"] { + 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"] { + 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/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index 4aff46a7a7e..7275ebd6878 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -1,5 +1,6 @@ pub const MIME_APPLICATION_PDF: &str = "application/pdf"; pub const MIME_APPLICATION_OCTET_STREAM: &str = "application/octet-stream"; +pub const MIME_BINARY_OCTET_STREAM: &str = "binary/octet-stream"; pub const MIME_IMAGE_PNG: &str = "image/png"; pub const MIME_IMAGE_JPEG: &str = "image/jpeg"; pub const MIME_IMAGE_GIF: &str = "image/gif"; diff --git a/litellm-rust/crates/core/src/ocr/mime.rs b/litellm-rust/crates/core/src/ocr/mime.rs index 79e6deb95bf..1c2d8981def 100644 --- a/litellm-rust/crates/core/src/ocr/mime.rs +++ b/litellm-rust/crates/core/src/ocr/mime.rs @@ -4,7 +4,7 @@ use crate::constants::{ }; pub fn sniff_mime(bytes: &[u8]) -> Option<&'static str> { - if bytes.starts_with(b"%PDF") { + if bytes.starts_with(b"%PDF-") { return Some(MIME_APPLICATION_PDF); } if bytes.starts_with(&[0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A]) { @@ -38,6 +38,11 @@ mod tests { assert_eq!(sniff_mime(b"%PDF-1.7\n..."), Some(MIME_APPLICATION_PDF)); } + #[test] + fn does_not_sniff_pdf_without_version_marker() { + assert_eq!(sniff_mime(b"%PDFxx"), None); + } + #[test] fn sniffs_png() { assert_eq!( @@ -59,18 +64,15 @@ mod tests { #[test] fn sniffs_webp() { - let mut bytes = b"RIFF".to_vec(); - bytes.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); - bytes.extend_from_slice(b"WEBP"); - assert_eq!(sniff_mime(&bytes), Some(MIME_IMAGE_WEBP)); + assert_eq!( + sniff_mime(b"RIFF\x00\x00\x00\x00WEBP"), + Some(MIME_IMAGE_WEBP) + ); } #[test] fn does_not_sniff_riff_without_webp() { - let mut bytes = b"RIFF".to_vec(); - bytes.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); - bytes.extend_from_slice(b"WAVE"); - assert_eq!(sniff_mime(&bytes), None); + assert_eq!(sniff_mime(b"RIFF\x00\x00\x00\x00WAVE"), None); } #[test]