From 698072308bcd6b56b9455e49e0a69c3680482c12 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 18 Jul 2026 23:51:30 +0000 Subject: [PATCH] feat(litellm-rust): serve /v1/messages natively on the ai-gateway axum server Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm-rust/Cargo.lock | 1 + litellm-rust/crates/ai-gateway/Cargo.toml | 1 + .../crates/ai-gateway/src/constants.rs | 5 + litellm-rust/crates/ai-gateway/src/main.rs | 11 + .../crates/ai-gateway/src/messages/handler.rs | 31 +++ .../crates/ai-gateway/src/messages/mod.rs | 8 + .../crates/ai-gateway/src/messages/prepare.rs | 2 +- .../ai-gateway/src/routes/messages/mod.rs | 110 ++++++++++ .../ai-gateway/src/routes/messages/service.rs | 191 ++++++++++++++++++ .../crates/ai-gateway/src/routes/mod.rs | 2 + 10 files changed, 361 insertions(+), 1 deletion(-) create mode 100644 litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs create mode 100644 litellm-rust/crates/ai-gateway/src/routes/messages/service.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index f563c18ea14..7daff1dfc91 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -621,6 +621,7 @@ dependencies = [ "subtle", "tokio", "tokio-tungstenite", + "tower", ] [[package]] diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index 4055be36785..24b5846312d 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -41,3 +41,4 @@ python-config = ["dep:pyo3"] [dev-dependencies] futures-channel = "0.3" +tower = "0.5" diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 557fe5d53d4..cd4c1626519 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -40,3 +40,8 @@ pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10; /// Max characters of an upstream error body echoed across the host boundary /// before truncation, so provider bodies are bounded and data-minimized. pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256; + +#[cfg(feature = "server")] +pub(crate) const MESSAGES_STREAM_ACCEPT_HEADER: &str = "accept"; +#[cfg(feature = "server")] +pub(crate) const MESSAGES_STREAM_ACCEPT_VALUE: &str = "text/event-stream"; diff --git a/litellm-rust/crates/ai-gateway/src/main.rs b/litellm-rust/crates/ai-gateway/src/main.rs index f9ce97801d3..9c715d9ae41 100644 --- a/litellm-rust/crates/ai-gateway/src/main.rs +++ b/litellm-rust/crates/ai-gateway/src/main.rs @@ -142,6 +142,17 @@ fn build_router() -> Router { /// A real deployment loads `model_list` from config; this is the minimal stand-in /// so the gateway has one OpenAI deployment to route to. fn build_router_from_env() -> Router { + if let Ok(model_name) = std::env::var("LITELLM_MODEL_NAME") { + let model = std::env::var("LITELLM_MODEL").unwrap_or_else(|_| model_name.clone()); + return Router::new(vec![Deployment { + model_name, + litellm_params: LiteLLMParams { + model, + api_key: std::env::var("LITELLM_API_KEY").ok(), + api_base: std::env::var("LITELLM_API_BASE").ok(), + }, + }]); + } let model = std::env::var("OPENAI_REALTIME_MODEL").unwrap_or_else(|_| "gpt-realtime".to_string()); let api_key = std::env::var("OPENAI_API_KEY").ok(); diff --git a/litellm-rust/crates/ai-gateway/src/messages/handler.rs b/litellm-rust/crates/ai-gateway/src/messages/handler.rs index dd4a2f22aa7..eb372cd1e94 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/handler.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/handler.rs @@ -45,3 +45,34 @@ pub(super) async fn execute_messages_provider_call( CoreError::InvalidResponse(format!("failed to serialize messages response: {err}")) }) } + +#[cfg(feature = "server")] +pub(super) async fn execute_messages_provider_stream( + request: ProviderMessagesRequest, +) -> CoreResult { + let mut request_builder = http_client().post(&request.url).json(&request.body); + for (key, value) in &request.upstream_headers { + request_builder = request_builder.header(key, value); + } + if let Some(duration) = request.timeout { + request_builder = request_builder.timeout(duration); + } + + let response = request_builder + .send() + .await + .map_err(|err| CoreError::Network(err.to_string()))?; + if response.status().is_success() { + return Ok(response); + } + + let status = response.status(); + let body = response + .text() + .await + .map_err(|err| CoreError::Network(err.to_string()))?; + Err(CoreError::Http { + status: status.as_u16(), + body: truncate_error_body(&body), + }) +} diff --git a/litellm-rust/crates/ai-gateway/src/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/messages/mod.rs index 7ed81474c47..feb36b2166e 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/mod.rs @@ -10,6 +10,8 @@ mod types; pub use types::MessagesRequest; use handler::execute_messages_provider_call; +#[cfg(feature = "server")] +use handler::execute_messages_provider_stream; use prepare::prepare_messages_call; pub async fn messages(request: MessagesRequest<'_>) -> CoreResult { @@ -17,5 +19,11 @@ pub async fn messages(request: MessagesRequest<'_>) -> CoreResult { execute_messages_provider_call(prepared).await } +#[cfg(feature = "server")] +pub(crate) async fn stream_messages(request: MessagesRequest<'_>) -> CoreResult { + let prepared = prepare_messages_call(request)?; + execute_messages_provider_stream(prepared).await +} + #[cfg(test)] mod tests; diff --git a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs index 47105b39954..c27d368670e 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs @@ -6,7 +6,7 @@ use litellm_core::CoreResult; use super::common_utils::{has_header, messages_provider_config, string_headers}; use super::types::{MessagesRequest, ProviderMessagesRequest}; -pub(super) fn prepare_messages_call( +pub(crate) fn prepare_messages_call( request: MessagesRequest<'_>, ) -> CoreResult { let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider) diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs new file mode 100644 index 00000000000..eead05318e7 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -0,0 +1,110 @@ +mod service; + +use axum::body::Body; +use axum::extract::{Json, State}; +use axum::http::header::CONTENT_TYPE; +use axum::http::{Response, StatusCode}; +use axum::response::IntoResponse; +use axum::routing::post; +use axum::Router; +use serde_json::Value; + +use crate::auth::RequireMasterKey; +use crate::state::AppState; + +pub fn router() -> Router { + Router::new().route("/v1/messages", post(handle)) +} + +async fn handle( + _auth: RequireMasterKey, + State(state): State, + Json(body): Json, +) -> Result, (StatusCode, String)> { + match service::run(&state.router, body).await.map_err(map_error)? { + service::MessagesResponse::Json(body) => { + Ok((StatusCode::OK, axum::Json(body)).into_response()) + } + service::MessagesResponse::Stream(response) => { + let content_type = response + .headers() + .get(CONTENT_TYPE) + .cloned() + .unwrap_or_else(|| axum::http::HeaderValue::from_static("text/event-stream")); + let mut result = Response::new(Body::from_stream(response.bytes_stream())); + result.headers_mut().insert(CONTENT_TYPE, content_type); + Ok(result) + } + } +} + +fn map_error(error: litellm_core::CoreError) -> (StatusCode, String) { + match error { + litellm_core::CoreError::Http { status, body } => ( + StatusCode::from_u16(status).unwrap_or(StatusCode::BAD_GATEWAY), + body, + ), + litellm_core::CoreError::InvalidRequest(message) + | litellm_core::CoreError::InvalidProvider(message) + | litellm_core::CoreError::Routing(message) => (StatusCode::BAD_REQUEST, message), + litellm_core::CoreError::Network(message) => (StatusCode::BAD_GATEWAY, message), + error => (StatusCode::BAD_REQUEST, error.to_string()), + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use axum::body::Body; + use axum::http::{Request, StatusCode}; + use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter}; + use serde_json::json; + use tower::ServiceExt; + + use super::router; + use crate::io::realtime_pool::RealtimePool; + use crate::state::AppState; + + fn app() -> axum::Router { + let state = AppState { + router: Arc::new(ModelRouter::new(vec![Deployment { + model_name: "rust-model".to_string(), + litellm_params: LiteLLMParams { + model: "azure_ai/claude-loadtest".to_string(), + api_key: Some("sk-upstream".to_string()), + api_base: Some("http://127.0.0.1:1".to_string()), + }, + }])), + master_key: Some(Arc::from("sk-1234")), + loggers: Arc::new(Vec::new()), + realtime_pool: RealtimePool::disabled(), + }; + router().with_state(state) + } + + #[tokio::test] + async fn requires_authentication() { + let request = Request::builder() + .method("POST") + .uri("/v1/messages") + .header("content-type", "application/json") + .body(Body::from(json!({"model": "rust-model"}).to_string())) + .expect("request builds"); + let response = app().oneshot(request).await.expect("response"); + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn rejects_unknown_model() { + let request = Request::builder() + .method("POST") + .uri("/v1/messages") + .header("authorization", "Bearer sk-1234") + .header("content-type", "application/json") + .body(Body::from(json!({"model": "missing"}).to_string())) + .expect("request builds"); + let response = app().oneshot(request).await.expect("response"); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs new file mode 100644 index 00000000000..54a2e4d64d7 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -0,0 +1,191 @@ +use litellm_core::error::CoreError; +use litellm_core::router::Router; +use serde_json::{Map, Value}; + +pub enum MessagesResponse { + Json(Value), + Stream(reqwest::Response), +} + +pub async fn run(router: &Router, mut body: Value) -> Result { + let model = body + .get("model") + .and_then(Value::as_str) + .filter(|model| !model.trim().is_empty()) + .ok_or_else(|| CoreError::InvalidRequest("missing 'model' in request body".to_string()))?; + let deployment = router.get_available_deployment(model).ok_or_else(|| { + CoreError::InvalidRequest(format!("no deployment registered for model '{model}'")) + })?; + let params = &deployment.litellm_params; + let provider_model = params.model.clone(); + let is_stream = body.get("stream").and_then(Value::as_bool).unwrap_or(false); + let object = body.as_object_mut().ok_or_else(|| { + CoreError::InvalidRequest("Anthropic messages request must be a JSON object".to_string()) + })?; + object.insert("model".to_string(), Value::String(provider_model.clone())); + let extra_headers = is_stream.then(|| { + Map::from_iter([( + crate::constants::MESSAGES_STREAM_ACCEPT_HEADER.to_string(), + Value::String(crate::constants::MESSAGES_STREAM_ACCEPT_VALUE.to_string()), + )]) + }); + + let request = crate::messages::MessagesRequest { + model: &provider_model, + body, + api_key: params.api_key.as_deref(), + api_base: params.api_base.as_deref(), + custom_llm_provider: None, + extra_headers, + timeout: None, + }; + if is_stream { + crate::messages::stream_messages(request) + .await + .map(MessagesResponse::Stream) + } else { + crate::messages::messages(request) + .await + .map(MessagesResponse::Json) + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use serde_json::json; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + use super::*; + use litellm_core::router::{Deployment, LiteLLMParams}; + + fn router(api_base: String) -> Router { + Router::new(vec![Deployment { + model_name: "rust-model".to_string(), + litellm_params: LiteLLMParams { + model: "azure_ai/claude-loadtest".to_string(), + api_key: Some("sk-upstream".to_string()), + api_base: Some(api_base), + }, + }]) + } + + async fn accept_request(listener: TcpListener, response: String) -> String { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let mut request = Vec::new(); + let mut buffer = [0_u8; 1024]; + loop { + let count = socket.read(&mut buffer).await.expect("reads request"); + if count == 0 { + break; + } + request.extend_from_slice(&buffer[..count]); + if request.windows(4).any(|window| window == b"\r\n\r\n") { + break; + } + } + let header_end = request + .windows(4) + .position(|window| window == b"\r\n\r\n") + .map(|position| position + 4) + .expect("request headers"); + let headers = String::from_utf8_lossy(&request[..header_end]); + let content_length = headers + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().ok()) + .flatten() + }) + .unwrap_or(0); + while request.len().saturating_sub(header_end) < content_length { + let count = socket.read(&mut buffer).await.expect("reads body"); + if count == 0 { + break; + } + request.extend_from_slice(&buffer[..count]); + } + socket + .write_all(response.as_bytes()) + .await + .expect("writes response"); + String::from_utf8(request).expect("request is utf8") + } + + #[tokio::test] + async fn runs_non_stream_request_against_selected_deployment() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let address = listener.local_addr().expect("address"); + let body = br#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-loadtest","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}"#; + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + body.len(), + String::from_utf8_lossy(body) + ); + let server = tokio::spawn(accept_request(listener, response)); + + let result = run( + &router(format!("http://{address}")), + json!({ + "model": "rust-model", + "max_tokens": 8, + "messages": [{"role": "user", "content": "ping"}] + }), + ) + .await + .expect("request succeeds"); + let MessagesResponse::Json(result) = result else { + panic!("expected JSON response"); + }; + assert_eq!(result["id"], "msg_1"); + let request = server.await.expect("server completes"); + assert!(request.contains("POST /anthropic/v1/messages ")); + assert!(request.contains("claude-loadtest"), "{request}"); + } + + #[tokio::test] + async fn passes_stream_bytes_through_unchanged() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let address = listener.local_addr().expect("address"); + let body = b"event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ncontent-length: {}\r\nconnection: close\r\n\r\n", + body.len() + ); + let response = Arc::new([response.as_bytes(), body].concat()); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let mut request = [0_u8; 4096]; + let count = socket.read(&mut request).await.expect("reads request"); + socket.write_all(&response).await.expect("writes response"); + String::from_utf8_lossy(&request[..count]).into_owned() + }); + + let result = run( + &router(format!("http://{address}")), + json!({ + "model": "rust-model", + "stream": true, + "max_tokens": 8, + "messages": [{"role": "user", "content": "ping"}] + }), + ) + .await + .expect("request succeeds"); + let MessagesResponse::Stream(response) = result else { + panic!("expected stream response"); + }; + assert_eq!( + response.bytes().await.expect("reads stream").as_ref(), + body.as_slice() + ); + let request = server.await.expect("server completes"); + assert!(request.contains("\"stream\":true")); + assert!(request + .to_ascii_lowercase() + .contains("accept: text/event-stream")); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/routes/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/mod.rs index c6b9573781a..8872e9c8131 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/mod.rs @@ -7,6 +7,7 @@ pub mod gil; pub mod health; +pub mod messages; pub mod realtime; use axum::Router; @@ -18,6 +19,7 @@ pub fn app(state: AppState) -> Router { Router::new() .merge(health::router()) .merge(gil::router()) + .merge(messages::router()) .merge(realtime::router()) .with_state(state) }