diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index 1b6a7984b7f..c2f3ea6adc2 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -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 diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 5eba7d36442..22f29f566c4 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1428,6 +1428,7 @@ dependencies = [ "aws-sigv4", "aws-smithy-runtime-api", "aws-types", + "futures-util", "rand 0.8.7", "reqwest", "serde", diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index ab8050734f2..2b7c76a9bec 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +futures-util.workspace = true rand.workspace = true reqwest.workspace = true serde.workspace = true diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index fd443f77a01..07298e33a4e 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -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, Error> { + Err(Error::Unsupported( + "chat completions streaming provider registration", + )) +} + #[cfg(test)] mod tests; diff --git a/litellm-rust/crates/core/src/chat_completions/transformation.rs b/litellm-rust/crates/core/src/chat_completions/transformation.rs index 8e0e31401ad..d9013284cb5 100644 --- a/litellm-rust/crates/core/src/chat_completions/transformation.rs +++ b/litellm-rust/crates/core/src/chat_completions/transformation.rs @@ -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 +{ +} + +impl ChatCompletionsStreamProvider for T where + T: StreamProvider +{ +} /// How the upstream call is authenticated. API-key strategies are resolved in /// `prepare`; SigV4 needs the serialized body, so the handler signs it. diff --git a/litellm-rust/crates/core/src/chat_completions/types.rs b/litellm-rust/crates/core/src/chat_completions/types.rs index 35dd543a986..ff0273604fb 100644 --- a/litellm-rust/crates/core/src/chat_completions/types.rs +++ b/litellm-rust/crates/core/src/chat_completions/types.rs @@ -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, 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), +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ChatStreamMessage { + pub role: ChatStreamRole, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(untagged)] +pub enum ChatStreamStop { + One(String), + Many(Vec), +} + +#[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, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub top_p: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_completion_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub stop: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub stream: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub stream_options: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tools: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_choice: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub response_format: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub reasoning_effort: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub thinking: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub user: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ChatCompletionsStreamRequestBody { + pub model: String, + pub messages: Vec, + #[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, + pub arguments: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub provider_specific_fields: Option, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ChatStreamToolCallChunk { + pub id: Option, + #[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, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub completion_tokens_details: Option, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ChatStreamEvent { + pub text: String, + pub tool_use: Option, + pub is_finished: bool, + pub finish_reason: String, + pub usage: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub index: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub provider_specific_fields: Option, +} + +#[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 + }) + ); + } +} diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index 0e18d24e5d8..3a651111f8d 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -12,5 +12,6 @@ pub mod realtime; pub mod responses; pub mod router; pub mod routing_utils; +pub mod streaming; pub use error::Error; diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index b4ee4b247f6..a2b501743a6 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -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 { execute_messages_provider_call(prepare_messages_call(request)?).await @@ -28,5 +31,13 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result Result, Error> { + Err(Error::Unsupported( + "messages event streaming provider registration", + )) +} + #[cfg(test)] mod tests; diff --git a/litellm-rust/crates/core/src/messages/transformation.rs b/litellm-rust/crates/core/src/messages/transformation.rs index 6b0a62acd0f..9453b0290cc 100644 --- a/litellm-rust/crates/core/src/messages/transformation.rs +++ b/litellm-rust/crates/core/src/messages/transformation.rs @@ -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 +{ +} + +impl MessagesStreamProvider for T where + T: StreamProvider +{ +} #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum MessagesAuthStrategy { diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index b9f807c29fd..a28dc04ee05 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -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, } + +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, + }, + 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"} + }) + ); + } +} diff --git a/litellm-rust/crates/core/src/responses/mod.rs b/litellm-rust/crates/core/src/responses/mod.rs index 5ec5a2caef8..e65857814c5 100644 --- a/litellm-rust/crates/core/src/responses/mod.rs +++ b/litellm-rust/crates/core/src/responses/mod.rs @@ -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, Error> { + Err(Error::Unsupported( + "responses HTTP streaming provider registration", + )) +} + +pub async fn responses_websocket( + _request: ResponsesWebSocketRequest, +) -> Result, Error> { + Err(Error::Unsupported( + "responses WebSocket streaming provider registration", + )) +} diff --git a/litellm-rust/crates/core/src/responses/types.rs b/litellm-rust/crates/core/src/responses/types.rs index 4942309992e..c05408c2aa7 100644 --- a/litellm-rust/crates/core/src/responses/types.rs +++ b/litellm-rust/crates/core/src/responses/types.rs @@ -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, } +pub type ResponsesStreamEvent = ResponsesWsEvent; +pub type ResponseCommand = ResponsesWsEvent; + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(untagged)] +pub enum ResponsesInput { + Text(String), + Items(Vec), +} + +#[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, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_output_tokens: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub previous_response_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub stream: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub store: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tools: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_choice: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub metadata: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub include: Option>, + #[serde(flatten)] + pub extra: Map, +} + +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" + }) + ); + } } diff --git a/litellm-rust/crates/core/src/responses/websocket.rs b/litellm-rust/crates/core/src/responses/websocket.rs index 5d037e9cf1b..279a94e2a7d 100644 --- a/litellm-rust/crates/core/src/responses/websocket.rs +++ b/litellm-rust/crates/core/src/responses/websocket.rs @@ -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, Error>>; + + fn close(&self) -> BoxFuture<'_, Result<(), Error>>; +} pub trait ResponsesWebSocketProviderConfig: Sync { fn supports_native_websocket(&self) -> bool { diff --git a/litellm-rust/crates/core/src/streaming.rs b/litellm-rust/crates/core/src/streaming.rs new file mode 100644 index 00000000000..ec12b828e8b --- /dev/null +++ b/litellm-rust/crates/core/src/streaming.rs @@ -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 = Pin> + Send + 'static>>; +pub type ProviderChunkStream = + Pin> + 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 { + 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); + +#[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, + aws_access_key_id: Option, + aws_secret_access_key: Option, + aws_session_token: Option, +} + +impl ProviderCredentials { + pub fn new( + api_key: Option, + aws_access_key_id: Option, + aws_secret_access_key: Option, + aws_session_token: Option, + ) -> 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() {} +/// assert_serialize::(); +/// assert_serialize::(); +/// assert_serialize::(); +/// ``` +/// +/// ```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, +} + +impl StreamTarget { + pub fn new( + provider: StreamProviderId, + credentials: ProviderCredentials, + api_base: Option, + ) -> 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
, + timeout: Option, +} + +impl StreamTransportOptions { + pub fn new(forwarded_headers: Vec
, timeout: Option) -> Self { + Self { + forwarded_headers, + timeout, + } + } + + pub fn forwarded_headers(&self) -> &[Header] { + &self.forwarded_headers + } + + pub fn timeout(&self) -> Option { + 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
, +} + +pub struct OpenedStream { + pub metadata: StreamMetadata, + pub events: EventStream, +} + +pub struct OpenedWireStream { + pub metadata: StreamMetadata, + pub chunks: ProviderChunkStream, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ProviderStreamChunk(Vec); + +impl ProviderStreamChunk { + pub fn new(bytes: impl Into>) -> 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, Error>; + + fn finish(&mut self) -> Result, Error> { + Ok(Vec::new()) + } +} + +pub trait StreamProvider: Send + Sync + 'static { + type PreparedRequest: Send + 'static; + type WireEvent: Send + 'static; + type Decoder: StreamDecoder; + + fn transform_request(&self, request: R) -> Result; + + fn call( + &'static self, + request: Self::PreparedRequest, + ) -> BoxFuture<'static, Result>; + + fn decoder(&self) -> Self::Decoder; + + fn normalize(&self, event: Self::WireEvent) -> Result, Error>; +} + +struct PipelineState +where + D: StreamDecoder, +{ + provider: &'static P, + decoder: D, + chunks: ProviderChunkStream, + pending: VecDeque>, + finished: bool, + request: PhantomData, +} + +pub async fn open_provider_stream( + provider: &'static P, + request: R, +) -> Result, Error> +where + P: StreamProvider, + 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(state: &mut PipelineState, events: Vec) +where + P: StreamProvider, + D: StreamDecoder>::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>>, + } + + struct FakeDecoder { + calls: Arc>>, + pending: String, + } + + impl StreamDecoder for FakeDecoder { + type WireEvent = String; + + fn push(&mut self, chunk: ProviderStreamChunk) -> Result, 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::>(); + self.pending = parts.pop().expect("split always returns one item"); + Ok(parts) + } + + fn finish(&mut self) -> Result, Error> { + if self.pending.is_empty() { + return Ok(Vec::new()); + } + Ok(vec![std::mem::take(&mut self.pending)]) + } + } + + impl StreamProvider for FakeProvider { + type PreparedRequest = PreparedRequest; + type WireEvent = String; + type Decoder = FakeDecoder; + + fn transform_request(&self, _request: FakeRequest) -> Result { + self.calls.lock().expect("call log").push("transform"); + Ok(PreparedRequest) + } + + fn call( + &'static self, + _request: Self::PreparedRequest, + ) -> BoxFuture<'static, Result> { + 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, 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::>() + .await + .into_iter() + .collect::, _>>() + .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", + ] + ); + } +} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 5fcdc145187..976a2596bf1 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -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", ]; diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index d9580cc8ba9..d6675a2f1f2 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -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)] diff --git a/litellm-rust/crates/python-bridge/src/routes/receiver.rs b/litellm-rust/crates/python-bridge/src/routes/receiver.rs new file mode 100644 index 00000000000..a97a256cfa2 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/receiver.rs @@ -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 { + receiver: Mutex>>, + reading: AtomicBool, + closed: AtomicBool, + producer: std::sync::Mutex>>, +} + +impl Drop for ReceiverState { + 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 { + state: Arc>, +} + +struct ReadGuard<'a>(&'a AtomicBool); + +impl Drop for ReadGuard<'_> { + fn drop(&mut self) { + self.0.store(false, Ordering::Release); + } +} + +impl BridgeReceiver { + pub(super) fn from_stream(stream: S) -> Self + where + S: Stream> + 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, 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); + + 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::>::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::>::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::, 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"); + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/runtime.rs b/litellm-rust/crates/python-bridge/src/routes/runtime.rs index 4117e5fd16c..ad447cb8f2c 100644 --- a/litellm-rust/crates/python-bridge/src/routes/runtime.rs +++ b/litellm-rust/crates/python-bridge/src/routes/runtime.rs @@ -65,6 +65,45 @@ where }) } +pub(super) fn run_sync_with(py: Python<'_>, future: F, convert: C) -> PyResult> +where + T: Send + 'static, + F: Future> + Send + 'static, + C: FnOnce(Python<'_>, T) -> PyResult>, +{ + 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( + py: Python<'_>, + future: F, + convert: C, +) -> PyResult> +where + T: Send + 'static, + F: Future> + Send + 'static, + C: FnOnce(Python<'_>, T) -> PyResult> + 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(result: Result, map_error: fn(Error) -> PyErr) -> PyResult { match result { Ok(value) => Ok(value), diff --git a/litellm-rust/crates/python-bridge/src/routes/streaming.rs b/litellm-rust/crates/python-bridge/src/routes/streaming.rs new file mode 100644 index 00000000000..88fa4287851 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/streaming.rs @@ -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 { + metadata: StreamMetadata, + receiver: BridgeReceiver, +} + +impl TypedEventReceiver +where + E: Send + 'static, +{ + fn from_opened(opened: OpenedStream) -> Self { + Self { + metadata: opened.metadata, + receiver: BridgeReceiver::from_stream(opened.events), + } + } +} + +fn next_event( + py: Python<'_>, + receiver: BridgeReceiver, + stop_iteration: bool, +) -> PyResult> +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( + py: Python<'_>, + receiver: BridgeReceiver, + stop_iteration: bool, +) -> PyResult> +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> for $class { + fn from(opened: OpenedStream<$event>) -> Self { + Self { + inner: TypedEventReceiver::from_opened(opened), + } + } + } + + #[pymethods] + impl $class { + #[getter] + fn metadata(&self, py: Python<'_>) -> PyResult> { + to_py(py, &self.inner.metadata) + } + + fn next_event(&self, py: Python<'_>) -> PyResult> { + next_event(py, self.inner.receiver.clone(), false) + } + + fn anext_event<'py>(&self, py: Python<'py>) -> PyResult> { + anext_event(py, self.inner.receiver.clone(), false) + } + + fn __iter__(slf: PyRef<'_, Self>) -> PyRef<'_, Self> { + slf + } + + fn __next__(&self, py: Python<'_>) -> PyResult> { + next_event(py, self.inner.receiver.clone(), true) + } + + fn __aiter__(slf: PyRef<'_, Self>) -> PyRef<'_, Self> { + slf + } + + fn __anext__<'py>(&self, py: Python<'py>) -> PyResult> { + anext_event(py, self.inner.receiver.clone(), true) + } + + fn close(&self) { + self.inner.receiver.close(); + } + + fn aclose<'py>(&self, py: Python<'py>) -> PyResult> { + 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, + aws_access_key_id: Option, + aws_secret_access_key: Option, + aws_session_token: Option, +} + +struct PythonStreamTarget { + provider: String, + credentials: Option>, + api_base: Option, +} + +struct PythonStreamTransport { + extra_headers: Option>, + timeout_seconds: Option, +} + +fn parse_call( + py: Python<'_>, + request: Py, + 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::(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, + provider: String, + credentials: Option>, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + let (body, target, transport) = parse_call::( + 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, + provider: String, + credentials: Option>, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + let (body, target, transport) = parse_call::( + 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, + provider: String, + credentials: Option>, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + let (body, target, transport) = parse_call::( + 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, + provider: String, + credentials: Option>, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + let (body, target, transport) = parse_call::( + 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, + provider: String, + credentials: Option>, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + let (body, target, transport) = parse_call::( + 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, + provider: String, + credentials: Option>, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + let (body, target, transport) = parse_call::( + 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, +} + +#[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>, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, + ) -> PyResult> { + let empty_request = pyo3::types::PyDict::new(py).unbind().into_any(); + let (_, target, transport) = parse_call::( + 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) -> PyResult> { + 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> { + 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> { + 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::()?; + module.add_class::()?; + module.add_class::()?; + module.add_class::()?; + 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)?) +} diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index c82be07a5c5..877aab0ed00 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -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, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 78f88a0b305..bd28d121590 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, ) diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 82297d35170..b790526c0d2 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -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: diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 0634867af1c..b4079b3d37e 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -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) diff --git a/litellm/rust_bridge/streaming.py b/litellm/rust_bridge/streaming.py new file mode 100644 index 00000000000..33ca3186d13 --- /dev/null +++ b/litellm/rust_bridge/streaming.py @@ -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) diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 26f841c1146..2b63f48e036 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -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 diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index c9a5b988be6..62cd7e62f63 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -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() diff --git a/tests/test_litellm/rust_bridge/test_native_integration.py b/tests/test_litellm/rust_bridge/test_native_integration.py new file mode 100644 index 00000000000..cabfbb7558d --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_native_integration.py @@ -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") diff --git a/tests/test_litellm/rust_bridge/test_streaming.py b/tests/test_litellm/rust_bridge/test_streaming.py new file mode 100644 index 00000000000..04624425dc4 --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_streaming.py @@ -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