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>
This commit is contained in:
Devin AI 2026-07-18 23:51:30 +00:00
parent 366ec6f487
commit 698072308b
10 changed files with 361 additions and 1 deletions

View file

@ -621,6 +621,7 @@ dependencies = [
"subtle",
"tokio",
"tokio-tungstenite",
"tower",
]
[[package]]

View file

@ -41,3 +41,4 @@ python-config = ["dep:pyo3"]
[dev-dependencies]
futures-channel = "0.3"
tower = "0.5"

View file

@ -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";

View file

@ -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();

View file

@ -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<reqwest::Response> {
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),
})
}

View file

@ -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<Value> {
@ -17,5 +19,11 @@ pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<Value> {
execute_messages_provider_call(prepared).await
}
#[cfg(feature = "server")]
pub(crate) async fn stream_messages(request: MessagesRequest<'_>) -> CoreResult<reqwest::Response> {
let prepared = prepare_messages_call(request)?;
execute_messages_provider_stream(prepared).await
}
#[cfg(test)]
mod tests;

View file

@ -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<ProviderMessagesRequest> {
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)

View file

@ -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<AppState> {
Router::new().route("/v1/messages", post(handle))
}
async fn handle(
_auth: RequireMasterKey,
State(state): State<AppState>,
Json(body): Json<Value>,
) -> Result<Response<Body>, (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);
}
}

View file

@ -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<MessagesResponse, CoreError> {
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::<usize>().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"));
}
}

View file

@ -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)
}