From 95332584fdb0162f39aea32f400ca8b7d18d18c2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 19 Jul 2026 01:41:37 +0000 Subject: [PATCH 1/2] feat(rust): expose Anthropic Messages route (POST /v1/messages) on the axum gateway (#33880) * feat(rust): expose anthropic messages route Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * fix(rust): use provider model for messages upstream Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * feat(rust): stream Anthropic Messages SSE on POST /v1/messages Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * test(rust): prove alias is substituted with provider model on /v1/messages Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * fix(rust): make anthropic messages provider constant available without server feature Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> 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 | 13 + .../ai-gateway/src/messages/common_utils.rs | 2 +- .../crates/ai-gateway/src/messages/handler.rs | 36 ++ .../crates/ai-gateway/src/messages/mod.rs | 32 +- .../crates/ai-gateway/src/messages/prepare.rs | 1 + .../crates/ai-gateway/src/messages/tests.rs | 2 +- .../crates/ai-gateway/src/messages/types.rs | 1 + .../ai-gateway/src/routes/messages/mod.rs | 513 ++++++++++++++++++ .../ai-gateway/src/routes/messages/service.rs | 64 +++ .../crates/ai-gateway/src/routes/mod.rs | 2 + 12 files changed, 664 insertions(+), 4 deletions(-) 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..c15af4cc478 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 = { version = "0.5.3", features = ["util"] } diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 557fe5d53d4..ac038b790ae 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -40,3 +40,16 @@ 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; + +/// HTTP path for the non-streaming Anthropic Messages route. +#[cfg(feature = "server")] +pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages"; + +/// Provider name used by the Anthropic Messages route when a deployment's +/// provider model does not carry an explicit provider prefix. +pub(crate) const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; + +/// Request headers owned by the gateway and never forwarded upstream. +#[cfg(feature = "server")] +pub(crate) const MESSAGES_HEADERS_NOT_FORWARDED: &[&str] = + &["authorization", "connection", "content-length", "host"]; diff --git a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs b/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs index bc4b13bfe3d..33894d0ee64 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs @@ -19,8 +19,8 @@ pub(super) fn messages_provider_config( provider: &str, ) -> Option<&'static dyn AnthropicMessagesProviderConfig> { match provider { - "azure_ai" => Some(&AZURE_ANTHROPIC_MESSAGES_CONFIG), "anthropic" => Some(&ANTHROPIC_MESSAGES_CONFIG), + "azure_ai" => Some(&AZURE_ANTHROPIC_MESSAGES_CONFIG), _ => None, } } diff --git a/litellm-rust/crates/ai-gateway/src/messages/handler.rs b/litellm-rust/crates/ai-gateway/src/messages/handler.rs index dd4a2f22aa7..d3b9d3b3fba 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/handler.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/handler.rs @@ -5,6 +5,7 @@ use serde_json::Value; use super::client::http_client; use super::common_utils::truncate_error_body; use super::types::ProviderMessagesRequest; +use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; pub(super) async fn execute_messages_provider_call( request: ProviderMessagesRequest, @@ -45,3 +46,38 @@ pub(super) async fn execute_messages_provider_call( CoreError::InvalidResponse(format!("failed to serialize messages response: {err}")) }) } + +pub(super) async fn execute_messages_provider_stream( + request: ProviderMessagesRequest, +) -> CoreResult { + if request.provider != ANTHROPIC_MESSAGES_PROVIDER { + return Err(CoreError::InvalidRequest( + "streaming messages is not supported for this provider".to_string(), + )); + } + + 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()))?; + let status = response.status(); + if !status.is_success() { + let text = response + .text() + .await + .map_err(|err| CoreError::Network(err.to_string()))?; + return Err(CoreError::Http { + status: status.as_u16(), + body: truncate_error_body(&text), + }); + } + Ok(response) +} diff --git a/litellm-rust/crates/ai-gateway/src/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/messages/mod.rs index 7ed81474c47..fd2dd546941 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/mod.rs @@ -9,12 +9,40 @@ mod types; pub use types::MessagesRequest; -use handler::execute_messages_provider_call; +use handler::{execute_messages_provider_call, execute_messages_provider_stream}; use prepare::prepare_messages_call; pub async fn messages(request: MessagesRequest<'_>) -> CoreResult { + match execute_messages(request, false).await? { + MessagesResponse::Json(body) => Ok(body), + MessagesResponse::Stream(response) => { + drop(response); + Err(litellm_core::CoreError::InvalidResponse( + "non-streaming messages execution returned a stream".to_string(), + )) + } + } +} + +pub(crate) enum MessagesResponse { + Json(Value), + Stream(reqwest::Response), +} + +pub(crate) async fn execute_messages( + request: MessagesRequest<'_>, + stream: bool, +) -> CoreResult { let prepared = prepare_messages_call(request)?; - execute_messages_provider_call(prepared).await + if stream { + execute_messages_provider_stream(prepared) + .await + .map(MessagesResponse::Stream) + } else { + execute_messages_provider_call(prepared) + .await + .map(MessagesResponse::Json) + } } #[cfg(test)] diff --git a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs index 47105b39954..6176f9cb67f 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs @@ -62,6 +62,7 @@ pub(super) fn prepare_messages_call( })?; Ok(ProviderMessagesRequest { + provider: provider.to_string(), model, config, url, diff --git a/litellm-rust/crates/ai-gateway/src/messages/tests.rs b/litellm-rust/crates/ai-gateway/src/messages/tests.rs index 1175d800e1c..9b1cc45aacb 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/tests.rs @@ -53,8 +53,8 @@ fn write_response(body: &str) -> String { #[test] fn provider_config_resolves_anthropic_and_azure_ai() { - assert!(messages_provider_config("azure_ai").is_some()); assert!(messages_provider_config("anthropic").is_some()); + assert!(messages_provider_config("azure_ai").is_some()); assert!(messages_provider_config("openai").is_none()); } diff --git a/litellm-rust/crates/ai-gateway/src/messages/types.rs b/litellm-rust/crates/ai-gateway/src/messages/types.rs index 6840ff57cc4..848fadb4b02 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/types.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/types.rs @@ -14,6 +14,7 @@ pub struct MessagesRequest<'a> { } pub(crate) struct ProviderMessagesRequest { + pub(crate) provider: String, pub(crate) model: String, pub(crate) config: &'static dyn AnthropicMessagesProviderConfig, pub(crate) url: String, 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..933386282fa --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -0,0 +1,513 @@ +//! `POST /v1/messages`, the Anthropic Messages HTTP surface. + +mod service; + +use axum::body::Body; +use axum::extract::{Json, State}; +use axum::http::header::{HeaderMap, HeaderValue, CACHE_CONTROL, CONTENT_TYPE}; +use axum::http::StatusCode; +use axum::response::{IntoResponse, Response}; +use axum::routing::post; +use axum::Router; +use litellm_core::CoreError; +use serde_json::{Map, Value}; + +use crate::auth::RequireMasterKey; +use crate::constants::{MESSAGES_HEADERS_NOT_FORWARDED, MESSAGES_ROUTE_PATH}; +use crate::state::AppState; + +/// This route's contribution to the app router. +pub fn router() -> Router { + Router::new().route(MESSAGES_ROUTE_PATH, post(handle)) +} + +async fn handle( + _auth: RequireMasterKey, + State(state): State, + headers: HeaderMap, + Json(body): Json, +) -> Result { + let extra_headers = forwarded_headers(&headers)?; + match service::run(&state.router, body, extra_headers) + .await + .map_err(MessagesRouteError::from)? + { + service::MessagesResponse::Json(body) => Ok(Json(body).into_response()), + service::MessagesResponse::Stream(upstream) => stream_response(upstream), + } +} + +fn stream_response(upstream: reqwest::Response) -> Result { + let content_type = upstream + .headers() + .get(CONTENT_TYPE) + .cloned() + .unwrap_or_else(|| HeaderValue::from_static("text/event-stream")); + let mut response = Response::builder() + .status( + StatusCode::from_u16(upstream.status().as_u16()).map_err(|error| { + MessagesRouteError(CoreError::InvalidResponse(format!( + "invalid upstream response status: {error}" + ))) + })?, + ) + .header(CONTENT_TYPE, content_type); + if let Some(value) = upstream.headers().get(CACHE_CONTROL) { + response = response.header(CACHE_CONTROL, value); + } + response + .body(Body::from_stream(upstream.bytes_stream())) + .map_err(|error| { + MessagesRouteError(CoreError::InvalidResponse(format!( + "failed to build streaming response: {error}" + ))) + }) +} + +fn forwarded_headers(headers: &HeaderMap) -> Result>, CoreError> { + let forwarded = headers + .iter() + .filter(|(name, _)| { + !MESSAGES_HEADERS_NOT_FORWARDED + .iter() + .any(|excluded| name.as_str().eq_ignore_ascii_case(excluded)) + }) + .map(|(name, value)| { + let value = value.to_str().map_err(|_| { + CoreError::InvalidRequest(format!("invalid value for header {}", name.as_str())) + })?; + Ok((name.to_string(), Value::String(value.to_string()))) + }) + .collect::, CoreError>>()?; + Ok((!forwarded.is_empty()).then_some(forwarded)) +} + +#[derive(Debug)] +struct MessagesRouteError(CoreError); + +impl From for MessagesRouteError { + fn from(error: CoreError) -> Self { + Self(error) + } +} + +impl IntoResponse for MessagesRouteError { + fn into_response(self) -> Response { + let (status, message) = match self.0 { + CoreError::InvalidRequest(message) => (StatusCode::BAD_REQUEST, message), + CoreError::InvalidProvider(_) | CoreError::Routing(_) => ( + StatusCode::NOT_FOUND, + "no messages deployment is configured for this model".to_string(), + ), + CoreError::Auth(_) => ( + StatusCode::BAD_GATEWAY, + "messages provider authentication failed".to_string(), + ), + CoreError::Http { .. } + | CoreError::Network(_) + | CoreError::InvalidResponse(_) + | CoreError::InvalidType { .. } + | CoreError::MissingField(_) => ( + StatusCode::BAD_GATEWAY, + "messages provider request failed".to_string(), + ), + }; + ( + status, + Json(serde_json::json!({"error": {"message": message}})), + ) + .into_response() + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use axum::body::Body; + use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE}; + use axum::http::Request; + use axum::http::StatusCode; + use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter}; + use serde_json::json; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + use tower::ServiceExt; + + use super::super::app; + use crate::io::realtime_pool::RealtimePool; + use crate::state::AppState; + + fn state(model: &str, api_base: String, master_key: Option<&str>) -> AppState { + state_with_provider(model, model, api_base, master_key) + } + + fn state_with_provider( + model_alias: &str, + provider_model: &str, + api_base: String, + master_key: Option<&str>, + ) -> AppState { + AppState { + router: Arc::new(ModelRouter::new(vec![Deployment { + model_name: model_alias.to_string(), + litellm_params: LiteLLMParams { + model: format!("anthropic/{provider_model}"), + api_key: Some("upstream-key".to_string()), + api_base: Some(api_base), + }, + }])), + master_key: master_key.map(Arc::from), + loggers: Arc::new(Vec::new()), + realtime_pool: RealtimePool::disabled(), + } + } + + async fn upstream(listener: TcpListener) -> (String, tokio::task::JoinHandle) { + let address = listener.local_addr().expect("listener has address"); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let mut request = Vec::new(); + let mut buffer = [0_u8; 4096]; + loop { + let read = socket.read(&mut buffer).await.expect("reads request"); + request.extend_from_slice(&buffer[..read]); + if request.windows(4).any(|window| window == b"\r\n\r\n") { + break; + } + } + let request = String::from_utf8(request).expect("request is utf8"); + let content_length = request + .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); + let header_end = request.find("\r\n\r\n").expect("request has headers") + 4; + let mut full_request = request.into_bytes(); + while full_request.len().saturating_sub(header_end) < content_length { + let read = socket.read(&mut buffer).await.expect("reads body"); + full_request.extend_from_slice(&buffer[..read]); + } + let request = String::from_utf8(full_request).expect("request is utf8"); + let body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-test"}"#; + 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(), + body + ); + socket + .write_all(response.as_bytes()) + .await + .expect("writes response"); + request + }); + (format!("http://{address}"), server) + } + + async fn streaming_upstream( + listener: TcpListener, + status: u16, + content_type: &'static str, + body: &'static str, + ) -> (String, tokio::task::JoinHandle) { + let address = listener.local_addr().expect("listener has address"); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let mut request = Vec::new(); + let mut buffer = [0_u8; 4096]; + loop { + let read = socket.read(&mut buffer).await.expect("reads request"); + request.extend_from_slice(&buffer[..read]); + if request.windows(4).any(|window| window == b"\r\n\r\n") { + break; + } + } + let request_text = String::from_utf8(request).expect("request is utf8"); + let content_length = request_text + .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); + let header_end = request_text.find("\r\n\r\n").expect("request has headers") + 4; + let mut full_request = request_text.into_bytes(); + while full_request.len().saturating_sub(header_end) < content_length { + let read = socket.read(&mut buffer).await.expect("reads body"); + full_request.extend_from_slice(&buffer[..read]); + } + let response = format!( + "HTTP/1.1 {status} OK\r\ncontent-type: {content_type}\r\ncache-control: no-cache\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}", + body.len() + ); + socket + .write_all(response.as_bytes()) + .await + .expect("writes response"); + String::from_utf8(full_request).expect("request is utf8") + }); + (format!("http://{address}"), server) + } + + #[tokio::test] + async fn route_constructs_anthropic_upstream_request() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let (api_base, server) = upstream(listener).await; + let app = app(state("claude-test", api_base, Some("master-key"))); + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/messages") + .header("authorization", "Bearer master-key") + .header("x-api-key", "request-upstream-key") + .header("anthropic-beta", "beta-feature") + .header("content-type", "application/json") + .body(Body::from( + json!({ + "model": "claude-test", + "max_tokens": 16, + "messages": [{"role": "user", "content": "hello"}] + }) + .to_string(), + )) + .expect("request builds"), + ) + .await + .expect("route responds"); + assert_eq!(response.status(), StatusCode::OK); + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .expect("response body reads"); + assert_eq!( + serde_json::from_slice::(&body).expect("json")["id"], + "msg_1" + ); + let upstream_request = server.await.expect("upstream task completes"); + let (head, body) = upstream_request + .split_once("\r\n\r\n") + .expect("upstream request has body"); + let head = head.to_ascii_lowercase(); + assert!(head.contains("x-api-key: request-upstream-key")); + assert!(head.contains("anthropic-beta: beta-feature")); + assert!(!head.contains("authorization: bearer master-key")); + let body: serde_json::Value = serde_json::from_str(body).expect("upstream body is json"); + assert_eq!(body["model"], "claude-test"); + assert_eq!(body["messages"][0]["content"], "hello"); + } + + #[tokio::test] + async fn route_substitutes_model_alias_with_provider_model_upstream() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let (api_base, server) = upstream(listener).await; + let app = app(state_with_provider( + "production", + "claude-sonnet-4-5", + api_base, + Some("master-key"), + )); + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/messages") + .header("authorization", "Bearer master-key") + .header("content-type", "application/json") + .body(Body::from( + json!({ + "model": "production", + "max_tokens": 16, + "messages": [{"role": "user", "content": "hello"}] + }) + .to_string(), + )) + .expect("request builds"), + ) + .await + .expect("route responds"); + assert_eq!(response.status(), StatusCode::OK); + let upstream_request = server.await.expect("upstream task completes"); + let (_, upstream_body) = upstream_request + .split_once("\r\n\r\n") + .expect("upstream request has body"); + let upstream_body: serde_json::Value = + serde_json::from_str(upstream_body).expect("upstream body is json"); + assert_eq!(upstream_body["model"], "claude-sonnet-4-5"); + assert_ne!(upstream_body["model"], "production"); + } + + #[tokio::test] + async fn route_streams_anthropic_events_without_buffering_or_reordering() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let events = "event: message_start\ndata: {\"type\":\"message_start\"}\n\nevent: content_block_delta\ndata: {\"type\":\"content_block_delta\"}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + let (api_base, server) = + streaming_upstream(listener, 200, "text/event-stream", events).await; + let app = app(state("claude-test", api_base, Some("master-key"))); + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/messages") + .header("authorization", "Bearer master-key") + .header("content-type", "application/json") + .body(Body::from( + json!({ + "model": "claude-test", + "max_tokens": 16, + "stream": true, + "messages": [{"role": "user", "content": "hello"}] + }) + .to_string(), + )) + .expect("request builds"), + ) + .await + .expect("route responds"); + assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response + .headers() + .get(CONTENT_TYPE) + .unwrap() + .to_str() + .unwrap(), + "text/event-stream" + ); + assert_eq!( + response + .headers() + .get(CACHE_CONTROL) + .unwrap() + .to_str() + .unwrap(), + "no-cache" + ); + let response_body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .expect("response body reads"); + assert_eq!(response_body, events.as_bytes()); + let upstream_request = server.await.expect("upstream task completes"); + let (_, upstream_body) = upstream_request + .split_once("\r\n\r\n") + .expect("upstream request has body"); + assert_eq!( + serde_json::from_str::(upstream_body) + .expect("upstream body is json")["stream"], + true + ); + } + + #[tokio::test] + async fn route_maps_streaming_upstream_errors_before_starting_response() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let (api_base, server) = streaming_upstream( + listener, + 429, + "application/json", + r#"{"error":"rate limited"}"#, + ) + .await; + let app = app(state("claude-test", api_base, Some("master-key"))); + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/messages") + .header("authorization", "Bearer master-key") + .header("content-type", "application/json") + .body(Body::from( + json!({ + "model": "claude-test", + "max_tokens": 16, + "stream": true, + "messages": [{"role": "user", "content": "hello"}] + }) + .to_string(), + )) + .expect("request builds"), + ) + .await + .expect("route responds"); + assert_eq!(response.status(), StatusCode::BAD_GATEWAY); + let response_body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .expect("response body reads"); + assert_eq!( + serde_json::from_slice::(&response_body).expect("error is json") + ["error"]["message"], + "messages provider request failed" + ); + server.await.expect("upstream task completes"); + } + + #[tokio::test] + async fn route_rejects_missing_master_key() { + let app = app(state( + "claude-test", + "http://127.0.0.1:1".to_string(), + Some("master-key"), + )); + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/messages") + .header("content-type", "application/json") + .body(Body::from("{}")) + .expect("request builds"), + ) + .await + .expect("route responds"); + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn route_rejects_invalid_master_key() { + let app = app(state( + "claude-test", + "http://127.0.0.1:1".to_string(), + Some("master-key"), + )); + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/messages") + .header("authorization", "Bearer wrong-key") + .header("content-type", "application/json") + .body(Body::from("{}")) + .expect("request builds"), + ) + .await + .expect("route responds"); + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn route_rejects_malformed_json_without_panicking() { + let app = app(state( + "claude-test", + "http://127.0.0.1:1".to_string(), + Some("master-key"), + )); + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/messages") + .header("authorization", "Bearer master-key") + .header("content-type", "application/json") + .body(Body::from("{not-json")) + .expect("request builds"), + ) + .await + .expect("route responds"); + 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..7f00123ca39 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -0,0 +1,64 @@ +use std::sync::Arc; + +use litellm_core::router::Router; +use litellm_core::{CoreError, CoreResult}; +use serde_json::{Map, Value}; + +use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; +use crate::messages::{execute_messages, MessagesRequest}; + +pub(crate) enum MessagesResponse { + Json(Value), + Stream(reqwest::Response), +} + +pub async fn run( + router: &Arc, + body: Value, + extra_headers: Option>, +) -> CoreResult { + let model = body + .get("model") + .and_then(Value::as_str) + .map(str::trim) + .filter(|model| !model.is_empty()) + .ok_or_else(|| CoreError::InvalidRequest("messages body requires a model".to_string()))?; + let deployment = router.get_available_deployment(model).ok_or_else(|| { + CoreError::Routing(format!("no deployment available for model '{model}'")) + })?; + let provider_model = deployment.litellm_params.model.as_str(); + let upstream_model = provider_model + .split_once('/') + .map_or(provider_model, |(_, model)| model); + let custom_llm_provider = if provider_model.contains('/') { + None + } else { + Some(ANTHROPIC_MESSAGES_PROVIDER) + }; + let mut body = body; + body.as_object_mut() + .ok_or_else(|| CoreError::InvalidRequest("messages body must be an object".to_string()))? + .insert( + "model".to_string(), + Value::String(upstream_model.to_string()), + ); + + let request = MessagesRequest { + model: provider_model, + body, + api_key: deployment.litellm_params.api_key.as_deref(), + api_base: deployment.litellm_params.api_base.as_deref(), + custom_llm_provider, + extra_headers, + timeout: None, + }; + let stream = request.body.get("stream").and_then(Value::as_bool) == Some(true); + execute_messages(request, stream) + .await + .map(|response| match response { + crate::messages::MessagesResponse::Json(body) => MessagesResponse::Json(body), + crate::messages::MessagesResponse::Stream(upstream) => { + MessagesResponse::Stream(upstream) + } + }) +} 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) } From 7891388975d209d07c0dd1808c18de98a97d1541 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 19 Jul 2026 01:55:35 +0000 Subject: [PATCH 2/2] feat(rust): 1:1 port of OpenAI Responses API WebSockets to litellm-rust (#33849) * feat(rust): add OpenAI Responses WebSocket gateway Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * test(rust): cover Responses WebSocket gateway behavior Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * fix(rust): align Responses WebSocket parity Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * feat(rust): expose Responses WebSockets through bridge Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * fix(rust): reject non-openai responses deployments early Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * fix(rust): align Responses WebSocket bridge semantics Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * refactor(rust): move Responses instrumentation into core Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * fix(rust): preserve Responses WebSocket callback dispatch Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * build(deps): authorize vcrpy and locust licenses in liccheck Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm-rust/CLAUDE.md | 5 + .../crates/ai-gateway/src/constants.rs | 3 + litellm-rust/crates/ai-gateway/src/io/mod.rs | 1 + .../crates/ai-gateway/src/io/responses_ws.rs | 548 ++++++++++++++++++ litellm-rust/crates/ai-gateway/src/lib.rs | 3 - .../crates/ai-gateway/src/routes/mod.rs | 2 + .../ai-gateway/src/routes/responses/mod.rs | 348 +++++++++++ .../src/routes/responses/service.rs | 156 +++++ litellm-rust/crates/core/src/constants.rs | 3 + litellm-rust/crates/core/src/lib.rs | 2 + .../crates/core/src/providers/openai/mod.rs | 1 + .../src/providers/openai/responses/mod.rs | 1 + .../openai/responses/transformation.rs | 48 ++ .../core/src/responses/instrumentation.rs | 365 ++++++++++++ litellm-rust/crates/core/src/responses/mod.rs | 3 + .../crates/core/src/responses/types.rs | 166 ++++++ .../crates/core/src/responses/websocket.rs | 188 ++++++ litellm-rust/crates/python-bridge/src/lib.rs | 73 +++ litellm/llms/custom_httpx/llm_http_handler.py | 38 +- litellm/rust_bridge/ocr.py | 27 +- litellm/rust_bridge/responses_websocket.py | 95 +++ tests/code_coverage_tests/liccheck.ini | 2 + .../responses/test_rust_bridge_websocket.py | 86 +++ 23 files changed, 2145 insertions(+), 19 deletions(-) create mode 100644 litellm-rust/crates/ai-gateway/src/io/responses_ws.rs create mode 100644 litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs create mode 100644 litellm-rust/crates/ai-gateway/src/routes/responses/service.rs create mode 100644 litellm-rust/crates/core/src/constants.rs create mode 100644 litellm-rust/crates/core/src/providers/openai/responses/mod.rs create mode 100644 litellm-rust/crates/core/src/providers/openai/responses/transformation.rs create mode 100644 litellm-rust/crates/core/src/responses/instrumentation.rs create mode 100644 litellm-rust/crates/core/src/responses/mod.rs create mode 100644 litellm-rust/crates/core/src/responses/types.rs create mode 100644 litellm-rust/crates/core/src/responses/websocket.rs create mode 100644 litellm/rust_bridge/responses_websocket.py create mode 100644 tests/test_litellm/responses/test_rust_bridge_websocket.py diff --git a/litellm-rust/CLAUDE.md b/litellm-rust/CLAUDE.md index 4468a369e41..519b1d205ef 100644 --- a/litellm-rust/CLAUDE.md +++ b/litellm-rust/CLAUDE.md @@ -39,6 +39,11 @@ Route-level Rust structure mirrors LiteLLM's Python responsibilities: - Network execution lives in the host crate `ai-gateway` (`ai-gateway/src/io/`), never inside `core`. +Call-hook and lifecycle instrumentation, including phase timing, usage +accumulation, and callback payload construction, always lives in `core`. +Hosts feed observed events into core and dispatch the completed payloads through +their I/O logger; hosts must not own callback orchestration. + Allowed in `core`: - Pure request transforms - Pure response transforms diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index ac038b790ae..74808cf1ce6 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -41,6 +41,9 @@ pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10; /// before truncation, so provider bodies are bounded and data-minimized. pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256; +pub(crate) const DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10; +pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300; + /// HTTP path for the non-streaming Anthropic Messages route. #[cfg(feature = "server")] pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages"; diff --git a/litellm-rust/crates/ai-gateway/src/io/mod.rs b/litellm-rust/crates/ai-gateway/src/io/mod.rs index 7bc9642d192..9cbfa568121 100644 --- a/litellm-rust/crates/ai-gateway/src/io/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/io/mod.rs @@ -2,3 +2,4 @@ pub mod messages; pub mod ocr; pub mod realtime; pub mod realtime_pool; +pub mod responses_ws; diff --git a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs new file mode 100644 index 00000000000..ae6ad150bcf --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs @@ -0,0 +1,548 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use futures_util::stream::{SplitSink, SplitStream}; +use futures_util::{Sink, SinkExt, Stream, StreamExt}; +use litellm_core::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG; +use litellm_core::responses::types::ResponsesWsEvent; +use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig; +use litellm_core::{CoreError, CoreResult}; +use tokio::net::TcpStream; +use tokio::sync::Mutex; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::http::header::{HeaderName, AUTHORIZATION}; +use tokio_tungstenite::tungstenite::http::HeaderValue; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::{connect_async, MaybeTlsStream, WebSocketStream}; + +use crate::constants::{ + DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS, DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS, +}; + +const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY"; +const MISSING_KEY_MESSAGE: &str = + "Missing OpenAI API Key - a Responses WebSocket call is being made but no key was passed via params or the OPENAI_API_KEY environment variable"; + +pub type ResponsesUpstreamWs = WebSocketStream>; +type UpstreamTx = SplitSink; +type UpstreamRx = SplitStream; + +#[derive(Clone)] +pub struct ResponsesWebSocketConnection { + socket: Arc>>, +} + +impl ResponsesWebSocketConnection { + pub async fn connect_url( + url: &str, + headers: &HashMap, + timeout: Option, + ) -> CoreResult { + let mut request = url + .into_client_request() + .map_err(|error| CoreError::Network(error.to_string()))?; + for (name, value) in headers { + let header_name = name + .parse::() + .map_err(|error| CoreError::InvalidRequest(error.to_string()))?; + let header_value = HeaderValue::from_str(value) + .map_err(|error| CoreError::InvalidRequest(error.to_string()))?; + request.headers_mut().insert(header_name, header_value); + } + let connect = connect_async(request); + let result = match timeout { + Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| { + CoreError::Network("Responses WebSocket connection timed out".to_string()) + })?, + None => connect.await, + }; + let (socket, _) = result.map_err(|error| match error { + tokio_tungstenite::tungstenite::Error::Http(response) => CoreError::Http { + status: response.status().as_u16(), + body: String::new(), + }, + other => CoreError::Network(other.to_string()), + })?; + Ok(Self { + socket: Arc::new(Mutex::new(Some(socket))), + }) + } + + pub async fn send_text(&self, text: String) -> CoreResult<()> { + let mut socket = self.socket.lock().await; + let Some(socket) = socket.as_mut() else { + return Err(CoreError::Network( + "Responses WebSocket is closed".to_string(), + )); + }; + socket + .send(Message::Text(text)) + .await + .map_err(|error| CoreError::Network(error.to_string())) + } + + pub async fn recv_text(&self) -> CoreResult> { + let mut socket = self.socket.lock().await; + let Some(socket) = socket.as_mut() else { + return Ok(None); + }; + match socket.next().await { + Some(Ok(Message::Text(text))) => Ok(Some(text)), + Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec()) + .map(Some) + .map_err(|error| CoreError::InvalidResponse(error.to_string())), + Some(Ok(Message::Close(_))) | None => Ok(None), + Some(Ok(_)) => Ok(None), + Some(Err(error)) => Err(CoreError::Network(error.to_string())), + } + } + + pub async fn close(&self) -> CoreResult<()> { + let mut socket = self.socket.lock().await; + if let Some(socket) = socket.as_mut() { + socket + .close(None) + .await + .map_err(|error| CoreError::Network(error.to_string()))?; + } + *socket = None; + Ok(()) + } +} + +pub(crate) fn resolve_api_key(api_key: Option<&str>) -> CoreResult { + api_key + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .or_else(|| { + std::env::var(OPENAI_API_KEY_ENV) + .ok() + .filter(|value| !value.trim().is_empty()) + }) + .ok_or_else(|| CoreError::Auth(MISSING_KEY_MESSAGE.to_string())) +} + +async fn dial_upstream( + model: &str, + api_key: &str, + api_base: Option<&str>, +) -> CoreResult { + let url = OPENAI_RESPONSES_WS_CONFIG.complete_websocket_url(api_base, model); + let mut request = url + .as_str() + .into_client_request() + .map_err(|error| CoreError::Network(error.to_string()))?; + request.headers_mut().insert( + AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {api_key}")) + .map_err(|error| CoreError::Auth(error.to_string()))?, + ); + let result = tokio::time::timeout( + Duration::from_secs(DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS), + connect_async(request), + ) + .await + .map_err(|_| CoreError::Network("Responses WebSocket connection timed out".to_string()))?; + result + .map(|(socket, _)| socket) + .map_err(|error| match error { + tokio_tungstenite::tungstenite::Error::Http(response) => CoreError::Http { + status: response.status().as_u16(), + body: String::new(), + }, + other => CoreError::Network(other.to_string()), + }) +} + +pub struct ResponsesWebSocketStreaming; + +impl ResponsesWebSocketStreaming { + pub async fn bidirectional_forward( + model: &str, + upstream_tx: UpstreamTx, + upstream_rx: UpstreamRx, + idle_timeout: Option, + observe: impl FnMut(&ResponsesWsEvent) + Send, + client_in: In, + client_out: Out, + ) -> CoreResult<()> + where + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, + { + splice( + model, + upstream_tx, + upstream_rx, + idle_timeout, + observe, + client_in, + client_out, + ) + .await + } +} + +pub(crate) async fn splice( + model: &str, + mut upstream_tx: UpstreamTx, + mut upstream_rx: UpstreamRx, + idle_timeout: Option, + mut observe: impl FnMut(&ResponsesWsEvent) + Send, + mut client_in: In, + mut client_out: Out, +) -> CoreResult<()> +where + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, +{ + let idle = + idle_timeout.unwrap_or_else(|| Duration::from_secs(DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS)); + loop { + tokio::select! { + event = client_in.next() => { + let Some(event) = event else { break }; + for outbound in OPENAI_RESPONSES_WS_CONFIG + .transform_ws_request(&event, model)? + .events + { + let payload = serde_json::to_string(&outbound) + .map_err(|error| CoreError::InvalidResponse(error.to_string()))?; + upstream_tx.send(Message::Text(payload)) + .await + .map_err(|error| CoreError::Network(error.to_string()))?; + } + } + message = upstream_rx.next() => { + let Some(message) = message else { break }; + match message.map_err(|error| CoreError::Network(error.to_string()))? { + Message::Text(text) => { + let event = serde_json::from_str::(&text) + .map_err(|error| CoreError::InvalidResponse(error.to_string()))?; + observe(&event); + for outbound in OPENAI_RESPONSES_WS_CONFIG + .transform_ws_response(&event, model)? + .events + { + client_out.send(outbound) + .await + .map_err(|error| CoreError::Network(error.to_string()))?; + } + } + Message::Close(_) => break, + _ => {} + } + } + _ = tokio::time::sleep(idle) => break, + } + } + Ok(()) +} + +#[allow(clippy::too_many_arguments)] +pub async fn async_responses_websocket( + model: &str, + api_key: Option<&str>, + api_base: Option<&str>, + first_frame: Option, + idle_timeout: Option, + mut observe: impl FnMut(&ResponsesWsEvent) + Send, + client_in: In, + client_out: Out, +) -> CoreResult<()> +where + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, +{ + let key = resolve_api_key(api_key)?; + let upstream = dial_upstream(model, &key, api_base).await?; + let (mut upstream_tx, upstream_rx) = upstream.split(); + if let Some(first_frame) = first_frame { + for outbound in OPENAI_RESPONSES_WS_CONFIG + .transform_ws_request(&first_frame, model)? + .events + { + let payload = serde_json::to_string(&outbound) + .map_err(|error| CoreError::InvalidResponse(error.to_string()))?; + upstream_tx + .send(Message::Text(payload)) + .await + .map_err(|error| CoreError::Network(error.to_string()))?; + } + } + ResponsesWebSocketStreaming::bidirectional_forward( + model, + upstream_tx, + upstream_rx, + idle_timeout, + &mut observe, + client_in, + client_out, + ) + .await +} + +#[allow(clippy::too_many_arguments)] +pub async fn responses_ws( + model: &str, + api_key: Option<&str>, + api_base: Option<&str>, + first_frame: Option, + idle_timeout: Option, + observe: impl FnMut(&ResponsesWsEvent) + Send, + client_in: In, + client_out: Out, +) -> CoreResult<()> +where + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, +{ + async_responses_websocket( + model, + api_key, + api_base, + first_frame, + idle_timeout, + observe, + client_in, + client_out, + ) + .await +} + +#[cfg(test)] +mod tests { + use super::*; + use futures_channel::mpsc; + use futures_util::{SinkExt, StreamExt}; + use litellm_core::responses::types::ResponsesWsEventType; + use serde_json::json; + use tokio::io::AsyncWriteExt; + use tokio::net::TcpListener; + use tokio_tungstenite::accept_async; + + async fn websocket_base() -> (String, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("local address"); + let task = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept"); + let mut socket = accept_async(stream).await.expect("websocket handshake"); + while let Some(Ok(Message::Text(text))) = socket.next().await { + let request: serde_json::Value = serde_json::from_str(&text).expect("request json"); + let model = request + .get("model") + .and_then(serde_json::Value::as_str) + .or_else(|| { + request + .get("response") + .and_then(serde_json::Value::as_object) + .and_then(|response| { + response.get("model").and_then(serde_json::Value::as_str) + }) + }) + .expect("enforced model"); + socket + .send(Message::Text( + json!({ + "type": "response.created", + "response": { + "id": format!("resp-{model}"), + "model": model, + "extra": "preserved" + } + }) + .to_string(), + )) + .await + .expect("created event"); + socket + .send(Message::Text( + json!({ + "type": "response.completed", + "response": { + "id": format!("resp-{model}"), + "model": model, + "usage": { + "input_tokens": 1, + "output_tokens": 2, + "total_tokens": 3 + } + } + }) + .to_string(), + )) + .await + .expect("completed event"); + } + }); + (format!("http://{address}"), task) + } + + fn event(value: serde_json::Value) -> ResponsesWsEvent { + serde_json::from_value(value).expect("event") + } + + #[test] + fn explicit_nonblank_key_wins() { + assert_eq!( + resolve_api_key(Some(" explicit ")).expect("key"), + "explicit" + ); + } + + #[test] + fn blank_key_is_not_accepted_without_environment_key() { + if std::env::var(OPENAI_API_KEY_ENV).is_err() { + assert!(resolve_api_key(Some(" ")).is_err()); + } + } + + #[tokio::test] + async fn forwards_events_sequentially_and_enforces_model() { + let (api_base, server) = websocket_base().await; + let (client_tx, client_rx) = mpsc::unbounded(); + let (output_tx, mut output_rx) = mpsc::unbounded(); + let (observed_tx, observed_rx) = mpsc::unbounded(); + client_tx + .unbounded_send(event(json!({ + "type": "response.create", + "model": "wrong" + }))) + .expect("first request"); + client_tx + .unbounded_send(event(json!({ + "type": "response.create", + "response": {"model": "also-wrong"} + }))) + .expect("second request"); + + let task = tokio::spawn(async move { + responses_ws( + "authorized-model", + Some("test-key"), + Some(&api_base), + None, + Some(Duration::from_secs(1)), + move |event| { + observed_tx + .unbounded_send(event.clone()) + .expect("observe event"); + }, + client_rx, + output_tx, + ) + .await + }); + + let first = output_rx.next().await.expect("first output"); + let second = output_rx.next().await.expect("second output"); + let third = output_rx.next().await.expect("third output"); + let fourth = output_rx.next().await.expect("fourth output"); + drop(client_tx); + task.await.expect("splice task").expect("successful splice"); + server.await.expect("server task"); + + assert_eq!(first.event_type, ResponsesWsEventType::ResponseCreated); + assert_eq!(first.model(), Some("authorized-model")); + assert_eq!(first.data["response"]["extra"], "preserved"); + assert_eq!(second.event_type, ResponsesWsEventType::ResponseCompleted); + assert_eq!(third.event_type, ResponsesWsEventType::ResponseCreated); + assert_eq!(fourth.event_type, ResponsesWsEventType::ResponseCompleted); + let observed: Vec<_> = observed_rx.collect().await; + assert_eq!(observed.len(), 4); + assert!(observed + .iter() + .all(|event| event.event_type != ResponsesWsEventType::ResponseCreate)); + } + + #[tokio::test] + async fn idle_timeout_ends_without_upstream_events() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept"); + let _socket = accept_async(stream).await.expect("handshake"); + tokio::time::sleep(Duration::from_secs(1)).await; + }); + let (_client_tx, client_rx) = mpsc::unbounded::(); + let (output_tx, mut output_rx) = mpsc::unbounded(); + let result = responses_ws( + "model", + Some("key"), + Some(&format!("http://{address}")), + None, + Some(Duration::from_millis(20)), + |_| {}, + client_rx, + output_tx, + ) + .await; + assert!(result.is_ok()); + assert!(output_rx.next().await.is_none()); + server.abort(); + } + + #[tokio::test] + async fn dial_http_status_is_preserved() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("accept"); + stream + .write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\n\r\n") + .await + .expect("response"); + }); + let (_client_tx, client_rx) = mpsc::unbounded::(); + let (output_tx, _output_rx) = mpsc::unbounded(); + let error = responses_ws( + "model", + Some("key"), + Some(&format!("http://{address}")), + None, + Some(Duration::from_millis(20)), + |_| {}, + client_rx, + output_tx, + ) + .await + .expect_err("status error"); + assert!(matches!(error, CoreError::Http { status: 401, .. })); + server.await.expect("server task"); + } + + #[tokio::test] + async fn dial_http_500_status_is_preserved() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let address = listener.local_addr().expect("address"); + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("accept"); + stream + .write_all(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 0\r\n\r\n") + .await + .expect("response"); + }); + let (_client_tx, client_rx) = mpsc::unbounded::(); + let (output_tx, _output_rx) = mpsc::unbounded(); + let error = responses_ws( + "model", + Some("key"), + Some(&format!("http://{address}")), + None, + Some(Duration::from_millis(20)), + |_| {}, + client_rx, + output_tx, + ) + .await + .expect_err("status error"); + assert!(matches!(error, CoreError::Http { status: 500, .. })); + server.await.expect("server task"); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index db4c8211a5a..25aac3c495b 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -26,9 +26,6 @@ pub mod routes; #[cfg(feature = "server")] pub mod state; -// Realtime request logging. Only the server serves realtime, so these are -// `server`-gated; `io::realtime` exposes the generic `observe` hook while the -// collector and callback fan-out live here. mod constants; pub mod integrations; #[cfg(feature = "server")] diff --git a/litellm-rust/crates/ai-gateway/src/routes/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/mod.rs index 8872e9c8131..c26be8ffee3 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/mod.rs @@ -9,6 +9,7 @@ pub mod gil; pub mod health; pub mod messages; pub mod realtime; +pub mod responses; use axum::Router; @@ -21,5 +22,6 @@ pub fn app(state: AppState) -> Router { .merge(gil::router()) .merge(messages::router()) .merge(realtime::router()) + .merge(responses::router()) .with_state(state) } diff --git a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs new file mode 100644 index 00000000000..bdaffc97afb --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs @@ -0,0 +1,348 @@ +mod service; + +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::Arc; +use std::time::{SystemTime, UNIX_EPOCH}; + +use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade}; +use axum::extract::{Query, State}; +use axum::http::StatusCode; +use axum::response::Response; +use axum::routing::get; +use axum::Router; +use futures_util::{Sink, SinkExt, StreamExt}; +use litellm_core::responses::types::{ResponsesErrorFrame, ResponsesWsEvent, ResponsesWsEventType}; +use litellm_core::router::Router as ModelRouter; +use serde::Deserialize; + +use crate::auth::RequireMasterKey; +use crate::integrations::custom_logger::CustomLogger; +use crate::integrations::types::RequestMetadata; +use crate::state::AppState; + +static CALL_SEQ: AtomicU64 = AtomicU64::new(0); + +fn new_call_id() -> String { + let nanos = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_nanos()) + .unwrap_or(0); + let sequence = CALL_SEQ.fetch_add(1, Ordering::Relaxed); + format!("respws-{nanos:x}-{sequence:x}") +} + +pub fn router() -> Router { + Router::new() + .route("/v1/responses", get(handle)) + .route("/responses", get(handle)) +} + +#[derive(Debug, Deserialize)] +struct ResponsesQuery { + model: Option, +} + +async fn handle( + _auth: RequireMasterKey, + ws: WebSocketUpgrade, + State(state): State, + Query(query): Query, +) -> Result { + if let Some(model) = query.model.as_deref() { + validate_model(&state.router, model)?; + } + let router = state.router.clone(); + let loggers = state.loggers.clone(); + let master_key = state.master_key.clone(); + Ok(ws.on_upgrade(move |socket| bridge(socket, router, loggers, master_key, query.model))) +} + +fn validate_model(router: &ModelRouter, model: &str) -> Result<(), (StatusCode, String)> { + if model.trim().is_empty() { + return Err(( + StatusCode::BAD_REQUEST, + "missing 'model' query param".to_string(), + )); + } + let Some(deployment) = router.get_available_deployment(model) else { + return Err(( + StatusCode::NOT_FOUND, + format!("no deployment for model '{model}'"), + )); + }; + if deployment.litellm_params.model.contains('/') + && !deployment.litellm_params.model.starts_with("openai/") + { + return Err(( + StatusCode::BAD_REQUEST, + "Responses WebSocket route supports OpenAI deployments only".to_string(), + )); + } + Ok(()) +} + +async fn send_error_and_close(sink: &mut S, message: String) +where + S: futures_util::Sink + Unpin, + S::Error: std::fmt::Display, +{ + if let Ok(payload) = serde_json::to_string(&ResponsesErrorFrame::invalid_request(message)) { + let _ = sink.send(Message::Text(payload)).await; + } + let _ = sink + .send(Message::Close(Some(axum::extract::ws::CloseFrame { + code: 1008, + reason: "Pre-call error".into(), + }))) + .await; + let _ = sink.close().await; +} + +struct ResponseClientSink { + sink: futures_util::stream::SplitSink, +} + +impl Sink for ResponseClientSink { + type Error = axum::Error; + + fn poll_ready( + mut self: std::pin::Pin<&mut Self>, + context: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::pin::Pin::new(&mut self.sink).poll_ready(context) + } + + fn start_send( + mut self: std::pin::Pin<&mut Self>, + item: ResponsesWsEvent, + ) -> Result<(), Self::Error> { + let payload = serde_json::to_string(&item).map_err(axum::Error::new)?; + std::pin::Pin::new(&mut self.sink).start_send(Message::Text(payload)) + } + + fn poll_flush( + mut self: std::pin::Pin<&mut Self>, + context: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::pin::Pin::new(&mut self.sink).poll_flush(context) + } + + fn poll_close( + mut self: std::pin::Pin<&mut Self>, + context: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::pin::Pin::new(&mut self.sink).poll_close(context) + } +} + +impl ResponseClientSink { + async fn close_with_code(&mut self, code: u16, reason: &'static str) { + let _ = self + .sink + .send(Message::Close(Some(axum::extract::ws::CloseFrame { + code, + reason: reason.into(), + }))) + .await; + let _ = self.sink.close().await; + } +} + +async fn bridge( + socket: WebSocket, + router: Arc, + loggers: Arc>>, + master_key: Option>, + requested_model: Option, +) { + let (mut ws_sink, ws_stream) = socket.split(); + let (model, first_frame, stream) = if let Some(model) = requested_model { + (model, None, ws_stream) + } else { + let mut stream = ws_stream; + let first = match stream.next().await { + Some(Ok(Message::Text(text))) => { + match serde_json::from_str::(&text) { + Ok(event) => event, + Err(_) => { + send_error_and_close( + &mut ws_sink, + "Invalid JSON in response.create event".to_string(), + ) + .await; + return; + } + } + } + _ => { + send_error_and_close(&mut ws_sink, "Missing response.create event".to_string()) + .await; + return; + } + }; + let Some(model) = first.model().filter(|value| !value.trim().is_empty()) else { + send_error_and_close( + &mut ws_sink, + "Missing model in response.create event".to_string(), + ) + .await; + return; + }; + if first.event_type != ResponsesWsEventType::ResponseCreate { + send_error_and_close( + &mut ws_sink, + "First frame must be a response.create event".to_string(), + ) + .await; + return; + } + (model.to_string(), Some(first), stream) + }; + if let Err((status, message)) = validate_model(&router, &model) { + let _ = status; + let _ = message; + send_error_and_close(&mut ws_sink, "Unknown model deployment".to_string()).await; + return; + } + + let call_id = new_call_id(); + let metadata = RequestMetadata { + user_api_key_hash: master_key.as_deref().map(crate::auth::hash_token), + ..RequestMetadata::default() + }; + let client_in = Box::pin(stream.filter_map(|message| async move { + match message { + Ok(Message::Text(text)) => serde_json::from_str::(&text).ok(), + _ => None, + } + })); + let mut client_out = ResponseClientSink { sink: ws_sink }; + let result = service::run( + &router, + &model, + first_frame, + None, + loggers, + call_id, + metadata, + client_in, + &mut client_out, + ) + .await; + if result.is_err() { + client_out + .close_with_code(1011, "Internal server error") + .await; + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::io::realtime_pool::RealtimePool; + use crate::state::AppState; + use axum::body::Body; + use axum::http::Request; + use litellm_core::router::Router as ModelRouter; + use serde_json::json; + use std::pin::Pin; + use std::sync::Arc; + use std::task::{Context, Poll}; + use tower::ServiceExt; + + struct RecordingSink { + messages: Vec, + } + + impl Sink for RecordingSink { + type Error = std::convert::Infallible; + + fn poll_ready( + self: Pin<&mut Self>, + _context: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } + + fn start_send(mut self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> { + self.messages.push(item); + Ok(()) + } + + fn poll_flush( + self: Pin<&mut Self>, + _context: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close( + self: Pin<&mut Self>, + _context: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } + } + + #[tokio::test] + async fn pre_call_error_matches_python_frame_and_close() { + let mut sink = RecordingSink { + messages: Vec::new(), + }; + send_error_and_close(&mut sink, "missing model".to_string()).await; + let Message::Text(payload) = &sink.messages[0] else { + panic!("expected error text frame"); + }; + assert_eq!( + serde_json::from_str::(payload).expect("error json"), + json!({ + "type": "error", + "error": { + "type": "invalid_request_error", + "message": "missing model" + } + }) + ); + assert_eq!( + sink.messages[1], + Message::Close(Some(axum::extract::ws::CloseFrame { + code: 1008, + reason: "Pre-call error".into(), + })) + ); + } + + fn state() -> AppState { + AppState { + router: Arc::new(ModelRouter::default()), + master_key: Some(Arc::from("master-key")), + loggers: Arc::new(Vec::new()), + realtime_pool: RealtimePool::disabled(), + } + } + + #[tokio::test] + async fn auth_rejects_responses_upgrade_before_handler() { + let request = Request::builder() + .uri("/responses?model=known") + .body(Body::empty()) + .expect("request"); + let response = router() + .with_state(state()) + .oneshot(request) + .await + .expect("response"); + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + } + + #[test] + fn unknown_query_model_is_rejected_before_upgrade() { + assert_eq!( + validate_model(&ModelRouter::default(), "unknown").expect_err("unknown model"), + ( + StatusCode::NOT_FOUND, + "no deployment for model 'unknown'".to_string() + ) + ); + } +} diff --git a/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs b/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs new file mode 100644 index 00000000000..165c95695d3 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs @@ -0,0 +1,156 @@ +use std::sync::Arc; +use std::time::Duration; + +use futures_util::{Sink, Stream}; +use litellm_core::call_lifecycle::{CallLifecycle, CallLifecycleContext}; +use litellm_core::responses::instrumentation::{ + ResponsesWsCallbackPayload, ResponsesWsInstrumentation, ResponsesWsLogOutcome, + ResponsesWsMetadata, +}; +use litellm_core::responses::types::ResponsesWsEvent; +use litellm_core::{CoreError, CoreResult}; + +use crate::integrations::custom_logger::{ + CallbackTiming, CallbackValue, CustomLogger, CustomLoggerRunner, LoggingError, ModelCallDetails, +}; +use crate::integrations::types::RequestMetadata; + +#[allow(clippy::too_many_arguments)] +pub async fn run( + router: &litellm_core::router::Router, + model: &str, + first_frame: Option, + idle_timeout: Option, + loggers: Arc>>, + call_id: String, + metadata: RequestMetadata, + client_in: In, + client_out: Out, +) -> CoreResult<()> +where + In: Stream + Unpin + Send, + Out: Sink + Unpin + Send, + Out::Error: std::fmt::Display, +{ + let deployment = router.get_available_deployment(model).ok_or_else(|| { + CoreError::Routing(format!("no deployment available for model '{model}'")) + })?; + let params = &deployment.litellm_params; + let provider_model = params + .model + .strip_prefix("openai/") + .unwrap_or(¶ms.model); + if params.model.contains('/') && !params.model.starts_with("openai/") { + return Err(CoreError::InvalidProvider( + "Responses WebSocket route supports OpenAI deployments only".to_string(), + )); + } + let instrumentation = Arc::new(ResponsesWsInstrumentation::new( + call_id.clone(), + model, + ResponsesWsMetadata { + user_api_key_hash: metadata.user_api_key_hash, + user_api_key_user_id: metadata.user_api_key_user_id, + user_api_key_team_id: metadata.user_api_key_team_id, + }, + )); + let observer_instrumentation = Arc::clone(&instrumentation); + let context = CallLifecycleContext::new("responses_websocket", model, "openai", call_id); + let result = CallLifecycle::default() + .run(context, (), instrumentation.as_ref(), |_| async move { + crate::io::responses_ws::async_responses_websocket( + provider_model, + params.api_key.as_deref(), + params.api_base.as_deref(), + first_frame, + idle_timeout, + move |event| { + observer_instrumentation.observe(event); + }, + client_in, + client_out, + ) + .await + }) + .await; + let outcome = instrumentation.take_or_build_outcome(result.is_ok()); + dispatch_outcome(loggers, outcome).await; + result +} + +async fn dispatch_outcome( + loggers: Arc>>, + outcome: ResponsesWsLogOutcome, +) { + let runner = CustomLoggerRunner::new(loggers.as_ref().clone()); + match outcome { + ResponsesWsLogOutcome::Success { payload, callback } => { + let (details, response, start_time, end_time) = logging_values(payload, callback, None); + let _ = runner + .async_log_success_event( + &details, + &response, + CallbackTiming::new(start_time, end_time), + ) + .await; + } + ResponsesWsLogOutcome::Failure { + payload, + callback, + error_message, + error_kind, + } => { + let error = LoggingError { + message: error_message, + kind: error_kind, + }; + let (details, response, start_time, end_time) = + logging_values(payload, callback, Some(error)); + let _ = runner + .async_log_failure_event( + &details, + Some(&response), + CallbackTiming::new(start_time, end_time), + ) + .await; + } + } +} + +fn logging_values( + payload: litellm_core::responses::instrumentation::ResponsesWsLogPayload, + callback: ResponsesWsCallbackPayload, + error: Option, +) -> (ModelCallDetails, CallbackValue, f64, f64) { + let start_time = payload.start_time; + let end_time = payload.end_time; + let callback = CallbackValue::new(callback.object, callback.value); + let details = ModelCallDetails::from_standard_logging_payload( + crate::integrations::types::StandardLoggingPayload { + id: payload.id, + litellm_call_id: payload.litellm_call_id, + call_type: payload.call_type, + model: payload.model, + custom_llm_provider: payload.custom_llm_provider, + response_cost: payload.response_cost, + prompt_tokens: payload.usage.prompt_tokens, + completion_tokens: payload.usage.completion_tokens, + total_tokens: payload.usage.total_tokens, + start_time: payload.start_time, + end_time: payload.end_time, + stream: payload.stream, + metadata: crate::integrations::types::StandardLoggingMetadata { + user_api_key_hash: payload.metadata.user_api_key_hash, + user_api_key_user_id: payload.metadata.user_api_key_user_id, + user_api_key_team_id: payload.metadata.user_api_key_team_id, + ..Default::default() + }, + messages: None, + }, + ); + let details = match error { + Some(error) => details.with_failure_error(error), + None => details, + }; + (details, callback, start_time, end_time) +} diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs new file mode 100644 index 00000000000..5826a5bc9c1 --- /dev/null +++ b/litellm-rust/crates/core/src/constants.rs @@ -0,0 +1,3 @@ +pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com"; +pub const OPENAI_RESPONSES_DEFAULT_API_BASE: &str = "https://api.openai.com/v1"; +pub const OPENAI_RESPONSES_PATH: &str = "/responses"; diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index 27154f5a08b..117596f53d5 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -1,9 +1,11 @@ pub mod call_lifecycle; +pub mod constants; pub mod error; pub mod messages; pub mod ocr; pub mod providers; pub mod realtime; +pub mod responses; pub mod router; pub mod routing_utils; diff --git a/litellm-rust/crates/core/src/providers/openai/mod.rs b/litellm-rust/crates/core/src/providers/openai/mod.rs index 403e32975cf..62fcc50f2ac 100644 --- a/litellm-rust/crates/core/src/providers/openai/mod.rs +++ b/litellm-rust/crates/core/src/providers/openai/mod.rs @@ -1 +1,2 @@ pub mod realtime; +pub mod responses; diff --git a/litellm-rust/crates/core/src/providers/openai/responses/mod.rs b/litellm-rust/crates/core/src/providers/openai/responses/mod.rs new file mode 100644 index 00000000000..f239b6921fa --- /dev/null +++ b/litellm-rust/crates/core/src/providers/openai/responses/mod.rs @@ -0,0 +1 @@ +pub mod transformation; diff --git a/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs b/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs new file mode 100644 index 00000000000..ece10971806 --- /dev/null +++ b/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs @@ -0,0 +1,48 @@ +use crate::responses::types::{ResponsesWsEvent, ResponsesWsTransformResult}; +use crate::responses::websocket::{enforce_model, ResponsesWebSocketProviderConfig}; +use crate::CoreResult; + +pub struct OpenAIResponsesWsConfig; + +pub const OPENAI_RESPONSES_WS_CONFIG: OpenAIResponsesWsConfig = OpenAIResponsesWsConfig; + +impl ResponsesWebSocketProviderConfig for OpenAIResponsesWsConfig { + fn supports_native_websocket(&self) -> bool { + true + } + + fn transform_ws_request( + &self, + event: &ResponsesWsEvent, + model: &str, + ) -> CoreResult { + Ok(ResponsesWsTransformResult::passthrough(enforce_model( + event, model, + ))) + } + + fn transform_ws_response( + &self, + event: &ResponsesWsEvent, + _model: &str, + ) -> CoreResult { + Ok(ResponsesWsTransformResult::passthrough(event.clone())) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn openai_config_is_native_and_enforces_model() { + let event: ResponsesWsEvent = + serde_json::from_value(serde_json::json!({"type":"response.create"})) + .expect("valid event"); + let result = OPENAI_RESPONSES_WS_CONFIG + .transform_ws_request(&event, "gpt-5") + .expect("valid transform"); + assert_eq!(result.events[0].model(), Some("gpt-5")); + assert!(OPENAI_RESPONSES_WS_CONFIG.supports_native_websocket()); + } +} diff --git a/litellm-rust/crates/core/src/responses/instrumentation.rs b/litellm-rust/crates/core/src/responses/instrumentation.rs new file mode 100644 index 00000000000..ec04571da14 --- /dev/null +++ b/litellm-rust/crates/core/src/responses/instrumentation.rs @@ -0,0 +1,365 @@ +use std::future::Future; +use std::pin::Pin; +use std::sync::Mutex; +use std::time::{SystemTime, UNIX_EPOCH}; + +use serde_json::Value; + +use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; +use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType}; +use crate::{CoreError, CoreResult}; + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct ResponsesWsUsage { + pub prompt_tokens: u64, + pub completion_tokens: u64, + pub total_tokens: u64, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct ResponsesWsMetadata { + pub user_api_key_hash: Option, + pub user_api_key_user_id: Option, + pub user_api_key_team_id: Option, +} + +#[derive(Clone, Debug, PartialEq)] +pub struct ResponsesWsLogPayload { + pub id: String, + pub litellm_call_id: String, + pub call_type: String, + pub model: String, + pub custom_llm_provider: String, + pub response_cost: f64, + pub usage: ResponsesWsUsage, + pub start_time: f64, + pub end_time: f64, + pub stream: bool, + pub metadata: ResponsesWsMetadata, +} + +#[derive(Clone, Debug, PartialEq)] +pub enum ResponsesWsLogOutcome { + Success { + payload: ResponsesWsLogPayload, + callback: ResponsesWsCallbackPayload, + }, + Failure { + payload: ResponsesWsLogPayload, + callback: ResponsesWsCallbackPayload, + error_message: String, + error_kind: String, + }, +} + +#[derive(Clone, Debug, PartialEq)] +pub struct ResponsesWsCallbackPayload { + pub object: String, + pub value: Value, +} + +struct InstrumentationState { + litellm_call_id: String, + id: String, + model: String, + usage: ResponsesWsUsage, + start_time: f64, + end_time: f64, + metadata: ResponsesWsMetadata, + outcome: Option, +} + +pub struct ResponsesWsInstrumentation { + state: Mutex, +} + +impl ResponsesWsInstrumentation { + pub fn new( + litellm_call_id: impl Into, + model: impl Into, + metadata: ResponsesWsMetadata, + ) -> Self { + let litellm_call_id = litellm_call_id.into(); + let now = epoch_seconds(); + Self { + state: Mutex::new(InstrumentationState { + id: litellm_call_id.clone(), + litellm_call_id, + model: model.into(), + usage: ResponsesWsUsage::default(), + start_time: now, + end_time: now, + metadata, + outcome: None, + }), + } + } + + pub fn observe(&self, event: &ResponsesWsEvent) { + if !matches!( + event.event_type, + ResponsesWsEventType::ResponseCreated + | ResponsesWsEventType::ResponseCompleted + | ResponsesWsEventType::ResponseFailed + | ResponsesWsEventType::ResponseIncomplete + | ResponsesWsEventType::Error + ) { + return; + } + let Ok(mut state) = self.state.lock() else { + return; + }; + let Some(response) = event.data.get("response").and_then(Value::as_object) else { + return; + }; + if let Some(id) = response + .get("id") + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) + { + state.id = id.to_string(); + state.litellm_call_id = id.to_string(); + } + if let Some(model) = response + .get("model") + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) + { + state.model = model.to_string(); + } + let Some(usage) = response.get("usage").and_then(Value::as_object) else { + return; + }; + if let Some(input) = usage.get("input_tokens").and_then(Value::as_u64) { + state.usage.prompt_tokens += input; + } + if let Some(output) = usage.get("output_tokens").and_then(Value::as_u64) { + state.usage.completion_tokens += output; + } + state.usage.total_tokens += usage + .get("total_tokens") + .and_then(Value::as_u64) + .unwrap_or_else(|| { + usage + .get("input_tokens") + .and_then(Value::as_u64) + .unwrap_or(0) + + usage + .get("output_tokens") + .and_then(Value::as_u64) + .unwrap_or(0) + }); + } + + pub fn success_outcome(&self) -> ResponsesWsLogOutcome { + let mut state = self + .state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + state.end_time = epoch_seconds(); + ResponsesWsLogOutcome::Success { + payload: build_payload(&state), + callback: ResponsesWsCallbackPayload { + object: "responses_websocket".to_string(), + value: Value::Null, + }, + } + } + + pub fn failure_outcome(&self) -> ResponsesWsLogOutcome { + let mut state = self + .state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + state.end_time = epoch_seconds(); + ResponsesWsLogOutcome::Failure { + payload: build_payload(&state), + callback: ResponsesWsCallbackPayload { + object: "error".to_string(), + value: serde_json::json!({ + "message": "Responses WebSocket session ended in failure", + "kind": "ResponsesWebSocketError", + }), + }, + error_message: "Responses WebSocket session ended in failure".to_string(), + error_kind: "ResponsesWebSocketError".to_string(), + } + } + + pub fn take_outcome(&self) -> Option { + self.state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .outcome + .take() + } + + pub fn take_or_build_outcome(&self, success: bool) -> ResponsesWsLogOutcome { + self.take_outcome().unwrap_or_else(|| { + if success { + self.success_outcome() + } else { + self.failure_outcome() + } + }) + } +} + +type LifecycleFuture<'a, T> = Pin> + Send + 'a>>; + +impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation { + type PreCallFuture<'a> = LifecycleFuture<'a, ()>; + type DuringCallFuture<'a> = LifecycleFuture<'a, ()>; + type SuccessFuture<'a> = Pin + Send + 'a>>; + type FailureFuture<'a> = Pin + Send + 'a>>; + + fn async_pre_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + request: (), + ) -> Self::PreCallFuture<'a> { + Box::pin(async move { Ok(request) }) + } + + fn async_during_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + request: (), + ) -> Self::DuringCallFuture<'a> { + Box::pin(async move { Ok(request) }) + } + + fn async_log_success_event<'a>( + &'a self, + _context: &'a CallLifecycleContext, + _response: &'a (), + _timing: &'a CallLifecycleTiming, + ) -> Self::SuccessFuture<'a> { + Box::pin(async move { + let outcome = self.success_outcome(); + if let Ok(mut state) = self.state.lock() { + state.outcome = Some(outcome); + } + }) + } + + fn async_log_failure_event<'a>( + &'a self, + _context: &'a CallLifecycleContext, + _error: &'a CoreError, + _timing: &'a CallLifecycleTiming, + ) -> Self::FailureFuture<'a> { + Box::pin(async move { + let outcome = self.failure_outcome(); + if let Ok(mut state) = self.state.lock() { + state.outcome = Some(outcome); + } + }) + } +} + +fn build_payload(state: &InstrumentationState) -> ResponsesWsLogPayload { + ResponsesWsLogPayload { + id: state.id.clone(), + litellm_call_id: state.litellm_call_id.clone(), + call_type: "responses_websocket".to_string(), + model: state.model.clone(), + custom_llm_provider: "openai".to_string(), + response_cost: 0.0, + usage: state.usage.clone(), + start_time: state.start_time, + end_time: state.end_time, + stream: true, + metadata: state.metadata.clone(), + } +} + +fn epoch_seconds() -> f64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_secs_f64()) + .unwrap_or(0.0) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn event(value: Value) -> ResponsesWsEvent { + serde_json::from_value(value).expect("valid Responses WebSocket event") + } + + #[test] + fn accumulates_upstream_usage_and_identity() { + let instrumentation = + ResponsesWsInstrumentation::new("call-1", "gpt-5", ResponsesWsMetadata::default()); + instrumentation.observe(&event(serde_json::json!({ + "type": "response.completed", + "response": { + "id": "resp-1", + "model": "gpt-5-mini", + "usage": { + "input_tokens": 3, + "output_tokens": 5, + "total_tokens": 8 + } + } + }))); + + let ResponsesWsLogOutcome::Success { payload, .. } = instrumentation.success_outcome() + else { + panic!("expected success outcome"); + }; + assert_eq!(payload.id, "resp-1"); + assert_eq!(payload.model, "gpt-5-mini"); + assert_eq!(payload.usage.prompt_tokens, 3); + assert_eq!(payload.usage.completion_tokens, 5); + assert_eq!(payload.usage.total_tokens, 8); + assert!(payload.end_time >= payload.start_time); + } + + #[test] + fn builds_failure_payload_without_dispatching_callbacks() { + let instrumentation = + ResponsesWsInstrumentation::new("call-1", "gpt-5", ResponsesWsMetadata::default()); + assert!(matches!( + instrumentation.failure_outcome(), + ResponsesWsLogOutcome::Failure { .. } + )); + } + + #[tokio::test] + async fn lifecycle_records_success_outcome_for_provider_completion() { + let instrumentation = + ResponsesWsInstrumentation::new("call-1", "gpt-5", ResponsesWsMetadata::default()); + let result = crate::call_lifecycle::CallLifecycle::default() + .run( + crate::call_lifecycle::CallLifecycleContext::new( + "responses_websocket", + "gpt-5", + "openai", + "call-1", + ), + (), + &instrumentation, + |_| async { Ok::<(), CoreError>(()) }, + ) + .await; + + assert!(result.is_ok()); + assert!(matches!( + instrumentation.take_outcome(), + Some(ResponsesWsLogOutcome::Success { .. }) + )); + } + + #[test] + fn builds_outcome_when_lifecycle_did_not_record_one() { + let instrumentation = + ResponsesWsInstrumentation::new("call-1", "gpt-5", ResponsesWsMetadata::default()); + assert!(matches!( + instrumentation.take_or_build_outcome(true), + ResponsesWsLogOutcome::Success { .. } + )); + } +} diff --git a/litellm-rust/crates/core/src/responses/mod.rs b/litellm-rust/crates/core/src/responses/mod.rs new file mode 100644 index 00000000000..5ec5a2caef8 --- /dev/null +++ b/litellm-rust/crates/core/src/responses/mod.rs @@ -0,0 +1,3 @@ +pub mod instrumentation; +pub mod types; +pub mod websocket; diff --git a/litellm-rust/crates/core/src/responses/types.rs b/litellm-rust/crates/core/src/responses/types.rs new file mode 100644 index 00000000000..4942309992e --- /dev/null +++ b/litellm-rust/crates/core/src/responses/types.rs @@ -0,0 +1,166 @@ +use serde::{Deserialize, Deserializer, Serialize, Serializer}; +use serde_json::{Map, Value}; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum ResponsesWsEventType { + ResponseCreate, + ResponseCreated, + ResponseCompleted, + ResponseFailed, + ResponseIncomplete, + Error, + Other(String), +} + +impl ResponsesWsEventType { + pub fn as_str(&self) -> &str { + match self { + Self::ResponseCreate => "response.create", + Self::ResponseCreated => "response.created", + Self::ResponseCompleted => "response.completed", + Self::ResponseFailed => "response.failed", + Self::ResponseIncomplete => "response.incomplete", + Self::Error => "error", + Self::Other(value) => value, + } + } +} + +impl Serialize for ResponsesWsEventType { + fn serialize(&self, serializer: S) -> Result + where + S: Serializer, + { + serializer.serialize_str(self.as_str()) + } +} + +impl<'de> Deserialize<'de> for ResponsesWsEventType { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + let value = String::deserialize(deserializer)?; + Ok(match value.as_str() { + "response.create" => Self::ResponseCreate, + "response.created" => Self::ResponseCreated, + "response.completed" => Self::ResponseCompleted, + "response.failed" => Self::ResponseFailed, + "response.incomplete" => Self::ResponseIncomplete, + "error" => Self::Error, + _ => Self::Other(value), + }) + } +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ResponsesWsEvent { + #[serde(rename = "type")] + pub event_type: ResponsesWsEventType, + #[serde(flatten)] + pub data: Map, +} + +impl ResponsesWsEvent { + pub fn model(&self) -> Option<&str> { + let model = self.data.get("model").and_then(Value::as_str); + if model.is_some() { + return model; + } + self.data + .get("response") + .and_then(Value::as_object) + .and_then(|response| response.get("model")) + .and_then(Value::as_str) + } + + pub fn is_response_create(&self) -> bool { + self.event_type == ResponsesWsEventType::ResponseCreate + } +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ResponsesWsTransformResult { + pub events: Vec, +} + +impl ResponsesWsTransformResult { + pub fn passthrough(event: ResponsesWsEvent) -> Self { + Self { + events: vec![event], + } + } +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct ResponsesErrorFrame { + #[serde(rename = "type")] + pub frame_type: &'static str, + pub error: ResponsesErrorBody, +} + +impl ResponsesErrorFrame { + pub fn invalid_request(message: impl Into) -> Self { + Self { + frame_type: "error", + error: ResponsesErrorBody { + error_type: "invalid_request_error", + message: message.into(), + }, + } + } +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct ResponsesErrorBody { + #[serde(rename = "type")] + pub error_type: &'static str, + pub message: String, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn event_type_round_trips_known_and_unknown_values() { + let known: ResponsesWsEventType = + serde_json::from_str("\"response.completed\"").expect("valid event type"); + assert_eq!(known, ResponsesWsEventType::ResponseCompleted); + let unknown: ResponsesWsEventType = + serde_json::from_str("\"response.output_text.delta\"").expect("valid event type"); + assert_eq!( + unknown, + ResponsesWsEventType::Other("response.output_text.delta".to_string()) + ); + } + + #[test] + fn error_frame_matches_proxy_shape() { + let frame = ResponsesErrorFrame::invalid_request("missing model"); + assert_eq!( + serde_json::to_value(frame).expect("serializable"), + serde_json::json!({ + "type": "error", + "error": { + "type": "invalid_request_error", + "message": "missing model" + } + }) + ); + } + + #[test] + fn model_reads_flat_and_nested_create_shapes() { + let flat: ResponsesWsEvent = + serde_json::from_value(serde_json::json!({"type":"response.create","model":"gpt-5"})) + .expect("valid event"); + let nested: ResponsesWsEvent = serde_json::from_value(serde_json::json!({ + "type":"response.create", + "response":{"model":"gpt-5-mini"} + })) + .expect("valid event"); + assert_eq!(flat.model(), Some("gpt-5")); + assert_eq!(nested.model(), Some("gpt-5-mini")); + } +} diff --git a/litellm-rust/crates/core/src/responses/websocket.rs b/litellm-rust/crates/core/src/responses/websocket.rs new file mode 100644 index 00000000000..1edffd44985 --- /dev/null +++ b/litellm-rust/crates/core/src/responses/websocket.rs @@ -0,0 +1,188 @@ +use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH}; +use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult}; +use crate::CoreResult; + +pub trait ResponsesWebSocketProviderConfig: Sync { + fn supports_native_websocket(&self) -> bool { + false + } + + fn model_in_websocket_url(&self) -> bool { + true + } + + fn complete_websocket_url(&self, api_base: Option<&str>, model: &str) -> String { + complete_websocket_url(api_base, model, self.model_in_websocket_url()) + } + + fn transform_ws_request( + &self, + event: &ResponsesWsEvent, + model: &str, + ) -> CoreResult; + + fn transform_ws_response( + &self, + event: &ResponsesWsEvent, + model: &str, + ) -> CoreResult; +} + +pub fn complete_websocket_url( + api_base: Option<&str>, + model: &str, + model_in_websocket_url: bool, +) -> String { + let base = api_base + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or(OPENAI_RESPONSES_DEFAULT_API_BASE); + let (base_without_query, query) = base + .split_once('?') + .map_or((base, None), |(value, query)| (value, Some(query))); + let response_url = format!( + "{}{}", + base_without_query.trim_end_matches('/'), + OPENAI_RESPONSES_PATH + ); + let scheme_flipped = if let Some(rest) = response_url.strip_prefix("https://") { + format!("wss://{rest}") + } else if let Some(rest) = response_url.strip_prefix("http://") { + format!("ws://{rest}") + } else { + response_url + }; + let url = query.map_or(scheme_flipped.clone(), |value| { + format!("{scheme_flipped}?{value}") + }); + if !model_in_websocket_url + || query.is_some_and(|value| { + value + .split('&') + .any(|part| part.split('=').next() == Some("model")) + }) + { + return url; + } + format!( + "{url}{}model={}", + if query.is_some() { "&" } else { "?" }, + percent_encode(model) + ) +} + +fn percent_encode(value: &str) -> String { + value + .bytes() + .map(|byte| { + if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~') { + format!("{}", byte as char) + } else { + format!("%{byte:02X}") + } + }) + .collect() +} + +pub fn enforce_model(event: &ResponsesWsEvent, model: &str) -> ResponsesWsEvent { + if !event.is_response_create() { + return event.clone(); + } + let mut enforced = event.clone(); + let has_flat_model = enforced.data.contains_key("model"); + if let Some(response) = enforced + .data + .get_mut("response") + .and_then(serde_json::Value::as_object_mut) + { + response.insert( + "model".to_string(), + serde_json::Value::String(model.to_string()), + ); + if has_flat_model { + enforced.data.insert( + "model".to_string(), + serde_json::Value::String(model.to_string()), + ); + } + } else { + enforced.data.insert( + "model".to_string(), + serde_json::Value::String(model.to_string()), + ); + } + enforced +} + +pub fn is_terminal_event(event_type: &ResponsesWsEventType) -> bool { + matches!( + event_type, + ResponsesWsEventType::ResponseCreated + | ResponsesWsEventType::ResponseCompleted + | ResponsesWsEventType::ResponseFailed + | ResponsesWsEventType::ResponseIncomplete + | ResponsesWsEventType::Error + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn event(value: serde_json::Value) -> ResponsesWsEvent { + serde_json::from_value(value).expect("valid event") + } + + #[test] + fn url_construction_matches_python_defaults_and_query_behavior() { + assert_eq!( + complete_websocket_url(None, "gpt-5", true), + "wss://api.openai.com/v1/responses?model=gpt-5" + ); + assert_eq!( + complete_websocket_url(Some("http://localhost:8080/"), "gpt 5", true), + "ws://localhost:8080/responses?model=gpt%205" + ); + assert_eq!( + complete_websocket_url(Some("https://example.test/v1?foo=bar"), "gpt-5", true), + "wss://example.test/v1/responses?foo=bar&model=gpt-5" + ); + assert_eq!( + complete_websocket_url(Some("https://example.test?model=existing"), "gpt-5", true), + "wss://example.test/responses?model=existing" + ); + } + + #[test] + fn enforce_model_overrides_flat_and_nested_values() { + let flat = enforce_model( + &event(serde_json::json!({"type":"response.create","model":"wrong"})), + "gpt-5", + ); + assert_eq!(flat.model(), Some("gpt-5")); + let nested = enforce_model( + &event(serde_json::json!({ + "type":"response.create", + "model":"wrong", + "response":{"model":"also-wrong"} + })), + "gpt-5", + ); + assert_eq!(nested.model(), Some("gpt-5")); + assert_eq!( + nested + .data + .get("response") + .and_then(|value| value.get("model")), + Some(&serde_json::json!("gpt-5")) + ); + let nested_without_flat = enforce_model( + &event(serde_json::json!({ + "type":"response.create", + "response":{"model":"also-wrong"} + })), + "gpt-5", + ); + assert!(!nested_without_flat.data.contains_key("model")); + } +} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 77f4427127a..1decb789a22 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -1,7 +1,9 @@ +use std::collections::HashMap; use std::time::Duration; use litellm_ai_gateway::io::messages::{messages as run_messages, MessagesRequest}; use litellm_ai_gateway::io::ocr::{ocr as run_ocr, OcrRequest}; +use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use litellm_core::error::CoreError; use pyo3::exceptions::{PyRuntimeError, PyValueError}; use pyo3::prelude::*; @@ -65,6 +67,76 @@ fn optional_timeout(timeout_seconds: Option) -> Option { }) } +fn marshal_headers( + py: Python<'_>, + headers: Option>, +) -> PyResult> { + let value = match headers { + Some(headers) => py_to_json(py, headers.bind(py))?, + None => Value::Object(Map::new()), + }; + let Value::Object(headers) = value else { + return Err(PyValueError::new_err("headers must be a dict")); + }; + headers + .into_iter() + .map(|(name, value)| { + value + .as_str() + .map(|value| (name, value.to_string())) + .ok_or_else(|| PyValueError::new_err("header values must be strings")) + }) + .collect() +} + +#[pyclass] +struct ResponsesWebSocketConnection { + inner: RustResponsesWebSocketConnection, +} + +#[pymethods] +impl ResponsesWebSocketConnection { + #[classmethod] + #[pyo3(signature = (url, headers=None, timeout_seconds=None))] + fn connect<'py>( + _cls: &Bound<'py, pyo3::types::PyType>, + py: Python<'py>, + url: String, + headers: Option>, + timeout_seconds: Option, + ) -> PyResult> { + let headers = marshal_headers(py, headers)?; + let timeout = optional_timeout(timeout_seconds); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout) + .await + .map_err(core_error_to_pyerr)?; + Python::attach(|py| Py::new(py, ResponsesWebSocketConnection { inner })) + }) + } + + fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult> { + let inner = self.inner.clone(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + inner.send_text(text).await.map_err(core_error_to_pyerr) + }) + } + + fn recv_text<'py>(&self, py: Python<'py>) -> PyResult> { + let inner = self.inner.clone(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + inner.recv_text().await.map_err(core_error_to_pyerr) + }) + } + + fn close<'py>(&self, py: Python<'py>) -> PyResult> { + let inner = self.inner.clone(); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + inner.close().await.map_err(core_error_to_pyerr) + }) + } +} + fn marshal_inputs( py: Python<'_>, document: Py, @@ -271,6 +343,7 @@ fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> { module.add_function(wrap_pyfunction!(aocr, module)?)?; module.add_function(wrap_pyfunction!(messages, module)?)?; module.add_function(wrap_pyfunction!(amessages, module)?)?; + module.add_class::()?; module.add_function(wrap_pyfunction!(gil_stats, module)?)?; Ok(()) } diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 116ab905b88..c48d75439a7 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2,6 +2,7 @@ import asyncio import json import os import ssl +from contextlib import asynccontextmanager from functools import lru_cache from typing import ( TYPE_CHECKING, @@ -148,6 +149,14 @@ from litellm.utils import ( async_pre_call_deployment_hook, ) + +def _rust_responses_websocket_enabled( + custom_llm_provider: str | None, + litellm_params: GenericLiteLLMParams, +) -> bool: + return custom_llm_provider == "openai" and litellm_params.get("rust") is True + + from .http_handler import get_shared_realtime_ssl_context if TYPE_CHECKING: @@ -6221,12 +6230,29 @@ class BaseLLMHTTPHandler: }, ) - async with websockets.connect( # type: ignore - ws_url, - additional_headers=headers, - max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, - ssl=ssl_context, - ) as backend_ws: + @asynccontextmanager + async def _backend_connection(): + if _rust_responses_websocket_enabled(custom_llm_provider, litellm_params): + from litellm.rust_bridge import responses_websocket as rust_responses_websocket + + rust_backend = await rust_responses_websocket.connect( + url=ws_url, + headers={str(key): str(value) for key, value in headers.items()}, + timeout=timeout, + ) + if rust_backend is not None: + yield rust_backend + return + + async with websockets.connect( # type: ignore + ws_url, + additional_headers=headers, + max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, + ssl=ssl_context, + ) as backend: + yield backend + + async with _backend_connection() as backend_ws: _request_data: Dict[str, Any] = {} if litellm_metadata: _request_data["litellm_metadata"] = litellm_metadata diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 35de2eb9727..e9139a634f1 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -3,7 +3,7 @@ from __future__ import annotations import os -from typing import TYPE_CHECKING, Awaitable, Final, Protocol, Union, cast +from typing import TYPE_CHECKING, Any, Awaitable, Final, Protocol, Union, cast import httpx @@ -71,26 +71,33 @@ def use_litellm_rust( aocr: RustAocr | None | _Unset = _UNSET, messages: RustMessages | None | _Unset = _UNSET, amessages: RustAmessages | None | _Unset = _UNSET, + responses_websocket: Any | None | _Unset = _UNSET, ) -> None: global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl configuring_ocr = not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset) configuring_messages = not isinstance(messages, _Unset) or not isinstance(amessages, _Unset) - if configuring_ocr or not configuring_messages: + configuring_responses_websocket = not isinstance(responses_websocket, _Unset) + if configuring_ocr or (not configuring_messages and not configuring_responses_websocket): _rust_ocr_enabled = enabled if not isinstance(ocr, _Unset): _rust_ocr_impl = ocr if not isinstance(aocr, _Unset): _rust_aocr_impl = aocr - if not configuring_messages: + if not configuring_messages and not configuring_responses_websocket: return - from litellm.rust_bridge.messages import set_rust_messages + if configuring_messages: + from litellm.rust_bridge.messages import set_rust_messages - if not isinstance(messages, _Unset) and not isinstance(amessages, _Unset): - set_rust_messages(messages=messages, amessages=amessages) - elif not isinstance(messages, _Unset): - set_rust_messages(messages=messages) - else: - set_rust_messages(amessages=amessages) + if not isinstance(messages, _Unset) and not isinstance(amessages, _Unset): + set_rust_messages(messages=messages, amessages=amessages) + elif not isinstance(messages, _Unset): + set_rust_messages(messages=messages) + else: + set_rust_messages(amessages=amessages) + if configuring_responses_websocket: + from litellm.rust_bridge.responses_websocket import set_rust_responses_websocket + + set_rust_responses_websocket(connection=responses_websocket) def rust_ocr_enabled() -> bool: diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py new file mode 100644 index 00000000000..5b3d486e8d3 --- /dev/null +++ b/litellm/rust_bridge/responses_websocket.py @@ -0,0 +1,95 @@ +"""Thin Python wrapper for the native Rust Responses WebSocket bridge.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Final, Protocol + +import httpx +from websockets.exceptions import ConnectionClosedOK + +from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.timeouts import timeout_to_seconds + + +class RustResponsesWebSocketConnection(Protocol): + @classmethod + def connect( + cls, + url: str, + headers: dict[str, str], + timeout_seconds: float | None, + ) -> Any: + raise NotImplementedError + + +class _Unset: + pass + + +_UNSET: Final[_Unset] = _Unset() + + +@dataclass(slots=True) +class _RustResponsesWebSocketState: + connection: Any = None + + +_STATE: Final[_RustResponsesWebSocketState] = _RustResponsesWebSocketState() + + +def set_rust_responses_websocket( + *, + connection: Any = _UNSET, +) -> None: + if not isinstance(connection, _Unset): + _STATE.connection = connection + + +def load_rust_responses_websocket() -> Any: + if _STATE.connection is not None: + return _STATE.connection + native_bridge = get_native_bridge() + if native_bridge is None: + return None + try: + return native_bridge.ResponsesWebSocketConnection + except AttributeError: + return None + + +class _ConnectionAdapter: + def __init__(self, connection: Any): + self._connection = connection + + async def send(self, text: str) -> None: + await self._connection.send_text(text) + + async def recv(self) -> str: + message = await self._connection.recv_text() + if message is None: + raise ConnectionClosedOK(None, None) + return message + + async def close(self) -> None: + await self._connection.close() + + +async def connect( + *, + url: str, + headers: dict[str, str], + timeout: float | httpx.Timeout | None, +) -> _ConnectionAdapter | None: + connection_type = load_rust_responses_websocket() + if connection_type is None: + return None + try: + connection = await connection_type.connect( + url=url, + headers=headers, + timeout_seconds=timeout_to_seconds(timeout), + ) + except Exception: # noqa: BLE001 # bridge failures must fall back to Python + return None + return _ConnectionAdapter(connection) diff --git a/tests/code_coverage_tests/liccheck.ini b/tests/code_coverage_tests/liccheck.ini index 4d6e0528ed4..f5a2fb4b14c 100644 --- a/tests/code_coverage_tests/liccheck.ini +++ b/tests/code_coverage_tests/liccheck.ini @@ -114,6 +114,8 @@ apscheduler: >=3.10.4 # Unknown license fastapi-sso: >=0.16.0 # Unknown license filelock: >=3.20.0 # Unlicense (public domain) - https://unlicense.org / https://github.com/tox-dev/filelock pyjwt: >=2.9.0 # Unknown license +vcrpy: >=8.2.1 # MIT License - https://github.com/kevin1024/vcrpy/blob/master/LICENSE.txt +locust: >=2.45.0 # MIT License - https://github.com/locustio/locust/blob/master/LICENSE python-multipart: >=0.0.20 # Unknown license pillow: >=11.0.0 # Unknown license azure-ai-contentsafety: >=1.0.0 # Unknown license diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py new file mode 100644 index 00000000000..c9a5b988be6 --- /dev/null +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +import pytest + +from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled +from litellm.rust_bridge import responses_websocket +from litellm.types.router import GenericLiteLLMParams + + +class _FakeNativeConnection: + def __init__(self) -> None: + self.sent: list[str] = [] + self.closed = False + + async def send_text(self, text: str) -> None: + self.sent.append(text) + + async def recv_text(self) -> str: + return "response.completed" + + async def close(self) -> None: + self.closed = True + + +class _ClosedNativeConnection: + async def recv_text(self) -> None: + return None + + +class _FakeNativeBridge: + @classmethod + async def connect( + cls, + *, + url: str, + headers: dict[str, str], + timeout_seconds: float | None, + ) -> _FakeNativeConnection: + return _FakeNativeConnection() + + +def test_rust_websocket_bridge_is_disabled_without_flag() -> None: + assert not _rust_responses_websocket_enabled("openai", GenericLiteLLMParams()) + assert not _rust_responses_websocket_enabled("anthropic", GenericLiteLLMParams(rust=True)) + assert _rust_responses_websocket_enabled("openai", GenericLiteLLMParams(rust=True)) + + +@pytest.mark.asyncio +async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None: + adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection()) + + with pytest.raises(responses_websocket.ConnectionClosedOK): + await adapter.recv() + + +@pytest.mark.asyncio +async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(responses_websocket, "_STATE", responses_websocket._RustResponsesWebSocketState()) + monkeypatch.setattr(responses_websocket, "get_native_bridge", lambda: None) + + assert ( + await responses_websocket.connect( + url="wss://example.test/responses", + headers={}, + timeout=None, + ) + is None + ) + + +@pytest.mark.asyncio +async def test_enabled_bridge_connects_and_adapts_socket( + monkeypatch: pytest.MonkeyPatch, +) -> None: + responses_websocket.set_rust_responses_websocket(connection=_FakeNativeBridge) + + connection = await responses_websocket.connect( + url="wss://example.test/responses", + headers={"Authorization": "Bearer key"}, + timeout=1.0, + ) + + assert connection is not None + await connection.send("response.create") + assert await connection.recv() == "response.completed" + await connection.close()