mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
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:
parent
438bffc26e
commit
1ceeefbf84
19 changed files with 870 additions and 292 deletions
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -3262,6 +3262,7 @@ dependencies = [
|
|||
"futures-util",
|
||||
"litellm-auth",
|
||||
"litellm-core",
|
||||
"litellm-host-http",
|
||||
"litellm-http",
|
||||
"litellm-llms",
|
||||
"litellm-router",
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
pub mod route;
|
||||
pub mod types;
|
||||
pub use crate::error::RouteError as Error;
|
||||
mod common_utils;
|
||||
|
|
|
|||
44
litellm-rust/crates/core/src/chat_completions/route.rs
Normal file
44
litellm-rust/crates/core/src/chat_completions/route.rs
Normal 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)
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
|
|
|||
87
litellm-rust/crates/gateway-inference/src/messages.rs
Normal file
87
litellm-rust/crates/gateway-inference/src/messages.rs
Normal 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,
|
||||
})
|
||||
})
|
||||
}
|
||||
|
|
@ -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()
|
||||
}
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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>(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue