mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
366ec6f487
commit
698072308b
10 changed files with 361 additions and 1 deletions
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -621,6 +621,7 @@ dependencies = [
|
|||
"subtle",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
"tower",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
|
|
@ -41,3 +41,4 @@ python-config = ["dep:pyo3"]
|
|||
|
||||
[dev-dependencies]
|
||||
futures-channel = "0.3"
|
||||
tower = "0.5"
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
110
litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs
Normal file
110
litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
191
litellm-rust/crates/ai-gateway/src/routes/messages/service.rs
Normal file
191
litellm-rust/crates/ai-gateway/src/routes/messages/service.rs
Normal 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"));
|
||||
}
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue