diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 422dfb20065..3de0376aebc 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -45,6 +45,7 @@ dependencies = [ "matchit", "memchr", "mime", + "multer", "percent-encoding", "pin-project-lite", "rustversion", @@ -206,6 +207,15 @@ dependencies = [ "syn", ] +[[package]] +name = "encoding_rs" +version = "0.8.35" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3" +dependencies = [ + "cfg-if", +] + [[package]] name = "equivalent" version = "1.0.2" @@ -753,6 +763,23 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "multer" +version = "3.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "83e87776546dc87511aa5ee218730c92b666d7264ab6ed41f9d215af9cd5224b" +dependencies = [ + "bytes", + "encoding_rs", + "futures-util", + "http", + "httparse", + "memchr", + "mime", + "spin", + "version_check", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -1286,6 +1313,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "spin" +version = "0.9.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e" + [[package]] name = "stable_deref_trait" version = "1.2.1" diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index 2f414159158..ac4b4340999 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -24,7 +24,7 @@ tokio-tungstenite.workspace = true futures-util.workspace = true serde_json.workspace = true base64.workspace = true -axum = { workspace = true, features = ["ws"], optional = true } +axum = { workspace = true, features = ["ws", "multipart"], optional = true } serde = { workspace = true, optional = true } subtle = { workspace = true, optional = true } # sha2 hashes the master key into user_api_key_hash (matches the proxy's diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 3116a4c9932..1ef19ae1114 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -27,3 +27,26 @@ pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500; /// Provider attributed to realtime sessions in the logging payload. pub(crate) const DEFAULT_PROVIDER: &str = "openai"; + +/// MIME type used when an OCR upload declares none and its filename extension is +/// unrecognized. Mirrors the Python proxy's `_build_document_from_upload`. +pub(crate) const DEFAULT_UPLOAD_MIME_TYPE: &str = "application/octet-stream"; + +/// Upper bound on an OCR request body (JSON or multipart). Base64-encoded PDFs +/// inflate roughly 1.33x over the raw bytes, so this stays well above the +/// default 50MB document download cap while still rejecting absurd payloads. +pub(crate) const MAX_OCR_REQUEST_BYTES: usize = 100 * 1024 * 1024; + +/// Filename-extension to MIME map for OCR uploads whose transport declares no +/// usable content type. Mirrors `litellm.ocr.main.get_mime_type`'s explicit map. +pub(crate) const OCR_UPLOAD_MIME_BY_EXTENSION: &[(&str, &str)] = &[ + ("pdf", "application/pdf"), + ("png", "image/png"), + ("jpg", "image/jpeg"), + ("jpeg", "image/jpeg"), + ("gif", "image/gif"), + ("webp", "image/webp"), + ("tiff", "image/tiff"), + ("tif", "image/tiff"), + ("bmp", "image/bmp"), +]; diff --git a/litellm-rust/crates/ai-gateway/src/routes/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/mod.rs index c6b9573781a..af40e7db189 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/mod.rs @@ -7,6 +7,7 @@ pub mod gil; pub mod health; +pub mod ocr; pub mod realtime; use axum::Router; @@ -18,6 +19,7 @@ pub fn app(state: AppState) -> Router { Router::new() .merge(health::router()) .merge(gil::router()) + .merge(ocr::router()) .merge(realtime::router()) .with_state(state) } diff --git a/litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs new file mode 100644 index 00000000000..2d35dbd3ed3 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs @@ -0,0 +1,414 @@ +mod service; +mod transport; + +use axum::body::{to_bytes, Body}; +use axum::extract::{DefaultBodyLimit, FromRequest, Multipart, Request, State}; +use axum::http::header::CONTENT_TYPE; +use axum::http::StatusCode; +use axum::response::{IntoResponse, Response}; +use axum::routing::post; +use axum::{Json, Router}; +use litellm_core::error::CoreError; +use litellm_core::CoreResult; +use serde_json::{json, Value}; + +use crate::auth::RequireMasterKey; +use crate::constants::MAX_OCR_REQUEST_BYTES; +use crate::state::AppState; + +use transport::OcrCall; + +pub fn router() -> Router { + Router::new() + .route("/v1/ocr", post(handle)) + .route("/ocr", post(handle)) + .layer(DefaultBodyLimit::max(MAX_OCR_REQUEST_BYTES)) +} + +async fn handle( + _auth: RequireMasterKey, + State(state): State, + request: Request, +) -> Response { + let call = match parse_request(request, &state).await { + Ok(call) => call, + Err(err) => return error_response(&err), + }; + match service::run_ocr(&state.router, call).await { + Ok(value) => (StatusCode::OK, Json(value)).into_response(), + Err(err) => error_response(&err), + } +} + +fn is_multipart(request: &Request) -> bool { + request + .headers() + .get(CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .map(|value| value.to_ascii_lowercase().contains("multipart/form-data")) + .unwrap_or(false) +} + +async fn parse_request(request: Request, state: &AppState) -> CoreResult { + if is_multipart(&request) { + parse_multipart(request, state).await + } else { + parse_json(request).await + } +} + +async fn parse_json(request: Request) -> CoreResult { + let bytes = read_body(request.into_body()).await?; + transport::parse_json_body(&bytes) +} + +async fn read_body(body: Body) -> CoreResult> { + to_bytes(body, MAX_OCR_REQUEST_BYTES) + .await + .map(|bytes| bytes.to_vec()) + .map_err(|err| CoreError::InvalidRequest(format!("could not read request body: {err}"))) +} + +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}")) + })?; + + let mut file: Option<(Vec, Option, Option)> = None; + let mut text_fields: Vec<(String, String)> = Vec::new(); + + while let Some(field) = multipart + .next_field() + .await + .map_err(|err| CoreError::InvalidRequest(format!("invalid multipart field: {err}")))? + { + let name = field.name().map(str::to_string); + match name.as_deref() { + Some("file") => { + 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}")) + })?; + 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}")) + })?; + text_fields.push((name, text)); + } + None => {} + } + } + + let (bytes, filename, content_type) = file.ok_or_else(|| { + CoreError::InvalidRequest( + "multipart OCR request must include a 'file' field with the document to process" + .to_string(), + ) + })?; + let document = + transport::build_upload_document(bytes, filename.as_deref(), content_type.as_deref())?; + transport::assemble_multipart_call(document, &text_fields) +} + +fn error_response(error: &CoreError) -> Response { + let (status, error_type, message) = match error { + CoreError::InvalidRequest(_) + | CoreError::InvalidType { .. } + | CoreError::MissingField(_) + | CoreError::InvalidProvider(_) => ( + StatusCode::BAD_REQUEST, + "invalid_request_error", + error.to_string(), + ), + CoreError::Auth(_) => ( + StatusCode::UNAUTHORIZED, + "authentication_error", + error.to_string(), + ), + CoreError::Routing(_) => (StatusCode::NOT_FOUND, "not_found_error", error.to_string()), + CoreError::Http { status, .. } => { + let status = StatusCode::from_u16(*status).unwrap_or(StatusCode::BAD_GATEWAY); + ( + status, + "upstream_error", + format!( + "the OCR provider returned an error (status {})", + status.as_u16() + ), + ) + } + CoreError::Network(_) => ( + StatusCode::BAD_GATEWAY, + "upstream_error", + "the OCR provider could not be reached".to_string(), + ), + CoreError::InvalidResponse(_) => ( + StatusCode::BAD_GATEWAY, + "upstream_error", + "the OCR provider returned an unexpected response".to_string(), + ), + }; + (status, Json(error_body(&message, error_type))).into_response() +} + +fn error_body(message: &str, error_type: &str) -> Value { + json!({ + "error": { + "message": message, + "type": error_type, + } + }) +} + +#[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 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 { + let router = ModelRouter::new(vec![Deployment { + model_name: "rust-ocr-mistral".to_string(), + litellm_params: LiteLLMParams { + model: "mistral/mistral-ocr-latest".to_string(), + api_key: Some("sk-upstream".to_string()), + api_base: Some(api_base.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"], "mistral-ocr-latest", "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 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 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/service.rs b/litellm-rust/crates/ai-gateway/src/routes/ocr/service.rs new file mode 100644 index 00000000000..2b011ac9d0f --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/ocr/service.rs @@ -0,0 +1,87 @@ +use litellm_core::error::CoreError; +use litellm_core::router::Router; +use litellm_core::CoreResult; +use serde_json::Value; + +use crate::io::ocr::{ocr, OcrRequest}; + +use super::transport::OcrCall; + +fn present(value: Option<&str>) -> Option<&str> { + value.map(str::trim).filter(|value| !value.is_empty()) +} + +fn split_provider(model: &str) -> CoreResult<(&str, &str)> { + model + .split_once('/') + .filter(|(provider, rest)| !provider.is_empty() && !rest.is_empty()) + .ok_or_else(|| { + CoreError::InvalidProvider(format!( + "deployment model '{model}' must be prefixed with an OCR provider, e.g. \ + 'mistral/mistral-ocr-latest'" + )) + }) +} + +pub async fn run_ocr(router: &Router, call: OcrCall) -> CoreResult { + let deployment = router + .get_available_deployment(&call.model) + .ok_or_else(|| { + CoreError::Routing(format!( + "no deployment available for model '{}'", + call.model + )) + })?; + let params = &deployment.litellm_params; + let (provider, provider_model) = split_provider(¶ms.model)?; + + ocr(OcrRequest { + model: provider_model, + document: call.document, + api_key: present(params.api_key.as_deref()), + api_base: present(params.api_base.as_deref()), + custom_llm_provider: provider, + extra_headers: None, + optional_params: call.optional_params, + timeout: call.timeout, + }) + .await +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn splits_provider_prefix() { + assert_eq!( + split_provider("mistral/mistral-ocr-latest").expect("splits"), + ("mistral", "mistral-ocr-latest") + ); + assert_eq!( + split_provider("azure_ai/doc-intelligence/prebuilt-layout").expect("splits"), + ("azure_ai", "doc-intelligence/prebuilt-layout") + ); + } + + #[test] + fn present_treats_empty_and_whitespace_as_absent() { + assert_eq!(present(Some("sk-key")), Some("sk-key")); + assert_eq!(present(Some(" sk-key ")), Some("sk-key")); + assert_eq!(present(Some("")), None); + assert_eq!(present(Some(" ")), None); + assert_eq!(present(None), None); + } + + #[test] + fn rejects_model_without_provider_prefix() { + assert!(matches!( + split_provider("mistral-ocr-latest"), + Err(CoreError::InvalidProvider(_)) + )); + assert!(matches!( + split_provider("mistral/"), + Err(CoreError::InvalidProvider(_)) + )); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs b/litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs new file mode 100644 index 00000000000..dcdf73a5a69 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs @@ -0,0 +1,318 @@ +use std::time::Duration; + +use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; +use base64::Engine; +use litellm_core::error::CoreError; +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}; + +#[derive(Debug, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum OcrDocument { + DocumentUrl { document_url: String }, + ImageUrl { image_url: String }, + File {}, +} + +impl OcrDocument { + fn into_value(self) -> CoreResult { + let (field, url) = match self { + OcrDocument::DocumentUrl { document_url } => ("document_url", document_url), + OcrDocument::ImageUrl { image_url } => ("image_url", image_url), + OcrDocument::File {} => { + return Err(CoreError::InvalidRequest( + "document type 'file' is not supported through the JSON API; upload the \ + file via multipart/form-data with a 'file' field, or use a 'document_url' \ + or 'image_url' document type" + .to_string(), + )) + } + }; + if url.starts_with("reducto://") { + return Err(CoreError::InvalidRequest( + "reducto:// file IDs are not accepted through the OCR API; upload the file in \ + the same request via multipart/form-data with a 'file' field, or pass an \ + inline base64 data URI as the document URL" + .to_string(), + )); + } + Ok(json!({ "type": field, field: url })) + } +} + +#[derive(Debug, Deserialize)] +pub struct OcrJsonRequest { + pub model: String, + pub document: OcrDocument, + #[serde(default)] + pub timeout: Option, + #[serde(flatten)] + pub optional_params: Map, +} + +#[derive(Debug)] +pub struct OcrCall { + pub model: String, + pub document: Value, + pub optional_params: Map, + pub timeout: Option, +} + +fn positive_duration(seconds: Option) -> Option { + seconds + .filter(|secs| secs.is_finite() && *secs > 0.0) + .map(Duration::from_secs_f64) +} + +pub fn parse_json_body(body: &[u8]) -> CoreResult { + if body.is_empty() { + return Err(CoreError::InvalidRequest( + "empty request body; send a JSON body with 'model' and 'document', or use \ + multipart/form-data for file uploads" + .to_string(), + )); + } + let request: OcrJsonRequest = serde_json::from_slice(body) + .map_err(|err| CoreError::InvalidRequest(format!("invalid OCR request body: {err}")))?; + Ok(OcrCall { + model: request.model, + document: request.document.into_value()?, + optional_params: request.optional_params, + timeout: positive_duration(request.timeout), + }) +} + +fn mime_from_filename(filename: &str) -> Option<&'static str> { + let extension = filename + .rsplit_once('.') + .map(|(_, ext)| ext.to_ascii_lowercase())?; + OCR_UPLOAD_MIME_BY_EXTENSION + .iter() + .find(|(candidate, _)| *candidate == extension) + .map(|(_, mime)| *mime) +} + +fn resolve_upload_mime(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); + if let Some(declared) = declared { + return declared.to_string(); + } + filename + .and_then(mime_from_filename) + .unwrap_or(DEFAULT_UPLOAD_MIME_TYPE) + .to_string() +} + +pub fn build_upload_document( + bytes: Vec, + filename: Option<&str>, + content_type: Option<&str>, +) -> CoreResult { + if bytes.is_empty() { + return Err(CoreError::InvalidRequest( + "uploaded file is empty".to_string(), + )); + } + let mime = resolve_upload_mime(content_type, filename); + let data_uri = format!("data:{mime};base64,{}", BASE64_STANDARD.encode(&bytes)); + let field = if mime.starts_with("image/") { + "image_url" + } else { + "document_url" + }; + Ok(json!({ "type": field, field: data_uri })) +} + +fn coerce_form_field(value: &str) -> Value { + serde_json::from_str(value).unwrap_or_else(|_| Value::String(value.to_string())) +} + +pub fn assemble_multipart_call( + document: Value, + text_fields: &[(String, String)], +) -> CoreResult { + let model = text_fields + .iter() + .find(|(name, _)| name == "model") + .map(|(_, value)| value.clone()) + .ok_or_else(|| { + CoreError::InvalidRequest( + "multipart OCR request must include a 'model' form field".to_string(), + ) + })?; + + let timeout = match text_fields.iter().find(|(name, _)| name == "timeout") { + Some((_, value)) => positive_duration(Some(value.parse::().map_err(|_| { + CoreError::InvalidRequest(format!("invalid 'timeout' form field: {value:?}")) + })?)), + None => None, + }; + + let optional_params: Map = text_fields + .iter() + .filter(|(name, _)| !matches!(name.as_str(), "model" | "timeout" | "file" | "document")) + .map(|(name, value)| (name.clone(), coerce_form_field(value))) + .collect(); + + Ok(OcrCall { + model, + document, + optional_params, + timeout, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_document_url_and_flattens_provider_params() { + let call = parse_json_body( + br#"{"model":"rust-ocr","document":{"type":"document_url","document_url":"https://x/doc.pdf"},"include_image_base64":true}"#, + ) + .expect("valid body parses"); + + assert_eq!(call.model, "rust-ocr"); + assert_eq!(call.document["type"], "document_url"); + assert_eq!(call.document["document_url"], "https://x/doc.pdf"); + assert_eq!( + call.optional_params["include_image_base64"], + Value::Bool(true) + ); + assert!(call.timeout.is_none()); + } + + #[test] + fn parses_image_url_document() { + let call = parse_json_body( + br#"{"model":"m","document":{"type":"image_url","image_url":"https://x/i.png"}}"#, + ) + .expect("valid body parses"); + assert_eq!(call.document["type"], "image_url"); + assert_eq!(call.document["image_url"], "https://x/i.png"); + } + + #[test] + fn extracts_timeout_and_keeps_it_out_of_provider_params() { + let call = parse_json_body( + br#"{"model":"m","document":{"type":"document_url","document_url":"https://x"},"timeout":12.5}"#, + ) + .expect("valid body parses"); + assert_eq!(call.timeout, Some(Duration::from_secs_f64(12.5))); + assert!(!call.optional_params.contains_key("timeout")); + } + + #[test] + fn rejects_file_document_over_json() { + let err = + parse_json_body(br#"{"model":"m","document":{"type":"file","file":"/etc/passwd"}}"#) + .expect_err("file type rejected"); + match err { + CoreError::InvalidRequest(message) => assert!(message.contains("multipart/form-data")), + other => panic!("expected InvalidRequest, got {other:?}"), + } + } + + #[test] + fn rejects_reducto_file_id_over_json() { + let err = parse_json_body( + br#"{"model":"m","document":{"type":"document_url","document_url":"reducto://abc"}}"#, + ) + .expect_err("reducto id rejected"); + match err { + CoreError::InvalidRequest(message) => assert!(message.contains("reducto://")), + other => panic!("expected InvalidRequest, got {other:?}"), + } + } + + #[test] + fn empty_body_is_rejected() { + assert!(matches!( + parse_json_body(b""), + Err(CoreError::InvalidRequest(_)) + )); + } + + #[test] + fn upload_prefers_declared_content_type() { + let document = build_upload_document( + b"%PDF-1.4".to_vec(), + Some("scan.bin"), + Some("application/pdf"), + ) + .expect("builds document"); + assert_eq!(document["type"], "document_url"); + assert!(document["document_url"] + .as_str() + .expect("data uri") + .starts_with("data:application/pdf;base64,")); + } + + #[test] + fn upload_infers_mime_from_filename_when_octet_stream() { + let document = build_upload_document( + vec![0x89, b'P', b'N', b'G'], + Some("photo.PNG"), + Some("application/octet-stream"), + ) + .expect("builds document"); + assert_eq!(document["type"], "image_url"); + assert!(document["image_url"] + .as_str() + .expect("data uri") + .starts_with("data:image/png;base64,")); + } + + #[test] + fn upload_falls_back_to_octet_stream() { + let document = build_upload_document(vec![1, 2, 3], Some("data.unknown"), None) + .expect("builds document"); + assert_eq!(document["type"], "document_url"); + assert!(document["document_url"] + .as_str() + .expect("data uri") + .starts_with("data:application/octet-stream;base64,")); + } + + #[test] + fn empty_upload_is_rejected() { + assert!(matches!( + build_upload_document(Vec::new(), Some("a.pdf"), Some("application/pdf")), + Err(CoreError::InvalidRequest(_)) + )); + } + + #[test] + fn multipart_extracts_model_timeout_and_json_params() { + let document = + json!({"type": "document_url", "document_url": "data:application/pdf;base64,AA=="}); + let fields = vec![ + ("model".to_string(), "rust-ocr".to_string()), + ("timeout".to_string(), "30".to_string()), + ("pages".to_string(), "[0,1,2]".to_string()), + ("id".to_string(), "abc".to_string()), + ]; + let call = assemble_multipart_call(document, &fields).expect("assembles call"); + + assert_eq!(call.model, "rust-ocr"); + assert_eq!(call.timeout, Some(Duration::from_secs(30))); + assert_eq!(call.optional_params["pages"], json!([0, 1, 2])); + assert_eq!(call.optional_params["id"], Value::String("abc".to_string())); + assert!(!call.optional_params.contains_key("timeout")); + assert!(!call.optional_params.contains_key("model")); + } + + #[test] + fn multipart_requires_model() { + let document = json!({"type": "document_url", "document_url": "data:x"}); + let err = assemble_multipart_call(document, &[]).expect_err("model required"); + assert!(matches!(err, CoreError::InvalidRequest(_))); + } +}