mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_otel_v2_admin_owned_destinations
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled
This commit is contained in:
commit
749d700571
33 changed files with 2809 additions and 23 deletions
|
|
@ -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
|
||||
|
|
|
|||
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -621,6 +621,7 @@ dependencies = [
|
|||
"subtle",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
"tower",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
|
|
@ -41,3 +41,4 @@ python-config = ["dep:pyo3"]
|
|||
|
||||
[dev-dependencies]
|
||||
futures-channel = "0.3"
|
||||
tower = { version = "0.5.3", features = ["util"] }
|
||||
|
|
|
|||
|
|
@ -40,3 +40,19 @@ 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;
|
||||
|
||||
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";
|
||||
|
||||
/// 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"];
|
||||
|
|
|
|||
|
|
@ -2,3 +2,4 @@ pub mod messages;
|
|||
pub mod ocr;
|
||||
pub mod realtime;
|
||||
pub mod realtime_pool;
|
||||
pub mod responses_ws;
|
||||
|
|
|
|||
548
litellm-rust/crates/ai-gateway/src/io/responses_ws.rs
Normal file
548
litellm-rust/crates/ai-gateway/src/io/responses_ws.rs
Normal file
|
|
@ -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<MaybeTlsStream<TcpStream>>;
|
||||
type UpstreamTx = SplitSink<ResponsesUpstreamWs, Message>;
|
||||
type UpstreamRx = SplitStream<ResponsesUpstreamWs>;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponsesWebSocketConnection {
|
||||
socket: Arc<Mutex<Option<ResponsesUpstreamWs>>>,
|
||||
}
|
||||
|
||||
impl ResponsesWebSocketConnection {
|
||||
pub async fn connect_url(
|
||||
url: &str,
|
||||
headers: &HashMap<String, String>,
|
||||
timeout: Option<Duration>,
|
||||
) -> CoreResult<Self> {
|
||||
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::<HeaderName>()
|
||||
.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<Option<String>> {
|
||||
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<String> {
|
||||
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<ResponsesUpstreamWs> {
|
||||
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<In, Out>(
|
||||
model: &str,
|
||||
upstream_tx: UpstreamTx,
|
||||
upstream_rx: UpstreamRx,
|
||||
idle_timeout: Option<Duration>,
|
||||
observe: impl FnMut(&ResponsesWsEvent) + Send,
|
||||
client_in: In,
|
||||
client_out: Out,
|
||||
) -> CoreResult<()>
|
||||
where
|
||||
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
|
||||
Out: Sink<ResponsesWsEvent> + 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<In, Out>(
|
||||
model: &str,
|
||||
mut upstream_tx: UpstreamTx,
|
||||
mut upstream_rx: UpstreamRx,
|
||||
idle_timeout: Option<Duration>,
|
||||
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
|
||||
mut client_in: In,
|
||||
mut client_out: Out,
|
||||
) -> CoreResult<()>
|
||||
where
|
||||
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
|
||||
Out: Sink<ResponsesWsEvent> + 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::<ResponsesWsEvent>(&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<In, Out>(
|
||||
model: &str,
|
||||
api_key: Option<&str>,
|
||||
api_base: Option<&str>,
|
||||
first_frame: Option<ResponsesWsEvent>,
|
||||
idle_timeout: Option<Duration>,
|
||||
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
|
||||
client_in: In,
|
||||
client_out: Out,
|
||||
) -> CoreResult<()>
|
||||
where
|
||||
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
|
||||
Out: Sink<ResponsesWsEvent> + 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<In, Out>(
|
||||
model: &str,
|
||||
api_key: Option<&str>,
|
||||
api_base: Option<&str>,
|
||||
first_frame: Option<ResponsesWsEvent>,
|
||||
idle_timeout: Option<Duration>,
|
||||
observe: impl FnMut(&ResponsesWsEvent) + Send,
|
||||
client_in: In,
|
||||
client_out: Out,
|
||||
) -> CoreResult<()>
|
||||
where
|
||||
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
|
||||
Out: Sink<ResponsesWsEvent> + 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::<ResponsesWsEvent>();
|
||||
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::<ResponsesWsEvent>();
|
||||
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::<ResponsesWsEvent>();
|
||||
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");
|
||||
}
|
||||
}
|
||||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<reqwest::Response> {
|
||||
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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<Value> {
|
||||
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<MessagesResponse> {
|
||||
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)]
|
||||
|
|
|
|||
|
|
@ -62,6 +62,7 @@ pub(super) fn prepare_messages_call(
|
|||
})?;
|
||||
|
||||
Ok(ProviderMessagesRequest {
|
||||
provider: provider.to_string(),
|
||||
model,
|
||||
config,
|
||||
url,
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
513
litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs
Normal file
513
litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs
Normal file
|
|
@ -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<AppState> {
|
||||
Router::new().route(MESSAGES_ROUTE_PATH, post(handle))
|
||||
}
|
||||
|
||||
async fn handle(
|
||||
_auth: RequireMasterKey,
|
||||
State(state): State<AppState>,
|
||||
headers: HeaderMap,
|
||||
Json(body): Json<Value>,
|
||||
) -> Result<Response, MessagesRouteError> {
|
||||
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<Response, MessagesRouteError> {
|
||||
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<Option<Map<String, Value>>, 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::<Result<Map<_, _>, CoreError>>()?;
|
||||
Ok((!forwarded.is_empty()).then_some(forwarded))
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct MessagesRouteError(CoreError);
|
||||
|
||||
impl From<CoreError> 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<String>) {
|
||||
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::<usize>().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<String>) {
|
||||
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::<usize>().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::<serde_json::Value>(&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::<serde_json::Value>(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::<serde_json::Value>(&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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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<Router>,
|
||||
body: Value,
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
) -> CoreResult<MessagesResponse> {
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
|
@ -7,7 +7,9 @@
|
|||
|
||||
pub mod gil;
|
||||
pub mod health;
|
||||
pub mod messages;
|
||||
pub mod realtime;
|
||||
pub mod responses;
|
||||
|
||||
use axum::Router;
|
||||
|
||||
|
|
@ -18,6 +20,8 @@ pub fn app(state: AppState) -> Router {
|
|||
Router::new()
|
||||
.merge(health::router())
|
||||
.merge(gil::router())
|
||||
.merge(messages::router())
|
||||
.merge(realtime::router())
|
||||
.merge(responses::router())
|
||||
.with_state(state)
|
||||
}
|
||||
|
|
|
|||
348
litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs
Normal file
348
litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs
Normal file
|
|
@ -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<AppState> {
|
||||
Router::new()
|
||||
.route("/v1/responses", get(handle))
|
||||
.route("/responses", get(handle))
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct ResponsesQuery {
|
||||
model: Option<String>,
|
||||
}
|
||||
|
||||
async fn handle(
|
||||
_auth: RequireMasterKey,
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<AppState>,
|
||||
Query(query): Query<ResponsesQuery>,
|
||||
) -> Result<Response, (StatusCode, String)> {
|
||||
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<S>(sink: &mut S, message: String)
|
||||
where
|
||||
S: futures_util::Sink<Message> + 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<WebSocket, Message>,
|
||||
}
|
||||
|
||||
impl Sink<ResponsesWsEvent> for ResponseClientSink {
|
||||
type Error = axum::Error;
|
||||
|
||||
fn poll_ready(
|
||||
mut self: std::pin::Pin<&mut Self>,
|
||||
context: &mut std::task::Context<'_>,
|
||||
) -> std::task::Poll<Result<(), Self::Error>> {
|
||||
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<Result<(), Self::Error>> {
|
||||
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<Result<(), Self::Error>> {
|
||||
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<ModelRouter>,
|
||||
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
|
||||
master_key: Option<Arc<str>>,
|
||||
requested_model: Option<String>,
|
||||
) {
|
||||
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::<ResponsesWsEvent>(&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::<ResponsesWsEvent>(&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<Message>,
|
||||
}
|
||||
|
||||
impl Sink<Message> for RecordingSink {
|
||||
type Error = std::convert::Infallible;
|
||||
|
||||
fn poll_ready(
|
||||
self: Pin<&mut Self>,
|
||||
_context: &mut Context<'_>,
|
||||
) -> Poll<Result<(), Self::Error>> {
|
||||
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<Result<(), Self::Error>> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
|
||||
fn poll_close(
|
||||
self: Pin<&mut Self>,
|
||||
_context: &mut Context<'_>,
|
||||
) -> Poll<Result<(), Self::Error>> {
|
||||
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::<serde_json::Value>(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()
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
156
litellm-rust/crates/ai-gateway/src/routes/responses/service.rs
Normal file
156
litellm-rust/crates/ai-gateway/src/routes/responses/service.rs
Normal file
|
|
@ -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<In, Out>(
|
||||
router: &litellm_core::router::Router,
|
||||
model: &str,
|
||||
first_frame: Option<ResponsesWsEvent>,
|
||||
idle_timeout: Option<Duration>,
|
||||
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
|
||||
call_id: String,
|
||||
metadata: RequestMetadata,
|
||||
client_in: In,
|
||||
client_out: Out,
|
||||
) -> CoreResult<()>
|
||||
where
|
||||
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
|
||||
Out: Sink<ResponsesWsEvent> + 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<Vec<Arc<dyn CustomLogger>>>,
|
||||
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<LoggingError>,
|
||||
) -> (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)
|
||||
}
|
||||
3
litellm-rust/crates/core/src/constants.rs
Normal file
3
litellm-rust/crates/core/src/constants.rs
Normal file
|
|
@ -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";
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -1 +1,2 @@
|
|||
pub mod realtime;
|
||||
pub mod responses;
|
||||
|
|
|
|||
|
|
@ -0,0 +1 @@
|
|||
pub mod transformation;
|
||||
|
|
@ -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<ResponsesWsTransformResult> {
|
||||
Ok(ResponsesWsTransformResult::passthrough(enforce_model(
|
||||
event, model,
|
||||
)))
|
||||
}
|
||||
|
||||
fn transform_ws_response(
|
||||
&self,
|
||||
event: &ResponsesWsEvent,
|
||||
_model: &str,
|
||||
) -> CoreResult<ResponsesWsTransformResult> {
|
||||
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());
|
||||
}
|
||||
}
|
||||
365
litellm-rust/crates/core/src/responses/instrumentation.rs
Normal file
365
litellm-rust/crates/core/src/responses/instrumentation.rs
Normal file
|
|
@ -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<String>,
|
||||
pub user_api_key_user_id: Option<String>,
|
||||
pub user_api_key_team_id: Option<String>,
|
||||
}
|
||||
|
||||
#[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<ResponsesWsLogOutcome>,
|
||||
}
|
||||
|
||||
pub struct ResponsesWsInstrumentation {
|
||||
state: Mutex<InstrumentationState>,
|
||||
}
|
||||
|
||||
impl ResponsesWsInstrumentation {
|
||||
pub fn new(
|
||||
litellm_call_id: impl Into<String>,
|
||||
model: impl Into<String>,
|
||||
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<ResponsesWsLogOutcome> {
|
||||
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<Box<dyn Future<Output = CoreResult<T>> + Send + 'a>>;
|
||||
|
||||
impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation {
|
||||
type PreCallFuture<'a> = LifecycleFuture<'a, ()>;
|
||||
type DuringCallFuture<'a> = LifecycleFuture<'a, ()>;
|
||||
type SuccessFuture<'a> = Pin<Box<dyn Future<Output = ()> + Send + 'a>>;
|
||||
type FailureFuture<'a> = Pin<Box<dyn Future<Output = ()> + 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 { .. }
|
||||
));
|
||||
}
|
||||
}
|
||||
3
litellm-rust/crates/core/src/responses/mod.rs
Normal file
3
litellm-rust/crates/core/src/responses/mod.rs
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
pub mod instrumentation;
|
||||
pub mod types;
|
||||
pub mod websocket;
|
||||
166
litellm-rust/crates/core/src/responses/types.rs
Normal file
166
litellm-rust/crates/core/src/responses/types.rs
Normal file
|
|
@ -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<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
serializer.serialize_str(self.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for ResponsesWsEventType {
|
||||
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
||||
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<String, Value>,
|
||||
}
|
||||
|
||||
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<ResponsesWsEvent>,
|
||||
}
|
||||
|
||||
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<String>) -> 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"));
|
||||
}
|
||||
}
|
||||
188
litellm-rust/crates/core/src/responses/websocket.rs
Normal file
188
litellm-rust/crates/core/src/responses/websocket.rs
Normal file
|
|
@ -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<ResponsesWsTransformResult>;
|
||||
|
||||
fn transform_ws_response(
|
||||
&self,
|
||||
event: &ResponsesWsEvent,
|
||||
model: &str,
|
||||
) -> CoreResult<ResponsesWsTransformResult>;
|
||||
}
|
||||
|
||||
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"));
|
||||
}
|
||||
}
|
||||
|
|
@ -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<f64>) -> Option<Duration> {
|
|||
})
|
||||
}
|
||||
|
||||
fn marshal_headers(
|
||||
py: Python<'_>,
|
||||
headers: Option<Py<PyAny>>,
|
||||
) -> PyResult<HashMap<String, String>> {
|
||||
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<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
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<Bound<'py, PyAny>> {
|
||||
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<Bound<'py, PyAny>> {
|
||||
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<Bound<'py, PyAny>> {
|
||||
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<PyAny>,
|
||||
|
|
@ -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::<ResponsesWebSocketConnection>()?;
|
||||
module.add_function(wrap_pyfunction!(gil_stats, module)?)?;
|
||||
Ok(())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
95
litellm/rust_bridge/responses_websocket.py
Normal file
95
litellm/rust_bridge/responses_websocket.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
86
tests/test_litellm/responses/test_rust_bridge_websocket.py
Normal file
86
tests/test_litellm/responses/test_rust_bridge_websocket.py
Normal file
|
|
@ -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()
|
||||
Loading…
Add table
Reference in a new issue