mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
c264758eb8
commit
e55d2aae95
7 changed files with 878 additions and 1 deletions
33
litellm-rust/Cargo.lock
generated
33
litellm-rust/Cargo.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
];
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
414
litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs
Normal file
414
litellm-rust/crates/ai-gateway/src/routes/ocr/mod.rs
Normal 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");
|
||||
}
|
||||
}
|
||||
87
litellm-rust/crates/ai-gateway/src/routes/ocr/service.rs
Normal file
87
litellm-rust/crates/ai-gateway/src/routes/ocr/service.rs
Normal 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(¶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(_))
|
||||
));
|
||||
}
|
||||
}
|
||||
318
litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs
Normal file
318
litellm-rust/crates/ai-gateway/src/routes/ocr/transport.rs
Normal 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(_)));
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue