From 1ceeefbf84c291e6d544c5cc44641ea0671d7e8b Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 16:50:05 -0700 Subject: [PATCH] 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 Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/Cargo.lock | 1 + .../crates/core/src/chat_completions/mod.rs | 1 + .../crates/core/src/chat_completions/route.rs | 44 +++ .../crates/core/src/chat_completions/types.rs | 26 ++ .../crates/core/tests/chat_completions.rs | 111 ++++++++ .../crates/gateway-inference/AGENTS.md | 3 +- .../crates/gateway-inference/Cargo.toml | 4 +- .../src/audio_transcription.rs | 43 ++- .../gateway-inference/src/chat_completions.rs | 97 +++---- .../crates/gateway-inference/src/error.rs | 15 ++ .../crates/gateway-inference/src/lib.rs | 18 +- .../crates/gateway-inference/src/messages.rs | 87 ++++++ .../gateway-inference/src/messages/mod.rs | 113 -------- .../crates/gateway-inference/src/ocr.rs | 44 +-- .../crates/gateway-inference/src/request.rs | 160 ++++++++--- .../gateway-inference/tests/messages.rs | 71 ++++- .../crates/gateway-inference/tests/ocr.rs | 70 +++++ .../crates/gateway-inference/tests/routes.rs | 252 ++++++++++++++++-- .../gateway-inference/tests/support/mod.rs | 2 +- 19 files changed, 870 insertions(+), 292 deletions(-) create mode 100644 litellm-rust/crates/core/src/chat_completions/route.rs create mode 100644 litellm-rust/crates/gateway-inference/src/messages.rs delete mode 100644 litellm-rust/crates/gateway-inference/src/messages/mod.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index fcf9df25cc8..f3682dacc61 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3262,6 +3262,7 @@ dependencies = [ "futures-util", "litellm-auth", "litellm-core", + "litellm-host-http", "litellm-http", "litellm-llms", "litellm-router", diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 62e9293c56e..c6b43ff9eea 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -1,3 +1,4 @@ +pub mod route; pub mod types; pub use crate::error::RouteError as Error; mod common_utils; diff --git a/litellm-rust/crates/core/src/chat_completions/route.rs b/litellm-rust/crates/core/src/chat_completions/route.rs new file mode 100644 index 00000000000..3f10a588692 --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/route.rs @@ -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 { + 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) + }, + ) + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/types.rs b/litellm-rust/crates/core/src/chat_completions/types.rs index ffd5e25d252..73d6378fc92 100644 --- a/litellm-rust/crates/core/src/chat_completions/types.rs +++ b/litellm-rust/crates/core/src/chat_completions/types.rs @@ -23,6 +23,32 @@ pub struct ChatCompletionsRequest<'a> { pub timeout: Option, } +pub struct ChatCompletionsCall { + pub model: String, + pub messages: Value, + pub optional_params: Map, + pub api_key: Option, + pub api_base: Option, + pub custom_llm_provider: Option, + pub extra_headers: Option>, + pub timeout: Option, +} + +impl From> 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, diff --git a/litellm-rust/crates/core/tests/chat_completions.rs b/litellm-rust/crates/core/tests/chat_completions.rs index bbd4b083a4c..1b0cba12ce1 100644 --- a/litellm-rust/crates/core/tests/chat_completions.rs +++ b/litellm-rust/crates/core/tests/chat_completions.rs @@ -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::::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 for FailingHook { + async fn before_provider_request( + &self, + wire: WireRequest, + _: RequestContext, + ) -> Result { + 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); +} diff --git a/litellm-rust/crates/gateway-inference/AGENTS.md b/litellm-rust/crates/gateway-inference/AGENTS.md index 7dc57380083..a03cd9cbbb6 100644 --- a/litellm-rust/crates/gateway-inference/AGENTS.md +++ b/litellm-rust/crates/gateway-inference/AGENTS.md @@ -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 diff --git a/litellm-rust/crates/gateway-inference/Cargo.toml b/litellm-rust/crates/gateway-inference/Cargo.toml index f5ee5e81ba6..8217c273dbf 100644 --- a/litellm-rust/crates/gateway-inference/Cargo.toml +++ b/litellm-rust/crates/gateway-inference/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/gateway-inference/src/audio_transcription.rs b/litellm-rust/crates/gateway-inference/src/audio_transcription.rs index 1f3d32dac80..94b4d6ce91f 100644 --- a/litellm-rust/crates/gateway-inference/src/audio_transcription.rs +++ b/litellm-rust/crates/gateway-inference/src/audio_transcription.rs @@ -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>, 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>, + body: InferenceBody, +) -> Result { + handle(&gateway, body).await.map(Json) } -async fn handle(gateway: &Gateway, request: Request) -> Result { - let (body, upload) = request::parse(request).await?; - let deployment = request::deployment(gateway, &body)?; +async fn handle( + gateway: &Gateway, + InferenceBody { + fields: body, + upload, + }: InferenceBody, +) -> Result { + 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 { .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 diff --git a/litellm-rust/crates/gateway-inference/src/chat_completions.rs b/litellm-rust/crates/gateway-inference/src/chat_completions.rs index 814a368ea24..16022e2f103 100644 --- a/litellm-rust/crates/gateway-inference/src/chat_completions.rs +++ b/litellm-rust/crates/gateway-inference/src/chat_completions.rs @@ -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>, body: Bytes) -> Response { - respond(&gateway, request::object(&body)).await -} - -pub(crate) async fn deployment( +pub(crate) async fn create( State(gateway): State>, - Path(path): Path, - 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 { + handle(&gateway, body).await } -async fn respond(gateway: &Gateway, body: Result, 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>, + Path(model): Path, + JsonObject(body): JsonObject, +) -> Result { + 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) -> Result { - 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) } diff --git a/litellm-rust/crates/gateway-inference/src/error.rs b/litellm-rust/crates/gateway-inference/src/error.rs index 1d40857d962..96edc3a94c5 100644 --- a/litellm-rust/crates/gateway-inference/src/error.rs +++ b/litellm-rust/crates/gateway-inference/src/error.rs @@ -28,6 +28,21 @@ pub enum Error { Internal(String), } +impl IntoResponse for Error { + fn into_response(self) -> Response { + self.openai_response() + } +} + +impl From> for Error { + fn from(error: litellm_host_http::Error) -> 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 { diff --git a/litellm-rust/crates/gateway-inference/src/lib.rs b/litellm-rust/crates/gateway-inference/src/lib.rs index cd6e53ba102..a0ef7563b8a 100644 --- a/litellm-rust/crates/gateway-inference/src/lib.rs +++ b/litellm-rust/crates/gateway-inference/src/lib.rs @@ -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) -> 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) -> Router { )) .with_state(gateway) } + +fn model_routes() -> Router> { + Router::new() + .route( + "/chat/completions", + post(chat_completions::create_from_model_path), + ) + .route("/embeddings", post(request::unsupported)) + .route("/completions", post(request::unsupported)) +} diff --git a/litellm-rust/crates/gateway-inference/src/messages.rs b/litellm-rust/crates/gateway-inference/src/messages.rs new file mode 100644 index 00000000000..0338640e741 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/messages.rs @@ -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>, + RequestId(request_id): RequestId, + headers: HeaderMap, + body: Result, +) -> 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, +) -> Result { + let deployment = request::resolve_deployment(gateway, &body)?; + let call = project(deployment, body, headers)?; + let machine = gateway.messages.clone().machine(call); + let stream = + Sse::::new(Json, |error| Bytes::from(Error::from(error).sse_frame())); + Ok(litellm_host_http::serve(machine, (), (), stream).await?) +} + +fn project( + deployment: &Deployment, + body: Map, + headers: &HeaderMap, +) -> Result { + 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 { + let extra_headers: Map = 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, + }) + }) +} diff --git a/litellm-rust/crates/gateway-inference/src/messages/mod.rs b/litellm-rust/crates/gateway-inference/src/messages/mod.rs deleted file mode 100644 index 837757aefa8..00000000000 --- a/litellm-rust/crates/gateway-inference/src/messages/mod.rs +++ /dev/null @@ -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>, - 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 { - 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, - headers: &HeaderMap, -) -> Result { - 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 { - let extra_headers: Map = 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>) -> 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() -} diff --git a/litellm-rust/crates/gateway-inference/src/ocr.rs b/litellm-rust/crates/gateway-inference/src/ocr.rs index eb8f0e7307e..787183ccb33 100644 --- a/litellm-rust/crates/gateway-inference/src/ocr.rs +++ b/litellm-rust/crates/gateway-inference/src/ocr.rs @@ -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>, 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>, + headers: HeaderMap, + body: InferenceBody, +) -> Result { + handle(&gateway, &headers, body).await.map(Json) } -async fn handle(gateway: &Gateway, request: Request) -> Result { - let header_format = request - .headers() +async fn handle( + gateway: &Gateway, + headers: &HeaderMap, + InferenceBody { + fields: body, + upload, + }: InferenceBody, +) -> Result { + 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") diff --git a/litellm-rust/crates/gateway-inference/src/request.rs b/litellm-rust/crates/gateway-inference/src/request.rs index f58c7b3ed79..dfedef93a76 100644 --- a/litellm-rust/crates/gateway-inference/src/request.rs +++ b/litellm-rust/crates/gateway-inference/src/request.rs @@ -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, } -pub(crate) fn object(body: &[u8]) -> Result, Error> { +pub struct JsonObject(pub Map); + +pub struct RequestId(pub Option); + +impl FromRequestParts for RequestId { + type Rejection = std::convert::Infallible; + + async fn from_request_parts(parts: &mut Parts, _: &S) -> Result { + Ok(Self( + parts + .headers + .get("x-request-id") + .and_then(|value| value.to_str().ok()) + .map(str::to_owned), + )) + } +} + +impl FromRequest for JsonObject { + type Rejection = Error; + + async fn from_request(request: Request, state: &S) -> Result { + 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, + pub upload: Option, +} + +impl FromRequest for InferenceBody { + type Rejection = Error; + + async fn from_request(request: Request, state: &S) -> Result { + 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, 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, Error> { } } -pub(crate) fn deployment<'a>( +pub(crate) fn resolve_deployment<'a>( gateway: &'a Gateway, body: &Map, ) -> 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, Option), 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 { 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, 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, 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); + } } diff --git a/litellm-rust/crates/gateway-inference/tests/messages.rs b/litellm-rust/crates/gateway-inference/tests/messages.rs index 30836498d49..8d2578a954b 100644 --- a/litellm-rust/crates/gateway-inference/tests/messages.rs +++ b/litellm-rust/crates/gateway-inference/tests/messages.rs @@ -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"); +} diff --git a/litellm-rust/crates/gateway-inference/tests/ocr.rs b/litellm-rust/crates/gateway-inference/tests/ocr.rs index 3d5a2d22b09..dba8099b685 100644 --- a/litellm-rust/crates/gateway-inference/tests/ocr.rs +++ b/litellm-rust/crates/gateway-inference/tests/ocr.rs @@ -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()); +} diff --git a/litellm-rust/crates/gateway-inference/tests/routes.rs b/litellm-rust/crates/gateway-inference/tests/routes.rs index 5d3b06ccadb..48f73358ce4 100644 --- a/litellm-rust/crates/gateway-inference/tests/routes.rs +++ b/litellm-rust/crates/gateway-inference/tests/routes.rs @@ -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, ) { 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::(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); +} diff --git a/litellm-rust/crates/gateway-inference/tests/support/mod.rs b/litellm-rust/crates/gateway-inference/tests/support/mod.rs index 5fa0b044e3c..7c35f73b82b 100644 --- a/litellm-rust/crates/gateway-inference/tests/support/mod.rs +++ b/litellm-rust/crates/gateway-inference/tests/support/mod.rs @@ -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>(