feat(rust-gateway): expose OCR over Axum

Mount authenticated POST /v1/ocr and /ocr on the standalone Axum gateway over the existing Rust OCR I/O. Strongly typed JSON and multipart adapters build a provider-compatible document, resolve the public model alias through core::router, propagate the request timeout, and return the normalized OCR response. Public errors are stable and sanitized so upstream provider bodies never cross the gateway boundary.

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-16 22:13:39 +00:00
parent c264758eb8
commit e55d2aae95
7 changed files with 878 additions and 1 deletions

View file

@ -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"

View file

@ -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

View file

@ -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"),
];

View file

@ -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)
}

View file

@ -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<AppState> {
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<AppState>,
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<OcrCall> {
if is_multipart(&request) {
parse_multipart(request, state).await
} else {
parse_json(request).await
}
}
async fn parse_json(request: Request) -> CoreResult<OcrCall> {
let bytes = read_body(request.into_body()).await?;
transport::parse_json_body(&bytes)
}
async fn read_body(body: Body) -> CoreResult<Vec<u8>> {
to_bytes(body, MAX_OCR_REQUEST_BYTES)
.await
.map(|bytes| bytes.to_vec())
.map_err(|err| CoreError::InvalidRequest(format!("could not read request body: {err}")))
}
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}"))
})?;
let mut file: Option<(Vec<u8>, Option<String>, Option<String>)> = 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<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 {
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");
}
}

View file

@ -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<Value> {
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(&params.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(_))
));
}
}

View file

@ -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<Value> {
let (field, url) = match self {
OcrDocument::DocumentUrl { document_url } => ("document_url", document_url),
OcrDocument::ImageUrl { image_url } => ("image_url", image_url),
OcrDocument::File {} => {
return Err(CoreError::InvalidRequest(
"document type 'file' is not supported through the JSON API; upload the \
file via multipart/form-data with a 'file' field, or use a 'document_url' \
or 'image_url' document type"
.to_string(),
))
}
};
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<f64>,
#[serde(flatten)]
pub optional_params: Map<String, Value>,
}
#[derive(Debug)]
pub struct OcrCall {
pub model: String,
pub document: Value,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
}
fn positive_duration(seconds: Option<f64>) -> Option<Duration> {
seconds
.filter(|secs| secs.is_finite() && *secs > 0.0)
.map(Duration::from_secs_f64)
}
pub fn parse_json_body(body: &[u8]) -> CoreResult<OcrCall> {
if body.is_empty() {
return Err(CoreError::InvalidRequest(
"empty request body; send a JSON body with 'model' and 'document', or use \
multipart/form-data for file uploads"
.to_string(),
));
}
let request: OcrJsonRequest = serde_json::from_slice(body)
.map_err(|err| CoreError::InvalidRequest(format!("invalid OCR request body: {err}")))?;
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<u8>,
filename: Option<&str>,
content_type: Option<&str>,
) -> CoreResult<Value> {
if bytes.is_empty() {
return Err(CoreError::InvalidRequest(
"uploaded file is empty".to_string(),
));
}
let mime = resolve_upload_mime(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<OcrCall> {
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::<f64>().map_err(|_| {
CoreError::InvalidRequest(format!("invalid 'timeout' form field: {value:?}"))
})?)),
None => None,
};
let optional_params: Map<String, Value> = text_fields
.iter()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "timeout" | "file" | "document"))
.map(|(name, value)| (name.clone(), coerce_form_field(value)))
.collect();
Ok(OcrCall {
model,
document,
optional_params,
timeout,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_document_url_and_flattens_provider_params() {
let call = parse_json_body(
br#"{"model":"rust-ocr","document":{"type":"document_url","document_url":"https://x/doc.pdf"},"include_image_base64":true}"#,
)
.expect("valid body parses");
assert_eq!(call.model, "rust-ocr");
assert_eq!(call.document["type"], "document_url");
assert_eq!(call.document["document_url"], "https://x/doc.pdf");
assert_eq!(
call.optional_params["include_image_base64"],
Value::Bool(true)
);
assert!(call.timeout.is_none());
}
#[test]
fn parses_image_url_document() {
let call = parse_json_body(
br#"{"model":"m","document":{"type":"image_url","image_url":"https://x/i.png"}}"#,
)
.expect("valid body parses");
assert_eq!(call.document["type"], "image_url");
assert_eq!(call.document["image_url"], "https://x/i.png");
}
#[test]
fn extracts_timeout_and_keeps_it_out_of_provider_params() {
let call = parse_json_body(
br#"{"model":"m","document":{"type":"document_url","document_url":"https://x"},"timeout":12.5}"#,
)
.expect("valid body parses");
assert_eq!(call.timeout, Some(Duration::from_secs_f64(12.5)));
assert!(!call.optional_params.contains_key("timeout"));
}
#[test]
fn rejects_file_document_over_json() {
let err =
parse_json_body(br#"{"model":"m","document":{"type":"file","file":"/etc/passwd"}}"#)
.expect_err("file type rejected");
match err {
CoreError::InvalidRequest(message) => assert!(message.contains("multipart/form-data")),
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
#[test]
fn 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(_)));
}
}