feat(rust): add typed streaming boundary

This commit is contained in:
Yujong Lee 2026-09-01 10:41:26 -07:00
parent b0ae5bc4e3
commit ceaf967da8
28 changed files with 2640 additions and 50 deletions

View file

@ -126,3 +126,15 @@ jobs:
- name: Test native route wheel
run: python .github/scripts/test_native_routes.py dist/*.whl
- name: Test Python-Rust bridge contract from release wheel
env:
LITELLM_REQUIRE_NATIVE_BRIDGE: "1"
run: |
uv venv --python python3.12 .native-test
uv pip install --python .native-test/bin/python dist/*.whl pytest==9.0.3 pytest-asyncio==1.3.0
cp tests/test_litellm/rust_bridge/test_native_integration.py /tmp/test_native_integration.py
cd /tmp
"$GITHUB_WORKSPACE/.native-test/bin/python" -m pytest \
/tmp/test_native_integration.py \
--rootdir=/tmp -q

View file

@ -1428,6 +1428,7 @@ dependencies = [
"aws-sigv4",
"aws-smithy-runtime-api",
"aws-types",
"futures-util",
"rand 0.8.7",
"reqwest",
"serde",

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
futures-util.workspace = true
rand.workspace = true
reqwest.workspace = true
serde.workspace = true

View file

@ -19,9 +19,12 @@ use serde_json::{Map, Value};
use crate::error::Error;
use crate::streaming::OpenedStream;
use handler::execute_chat_completions_provider_call;
use prepare::{parse_messages, prepare_chat_completions_call, resolve_provider_config};
use types::{ChatCompletionsRequest, ChatCompletionsResponse};
use types::{
ChatCompletionsRequest, ChatCompletionsResponse, ChatCompletionsStreamRequest, ChatStreamEvent,
};
pub async fn chat_completions(
request: ChatCompletionsRequest<'_>,
@ -55,5 +58,13 @@ pub fn chat_completions_decline_reason(
.map(|reason| reason.0)
}
pub async fn chat_completions_stream(
_request: ChatCompletionsStreamRequest,
) -> Result<OpenedStream<ChatStreamEvent>, Error> {
Err(Error::Unsupported(
"chat completions streaming provider registration",
))
}
#[cfg(test)]
mod tests;

View file

@ -6,6 +6,18 @@ use super::types::{
ChatCompletionsResponse, ChatMessage, ChatMessageContent, ProviderChatRequestData,
ProviderChatResponseData,
};
use super::types::{ChatCompletionsStreamRequest, ChatStreamEvent};
use crate::streaming::StreamProvider;
pub trait ChatCompletionsStreamProvider:
StreamProvider<ChatCompletionsStreamRequest, ChatStreamEvent>
{
}
impl<T> ChatCompletionsStreamProvider for T where
T: StreamProvider<ChatCompletionsStreamRequest, ChatStreamEvent>
{
}
/// How the upstream call is authenticated. API-key strategies are resolved in
/// `prepare`; SigV4 needs the serialized body, so the handler signs it.

View file

@ -4,6 +4,7 @@ use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use super::transformation::{ChatCompletionsAuth, ChatCompletionsProviderConfig};
use crate::streaming::{JsonObject, StreamTarget, StreamTransportOptions};
/// A `/chat/completions` call as it crosses into the core.
///
@ -110,3 +111,183 @@ pub struct ChatCompletionsResponse {
pub choices: Vec<ChatCompletionsChoice>,
pub usage: ChatCompletionsUsage,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ChatStreamRole {
Assistant,
Developer,
Function,
System,
Tool,
User,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ChatStreamMessageContent {
Text(String),
Parts(Vec<JsonObject>),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ChatStreamMessage {
pub role: ChatStreamRole,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub content: Option<ChatStreamMessageContent>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(flatten)]
pub extra: Map<String, Value>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ChatStreamStop {
One(String),
Many(Vec<String>),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ChatStreamStringOrObject {
Name(String),
Definition(JsonObject),
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct ChatCompletionsStreamParameters {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub temperature: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub top_p: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_completion_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stop: Option<ChatStreamStop>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stream: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stream_options: Option<JsonObject>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<JsonObject>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<ChatStreamStringOrObject>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub response_format: Option<JsonObject>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reasoning_effort: Option<ChatStreamStringOrObject>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub thinking: Option<JsonObject>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(flatten)]
pub extra: Map<String, Value>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ChatCompletionsStreamRequestBody {
pub model: String,
pub messages: Vec<ChatStreamMessage>,
#[serde(flatten)]
pub parameters: ChatCompletionsStreamParameters,
}
pub struct ChatCompletionsStreamRequest {
pub body: ChatCompletionsStreamRequestBody,
pub target: StreamTarget,
pub transport: StreamTransportOptions,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ChatStreamToolFunctionChunk {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
pub arguments: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_specific_fields: Option<JsonObject>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ChatStreamToolCallChunk {
pub id: Option<String>,
#[serde(rename = "type")]
pub tool_type: String,
pub function: ChatStreamToolFunctionChunk,
pub index: u64,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ChatStreamUsage {
pub prompt_tokens: u64,
pub completion_tokens: u64,
pub total_tokens: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prompt_tokens_details: Option<JsonObject>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub completion_tokens_details: Option<JsonObject>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ChatStreamEvent {
pub text: String,
pub tool_use: Option<ChatStreamToolCallChunk>,
pub is_finished: bool,
pub finish_reason: String,
pub usage: Option<ChatStreamUsage>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub index: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_specific_fields: Option<JsonObject>,
}
#[cfg(test)]
mod stream_contract_tests {
use super::*;
#[test]
fn request_uses_public_chat_completion_parameter_names() {
let request: ChatCompletionsStreamRequestBody = serde_json::from_value(serde_json::json!({
"model": "claude-sonnet",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 32,
"stream": true,
"tool_choice": "auto"
}))
.expect("public request shape");
assert_eq!(request.parameters.max_tokens, Some(32));
assert_eq!(request.parameters.stream, Some(true));
assert!(matches!(
request.parameters.tool_choice,
Some(ChatStreamStringOrObject::Name(ref value)) if value == "auto"
));
}
#[test]
fn event_matches_python_generic_streaming_chunk_shape() {
let event = ChatStreamEvent {
text: "hello".to_string(),
tool_use: None,
is_finished: false,
finish_reason: String::new(),
usage: None,
index: Some(0),
provider_specific_fields: None,
};
assert_eq!(
serde_json::to_value(event).expect("serializable event"),
serde_json::json!({
"text": "hello",
"tool_use": null,
"is_finished": false,
"finish_reason": "",
"usage": null,
"index": 0
})
);
}
}

View file

@ -12,5 +12,6 @@ pub mod realtime;
pub mod responses;
pub mod router;
pub mod routing_utils;
pub mod streaming;
pub use error::Error;

View file

@ -15,10 +15,13 @@ pub mod transformation;
pub mod types;
use crate::error::Error;
use crate::streaming::OpenedStream;
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
use prepare::prepare_messages_call;
use types::{AnthropicMessagesResponse, MessagesRequest};
use types::{
AnthropicMessagesResponse, MessagesRequest, MessagesStreamEvent, MessagesStreamRequest,
};
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
execute_messages_provider_call(prepare_messages_call(request)?).await
@ -28,5 +31,13 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Re
execute_messages_provider_stream(prepare_messages_call(request)?).await
}
pub async fn messages_event_stream(
_request: MessagesStreamRequest,
) -> Result<OpenedStream<MessagesStreamEvent>, Error> {
Err(Error::Unsupported(
"messages event streaming provider registration",
))
}
#[cfg(test)]
mod tests;

View file

@ -1,6 +1,19 @@
use crate::error::Error;
use crate::streaming::StreamProvider;
use super::types::{AnthropicMessagesRequest, AnthropicMessagesResponse};
use super::types::{
AnthropicMessagesRequest, AnthropicMessagesResponse, MessagesStreamEvent, MessagesStreamRequest,
};
pub trait MessagesStreamProvider:
StreamProvider<MessagesStreamRequest, MessagesStreamEvent>
{
}
impl<T> MessagesStreamProvider for T where
T: StreamProvider<MessagesStreamRequest, MessagesStreamEvent>
{
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MessagesAuthStrategy {

View file

@ -4,6 +4,7 @@ use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use super::transformation::AnthropicMessagesProviderConfig;
use crate::streaming::{JsonObject, StreamTarget, StreamTransportOptions};
pub struct MessagesRequest<'a> {
pub model: &'a str,
@ -132,3 +133,73 @@ pub struct AnthropicMessagesResponse {
#[serde(flatten)]
pub extra: Map<String, Value>,
}
pub struct MessagesStreamRequest {
pub body: AnthropicMessagesRequest,
pub target: StreamTarget,
pub transport: StreamTransportOptions,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum MessagesStreamEvent {
MessageStart {
message: AnthropicMessagesResponse,
},
ContentBlockStart {
index: u64,
content_block: JsonObject,
},
ContentBlockDelta {
index: u64,
delta: JsonObject,
},
ContentBlockStop {
index: u64,
},
MessageDelta {
delta: JsonObject,
#[serde(default, skip_serializing_if = "Option::is_none")]
usage: Option<JsonObject>,
},
MessageStop,
Ping,
Error {
error: JsonObject,
},
}
#[cfg(test)]
mod stream_contract_tests {
use super::*;
#[test]
fn message_stop_serializes_as_anthropic_event() {
assert_eq!(
serde_json::to_value(MessagesStreamEvent::MessageStop).expect("serializable event"),
serde_json::json!({"type": "message_stop"})
);
}
#[test]
fn content_delta_keeps_typed_event_fields() {
let event = MessagesStreamEvent::ContentBlockDelta {
index: 0,
delta: JsonObject(
serde_json::json!({"type": "text_delta", "text": "hello"})
.as_object()
.expect("object")
.clone(),
),
};
assert_eq!(
serde_json::to_value(event).expect("serializable event"),
serde_json::json!({
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": "hello"}
})
);
}
}

View file

@ -1,3 +1,24 @@
pub mod instrumentation;
pub mod types;
pub mod websocket;
use crate::error::Error;
use crate::streaming::OpenedStream;
use types::{ResponsesStreamEvent, ResponsesStreamRequest, ResponsesWebSocketRequest};
use websocket::TypedResponsesWebSocketSession;
pub async fn responses_stream(
_request: ResponsesStreamRequest,
) -> Result<OpenedStream<ResponsesStreamEvent>, Error> {
Err(Error::Unsupported(
"responses HTTP streaming provider registration",
))
}
pub async fn responses_websocket(
_request: ResponsesWebSocketRequest,
) -> Result<Box<dyn TypedResponsesWebSocketSession>, Error> {
Err(Error::Unsupported(
"responses WebSocket streaming provider registration",
))
}

View file

@ -1,13 +1,46 @@
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use serde_json::{Map, Value};
use crate::streaming::{JsonObject, StreamTarget, StreamTransportOptions};
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ResponsesWsEventType {
ResponseCreate,
ResponseCreated,
ResponseInProgress,
ResponseReasoningSummaryPartAdded,
ResponseReasoningSummaryTextDelta,
ResponseReasoningSummaryTextDone,
ResponseReasoningSummaryPartDone,
ResponseOutputItemAdded,
ResponseOutputTextDelta,
ResponseOutputTextAnnotationAdded,
ResponseOutputTextDone,
ResponseRefusalDelta,
ResponseRefusalDone,
ResponseFunctionCallArgumentsDelta,
ResponseFunctionCallArgumentsDone,
ResponseFileSearchCallInProgress,
ResponseFileSearchCallSearching,
ResponseFileSearchCallCompleted,
ResponseWebSearchCallInProgress,
ResponseWebSearchCallSearching,
ResponseWebSearchCallCompleted,
ResponseMcpListToolsInProgress,
ResponseMcpListToolsCompleted,
ResponseMcpListToolsFailed,
ResponseMcpCallInProgress,
ResponseMcpCallArgumentsDelta,
ResponseMcpCallArgumentsDone,
ResponseMcpCallCompleted,
ResponseMcpCallFailed,
ResponseContentPartAdded,
ResponseContentPartDone,
ResponseOutputItemDone,
ResponseCompleted,
ResponseFailed,
ResponseIncomplete,
ImageGenerationPartialImage,
Error,
Other(String),
}
@ -17,9 +50,40 @@ impl ResponsesWsEventType {
match self {
Self::ResponseCreate => "response.create",
Self::ResponseCreated => "response.created",
Self::ResponseInProgress => "response.in_progress",
Self::ResponseReasoningSummaryPartAdded => "response.reasoning_summary_part.added",
Self::ResponseReasoningSummaryTextDelta => "response.reasoning_summary_text.delta",
Self::ResponseReasoningSummaryTextDone => "response.reasoning_summary_text.done",
Self::ResponseReasoningSummaryPartDone => "response.reasoning_summary_part.done",
Self::ResponseOutputItemAdded => "response.output_item.added",
Self::ResponseOutputTextDelta => "response.output_text.delta",
Self::ResponseOutputTextAnnotationAdded => "response.output_text.annotation.added",
Self::ResponseOutputTextDone => "response.output_text.done",
Self::ResponseRefusalDelta => "response.refusal.delta",
Self::ResponseRefusalDone => "response.refusal.done",
Self::ResponseFunctionCallArgumentsDelta => "response.function_call_arguments.delta",
Self::ResponseFunctionCallArgumentsDone => "response.function_call_arguments.done",
Self::ResponseFileSearchCallInProgress => "response.file_search_call.in_progress",
Self::ResponseFileSearchCallSearching => "response.file_search_call.searching",
Self::ResponseFileSearchCallCompleted => "response.file_search_call.completed",
Self::ResponseWebSearchCallInProgress => "response.web_search_call.in_progress",
Self::ResponseWebSearchCallSearching => "response.web_search_call.searching",
Self::ResponseWebSearchCallCompleted => "response.web_search_call.completed",
Self::ResponseMcpListToolsInProgress => "response.mcp_list_tools.in_progress",
Self::ResponseMcpListToolsCompleted => "response.mcp_list_tools.completed",
Self::ResponseMcpListToolsFailed => "response.mcp_list_tools.failed",
Self::ResponseMcpCallInProgress => "response.mcp_call.in_progress",
Self::ResponseMcpCallArgumentsDelta => "response.mcp_call_arguments.delta",
Self::ResponseMcpCallArgumentsDone => "response.mcp_call_arguments.done",
Self::ResponseMcpCallCompleted => "response.mcp_call.completed",
Self::ResponseMcpCallFailed => "response.mcp_call.failed",
Self::ResponseContentPartAdded => "response.content_part.added",
Self::ResponseContentPartDone => "response.content_part.done",
Self::ResponseOutputItemDone => "response.output_item.done",
Self::ResponseCompleted => "response.completed",
Self::ResponseFailed => "response.failed",
Self::ResponseIncomplete => "response.incomplete",
Self::ImageGenerationPartialImage => "image_generation.partial_image",
Self::Error => "error",
Self::Other(value) => value,
}
@ -44,9 +108,40 @@ impl<'de> Deserialize<'de> for ResponsesWsEventType {
Ok(match value.as_str() {
"response.create" => Self::ResponseCreate,
"response.created" => Self::ResponseCreated,
"response.in_progress" => Self::ResponseInProgress,
"response.reasoning_summary_part.added" => Self::ResponseReasoningSummaryPartAdded,
"response.reasoning_summary_text.delta" => Self::ResponseReasoningSummaryTextDelta,
"response.reasoning_summary_text.done" => Self::ResponseReasoningSummaryTextDone,
"response.reasoning_summary_part.done" => Self::ResponseReasoningSummaryPartDone,
"response.output_item.added" => Self::ResponseOutputItemAdded,
"response.output_text.delta" => Self::ResponseOutputTextDelta,
"response.output_text.annotation.added" => Self::ResponseOutputTextAnnotationAdded,
"response.output_text.done" => Self::ResponseOutputTextDone,
"response.refusal.delta" => Self::ResponseRefusalDelta,
"response.refusal.done" => Self::ResponseRefusalDone,
"response.function_call_arguments.delta" => Self::ResponseFunctionCallArgumentsDelta,
"response.function_call_arguments.done" => Self::ResponseFunctionCallArgumentsDone,
"response.file_search_call.in_progress" => Self::ResponseFileSearchCallInProgress,
"response.file_search_call.searching" => Self::ResponseFileSearchCallSearching,
"response.file_search_call.completed" => Self::ResponseFileSearchCallCompleted,
"response.web_search_call.in_progress" => Self::ResponseWebSearchCallInProgress,
"response.web_search_call.searching" => Self::ResponseWebSearchCallSearching,
"response.web_search_call.completed" => Self::ResponseWebSearchCallCompleted,
"response.mcp_list_tools.in_progress" => Self::ResponseMcpListToolsInProgress,
"response.mcp_list_tools.completed" => Self::ResponseMcpListToolsCompleted,
"response.mcp_list_tools.failed" => Self::ResponseMcpListToolsFailed,
"response.mcp_call.in_progress" => Self::ResponseMcpCallInProgress,
"response.mcp_call_arguments.delta" => Self::ResponseMcpCallArgumentsDelta,
"response.mcp_call_arguments.done" => Self::ResponseMcpCallArgumentsDone,
"response.mcp_call.completed" => Self::ResponseMcpCallCompleted,
"response.mcp_call.failed" => Self::ResponseMcpCallFailed,
"response.content_part.added" => Self::ResponseContentPartAdded,
"response.content_part.done" => Self::ResponseContentPartDone,
"response.output_item.done" => Self::ResponseOutputItemDone,
"response.completed" => Self::ResponseCompleted,
"response.failed" => Self::ResponseFailed,
"response.incomplete" => Self::ResponseIncomplete,
"image_generation.partial_image" => Self::ImageGenerationPartialImage,
"error" => Self::Error,
_ => Self::Other(value),
})
@ -84,6 +179,60 @@ pub struct ResponsesWsTransformResult {
pub events: Vec<ResponsesWsEvent>,
}
pub type ResponsesStreamEvent = ResponsesWsEvent;
pub type ResponseCommand = ResponsesWsEvent;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ResponsesInput {
Text(String),
Items(Vec<JsonObject>),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ResponsesToolChoice {
Name(String),
Definition(JsonObject),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ResponsesStreamRequestBody {
pub model: String,
pub input: ResponsesInput,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub instructions: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_output_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub previous_response_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stream: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub store: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<JsonObject>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<ResponsesToolChoice>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub metadata: Option<JsonObject>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub include: Option<Vec<String>>,
#[serde(flatten)]
pub extra: Map<String, Value>,
}
pub struct ResponsesStreamRequest {
pub body: ResponsesStreamRequestBody,
pub target: StreamTarget,
pub transport: StreamTransportOptions,
}
pub struct ResponsesWebSocketRequest {
pub target: StreamTarget,
pub transport: StreamTransportOptions,
}
impl ResponsesWsTransformResult {
pub fn passthrough(event: ResponsesWsEvent) -> Self {
Self {
@ -127,11 +276,14 @@ mod tests {
let known: ResponsesWsEventType =
serde_json::from_str("\"response.completed\"").expect("valid event type");
assert_eq!(known, ResponsesWsEventType::ResponseCompleted);
let unknown: ResponsesWsEventType =
let output_delta: ResponsesWsEventType =
serde_json::from_str("\"response.output_text.delta\"").expect("valid event type");
assert_eq!(output_delta, ResponsesWsEventType::ResponseOutputTextDelta);
let unknown: ResponsesWsEventType =
serde_json::from_str("\"response.future_event\"").expect("valid event type");
assert_eq!(
unknown,
ResponsesWsEventType::Other("response.output_text.delta".to_string())
ResponsesWsEventType::Other("response.future_event".to_string())
);
}
@ -163,4 +315,42 @@ mod tests {
assert_eq!(flat.model(), Some("gpt-5"));
assert_eq!(nested.model(), Some("gpt-5-mini"));
}
#[test]
fn stream_request_deserializes_the_public_responses_shape() {
let request: ResponsesStreamRequestBody = serde_json::from_value(serde_json::json!({
"model": "gpt-5",
"input": "hello",
"stream": true,
"max_output_tokens": 32
}))
.expect("public request shape");
assert_eq!(request.model, "gpt-5");
assert_eq!(request.stream, Some(true));
assert!(matches!(request.input, ResponsesInput::Text(ref text) if text == "hello"));
}
#[test]
fn unknown_stream_events_round_trip_for_forward_compatibility() {
let event: ResponsesStreamEvent = serde_json::from_value(serde_json::json!({
"type": "response.future_event",
"sequence_number": 7,
"future_field": "value"
}))
.expect("unknown event");
assert_eq!(
event.event_type,
ResponsesWsEventType::Other("response.future_event".to_string())
);
assert_eq!(
serde_json::to_value(event).expect("serializable event"),
serde_json::json!({
"type": "response.future_event",
"sequence_number": 7,
"future_field": "value"
})
);
}
}

View file

@ -1,6 +1,19 @@
use crate::Error;
use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH};
use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult};
use futures_util::future::BoxFuture;
use crate::responses::types::{
ResponseCommand, ResponsesStreamEvent, ResponsesWsEvent, ResponsesWsEventType,
ResponsesWsTransformResult,
};
pub trait TypedResponsesWebSocketSession: Send + Sync {
fn send(&self, command: ResponseCommand) -> BoxFuture<'_, Result<(), Error>>;
fn recv(&self) -> BoxFuture<'_, Result<Option<ResponsesStreamEvent>, Error>>;
fn close(&self) -> BoxFuture<'_, Result<(), Error>>;
}
pub trait ResponsesWebSocketProviderConfig: Sync {
fn supports_native_websocket(&self) -> bool {

View file

@ -0,0 +1,435 @@
use std::collections::VecDeque;
use std::marker::PhantomData;
use std::pin::Pin;
use std::time::Duration;
use crate::error::Error;
use futures_util::future::BoxFuture;
use futures_util::{Stream, StreamExt, stream};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
pub type EventStream<E> = Pin<Box<dyn Stream<Item = Result<E, Error>> + Send + 'static>>;
pub type ProviderChunkStream =
Pin<Box<dyn Stream<Item = Result<ProviderStreamChunk, Error>> + Send + 'static>>;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum StreamTransport {
Http,
WebSocket,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum StreamProviderId {
Anthropic,
AzureAi,
BedrockConverse,
OpenAi,
}
impl TryFrom<&str> for StreamProviderId {
type Error = Error;
fn try_from(value: &str) -> Result<Self, Self::Error> {
match value {
"anthropic" => Ok(Self::Anthropic),
"azure_ai" => Ok(Self::AzureAi),
"bedrock" | "bedrock_converse" => Ok(Self::BedrockConverse),
"openai" => Ok(Self::OpenAi),
_ => Err(Error::InvalidProvider(value.to_string())),
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(transparent)]
pub struct JsonObject(pub Map<String, Value>);
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct Header {
pub name: String,
pub value: String,
}
#[derive(Clone, Default, PartialEq)]
pub struct ProviderCredentials {
api_key: Option<String>,
aws_access_key_id: Option<String>,
aws_secret_access_key: Option<String>,
aws_session_token: Option<String>,
}
impl ProviderCredentials {
pub fn new(
api_key: Option<String>,
aws_access_key_id: Option<String>,
aws_secret_access_key: Option<String>,
aws_session_token: Option<String>,
) -> Self {
Self {
api_key,
aws_access_key_id,
aws_secret_access_key,
aws_session_token,
}
}
pub fn api_key(&self) -> Option<&str> {
self.api_key.as_deref()
}
pub fn aws_access_key_id(&self) -> Option<&str> {
self.aws_access_key_id.as_deref()
}
pub fn aws_secret_access_key(&self) -> Option<&str> {
self.aws_secret_access_key.as_deref()
}
pub fn aws_session_token(&self) -> Option<&str> {
self.aws_session_token.as_deref()
}
}
/// ```compile_fail
/// fn assert_serialize<T: serde::Serialize>() {}
/// assert_serialize::<litellm_core::streaming::ProviderCredentials>();
/// assert_serialize::<litellm_core::streaming::StreamTarget>();
/// assert_serialize::<litellm_core::streaming::StreamTransportOptions>();
/// ```
///
/// ```compile_fail
/// use litellm_core::streaming::{JsonObject, ProviderCredentials, StreamProviderId, StreamTarget};
/// let mut target = StreamTarget::new(
/// StreamProviderId::OpenAi,
/// ProviderCredentials::default(),
/// None,
/// );
/// target.metadata = JsonObject::default();
/// ```
#[derive(Clone, PartialEq)]
pub struct StreamTarget {
provider: StreamProviderId,
credentials: ProviderCredentials,
api_base: Option<String>,
}
impl StreamTarget {
pub fn new(
provider: StreamProviderId,
credentials: ProviderCredentials,
api_base: Option<String>,
) -> Self {
Self {
provider,
credentials,
api_base,
}
}
pub fn provider(&self) -> StreamProviderId {
self.provider
}
pub fn credentials(&self) -> &ProviderCredentials {
&self.credentials
}
pub fn api_base(&self) -> Option<&str> {
self.api_base.as_deref()
}
}
#[derive(Clone, Default, PartialEq)]
pub struct StreamTransportOptions {
forwarded_headers: Vec<Header>,
timeout: Option<Duration>,
}
impl StreamTransportOptions {
pub fn new(forwarded_headers: Vec<Header>, timeout: Option<Duration>) -> Self {
Self {
forwarded_headers,
timeout,
}
}
pub fn forwarded_headers(&self) -> &[Header] {
&self.forwarded_headers
}
pub fn timeout(&self) -> Option<Duration> {
self.timeout
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct StreamMetadata {
pub status_code: u16,
pub provider: StreamProviderId,
pub transport: StreamTransport,
pub response_headers: Vec<Header>,
}
pub struct OpenedStream<E> {
pub metadata: StreamMetadata,
pub events: EventStream<E>,
}
pub struct OpenedWireStream {
pub metadata: StreamMetadata,
pub chunks: ProviderChunkStream,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ProviderStreamChunk(Vec<u8>);
impl ProviderStreamChunk {
pub fn new(bytes: impl Into<Vec<u8>>) -> Self {
Self(bytes.into())
}
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
}
pub trait StreamDecoder: Send + 'static {
type WireEvent: Send + 'static;
fn push(&mut self, chunk: ProviderStreamChunk) -> Result<Vec<Self::WireEvent>, Error>;
fn finish(&mut self) -> Result<Vec<Self::WireEvent>, Error> {
Ok(Vec::new())
}
}
pub trait StreamProvider<R, E>: Send + Sync + 'static {
type PreparedRequest: Send + 'static;
type WireEvent: Send + 'static;
type Decoder: StreamDecoder<WireEvent = Self::WireEvent>;
fn transform_request(&self, request: R) -> Result<Self::PreparedRequest, Error>;
fn call(
&'static self,
request: Self::PreparedRequest,
) -> BoxFuture<'static, Result<OpenedWireStream, Error>>;
fn decoder(&self) -> Self::Decoder;
fn normalize(&self, event: Self::WireEvent) -> Result<Vec<E>, Error>;
}
struct PipelineState<P: 'static, D, E, R>
where
D: StreamDecoder,
{
provider: &'static P,
decoder: D,
chunks: ProviderChunkStream,
pending: VecDeque<Result<E, Error>>,
finished: bool,
request: PhantomData<fn(R)>,
}
pub async fn open_provider_stream<P, R, E>(
provider: &'static P,
request: R,
) -> Result<OpenedStream<E>, Error>
where
P: StreamProvider<R, E>,
R: Send + 'static,
E: Send + 'static,
{
let prepared = provider.transform_request(request)?;
let opened = provider.call(prepared).await?;
let state = PipelineState {
provider,
decoder: provider.decoder(),
chunks: opened.chunks,
pending: VecDeque::new(),
finished: false,
request: PhantomData,
};
let events = stream::unfold(state, |mut state| async move {
loop {
if let Some(event) = state.pending.pop_front() {
return Some((event, state));
}
if state.finished {
return None;
}
match state.chunks.next().await {
Some(Ok(chunk)) => match state.decoder.push(chunk) {
Ok(events) => queue_normalized(&mut state, events),
Err(error) => {
state.finished = true;
return Some((Err(error), state));
}
},
Some(Err(error)) => {
state.finished = true;
return Some((Err(error), state));
}
None => {
state.finished = true;
match state.decoder.finish() {
Ok(events) => queue_normalized(&mut state, events),
Err(error) => return Some((Err(error), state)),
}
}
}
}
});
Ok(OpenedStream {
metadata: opened.metadata,
events: Box::pin(events),
})
}
fn queue_normalized<P, D, E, R>(state: &mut PipelineState<P, D, E, R>, events: Vec<D::WireEvent>)
where
P: StreamProvider<R, E, Decoder = D>,
D: StreamDecoder<WireEvent = <P as StreamProvider<R, E>>::WireEvent>,
{
for event in events {
match state.provider.normalize(event) {
Ok(normalized) => state.pending.extend(normalized.into_iter().map(Ok)),
Err(error) => {
state.pending.push_back(Err(error));
state.finished = true;
return;
}
}
}
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
use futures_util::future::FutureExt;
use super::*;
struct FakeRequest;
struct PreparedRequest;
struct FakeProvider {
calls: Arc<Mutex<Vec<&'static str>>>,
}
struct FakeDecoder {
calls: Arc<Mutex<Vec<&'static str>>>,
pending: String,
}
impl StreamDecoder for FakeDecoder {
type WireEvent = String;
fn push(&mut self, chunk: ProviderStreamChunk) -> Result<Vec<Self::WireEvent>, Error> {
self.calls.lock().expect("call log").push("decode");
self.pending
.push_str(std::str::from_utf8(chunk.as_bytes()).expect("test utf-8"));
let mut parts = self
.pending
.split('|')
.map(str::to_string)
.collect::<Vec<_>>();
self.pending = parts.pop().expect("split always returns one item");
Ok(parts)
}
fn finish(&mut self) -> Result<Vec<Self::WireEvent>, Error> {
if self.pending.is_empty() {
return Ok(Vec::new());
}
Ok(vec![std::mem::take(&mut self.pending)])
}
}
impl StreamProvider<FakeRequest, String> for FakeProvider {
type PreparedRequest = PreparedRequest;
type WireEvent = String;
type Decoder = FakeDecoder;
fn transform_request(&self, _request: FakeRequest) -> Result<Self::PreparedRequest, Error> {
self.calls.lock().expect("call log").push("transform");
Ok(PreparedRequest)
}
fn call(
&'static self,
_request: Self::PreparedRequest,
) -> BoxFuture<'static, Result<OpenedWireStream, Error>> {
self.calls.lock().expect("call log").push("call");
async move {
Ok(OpenedWireStream {
metadata: StreamMetadata {
status_code: 200,
provider: StreamProviderId::Anthropic,
transport: StreamTransport::Http,
response_headers: vec![Header {
name: "x-test".to_string(),
value: "ready".to_string(),
}],
},
chunks: Box::pin(stream::iter([
Ok(ProviderStreamChunk::new(b"one|tw".to_vec())),
Ok(ProviderStreamChunk::new(b"o|three".to_vec())),
])),
})
}
.boxed()
}
fn decoder(&self) -> Self::Decoder {
FakeDecoder {
calls: self.calls.clone(),
pending: String::new(),
}
}
fn normalize(&self, event: Self::WireEvent) -> Result<Vec<String>, Error> {
self.calls.lock().expect("call log").push("normalize");
Ok(vec![event.to_uppercase()])
}
}
#[tokio::test]
async fn fake_provider_proves_pipeline_order_and_fragmentation() {
let calls = Arc::new(Mutex::new(Vec::new()));
let provider = Box::leak(Box::new(FakeProvider {
calls: calls.clone(),
}));
let mut opened = open_provider_stream(provider, FakeRequest)
.await
.expect("stream opens");
let events = opened
.events
.by_ref()
.collect::<Vec<_>>()
.await
.into_iter()
.collect::<Result<Vec<_>, _>>()
.expect("events normalize");
assert_eq!(events, ["ONE", "TWO", "THREE"]);
assert_eq!(opened.metadata.response_headers[0].name, "x-test");
assert_eq!(
*calls.lock().expect("call log"),
[
"transform",
"call",
"decode",
"normalize",
"decode",
"normalize",
"normalize",
]
);
}
}

View file

@ -116,6 +116,16 @@ mod tests {
"chat_completions_decline",
"chat_completions",
"achat_completions",
"ChatCompletionsEventStream",
"MessagesEventStream",
"ResponsesEventStream",
"ResponsesWebSocketSession",
"chat_completions_stream",
"achat_completions_stream",
"messages_stream",
"amessages_stream",
"responses_stream",
"aresponses_stream",
"ResponsesWebSocketConnection",
"gil_stats",
];

View file

@ -78,12 +78,15 @@ mod audio_transcription;
mod chat_completions;
mod messages;
mod ocr;
mod receiver;
mod streaming;
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
ocr::register(module)?;
audio_transcription::register(module)?;
messages::register(module)?;
chat_completions::register(module)
chat_completions::register(module)?;
streaming::register(module)
}
#[cfg(test)]

View file

@ -0,0 +1,206 @@
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use futures_util::{Stream, StreamExt, pin_mut};
use litellm_core::Error;
use tokio::sync::{Mutex, mpsc};
use tokio::task::JoinHandle;
const BRIDGE_CHANNEL_CAPACITY: usize = 1;
struct ReceiverState<T> {
receiver: Mutex<mpsc::Receiver<Result<T, Error>>>,
reading: AtomicBool,
closed: AtomicBool,
producer: std::sync::Mutex<Option<JoinHandle<()>>>,
}
impl<T> Drop for ReceiverState<T> {
fn drop(&mut self) {
if let Ok(producer) = self.producer.get_mut()
&& let Some(producer) = producer.take()
{
producer.abort();
}
}
}
#[derive(Clone)]
pub(super) struct BridgeReceiver<T> {
state: Arc<ReceiverState<T>>,
}
struct ReadGuard<'a>(&'a AtomicBool);
impl Drop for ReadGuard<'_> {
fn drop(&mut self) {
self.0.store(false, Ordering::Release);
}
}
impl<T: Send + 'static> BridgeReceiver<T> {
pub(super) fn from_stream<S>(stream: S) -> Self
where
S: Stream<Item = Result<T, Error>> + Send + 'static,
{
let (sender, receiver) = mpsc::channel(BRIDGE_CHANNEL_CAPACITY);
let producer = pyo3_async_runtimes::tokio::get_runtime().spawn(async move {
pin_mut!(stream);
while let Some(item) = stream.next().await {
let terminal = item.is_err();
if sender.send(item).await.is_err() || terminal {
return;
}
}
});
Self {
state: Arc::new(ReceiverState {
receiver: Mutex::new(receiver),
reading: AtomicBool::new(false),
closed: AtomicBool::new(false),
producer: std::sync::Mutex::new(Some(producer)),
}),
}
}
pub(super) async fn next(&self) -> Result<Option<T>, Error> {
if self.state.closed.load(Ordering::Acquire) {
return Ok(None);
}
if self.state.reading.swap(true, Ordering::AcqRel) {
return Err(Error::InvalidRequest(
"native stream does not support concurrent reads".to_string(),
));
}
let _guard = ReadGuard(&self.state.reading);
match self.state.receiver.lock().await.recv().await {
Some(Ok(item)) => Ok(Some(item)),
Some(Err(error)) => Err(error),
None => Ok(None),
}
}
pub(super) fn close(&self) {
if self.state.closed.swap(true, Ordering::AcqRel) {
return;
}
if let Ok(mut producer) = self.state.producer.lock()
&& let Some(producer) = producer.take()
{
producer.abort();
}
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use futures_util::stream;
use super::*;
struct DropFlag(Arc<AtomicBool>);
impl Drop for DropFlag {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
#[tokio::test]
async fn receiver_preserves_items_and_terminal_error() {
let receiver = BridgeReceiver::from_stream(stream::iter([
Ok(vec![1_u8]),
Err(Error::Network("broken".to_string())),
Ok(vec![2_u8]),
]));
assert_eq!(receiver.next().await.expect("first item"), Some(vec![1]));
assert!(matches!(
receiver.next().await,
Err(Error::Network(message)) if message == "broken"
));
assert_eq!(receiver.next().await.expect("closed after error"), None);
}
#[tokio::test]
async fn close_unblocks_a_pending_read() {
let receiver = BridgeReceiver::<Vec<u8>>::from_stream(stream::pending());
let pending = {
let receiver = receiver.clone();
tokio::spawn(async move { receiver.next().await })
};
tokio::time::sleep(Duration::from_millis(10)).await;
receiver.close();
assert_eq!(
tokio::time::timeout(Duration::from_secs(1), pending)
.await
.expect("read should unblock")
.expect("task should finish")
.expect("close is clean"),
None
);
}
#[tokio::test]
async fn capacity_one_stops_the_producer_from_draining_the_source() {
let polls = Arc::new(AtomicUsize::new(0));
let source_polls = polls.clone();
let source = stream::poll_fn(move |_| {
let item = source_polls.fetch_add(1, Ordering::SeqCst);
std::task::Poll::Ready(Some(Ok(item)))
});
let receiver = BridgeReceiver::from_stream(source);
tokio::time::sleep(Duration::from_millis(20)).await;
assert_eq!(polls.load(Ordering::SeqCst), 2);
assert_eq!(receiver.next().await.expect("first item"), Some(0));
tokio::time::sleep(Duration::from_millis(20)).await;
assert_eq!(polls.load(Ordering::SeqCst), 3);
}
#[tokio::test]
async fn concurrent_reads_are_rejected() {
let receiver = BridgeReceiver::<Vec<u8>>::from_stream(stream::pending());
let pending = {
let receiver = receiver.clone();
tokio::spawn(async move { receiver.next().await })
};
tokio::time::sleep(Duration::from_millis(10)).await;
assert!(matches!(
receiver.next().await,
Err(Error::InvalidRequest(message))
if message == "native stream does not support concurrent reads"
));
receiver.close();
pending
.await
.expect("pending read task")
.expect("clean close");
}
#[tokio::test]
async fn dropping_the_last_receiver_cancels_the_producer() {
let dropped = Arc::new(AtomicBool::new(false));
let producer_dropped = dropped.clone();
let source = stream::once(async move {
let _flag = DropFlag(producer_dropped);
std::future::pending::<Result<Vec<u8>, Error>>().await
});
let receiver = BridgeReceiver::from_stream(source);
tokio::time::sleep(Duration::from_millis(10)).await;
drop(receiver);
tokio::time::timeout(Duration::from_secs(1), async {
while !dropped.load(Ordering::SeqCst) {
tokio::task::yield_now().await;
}
})
.await
.expect("producer future should be dropped");
}
}

View file

@ -65,6 +65,45 @@ where
})
}
pub(super) fn run_sync_with<T, F, C>(py: Python<'_>, future: F, convert: C) -> PyResult<Py<PyAny>>
where
T: Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
C: FnOnce(Python<'_>, T) -> PyResult<Py<PyAny>>,
{
if Handle::try_current().is_ok() {
return Err(PyRuntimeError::new_err(
"synchronous native routes cannot run from a Tokio context; use the async route",
));
}
let result = release_gil(py, move || {
pyo3_async_runtimes::tokio::get_runtime().block_on(wait_for_sync_result(future))
})?;
let result = map_core_result(result, crate::errors::fallback_route_error_to_pyerr)?;
std::panic::catch_unwind(AssertUnwindSafe(|| convert(py, result))).map_err(panic_to_pyerr)?
}
pub(super) fn run_async_with<T, F, C>(
py: Python<'_>,
future: F,
convert: C,
) -> PyResult<Bound<'_, PyAny>>
where
T: Send + 'static,
F: Future<Output = Result<T, Error>> + Send + 'static,
C: FnOnce(Python<'_>, T) -> PyResult<Py<PyAny>> + Send + 'static,
{
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let result = catch_route_panic(future).await?;
let result = map_core_result(result, crate::errors::fallback_route_error_to_pyerr)?;
Python::attach(|py| {
std::panic::catch_unwind(AssertUnwindSafe(|| convert(py, result)))
.map_err(panic_to_pyerr)?
})
})
}
fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -> PyResult<T> {
match result {
Ok(value) => Ok(value),

View file

@ -0,0 +1,541 @@
use std::sync::Arc;
use litellm_core::chat_completions::chat_completions_stream as run_chat_completions_stream;
use litellm_core::chat_completions::types::{
ChatCompletionsStreamRequest, ChatCompletionsStreamRequestBody, ChatStreamEvent,
};
use litellm_core::messages::messages_event_stream as run_messages_stream;
use litellm_core::messages::types::{
AnthropicMessagesRequest, MessagesStreamEvent, MessagesStreamRequest,
};
use litellm_core::responses::responses_stream as run_responses_stream;
use litellm_core::responses::responses_websocket as run_responses_websocket;
use litellm_core::responses::types::{
ResponseCommand, ResponsesStreamEvent, ResponsesStreamRequest, ResponsesStreamRequestBody,
ResponsesWebSocketRequest,
};
use litellm_core::responses::websocket::TypedResponsesWebSocketSession;
use litellm_core::streaming::{
JsonObject, OpenedStream, ProviderCredentials, StreamMetadata, StreamProviderId, StreamTarget,
StreamTransportOptions,
};
use litellm_python_interop::{from_py, to_py};
use pyo3::exceptions::{PyStopAsyncIteration, PyStopIteration};
use pyo3::prelude::*;
use pyo3::types::{PyModule, PyType};
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use crate::marshal::{marshal_headers, optional_timeout};
use crate::routes::receiver::BridgeReceiver;
use crate::routes::runtime::{run_async, run_async_with, run_sync_with};
struct TypedEventReceiver<E> {
metadata: StreamMetadata,
receiver: BridgeReceiver<E>,
}
impl<E> TypedEventReceiver<E>
where
E: Send + 'static,
{
fn from_opened(opened: OpenedStream<E>) -> Self {
Self {
metadata: opened.metadata,
receiver: BridgeReceiver::from_stream(opened.events),
}
}
}
fn next_event<E>(
py: Python<'_>,
receiver: BridgeReceiver<E>,
stop_iteration: bool,
) -> PyResult<Py<PyAny>>
where
E: Serialize + Send + 'static,
{
run_sync_with(
py,
async move { receiver.next().await },
move |py, event| match event {
Some(event) => to_py(py, &event),
None if stop_iteration => Err(PyStopIteration::new_err(())),
None => Ok(py.None()),
},
)
}
fn anext_event<E>(
py: Python<'_>,
receiver: BridgeReceiver<E>,
stop_iteration: bool,
) -> PyResult<Bound<'_, PyAny>>
where
E: Serialize + Send + 'static,
{
run_async_with(
py,
async move { receiver.next().await },
move |py, event| match event {
Some(event) => to_py(py, &event),
None if stop_iteration => Err(PyStopAsyncIteration::new_err(())),
None => Ok(py.None()),
},
)
}
macro_rules! event_stream_class {
($class:ident, $event:ty) => {
#[pyclass]
struct $class {
inner: TypedEventReceiver<$event>,
}
impl From<OpenedStream<$event>> for $class {
fn from(opened: OpenedStream<$event>) -> Self {
Self {
inner: TypedEventReceiver::from_opened(opened),
}
}
}
#[pymethods]
impl $class {
#[getter]
fn metadata(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
to_py(py, &self.inner.metadata)
}
fn next_event(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
next_event(py, self.inner.receiver.clone(), false)
}
fn anext_event<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
anext_event(py, self.inner.receiver.clone(), false)
}
fn __iter__(slf: PyRef<'_, Self>) -> PyRef<'_, Self> {
slf
}
fn __next__(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
next_event(py, self.inner.receiver.clone(), true)
}
fn __aiter__(slf: PyRef<'_, Self>) -> PyRef<'_, Self> {
slf
}
fn __anext__<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
anext_event(py, self.inner.receiver.clone(), true)
}
fn close(&self) {
self.inner.receiver.close();
}
fn aclose<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let receiver = self.inner.receiver.clone();
run_async(
py,
async move {
receiver.close();
Ok(())
},
crate::errors::fallback_route_error_to_pyerr,
)
}
}
};
}
event_stream_class!(ChatCompletionsEventStream, ChatStreamEvent);
event_stream_class!(MessagesEventStream, MessagesStreamEvent);
event_stream_class!(ResponsesEventStream, ResponsesStreamEvent);
#[derive(Default, Deserialize)]
struct PythonProviderCredentials {
api_key: Option<String>,
aws_access_key_id: Option<String>,
aws_secret_access_key: Option<String>,
aws_session_token: Option<String>,
}
struct PythonStreamTarget {
provider: String,
credentials: Option<Py<PyAny>>,
api_base: Option<String>,
}
struct PythonStreamTransport {
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
}
fn parse_call<B>(
py: Python<'_>,
request: Py<PyAny>,
target: PythonStreamTarget,
transport: PythonStreamTransport,
) -> PyResult<(B, StreamTarget, StreamTransportOptions)>
where
B: DeserializeOwned,
{
let body = from_py(request.bind(py))?;
let credentials = target
.credentials
.map(|value| from_py::<PythonProviderCredentials>(value.bind(py)))
.transpose()?
.unwrap_or_default();
let provider = StreamProviderId::try_from(target.provider.as_str())
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?;
let target = StreamTarget::new(
provider,
ProviderCredentials::new(
credentials.api_key,
credentials.aws_access_key_id,
credentials.aws_secret_access_key,
credentials.aws_session_token,
),
target.api_base,
);
let forwarded_headers = marshal_headers(py, transport.extra_headers)?
.into_iter()
.map(|(name, value)| litellm_core::streaming::Header { name, value })
.collect();
let transport = StreamTransportOptions::new(
forwarded_headers,
optional_timeout(transport.timeout_seconds),
);
Ok((body, target, transport))
}
#[pyfunction]
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn chat_completions_stream(
py: Python<'_>,
request: Py<PyAny>,
provider: String,
credentials: Option<Py<PyAny>>,
api_base: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let (body, target, transport) = parse_call::<ChatCompletionsStreamRequestBody>(
py,
request,
PythonStreamTarget {
provider,
credentials,
api_base,
},
PythonStreamTransport {
extra_headers,
timeout_seconds,
},
)?;
run_sync_with(
py,
async move {
run_chat_completions_stream(ChatCompletionsStreamRequest {
body,
target,
transport,
})
.await
},
|py, opened| Ok(Py::new(py, ChatCompletionsEventStream::from(opened))?.into_any()),
)
}
#[pyfunction]
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn achat_completions_stream(
py: Python<'_>,
request: Py<PyAny>,
provider: String,
credentials: Option<Py<PyAny>>,
api_base: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let (body, target, transport) = parse_call::<ChatCompletionsStreamRequestBody>(
py,
request,
PythonStreamTarget {
provider,
credentials,
api_base,
},
PythonStreamTransport {
extra_headers,
timeout_seconds,
},
)?;
run_async_with(
py,
async move {
run_chat_completions_stream(ChatCompletionsStreamRequest {
body,
target,
transport,
})
.await
},
|py, opened| Ok(Py::new(py, ChatCompletionsEventStream::from(opened))?.into_any()),
)
}
#[pyfunction]
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn messages_stream(
py: Python<'_>,
request: Py<PyAny>,
provider: String,
credentials: Option<Py<PyAny>>,
api_base: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let (body, target, transport) = parse_call::<AnthropicMessagesRequest>(
py,
request,
PythonStreamTarget {
provider,
credentials,
api_base,
},
PythonStreamTransport {
extra_headers,
timeout_seconds,
},
)?;
run_sync_with(
py,
async move {
run_messages_stream(MessagesStreamRequest {
body,
target,
transport,
})
.await
},
|py, opened| Ok(Py::new(py, MessagesEventStream::from(opened))?.into_any()),
)
}
#[pyfunction]
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn amessages_stream(
py: Python<'_>,
request: Py<PyAny>,
provider: String,
credentials: Option<Py<PyAny>>,
api_base: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let (body, target, transport) = parse_call::<AnthropicMessagesRequest>(
py,
request,
PythonStreamTarget {
provider,
credentials,
api_base,
},
PythonStreamTransport {
extra_headers,
timeout_seconds,
},
)?;
run_async_with(
py,
async move {
run_messages_stream(MessagesStreamRequest {
body,
target,
transport,
})
.await
},
|py, opened| Ok(Py::new(py, MessagesEventStream::from(opened))?.into_any()),
)
}
#[pyfunction]
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn responses_stream(
py: Python<'_>,
request: Py<PyAny>,
provider: String,
credentials: Option<Py<PyAny>>,
api_base: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let (body, target, transport) = parse_call::<ResponsesStreamRequestBody>(
py,
request,
PythonStreamTarget {
provider,
credentials,
api_base,
},
PythonStreamTransport {
extra_headers,
timeout_seconds,
},
)?;
run_sync_with(
py,
async move {
run_responses_stream(ResponsesStreamRequest {
body,
target,
transport,
})
.await
},
|py, opened| Ok(Py::new(py, ResponsesEventStream::from(opened))?.into_any()),
)
}
#[pyfunction]
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn aresponses_stream(
py: Python<'_>,
request: Py<PyAny>,
provider: String,
credentials: Option<Py<PyAny>>,
api_base: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let (body, target, transport) = parse_call::<ResponsesStreamRequestBody>(
py,
request,
PythonStreamTarget {
provider,
credentials,
api_base,
},
PythonStreamTransport {
extra_headers,
timeout_seconds,
},
)?;
run_async_with(
py,
async move {
run_responses_stream(ResponsesStreamRequest {
body,
target,
transport,
})
.await
},
|py, opened| Ok(Py::new(py, ResponsesEventStream::from(opened))?.into_any()),
)
}
#[pyclass]
struct ResponsesWebSocketSession {
session: Arc<dyn TypedResponsesWebSocketSession>,
}
#[pymethods]
impl ResponsesWebSocketSession {
#[classmethod]
#[pyo3(signature = (provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn connect<'py>(
_cls: &Bound<'py, PyType>,
py: Python<'py>,
provider: String,
credentials: Option<Py<PyAny>>,
api_base: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'py, PyAny>> {
let empty_request = pyo3::types::PyDict::new(py).unbind().into_any();
let (_, target, transport) = parse_call::<JsonObject>(
py,
empty_request,
PythonStreamTarget {
provider,
credentials,
api_base,
},
PythonStreamTransport {
extra_headers,
timeout_seconds,
},
)?;
run_async_with(
py,
async move {
run_responses_websocket(ResponsesWebSocketRequest { target, transport }).await
},
|py, session| {
Ok(Py::new(
py,
ResponsesWebSocketSession {
session: Arc::from(session),
},
)?
.into_any())
},
)
}
fn send_event<'py>(&self, py: Python<'py>, command: Py<PyAny>) -> PyResult<Bound<'py, PyAny>> {
let command: ResponseCommand = from_py(command.bind(py))?;
let session = self.session.clone();
run_async(
py,
async move { session.send(command).await },
crate::errors::fallback_route_error_to_pyerr,
)
}
fn recv_event<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let session = self.session.clone();
run_async_with(
py,
async move { session.recv().await },
|py, event| match event {
Some(event) => to_py(py, &event),
None => Ok(py.None()),
},
)
}
fn close<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let session = self.session.clone();
run_async(
py,
async move { session.close().await },
crate::errors::fallback_route_error_to_pyerr,
)
}
}
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_class::<ChatCompletionsEventStream>()?;
module.add_class::<MessagesEventStream>()?;
module.add_class::<ResponsesEventStream>()?;
module.add_class::<ResponsesWebSocketSession>()?;
module.add_function(wrap_pyfunction!(chat_completions_stream, module)?)?;
module.add_function(wrap_pyfunction!(achat_completions_stream, module)?)?;
module.add_function(wrap_pyfunction!(messages_stream, module)?)?;
module.add_function(wrap_pyfunction!(amessages_stream, module)?)?;
module.add_function(wrap_pyfunction!(responses_stream, module)?)?;
module.add_function(wrap_pyfunction!(aresponses_stream, module)?)
}

View file

@ -519,7 +519,6 @@ class AnthropicChatCompletion(BaseLLM):
stream=stream,
_is_function_call=_is_function_call,
json_mode=json_mode,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=headers,
timeout=timeout,

View file

@ -6506,7 +6506,9 @@ class BaseLLMHTTPHandler:
from litellm.rust_bridge import responses_websocket as rust_responses_websocket
rust_backend: Final = await rust_responses_websocket.connect(
url=ws_url,
provider="openai",
api_key=api_key,
api_base=api_base,
headers={str(key): str(value) for key, value in headers.items()},
timeout=timeout,
)

View file

@ -109,7 +109,7 @@ def use_litellm_rust(
if configuring_responses_websocket:
from litellm.rust_bridge.responses_websocket import set_rust_responses_websocket
set_rust_responses_websocket(connection=responses_websocket)
set_rust_responses_websocket(connection=responses_websocket if enabled else None)
def rust_ocr_enabled() -> bool:

View file

@ -1,21 +1,25 @@
"""Thin Python wrapper for the native Rust Responses WebSocket bridge."""
from __future__ import annotations
import json
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Protocol
import httpx
from pydantic import TypeAdapter
from websockets.exceptions import ConnectionClosedOK
from litellm.rust_bridge.loader import get_native_bridge
from litellm.rust_bridge.timeouts import timeout_to_seconds
_EVENT_ADAPTER: Final = TypeAdapter(Mapping[str, object])
class RustResponsesWebSocket(Protocol):
async def send_text(self, text: str) -> None: ...
async def send_event(self, event: Mapping[str, object]) -> None: ...
async def recv_text(self) -> str | None: ...
async def recv_event(self) -> Mapping[str, object] | None: ...
async def close(self) -> None: ...
@ -24,8 +28,10 @@ class RustResponsesWebSocketConnection(Protocol):
@classmethod
async def connect(
cls,
url: str,
headers: dict[str, str],
provider: str,
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout_seconds: float | None,
) -> RustResponsesWebSocket: ...
@ -34,49 +40,50 @@ class _Unset:
pass
_UNSET: Final[_Unset] = _Unset()
_UNSET: Final = _Unset()
@dataclass(slots=True)
class _RustResponsesWebSocketState:
connection: RustResponsesWebSocketConnection | None = None
connection: type[RustResponsesWebSocketConnection] | None = None
_STATE: Final[_RustResponsesWebSocketState] = _RustResponsesWebSocketState()
_STATE: Final = _RustResponsesWebSocketState()
def set_rust_responses_websocket(
*,
connection: RustResponsesWebSocketConnection | None | _Unset = _UNSET,
connection: type[RustResponsesWebSocketConnection] | None | _Unset = _UNSET,
) -> None:
if not isinstance(connection, _Unset):
_STATE.connection = connection
def load_rust_responses_websocket() -> RustResponsesWebSocketConnection | None:
def load_rust_responses_websocket() -> type[RustResponsesWebSocketConnection] | None:
if _STATE.connection is not None:
return _STATE.connection
native_bridge: Final = get_native_bridge()
if native_bridge is None:
return None
connection_type: Final[RustResponsesWebSocketConnection | None] = getattr(
native_bridge, "ResponsesWebSocketConnection", None
connection_type: Final[type[RustResponsesWebSocketConnection] | None] = getattr(
native_bridge, "ResponsesWebSocketSession", None
)
return connection_type
class _ConnectionAdapter:
def __init__(self, connection: RustResponsesWebSocket):
self._connection: Final[RustResponsesWebSocket] = connection
self._connection: Final = connection
async def send(self, text: str) -> None:
await self._connection.send_text(text)
event: Final = _EVENT_ADAPTER.validate_json(text)
await self._connection.send_event(event)
async def recv(self) -> str:
message: Final = await self._connection.recv_text()
if message is None:
event: Final = await self._connection.recv_event()
if event is None:
raise ConnectionClosedOK(None, None)
return message
return json.dumps(dict(event), separators=(",", ":")) # mutable-ok: JSON requires a concrete dict
async def close(self) -> None:
await self._connection.close()
@ -84,19 +91,28 @@ class _ConnectionAdapter:
async def connect(
*,
url: str,
headers: dict[str, str],
provider: str,
api_key: str | None,
api_base: str | None,
headers: Mapping[str, str],
timeout: float | httpx.Timeout | None,
) -> _ConnectionAdapter | None:
connection_type: Final = load_rust_responses_websocket()
if connection_type is None:
return None
credentials: Final = None if api_key is None else MappingProxyType({"api_key": api_key})
try:
connection: Final = await connection_type.connect(
url=url,
headers=headers,
timeout_seconds=timeout_to_seconds(timeout),
provider,
credentials,
api_base,
headers,
timeout_to_seconds(timeout),
)
except Exception: # noqa: BLE001 # bridge failures must fall back to Python
return None
except Exception as error:
native: Final = get_native_bridge()
declined: Final = None if native is None else getattr(native, "RustBridgeDeclined", None)
if isinstance(declined, type) and isinstance(error, declined):
return None
raise
return _ConnectionAdapter(connection)

View file

@ -0,0 +1,431 @@
from __future__ import annotations
import json
from collections.abc import AsyncIterator, Generator, Iterator, Mapping
from dataclasses import dataclass
from functools import lru_cache
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, NoReturn, Protocol, TypeAlias, runtime_checkable
import httpx
from litellm.exceptions import APIError
from litellm.rust_bridge.loader import get_native_bridge
from litellm.rust_bridge.timeouts import timeout_to_seconds
if TYPE_CHECKING:
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponsesAPIStreamingResponse
StreamApi: TypeAlias = Literal["chat_completions", "messages", "responses"]
StreamTransport: TypeAlias = Literal["http", "websocket"]
Event: TypeAlias = Mapping[str, object]
@runtime_checkable
class RustEventStream(Protocol):
@property
def metadata(self) -> Mapping[str, object]: ...
def next_event(self) -> Event | None: ...
async def anext_event(self) -> Event | None: ...
def close(self) -> None: ...
async def aclose(self) -> None: ...
class RustStreamOpen(Protocol):
def __call__(
self,
request: Mapping[str, object],
provider: str,
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout_seconds: float | None,
) -> RustEventStream: ...
class RustAsyncStreamOpen(Protocol):
async def __call__(
self,
request: Mapping[str, object],
provider: str,
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout_seconds: float | None,
) -> RustEventStream: ...
@runtime_checkable
class ObjectAwaitable(Protocol):
def __await__(self) -> Generator[object, None, object]: ...
class _Unset:
pass
_UNSET: Final = _Unset()
@dataclass(slots=True)
class _RustStreamingState:
chat: RustStreamOpen | None = None
achat: RustAsyncStreamOpen | None = None
messages: RustStreamOpen | None = None
amessages: RustAsyncStreamOpen | None = None
responses: RustStreamOpen | None = None
aresponses: RustAsyncStreamOpen | None = None
_STATE: Final = _RustStreamingState()
def set_rust_streaming(
*,
chat: RustStreamOpen | None | _Unset = _UNSET,
achat: RustAsyncStreamOpen | None | _Unset = _UNSET,
messages: RustStreamOpen | None | _Unset = _UNSET,
amessages: RustAsyncStreamOpen | None | _Unset = _UNSET,
responses: RustStreamOpen | None | _Unset = _UNSET,
aresponses: RustAsyncStreamOpen | None | _Unset = _UNSET,
) -> None:
if not isinstance(chat, _Unset):
_STATE.chat = chat
if not isinstance(achat, _Unset):
_STATE.achat = achat
if not isinstance(messages, _Unset):
_STATE.messages = messages
if not isinstance(amessages, _Unset):
_STATE.amessages = amessages
if not isinstance(responses, _Unset):
_STATE.responses = responses
if not isinstance(aresponses, _Unset):
_STATE.aresponses = aresponses
def _native_attribute(name: str) -> object | None:
native: Final = get_native_bridge()
return None if native is None else getattr(native, name, None)
def _native_sync_opener(name: str) -> RustStreamOpen | None:
opener: Final = _native_attribute(name)
if not callable(opener):
return None
def open_stream(
request: Mapping[str, object],
provider: str,
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout_seconds: float | None,
) -> RustEventStream:
stream: Final = opener(
request,
provider,
credentials,
api_base,
extra_headers,
timeout_seconds,
)
if not isinstance(stream, RustEventStream):
raise TypeError("native stream opener returned an invalid stream")
return stream
return open_stream
def _native_async_opener(name: str) -> RustAsyncStreamOpen | None:
opener: Final = _native_attribute(name)
if not callable(opener):
return None
async def open_stream(
request: Mapping[str, object],
provider: str,
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout_seconds: float | None,
) -> RustEventStream:
pending: Final = opener(
request,
provider,
credentials,
api_base,
extra_headers,
timeout_seconds,
)
if not isinstance(pending, ObjectAwaitable):
raise TypeError("native async stream opener returned a non-awaitable")
stream: Final[object] = await pending
if not isinstance(stream, RustEventStream):
raise TypeError("native async stream opener returned an invalid stream")
return stream
return open_stream
def _sync_opener(api: StreamApi) -> RustStreamOpen | None:
match api:
case "chat_completions":
return _STATE.chat or _native_sync_opener("chat_completions_stream")
case "messages":
return _STATE.messages or _native_sync_opener("messages_stream")
case "responses":
return _STATE.responses or _native_sync_opener("responses_stream")
def _async_opener(api: StreamApi) -> RustAsyncStreamOpen | None:
match api:
case "chat_completions":
return _STATE.achat or _native_async_opener("achat_completions_stream")
case "messages":
return _STATE.amessages or _native_async_opener("amessages_stream")
case "responses":
return _STATE.aresponses or _native_async_opener("aresponses_stream")
def _native_exceptions() -> tuple[type[BaseException], type[BaseException]] | None:
native: Final = get_native_bridge()
if native is None:
return None
declined: Final = getattr(native, "RustBridgeDeclined", None)
upstream: Final = getattr(native, "RustUpstreamError", None)
if not isinstance(declined, type) or not isinstance(upstream, type):
return None
return declined, upstream
def _handle_open_error(error: Exception, provider: str) -> None:
exceptions: Final = _native_exceptions()
if exceptions is not None and isinstance(error, exceptions[0]):
return
_raise_stream_error(error, provider)
def _raise_stream_error(error: Exception, provider: str) -> NoReturn:
exceptions: Final = _native_exceptions()
if exceptions is None or not isinstance(error, exceptions[1]):
raise error
args: Final[tuple[object, ...]] = error.args
status_value: Final = args[0] if args else 0
message_value: Final = args[1] if len(args) > 1 else str(error)
status: Final = status_value if isinstance(status_value, int) else 0
message: Final = message_value if isinstance(message_value, str) else str(message_value)
raise APIError(
status_code=status or 500,
message=f"litellm rust typed stream: {message}",
llm_provider=provider,
model="",
) from error
class TypedEventStreamAdapter:
def __init__(self, stream: RustEventStream, provider: str) -> None:
self._stream: Final = stream
self._provider: Final = provider
self.metadata: Final = stream.metadata
self._mode: Literal["sync", "async"] | None = None
def _claim(self, mode: Literal["sync", "async"]) -> None:
if self._mode is None:
self._mode = mode
return
if self._mode != mode:
raise RuntimeError("native stream cannot mix synchronous and asynchronous consumption")
def __iter__(self) -> Iterator[Event]:
self._claim("sync")
return self
def __next__(self) -> Event:
self._claim("sync")
event: Final = self._next_event()
if event is None:
raise StopIteration
return event
def _next_event(self) -> Event | None:
try:
return self._stream.next_event()
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
_raise_stream_error(error, self._provider)
def __aiter__(self) -> AsyncIterator[Event]:
self._claim("async")
return self
async def __anext__(self) -> Event:
self._claim("async")
try:
event: Final = await self._stream.anext_event()
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
_raise_stream_error(error, self._provider)
if event is None:
raise StopAsyncIteration
return event
def close(self) -> None:
self._stream.close()
async def aclose(self) -> None:
await self._stream.aclose()
class MessagesSseStreamAdapter:
def __init__(self, events: TypedEventStreamAdapter) -> None:
self._events: Final = events
self.metadata: Final = events.metadata
def __iter__(self) -> Iterator[bytes]:
return (_event_to_sse(event) for event in self._events)
async def __aiter__(self) -> AsyncIterator[bytes]:
async for event in self._events:
yield _event_to_sse(event)
def close(self) -> None:
self._events.close()
async def aclose(self) -> None:
await self._events.aclose()
class ResponsesSdkEventStreamAdapter:
def __init__(self, events: TypedEventStreamAdapter) -> None:
self._events: Final = events
self.metadata: Final = events.metadata
def __iter__(self) -> Iterator[ResponsesAPIStreamingResponse]:
return (_responses_event_to_sdk(event) for event in self._events)
async def __aiter__(self) -> AsyncIterator[ResponsesAPIStreamingResponse]:
async for event in self._events:
yield _responses_event_to_sdk(event)
def close(self) -> None:
self._events.close()
async def aclose(self) -> None:
await self._events.aclose()
def _responses_event_to_sdk(event: Event) -> ResponsesAPIStreamingResponse:
from litellm.types.llms.openai import GenericEvent
event_type: Final = event.get("type")
model: Final = _responses_event_models().get(event_type) if isinstance(event_type, str) else None
return (model or GenericEvent).model_validate(event)
@lru_cache(maxsize=1)
def _responses_event_models() -> Mapping[str, type[BaseLiteLLMOpenAIResponseObject]]:
from litellm.types.llms import openai as openai_types
return MappingProxyType(
{
"response.created": openai_types.ResponseCreatedEvent,
"response.in_progress": openai_types.ResponseInProgressEvent,
"response.completed": openai_types.ResponseCompletedEvent,
"response.failed": openai_types.ResponseFailedEvent,
"response.incomplete": openai_types.ResponseIncompleteEvent,
"response.reasoning_summary_part.added": openai_types.ResponsePartAddedEvent,
"response.reasoning_summary_text.delta": openai_types.ReasoningSummaryTextDeltaEvent,
"response.reasoning_summary_text.done": openai_types.ReasoningSummaryTextDoneEvent,
"response.reasoning_summary_part.done": openai_types.ReasoningSummaryPartDoneEvent,
"response.output_item.added": openai_types.OutputItemAddedEvent,
"response.output_item.done": openai_types.OutputItemDoneEvent,
"response.content_part.added": openai_types.ContentPartAddedEvent,
"response.content_part.done": openai_types.ContentPartDoneEvent,
"response.output_text.delta": openai_types.OutputTextDeltaEvent,
"response.output_text.annotation.added": openai_types.OutputTextAnnotationAddedEvent,
"response.output_text.done": openai_types.OutputTextDoneEvent,
"response.refusal.delta": openai_types.RefusalDeltaEvent,
"response.refusal.done": openai_types.RefusalDoneEvent,
"response.function_call_arguments.delta": openai_types.FunctionCallArgumentsDeltaEvent,
"response.function_call_arguments.done": openai_types.FunctionCallArgumentsDoneEvent,
"response.file_search_call.in_progress": openai_types.FileSearchCallInProgressEvent,
"response.file_search_call.searching": openai_types.FileSearchCallSearchingEvent,
"response.file_search_call.completed": openai_types.FileSearchCallCompletedEvent,
"response.web_search_call.in_progress": openai_types.WebSearchCallInProgressEvent,
"response.web_search_call.searching": openai_types.WebSearchCallSearchingEvent,
"response.web_search_call.completed": openai_types.WebSearchCallCompletedEvent,
"response.mcp_list_tools.in_progress": openai_types.MCPListToolsInProgressEvent,
"response.mcp_list_tools.completed": openai_types.MCPListToolsCompletedEvent,
"response.mcp_list_tools.failed": openai_types.MCPListToolsFailedEvent,
"response.mcp_call.in_progress": openai_types.MCPCallInProgressEvent,
"response.mcp_call_arguments.delta": openai_types.MCPCallArgumentsDeltaEvent,
"response.mcp_call_arguments.done": openai_types.MCPCallArgumentsDoneEvent,
"response.mcp_call.completed": openai_types.MCPCallCompletedEvent,
"response.mcp_call.failed": openai_types.MCPCallFailedEvent,
"image_generation.partial_image": openai_types.ImageGenerationPartialImageEvent,
"error": openai_types.ErrorEvent,
}
)
def _event_to_sse(event: Event) -> bytes:
payload: Final = dict(event) # mutable-ok: JSON requires a concrete dict
return f"data: {json.dumps(payload, separators=(',', ':'))}\n\n".encode()
def open_stream(
*,
api: StreamApi,
provider: str,
request: Mapping[str, object],
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout: float | httpx.Timeout | None,
) -> TypedEventStreamAdapter | None:
opener: Final = _sync_opener(api)
if opener is None:
return None
try:
stream: Final = opener(
request,
provider,
credentials,
api_base,
extra_headers,
timeout_to_seconds(timeout),
)
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
_handle_open_error(error, provider)
return None
return TypedEventStreamAdapter(stream, provider)
async def aopen_stream(
*,
api: StreamApi,
provider: str,
request: Mapping[str, object],
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout: float | httpx.Timeout | None,
) -> TypedEventStreamAdapter | None:
opener: Final = _async_opener(api)
if opener is None:
return None
try:
stream: Final = await opener(
request,
provider,
credentials,
api_base,
extra_headers,
timeout_to_seconds(timeout),
)
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
_handle_open_error(error, provider)
return None
return TypedEventStreamAdapter(stream, provider)

View file

@ -2672,7 +2672,7 @@ async def test_generic_http_handler_async_streaming_forwards_provider_response_h
@pytest.mark.parametrize(
"custom_llm_provider, litellm_params, expected",
[
("openai", GenericLiteLLMParams(rust=True), True),
("openai", GenericLiteLLMParams(rust=True), False),
("openai", GenericLiteLLMParams(), False),
("openai", GenericLiteLLMParams(rust=False), False),
("azure", GenericLiteLLMParams(rust=True), False),
@ -2680,7 +2680,7 @@ async def test_generic_http_handler_async_streaming_forwards_provider_response_h
(None, GenericLiteLLMParams(rust=True), False),
],
)
def test_the_rust_responses_websocket_needs_both_openai_and_the_rust_flag(
def test_the_rust_responses_websocket_stays_disabled_without_a_typed_capability(
custom_llm_provider, litellm_params, expected
):
assert _rust_responses_websocket_enabled(custom_llm_provider, litellm_params) is expected

View file

@ -1,5 +1,7 @@
from __future__ import annotations
from collections.abc import Mapping
import pytest
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
@ -9,21 +11,27 @@ from litellm.types.router import GenericLiteLLMParams
class _FakeNativeConnection:
def __init__(self) -> None:
self.sent: list[str] = []
self.sent: list[dict[str, object]] = []
self.closed = False
async def send_text(self, text: str) -> None:
self.sent.append(text)
async def send_event(self, event: Mapping[str, object]) -> None:
self.sent.append(dict(event))
async def recv_text(self) -> str:
return "response.completed"
async def recv_event(self) -> dict[str, object]:
return {"type": "response.completed"}
async def close(self) -> None:
self.closed = True
class _ClosedNativeConnection:
async def recv_text(self) -> None:
async def send_event(self, event: Mapping[str, object]) -> None:
return None
async def recv_event(self) -> None:
return None
async def close(self) -> None:
return None
@ -31,14 +39,22 @@ class _FakeNativeBridge:
@classmethod
async def connect(
cls,
*,
url: str,
headers: dict[str, str],
provider: str,
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout_seconds: float | None,
) -> _FakeNativeConnection:
return _FakeNativeConnection()
@pytest.fixture(autouse=True)
def reset_responses_websocket():
responses_websocket.set_rust_responses_websocket(connection=None)
yield
responses_websocket.set_rust_responses_websocket(connection=None)
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))
@ -60,7 +76,9 @@ async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch)
assert (
await responses_websocket.connect(
url="wss://example.test/responses",
provider="openai",
api_key=None,
api_base="https://example.test",
headers={},
timeout=None,
)
@ -75,12 +93,14 @@ async def test_enabled_bridge_connects_and_adapts_socket(
responses_websocket.set_rust_responses_websocket(connection=_FakeNativeBridge)
connection = await responses_websocket.connect(
url="wss://example.test/responses",
provider="openai",
api_key="key",
api_base="https://example.test",
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.send('{"type":"response.create","model":"gpt-5"}')
assert await connection.recv() == '{"type":"response.completed"}'
await connection.close()

View file

@ -0,0 +1,58 @@
from __future__ import annotations
import os
from importlib import import_module
from typing import Final, Protocol, cast
import pytest
class StreamRoute(Protocol):
def __call__(self, *, request: object, provider: str) -> object: ...
class ResponsesWebSocketSession(Protocol):
@classmethod
def connect(cls, *, provider: str) -> object: ...
class NativeBridge(Protocol):
RustBridgeDeclined: type[Exception]
ResponsesWebSocketSession: type[ResponsesWebSocketSession]
chat_completions_stream: StreamRoute
achat_completions_stream: StreamRoute
messages_stream: StreamRoute
amessages_stream: StreamRoute
responses_stream: StreamRoute
aresponses_stream: StreamRoute
try:
native: Final = cast(NativeBridge, import_module("litellm.rust_bridge._native"))
except ImportError:
if os.getenv("LITELLM_REQUIRE_NATIVE_BRIDGE") == "1":
raise
pytest.skip("the native bridge has not been built", allow_module_level=True)
STREAM_ROUTES: Final[tuple[tuple[StreamRoute, str], ...]] = (
(native.chat_completions_stream, "anthropic"),
(native.achat_completions_stream, "anthropic"),
(native.messages_stream, "anthropic"),
(native.amessages_stream, "anthropic"),
(native.responses_stream, "openai"),
(native.aresponses_stream, "openai"),
)
@pytest.mark.parametrize(("function", "provider"), STREAM_ROUTES)
def test_stream_routes_preserve_validation_errors(
function: StreamRoute,
provider: str,
) -> None:
with pytest.raises(ValueError):
function(request=object(), provider=provider)
def test_native_websocket_connect_preserves_decline() -> None:
with pytest.raises(native.RustBridgeDeclined):
native.ResponsesWebSocketSession.connect(provider="unsupported")

View file

@ -0,0 +1,292 @@
from __future__ import annotations
from collections.abc import Iterator, Mapping
from typing import Final
from unittest.mock import MagicMock
import pytest
from litellm.exceptions import APIError
from litellm.rust_bridge import streaming
class _FakeEventStream:
def __init__(self, events: tuple[Mapping[str, object], ...]) -> None:
self.metadata: Final = {
"status_code": 200,
"provider": "anthropic",
"transport": "http",
"response_headers": [{"name": "x-test", "value": "ready"}],
}
self._events: Final = iter(events)
self.closed = False
def next_event(self) -> Mapping[str, object] | None:
if self.closed:
return None
return next(self._events, None)
async def anext_event(self) -> Mapping[str, object] | None:
return self.next_event()
def close(self) -> None:
self.closed = True
async def aclose(self) -> None:
self.close()
class _RecordingOpen:
def __init__(self, events: tuple[Mapping[str, object], ...]) -> None:
self._events: Final = events
self.calls = 0
def __call__(
self,
request: Mapping[str, object],
provider: str,
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout_seconds: float | None,
) -> _FakeEventStream:
self.calls += 1
return _FakeEventStream(self._events)
class _RecordingAsyncOpen:
def __init__(self, events: tuple[Mapping[str, object], ...]) -> None:
self._events: Final = events
self.calls = 0
async def __call__(
self,
request: Mapping[str, object],
provider: str,
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout_seconds: float | None,
) -> _FakeEventStream:
self.calls += 1
return _FakeEventStream(self._events)
class _Declined(Exception):
pass
class _Upstream(Exception):
pass
class _NativeErrors:
RustBridgeDeclined = _Declined
RustUpstreamError = _Upstream
class _FailingOpen:
def __init__(self, error: Exception) -> None:
self._error: Final = error
def __call__(
self,
request: Mapping[str, object],
provider: str,
credentials: Mapping[str, str] | None,
api_base: str | None,
extra_headers: Mapping[str, str] | None,
timeout_seconds: float | None,
) -> _FakeEventStream:
raise self._error
@pytest.fixture(autouse=True)
def reset_bridge() -> Iterator[None]:
streaming.set_rust_streaming(
chat=None,
achat=None,
messages=None,
amessages=None,
responses=None,
aresponses=None,
)
yield
streaming.set_rust_streaming(
chat=None,
achat=None,
messages=None,
amessages=None,
responses=None,
aresponses=None,
)
def _chat_event(text: str) -> Mapping[str, object]:
return {
"text": text,
"tool_use": None,
"is_finished": False,
"finish_reason": "",
"usage": None,
}
def test_unavailable_native_bridge_falls_back(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(streaming, "get_native_bridge", lambda: None)
result: Final = streaming.open_stream(
api="chat_completions",
provider="anthropic",
request={"model": "claude", "messages": []},
credentials={"api_key": "test"},
api_base=None,
extra_headers=None,
timeout=None,
)
assert result is None
def test_declined_open_failure_falls_back(
monkeypatch: pytest.MonkeyPatch,
) -> None:
python_calls = 0
monkeypatch.setattr(streaming, "get_native_bridge", lambda: _NativeErrors())
streaming.set_rust_streaming(
chat=_FailingOpen(_Declined("unsupported request")),
)
result: Final = streaming.open_stream(
api="chat_completions",
provider="anthropic",
request={"model": "claude", "messages": []},
credentials=None,
api_base=None,
extra_headers=None,
timeout=None,
)
if result is None:
python_calls += 1
assert python_calls == 1
def test_upstream_open_failure_never_falls_back(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(streaming, "get_native_bridge", lambda: _NativeErrors())
streaming.set_rust_streaming(
chat=_FailingOpen(_Upstream(503, "connection closed after request")),
)
with pytest.raises(APIError, match="connection closed after request"):
streaming.open_stream(
api="chat_completions",
provider="anthropic",
request={"model": "claude", "messages": []},
credentials=None,
api_base=None,
extra_headers=None,
timeout=None,
)
def test_sync_typed_events_preserve_shape_metadata_and_close() -> None:
opener: Final = _RecordingOpen((_chat_event("one"), _chat_event("two")))
streaming.set_rust_streaming(chat=opener)
result: Final = streaming.open_stream(
api="chat_completions",
provider="anthropic",
request={"model": "claude", "messages": []},
credentials=None,
api_base=None,
extra_headers=None,
timeout=1.0,
)
assert result is not None
assert tuple(event["text"] for event in result) == ("one", "two")
assert result.metadata["provider"] == "anthropic"
result.close()
assert opener.calls == 1
def test_chat_events_flow_through_custom_stream_wrapper() -> None:
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.types.utils import ModelResponseStream
native: Final = _FakeEventStream((_chat_event("hello"),))
events: Final = streaming.TypedEventStreamAdapter(native, "anthropic")
wrapper: Final = CustomStreamWrapper(
completion_stream=events,
model="claude",
logging_obj=MagicMock(),
custom_llm_provider="anthropic",
)
chunk: Final = next(wrapper)
assert isinstance(chunk, ModelResponseStream)
assert chunk.choices[0].delta.content == "hello"
@pytest.mark.asyncio
async def test_async_typed_events_and_cancellation() -> None:
opener: Final = _RecordingAsyncOpen((_chat_event("one"), _chat_event("two")))
streaming.set_rust_streaming(achat=opener)
result: Final = await streaming.aopen_stream(
api="chat_completions",
provider="anthropic",
request={"model": "claude", "messages": []},
credentials=None,
api_base=None,
extra_headers=None,
timeout=None,
)
assert result is not None
collected: Final = tuple([event async for event in result])
assert tuple(event["text"] for event in collected) == ("one", "two")
await result.aclose()
def test_messages_events_are_wrapped_in_existing_sse_bytes() -> None:
native: Final = _FakeEventStream(({"type": "message_stop"},))
events: Final = streaming.TypedEventStreamAdapter(native, "anthropic")
messages: Final = streaming.MessagesSseStreamAdapter(events)
assert tuple(messages) == (b'data: {"type":"message_stop"}\n\n',)
def test_responses_events_are_validated_into_existing_sdk_objects() -> None:
from litellm.types.llms.openai import OutputTextDeltaEvent
native: Final = _FakeEventStream(
(
{
"type": "response.output_text.delta",
"item_id": "item_1",
"output_index": 0,
"content_index": 0,
"delta": "hello",
},
)
)
events: Final = streaming.TypedEventStreamAdapter(native, "openai")
responses: Final = streaming.ResponsesSdkEventStreamAdapter(events)
event: Final = next(iter(responses))
assert isinstance(event, OutputTextDeltaEvent)
assert event.type == "response.output_text.delta"
assert event.delta == "hello"
@pytest.mark.asyncio
async def test_rejects_mixed_sync_and_async_consumption() -> None:
events: Final = streaming.TypedEventStreamAdapter(_FakeEventStream((_chat_event("one"),)), "anthropic")
assert tuple(events) == (_chat_event("one"),)
with pytest.raises(RuntimeError, match="cannot mix"):
async for _ in events:
pass