diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 0d009d36d16..283fccaebe6 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -16,11 +16,14 @@ pub mod response_utils; pub mod transformation; pub mod types; +use crate::streaming::OpenedStream; use serde_json::{Map, Value}; 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<'_>, @@ -54,5 +57,51 @@ pub fn chat_completions_decline_reason( .map(|reason| reason.0) } +pub async fn chat_completions_stream( + _request: ChatCompletionsStreamRequest, +) -> Result, Error> { + Err(crate::Error::Unsupported( + "chat completions streaming provider registration", + )) +} + #[cfg(test)] mod tests; + +#[cfg(test)] +mod stream_entrypoint_tests { + use serde_json::json; + + use super::*; + use crate::Error; + use crate::streaming::{ + ProviderCredentials, StreamProviderId, StreamTarget, StreamTransportOptions, + }; + + #[tokio::test] + async fn typed_stream_declines_until_a_provider_is_registered() { + let body = serde_json::from_value(json!({ + "model": "claude-sonnet", + "messages": [{"role": "user", "content": "hello"}], + "stream": true + })) + .expect("valid chat stream request"); + let result = chat_completions_stream(ChatCompletionsStreamRequest { + body, + target: StreamTarget::new( + StreamProviderId::Anthropic, + ProviderCredentials::default(), + None, + ), + transport: StreamTransportOptions::default(), + }) + .await; + + assert!(matches!( + result, + Err(Error::Unsupported( + "chat completions streaming provider registration" + )) + )); + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/transformation.rs b/litellm-rust/crates/core/src/chat_completions/transformation.rs index a0868209305..aec28c68886 100644 --- a/litellm-rust/crates/core/src/chat_completions/transformation.rs +++ b/litellm-rust/crates/core/src/chat_completions/transformation.rs @@ -5,6 +5,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/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index ee2877e61fc..f9a70a991d8 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -14,10 +14,13 @@ mod handler; mod prepare; pub mod transformation; pub mod types; +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 @@ -27,5 +30,52 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result Result, Error> { + Err(crate::Error::Unsupported( + "messages event streaming provider registration", + )) +} + #[cfg(test)] mod tests; + +#[cfg(test)] +mod stream_entrypoint_tests { + use serde_json::json; + + use super::*; + use crate::Error; + use crate::streaming::{ + ProviderCredentials, StreamProviderId, StreamTarget, StreamTransportOptions, + }; + + #[tokio::test] + async fn typed_event_stream_declines_until_a_provider_is_registered() { + let body = serde_json::from_value(json!({ + "model": "claude-sonnet", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 32, + "stream": true + })) + .expect("valid Messages stream request"); + let result = messages_event_stream(MessagesStreamRequest { + body, + target: StreamTarget::new( + StreamProviderId::Anthropic, + ProviderCredentials::default(), + None, + ), + transport: StreamTransportOptions::default(), + }) + .await; + + assert!(matches!( + result, + Err(Error::Unsupported( + "messages event streaming provider registration" + )) + )); + } +} diff --git a/litellm-rust/crates/core/src/messages/transformation.rs b/litellm-rust/crates/core/src/messages/transformation.rs index 673a5728aca..1c3913a6802 100644 --- a/litellm-rust/crates/core/src/messages/transformation.rs +++ b/litellm-rust/crates/core/src/messages/transformation.rs @@ -1,5 +1,19 @@ -use super::types::{AnthropicMessagesRequest, AnthropicMessagesResponse}; use crate::Error; +use crate::streaming::StreamProvider; + +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..94f979bb730 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, @@ -47,6 +48,76 @@ pub struct ContentBlock { 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"} + }) + ); + } +} + #[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] pub struct CacheControl { #[serde(rename = "type", skip_serializing_if = "Option::is_none")] diff --git a/litellm-rust/crates/core/src/responses/mod.rs b/litellm-rust/crates/core/src/responses/mod.rs index 5ec5a2caef8..7b211ba191b 100644 --- a/litellm-rust/crates/core/src/responses/mod.rs +++ b/litellm-rust/crates/core/src/responses/mod.rs @@ -1,3 +1,81 @@ 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", + )) +} + +#[cfg(test)] +mod stream_entrypoint_tests { + use serde_json::json; + + use super::*; + use crate::streaming::{ + ProviderCredentials, StreamProviderId, StreamTarget, StreamTransportOptions, + }; + + fn target() -> StreamTarget { + StreamTarget::new( + StreamProviderId::OpenAi, + ProviderCredentials::default(), + None, + ) + } + + #[tokio::test] + async fn typed_http_stream_declines_until_a_provider_is_registered() { + let body = serde_json::from_value(json!({ + "model": "gpt-5", + "input": "hello", + "stream": true + })) + .expect("valid Responses stream request"); + let result = responses_stream(ResponsesStreamRequest { + body, + target: target(), + transport: StreamTransportOptions::default(), + }) + .await; + + assert!(matches!( + result, + Err(Error::Unsupported( + "responses HTTP streaming provider registration" + )) + )); + } + + #[tokio::test] + async fn typed_websocket_declines_until_a_provider_is_registered() { + let result = responses_websocket(ResponsesWebSocketRequest { + target: target(), + transport: StreamTransportOptions::default(), + }) + .await; + + assert!(matches!( + result, + 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..2432e46e017 100644 --- a/litellm-rust/crates/core/src/responses/websocket.rs +++ b/litellm-rust/crates/core/src/responses/websocket.rs @@ -1,6 +1,18 @@ use crate::Error; use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH}; -use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult}; +use crate::responses::types::{ + ResponseCommand, ResponsesStreamEvent, ResponsesWsEvent, ResponsesWsEventType, + ResponsesWsTransformResult, +}; +use futures_util::future::BoxFuture; + +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 {