mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
feat(rust): add typed streaming boundary
This commit is contained in:
parent
b0ae5bc4e3
commit
ceaf967da8
28 changed files with 2640 additions and 50 deletions
12
.github/workflows/test-rust.yml
vendored
12
.github/workflows/test-rust.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -1428,6 +1428,7 @@ dependencies = [
|
|||
"aws-sigv4",
|
||||
"aws-smithy-runtime-api",
|
||||
"aws-types",
|
||||
"futures-util",
|
||||
"rand 0.8.7",
|
||||
"reqwest",
|
||||
"serde",
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
futures-util.workspace = true
|
||||
rand.workspace = true
|
||||
reqwest.workspace = true
|
||||
serde.workspace = true
|
||||
|
|
|
|||
|
|
@ -19,9 +19,12 @@ use serde_json::{Map, Value};
|
|||
|
||||
use crate::error::Error;
|
||||
|
||||
use crate::streaming::OpenedStream;
|
||||
use handler::execute_chat_completions_provider_call;
|
||||
use prepare::{parse_messages, prepare_chat_completions_call, resolve_provider_config};
|
||||
use types::{ChatCompletionsRequest, ChatCompletionsResponse};
|
||||
use types::{
|
||||
ChatCompletionsRequest, ChatCompletionsResponse, ChatCompletionsStreamRequest, ChatStreamEvent,
|
||||
};
|
||||
|
||||
pub async fn chat_completions(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
|
|
@ -55,5 +58,13 @@ pub fn chat_completions_decline_reason(
|
|||
.map(|reason| reason.0)
|
||||
}
|
||||
|
||||
pub async fn chat_completions_stream(
|
||||
_request: ChatCompletionsStreamRequest,
|
||||
) -> Result<OpenedStream<ChatStreamEvent>, Error> {
|
||||
Err(Error::Unsupported(
|
||||
"chat completions streaming provider registration",
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -6,6 +6,18 @@ use super::types::{
|
|||
ChatCompletionsResponse, ChatMessage, ChatMessageContent, ProviderChatRequestData,
|
||||
ProviderChatResponseData,
|
||||
};
|
||||
use super::types::{ChatCompletionsStreamRequest, ChatStreamEvent};
|
||||
use crate::streaming::StreamProvider;
|
||||
|
||||
pub trait ChatCompletionsStreamProvider:
|
||||
StreamProvider<ChatCompletionsStreamRequest, ChatStreamEvent>
|
||||
{
|
||||
}
|
||||
|
||||
impl<T> ChatCompletionsStreamProvider for T where
|
||||
T: StreamProvider<ChatCompletionsStreamRequest, ChatStreamEvent>
|
||||
{
|
||||
}
|
||||
|
||||
/// How the upstream call is authenticated. API-key strategies are resolved in
|
||||
/// `prepare`; SigV4 needs the serialized body, so the handler signs it.
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ use serde::{Deserialize, Serialize};
|
|||
use serde_json::{Map, Value};
|
||||
|
||||
use super::transformation::{ChatCompletionsAuth, ChatCompletionsProviderConfig};
|
||||
use crate::streaming::{JsonObject, StreamTarget, StreamTransportOptions};
|
||||
|
||||
/// A `/chat/completions` call as it crosses into the core.
|
||||
///
|
||||
|
|
@ -110,3 +111,183 @@ pub struct ChatCompletionsResponse {
|
|||
pub choices: Vec<ChatCompletionsChoice>,
|
||||
pub usage: ChatCompletionsUsage,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum ChatStreamRole {
|
||||
Assistant,
|
||||
Developer,
|
||||
Function,
|
||||
System,
|
||||
Tool,
|
||||
User,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum ChatStreamMessageContent {
|
||||
Text(String),
|
||||
Parts(Vec<JsonObject>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ChatStreamMessage {
|
||||
pub role: ChatStreamRole,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub content: Option<ChatStreamMessageContent>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub name: Option<String>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum ChatStreamStop {
|
||||
One(String),
|
||||
Many(Vec<String>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum ChatStreamStringOrObject {
|
||||
Name(String),
|
||||
Definition(JsonObject),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ChatCompletionsStreamParameters {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub temperature: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub top_p: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub max_tokens: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub max_completion_tokens: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub stop: Option<ChatStreamStop>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub stream: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub stream_options: Option<JsonObject>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tools: Option<Vec<JsonObject>>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tool_choice: Option<ChatStreamStringOrObject>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub response_format: Option<JsonObject>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning_effort: Option<ChatStreamStringOrObject>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub thinking: Option<JsonObject>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub user: Option<String>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ChatCompletionsStreamRequestBody {
|
||||
pub model: String,
|
||||
pub messages: Vec<ChatStreamMessage>,
|
||||
#[serde(flatten)]
|
||||
pub parameters: ChatCompletionsStreamParameters,
|
||||
}
|
||||
|
||||
pub struct ChatCompletionsStreamRequest {
|
||||
pub body: ChatCompletionsStreamRequestBody,
|
||||
pub target: StreamTarget,
|
||||
pub transport: StreamTransportOptions,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ChatStreamToolFunctionChunk {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub name: Option<String>,
|
||||
pub arguments: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub provider_specific_fields: Option<JsonObject>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ChatStreamToolCallChunk {
|
||||
pub id: Option<String>,
|
||||
#[serde(rename = "type")]
|
||||
pub tool_type: String,
|
||||
pub function: ChatStreamToolFunctionChunk,
|
||||
pub index: u64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ChatStreamUsage {
|
||||
pub prompt_tokens: u64,
|
||||
pub completion_tokens: u64,
|
||||
pub total_tokens: u64,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub prompt_tokens_details: Option<JsonObject>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub completion_tokens_details: Option<JsonObject>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ChatStreamEvent {
|
||||
pub text: String,
|
||||
pub tool_use: Option<ChatStreamToolCallChunk>,
|
||||
pub is_finished: bool,
|
||||
pub finish_reason: String,
|
||||
pub usage: Option<ChatStreamUsage>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub index: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub provider_specific_fields: Option<JsonObject>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod stream_contract_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn request_uses_public_chat_completion_parameter_names() {
|
||||
let request: ChatCompletionsStreamRequestBody = serde_json::from_value(serde_json::json!({
|
||||
"model": "claude-sonnet",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"max_tokens": 32,
|
||||
"stream": true,
|
||||
"tool_choice": "auto"
|
||||
}))
|
||||
.expect("public request shape");
|
||||
|
||||
assert_eq!(request.parameters.max_tokens, Some(32));
|
||||
assert_eq!(request.parameters.stream, Some(true));
|
||||
assert!(matches!(
|
||||
request.parameters.tool_choice,
|
||||
Some(ChatStreamStringOrObject::Name(ref value)) if value == "auto"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_matches_python_generic_streaming_chunk_shape() {
|
||||
let event = ChatStreamEvent {
|
||||
text: "hello".to_string(),
|
||||
tool_use: None,
|
||||
is_finished: false,
|
||||
finish_reason: String::new(),
|
||||
usage: None,
|
||||
index: Some(0),
|
||||
provider_specific_fields: None,
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
serde_json::to_value(event).expect("serializable event"),
|
||||
serde_json::json!({
|
||||
"text": "hello",
|
||||
"tool_use": null,
|
||||
"is_finished": false,
|
||||
"finish_reason": "",
|
||||
"usage": null,
|
||||
"index": 0
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -12,5 +12,6 @@ pub mod realtime;
|
|||
pub mod responses;
|
||||
pub mod router;
|
||||
pub mod routing_utils;
|
||||
pub mod streaming;
|
||||
|
||||
pub use error::Error;
|
||||
|
|
|
|||
|
|
@ -15,10 +15,13 @@ pub mod transformation;
|
|||
pub mod types;
|
||||
|
||||
use crate::error::Error;
|
||||
use crate::streaming::OpenedStream;
|
||||
|
||||
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
|
||||
use prepare::prepare_messages_call;
|
||||
use types::{AnthropicMessagesResponse, MessagesRequest};
|
||||
use types::{
|
||||
AnthropicMessagesResponse, MessagesRequest, MessagesStreamEvent, MessagesStreamRequest,
|
||||
};
|
||||
|
||||
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
|
||||
execute_messages_provider_call(prepare_messages_call(request)?).await
|
||||
|
|
@ -28,5 +31,13 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Re
|
|||
execute_messages_provider_stream(prepare_messages_call(request)?).await
|
||||
}
|
||||
|
||||
pub async fn messages_event_stream(
|
||||
_request: MessagesStreamRequest,
|
||||
) -> Result<OpenedStream<MessagesStreamEvent>, Error> {
|
||||
Err(Error::Unsupported(
|
||||
"messages event streaming provider registration",
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -1,6 +1,19 @@
|
|||
use crate::error::Error;
|
||||
use crate::streaming::StreamProvider;
|
||||
|
||||
use super::types::{AnthropicMessagesRequest, AnthropicMessagesResponse};
|
||||
use super::types::{
|
||||
AnthropicMessagesRequest, AnthropicMessagesResponse, MessagesStreamEvent, MessagesStreamRequest,
|
||||
};
|
||||
|
||||
pub trait MessagesStreamProvider:
|
||||
StreamProvider<MessagesStreamRequest, MessagesStreamEvent>
|
||||
{
|
||||
}
|
||||
|
||||
impl<T> MessagesStreamProvider for T where
|
||||
T: StreamProvider<MessagesStreamRequest, MessagesStreamEvent>
|
||||
{
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum MessagesAuthStrategy {
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ use serde::{Deserialize, Serialize};
|
|||
use serde_json::{Map, Value};
|
||||
|
||||
use super::transformation::AnthropicMessagesProviderConfig;
|
||||
use crate::streaming::{JsonObject, StreamTarget, StreamTransportOptions};
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
pub model: &'a str,
|
||||
|
|
@ -132,3 +133,73 @@ pub struct AnthropicMessagesResponse {
|
|||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
pub struct MessagesStreamRequest {
|
||||
pub body: AnthropicMessagesRequest,
|
||||
pub target: StreamTarget,
|
||||
pub transport: StreamTransportOptions,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum MessagesStreamEvent {
|
||||
MessageStart {
|
||||
message: AnthropicMessagesResponse,
|
||||
},
|
||||
ContentBlockStart {
|
||||
index: u64,
|
||||
content_block: JsonObject,
|
||||
},
|
||||
ContentBlockDelta {
|
||||
index: u64,
|
||||
delta: JsonObject,
|
||||
},
|
||||
ContentBlockStop {
|
||||
index: u64,
|
||||
},
|
||||
MessageDelta {
|
||||
delta: JsonObject,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
usage: Option<JsonObject>,
|
||||
},
|
||||
MessageStop,
|
||||
Ping,
|
||||
Error {
|
||||
error: JsonObject,
|
||||
},
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod stream_contract_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn message_stop_serializes_as_anthropic_event() {
|
||||
assert_eq!(
|
||||
serde_json::to_value(MessagesStreamEvent::MessageStop).expect("serializable event"),
|
||||
serde_json::json!({"type": "message_stop"})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn content_delta_keeps_typed_event_fields() {
|
||||
let event = MessagesStreamEvent::ContentBlockDelta {
|
||||
index: 0,
|
||||
delta: JsonObject(
|
||||
serde_json::json!({"type": "text_delta", "text": "hello"})
|
||||
.as_object()
|
||||
.expect("object")
|
||||
.clone(),
|
||||
),
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
serde_json::to_value(event).expect("serializable event"),
|
||||
serde_json::json!({
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "hello"}
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,24 @@
|
|||
pub mod instrumentation;
|
||||
pub mod types;
|
||||
pub mod websocket;
|
||||
|
||||
use crate::error::Error;
|
||||
use crate::streaming::OpenedStream;
|
||||
use types::{ResponsesStreamEvent, ResponsesStreamRequest, ResponsesWebSocketRequest};
|
||||
use websocket::TypedResponsesWebSocketSession;
|
||||
|
||||
pub async fn responses_stream(
|
||||
_request: ResponsesStreamRequest,
|
||||
) -> Result<OpenedStream<ResponsesStreamEvent>, Error> {
|
||||
Err(Error::Unsupported(
|
||||
"responses HTTP streaming provider registration",
|
||||
))
|
||||
}
|
||||
|
||||
pub async fn responses_websocket(
|
||||
_request: ResponsesWebSocketRequest,
|
||||
) -> Result<Box<dyn TypedResponsesWebSocketSession>, Error> {
|
||||
Err(Error::Unsupported(
|
||||
"responses WebSocket streaming provider registration",
|
||||
))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,13 +1,46 @@
|
|||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::streaming::{JsonObject, StreamTarget, StreamTransportOptions};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum ResponsesWsEventType {
|
||||
ResponseCreate,
|
||||
ResponseCreated,
|
||||
ResponseInProgress,
|
||||
ResponseReasoningSummaryPartAdded,
|
||||
ResponseReasoningSummaryTextDelta,
|
||||
ResponseReasoningSummaryTextDone,
|
||||
ResponseReasoningSummaryPartDone,
|
||||
ResponseOutputItemAdded,
|
||||
ResponseOutputTextDelta,
|
||||
ResponseOutputTextAnnotationAdded,
|
||||
ResponseOutputTextDone,
|
||||
ResponseRefusalDelta,
|
||||
ResponseRefusalDone,
|
||||
ResponseFunctionCallArgumentsDelta,
|
||||
ResponseFunctionCallArgumentsDone,
|
||||
ResponseFileSearchCallInProgress,
|
||||
ResponseFileSearchCallSearching,
|
||||
ResponseFileSearchCallCompleted,
|
||||
ResponseWebSearchCallInProgress,
|
||||
ResponseWebSearchCallSearching,
|
||||
ResponseWebSearchCallCompleted,
|
||||
ResponseMcpListToolsInProgress,
|
||||
ResponseMcpListToolsCompleted,
|
||||
ResponseMcpListToolsFailed,
|
||||
ResponseMcpCallInProgress,
|
||||
ResponseMcpCallArgumentsDelta,
|
||||
ResponseMcpCallArgumentsDone,
|
||||
ResponseMcpCallCompleted,
|
||||
ResponseMcpCallFailed,
|
||||
ResponseContentPartAdded,
|
||||
ResponseContentPartDone,
|
||||
ResponseOutputItemDone,
|
||||
ResponseCompleted,
|
||||
ResponseFailed,
|
||||
ResponseIncomplete,
|
||||
ImageGenerationPartialImage,
|
||||
Error,
|
||||
Other(String),
|
||||
}
|
||||
|
|
@ -17,9 +50,40 @@ impl ResponsesWsEventType {
|
|||
match self {
|
||||
Self::ResponseCreate => "response.create",
|
||||
Self::ResponseCreated => "response.created",
|
||||
Self::ResponseInProgress => "response.in_progress",
|
||||
Self::ResponseReasoningSummaryPartAdded => "response.reasoning_summary_part.added",
|
||||
Self::ResponseReasoningSummaryTextDelta => "response.reasoning_summary_text.delta",
|
||||
Self::ResponseReasoningSummaryTextDone => "response.reasoning_summary_text.done",
|
||||
Self::ResponseReasoningSummaryPartDone => "response.reasoning_summary_part.done",
|
||||
Self::ResponseOutputItemAdded => "response.output_item.added",
|
||||
Self::ResponseOutputTextDelta => "response.output_text.delta",
|
||||
Self::ResponseOutputTextAnnotationAdded => "response.output_text.annotation.added",
|
||||
Self::ResponseOutputTextDone => "response.output_text.done",
|
||||
Self::ResponseRefusalDelta => "response.refusal.delta",
|
||||
Self::ResponseRefusalDone => "response.refusal.done",
|
||||
Self::ResponseFunctionCallArgumentsDelta => "response.function_call_arguments.delta",
|
||||
Self::ResponseFunctionCallArgumentsDone => "response.function_call_arguments.done",
|
||||
Self::ResponseFileSearchCallInProgress => "response.file_search_call.in_progress",
|
||||
Self::ResponseFileSearchCallSearching => "response.file_search_call.searching",
|
||||
Self::ResponseFileSearchCallCompleted => "response.file_search_call.completed",
|
||||
Self::ResponseWebSearchCallInProgress => "response.web_search_call.in_progress",
|
||||
Self::ResponseWebSearchCallSearching => "response.web_search_call.searching",
|
||||
Self::ResponseWebSearchCallCompleted => "response.web_search_call.completed",
|
||||
Self::ResponseMcpListToolsInProgress => "response.mcp_list_tools.in_progress",
|
||||
Self::ResponseMcpListToolsCompleted => "response.mcp_list_tools.completed",
|
||||
Self::ResponseMcpListToolsFailed => "response.mcp_list_tools.failed",
|
||||
Self::ResponseMcpCallInProgress => "response.mcp_call.in_progress",
|
||||
Self::ResponseMcpCallArgumentsDelta => "response.mcp_call_arguments.delta",
|
||||
Self::ResponseMcpCallArgumentsDone => "response.mcp_call_arguments.done",
|
||||
Self::ResponseMcpCallCompleted => "response.mcp_call.completed",
|
||||
Self::ResponseMcpCallFailed => "response.mcp_call.failed",
|
||||
Self::ResponseContentPartAdded => "response.content_part.added",
|
||||
Self::ResponseContentPartDone => "response.content_part.done",
|
||||
Self::ResponseOutputItemDone => "response.output_item.done",
|
||||
Self::ResponseCompleted => "response.completed",
|
||||
Self::ResponseFailed => "response.failed",
|
||||
Self::ResponseIncomplete => "response.incomplete",
|
||||
Self::ImageGenerationPartialImage => "image_generation.partial_image",
|
||||
Self::Error => "error",
|
||||
Self::Other(value) => value,
|
||||
}
|
||||
|
|
@ -44,9 +108,40 @@ impl<'de> Deserialize<'de> for ResponsesWsEventType {
|
|||
Ok(match value.as_str() {
|
||||
"response.create" => Self::ResponseCreate,
|
||||
"response.created" => Self::ResponseCreated,
|
||||
"response.in_progress" => Self::ResponseInProgress,
|
||||
"response.reasoning_summary_part.added" => Self::ResponseReasoningSummaryPartAdded,
|
||||
"response.reasoning_summary_text.delta" => Self::ResponseReasoningSummaryTextDelta,
|
||||
"response.reasoning_summary_text.done" => Self::ResponseReasoningSummaryTextDone,
|
||||
"response.reasoning_summary_part.done" => Self::ResponseReasoningSummaryPartDone,
|
||||
"response.output_item.added" => Self::ResponseOutputItemAdded,
|
||||
"response.output_text.delta" => Self::ResponseOutputTextDelta,
|
||||
"response.output_text.annotation.added" => Self::ResponseOutputTextAnnotationAdded,
|
||||
"response.output_text.done" => Self::ResponseOutputTextDone,
|
||||
"response.refusal.delta" => Self::ResponseRefusalDelta,
|
||||
"response.refusal.done" => Self::ResponseRefusalDone,
|
||||
"response.function_call_arguments.delta" => Self::ResponseFunctionCallArgumentsDelta,
|
||||
"response.function_call_arguments.done" => Self::ResponseFunctionCallArgumentsDone,
|
||||
"response.file_search_call.in_progress" => Self::ResponseFileSearchCallInProgress,
|
||||
"response.file_search_call.searching" => Self::ResponseFileSearchCallSearching,
|
||||
"response.file_search_call.completed" => Self::ResponseFileSearchCallCompleted,
|
||||
"response.web_search_call.in_progress" => Self::ResponseWebSearchCallInProgress,
|
||||
"response.web_search_call.searching" => Self::ResponseWebSearchCallSearching,
|
||||
"response.web_search_call.completed" => Self::ResponseWebSearchCallCompleted,
|
||||
"response.mcp_list_tools.in_progress" => Self::ResponseMcpListToolsInProgress,
|
||||
"response.mcp_list_tools.completed" => Self::ResponseMcpListToolsCompleted,
|
||||
"response.mcp_list_tools.failed" => Self::ResponseMcpListToolsFailed,
|
||||
"response.mcp_call.in_progress" => Self::ResponseMcpCallInProgress,
|
||||
"response.mcp_call_arguments.delta" => Self::ResponseMcpCallArgumentsDelta,
|
||||
"response.mcp_call_arguments.done" => Self::ResponseMcpCallArgumentsDone,
|
||||
"response.mcp_call.completed" => Self::ResponseMcpCallCompleted,
|
||||
"response.mcp_call.failed" => Self::ResponseMcpCallFailed,
|
||||
"response.content_part.added" => Self::ResponseContentPartAdded,
|
||||
"response.content_part.done" => Self::ResponseContentPartDone,
|
||||
"response.output_item.done" => Self::ResponseOutputItemDone,
|
||||
"response.completed" => Self::ResponseCompleted,
|
||||
"response.failed" => Self::ResponseFailed,
|
||||
"response.incomplete" => Self::ResponseIncomplete,
|
||||
"image_generation.partial_image" => Self::ImageGenerationPartialImage,
|
||||
"error" => Self::Error,
|
||||
_ => Self::Other(value),
|
||||
})
|
||||
|
|
@ -84,6 +179,60 @@ pub struct ResponsesWsTransformResult {
|
|||
pub events: Vec<ResponsesWsEvent>,
|
||||
}
|
||||
|
||||
pub type ResponsesStreamEvent = ResponsesWsEvent;
|
||||
pub type ResponseCommand = ResponsesWsEvent;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum ResponsesInput {
|
||||
Text(String),
|
||||
Items(Vec<JsonObject>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum ResponsesToolChoice {
|
||||
Name(String),
|
||||
Definition(JsonObject),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ResponsesStreamRequestBody {
|
||||
pub model: String,
|
||||
pub input: ResponsesInput,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub instructions: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub max_output_tokens: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub previous_response_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub stream: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub store: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tools: Option<Vec<JsonObject>>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tool_choice: Option<ResponsesToolChoice>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub metadata: Option<JsonObject>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub include: Option<Vec<String>>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
pub struct ResponsesStreamRequest {
|
||||
pub body: ResponsesStreamRequestBody,
|
||||
pub target: StreamTarget,
|
||||
pub transport: StreamTransportOptions,
|
||||
}
|
||||
|
||||
pub struct ResponsesWebSocketRequest {
|
||||
pub target: StreamTarget,
|
||||
pub transport: StreamTransportOptions,
|
||||
}
|
||||
|
||||
impl ResponsesWsTransformResult {
|
||||
pub fn passthrough(event: ResponsesWsEvent) -> Self {
|
||||
Self {
|
||||
|
|
@ -127,11 +276,14 @@ mod tests {
|
|||
let known: ResponsesWsEventType =
|
||||
serde_json::from_str("\"response.completed\"").expect("valid event type");
|
||||
assert_eq!(known, ResponsesWsEventType::ResponseCompleted);
|
||||
let unknown: ResponsesWsEventType =
|
||||
let output_delta: ResponsesWsEventType =
|
||||
serde_json::from_str("\"response.output_text.delta\"").expect("valid event type");
|
||||
assert_eq!(output_delta, ResponsesWsEventType::ResponseOutputTextDelta);
|
||||
let unknown: ResponsesWsEventType =
|
||||
serde_json::from_str("\"response.future_event\"").expect("valid event type");
|
||||
assert_eq!(
|
||||
unknown,
|
||||
ResponsesWsEventType::Other("response.output_text.delta".to_string())
|
||||
ResponsesWsEventType::Other("response.future_event".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -163,4 +315,42 @@ mod tests {
|
|||
assert_eq!(flat.model(), Some("gpt-5"));
|
||||
assert_eq!(nested.model(), Some("gpt-5-mini"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stream_request_deserializes_the_public_responses_shape() {
|
||||
let request: ResponsesStreamRequestBody = serde_json::from_value(serde_json::json!({
|
||||
"model": "gpt-5",
|
||||
"input": "hello",
|
||||
"stream": true,
|
||||
"max_output_tokens": 32
|
||||
}))
|
||||
.expect("public request shape");
|
||||
|
||||
assert_eq!(request.model, "gpt-5");
|
||||
assert_eq!(request.stream, Some(true));
|
||||
assert!(matches!(request.input, ResponsesInput::Text(ref text) if text == "hello"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_stream_events_round_trip_for_forward_compatibility() {
|
||||
let event: ResponsesStreamEvent = serde_json::from_value(serde_json::json!({
|
||||
"type": "response.future_event",
|
||||
"sequence_number": 7,
|
||||
"future_field": "value"
|
||||
}))
|
||||
.expect("unknown event");
|
||||
|
||||
assert_eq!(
|
||||
event.event_type,
|
||||
ResponsesWsEventType::Other("response.future_event".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::to_value(event).expect("serializable event"),
|
||||
serde_json::json!({
|
||||
"type": "response.future_event",
|
||||
"sequence_number": 7,
|
||||
"future_field": "value"
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,6 +1,19 @@
|
|||
use crate::Error;
|
||||
use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH};
|
||||
use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult};
|
||||
use futures_util::future::BoxFuture;
|
||||
|
||||
use crate::responses::types::{
|
||||
ResponseCommand, ResponsesStreamEvent, ResponsesWsEvent, ResponsesWsEventType,
|
||||
ResponsesWsTransformResult,
|
||||
};
|
||||
|
||||
pub trait TypedResponsesWebSocketSession: Send + Sync {
|
||||
fn send(&self, command: ResponseCommand) -> BoxFuture<'_, Result<(), Error>>;
|
||||
|
||||
fn recv(&self) -> BoxFuture<'_, Result<Option<ResponsesStreamEvent>, Error>>;
|
||||
|
||||
fn close(&self) -> BoxFuture<'_, Result<(), Error>>;
|
||||
}
|
||||
|
||||
pub trait ResponsesWebSocketProviderConfig: Sync {
|
||||
fn supports_native_websocket(&self) -> bool {
|
||||
|
|
|
|||
435
litellm-rust/crates/core/src/streaming.rs
Normal file
435
litellm-rust/crates/core/src/streaming.rs
Normal file
|
|
@ -0,0 +1,435 @@
|
|||
use std::collections::VecDeque;
|
||||
use std::marker::PhantomData;
|
||||
use std::pin::Pin;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::error::Error;
|
||||
use futures_util::future::BoxFuture;
|
||||
use futures_util::{Stream, StreamExt, stream};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub type EventStream<E> = Pin<Box<dyn Stream<Item = Result<E, Error>> + Send + 'static>>;
|
||||
pub type ProviderChunkStream =
|
||||
Pin<Box<dyn Stream<Item = Result<ProviderStreamChunk, Error>> + Send + 'static>>;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum StreamTransport {
|
||||
Http,
|
||||
WebSocket,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum StreamProviderId {
|
||||
Anthropic,
|
||||
AzureAi,
|
||||
BedrockConverse,
|
||||
OpenAi,
|
||||
}
|
||||
|
||||
impl TryFrom<&str> for StreamProviderId {
|
||||
type Error = Error;
|
||||
|
||||
fn try_from(value: &str) -> Result<Self, Self::Error> {
|
||||
match value {
|
||||
"anthropic" => Ok(Self::Anthropic),
|
||||
"azure_ai" => Ok(Self::AzureAi),
|
||||
"bedrock" | "bedrock_converse" => Ok(Self::BedrockConverse),
|
||||
"openai" => Ok(Self::OpenAi),
|
||||
_ => Err(Error::InvalidProvider(value.to_string())),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct JsonObject(pub Map<String, Value>);
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct Header {
|
||||
pub name: String,
|
||||
pub value: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Default, PartialEq)]
|
||||
pub struct ProviderCredentials {
|
||||
api_key: Option<String>,
|
||||
aws_access_key_id: Option<String>,
|
||||
aws_secret_access_key: Option<String>,
|
||||
aws_session_token: Option<String>,
|
||||
}
|
||||
|
||||
impl ProviderCredentials {
|
||||
pub fn new(
|
||||
api_key: Option<String>,
|
||||
aws_access_key_id: Option<String>,
|
||||
aws_secret_access_key: Option<String>,
|
||||
aws_session_token: Option<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
api_key,
|
||||
aws_access_key_id,
|
||||
aws_secret_access_key,
|
||||
aws_session_token,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn api_key(&self) -> Option<&str> {
|
||||
self.api_key.as_deref()
|
||||
}
|
||||
|
||||
pub fn aws_access_key_id(&self) -> Option<&str> {
|
||||
self.aws_access_key_id.as_deref()
|
||||
}
|
||||
|
||||
pub fn aws_secret_access_key(&self) -> Option<&str> {
|
||||
self.aws_secret_access_key.as_deref()
|
||||
}
|
||||
|
||||
pub fn aws_session_token(&self) -> Option<&str> {
|
||||
self.aws_session_token.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
/// ```compile_fail
|
||||
/// fn assert_serialize<T: serde::Serialize>() {}
|
||||
/// assert_serialize::<litellm_core::streaming::ProviderCredentials>();
|
||||
/// assert_serialize::<litellm_core::streaming::StreamTarget>();
|
||||
/// assert_serialize::<litellm_core::streaming::StreamTransportOptions>();
|
||||
/// ```
|
||||
///
|
||||
/// ```compile_fail
|
||||
/// use litellm_core::streaming::{JsonObject, ProviderCredentials, StreamProviderId, StreamTarget};
|
||||
/// let mut target = StreamTarget::new(
|
||||
/// StreamProviderId::OpenAi,
|
||||
/// ProviderCredentials::default(),
|
||||
/// None,
|
||||
/// );
|
||||
/// target.metadata = JsonObject::default();
|
||||
/// ```
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct StreamTarget {
|
||||
provider: StreamProviderId,
|
||||
credentials: ProviderCredentials,
|
||||
api_base: Option<String>,
|
||||
}
|
||||
|
||||
impl StreamTarget {
|
||||
pub fn new(
|
||||
provider: StreamProviderId,
|
||||
credentials: ProviderCredentials,
|
||||
api_base: Option<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn provider(&self) -> StreamProviderId {
|
||||
self.provider
|
||||
}
|
||||
|
||||
pub fn credentials(&self) -> &ProviderCredentials {
|
||||
&self.credentials
|
||||
}
|
||||
|
||||
pub fn api_base(&self) -> Option<&str> {
|
||||
self.api_base.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default, PartialEq)]
|
||||
pub struct StreamTransportOptions {
|
||||
forwarded_headers: Vec<Header>,
|
||||
timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
impl StreamTransportOptions {
|
||||
pub fn new(forwarded_headers: Vec<Header>, timeout: Option<Duration>) -> Self {
|
||||
Self {
|
||||
forwarded_headers,
|
||||
timeout,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn forwarded_headers(&self) -> &[Header] {
|
||||
&self.forwarded_headers
|
||||
}
|
||||
|
||||
pub fn timeout(&self) -> Option<Duration> {
|
||||
self.timeout
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct StreamMetadata {
|
||||
pub status_code: u16,
|
||||
pub provider: StreamProviderId,
|
||||
pub transport: StreamTransport,
|
||||
pub response_headers: Vec<Header>,
|
||||
}
|
||||
|
||||
pub struct OpenedStream<E> {
|
||||
pub metadata: StreamMetadata,
|
||||
pub events: EventStream<E>,
|
||||
}
|
||||
|
||||
pub struct OpenedWireStream {
|
||||
pub metadata: StreamMetadata,
|
||||
pub chunks: ProviderChunkStream,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct ProviderStreamChunk(Vec<u8>);
|
||||
|
||||
impl ProviderStreamChunk {
|
||||
pub fn new(bytes: impl Into<Vec<u8>>) -> Self {
|
||||
Self(bytes.into())
|
||||
}
|
||||
|
||||
pub fn as_bytes(&self) -> &[u8] {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
pub trait StreamDecoder: Send + 'static {
|
||||
type WireEvent: Send + 'static;
|
||||
|
||||
fn push(&mut self, chunk: ProviderStreamChunk) -> Result<Vec<Self::WireEvent>, Error>;
|
||||
|
||||
fn finish(&mut self) -> Result<Vec<Self::WireEvent>, Error> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
}
|
||||
|
||||
pub trait StreamProvider<R, E>: Send + Sync + 'static {
|
||||
type PreparedRequest: Send + 'static;
|
||||
type WireEvent: Send + 'static;
|
||||
type Decoder: StreamDecoder<WireEvent = Self::WireEvent>;
|
||||
|
||||
fn transform_request(&self, request: R) -> Result<Self::PreparedRequest, Error>;
|
||||
|
||||
fn call(
|
||||
&'static self,
|
||||
request: Self::PreparedRequest,
|
||||
) -> BoxFuture<'static, Result<OpenedWireStream, Error>>;
|
||||
|
||||
fn decoder(&self) -> Self::Decoder;
|
||||
|
||||
fn normalize(&self, event: Self::WireEvent) -> Result<Vec<E>, Error>;
|
||||
}
|
||||
|
||||
struct PipelineState<P: 'static, D, E, R>
|
||||
where
|
||||
D: StreamDecoder,
|
||||
{
|
||||
provider: &'static P,
|
||||
decoder: D,
|
||||
chunks: ProviderChunkStream,
|
||||
pending: VecDeque<Result<E, Error>>,
|
||||
finished: bool,
|
||||
request: PhantomData<fn(R)>,
|
||||
}
|
||||
|
||||
pub async fn open_provider_stream<P, R, E>(
|
||||
provider: &'static P,
|
||||
request: R,
|
||||
) -> Result<OpenedStream<E>, Error>
|
||||
where
|
||||
P: StreamProvider<R, E>,
|
||||
R: Send + 'static,
|
||||
E: Send + 'static,
|
||||
{
|
||||
let prepared = provider.transform_request(request)?;
|
||||
let opened = provider.call(prepared).await?;
|
||||
let state = PipelineState {
|
||||
provider,
|
||||
decoder: provider.decoder(),
|
||||
chunks: opened.chunks,
|
||||
pending: VecDeque::new(),
|
||||
finished: false,
|
||||
request: PhantomData,
|
||||
};
|
||||
let events = stream::unfold(state, |mut state| async move {
|
||||
loop {
|
||||
if let Some(event) = state.pending.pop_front() {
|
||||
return Some((event, state));
|
||||
}
|
||||
if state.finished {
|
||||
return None;
|
||||
}
|
||||
match state.chunks.next().await {
|
||||
Some(Ok(chunk)) => match state.decoder.push(chunk) {
|
||||
Ok(events) => queue_normalized(&mut state, events),
|
||||
Err(error) => {
|
||||
state.finished = true;
|
||||
return Some((Err(error), state));
|
||||
}
|
||||
},
|
||||
Some(Err(error)) => {
|
||||
state.finished = true;
|
||||
return Some((Err(error), state));
|
||||
}
|
||||
None => {
|
||||
state.finished = true;
|
||||
match state.decoder.finish() {
|
||||
Ok(events) => queue_normalized(&mut state, events),
|
||||
Err(error) => return Some((Err(error), state)),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
Ok(OpenedStream {
|
||||
metadata: opened.metadata,
|
||||
events: Box::pin(events),
|
||||
})
|
||||
}
|
||||
|
||||
fn queue_normalized<P, D, E, R>(state: &mut PipelineState<P, D, E, R>, events: Vec<D::WireEvent>)
|
||||
where
|
||||
P: StreamProvider<R, E, Decoder = D>,
|
||||
D: StreamDecoder<WireEvent = <P as StreamProvider<R, E>>::WireEvent>,
|
||||
{
|
||||
for event in events {
|
||||
match state.provider.normalize(event) {
|
||||
Ok(normalized) => state.pending.extend(normalized.into_iter().map(Ok)),
|
||||
Err(error) => {
|
||||
state.pending.push_back(Err(error));
|
||||
state.finished = true;
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use futures_util::future::FutureExt;
|
||||
|
||||
use super::*;
|
||||
|
||||
struct FakeRequest;
|
||||
struct PreparedRequest;
|
||||
|
||||
struct FakeProvider {
|
||||
calls: Arc<Mutex<Vec<&'static str>>>,
|
||||
}
|
||||
|
||||
struct FakeDecoder {
|
||||
calls: Arc<Mutex<Vec<&'static str>>>,
|
||||
pending: String,
|
||||
}
|
||||
|
||||
impl StreamDecoder for FakeDecoder {
|
||||
type WireEvent = String;
|
||||
|
||||
fn push(&mut self, chunk: ProviderStreamChunk) -> Result<Vec<Self::WireEvent>, Error> {
|
||||
self.calls.lock().expect("call log").push("decode");
|
||||
self.pending
|
||||
.push_str(std::str::from_utf8(chunk.as_bytes()).expect("test utf-8"));
|
||||
let mut parts = self
|
||||
.pending
|
||||
.split('|')
|
||||
.map(str::to_string)
|
||||
.collect::<Vec<_>>();
|
||||
self.pending = parts.pop().expect("split always returns one item");
|
||||
Ok(parts)
|
||||
}
|
||||
|
||||
fn finish(&mut self) -> Result<Vec<Self::WireEvent>, Error> {
|
||||
if self.pending.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
Ok(vec![std::mem::take(&mut self.pending)])
|
||||
}
|
||||
}
|
||||
|
||||
impl StreamProvider<FakeRequest, String> for FakeProvider {
|
||||
type PreparedRequest = PreparedRequest;
|
||||
type WireEvent = String;
|
||||
type Decoder = FakeDecoder;
|
||||
|
||||
fn transform_request(&self, _request: FakeRequest) -> Result<Self::PreparedRequest, Error> {
|
||||
self.calls.lock().expect("call log").push("transform");
|
||||
Ok(PreparedRequest)
|
||||
}
|
||||
|
||||
fn call(
|
||||
&'static self,
|
||||
_request: Self::PreparedRequest,
|
||||
) -> BoxFuture<'static, Result<OpenedWireStream, Error>> {
|
||||
self.calls.lock().expect("call log").push("call");
|
||||
async move {
|
||||
Ok(OpenedWireStream {
|
||||
metadata: StreamMetadata {
|
||||
status_code: 200,
|
||||
provider: StreamProviderId::Anthropic,
|
||||
transport: StreamTransport::Http,
|
||||
response_headers: vec![Header {
|
||||
name: "x-test".to_string(),
|
||||
value: "ready".to_string(),
|
||||
}],
|
||||
},
|
||||
chunks: Box::pin(stream::iter([
|
||||
Ok(ProviderStreamChunk::new(b"one|tw".to_vec())),
|
||||
Ok(ProviderStreamChunk::new(b"o|three".to_vec())),
|
||||
])),
|
||||
})
|
||||
}
|
||||
.boxed()
|
||||
}
|
||||
|
||||
fn decoder(&self) -> Self::Decoder {
|
||||
FakeDecoder {
|
||||
calls: self.calls.clone(),
|
||||
pending: String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize(&self, event: Self::WireEvent) -> Result<Vec<String>, Error> {
|
||||
self.calls.lock().expect("call log").push("normalize");
|
||||
Ok(vec![event.to_uppercase()])
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fake_provider_proves_pipeline_order_and_fragmentation() {
|
||||
let calls = Arc::new(Mutex::new(Vec::new()));
|
||||
let provider = Box::leak(Box::new(FakeProvider {
|
||||
calls: calls.clone(),
|
||||
}));
|
||||
let mut opened = open_provider_stream(provider, FakeRequest)
|
||||
.await
|
||||
.expect("stream opens");
|
||||
let events = opened
|
||||
.events
|
||||
.by_ref()
|
||||
.collect::<Vec<_>>()
|
||||
.await
|
||||
.into_iter()
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.expect("events normalize");
|
||||
|
||||
assert_eq!(events, ["ONE", "TWO", "THREE"]);
|
||||
assert_eq!(opened.metadata.response_headers[0].name, "x-test");
|
||||
assert_eq!(
|
||||
*calls.lock().expect("call log"),
|
||||
[
|
||||
"transform",
|
||||
"call",
|
||||
"decode",
|
||||
"normalize",
|
||||
"decode",
|
||||
"normalize",
|
||||
"normalize",
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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",
|
||||
];
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
206
litellm-rust/crates/python-bridge/src/routes/receiver.rs
Normal file
206
litellm-rust/crates/python-bridge/src/routes/receiver.rs
Normal file
|
|
@ -0,0 +1,206 @@
|
|||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
|
||||
use futures_util::{Stream, StreamExt, pin_mut};
|
||||
use litellm_core::Error;
|
||||
use tokio::sync::{Mutex, mpsc};
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
const BRIDGE_CHANNEL_CAPACITY: usize = 1;
|
||||
|
||||
struct ReceiverState<T> {
|
||||
receiver: Mutex<mpsc::Receiver<Result<T, Error>>>,
|
||||
reading: AtomicBool,
|
||||
closed: AtomicBool,
|
||||
producer: std::sync::Mutex<Option<JoinHandle<()>>>,
|
||||
}
|
||||
|
||||
impl<T> Drop for ReceiverState<T> {
|
||||
fn drop(&mut self) {
|
||||
if let Ok(producer) = self.producer.get_mut()
|
||||
&& let Some(producer) = producer.take()
|
||||
{
|
||||
producer.abort();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(super) struct BridgeReceiver<T> {
|
||||
state: Arc<ReceiverState<T>>,
|
||||
}
|
||||
|
||||
struct ReadGuard<'a>(&'a AtomicBool);
|
||||
|
||||
impl Drop for ReadGuard<'_> {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(false, Ordering::Release);
|
||||
}
|
||||
}
|
||||
|
||||
impl<T: Send + 'static> BridgeReceiver<T> {
|
||||
pub(super) fn from_stream<S>(stream: S) -> Self
|
||||
where
|
||||
S: Stream<Item = Result<T, Error>> + Send + 'static,
|
||||
{
|
||||
let (sender, receiver) = mpsc::channel(BRIDGE_CHANNEL_CAPACITY);
|
||||
let producer = pyo3_async_runtimes::tokio::get_runtime().spawn(async move {
|
||||
pin_mut!(stream);
|
||||
while let Some(item) = stream.next().await {
|
||||
let terminal = item.is_err();
|
||||
if sender.send(item).await.is_err() || terminal {
|
||||
return;
|
||||
}
|
||||
}
|
||||
});
|
||||
Self {
|
||||
state: Arc::new(ReceiverState {
|
||||
receiver: Mutex::new(receiver),
|
||||
reading: AtomicBool::new(false),
|
||||
closed: AtomicBool::new(false),
|
||||
producer: std::sync::Mutex::new(Some(producer)),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn next(&self) -> Result<Option<T>, Error> {
|
||||
if self.state.closed.load(Ordering::Acquire) {
|
||||
return Ok(None);
|
||||
}
|
||||
if self.state.reading.swap(true, Ordering::AcqRel) {
|
||||
return Err(Error::InvalidRequest(
|
||||
"native stream does not support concurrent reads".to_string(),
|
||||
));
|
||||
}
|
||||
let _guard = ReadGuard(&self.state.reading);
|
||||
match self.state.receiver.lock().await.recv().await {
|
||||
Some(Ok(item)) => Ok(Some(item)),
|
||||
Some(Err(error)) => Err(error),
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn close(&self) {
|
||||
if self.state.closed.swap(true, Ordering::AcqRel) {
|
||||
return;
|
||||
}
|
||||
if let Ok(mut producer) = self.state.producer.lock()
|
||||
&& let Some(producer) = producer.take()
|
||||
{
|
||||
producer.abort();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_util::stream;
|
||||
|
||||
use super::*;
|
||||
|
||||
struct DropFlag(Arc<AtomicBool>);
|
||||
|
||||
impl Drop for DropFlag {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn receiver_preserves_items_and_terminal_error() {
|
||||
let receiver = BridgeReceiver::from_stream(stream::iter([
|
||||
Ok(vec![1_u8]),
|
||||
Err(Error::Network("broken".to_string())),
|
||||
Ok(vec![2_u8]),
|
||||
]));
|
||||
|
||||
assert_eq!(receiver.next().await.expect("first item"), Some(vec![1]));
|
||||
assert!(matches!(
|
||||
receiver.next().await,
|
||||
Err(Error::Network(message)) if message == "broken"
|
||||
));
|
||||
assert_eq!(receiver.next().await.expect("closed after error"), None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn close_unblocks_a_pending_read() {
|
||||
let receiver = BridgeReceiver::<Vec<u8>>::from_stream(stream::pending());
|
||||
let pending = {
|
||||
let receiver = receiver.clone();
|
||||
tokio::spawn(async move { receiver.next().await })
|
||||
};
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
receiver.close();
|
||||
|
||||
assert_eq!(
|
||||
tokio::time::timeout(Duration::from_secs(1), pending)
|
||||
.await
|
||||
.expect("read should unblock")
|
||||
.expect("task should finish")
|
||||
.expect("close is clean"),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn capacity_one_stops_the_producer_from_draining_the_source() {
|
||||
let polls = Arc::new(AtomicUsize::new(0));
|
||||
let source_polls = polls.clone();
|
||||
let source = stream::poll_fn(move |_| {
|
||||
let item = source_polls.fetch_add(1, Ordering::SeqCst);
|
||||
std::task::Poll::Ready(Some(Ok(item)))
|
||||
});
|
||||
let receiver = BridgeReceiver::from_stream(source);
|
||||
|
||||
tokio::time::sleep(Duration::from_millis(20)).await;
|
||||
assert_eq!(polls.load(Ordering::SeqCst), 2);
|
||||
assert_eq!(receiver.next().await.expect("first item"), Some(0));
|
||||
tokio::time::sleep(Duration::from_millis(20)).await;
|
||||
assert_eq!(polls.load(Ordering::SeqCst), 3);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_reads_are_rejected() {
|
||||
let receiver = BridgeReceiver::<Vec<u8>>::from_stream(stream::pending());
|
||||
let pending = {
|
||||
let receiver = receiver.clone();
|
||||
tokio::spawn(async move { receiver.next().await })
|
||||
};
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
|
||||
assert!(matches!(
|
||||
receiver.next().await,
|
||||
Err(Error::InvalidRequest(message))
|
||||
if message == "native stream does not support concurrent reads"
|
||||
));
|
||||
receiver.close();
|
||||
pending
|
||||
.await
|
||||
.expect("pending read task")
|
||||
.expect("clean close");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn dropping_the_last_receiver_cancels_the_producer() {
|
||||
let dropped = Arc::new(AtomicBool::new(false));
|
||||
let producer_dropped = dropped.clone();
|
||||
let source = stream::once(async move {
|
||||
let _flag = DropFlag(producer_dropped);
|
||||
std::future::pending::<Result<Vec<u8>, Error>>().await
|
||||
});
|
||||
let receiver = BridgeReceiver::from_stream(source);
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
|
||||
drop(receiver);
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while !dropped.load(Ordering::SeqCst) {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("producer future should be dropped");
|
||||
}
|
||||
}
|
||||
|
|
@ -65,6 +65,45 @@ where
|
|||
})
|
||||
}
|
||||
|
||||
pub(super) fn run_sync_with<T, F, C>(py: Python<'_>, future: F, convert: C) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
T: Send + 'static,
|
||||
F: Future<Output = Result<T, Error>> + Send + 'static,
|
||||
C: FnOnce(Python<'_>, T) -> PyResult<Py<PyAny>>,
|
||||
{
|
||||
if Handle::try_current().is_ok() {
|
||||
return Err(PyRuntimeError::new_err(
|
||||
"synchronous native routes cannot run from a Tokio context; use the async route",
|
||||
));
|
||||
}
|
||||
|
||||
let result = release_gil(py, move || {
|
||||
pyo3_async_runtimes::tokio::get_runtime().block_on(wait_for_sync_result(future))
|
||||
})?;
|
||||
let result = map_core_result(result, crate::errors::fallback_route_error_to_pyerr)?;
|
||||
std::panic::catch_unwind(AssertUnwindSafe(|| convert(py, result))).map_err(panic_to_pyerr)?
|
||||
}
|
||||
|
||||
pub(super) fn run_async_with<T, F, C>(
|
||||
py: Python<'_>,
|
||||
future: F,
|
||||
convert: C,
|
||||
) -> PyResult<Bound<'_, PyAny>>
|
||||
where
|
||||
T: Send + 'static,
|
||||
F: Future<Output = Result<T, Error>> + Send + 'static,
|
||||
C: FnOnce(Python<'_>, T) -> PyResult<Py<PyAny>> + Send + 'static,
|
||||
{
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
let result = catch_route_panic(future).await?;
|
||||
let result = map_core_result(result, crate::errors::fallback_route_error_to_pyerr)?;
|
||||
Python::attach(|py| {
|
||||
std::panic::catch_unwind(AssertUnwindSafe(|| convert(py, result)))
|
||||
.map_err(panic_to_pyerr)?
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -> PyResult<T> {
|
||||
match result {
|
||||
Ok(value) => Ok(value),
|
||||
|
|
|
|||
541
litellm-rust/crates/python-bridge/src/routes/streaming.rs
Normal file
541
litellm-rust/crates/python-bridge/src/routes/streaming.rs
Normal file
|
|
@ -0,0 +1,541 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_core::chat_completions::chat_completions_stream as run_chat_completions_stream;
|
||||
use litellm_core::chat_completions::types::{
|
||||
ChatCompletionsStreamRequest, ChatCompletionsStreamRequestBody, ChatStreamEvent,
|
||||
};
|
||||
use litellm_core::messages::messages_event_stream as run_messages_stream;
|
||||
use litellm_core::messages::types::{
|
||||
AnthropicMessagesRequest, MessagesStreamEvent, MessagesStreamRequest,
|
||||
};
|
||||
use litellm_core::responses::responses_stream as run_responses_stream;
|
||||
use litellm_core::responses::responses_websocket as run_responses_websocket;
|
||||
use litellm_core::responses::types::{
|
||||
ResponseCommand, ResponsesStreamEvent, ResponsesStreamRequest, ResponsesStreamRequestBody,
|
||||
ResponsesWebSocketRequest,
|
||||
};
|
||||
use litellm_core::responses::websocket::TypedResponsesWebSocketSession;
|
||||
use litellm_core::streaming::{
|
||||
JsonObject, OpenedStream, ProviderCredentials, StreamMetadata, StreamProviderId, StreamTarget,
|
||||
StreamTransportOptions,
|
||||
};
|
||||
use litellm_python_interop::{from_py, to_py};
|
||||
use pyo3::exceptions::{PyStopAsyncIteration, PyStopIteration};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyModule, PyType};
|
||||
use serde::de::DeserializeOwned;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::marshal::{marshal_headers, optional_timeout};
|
||||
use crate::routes::receiver::BridgeReceiver;
|
||||
use crate::routes::runtime::{run_async, run_async_with, run_sync_with};
|
||||
|
||||
struct TypedEventReceiver<E> {
|
||||
metadata: StreamMetadata,
|
||||
receiver: BridgeReceiver<E>,
|
||||
}
|
||||
|
||||
impl<E> TypedEventReceiver<E>
|
||||
where
|
||||
E: Send + 'static,
|
||||
{
|
||||
fn from_opened(opened: OpenedStream<E>) -> Self {
|
||||
Self {
|
||||
metadata: opened.metadata,
|
||||
receiver: BridgeReceiver::from_stream(opened.events),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn next_event<E>(
|
||||
py: Python<'_>,
|
||||
receiver: BridgeReceiver<E>,
|
||||
stop_iteration: bool,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
E: Serialize + Send + 'static,
|
||||
{
|
||||
run_sync_with(
|
||||
py,
|
||||
async move { receiver.next().await },
|
||||
move |py, event| match event {
|
||||
Some(event) => to_py(py, &event),
|
||||
None if stop_iteration => Err(PyStopIteration::new_err(())),
|
||||
None => Ok(py.None()),
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn anext_event<E>(
|
||||
py: Python<'_>,
|
||||
receiver: BridgeReceiver<E>,
|
||||
stop_iteration: bool,
|
||||
) -> PyResult<Bound<'_, PyAny>>
|
||||
where
|
||||
E: Serialize + Send + 'static,
|
||||
{
|
||||
run_async_with(
|
||||
py,
|
||||
async move { receiver.next().await },
|
||||
move |py, event| match event {
|
||||
Some(event) => to_py(py, &event),
|
||||
None if stop_iteration => Err(PyStopAsyncIteration::new_err(())),
|
||||
None => Ok(py.None()),
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
macro_rules! event_stream_class {
|
||||
($class:ident, $event:ty) => {
|
||||
#[pyclass]
|
||||
struct $class {
|
||||
inner: TypedEventReceiver<$event>,
|
||||
}
|
||||
|
||||
impl From<OpenedStream<$event>> for $class {
|
||||
fn from(opened: OpenedStream<$event>) -> Self {
|
||||
Self {
|
||||
inner: TypedEventReceiver::from_opened(opened),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl $class {
|
||||
#[getter]
|
||||
fn metadata(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
|
||||
to_py(py, &self.inner.metadata)
|
||||
}
|
||||
|
||||
fn next_event(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
|
||||
next_event(py, self.inner.receiver.clone(), false)
|
||||
}
|
||||
|
||||
fn anext_event<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
anext_event(py, self.inner.receiver.clone(), false)
|
||||
}
|
||||
|
||||
fn __iter__(slf: PyRef<'_, Self>) -> PyRef<'_, Self> {
|
||||
slf
|
||||
}
|
||||
|
||||
fn __next__(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
|
||||
next_event(py, self.inner.receiver.clone(), true)
|
||||
}
|
||||
|
||||
fn __aiter__(slf: PyRef<'_, Self>) -> PyRef<'_, Self> {
|
||||
slf
|
||||
}
|
||||
|
||||
fn __anext__<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
anext_event(py, self.inner.receiver.clone(), true)
|
||||
}
|
||||
|
||||
fn close(&self) {
|
||||
self.inner.receiver.close();
|
||||
}
|
||||
|
||||
fn aclose<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
let receiver = self.inner.receiver.clone();
|
||||
run_async(
|
||||
py,
|
||||
async move {
|
||||
receiver.close();
|
||||
Ok(())
|
||||
},
|
||||
crate::errors::fallback_route_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
event_stream_class!(ChatCompletionsEventStream, ChatStreamEvent);
|
||||
event_stream_class!(MessagesEventStream, MessagesStreamEvent);
|
||||
event_stream_class!(ResponsesEventStream, ResponsesStreamEvent);
|
||||
|
||||
#[derive(Default, Deserialize)]
|
||||
struct PythonProviderCredentials {
|
||||
api_key: Option<String>,
|
||||
aws_access_key_id: Option<String>,
|
||||
aws_secret_access_key: Option<String>,
|
||||
aws_session_token: Option<String>,
|
||||
}
|
||||
|
||||
struct PythonStreamTarget {
|
||||
provider: String,
|
||||
credentials: Option<Py<PyAny>>,
|
||||
api_base: Option<String>,
|
||||
}
|
||||
|
||||
struct PythonStreamTransport {
|
||||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
}
|
||||
|
||||
fn parse_call<B>(
|
||||
py: Python<'_>,
|
||||
request: Py<PyAny>,
|
||||
target: PythonStreamTarget,
|
||||
transport: PythonStreamTransport,
|
||||
) -> PyResult<(B, StreamTarget, StreamTransportOptions)>
|
||||
where
|
||||
B: DeserializeOwned,
|
||||
{
|
||||
let body = from_py(request.bind(py))?;
|
||||
let credentials = target
|
||||
.credentials
|
||||
.map(|value| from_py::<PythonProviderCredentials>(value.bind(py)))
|
||||
.transpose()?
|
||||
.unwrap_or_default();
|
||||
let provider = StreamProviderId::try_from(target.provider.as_str())
|
||||
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?;
|
||||
let target = StreamTarget::new(
|
||||
provider,
|
||||
ProviderCredentials::new(
|
||||
credentials.api_key,
|
||||
credentials.aws_access_key_id,
|
||||
credentials.aws_secret_access_key,
|
||||
credentials.aws_session_token,
|
||||
),
|
||||
target.api_base,
|
||||
);
|
||||
let forwarded_headers = marshal_headers(py, transport.extra_headers)?
|
||||
.into_iter()
|
||||
.map(|(name, value)| litellm_core::streaming::Header { name, value })
|
||||
.collect();
|
||||
let transport = StreamTransportOptions::new(
|
||||
forwarded_headers,
|
||||
optional_timeout(transport.timeout_seconds),
|
||||
);
|
||||
Ok((body, target, transport))
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn chat_completions_stream(
|
||||
py: Python<'_>,
|
||||
request: Py<PyAny>,
|
||||
provider: String,
|
||||
credentials: Option<Py<PyAny>>,
|
||||
api_base: Option<String>,
|
||||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let (body, target, transport) = parse_call::<ChatCompletionsStreamRequestBody>(
|
||||
py,
|
||||
request,
|
||||
PythonStreamTarget {
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
},
|
||||
PythonStreamTransport {
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
run_sync_with(
|
||||
py,
|
||||
async move {
|
||||
run_chat_completions_stream(ChatCompletionsStreamRequest {
|
||||
body,
|
||||
target,
|
||||
transport,
|
||||
})
|
||||
.await
|
||||
},
|
||||
|py, opened| Ok(Py::new(py, ChatCompletionsEventStream::from(opened))?.into_any()),
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn achat_completions_stream(
|
||||
py: Python<'_>,
|
||||
request: Py<PyAny>,
|
||||
provider: String,
|
||||
credentials: Option<Py<PyAny>>,
|
||||
api_base: Option<String>,
|
||||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let (body, target, transport) = parse_call::<ChatCompletionsStreamRequestBody>(
|
||||
py,
|
||||
request,
|
||||
PythonStreamTarget {
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
},
|
||||
PythonStreamTransport {
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
run_async_with(
|
||||
py,
|
||||
async move {
|
||||
run_chat_completions_stream(ChatCompletionsStreamRequest {
|
||||
body,
|
||||
target,
|
||||
transport,
|
||||
})
|
||||
.await
|
||||
},
|
||||
|py, opened| Ok(Py::new(py, ChatCompletionsEventStream::from(opened))?.into_any()),
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn messages_stream(
|
||||
py: Python<'_>,
|
||||
request: Py<PyAny>,
|
||||
provider: String,
|
||||
credentials: Option<Py<PyAny>>,
|
||||
api_base: Option<String>,
|
||||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let (body, target, transport) = parse_call::<AnthropicMessagesRequest>(
|
||||
py,
|
||||
request,
|
||||
PythonStreamTarget {
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
},
|
||||
PythonStreamTransport {
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
run_sync_with(
|
||||
py,
|
||||
async move {
|
||||
run_messages_stream(MessagesStreamRequest {
|
||||
body,
|
||||
target,
|
||||
transport,
|
||||
})
|
||||
.await
|
||||
},
|
||||
|py, opened| Ok(Py::new(py, MessagesEventStream::from(opened))?.into_any()),
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn amessages_stream(
|
||||
py: Python<'_>,
|
||||
request: Py<PyAny>,
|
||||
provider: String,
|
||||
credentials: Option<Py<PyAny>>,
|
||||
api_base: Option<String>,
|
||||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let (body, target, transport) = parse_call::<AnthropicMessagesRequest>(
|
||||
py,
|
||||
request,
|
||||
PythonStreamTarget {
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
},
|
||||
PythonStreamTransport {
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
run_async_with(
|
||||
py,
|
||||
async move {
|
||||
run_messages_stream(MessagesStreamRequest {
|
||||
body,
|
||||
target,
|
||||
transport,
|
||||
})
|
||||
.await
|
||||
},
|
||||
|py, opened| Ok(Py::new(py, MessagesEventStream::from(opened))?.into_any()),
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn responses_stream(
|
||||
py: Python<'_>,
|
||||
request: Py<PyAny>,
|
||||
provider: String,
|
||||
credentials: Option<Py<PyAny>>,
|
||||
api_base: Option<String>,
|
||||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let (body, target, transport) = parse_call::<ResponsesStreamRequestBody>(
|
||||
py,
|
||||
request,
|
||||
PythonStreamTarget {
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
},
|
||||
PythonStreamTransport {
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
run_sync_with(
|
||||
py,
|
||||
async move {
|
||||
run_responses_stream(ResponsesStreamRequest {
|
||||
body,
|
||||
target,
|
||||
transport,
|
||||
})
|
||||
.await
|
||||
},
|
||||
|py, opened| Ok(Py::new(py, ResponsesEventStream::from(opened))?.into_any()),
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (request, provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn aresponses_stream(
|
||||
py: Python<'_>,
|
||||
request: Py<PyAny>,
|
||||
provider: String,
|
||||
credentials: Option<Py<PyAny>>,
|
||||
api_base: Option<String>,
|
||||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let (body, target, transport) = parse_call::<ResponsesStreamRequestBody>(
|
||||
py,
|
||||
request,
|
||||
PythonStreamTarget {
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
},
|
||||
PythonStreamTransport {
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
run_async_with(
|
||||
py,
|
||||
async move {
|
||||
run_responses_stream(ResponsesStreamRequest {
|
||||
body,
|
||||
target,
|
||||
transport,
|
||||
})
|
||||
.await
|
||||
},
|
||||
|py, opened| Ok(Py::new(py, ResponsesEventStream::from(opened))?.into_any()),
|
||||
)
|
||||
}
|
||||
|
||||
#[pyclass]
|
||||
struct ResponsesWebSocketSession {
|
||||
session: Arc<dyn TypedResponsesWebSocketSession>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl ResponsesWebSocketSession {
|
||||
#[classmethod]
|
||||
#[pyo3(signature = (provider, credentials=None, api_base=None, extra_headers=None, timeout_seconds=None))]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn connect<'py>(
|
||||
_cls: &Bound<'py, PyType>,
|
||||
py: Python<'py>,
|
||||
provider: String,
|
||||
credentials: Option<Py<PyAny>>,
|
||||
api_base: Option<String>,
|
||||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let empty_request = pyo3::types::PyDict::new(py).unbind().into_any();
|
||||
let (_, target, transport) = parse_call::<JsonObject>(
|
||||
py,
|
||||
empty_request,
|
||||
PythonStreamTarget {
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
},
|
||||
PythonStreamTransport {
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
run_async_with(
|
||||
py,
|
||||
async move {
|
||||
run_responses_websocket(ResponsesWebSocketRequest { target, transport }).await
|
||||
},
|
||||
|py, session| {
|
||||
Ok(Py::new(
|
||||
py,
|
||||
ResponsesWebSocketSession {
|
||||
session: Arc::from(session),
|
||||
},
|
||||
)?
|
||||
.into_any())
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn send_event<'py>(&self, py: Python<'py>, command: Py<PyAny>) -> PyResult<Bound<'py, PyAny>> {
|
||||
let command: ResponseCommand = from_py(command.bind(py))?;
|
||||
let session = self.session.clone();
|
||||
run_async(
|
||||
py,
|
||||
async move { session.send(command).await },
|
||||
crate::errors::fallback_route_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
|
||||
fn recv_event<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
let session = self.session.clone();
|
||||
run_async_with(
|
||||
py,
|
||||
async move { session.recv().await },
|
||||
|py, event| match event {
|
||||
Some(event) => to_py(py, &event),
|
||||
None => Ok(py.None()),
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn close<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
let session = self.session.clone();
|
||||
run_async(
|
||||
py,
|
||||
async move { session.close().await },
|
||||
crate::errors::fallback_route_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
module.add_class::<ChatCompletionsEventStream>()?;
|
||||
module.add_class::<MessagesEventStream>()?;
|
||||
module.add_class::<ResponsesEventStream>()?;
|
||||
module.add_class::<ResponsesWebSocketSession>()?;
|
||||
module.add_function(wrap_pyfunction!(chat_completions_stream, module)?)?;
|
||||
module.add_function(wrap_pyfunction!(achat_completions_stream, module)?)?;
|
||||
module.add_function(wrap_pyfunction!(messages_stream, module)?)?;
|
||||
module.add_function(wrap_pyfunction!(amessages_stream, module)?)?;
|
||||
module.add_function(wrap_pyfunction!(responses_stream, module)?)?;
|
||||
module.add_function(wrap_pyfunction!(aresponses_stream, module)?)
|
||||
}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
431
litellm/rust_bridge/streaming.py
Normal file
431
litellm/rust_bridge/streaming.py
Normal file
|
|
@ -0,0 +1,431 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Generator, Iterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal, NoReturn, Protocol, TypeAlias, runtime_checkable
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.rust_bridge.loader import get_native_bridge
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponsesAPIStreamingResponse
|
||||
|
||||
StreamApi: TypeAlias = Literal["chat_completions", "messages", "responses"]
|
||||
StreamTransport: TypeAlias = Literal["http", "websocket"]
|
||||
Event: TypeAlias = Mapping[str, object]
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class RustEventStream(Protocol):
|
||||
@property
|
||||
def metadata(self) -> Mapping[str, object]: ...
|
||||
|
||||
def next_event(self) -> Event | None: ...
|
||||
|
||||
async def anext_event(self) -> Event | None: ...
|
||||
|
||||
def close(self) -> None: ...
|
||||
|
||||
async def aclose(self) -> None: ...
|
||||
|
||||
|
||||
class RustStreamOpen(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
request: Mapping[str, object],
|
||||
provider: str,
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> RustEventStream: ...
|
||||
|
||||
|
||||
class RustAsyncStreamOpen(Protocol):
|
||||
async def __call__(
|
||||
self,
|
||||
request: Mapping[str, object],
|
||||
provider: str,
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> RustEventStream: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ObjectAwaitable(Protocol):
|
||||
def __await__(self) -> Generator[object, None, object]: ...
|
||||
|
||||
|
||||
class _Unset:
|
||||
pass
|
||||
|
||||
|
||||
_UNSET: Final = _Unset()
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _RustStreamingState:
|
||||
chat: RustStreamOpen | None = None
|
||||
achat: RustAsyncStreamOpen | None = None
|
||||
messages: RustStreamOpen | None = None
|
||||
amessages: RustAsyncStreamOpen | None = None
|
||||
responses: RustStreamOpen | None = None
|
||||
aresponses: RustAsyncStreamOpen | None = None
|
||||
|
||||
|
||||
_STATE: Final = _RustStreamingState()
|
||||
|
||||
|
||||
def set_rust_streaming(
|
||||
*,
|
||||
chat: RustStreamOpen | None | _Unset = _UNSET,
|
||||
achat: RustAsyncStreamOpen | None | _Unset = _UNSET,
|
||||
messages: RustStreamOpen | None | _Unset = _UNSET,
|
||||
amessages: RustAsyncStreamOpen | None | _Unset = _UNSET,
|
||||
responses: RustStreamOpen | None | _Unset = _UNSET,
|
||||
aresponses: RustAsyncStreamOpen | None | _Unset = _UNSET,
|
||||
) -> None:
|
||||
if not isinstance(chat, _Unset):
|
||||
_STATE.chat = chat
|
||||
if not isinstance(achat, _Unset):
|
||||
_STATE.achat = achat
|
||||
if not isinstance(messages, _Unset):
|
||||
_STATE.messages = messages
|
||||
if not isinstance(amessages, _Unset):
|
||||
_STATE.amessages = amessages
|
||||
if not isinstance(responses, _Unset):
|
||||
_STATE.responses = responses
|
||||
if not isinstance(aresponses, _Unset):
|
||||
_STATE.aresponses = aresponses
|
||||
|
||||
|
||||
def _native_attribute(name: str) -> object | None:
|
||||
native: Final = get_native_bridge()
|
||||
return None if native is None else getattr(native, name, None)
|
||||
|
||||
|
||||
def _native_sync_opener(name: str) -> RustStreamOpen | None:
|
||||
opener: Final = _native_attribute(name)
|
||||
if not callable(opener):
|
||||
return None
|
||||
|
||||
def open_stream(
|
||||
request: Mapping[str, object],
|
||||
provider: str,
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> RustEventStream:
|
||||
stream: Final = opener(
|
||||
request,
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
)
|
||||
if not isinstance(stream, RustEventStream):
|
||||
raise TypeError("native stream opener returned an invalid stream")
|
||||
return stream
|
||||
|
||||
return open_stream
|
||||
|
||||
|
||||
def _native_async_opener(name: str) -> RustAsyncStreamOpen | None:
|
||||
opener: Final = _native_attribute(name)
|
||||
if not callable(opener):
|
||||
return None
|
||||
|
||||
async def open_stream(
|
||||
request: Mapping[str, object],
|
||||
provider: str,
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> RustEventStream:
|
||||
pending: Final = opener(
|
||||
request,
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
)
|
||||
if not isinstance(pending, ObjectAwaitable):
|
||||
raise TypeError("native async stream opener returned a non-awaitable")
|
||||
stream: Final[object] = await pending
|
||||
if not isinstance(stream, RustEventStream):
|
||||
raise TypeError("native async stream opener returned an invalid stream")
|
||||
return stream
|
||||
|
||||
return open_stream
|
||||
|
||||
|
||||
def _sync_opener(api: StreamApi) -> RustStreamOpen | None:
|
||||
match api:
|
||||
case "chat_completions":
|
||||
return _STATE.chat or _native_sync_opener("chat_completions_stream")
|
||||
case "messages":
|
||||
return _STATE.messages or _native_sync_opener("messages_stream")
|
||||
case "responses":
|
||||
return _STATE.responses or _native_sync_opener("responses_stream")
|
||||
|
||||
|
||||
def _async_opener(api: StreamApi) -> RustAsyncStreamOpen | None:
|
||||
match api:
|
||||
case "chat_completions":
|
||||
return _STATE.achat or _native_async_opener("achat_completions_stream")
|
||||
case "messages":
|
||||
return _STATE.amessages or _native_async_opener("amessages_stream")
|
||||
case "responses":
|
||||
return _STATE.aresponses or _native_async_opener("aresponses_stream")
|
||||
|
||||
|
||||
def _native_exceptions() -> tuple[type[BaseException], type[BaseException]] | None:
|
||||
native: Final = get_native_bridge()
|
||||
if native is None:
|
||||
return None
|
||||
declined: Final = getattr(native, "RustBridgeDeclined", None)
|
||||
upstream: Final = getattr(native, "RustUpstreamError", None)
|
||||
if not isinstance(declined, type) or not isinstance(upstream, type):
|
||||
return None
|
||||
return declined, upstream
|
||||
|
||||
|
||||
def _handle_open_error(error: Exception, provider: str) -> None:
|
||||
exceptions: Final = _native_exceptions()
|
||||
if exceptions is not None and isinstance(error, exceptions[0]):
|
||||
return
|
||||
_raise_stream_error(error, provider)
|
||||
|
||||
|
||||
def _raise_stream_error(error: Exception, provider: str) -> NoReturn:
|
||||
exceptions: Final = _native_exceptions()
|
||||
if exceptions is None or not isinstance(error, exceptions[1]):
|
||||
raise error
|
||||
args: Final[tuple[object, ...]] = error.args
|
||||
status_value: Final = args[0] if args else 0
|
||||
message_value: Final = args[1] if len(args) > 1 else str(error)
|
||||
status: Final = status_value if isinstance(status_value, int) else 0
|
||||
message: Final = message_value if isinstance(message_value, str) else str(message_value)
|
||||
raise APIError(
|
||||
status_code=status or 500,
|
||||
message=f"litellm rust typed stream: {message}",
|
||||
llm_provider=provider,
|
||||
model="",
|
||||
) from error
|
||||
|
||||
|
||||
class TypedEventStreamAdapter:
|
||||
def __init__(self, stream: RustEventStream, provider: str) -> None:
|
||||
self._stream: Final = stream
|
||||
self._provider: Final = provider
|
||||
self.metadata: Final = stream.metadata
|
||||
self._mode: Literal["sync", "async"] | None = None
|
||||
|
||||
def _claim(self, mode: Literal["sync", "async"]) -> None:
|
||||
if self._mode is None:
|
||||
self._mode = mode
|
||||
return
|
||||
if self._mode != mode:
|
||||
raise RuntimeError("native stream cannot mix synchronous and asynchronous consumption")
|
||||
|
||||
def __iter__(self) -> Iterator[Event]:
|
||||
self._claim("sync")
|
||||
return self
|
||||
|
||||
def __next__(self) -> Event:
|
||||
self._claim("sync")
|
||||
event: Final = self._next_event()
|
||||
if event is None:
|
||||
raise StopIteration
|
||||
return event
|
||||
|
||||
def _next_event(self) -> Event | None:
|
||||
try:
|
||||
return self._stream.next_event()
|
||||
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
|
||||
_raise_stream_error(error, self._provider)
|
||||
|
||||
def __aiter__(self) -> AsyncIterator[Event]:
|
||||
self._claim("async")
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> Event:
|
||||
self._claim("async")
|
||||
try:
|
||||
event: Final = await self._stream.anext_event()
|
||||
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
|
||||
_raise_stream_error(error, self._provider)
|
||||
if event is None:
|
||||
raise StopAsyncIteration
|
||||
return event
|
||||
|
||||
def close(self) -> None:
|
||||
self._stream.close()
|
||||
|
||||
async def aclose(self) -> None:
|
||||
await self._stream.aclose()
|
||||
|
||||
|
||||
class MessagesSseStreamAdapter:
|
||||
def __init__(self, events: TypedEventStreamAdapter) -> None:
|
||||
self._events: Final = events
|
||||
self.metadata: Final = events.metadata
|
||||
|
||||
def __iter__(self) -> Iterator[bytes]:
|
||||
return (_event_to_sse(event) for event in self._events)
|
||||
|
||||
async def __aiter__(self) -> AsyncIterator[bytes]:
|
||||
async for event in self._events:
|
||||
yield _event_to_sse(event)
|
||||
|
||||
def close(self) -> None:
|
||||
self._events.close()
|
||||
|
||||
async def aclose(self) -> None:
|
||||
await self._events.aclose()
|
||||
|
||||
|
||||
class ResponsesSdkEventStreamAdapter:
|
||||
def __init__(self, events: TypedEventStreamAdapter) -> None:
|
||||
self._events: Final = events
|
||||
self.metadata: Final = events.metadata
|
||||
|
||||
def __iter__(self) -> Iterator[ResponsesAPIStreamingResponse]:
|
||||
return (_responses_event_to_sdk(event) for event in self._events)
|
||||
|
||||
async def __aiter__(self) -> AsyncIterator[ResponsesAPIStreamingResponse]:
|
||||
async for event in self._events:
|
||||
yield _responses_event_to_sdk(event)
|
||||
|
||||
def close(self) -> None:
|
||||
self._events.close()
|
||||
|
||||
async def aclose(self) -> None:
|
||||
await self._events.aclose()
|
||||
|
||||
|
||||
def _responses_event_to_sdk(event: Event) -> ResponsesAPIStreamingResponse:
|
||||
from litellm.types.llms.openai import GenericEvent
|
||||
|
||||
event_type: Final = event.get("type")
|
||||
model: Final = _responses_event_models().get(event_type) if isinstance(event_type, str) else None
|
||||
return (model or GenericEvent).model_validate(event)
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _responses_event_models() -> Mapping[str, type[BaseLiteLLMOpenAIResponseObject]]:
|
||||
from litellm.types.llms import openai as openai_types
|
||||
|
||||
return MappingProxyType(
|
||||
{
|
||||
"response.created": openai_types.ResponseCreatedEvent,
|
||||
"response.in_progress": openai_types.ResponseInProgressEvent,
|
||||
"response.completed": openai_types.ResponseCompletedEvent,
|
||||
"response.failed": openai_types.ResponseFailedEvent,
|
||||
"response.incomplete": openai_types.ResponseIncompleteEvent,
|
||||
"response.reasoning_summary_part.added": openai_types.ResponsePartAddedEvent,
|
||||
"response.reasoning_summary_text.delta": openai_types.ReasoningSummaryTextDeltaEvent,
|
||||
"response.reasoning_summary_text.done": openai_types.ReasoningSummaryTextDoneEvent,
|
||||
"response.reasoning_summary_part.done": openai_types.ReasoningSummaryPartDoneEvent,
|
||||
"response.output_item.added": openai_types.OutputItemAddedEvent,
|
||||
"response.output_item.done": openai_types.OutputItemDoneEvent,
|
||||
"response.content_part.added": openai_types.ContentPartAddedEvent,
|
||||
"response.content_part.done": openai_types.ContentPartDoneEvent,
|
||||
"response.output_text.delta": openai_types.OutputTextDeltaEvent,
|
||||
"response.output_text.annotation.added": openai_types.OutputTextAnnotationAddedEvent,
|
||||
"response.output_text.done": openai_types.OutputTextDoneEvent,
|
||||
"response.refusal.delta": openai_types.RefusalDeltaEvent,
|
||||
"response.refusal.done": openai_types.RefusalDoneEvent,
|
||||
"response.function_call_arguments.delta": openai_types.FunctionCallArgumentsDeltaEvent,
|
||||
"response.function_call_arguments.done": openai_types.FunctionCallArgumentsDoneEvent,
|
||||
"response.file_search_call.in_progress": openai_types.FileSearchCallInProgressEvent,
|
||||
"response.file_search_call.searching": openai_types.FileSearchCallSearchingEvent,
|
||||
"response.file_search_call.completed": openai_types.FileSearchCallCompletedEvent,
|
||||
"response.web_search_call.in_progress": openai_types.WebSearchCallInProgressEvent,
|
||||
"response.web_search_call.searching": openai_types.WebSearchCallSearchingEvent,
|
||||
"response.web_search_call.completed": openai_types.WebSearchCallCompletedEvent,
|
||||
"response.mcp_list_tools.in_progress": openai_types.MCPListToolsInProgressEvent,
|
||||
"response.mcp_list_tools.completed": openai_types.MCPListToolsCompletedEvent,
|
||||
"response.mcp_list_tools.failed": openai_types.MCPListToolsFailedEvent,
|
||||
"response.mcp_call.in_progress": openai_types.MCPCallInProgressEvent,
|
||||
"response.mcp_call_arguments.delta": openai_types.MCPCallArgumentsDeltaEvent,
|
||||
"response.mcp_call_arguments.done": openai_types.MCPCallArgumentsDoneEvent,
|
||||
"response.mcp_call.completed": openai_types.MCPCallCompletedEvent,
|
||||
"response.mcp_call.failed": openai_types.MCPCallFailedEvent,
|
||||
"image_generation.partial_image": openai_types.ImageGenerationPartialImageEvent,
|
||||
"error": openai_types.ErrorEvent,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _event_to_sse(event: Event) -> bytes:
|
||||
payload: Final = dict(event) # mutable-ok: JSON requires a concrete dict
|
||||
return f"data: {json.dumps(payload, separators=(',', ':'))}\n\n".encode()
|
||||
|
||||
|
||||
def open_stream(
|
||||
*,
|
||||
api: StreamApi,
|
||||
provider: str,
|
||||
request: Mapping[str, object],
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> TypedEventStreamAdapter | None:
|
||||
opener: Final = _sync_opener(api)
|
||||
if opener is None:
|
||||
return None
|
||||
try:
|
||||
stream: Final = opener(
|
||||
request,
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
extra_headers,
|
||||
timeout_to_seconds(timeout),
|
||||
)
|
||||
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
|
||||
_handle_open_error(error, provider)
|
||||
return None
|
||||
return TypedEventStreamAdapter(stream, provider)
|
||||
|
||||
|
||||
async def aopen_stream(
|
||||
*,
|
||||
api: StreamApi,
|
||||
provider: str,
|
||||
request: Mapping[str, object],
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> TypedEventStreamAdapter | None:
|
||||
opener: Final = _async_opener(api)
|
||||
if opener is None:
|
||||
return None
|
||||
try:
|
||||
stream: Final = await opener(
|
||||
request,
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
extra_headers,
|
||||
timeout_to_seconds(timeout),
|
||||
)
|
||||
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
|
||||
_handle_open_error(error, provider)
|
||||
return None
|
||||
return TypedEventStreamAdapter(stream, provider)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
58
tests/test_litellm/rust_bridge/test_native_integration.py
Normal file
58
tests/test_litellm/rust_bridge/test_native_integration.py
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from importlib import import_module
|
||||
from typing import Final, Protocol, cast
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
class StreamRoute(Protocol):
|
||||
def __call__(self, *, request: object, provider: str) -> object: ...
|
||||
|
||||
|
||||
class ResponsesWebSocketSession(Protocol):
|
||||
@classmethod
|
||||
def connect(cls, *, provider: str) -> object: ...
|
||||
|
||||
|
||||
class NativeBridge(Protocol):
|
||||
RustBridgeDeclined: type[Exception]
|
||||
ResponsesWebSocketSession: type[ResponsesWebSocketSession]
|
||||
chat_completions_stream: StreamRoute
|
||||
achat_completions_stream: StreamRoute
|
||||
messages_stream: StreamRoute
|
||||
amessages_stream: StreamRoute
|
||||
responses_stream: StreamRoute
|
||||
aresponses_stream: StreamRoute
|
||||
|
||||
|
||||
try:
|
||||
native: Final = cast(NativeBridge, import_module("litellm.rust_bridge._native"))
|
||||
except ImportError:
|
||||
if os.getenv("LITELLM_REQUIRE_NATIVE_BRIDGE") == "1":
|
||||
raise
|
||||
pytest.skip("the native bridge has not been built", allow_module_level=True)
|
||||
|
||||
STREAM_ROUTES: Final[tuple[tuple[StreamRoute, str], ...]] = (
|
||||
(native.chat_completions_stream, "anthropic"),
|
||||
(native.achat_completions_stream, "anthropic"),
|
||||
(native.messages_stream, "anthropic"),
|
||||
(native.amessages_stream, "anthropic"),
|
||||
(native.responses_stream, "openai"),
|
||||
(native.aresponses_stream, "openai"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("function", "provider"), STREAM_ROUTES)
|
||||
def test_stream_routes_preserve_validation_errors(
|
||||
function: StreamRoute,
|
||||
provider: str,
|
||||
) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
function(request=object(), provider=provider)
|
||||
|
||||
|
||||
def test_native_websocket_connect_preserves_decline() -> None:
|
||||
with pytest.raises(native.RustBridgeDeclined):
|
||||
native.ResponsesWebSocketSession.connect(provider="unsupported")
|
||||
292
tests/test_litellm/rust_bridge/test_streaming.py
Normal file
292
tests/test_litellm/rust_bridge/test_streaming.py
Normal file
|
|
@ -0,0 +1,292 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterator, Mapping
|
||||
from typing import Final
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.rust_bridge import streaming
|
||||
|
||||
|
||||
class _FakeEventStream:
|
||||
def __init__(self, events: tuple[Mapping[str, object], ...]) -> None:
|
||||
self.metadata: Final = {
|
||||
"status_code": 200,
|
||||
"provider": "anthropic",
|
||||
"transport": "http",
|
||||
"response_headers": [{"name": "x-test", "value": "ready"}],
|
||||
}
|
||||
self._events: Final = iter(events)
|
||||
self.closed = False
|
||||
|
||||
def next_event(self) -> Mapping[str, object] | None:
|
||||
if self.closed:
|
||||
return None
|
||||
return next(self._events, None)
|
||||
|
||||
async def anext_event(self) -> Mapping[str, object] | None:
|
||||
return self.next_event()
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.close()
|
||||
|
||||
|
||||
class _RecordingOpen:
|
||||
def __init__(self, events: tuple[Mapping[str, object], ...]) -> None:
|
||||
self._events: Final = events
|
||||
self.calls = 0
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
request: Mapping[str, object],
|
||||
provider: str,
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> _FakeEventStream:
|
||||
self.calls += 1
|
||||
return _FakeEventStream(self._events)
|
||||
|
||||
|
||||
class _RecordingAsyncOpen:
|
||||
def __init__(self, events: tuple[Mapping[str, object], ...]) -> None:
|
||||
self._events: Final = events
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
request: Mapping[str, object],
|
||||
provider: str,
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> _FakeEventStream:
|
||||
self.calls += 1
|
||||
return _FakeEventStream(self._events)
|
||||
|
||||
|
||||
class _Declined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _Upstream(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _NativeErrors:
|
||||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = _Upstream
|
||||
|
||||
|
||||
class _FailingOpen:
|
||||
def __init__(self, error: Exception) -> None:
|
||||
self._error: Final = error
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
request: Mapping[str, object],
|
||||
provider: str,
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> _FakeEventStream:
|
||||
raise self._error
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_bridge() -> Iterator[None]:
|
||||
streaming.set_rust_streaming(
|
||||
chat=None,
|
||||
achat=None,
|
||||
messages=None,
|
||||
amessages=None,
|
||||
responses=None,
|
||||
aresponses=None,
|
||||
)
|
||||
yield
|
||||
streaming.set_rust_streaming(
|
||||
chat=None,
|
||||
achat=None,
|
||||
messages=None,
|
||||
amessages=None,
|
||||
responses=None,
|
||||
aresponses=None,
|
||||
)
|
||||
|
||||
|
||||
def _chat_event(text: str) -> Mapping[str, object]:
|
||||
return {
|
||||
"text": text,
|
||||
"tool_use": None,
|
||||
"is_finished": False,
|
||||
"finish_reason": "",
|
||||
"usage": None,
|
||||
}
|
||||
|
||||
|
||||
def test_unavailable_native_bridge_falls_back(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(streaming, "get_native_bridge", lambda: None)
|
||||
result: Final = streaming.open_stream(
|
||||
api="chat_completions",
|
||||
provider="anthropic",
|
||||
request={"model": "claude", "messages": []},
|
||||
credentials={"api_key": "test"},
|
||||
api_base=None,
|
||||
extra_headers=None,
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_declined_open_failure_falls_back(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
python_calls = 0
|
||||
|
||||
monkeypatch.setattr(streaming, "get_native_bridge", lambda: _NativeErrors())
|
||||
streaming.set_rust_streaming(
|
||||
chat=_FailingOpen(_Declined("unsupported request")),
|
||||
)
|
||||
|
||||
result: Final = streaming.open_stream(
|
||||
api="chat_completions",
|
||||
provider="anthropic",
|
||||
request={"model": "claude", "messages": []},
|
||||
credentials=None,
|
||||
api_base=None,
|
||||
extra_headers=None,
|
||||
timeout=None,
|
||||
)
|
||||
if result is None:
|
||||
python_calls += 1
|
||||
|
||||
assert python_calls == 1
|
||||
|
||||
|
||||
def test_upstream_open_failure_never_falls_back(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(streaming, "get_native_bridge", lambda: _NativeErrors())
|
||||
streaming.set_rust_streaming(
|
||||
chat=_FailingOpen(_Upstream(503, "connection closed after request")),
|
||||
)
|
||||
|
||||
with pytest.raises(APIError, match="connection closed after request"):
|
||||
streaming.open_stream(
|
||||
api="chat_completions",
|
||||
provider="anthropic",
|
||||
request={"model": "claude", "messages": []},
|
||||
credentials=None,
|
||||
api_base=None,
|
||||
extra_headers=None,
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
|
||||
def test_sync_typed_events_preserve_shape_metadata_and_close() -> None:
|
||||
opener: Final = _RecordingOpen((_chat_event("one"), _chat_event("two")))
|
||||
streaming.set_rust_streaming(chat=opener)
|
||||
|
||||
result: Final = streaming.open_stream(
|
||||
api="chat_completions",
|
||||
provider="anthropic",
|
||||
request={"model": "claude", "messages": []},
|
||||
credentials=None,
|
||||
api_base=None,
|
||||
extra_headers=None,
|
||||
timeout=1.0,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert tuple(event["text"] for event in result) == ("one", "two")
|
||||
assert result.metadata["provider"] == "anthropic"
|
||||
result.close()
|
||||
assert opener.calls == 1
|
||||
|
||||
|
||||
def test_chat_events_flow_through_custom_stream_wrapper() -> None:
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.types.utils import ModelResponseStream
|
||||
|
||||
native: Final = _FakeEventStream((_chat_event("hello"),))
|
||||
events: Final = streaming.TypedEventStreamAdapter(native, "anthropic")
|
||||
wrapper: Final = CustomStreamWrapper(
|
||||
completion_stream=events,
|
||||
model="claude",
|
||||
logging_obj=MagicMock(),
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
|
||||
chunk: Final = next(wrapper)
|
||||
assert isinstance(chunk, ModelResponseStream)
|
||||
assert chunk.choices[0].delta.content == "hello"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_typed_events_and_cancellation() -> None:
|
||||
opener: Final = _RecordingAsyncOpen((_chat_event("one"), _chat_event("two")))
|
||||
streaming.set_rust_streaming(achat=opener)
|
||||
result: Final = await streaming.aopen_stream(
|
||||
api="chat_completions",
|
||||
provider="anthropic",
|
||||
request={"model": "claude", "messages": []},
|
||||
credentials=None,
|
||||
api_base=None,
|
||||
extra_headers=None,
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
collected: Final = tuple([event async for event in result])
|
||||
assert tuple(event["text"] for event in collected) == ("one", "two")
|
||||
await result.aclose()
|
||||
|
||||
|
||||
def test_messages_events_are_wrapped_in_existing_sse_bytes() -> None:
|
||||
native: Final = _FakeEventStream(({"type": "message_stop"},))
|
||||
events: Final = streaming.TypedEventStreamAdapter(native, "anthropic")
|
||||
messages: Final = streaming.MessagesSseStreamAdapter(events)
|
||||
|
||||
assert tuple(messages) == (b'data: {"type":"message_stop"}\n\n',)
|
||||
|
||||
|
||||
def test_responses_events_are_validated_into_existing_sdk_objects() -> None:
|
||||
from litellm.types.llms.openai import OutputTextDeltaEvent
|
||||
|
||||
native: Final = _FakeEventStream(
|
||||
(
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"item_id": "item_1",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": "hello",
|
||||
},
|
||||
)
|
||||
)
|
||||
events: Final = streaming.TypedEventStreamAdapter(native, "openai")
|
||||
responses: Final = streaming.ResponsesSdkEventStreamAdapter(events)
|
||||
|
||||
event: Final = next(iter(responses))
|
||||
assert isinstance(event, OutputTextDeltaEvent)
|
||||
assert event.type == "response.output_text.delta"
|
||||
assert event.delta == "hello"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rejects_mixed_sync_and_async_consumption() -> None:
|
||||
events: Final = streaming.TypedEventStreamAdapter(_FakeEventStream((_chat_event("one"),)), "anthropic")
|
||||
assert tuple(events) == (_chat_event("one"),)
|
||||
|
||||
with pytest.raises(RuntimeError, match="cannot mix"):
|
||||
async for _ in events:
|
||||
pass
|
||||
Loading…
Add table
Reference in a new issue