refactor(rust): use shared execution in gateway inference (#43463)

* feat(rust): add the HTTP host driver

* refactor(rust): use shared execution in gateway inference

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* style(rust): apply rustfmt

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Yujong Lee <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-27 16:50:05 -07:00 • committed by GitHub
parent 438bffc26e
commit 1ceeefbf84
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 870 additions and 292 deletions

View file

@ -3262,6 +3262,7 @@ dependencies = [
"futures-util",
"litellm-auth",
"litellm-core",
"litellm-host-http",
"litellm-http",
"litellm-llms",
"litellm-router",

View file

@ -1,3 +1,4 @@
pub mod route;
pub mod types;
pub use crate::error::RouteError as Error;
mod common_utils;

View file

@ -0,0 +1,44 @@
use std::convert::Infallible;
use litellm_host::{
call::{CallOutput, HostedMachine, hosted_call},
protocol::Protocol,
};
use litellm_types::utils::ChatCompletionsResponse;
use super::{
ChatCompletionsRoute, Error,
types::{ChatCompletionsCall, ChatCompletionsRequest},
};
pub struct ChatCompletions;
impl Protocol for ChatCompletions {
type Response = ChatCompletionsResponse;
type Error = Error;
type Request = ChatCompletionsCall;
type HostCall = Infallible;
type Chunk = Infallible;
type StreamHead = Infallible;
}
impl ChatCompletionsRoute {
pub fn machine(self, call: ChatCompletionsCall) -> HostedMachine<ChatCompletions> {
hosted_call(
call,
move |call: ChatCompletionsCall, _, hooks| async move {
let request = ChatCompletionsRequest {
model: &call.model,
messages: call.messages,
optional_params: call.optional_params,
api_key: call.api_key.as_deref(),
api_base: call.api_base.as_deref(),
custom_llm_provider: call.custom_llm_provider.as_deref(),
extra_headers: call.extra_headers,
timeout: call.timeout,
};
self.run(request, &hooks).await.map(CallOutput::Complete)
},
)
}
}

View file

@ -23,6 +23,32 @@ pub struct ChatCompletionsRequest<'a> {
pub timeout: Option<Duration>,
}
pub struct ChatCompletionsCall {
pub model: String,
pub messages: Value,
pub optional_params: Map<String, Value>,
pub api_key: Option<String>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
}
impl From<ChatCompletionsRequest<'_>> for ChatCompletionsCall {
fn from(request: ChatCompletionsRequest<'_>) -> Self {
Self {
model: request.model.into(),
messages: request.messages,
optional_params: request.optional_params,
api_key: request.api_key.map(str::to_owned),
api_base: request.api_base.map(str::to_owned),
custom_llm_provider: request.custom_llm_provider.map(str::to_owned),
extra_headers: request.extra_headers,
timeout: request.timeout,
}
}
}
pub struct ResolvedChatCompletionsRequest<'a> {
pub model: String,
pub custom_llm_provider: String,

View file

@ -323,3 +323,114 @@ async fn a_declined_request_fails_the_call_before_sending(
assert_eq!(error, Error::Unsupported("streaming"));
assert!(received(&upstream).await.is_empty());
}
#[rstest]
#[case::direct(false)]
#[case::hosted(true)]
#[tokio::test]
async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
request: ChatCompletionsRequest<'static>,
#[case] hosted: bool,
) {
use litellm_core::chat_completions::route::ChatCompletions;
use litellm_host::{call::HostedCompletion, event::CallEvent};
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
let host = RecordingCall::<ChatCompletions>::new(
ChatCompletionsRequest {
api_base: Some(&base),
..request
}
.into(),
);
let response = if hosted {
let result = litellm_host::in_process::run_hosted(
chat_completions_route().machine(host.request().unwrap()),
host.runtime(),
)
.await
.unwrap();
let HostedCompletion::Complete(response) = result else {
panic!("expected a complete response")
};
response
} else {
let call = host.request.lock().unwrap().take().unwrap();
chat_completions_route()
.execute(
ChatCompletionsRequest {
model: &call.model,
messages: call.messages,
optional_params: call.optional_params,
api_key: call.api_key.as_deref(),
api_base: call.api_base.as_deref(),
custom_llm_provider: call.custom_llm_provider.as_deref(),
extra_headers: call.extra_headers,
timeout: call.timeout,
},
&host,
)
.await
.unwrap()
};
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
assert_eq!(
only_request(&upstream).await.header("x-hook"),
Some("called")
);
let events = host.events.0.lock().unwrap();
assert!(matches!(
&events[..],
[
CallEvent::Started { .. },
CallEvent::Machine(_),
CallEvent::Succeeded { .. }
]
));
}
#[rstest]
#[tokio::test]
async fn a_post_call_hook_failure_never_looks_safe_to_retry(
request: ChatCompletionsRequest<'static>,
) {
use litellm_host::{
event::{MachineEvent, RequestContext, WireRequest},
hooks::RouteHooks,
};
struct FailingHook;
impl RouteHooks<Error> for FailingHook {
async fn before_provider_request(
&self,
wire: WireRequest,
_: RequestContext,
) -> Result<WireRequest, Error> {
Ok(wire)
}
async fn on_event(&self, _: MachineEvent) -> Result<(), Error> {
Err(Error::InvalidRequest("callback rejected".into()))
}
}
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();
let error = chat_completions_route()
.execute(
ChatCompletionsRequest {
api_base: Some(&base),
..request
},
&FailingHook,
)
.await
.unwrap_err();
assert_eq!(error.phase(), litellm_core::error::Phase::AfterSend);
let Error::PostCallHook(source) = error else {
panic!("expected retained callback error")
};
assert_eq!(*source, Error::InvalidRequest("callback rejected".into()));
assert_eq!(received(&upstream).await.len(), 1);
}

View file

@ -1,5 +1,6 @@
- Expose a mountable Axum router; listener binding, server lifecycle, and shared inbound middleware belong to `gateway`
- Own the public inference HTTP boundary: endpoint paths, request parsing, model alias resolution, response envelopes, and SSE delivery
- Own endpoint paths, request parsing, model alias resolution, and API-specific response and SSE error formats; delegate hosted call execution and HTTP body delivery to host-http
- Delegate inference execution to `core` and provider transformations and authentication to `llms` and the auth crates; do not duplicate them in handlers
- Let core validate inference fields and supported features, then map its errors to HTTP responses; do not add gateway checks for temporary core limitations
- Use injected deployments, HTTP pools, settings, and secret sources; do not load process configuration or construct independent clients in handlers
- Test HTTP contracts here, including status codes, forwarded headers, error envelopes, and streaming behavior; keep core and provider tests in their owning crates

View file

@ -6,12 +6,12 @@ license.workspace = true
repository.workspace = true
[dependencies]
axum = { workspace = true, features = ["json", "multipart"] }
axum = { workspace = true, features = ["json", "multipart", "original-uri"] }
base64.workspace = true
bytes.workspace = true
futures-util.workspace = true
litellm-auth.workspace = true
litellm-core.workspace = true
litellm-host-http.workspace = true
litellm-http.workspace = true
litellm-llms.workspace = true
litellm-router.workspace = true

View file

@ -1,26 +1,30 @@
use std::{path::Path, sync::Arc};
use axum::{
Json,
extract::{Request, State},
response::{IntoResponse, Response},
};
use axum::{Json, extract::State, response::IntoResponse};
use base64::{Engine, engine::general_purpose::STANDARD};
use litellm_core::audio_transcription::types::AudioTranscriptionRequest;
use serde_json::{Value, json};
use crate::{Error, Gateway, request};
use crate::{
Error, Gateway,
request::{self, InferenceBody},
};
pub(crate) async fn create(State(gateway): State<Arc<Gateway>>, request: Request) -> Response {
match handle(&gateway, request).await {
Ok(response) => Json(response).into_response(),
Err(error) => error.openai_response(),
}
pub(crate) async fn create(
State(gateway): State<Arc<Gateway>>,
body: InferenceBody,
) -> Result<impl IntoResponse, Error> {
handle(&gateway, body).await.map(Json)
}
async fn handle(gateway: &Gateway, request: Request) -> Result<Value, Error> {
let (body, upload) = request::parse(request).await?;
let deployment = request::deployment(gateway, &body)?;
async fn handle(
gateway: &Gateway,
InferenceBody {
fields: body,
upload,
}: InferenceBody,
) -> Result<Value, Error> {
let deployment = request::resolve_deployment(gateway, &body)?;
let audio = match upload {
Some(upload) => {
let format = upload
@ -28,15 +32,10 @@ async fn handle(gateway: &Gateway, request: Request) -> Result<Value, Error> {
.as_deref()
.and_then(|name| Path::new(name).extension())
.and_then(|extension| extension.to_str())
.ok_or_else(|| {
Error::InvalidBody("audio file requires a filename extension".into())
})?;
json!({"data": STANDARD.encode(upload.bytes), "format": format.to_ascii_lowercase()})
.map(str::to_ascii_lowercase);
json!({"data": STANDARD.encode(upload.bytes), "format": format})
}
None => body
.get("audio")
.cloned()
.ok_or_else(|| Error::InvalidBody("audio is required".into()))?,
None => body.get("audio").cloned().unwrap_or_default(),
};
Ok(gateway
.audio_transcription

View file

@ -2,83 +2,60 @@ use std::sync::Arc;
use axum::{
Json,
body::Bytes,
extract::{Path, State},
http::StatusCode,
response::{IntoResponse, Response},
};
use litellm_core::chat_completions::types::ChatCompletionsRequest;
use litellm_core::chat_completions::types::ChatCompletionsCall;
use serde_json::{Map, Value};
use crate::{Error, Gateway, request};
use crate::{Error, Gateway, JsonObject, request};
pub(crate) async fn create(State(gateway): State<Arc<Gateway>>, body: Bytes) -> Response {
respond(&gateway, request::object(&body)).await
}
pub(crate) async fn deployment(
pub(crate) async fn create(
State(gateway): State<Arc<Gateway>>,
Path(path): Path<String>,
body: Bytes,
) -> Response {
if let Some(model) = path
.strip_suffix("/chat/completions")
.filter(|model| !model.is_empty())
{
let body = request::object(&body).map(|body| {
if body.get("model").is_some_and(|model| !model.is_null()) {
return body;
}
body.into_iter()
.chain([("model".into(), Value::String(model.into()))])
.collect()
});
return respond(&gateway, body).await;
}
if path.ends_with("/embeddings") || path.ends_with("/completions") {
return Error::Unsupported(path).openai_response();
}
StatusCode::NOT_FOUND.into_response()
JsonObject(body): JsonObject,
) -> Result<impl IntoResponse, Error> {
handle(&gateway, body).await
}
async fn respond(gateway: &Gateway, body: Result<Map<String, Value>, Error>) -> Response {
let result = match body {
Ok(body) => handle(gateway, body).await,
Err(error) => Err(error),
pub(crate) async fn create_from_model_path(
State(gateway): State<Arc<Gateway>>,
Path(model): Path<String>,
JsonObject(body): JsonObject,
) -> Result<impl IntoResponse, Error> {
let body = match body.get("model") {
None | Some(Value::Null) => body
.into_iter()
.chain([("model".into(), Value::String(model))])
.collect(),
Some(_) => body,
};
match result {
Ok(response) => response,
Err(error) => error.openai_response(),
}
handle(&gateway, body).await
}
async fn handle(gateway: &Gateway, body: Map<String, Value>) -> Result<Response, Error> {
let deployment = request::deployment(gateway, &body)?;
if body.get("stream").and_then(Value::as_bool) == Some(true) {
return Err(Error::Unsupported("streaming chat completions".into()));
}
let messages = body
.get("messages")
.cloned()
.ok_or_else(|| Error::InvalidBody("messages is required".into()))?;
let response = gateway
.chat_completions
.execute(
ChatCompletionsRequest {
model: &deployment.model,
let deployment = request::resolve_deployment(gateway, &body)?;
let messages = body.get("messages").cloned().unwrap_or_default();
let response = litellm_host_http::serve_unary(
gateway
.chat_completions
.clone()
.machine(ChatCompletionsCall {
model: deployment.model.clone(),
messages,
optional_params: body
.into_iter()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages" | "stream"))
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages"))
.collect(),
api_key: deployment.api_key.as_deref(),
api_base: deployment.api_base.as_deref(),
custom_llm_provider: deployment.custom_llm_provider.as_deref(),
api_key: deployment.api_key.clone(),
api_base: deployment.api_base.clone(),
custom_llm_provider: deployment.custom_llm_provider.clone(),
extra_headers: None,
timeout: deployment.timeout,
},
&(),
)
.await?;
Ok(Json(response).into_response())
}),
(),
(),
litellm_host_http::Unary::new(Json),
)
.await?;
Ok(response)
}

View file

@ -28,6 +28,21 @@ pub enum Error {
Internal(String),
}
impl IntoResponse for Error {
fn into_response(self) -> Response {
self.openai_response()
}
}
impl From<litellm_host_http::Error<RouteError>> for Error {
fn from(error: litellm_host_http::Error<RouteError>) -> Self {
match error {
litellm_host_http::Error::Call(error) => Self::Route(error),
litellm_host_http::Error::Protocol => Self::Internal(error.to_string()),
}
}
}
impl Error {
pub fn status(&self) -> StatusCode {
match self {

View file

@ -23,6 +23,7 @@ use litellm_secrets::source::SecretSource;
pub use error::Error;
pub use litellm_router::{Deployment, Router as ModelList};
pub use request::{JsonObject, RequestId};
pub struct Gateway {
pub audio_transcription: AudioTranscriptionRoute,
@ -79,11 +80,8 @@ pub fn router(gateway: Arc<Gateway>) -> Router {
.route("/v1/ocr", post(ocr::create))
.route("/chat/completions", post(chat_completions::create))
.route("/v1/chat/completions", post(chat_completions::create))
.route("/engines/{*path}", post(chat_completions::deployment))
.route(
"/openai/deployments/{*path}",
post(chat_completions::deployment),
)
.nest("/engines/{model}", model_routes())
.nest("/openai/deployments/{model}", model_routes())
.route("/audio/transcriptions", post(audio_transcription::create))
.route(
"/v1/audio/transcriptions",
@ -100,3 +98,13 @@ pub fn router(gateway: Arc<Gateway>) -> Router {
))
.with_state(gateway)
}
fn model_routes() -> Router<Arc<Gateway>> {
Router::new()
.route(
"/chat/completions",
post(chat_completions::create_from_model_path),
)
.route("/embeddings", post(request::unsupported))
.route("/completions", post(request::unsupported))
}

View file

@ -0,0 +1,87 @@
//! `POST /v1/messages`, as the Python proxy's `anthropic_response` serves it.
use std::sync::Arc;
use axum::{
Json,
body::Bytes,
extract::State,
http::HeaderMap,
response::{IntoResponse, Response},
};
use litellm_core::messages::{MessagesCall, messages_body, route::Messages};
use litellm_host_http::Sse;
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
use serde_json::{Map, Value};
use crate::{Deployment, Error, Gateway, JsonObject, RequestId, request};
/// Client headers Python forwards to Anthropic-speaking providers on every call.
const ANTHROPIC_API_HEADERS: [&str; 2] = ["anthropic-version", "anthropic-beta"];
const ANTHROPIC_API_HEADER_PROVIDERS: &str = "anthropic,bedrock,bedrock_mantle,vertex_ai";
pub async fn create(
State(gateway): State<Arc<Gateway>>,
RequestId(request_id): RequestId,
headers: HeaderMap,
body: Result<JsonObject, Error>,
) -> impl IntoResponse {
let result = match body {
Ok(JsonObject(body)) => handle(&gateway, &headers, body).await,
Err(error) => Err(error),
};
result.map_err(|error| (error.status(), Json(error.body(request_id.as_deref()))))
}
async fn handle(
gateway: &Gateway,
headers: &HeaderMap,
body: Map<String, Value>,
) -> Result<Response, Error> {
let deployment = request::resolve_deployment(gateway, &body)?;
let call = project(deployment, body, headers)?;
let machine = gateway.messages.clone().machine(call);
let stream =
Sse::<Messages, _, _>::new(Json, |error| Bytes::from(Error::from(error).sse_frame()));
Ok(litellm_host_http::serve(machine, (), (), stream).await?)
}
fn project(
deployment: &Deployment,
body: Map<String, Value>,
headers: &HeaderMap,
) -> Result<MessagesCall, Error> {
let body = body
.into_iter()
.map(|(name, value)| match name.as_str() {
"model" => (name, Value::from(deployment.model.as_str())),
_ => (name, value),
})
.collect();
Ok(MessagesCall {
body: messages_body(body)?,
api_key: deployment.api_key.clone(),
api_base: deployment.api_base.clone(),
custom_llm_provider: deployment.custom_llm_provider.clone(),
extra_headers: None,
provider_specific_header: anthropic_api_headers(headers),
timeout: deployment.timeout,
shaping: deployment.shaping.clone(),
})
}
fn anthropic_api_headers(headers: &HeaderMap) -> Option<ProviderSpecificHeaders> {
let extra_headers: Map<String, Value> = ANTHROPIC_API_HEADERS
.into_iter()
.filter_map(|name| {
let value = headers.get(name)?.to_str().ok()?;
Some((name.to_owned(), Value::from(value)))
})
.collect();
(!extra_headers.is_empty()).then(|| {
ProviderSpecificHeaders::One(ProviderSpecificHeader {
custom_llm_provider: ANTHROPIC_API_HEADER_PROVIDERS.into(),
extra_headers,
})
})
}

View file

@ -1,113 +0,0 @@
//! `POST /v1/messages`, as the Python proxy's `anthropic_response` serves it.
use std::{convert::Infallible, sync::Arc};
use axum::{
Json,
body::{Body, Bytes},
extract::State,
http::{HeaderMap, StatusCode, header},
response::{IntoResponse, Response},
};
use futures_util::{StreamExt, stream::BoxStream};
use litellm_core::messages::{Error as RouteError, MessagesCall, MessagesResponse, messages_body};
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
use serde_json::{Map, Value};
use crate::{Deployment, Error, Gateway};
/// Client headers Python forwards to Anthropic-speaking providers on every call.
const ANTHROPIC_API_HEADERS: [&str; 2] = ["anthropic-version", "anthropic-beta"];
const ANTHROPIC_API_HEADER_PROVIDERS: &str = "anthropic,bedrock,bedrock_mantle,vertex_ai";
pub async fn create(
State(gateway): State<Arc<Gateway>>,
headers: HeaderMap,
body: Bytes,
) -> Response {
let request_id = headers
.get("x-request-id")
.and_then(|value| value.to_str().ok())
.map(str::to_owned);
match handle(&gateway, &headers, &body).await {
Ok(response) => response,
Err(error) => (error.status(), Json(error.body(request_id.as_deref()))).into_response(),
}
}
async fn handle(gateway: &Gateway, headers: &HeaderMap, body: &[u8]) -> Result<Response, Error> {
let body = match serde_json::from_slice(body) {
Ok(Value::Object(body)) => body,
Ok(_) => return Err(Error::InvalidBody("expected a JSON object".into())),
Err(error) => return Err(Error::InvalidBody(error.to_string())),
};
let model_name = body
.get("model")
.and_then(Value::as_str)
.ok_or_else(|| Error::InvalidBody("model is required".into()))?;
let deployment = gateway
.models
.get(model_name)
.ok_or_else(|| Error::UnknownModel(model_name.to_owned()))?;
let call = project(deployment, body, headers)?;
match gateway.messages.execute(call, &()).await? {
MessagesResponse::Complete(message) => Ok(Json(message).into_response()),
MessagesResponse::Stream { chunks, .. } => Ok(stream(chunks)),
}
}
fn project(
deployment: &Deployment,
body: Map<String, Value>,
headers: &HeaderMap,
) -> Result<MessagesCall, Error> {
let body = body
.into_iter()
.map(|(name, value)| match name.as_str() {
"model" => (name, Value::from(deployment.model.as_str())),
_ => (name, value),
})
.collect();
Ok(MessagesCall {
body: messages_body(body)?,
api_key: deployment.api_key.clone(),
api_base: deployment.api_base.clone(),
custom_llm_provider: deployment.custom_llm_provider.clone(),
extra_headers: None,
provider_specific_header: anthropic_api_headers(headers),
timeout: deployment.timeout,
shaping: deployment.shaping.clone(),
})
}
fn anthropic_api_headers(headers: &HeaderMap) -> Option<ProviderSpecificHeaders> {
let extra_headers: Map<String, Value> = ANTHROPIC_API_HEADERS
.into_iter()
.filter_map(|name| {
let value = headers.get(name)?.to_str().ok()?;
Some((name.to_owned(), Value::from(value)))
})
.collect();
(!extra_headers.is_empty()).then(|| {
ProviderSpecificHeaders::One(ProviderSpecificHeader {
custom_llm_provider: ANTHROPIC_API_HEADER_PROVIDERS.into(),
extra_headers,
})
})
}
/// A chunk that fails after the stream opened is delivered as an SSE error frame, since
/// the status line already went out; the stream ends on it.
fn stream(chunks: BoxStream<'static, Result<Bytes, RouteError>>) -> Response {
let body = chunks.map(|chunk| {
Ok::<_, Infallible>(
chunk.unwrap_or_else(|error| Bytes::from(Error::Route(error).sse_frame())),
)
});
(
StatusCode::OK,
[(header::CONTENT_TYPE, "text/event-stream")],
Body::from_stream(body),
)
.into_response()
}

View file

@ -1,44 +1,44 @@
use std::sync::Arc;
use axum::{
Json,
extract::{Request, State},
response::{IntoResponse, Response},
};
use axum::{Json, extract::State, http::HeaderMap, response::IntoResponse};
use litellm_auth::SecretValue;
use litellm_core::ocr::types::{LiteLLMOcrRequest, OcrConnectionInputs, OcrDocumentInput};
use litellm_llms::base_llm::ocr::transformation::OcrDocument;
use serde_json::Value;
use crate::{Error, Gateway, request};
use crate::{
Error, Gateway,
request::{self, InferenceBody},
};
pub(crate) async fn create(State(gateway): State<Arc<Gateway>>, request: Request) -> Response {
match handle(&gateway, request).await {
Ok(response) => Json(response).into_response(),
Err(error) => error.openai_response(),
}
pub(crate) async fn create(
State(gateway): State<Arc<Gateway>>,
headers: HeaderMap,
body: InferenceBody,
) -> Result<impl IntoResponse, Error> {
handle(&gateway, &headers, body).await.map(Json)
}
async fn handle(gateway: &Gateway, request: Request) -> Result<Value, Error> {
let header_format = request
.headers()
async fn handle(
gateway: &Gateway,
headers: &HeaderMap,
InferenceBody {
fields: body,
upload,
}: InferenceBody,
) -> Result<Value, Error> {
let header_format = headers
.get("x-req-format")
.and_then(|value| value.to_str().ok())
.map(str::to_owned);
let (body, upload) = request::parse(request).await?;
let deployment = request::deployment(gateway, &body)?;
let deployment = request::resolve_deployment(gateway, &body)?;
let document = match upload {
Some(upload) => OcrDocumentInput::Bytes {
bytes: upload.bytes,
file_name: upload.file_name,
mime_type: upload.mime_type,
},
None => OcrDocument::try_from(
body.get("document")
.cloned()
.ok_or_else(|| Error::InvalidBody("document is required".into()))?,
)?
.into(),
None => OcrDocument::try_from(body.get("document").cloned().unwrap_or_default())?.into(),
};
let format = body
.get("req_format")

View file

@ -1,8 +1,8 @@
use axum::{
body::{Bytes, to_bytes},
extract::{FromRequest, Multipart, Request},
http::Uri,
response::Response,
body::Bytes,
extract::{FromRequest, FromRequestParts, Multipart, OriginalUri, Request},
http::request::Parts,
response::IntoResponse,
};
use serde_json::{Map, Value};
@ -17,7 +17,72 @@ pub(crate) struct Upload {
pub mime_type: Option<String>,
}
pub(crate) fn object(body: &[u8]) -> Result<Map<String, Value>, Error> {
pub struct JsonObject(pub Map<String, Value>);
pub struct RequestId(pub Option<String>);
impl<S: Send + Sync> FromRequestParts<S> for RequestId {
type Rejection = std::convert::Infallible;
async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Self::Rejection> {
Ok(Self(
parts
.headers
.get("x-request-id")
.and_then(|value| value.to_str().ok())
.map(str::to_owned),
))
}
}
impl<S: Send + Sync> FromRequest<S> for JsonObject {
type Rejection = Error;
async fn from_request(request: Request, state: &S) -> Result<Self, Self::Rejection> {
let body = Bytes::from_request(request, state).await.map_err(|error| {
if error.status() == axum::http::StatusCode::PAYLOAD_TOO_LARGE {
Error::BodyTooLarge
} else {
Error::InvalidBody(error.to_string())
}
})?;
object(&body).map(Self)
}
}
pub(crate) struct InferenceBody {
pub fields: Map<String, Value>,
pub upload: Option<Upload>,
}
impl<S: Send + Sync> FromRequest<S> for InferenceBody {
type Rejection = Error;
async fn from_request(request: Request, state: &S) -> Result<Self, Self::Rejection> {
let multipart = request
.headers()
.get("content-type")
.and_then(|header| header.to_str().ok())
.is_some_and(|value| {
value
.to_ascii_lowercase()
.starts_with("multipart/form-data")
});
if !multipart {
let JsonObject(fields) = JsonObject::from_request(request, state).await?;
return Ok(Self {
fields,
upload: None,
});
}
let multipart = Multipart::from_request(request, state)
.await
.map_err(|error| Error::InvalidBody(error.to_string()))?;
parse_multipart(multipart).await
}
}
fn object(body: &[u8]) -> Result<Map<String, Value>, Error> {
match serde_json::from_slice(body) {
Ok(Value::Object(body)) => Ok(body),
Ok(_) => Err(Error::InvalidBody("expected a JSON object".into())),
@ -25,7 +90,7 @@ pub(crate) fn object(body: &[u8]) -> Result<Map<String, Value>, Error> {
}
}
pub(crate) fn deployment<'a>(
pub(crate) fn resolve_deployment<'a>(
gateway: &'a Gateway,
body: &Map<String, Value>,
) -> Result<&'a Deployment, Error> {
@ -39,25 +104,7 @@ pub(crate) fn deployment<'a>(
.ok_or_else(|| Error::UnknownModel(model.to_owned()))
}
pub(crate) async fn parse(request: Request) -> Result<(Map<String, Value>, Option<Upload>), Error> {
let multipart = request
.headers()
.get("content-type")
.and_then(|header| header.to_str().ok())
.is_some_and(|value| {
value
.to_ascii_lowercase()
.starts_with("multipart/form-data")
});
if !multipart {
let body = to_bytes(request.into_body(), MAX_BODY_BYTES)
.await
.map_err(|_| Error::BodyTooLarge)?;
return Ok((object(&body)?, None));
}
let mut multipart = Multipart::from_request(request, &())
.await
.map_err(|error| Error::InvalidBody(error.to_string()))?;
async fn parse_multipart(mut multipart: Multipart) -> Result<InferenceBody, Error> {
let mut fields = Map::new();
let mut upload = None;
while let Some(field) = multipart.next_field().await.map_err(multipart_error)? {
@ -74,9 +121,6 @@ pub(crate) async fn parse(request: Request) -> Result<(Map<String, Value>, Optio
if bytes.len() > MAX_FILE_BYTES {
return Err(Error::BodyTooLarge);
}
if bytes.is_empty() {
return Err(Error::InvalidBody("uploaded file is empty".into()));
}
upload = Some(Upload {
bytes,
file_name,
@ -88,12 +132,7 @@ pub(crate) async fn parse(request: Request) -> Result<(Map<String, Value>, Optio
fields.insert(name, value);
}
}
if upload.is_none() {
return Err(Error::InvalidBody(
"multipart request requires a file field".into(),
));
}
Ok((fields, upload))
Ok(InferenceBody { fields, upload })
}
fn multipart_error(error: axum::extract::multipart::MultipartError) -> Error {
@ -103,6 +142,55 @@ fn multipart_error(error: axum::extract::multipart::MultipartError) -> Error {
Error::InvalidBody(error.to_string())
}
pub(crate) async fn unsupported(uri: Uri) -> Response {
Error::Unsupported(uri.path().to_owned()).openai_response()
pub(crate) async fn unsupported(OriginalUri(uri): OriginalUri) -> impl IntoResponse {
Error::Unsupported(uri.path().to_owned())
}
#[cfg(test)]
mod tests {
use axum::{
Router,
body::{Body, to_bytes},
extract::DefaultBodyLimit,
routing::post,
};
use rstest::{fixture, rstest};
use tower::ServiceExt;
use super::*;
#[fixture]
fn limited_uploads() -> Router {
Router::new()
.route("/", post(|_: InferenceBody| async {}))
.layer(DefaultBodyLimit::max(64))
}
#[rstest]
#[case::json("application/json", format!("{{\"text\":\"{}\"}}", "x".repeat(64)))]
#[case::multipart(
"multipart/form-data; boundary=test",
format!("--test\r\nContent-Disposition: form-data; name=\"file\"; filename=\"file.pdf\"\r\n\r\n{}\r\n--test--\r\n", "x".repeat(64)),
)]
#[tokio::test]
async fn uploads_respect_the_configured_body_limit(
limited_uploads: Router,
#[case] content_type: &str,
#[case] payload: String,
) {
let response = limited_uploads
.oneshot(
Request::post("/")
.header("content-type", content_type)
.body(Body::from(payload))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), 413);
let body: Value =
serde_json::from_slice(&to_bytes(response.into_body(), 4096).await.unwrap()).unwrap();
assert_eq!(body["error"]["type"], "request_too_large");
assert_eq!(body["error"]["code"], 413);
}
}

View file

@ -14,10 +14,15 @@ use wiremock::{
};
#[rstest]
#[case(false)]
#[case(true)]
#[case::anthropic("anthropic/test-model", "/v1/messages", true)]
#[case::azure("azure_ai/test-model", "/anthropic/v1/messages", false)]
#[tokio::test]
async fn messages_reaches_the_provider_and_preserves_json_or_sse(#[case] streaming: bool) {
async fn messages_reaches_the_provider_and_preserves_json_or_sse(
#[case] model: &str,
#[case] upstream_path: &str,
#[case] anthropic_headers: bool,
#[values(false, true)] streaming: bool,
) {
let upstream = MockServer::start().await;
let message = json!({"id": "msg_test", "type": "message", "role": "assistant",
"model": "test-model", "content": [{"type": "text", "text": "hello"}],
@ -29,15 +34,29 @@ async fn messages_reaches_the_provider_and_preserves_json_or_sse(#[case] streami
ResponseTemplate::new(200).set_body_json(message.clone())
};
let messages = json!([{"role": "user", "content": "hi"}]);
Mock::given(method("POST")).and(path("/v1/messages"))
.and(header("x-api-key", "test-key"))
.and(header("anthropic-beta", "test-feature"))
.and(body_json(json!({"model": "test-model", "messages": messages, "max_tokens": 16, "stream": streaming})))
.respond_with(template).expect(1).mount(&upstream).await;
let mut mock = Mock::given(method("POST"))
.and(path(upstream_path))
.and(header("x-api-key", "test-key"));
if anthropic_headers {
mock = mock
.and(header("anthropic-beta", "test-feature"))
.and(header("anthropic-version", "test-version"));
}
mock.and(body_json(
json!({"model": "test-model", "messages": messages, "max_tokens": 16, "stream": streaming}),
))
.respond_with(template)
.expect(1)
.mount(&upstream)
.await;
let request = Request::post("/v1/messages")
.header("content-type", "application/json").header("anthropic-beta", "test-feature")
.header("anthropic-version", "test-version")
.header("x-api-key", "caller-key")
.header("authorization", "Bearer proxy-key")
.header("x-request-id", "caller-request")
.body(Body::from(json!({"model": "public/model", "messages": messages, "max_tokens": 16, "stream": streaming}).to_string())).unwrap();
let response = support::app("anthropic/test-model", &upstream.uri())
let response = support::app(model, &upstream.uri())
.oneshot(request)
.await
.unwrap();
@ -50,6 +69,10 @@ async fn messages_reaches_the_provider_and_preserves_json_or_sse(#[case] streami
assert_eq!(body["content"], message["content"]);
assert_eq!(body["usage"], message["usage"]);
}
let requests = upstream.received_requests().await.unwrap();
assert_eq!(requests.len(), 1);
assert!(!requests[0].headers.contains_key("authorization"));
assert!(!requests[0].headers.contains_key("x-request-id"));
}
#[tokio::test]
@ -119,3 +142,33 @@ async fn a_stream_that_fails_after_opening_ends_with_an_sse_error_frame() {
assert_eq!(error["type"], "error");
assert_eq!(error["error"]["type"], "api_error");
}
#[rstest]
#[tokio::test]
async fn hosted_provider_failure_preserves_status_body_and_request_id() {
let upstream = MockServer::start().await;
let error =
json!({"type": "error", "error": {"type": "rate_limit_error", "message": "retry later"}});
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(429).set_body_json(error.clone()))
.expect(1)
.mount(&upstream)
.await;
let request = Request::post("/v1/messages")
.header("content-type", "application/json")
.header("x-request-id", "host-http-request")
.body(Body::from(json!({
"model": "public/model", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16,
}).to_string()))
.unwrap();
let response = support::app("anthropic/test-model", &upstream.uri())
.oneshot(request)
.await
.unwrap();
assert_eq!(response.status(), 429);
let body = support::json(response).await;
assert_eq!(body["type"], error["type"]);
assert_eq!(body["error"], error["error"]);
assert_eq!(body["request_id"], "host-http-request");
}

View file

@ -1,6 +1,8 @@
mod support;
use axum::{body::Body, http::Request};
use litellm_gateway_inference::Error;
use litellm_llms::base_llm::ocr::{error::Error as OcrError, transformation::OcrDocument};
use rstest::rstest;
use serde_json::{Value, json};
use tower::ServiceExt;
@ -116,3 +118,71 @@ async fn invalid_ocr_requests_do_not_call_the_provider(#[case] body: Value) {
assert!(support::json(response).await["error"]["message"].is_string());
assert!(upstream.received_requests().await.unwrap().is_empty());
}
#[rstest]
#[tokio::test]
async fn malformed_multipart_uses_an_openai_error_envelope(
#[values("/v1/ocr", "/v1/audio/transcriptions")] route: &str,
) {
let upstream = MockServer::start().await;
let response = support::app("mistral/test-ocr", &upstream.uri())
.oneshot(
Request::post(route)
.header("content-type", "multipart/form-data")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), 400);
let body = support::json(response).await;
assert_eq!(body["error"]["type"], "invalid_request_error");
assert_eq!(body["error"]["code"], 400);
let error_message = body["error"]["message"].as_str().unwrap();
assert!(!error_message.is_empty());
assert!(upstream.received_requests().await.unwrap().is_empty());
}
#[rstest]
#[case::missing_document(
"/v1/ocr", "mistral/test-ocr", "",
Error::Ocr(OcrDocument::try_from(Value::Null).unwrap_err()),
)]
#[case::empty_document(
"/v1/ocr",
"mistral/test-ocr",
"--test\r\nContent-Disposition: form-data; name=\"file\"; filename=\"test.pdf\"\r\n\r\n\r\n",
Error::Ocr(OcrError::EmptyFile)
)]
#[case::empty_audio(
"/v1/audio/transcriptions", "bedrock/test-model",
"--test\r\nContent-Disposition: form-data; name=\"file\"; filename=\"test.wav\"\r\n\r\n\r\n",
Error::Route(litellm_llms::Error::MissingField("audio.data").into()),
)]
#[tokio::test]
async fn upload_validation_errors_come_from_core(
#[case] route: &str,
#[case] model: &str,
#[case] file: &str,
#[case] error: Error,
) {
let upstream = MockServer::start().await;
let payload = format!(
"--test\r\nContent-Disposition: form-data; name=\"model\"\r\n\r\npublic/model\r\n{file}--test--\r\n"
);
let response = support::app(model, &upstream.uri())
.oneshot(
Request::post(route)
.header("content-type", "multipart/form-data; boundary=test")
.body(Body::from(payload))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), 400);
assert_eq!(
support::json(response).await,
support::json(error.openai_response()).await
);
assert!(upstream.received_requests().await.unwrap().is_empty());
}

View file

@ -1,22 +1,38 @@
mod support;
use std::sync::Arc;
use axum::{
body::{Body, Bytes},
http::Request,
};
use litellm_core::{
chat_completions::{ChatCompletionsRoute, types::ChatCompletionsRequest},
resources::CoreResources,
};
use litellm_http::{
ClientVariant, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver,
};
use rstest::rstest;
use serde_json::{Value, json};
use tower::ServiceExt;
use wiremock::{
Mock, MockServer, ResponseTemplate,
matchers::{body_partial_json, method},
};
#[rstest]
#[case("/chat/completions", Some("public/model"))]
#[case("/v1/chat/completions", Some("public/model"))]
#[case("/engines/public/model/chat/completions", None)]
#[case("/openai/deployments/public/model/chat/completions", None)]
#[case("/openai/deployments/unused/chat/completions", Some("public/model"))]
#[case::chat("/chat/completions", Some("public/model"))]
#[case::versioned_chat("/v1/chat/completions", Some("public/model"))]
#[case::engine("/engines/public%2Fmodel/chat/completions", None)]
#[case::deployment("/openai/deployments/public%2Fmodel/chat/completions", None)]
#[case::body_model_wins("/openai/deployments/unused/chat/completions", Some("public/model"))]
#[tokio::test]
async fn chat_aliases_call_core_and_use_the_body_model_before_the_path(
#[case] route: &str,
#[case] model: Option<&str>,
#[values(None, Some("application/json"), Some("text/plain"))] content_type: Option<&str>,
#[values(None, Some(false))] stream: Option<bool>,
) {
let upstream = MockServer::start().await;
let messages = json!([{"role": "user", "content": "hi"}]);
@ -31,12 +47,21 @@ async fn chat_aliases_call_core_and_use_the_body_model_before_the_path(
.expect(1)
.mount(&upstream)
.await;
let response = support::post(
support::app("anthropic/test-model", &upstream.uri()),
route,
json!({"model": model, "messages": messages, "max_tokens": 16}),
)
.await;
let request = Request::post(route);
let request = match content_type {
Some(content_type) => request.header("content-type", content_type),
None => request,
};
let response = support::app("anthropic/test-model", &upstream.uri())
.oneshot(
request
.body(Body::from(
json!({"model": model, "messages": messages, "max_tokens": 16, "stream": stream}).to_string(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), 200);
assert_eq!(
support::json(response).await["choices"][0]["message"]["content"],
@ -45,14 +70,93 @@ async fn chat_aliases_call_core_and_use_the_body_model_before_the_path(
}
#[rstest]
#[case("/responses")]
#[case("/v1/responses")]
#[case::streaming(json!({"messages": [{"role": "user", "content": "hi"}], "stream": true}), 501)]
#[case::missing_messages(json!({}), 400)]
#[case::malformed_messages(json!({"messages": "hi"}), 400)]
#[case::invalid_streaming_request(json!({"messages": [], "stream": true}), 400)]
#[tokio::test]
async fn chat_errors_come_from_core(
#[case] fields: Value,
#[case] status: u16,
#[values("/v1/chat/completions", "/engines/public%2Fmodel/chat/completions")] path: &str,
) {
let upstream = MockServer::start().await;
let base = upstream.uri();
let resources = CoreResources::new(Arc::new(HttpClientPool::new(Arc::new(PublicDnsResolver))));
let http = Resolution::from(&HttpSettings::default()).config;
let provider = resources
.pool
.client(&http, ClientVariant::Provider)
.unwrap();
let fields = fields.as_object().unwrap();
let error = ChatCompletionsRoute::new(
provider,
resources.auth.clone(),
Arc::new(support::NoSecrets),
)
.execute(
ChatCompletionsRequest {
model: "anthropic/test-model",
messages: fields.get("messages").cloned().unwrap_or_default(),
optional_params: fields
.iter()
.filter(|(name, _)| name.as_str() != "messages")
.map(|(name, value)| (name.clone(), value.clone()))
.collect(),
api_key: Some("test-key"),
api_base: Some(&base),
custom_llm_provider: None,
extra_headers: None,
timeout: None,
},
&(),
)
.await
.unwrap_err();
let body = fields
.iter()
.map(|(name, value)| (name.clone(), value.clone()))
.chain([("model".into(), json!("public/model"))])
.collect();
let response = support::post(
support::app("anthropic/test-model", &base),
path,
Value::Object(body),
)
.await;
assert_eq!(response.status(), status);
let body = support::json(response).await;
assert_eq!(body["error"]["message"], error.to_string());
assert_eq!(body["error"]["code"], status);
assert!(upstream.received_requests().await.unwrap().is_empty());
}
#[rstest]
#[case::audio("/v1/audio/transcriptions", "bedrock/test-model", "audio")]
#[case::document("/v1/ocr", "mistral/test-ocr", "document")]
#[tokio::test]
async fn missing_inference_fields_use_the_same_validation_as_null(
#[case] path: &str,
#[case] model: &str,
#[case] field: &str,
) {
let upstream = MockServer::start().await;
let app = support::app(model, &upstream.uri());
let missing = support::post(app.clone(), path, json!({"model": "public/model"})).await;
let null = support::post(app, path, json!({"model": "public/model", field: null})).await;
assert_eq!(missing.status(), 400);
assert_eq!(missing.status(), null.status());
assert_eq!(support::json(missing).await, support::json(null).await);
assert!(upstream.received_requests().await.unwrap().is_empty());
}
#[rstest]
#[case("/embeddings")]
#[case("/v1/embeddings")]
#[case("/completions")]
#[case("/v1/completions")]
#[case("/engines/public/model/embeddings")]
#[case("/openai/deployments/public/model/completions")]
#[case("/engines/public%2Fmodel/embeddings")]
#[case("/openai/deployments/public%2Fmodel/completions")]
#[tokio::test]
async fn unimplemented_routes_return_an_explicit_error(#[case] path: &str) {
let response = support::post(
@ -62,12 +166,10 @@ async fn unimplemented_routes_return_an_explicit_error(#[case] path: &str) {
)
.await;
assert_eq!(response.status(), 501);
assert!(
support::json(response).await["error"]["message"]
.as_str()
.unwrap()
.contains("not implemented")
);
let body = support::json(response).await;
let message = body["error"]["message"].as_str().unwrap();
assert!(message.contains("not implemented"));
assert!(message.contains(path));
}
#[rstest]
@ -90,3 +192,111 @@ async fn transcription_aliases_reach_core_validation(#[case] path: &str) {
.contains("audio.format")
);
}
#[rstest]
#[case::no_extension("test")]
#[case::unsupported_extension("test.invalid")]
#[tokio::test]
async fn upload_audio_format_validation_matches_core(#[case] filename: &str) {
let upstream = MockServer::start().await;
let app = support::app("bedrock/test-model", &upstream.uri());
let path = "/v1/audio/transcriptions";
let expected = support::post(
app.clone(),
path,
json!({"model": "public/model", "audio": {"data": "YWJj", "format": null}}),
)
.await;
let payload = format!(
"--test\r\nContent-Disposition: form-data; name=\"model\"\r\n\r\npublic/model\r\n\
--test\r\nContent-Disposition: form-data; name=\"file\"; filename=\"{filename}\"\r\n\r\nabc\r\n--test--\r\n"
);
let response = app
.oneshot(
Request::post(path)
.header("content-type", "multipart/form-data; boundary=test")
.body(Body::from(payload))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), 400);
assert_eq!(support::json(response).await, support::json(expected).await);
assert!(upstream.received_requests().await.unwrap().is_empty());
}
#[rstest]
#[case::chat("/v1/chat/completions", false)]
#[case::deployment("/openai/deployments/public%2Fmodel/chat/completions", false)]
#[case::messages("/v1/messages", true)]
#[case::ocr("/v1/ocr", false)]
#[case::transcription("/v1/audio/transcriptions", false)]
#[tokio::test]
async fn json_extraction_rejections_use_the_endpoint_error_envelope(
#[case] path: &str,
#[case] anthropic: bool,
#[values("syntax", "array", "oversized", "read_failure")] failure: &str,
) {
let upstream = MockServer::start().await;
let payload = match failure {
"syntax" => Body::from("{"),
"array" => Body::from("[]"),
"oversized" => Body::from_stream(futures_util::stream::iter(
std::iter::repeat_n(Bytes::from(vec![b' '; 1024 * 1024]), 52)
.map(Ok::<_, std::io::Error>),
)),
"read_failure" => Body::from_stream(futures_util::stream::once(async {
Err::<Bytes, _>(std::io::Error::other("body read failed"))
})),
_ => unreachable!(),
};
let request = Request::post(path)
.header("x-request-id", "extractor-request")
.body(payload)
.unwrap();
let response = support::app("anthropic/test-model", &upstream.uri())
.oneshot(request)
.await
.unwrap();
let status = if failure == "oversized" { 413 } else { 400 };
assert_eq!(response.status(), status);
let body = support::json(response).await;
assert_eq!(
body["error"]["type"],
if status == 413 {
"request_too_large"
} else {
"invalid_request_error"
}
);
assert!(
body["error"]["message"]
.as_str()
.is_some_and(|message| !message.is_empty())
);
if anthropic {
assert_eq!(body["type"], "error");
assert_eq!(body["request_id"], "extractor-request");
} else {
assert_eq!(body["error"]["code"], status);
assert!(body["error"]["param"].is_null());
}
assert!(upstream.received_requests().await.unwrap().is_empty());
}
#[rstest]
#[case::unsupported("/engines/public%2Fmodel/embeddings", 501)]
#[case::unknown("/engines/public%2Fmodel/unknown", 404)]
#[case::unescaped_model("/engines/public/model/chat/completions", 404)]
#[case::extra_segment("/engines/public%2Fmodel/extra/chat/completions", 404)]
#[tokio::test]
async fn deployment_path_errors_take_precedence_over_invalid_json(
#[case] path: &str,
#[case] status: u16,
) {
let response = support::app("anthropic/test-model", "http://127.0.0.1:1")
.oneshot(Request::post(path).body(Body::from("{")).unwrap())
.await
.unwrap();
assert_eq!(response.status(), status);
}

View file

@ -14,7 +14,7 @@ use litellm_secrets::{SecretValue, source::SecretSource};
use serde_json::Value;
use tower::ServiceExt;
struct NoSecrets;
pub struct NoSecrets;
impl SecretSource for NoSecrets {
fn get_secret_str<'a>(