mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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>
This commit is contained in:
parent
f598cfe3a6
commit
47d8165a52
7 changed files with 737 additions and 432 deletions
|
|
@ -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)] = &[
|
||||
|
|
|
|||
|
|
@ -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)"
|
||||
);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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<Vec<u8>> {
|
|||
async fn parse_multipart(request: Request, state: &AppState) -> CoreResult<OcrCall> {
|
||||
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<u8>, Option<String>, Option<String>)> = None;
|
||||
let mut text_fields: Vec<(String, String)> = Vec::new();
|
||||
|
|
@ -82,22 +83,28 @@ async fn parse_multipart(request: Request, state: &AppState) -> CoreResult<OcrCa
|
|||
while let Some(field) = multipart
|
||||
.next_field()
|
||||
.await
|
||||
.map_err(|err| CoreError::InvalidRequest(format!("invalid multipart field: {err}")))?
|
||||
.map_err(|_| CoreError::InvalidRequest("could not read a multipart field".to_string()))?
|
||||
{
|
||||
let name = field.name().map(str::to_string);
|
||||
match name.as_deref() {
|
||||
Some("file") => {
|
||||
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<usize> = 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::<usize>().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<String>) {
|
||||
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<String>) {
|
||||
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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
493
litellm-rust/crates/ai-gateway/src/routes/ocr/tests.rs
Normal file
493
litellm-rust/crates/ai-gateway/src/routes/ocr/tests.rs
Normal file
|
|
@ -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<usize> = 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::<usize>().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<String>) {
|
||||
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<String>) {
|
||||
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");
|
||||
}
|
||||
|
|
@ -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<Duration>,
|
||||
}
|
||||
|
||||
fn positive_duration(seconds: Option<f64>) -> Option<Duration> {
|
||||
seconds
|
||||
.filter(|secs| secs.is_finite() && *secs > 0.0)
|
||||
.map(Duration::from_secs_f64)
|
||||
fn parse_timeout(seconds: Option<f64>) -> CoreResult<Option<Duration>> {
|
||||
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<Item = &'a str>) -> CoreResult<()> {
|
||||
let reserved = names.into_iter().find_map(|name| {
|
||||
OCR_RESERVED_PARAM_KEYS
|
||||
.iter()
|
||||
.copied()
|
||||
.find(|candidate| *candidate == name)
|
||||
});
|
||||
match reserved {
|
||||
None => Ok(()),
|
||||
Some(reserved) => Err(CoreError::InvalidRequest(format!(
|
||||
"the '{reserved}' parameter is not accepted on an OCR request; deployment \
|
||||
credentials, routing, and headers are server-controlled"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn parse_json_body(body: &[u8]) -> CoreResult<OcrCall> {
|
||||
|
|
@ -70,11 +93,12 @@ pub fn parse_json_body(body: &[u8]) -> CoreResult<OcrCall> {
|
|||
}
|
||||
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<OcrCall> {
|
||||
reject_reserved_params(text_fields.iter().map(|(name, _)| name.as_str()))?;
|
||||
reject_duplicate_unique_fields(text_fields)?;
|
||||
|
||||
let model = text_fields
|
||||
.iter()
|
||||
.find(|(name, _)| name == "model")
|
||||
|
|
@ -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::<f64>().map_err(|_| {
|
||||
CoreError::InvalidRequest(format!("invalid 'timeout' form field: {value:?}"))
|
||||
})?)),
|
||||
Some((_, raw)) => {
|
||||
let seconds = raw.parse::<f64>().map_err(|_| {
|
||||
CoreError::InvalidRequest(
|
||||
"'timeout' form field must be a number of seconds".to_string(),
|
||||
)
|
||||
})?;
|
||||
parse_timeout(Some(seconds))?
|
||||
}
|
||||
None => None,
|
||||
};
|
||||
|
||||
|
|
@ -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,"));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue