mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #22771 from BerriAI/litellm_responses_websocket_2
Add support for responses websocket for all providers
This commit is contained in:
commit
23d312dbd2
17 changed files with 1509 additions and 202 deletions
|
|
@ -14,6 +14,7 @@ Requests to /chat/completions may be bridged here automatically when the provide
|
|||
| Logging | ✅ | Works across all integrations |
|
||||
| End-user Tracking | ✅ | |
|
||||
| Streaming | ✅ | |
|
||||
| WebSocket Mode | ✅ | Lower-latency persistent connections for all providers |
|
||||
| Image Generation Streaming | ✅ | Progressive image generation with partial images (1-3) |
|
||||
| Fallbacks | ✅ | Works between supported models |
|
||||
| Loadbalancing | ✅ | Works between supported models |
|
||||
|
|
@ -810,6 +811,245 @@ for event in response:
|
|||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## WebSocket Mode
|
||||
|
||||
The Responses API supports **WebSocket mode** for lower-latency, persistent connections ideal for agentic workflows. WebSocket mode works with **all LiteLLM providers**, not just those with native WebSocket support.
|
||||
|
||||
### Architecture
|
||||
|
||||
LiteLLM provides two WebSocket modes:
|
||||
|
||||
1. **Native WebSocket**: Direct `wss://` connection to providers that support it (OpenAI, Azure)
|
||||
2. **Managed WebSocket**: HTTP streaming over WebSocket for all other providers (Anthropic, Gemini, Bedrock, etc.)
|
||||
|
||||
The system automatically selects the appropriate mode based on provider capabilities.
|
||||
|
||||
### Usage
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python (websocket-client)">
|
||||
|
||||
```python showLineNumbers title="WebSocket with Python"
|
||||
import json
|
||||
from websocket import create_connection # pip install websocket-client
|
||||
|
||||
# Connect to LiteLLM proxy WebSocket endpoint
|
||||
ws = create_connection(
|
||||
"ws://localhost:4000/v1/responses?model=gemini-2.5-flash",
|
||||
header=["Authorization: Bearer sk-1234"]
|
||||
)
|
||||
|
||||
try:
|
||||
# Send initial message
|
||||
ws.send(json.dumps({
|
||||
"type": "response.create",
|
||||
"model": "gemini-2.5-flash",
|
||||
"store": True,
|
||||
"input": [{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "My favorite color is blue."}]
|
||||
}]
|
||||
}))
|
||||
|
||||
# Collect response events
|
||||
response_id = None
|
||||
while True:
|
||||
event = json.loads(ws.recv())
|
||||
print(f"Event: {event['type']}")
|
||||
|
||||
if event["type"] == "response.completed":
|
||||
response_id = event["response"]["id"]
|
||||
break
|
||||
elif event["type"] == "response.output_text.delta":
|
||||
print(f"Text: {event.get('delta', '')}", end="", flush=True)
|
||||
|
||||
print(f"\nResponse ID: {response_id}")
|
||||
|
||||
# Send follow-up with previous_response_id for multi-turn
|
||||
ws.send(json.dumps({
|
||||
"type": "response.create",
|
||||
"model": "gemini-2.5-flash",
|
||||
"previous_response_id": response_id,
|
||||
"input": [{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "What is my favorite color?"}]
|
||||
}]
|
||||
}))
|
||||
|
||||
# Collect follow-up response
|
||||
while True:
|
||||
event = json.loads(ws.recv())
|
||||
if event["type"] == "response.completed":
|
||||
break
|
||||
elif event["type"] == "response.output_text.delta":
|
||||
print(event.get("delta", ""), end="", flush=True)
|
||||
|
||||
finally:
|
||||
ws.close()
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="javascript" label="JavaScript (ws)">
|
||||
|
||||
```javascript showLineNumbers title="WebSocket with JavaScript"
|
||||
const WebSocket = require('ws'); // npm install ws
|
||||
|
||||
const ws = new WebSocket(
|
||||
'ws://localhost:4000/v1/responses?model=gemini-2.5-flash',
|
||||
{
|
||||
headers: {
|
||||
'Authorization': 'Bearer sk-1234'
|
||||
}
|
||||
}
|
||||
);
|
||||
|
||||
ws.on('open', () => {
|
||||
// Send initial message
|
||||
ws.send(JSON.stringify({
|
||||
type: 'response.create',
|
||||
model: 'gemini-2.5-flash',
|
||||
store: true,
|
||||
input: [{
|
||||
type: 'message',
|
||||
role: 'user',
|
||||
content: [{ type: 'input_text', text: 'My favorite color is blue.' }]
|
||||
}]
|
||||
}));
|
||||
});
|
||||
|
||||
let responseId = null;
|
||||
|
||||
ws.on('message', (data) => {
|
||||
const event = JSON.parse(data.toString());
|
||||
console.log(`Event: ${event.type}`);
|
||||
|
||||
if (event.type === 'response.completed') {
|
||||
responseId = event.response.id;
|
||||
console.log(`Response ID: ${responseId}`);
|
||||
|
||||
// Send follow-up
|
||||
ws.send(JSON.stringify({
|
||||
type: 'response.create',
|
||||
model: 'gemini-2.5-flash',
|
||||
previous_response_id: responseId,
|
||||
input: [{
|
||||
type: 'message',
|
||||
role: 'user',
|
||||
content: [{ type: 'input_text', text: 'What is my favorite color?' }]
|
||||
}]
|
||||
}));
|
||||
} else if (event.type === 'response.output_text.delta') {
|
||||
process.stdout.write(event.delta || '');
|
||||
}
|
||||
});
|
||||
|
||||
ws.on('error', (error) => {
|
||||
console.error('WebSocket error:', error);
|
||||
});
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="curl" label="curl (websocat)">
|
||||
|
||||
```bash showLineNumbers title="WebSocket with websocat"
|
||||
# Install websocat: brew install websocat (macOS) or cargo install websocat
|
||||
|
||||
# Connect to WebSocket endpoint
|
||||
websocat "ws://localhost:4000/v1/responses?model=gemini-2.5-flash" \
|
||||
-H="Authorization: Bearer sk-1234"
|
||||
|
||||
# Then send JSON events (paste and press Enter):
|
||||
{"type":"response.create","model":"gemini-2.5-flash","input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"Hello!"}]}]}
|
||||
|
||||
# You'll receive streaming events back:
|
||||
# {"type":"response.created",...}
|
||||
# {"type":"response.in_progress",...}
|
||||
# {"type":"response.output_text.delta","delta":"Hello",...}
|
||||
# {"type":"response.completed",...}
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Event Types
|
||||
|
||||
WebSocket connections receive Server-Sent Events (SSE) formatted as JSON:
|
||||
|
||||
| Event Type | Description |
|
||||
|------------|-------------|
|
||||
| `response.created` | Response generation started |
|
||||
| `response.in_progress` | Response is being generated |
|
||||
| `response.output_item.added` | New output item (message, tool call, etc.) added |
|
||||
| `response.output_text.delta` | Incremental text chunk |
|
||||
| `response.output_text.done` | Text output completed |
|
||||
| `response.content_part.done` | Content part completed |
|
||||
| `response.output_item.done` | Output item completed |
|
||||
| `response.completed` | Full response completed successfully |
|
||||
| `response.failed` | Response generation failed |
|
||||
| `response.incomplete` | Response incomplete (e.g., max tokens reached) |
|
||||
| `error` | Error occurred |
|
||||
|
||||
### Multi-Turn Conversations
|
||||
|
||||
Use `previous_response_id` to maintain conversation context across multiple WebSocket messages:
|
||||
|
||||
```python showLineNumbers title="Multi-turn WebSocket Conversation"
|
||||
# Turn 1
|
||||
ws.send(json.dumps({
|
||||
"type": "response.create",
|
||||
"model": "gemini-2.5-flash",
|
||||
"store": True, # Required for multi-turn
|
||||
"input": [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": "Hello"}]}]
|
||||
}))
|
||||
|
||||
# ... collect events and get response_id from response.completed event ...
|
||||
|
||||
# Turn 2 - reference previous response
|
||||
ws.send(json.dumps({
|
||||
"type": "response.create",
|
||||
"model": "gemini-2.5-flash",
|
||||
"previous_response_id": response_id, # Links to previous turn
|
||||
"input": [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": "Continue"}]}]
|
||||
}))
|
||||
```
|
||||
|
||||
### Provider Support
|
||||
|
||||
| Provider | WebSocket Mode | Notes |
|
||||
|----------|----------------|-------|
|
||||
| OpenAI | Native | Direct `wss://` connection to OpenAI |
|
||||
| Azure OpenAI | Native | Direct `wss://` connection to Azure |
|
||||
| Anthropic | Managed | HTTP streaming over WebSocket |
|
||||
| Google AI Studio (Gemini) | Managed | HTTP streaming over WebSocket |
|
||||
| Vertex AI | Managed | HTTP streaming over WebSocket |
|
||||
| AWS Bedrock | Managed | HTTP streaming over WebSocket |
|
||||
| All other providers | Managed | HTTP streaming over WebSocket |
|
||||
|
||||
**Note**: Both native and managed modes provide the same event stream format. The difference is transparent to clients.
|
||||
|
||||
### Configuration
|
||||
|
||||
No special configuration needed. WebSocket mode is automatically available on the `/v1/responses` endpoint when accessed via WebSocket protocol (`ws://` or `wss://`).
|
||||
|
||||
For LiteLLM Proxy, ensure your models are configured normally:
|
||||
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
model_list:
|
||||
- model_name: gemini-2.5-flash
|
||||
litellm_params:
|
||||
model: gemini/gemini-2.5-flash
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
|
||||
- model_name: gpt-4o
|
||||
litellm_params:
|
||||
model: openai/gpt-4o
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
```
|
||||
|
||||
Both models will automatically support WebSocket mode at `ws://localhost:4000/v1/responses`.
|
||||
|
||||
## Response ID Security
|
||||
|
||||
By default, LiteLLM Proxy prevents users from accessing other users' response IDs.
|
||||
|
|
|
|||
|
|
@ -218,6 +218,18 @@ class BaseResponsesAPIConfig(ABC):
|
|||
"""Returns True if litellm should fake a stream for the given model and stream value"""
|
||||
return False
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""
|
||||
Returns True if the provider has a native WebSocket endpoint for Responses API.
|
||||
|
||||
Providers with native websocket support can connect directly to wss:// endpoints.
|
||||
Providers without native support will use the ManagedResponsesWebSocketHandler
|
||||
which makes HTTP streaming calls and forwards events over the websocket.
|
||||
|
||||
Default: False (use managed websocket handler)
|
||||
"""
|
||||
return False
|
||||
|
||||
#########################################################
|
||||
########## CANCEL RESPONSE API TRANSFORMATION ##########
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -1,14 +1,14 @@
|
|||
import json
|
||||
from typing import Any, Optional
|
||||
|
||||
from litellm.exceptions import AuthenticationError
|
||||
from litellm.constants import STREAM_SSE_DONE_STRING
|
||||
from litellm.exceptions import AuthenticationError
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
from litellm.llms.openai.common_utils import OpenAIError
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
_safe_convert_created_field,
|
||||
)
|
||||
from litellm.llms.openai.common_utils import OpenAIError
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.types.llms.openai import (
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
|
|
@ -200,3 +200,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
api_base = api_base or self.authenticator.get_api_base() or CHATGPT_API_BASE
|
||||
api_base = api_base.rstrip("/")
|
||||
return f"{api_base}/responses"
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""ChatGPT does not support native WebSocket for Responses API"""
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -4737,20 +4737,46 @@ class BaseLLMHTTPHandler:
|
|||
model: str,
|
||||
websocket: Any,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
responses_api_provider_config: BaseResponsesAPIConfig,
|
||||
responses_api_provider_config: Optional[BaseResponsesAPIConfig],
|
||||
api_base: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
timeout: Optional[float] = None,
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
litellm_metadata: Optional[Dict[str, Any]] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
Handles Responses API WebSocket mode.
|
||||
|
||||
Opens a persistent WebSocket to the provider's /v1/responses endpoint
|
||||
and proxies response.create events bidirectionally for lower-latency
|
||||
agentic workflows.
|
||||
For providers with native websocket support (OpenAI, Azure):
|
||||
- Opens a persistent WebSocket to the provider's /v1/responses endpoint
|
||||
- Proxies response.create events bidirectionally for lower-latency agentic workflows
|
||||
|
||||
For providers without native websocket support (all others):
|
||||
- Uses ManagedResponsesWebSocketHandler which makes HTTP streaming calls
|
||||
- Forwards events over the websocket connection
|
||||
"""
|
||||
if responses_api_provider_config is None or not responses_api_provider_config.supports_native_websocket():
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
handler = ManagedResponsesWebSocketHandler(
|
||||
websocket=websocket,
|
||||
model=model,
|
||||
logging_obj=logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_metadata=litellm_metadata,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
)
|
||||
await handler.run()
|
||||
return
|
||||
|
||||
import websockets
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
|
|
@ -4767,7 +4793,6 @@ class BaseLLMHTTPHandler:
|
|||
api_base=api_base,
|
||||
litellm_params={},
|
||||
)
|
||||
# /responses -> wss:// URL
|
||||
ws_url = http_url.replace("https://", "wss://").replace("http://", "ws://")
|
||||
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -98,3 +98,7 @@ class DatabricksResponsesAPIConfig(DatabricksBase, OpenAIResponsesAPIConfig):
|
|||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""Databricks does not support native WebSocket for Responses API"""
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -22,8 +22,8 @@ from litellm.types.utils import LlmProviders
|
|||
|
||||
from ..authenticator import Authenticator
|
||||
from ..common_utils import (
|
||||
GetAPIKeyError,
|
||||
GITHUB_COPILOT_API_BASE,
|
||||
GetAPIKeyError,
|
||||
get_copilot_default_headers,
|
||||
)
|
||||
|
||||
|
|
@ -329,3 +329,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
)
|
||||
|
||||
return False
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""GitHub Copilot does not support native WebSocket for Responses API"""
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -69,3 +69,7 @@ class HostedVLLMResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
return f"{api_base}/responses"
|
||||
|
||||
return f"{api_base}/v1/responses"
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""Hosted vLLM does not support native WebSocket for Responses API"""
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -46,3 +46,7 @@ class LiteLLMProxyResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
api_base = api_base.rstrip("/")
|
||||
|
||||
return f"{api_base}/responses"
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""LiteLLM Proxy does not support native WebSocket for Responses API"""
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -247,6 +247,10 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
response._hidden_params["headers"] = raw_response_headers
|
||||
return response
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""Manus does not support native WebSocket for Responses API"""
|
||||
return False
|
||||
|
||||
def transform_get_response_api_request(
|
||||
self,
|
||||
response_id: str,
|
||||
|
|
|
|||
|
|
@ -344,6 +344,10 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
)
|
||||
return False
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""OpenAI supports native WebSocket for Responses API"""
|
||||
return True
|
||||
|
||||
#########################################################
|
||||
########## DELETE RESPONSE API TRANSFORMATION ##############
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -75,3 +75,7 @@ class OpenRouterResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
api_base = api_base.rstrip("/")
|
||||
|
||||
return f"{api_base}/responses"
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""OpenRouter does not support native WebSocket for Responses API"""
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -490,3 +490,7 @@ class PerplexityResponsesConfig(OpenAIResponsesAPIConfig):
|
|||
verbose_logger.debug("Failed to transform Perplexity cost object: %s", e)
|
||||
|
||||
return chunk
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""Perplexity does not support native WebSocket for Responses API"""
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -16,16 +16,17 @@ from pydantic import fields as pyd_fields
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.llms.openai import ResponseInputParam, ResponsesAPIStreamingResponse
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
_safe_convert_created_field,
|
||||
)
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseInputParam,
|
||||
ResponsesAPIOptionalRequestParams,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamingResponse,
|
||||
)
|
||||
from litellm.types.responses.main import DeleteResponseResult
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -555,3 +556,7 @@ class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
|
||||
# Fall back to the first candidate
|
||||
return candidates[0]
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""VolcEngine does not support native WebSocket for Responses API"""
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -252,3 +252,7 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
|
||||
return f"{api_base}/responses"
|
||||
|
||||
def supports_native_websocket(self) -> bool:
|
||||
"""XAI does not support native WebSocket for Responses API"""
|
||||
return False
|
||||
|
||||
|
|
|
|||
|
|
@ -1729,11 +1729,6 @@ async def _aresponses_websocket(
|
|||
)
|
||||
)
|
||||
|
||||
if responses_api_provider_config is None:
|
||||
raise ValueError(
|
||||
f"Responses API WebSocket mode is not supported for provider: {_custom_llm_provider}"
|
||||
)
|
||||
|
||||
resolved_api_base = (
|
||||
dynamic_api_base
|
||||
or litellm_params.api_base
|
||||
|
|
@ -1748,6 +1743,9 @@ async def _aresponses_websocket(
|
|||
or get_secret_str("OPENAI_API_KEY")
|
||||
)
|
||||
|
||||
# Extract params that we're passing explicitly to avoid duplicates in **kwargs
|
||||
remaining_kwargs = {k: v for k, v in kwargs.items() if k not in {"user_api_key_dict", "litellm_metadata"}}
|
||||
|
||||
await base_llm_http_handler.async_responses_websocket(
|
||||
model=model,
|
||||
websocket=websocket,
|
||||
|
|
@ -1758,4 +1756,6 @@ async def _aresponses_websocket(
|
|||
timeout=timeout,
|
||||
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
||||
litellm_metadata=_build_litellm_metadata_for_ws(kwargs),
|
||||
custom_llm_provider=_custom_llm_provider,
|
||||
**remaining_kwargs,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from litellm.types.llms.openai import (
|
|||
OutputTextDeltaEvent,
|
||||
ResponseAPIUsage,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIRequestParams,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
ResponsesAPIStreamingResponse,
|
||||
|
|
@ -871,37 +872,17 @@ class ResponsesWebSocketStreaming:
|
|||
# Managed WebSocket mode (HTTP-backed, provider-agnostic)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_RESPONSE_CREATE_PARAMS = (
|
||||
"input",
|
||||
"model",
|
||||
"previous_response_id",
|
||||
"instructions",
|
||||
"max_output_tokens",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"temperature",
|
||||
"top_p",
|
||||
"store",
|
||||
"metadata",
|
||||
"truncation",
|
||||
"reasoning",
|
||||
"stream",
|
||||
"include",
|
||||
"parallel_tool_calls",
|
||||
"text",
|
||||
"user",
|
||||
"service_tier",
|
||||
"safety_identifier",
|
||||
"background",
|
||||
_RESPONSE_CREATE_PARAMS: frozenset = (
|
||||
ResponsesAPIRequestParams.__required_keys__ | ResponsesAPIRequestParams.__optional_keys__
|
||||
)
|
||||
|
||||
_MANAGED_WS_SKIP_KWARGS = frozenset(
|
||||
_MANAGED_WS_SKIP_KWARGS: frozenset = frozenset(
|
||||
{
|
||||
"litellm_logging_obj",
|
||||
"litellm_call_id",
|
||||
"aresponses",
|
||||
"_aresponses_websocket",
|
||||
"user_api_key_dict",
|
||||
"litellm_logging_obj",
|
||||
"litellm_call_id",
|
||||
"aresponses",
|
||||
"_aresponses_websocket",
|
||||
"user_api_key_dict",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -949,7 +930,7 @@ class ManagedResponsesWebSocketHandler:
|
|||
self.extra_kwargs: Dict[str, Any] = {
|
||||
k: v for k, v in kwargs.items() if k not in _MANAGED_WS_SKIP_KWARGS
|
||||
}
|
||||
# In-memory session history: response_id → list of input+output messages.
|
||||
# In-memory session history: response_id → full accumulated message list.
|
||||
# Keyed by the DECODED (pre-encoding) response ID from response.completed.
|
||||
# This avoids the async DB-write race condition where spend logs haven't
|
||||
# been committed yet when the next response.create arrives.
|
||||
|
|
@ -982,41 +963,27 @@ class ManagedResponsesWebSocketHandler:
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Core request handler
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _get_history_messages(self, previous_response_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Return accumulated message history for *previous_response_id*.
|
||||
|
||||
Checks the in-memory session store first (fast path, no DB round-trip).
|
||||
The key is the *decoded* response ID (the raw provider response ID before
|
||||
LiteLLM base64-encodes it into the ``resp_...`` format).
|
||||
"""
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(
|
||||
previous_response_id
|
||||
)
|
||||
raw_id = decoded.get("response_id", previous_response_id)
|
||||
return list(self._session_history.get(raw_id, []))
|
||||
|
||||
def _store_history(
|
||||
self,
|
||||
response_id: str,
|
||||
input_messages: List[Dict[str, Any]],
|
||||
output_messages: List[Dict[str, Any]],
|
||||
) -> None:
|
||||
def _store_history(self, response_id: str, messages: List[Dict[str, Any]]) -> None:
|
||||
"""
|
||||
Persist a turn's messages in the in-memory session store.
|
||||
Store the complete accumulated message history for *response_id*.
|
||||
|
||||
*response_id* is the raw (decoded) provider ID extracted from the
|
||||
``response.completed`` event so that the next turn can look it up via
|
||||
:meth:`_get_history_messages`.
|
||||
Replaces any prior value — callers are responsible for passing the full
|
||||
history (prior turns + current input + new output).
|
||||
"""
|
||||
prior: List[Dict[str, Any]] = self._session_history.get(response_id, [])
|
||||
self._session_history[response_id] = prior + input_messages + output_messages
|
||||
self._session_history[response_id] = messages
|
||||
|
||||
@staticmethod
|
||||
def _extract_response_id(completed_event: Dict[str, Any]) -> Optional[str]:
|
||||
|
|
@ -1024,8 +991,6 @@ class ManagedResponsesWebSocketHandler:
|
|||
Pull the raw (decoded) response ID out of a ``response.completed`` event.
|
||||
Returns *None* if the event doesn't contain a usable ID.
|
||||
"""
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
resp_obj = completed_event.get("response", {})
|
||||
encoded_id: Optional[str] = resp_obj.get("id") if isinstance(resp_obj, dict) else None
|
||||
if not encoded_id:
|
||||
|
|
@ -1037,7 +1002,7 @@ class ManagedResponsesWebSocketHandler:
|
|||
def _extract_output_messages(completed_event: Dict[str, Any]) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Convert the output items in a ``response.completed`` event into
|
||||
chat-completion style messages suitable for the next turn's ``input``.
|
||||
Responses API message dicts suitable for the next turn's ``input``.
|
||||
"""
|
||||
resp_obj = completed_event.get("response", {})
|
||||
if not isinstance(resp_obj, dict):
|
||||
|
|
@ -1074,6 +1039,169 @@ class ManagedResponsesWebSocketHandler:
|
|||
return [item for item in input_val if isinstance(item, dict)]
|
||||
return []
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# _process_response_create sub-methods
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _parse_message(self, raw_message: str) -> Optional[Dict[str, Any]]:
|
||||
"""Parse raw WS text; return the message dict or None (JSON error / ignored type)."""
|
||||
try:
|
||||
msg_obj = json.loads(raw_message)
|
||||
except json.JSONDecodeError:
|
||||
await self._send_error("Invalid JSON in response.create event", "invalid_request_error")
|
||||
return None
|
||||
if msg_obj.get("type") != "response.create":
|
||||
# Silently ignore non-response.create messages (e.g. warmup pings)
|
||||
return None
|
||||
return msg_obj
|
||||
|
||||
@staticmethod
|
||||
def _build_base_call_kwargs(msg_obj: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Extract Responses API params from the event, handling both wire formats:
|
||||
Nested: {"type": "response.create", "response": {"input": [...], ...}}
|
||||
Flat: {"type": "response.create", "input": [...], "model": "...", ...}
|
||||
"""
|
||||
nested = msg_obj.get("response")
|
||||
response_params: Dict[str, Any] = (
|
||||
nested
|
||||
if isinstance(nested, dict) and nested
|
||||
else {k: v for k, v in msg_obj.items() if k != "type"}
|
||||
)
|
||||
return {
|
||||
param: response_params[param]
|
||||
for param in _RESPONSE_CREATE_PARAMS
|
||||
if param in response_params and response_params[param] is not None
|
||||
}
|
||||
|
||||
def _apply_history(
|
||||
self,
|
||||
call_kwargs: Dict[str, Any],
|
||||
previous_response_id: Optional[str],
|
||||
current_messages: List[Dict[str, Any]],
|
||||
prior_history: List[Dict[str, Any]],
|
||||
) -> None:
|
||||
"""Prepend in-memory turn history, or fall back to DB-based reconstruction."""
|
||||
if not previous_response_id:
|
||||
return
|
||||
if prior_history:
|
||||
call_kwargs["input"] = prior_history + current_messages
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: prepended %d history messages for previous_response_id=%s",
|
||||
len(prior_history),
|
||||
previous_response_id,
|
||||
)
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: no in-memory history for previous_response_id=%s; "
|
||||
"falling back to DB-based session reconstruction",
|
||||
previous_response_id,
|
||||
)
|
||||
# Fall back to DB-based session reconstruction (may work for
|
||||
# cross-connection multi-turn when spend logs are committed)
|
||||
call_kwargs["previous_response_id"] = previous_response_id
|
||||
|
||||
def _inject_credentials(
|
||||
self, call_kwargs: Dict[str, Any], event_model: Optional[str]
|
||||
) -> None:
|
||||
"""Inject connection-level credentials and metadata into call_kwargs."""
|
||||
if self.api_key is not None:
|
||||
call_kwargs["api_key"] = self.api_key
|
||||
if self.api_base is not None:
|
||||
call_kwargs["api_base"] = self.api_base
|
||||
if self.timeout is not None:
|
||||
call_kwargs["timeout"] = self.timeout
|
||||
# Only propagate custom_llm_provider when no per-request model override exists.
|
||||
# If the payload specifies a different model, let litellm re-resolve the
|
||||
# provider so we don't accidentally force the wrong backend.
|
||||
if self.custom_llm_provider is not None and not event_model:
|
||||
call_kwargs["custom_llm_provider"] = self.custom_llm_provider
|
||||
if self.litellm_metadata:
|
||||
call_kwargs["litellm_metadata"] = dict(self.litellm_metadata)
|
||||
|
||||
@staticmethod
|
||||
def _update_proxy_request(call_kwargs: Dict[str, Any], model: str) -> None:
|
||||
"""Update proxy_server_request body so spend logs record the full request."""
|
||||
proxy_server_request = (call_kwargs.get("litellm_metadata") or {}).get(
|
||||
"proxy_server_request"
|
||||
) or {}
|
||||
if not isinstance(proxy_server_request, dict):
|
||||
return
|
||||
body = dict(proxy_server_request.get("body") or {})
|
||||
body["input"] = call_kwargs.get("input")
|
||||
body["store"] = call_kwargs.get("store")
|
||||
body["model"] = model
|
||||
for k in ("tools", "tool_choice", "instructions", "metadata"):
|
||||
if k in call_kwargs and call_kwargs[k] is not None:
|
||||
body[k] = call_kwargs[k]
|
||||
proxy_server_request = {**proxy_server_request, "body": body}
|
||||
if "litellm_metadata" not in call_kwargs:
|
||||
call_kwargs["litellm_metadata"] = {}
|
||||
call_kwargs["litellm_metadata"]["proxy_server_request"] = proxy_server_request
|
||||
call_kwargs.setdefault("litellm_params", {})
|
||||
call_kwargs["litellm_params"]["proxy_server_request"] = proxy_server_request
|
||||
|
||||
async def _stream_and_forward(
|
||||
self, model: str, call_kwargs: Dict[str, Any]
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Stream ``litellm.aresponses`` and forward every chunk over the WebSocket.
|
||||
|
||||
Captures the ``response.completed`` event type from the chunk object
|
||||
directly (before serialization) to avoid a redundant JSON round-trip on
|
||||
every chunk. Returns the completed event dict, or ``None``.
|
||||
"""
|
||||
completed_event: Optional[Dict[str, Any]] = None
|
||||
stream_response = await litellm.aresponses(model=model, **call_kwargs)
|
||||
async for chunk in stream_response: # type: ignore[union-attr]
|
||||
if chunk is None:
|
||||
continue
|
||||
# Read type from the object before serializing to avoid double JSON parse
|
||||
chunk_type = getattr(chunk, "type", None) or (
|
||||
chunk.get("type") if isinstance(chunk, dict) else None
|
||||
)
|
||||
serialized = self._serialize_chunk(chunk)
|
||||
if serialized is None:
|
||||
continue
|
||||
if chunk_type == "response.completed" and completed_event is None:
|
||||
try:
|
||||
completed_event = json.loads(serialized)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
await self.websocket.send_text(serialized)
|
||||
except Exception as send_exc:
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: error sending chunk to client: %s", send_exc
|
||||
)
|
||||
return completed_event # Client disconnected
|
||||
return completed_event
|
||||
|
||||
def _save_turn_history(
|
||||
self,
|
||||
completed_event: Optional[Dict[str, Any]],
|
||||
prior_history: List[Dict[str, Any]],
|
||||
current_messages: List[Dict[str, Any]],
|
||||
) -> None:
|
||||
"""Store this turn in in-memory history for future previous_response_id lookups."""
|
||||
if completed_event is None:
|
||||
return
|
||||
new_response_id = self._extract_response_id(completed_event)
|
||||
if not new_response_id:
|
||||
return
|
||||
output_msgs = self._extract_output_messages(completed_event)
|
||||
all_messages = prior_history + current_messages + output_msgs
|
||||
self._store_history(new_response_id, all_messages)
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: stored %d messages for response_id=%s",
|
||||
len(all_messages),
|
||||
new_response_id,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Core request handler
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _process_response_create(self, raw_message: str) -> None:
|
||||
"""
|
||||
Parse one ``response.create`` event, call ``litellm.aresponses(stream=True)``,
|
||||
|
|
@ -1094,157 +1222,41 @@ class ManagedResponsesWebSocketHandler:
|
|||
occurs when spend logs haven't been committed by the time the second
|
||||
``response.create`` arrives over the same WebSocket connection.
|
||||
"""
|
||||
import litellm as _litellm
|
||||
|
||||
try:
|
||||
msg_obj = json.loads(raw_message)
|
||||
except json.JSONDecodeError:
|
||||
await self._send_error("Invalid JSON in response.create event", "invalid_request_error")
|
||||
msg_obj = await self._parse_message(raw_message)
|
||||
if msg_obj is None:
|
||||
return
|
||||
|
||||
if msg_obj.get("type") != "response.create":
|
||||
# Silently ignore non-response.create messages (e.g. warmup pings)
|
||||
return
|
||||
|
||||
# Support two wire formats:
|
||||
# Nested : {"type": "response.create", "response": {"input": [...], ...}}
|
||||
# Flat : {"type": "response.create", "input": [...], "model": "...", ...}
|
||||
nested = msg_obj.get("response")
|
||||
if isinstance(nested, dict) and nested:
|
||||
response_params: Dict[str, Any] = nested
|
||||
else:
|
||||
response_params = {k: v for k, v in msg_obj.items() if k != "type"}
|
||||
|
||||
# Build kwargs for aresponses from the response.create payload
|
||||
call_kwargs: Dict[str, Any] = {}
|
||||
for param in _RESPONSE_CREATE_PARAMS:
|
||||
if param in response_params and response_params[param] is not None:
|
||||
call_kwargs[param] = response_params[param]
|
||||
|
||||
# Always stream
|
||||
call_kwargs = self._build_base_call_kwargs(msg_obj)
|
||||
call_kwargs["stream"] = True
|
||||
|
||||
# Use the model from the event if provided, otherwise fall back to the
|
||||
# model supplied at WebSocket connect time.
|
||||
event_model = call_kwargs.pop("model", None)
|
||||
event_model: Optional[str] = call_kwargs.pop("model", None)
|
||||
model = event_model or self.model
|
||||
|
||||
# ---- In-memory multi-turn: prepend history when previous_response_id set ----
|
||||
previous_response_id: Optional[str] = call_kwargs.pop("previous_response_id", None)
|
||||
current_input = call_kwargs.get("input")
|
||||
current_messages = self._input_to_messages(current_input)
|
||||
if previous_response_id:
|
||||
history = self._get_history_messages(previous_response_id)
|
||||
if history:
|
||||
# Prepend history; current messages are the new user turn
|
||||
call_kwargs["input"] = history + current_messages
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: prepended %d history messages for previous_response_id=%s",
|
||||
len(history),
|
||||
previous_response_id,
|
||||
)
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: no in-memory history for previous_response_id=%s; "
|
||||
"falling back to DB-based session reconstruction",
|
||||
previous_response_id,
|
||||
)
|
||||
# Fall back to DB-based session reconstruction (may work for
|
||||
# cross-connection multi-turn when spend logs are committed)
|
||||
call_kwargs["previous_response_id"] = previous_response_id
|
||||
# ---------------------------------------------------------------------------
|
||||
current_messages = self._input_to_messages(call_kwargs.get("input"))
|
||||
|
||||
# Inject connection-level credentials and metadata.
|
||||
# Only propagate custom_llm_provider when the request is using the
|
||||
# same model as the WebSocket connection (i.e. no per-request model
|
||||
# override). If the payload specifies a different model, let litellm
|
||||
# re-resolve the provider from the model name so we don't accidentally
|
||||
# force the wrong backend.
|
||||
if self.api_key is not None:
|
||||
call_kwargs["api_key"] = self.api_key
|
||||
if self.api_base is not None:
|
||||
call_kwargs["api_base"] = self.api_base
|
||||
if self.timeout is not None:
|
||||
call_kwargs["timeout"] = self.timeout
|
||||
if self.custom_llm_provider is not None and not event_model:
|
||||
call_kwargs["custom_llm_provider"] = self.custom_llm_provider
|
||||
if self.litellm_metadata:
|
||||
call_kwargs["litellm_metadata"] = dict(self.litellm_metadata)
|
||||
# Fetch history once; reused in both _apply_history and _save_turn_history
|
||||
prior_history = (
|
||||
self._get_history_messages(previous_response_id)
|
||||
if previous_response_id
|
||||
else []
|
||||
)
|
||||
|
||||
# Update proxy_server_request body so spend logs record the full request.
|
||||
proxy_server_request = (call_kwargs.get("litellm_metadata") or {}).get(
|
||||
"proxy_server_request"
|
||||
) or {}
|
||||
if isinstance(proxy_server_request, dict):
|
||||
body = dict(proxy_server_request.get("body") or {})
|
||||
body["input"] = call_kwargs.get("input")
|
||||
body["store"] = call_kwargs.get("store")
|
||||
body["model"] = model
|
||||
for k in ("tools", "tool_choice", "instructions", "metadata"):
|
||||
if k in call_kwargs and call_kwargs[k] is not None:
|
||||
body[k] = call_kwargs[k]
|
||||
proxy_server_request = dict(proxy_server_request)
|
||||
proxy_server_request["body"] = body
|
||||
if "litellm_metadata" not in call_kwargs:
|
||||
call_kwargs["litellm_metadata"] = {}
|
||||
call_kwargs["litellm_metadata"]["proxy_server_request"] = proxy_server_request
|
||||
call_kwargs.setdefault("litellm_params", {})
|
||||
call_kwargs["litellm_params"]["proxy_server_request"] = proxy_server_request
|
||||
|
||||
# Merge any safe pass-through kwargs (extra_headers, etc.)
|
||||
self._apply_history(call_kwargs, previous_response_id, current_messages, prior_history)
|
||||
self._inject_credentials(call_kwargs, event_model)
|
||||
self._update_proxy_request(call_kwargs, model)
|
||||
call_kwargs.update(self.extra_kwargs)
|
||||
|
||||
# Track the completed event to update in-memory history after the turn.
|
||||
completed_event: Optional[Dict[str, Any]] = None
|
||||
|
||||
try:
|
||||
stream_response = await _litellm.aresponses(model=model, **call_kwargs)
|
||||
|
||||
async for chunk in stream_response: # type: ignore[union-attr]
|
||||
if chunk is None:
|
||||
continue
|
||||
serialized = self._serialize_chunk(chunk)
|
||||
if serialized is not None:
|
||||
# Capture the completed event for history bookkeeping
|
||||
try:
|
||||
chunk_dict = json.loads(serialized) if isinstance(serialized, str) else {}
|
||||
if chunk_dict.get("type") == "response.completed":
|
||||
completed_event = chunk_dict
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
await self.websocket.send_text(serialized)
|
||||
except Exception as send_exc:
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: error sending chunk to client: %s", send_exc
|
||||
)
|
||||
return # Client disconnected
|
||||
|
||||
completed_event = await self._stream_and_forward(model, call_kwargs)
|
||||
except Exception as exc:
|
||||
verbose_logger.exception("ManagedResponsesWS: error processing response.create: %s", exc)
|
||||
verbose_logger.exception(
|
||||
"ManagedResponsesWS: error processing response.create: %s", exc
|
||||
)
|
||||
await self._send_error(str(exc))
|
||||
return
|
||||
|
||||
# ---- Store this turn in in-memory history for future previous_response_id lookups ----
|
||||
if completed_event is not None:
|
||||
new_response_id = self._extract_response_id(completed_event)
|
||||
if new_response_id:
|
||||
output_msgs = self._extract_output_messages(completed_event)
|
||||
# Accumulate: history from previous turn + current input + new output
|
||||
prior_history: List[Dict[str, Any]] = []
|
||||
if previous_response_id:
|
||||
prior_history = self._get_history_messages(previous_response_id)
|
||||
self._store_history(
|
||||
new_response_id,
|
||||
prior_history + current_messages,
|
||||
output_msgs,
|
||||
)
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: stored %d messages for response_id=%s",
|
||||
len(prior_history) + len(current_messages) + len(output_msgs),
|
||||
new_response_id,
|
||||
)
|
||||
# ---------------------------------------------------------------------------
|
||||
self._save_turn_history(completed_event, prior_history, current_messages)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Main entry point
|
||||
|
|
|
|||
|
|
@ -0,0 +1,973 @@
|
|||
"""
|
||||
Unit tests to verify that all providers support Responses API WebSocket mode.
|
||||
|
||||
Tests that:
|
||||
1. All providers with ResponsesAPIConfig support websocket mode
|
||||
2. Providers with native websocket support use direct connection
|
||||
3. Providers without native websocket support use ManagedResponsesWebSocketHandler
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.azure.responses.transformation import AzureOpenAIResponsesAPIConfig
|
||||
from litellm.llms.chatgpt.responses.transformation import ChatGPTResponsesAPIConfig
|
||||
from litellm.llms.databricks.responses.transformation import (
|
||||
DatabricksResponsesAPIConfig,
|
||||
)
|
||||
from litellm.llms.github_copilot.responses.transformation import (
|
||||
GithubCopilotResponsesAPIConfig,
|
||||
)
|
||||
from litellm.llms.hosted_vllm.responses.transformation import (
|
||||
HostedVLLMResponsesAPIConfig,
|
||||
)
|
||||
from litellm.llms.litellm_proxy.responses.transformation import (
|
||||
LiteLLMProxyResponsesAPIConfig,
|
||||
)
|
||||
from litellm.llms.manus.responses.transformation import ManusResponsesAPIConfig
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.llms.openrouter.responses.transformation import (
|
||||
OpenRouterResponsesAPIConfig,
|
||||
)
|
||||
from litellm.llms.perplexity.responses.transformation import PerplexityResponsesConfig
|
||||
from litellm.llms.volcengine.responses.transformation import (
|
||||
VolcEngineResponsesAPIConfig,
|
||||
)
|
||||
from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
|
||||
|
||||
|
||||
class TestResponsesAPIWebSocketSupport:
|
||||
"""Test that all providers have websocket support configured correctly"""
|
||||
|
||||
def test_openai_supports_native_websocket(self):
|
||||
"""OpenAI should support native websocket"""
|
||||
config = OpenAIResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is True
|
||||
), "OpenAI should support native websocket"
|
||||
|
||||
def test_azure_supports_native_websocket(self):
|
||||
"""Azure should support native websocket (inherits from OpenAI)"""
|
||||
config = AzureOpenAIResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is True
|
||||
), "Azure should support native websocket"
|
||||
|
||||
def test_xai_uses_managed_websocket(self):
|
||||
"""XAI should use managed websocket handler"""
|
||||
config = XAIResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is False
|
||||
), "XAI should use managed websocket handler"
|
||||
|
||||
def test_github_copilot_uses_managed_websocket(self):
|
||||
"""GitHub Copilot should use managed websocket handler"""
|
||||
config = GithubCopilotResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is False
|
||||
), "GitHub Copilot should use managed websocket handler"
|
||||
|
||||
def test_chatgpt_uses_managed_websocket(self):
|
||||
"""ChatGPT should use managed websocket handler"""
|
||||
config = ChatGPTResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is False
|
||||
), "ChatGPT should use managed websocket handler"
|
||||
|
||||
def test_litellm_proxy_uses_managed_websocket(self):
|
||||
"""LiteLLM Proxy should use managed websocket handler"""
|
||||
config = LiteLLMProxyResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is False
|
||||
), "LiteLLM Proxy should use managed websocket handler"
|
||||
|
||||
def test_volcengine_uses_managed_websocket(self):
|
||||
"""VolcEngine should use managed websocket handler"""
|
||||
config = VolcEngineResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is False
|
||||
), "VolcEngine should use managed websocket handler"
|
||||
|
||||
def test_manus_uses_managed_websocket(self):
|
||||
"""Manus should use managed websocket handler"""
|
||||
config = ManusResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is False
|
||||
), "Manus should use managed websocket handler"
|
||||
|
||||
def test_perplexity_uses_managed_websocket(self):
|
||||
"""Perplexity should use managed websocket handler"""
|
||||
config = PerplexityResponsesConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is False
|
||||
), "Perplexity should use managed websocket handler"
|
||||
|
||||
def test_databricks_uses_managed_websocket(self):
|
||||
"""Databricks should use managed websocket handler"""
|
||||
config = DatabricksResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is False
|
||||
), "Databricks should use managed websocket handler"
|
||||
|
||||
def test_openrouter_uses_managed_websocket(self):
|
||||
"""OpenRouter should use managed websocket handler"""
|
||||
config = OpenRouterResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is False
|
||||
), "OpenRouter should use managed websocket handler"
|
||||
|
||||
def test_hosted_vllm_uses_managed_websocket(self):
|
||||
"""Hosted vLLM should use managed websocket handler"""
|
||||
config = HostedVLLMResponsesAPIConfig()
|
||||
assert (
|
||||
config.supports_native_websocket() is False
|
||||
), "Hosted vLLM should use managed websocket handler"
|
||||
|
||||
|
||||
class TestManagedWebSocketHandlerIntegration:
|
||||
"""Test that ManagedResponsesWebSocketHandler is properly integrated"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_managed_handler_instantiation(self):
|
||||
"""Test that ManagedResponsesWebSocketHandler can be instantiated"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
mock_websocket = MagicMock()
|
||||
mock_logging_obj = Logging(
|
||||
model="test-model",
|
||||
messages=[],
|
||||
stream=True,
|
||||
call_type="aresponses",
|
||||
start_time=0,
|
||||
litellm_call_id="test-id",
|
||||
function_id="test-func",
|
||||
)
|
||||
|
||||
handler = ManagedResponsesWebSocketHandler(
|
||||
websocket=mock_websocket,
|
||||
model="test-model",
|
||||
logging_obj=mock_logging_obj,
|
||||
user_api_key_dict=None,
|
||||
litellm_metadata={},
|
||||
api_key="test-key",
|
||||
api_base="https://api.example.com",
|
||||
timeout=30.0,
|
||||
custom_llm_provider="test_provider",
|
||||
)
|
||||
|
||||
assert handler.model == "test-model"
|
||||
assert handler.api_key == "test-key"
|
||||
assert handler.api_base == "https://api.example.com"
|
||||
assert handler.timeout == 30.0
|
||||
assert handler.custom_llm_provider == "test_provider"
|
||||
|
||||
|
||||
class TestChunkTransformation:
|
||||
"""Test chunk serialization and transformation for WebSocket streaming"""
|
||||
|
||||
def test_serialize_chunk_with_dict(self):
|
||||
"""Test serialization of dict chunks"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.created",
|
||||
"response": {"id": "resp_456", "status": "in_progress"},
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.created" in serialized
|
||||
assert "resp_456" in serialized
|
||||
|
||||
def test_serialize_chunk_handles_invalid_json(self):
|
||||
"""Test that chunks with circular references are handled"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
# Create object with circular reference
|
||||
obj = {"a": 1}
|
||||
obj["self"] = obj # type: ignore
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(obj)
|
||||
assert serialized is None
|
||||
|
||||
def test_extract_output_messages_with_text_content(self):
|
||||
"""Test extraction of output messages with text content"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "Hello world"}],
|
||||
}
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._extract_output_messages(
|
||||
completed_event
|
||||
)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["type"] == "message"
|
||||
assert messages[0]["role"] == "assistant"
|
||||
assert messages[0]["content"][0]["text"] == "Hello world"
|
||||
|
||||
def test_extract_output_messages_with_multiple_content_parts(self):
|
||||
"""Test extraction with multiple content parts"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "output_text", "text": "Part 1. "},
|
||||
{"type": "output_text", "text": "Part 2."},
|
||||
],
|
||||
}
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._extract_output_messages(
|
||||
completed_event
|
||||
)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"][0]["text"] == "Part 1. Part 2."
|
||||
|
||||
def test_extract_output_messages_with_function_calls(self):
|
||||
"""Test that function calls are preserved"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"output": [
|
||||
{
|
||||
"type": "function_call",
|
||||
"id": "call_123",
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "Paris"}',
|
||||
}
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._extract_output_messages(
|
||||
completed_event
|
||||
)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["type"] == "function_call"
|
||||
assert messages[0]["id"] == "call_123"
|
||||
assert messages[0]["name"] == "get_weather"
|
||||
|
||||
def test_extract_output_messages_filters_empty_text(self):
|
||||
"""Test that messages with empty text are filtered out"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": ""}],
|
||||
},
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "Valid text"}],
|
||||
},
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._extract_output_messages(
|
||||
completed_event
|
||||
)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"][0]["text"] == "Valid text"
|
||||
|
||||
def test_extract_output_messages_handles_non_dict_items(self):
|
||||
"""Test that non-dict items are skipped"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"output": [
|
||||
"invalid_string",
|
||||
None,
|
||||
123,
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "Valid"}],
|
||||
},
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._extract_output_messages(
|
||||
completed_event
|
||||
)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"][0]["text"] == "Valid"
|
||||
|
||||
def test_input_to_messages_with_string(self):
|
||||
"""Test conversion of string input to messages"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._input_to_messages("Hello world")
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["type"] == "message"
|
||||
assert messages[0]["role"] == "user"
|
||||
assert messages[0]["content"][0]["type"] == "input_text"
|
||||
assert messages[0]["content"][0]["text"] == "Hello world"
|
||||
|
||||
def test_input_to_messages_with_list(self):
|
||||
"""Test conversion of list input to messages"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
input_list = [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "Question"}],
|
||||
}
|
||||
]
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._input_to_messages(input_list)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["type"] == "message"
|
||||
assert messages[0]["content"][0]["text"] == "Question"
|
||||
|
||||
def test_input_to_messages_filters_non_dict_items(self):
|
||||
"""Test that non-dict items in list input are filtered"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
input_list = [
|
||||
"invalid_string",
|
||||
None,
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "Valid"}],
|
||||
},
|
||||
]
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._input_to_messages(input_list)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"][0]["text"] == "Valid"
|
||||
|
||||
def test_input_to_messages_handles_empty_input(self):
|
||||
"""Test that empty input returns empty list"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
assert ManagedResponsesWebSocketHandler._input_to_messages(None) == []
|
||||
assert ManagedResponsesWebSocketHandler._input_to_messages([]) == []
|
||||
assert ManagedResponsesWebSocketHandler._input_to_messages({}) == []
|
||||
|
||||
|
||||
class TestWebSocketEventTypes:
|
||||
"""Test that all WebSocket event types are properly handled with dict-based chunks"""
|
||||
|
||||
def test_serialize_response_created_event_dict(self):
|
||||
"""Test serialization of response.created event as dict"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.created",
|
||||
"response_id": "resp_123",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"object": "response",
|
||||
"status": "in_progress",
|
||||
"created_at": 1234567890,
|
||||
},
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.created" in serialized
|
||||
assert "resp_123" in serialized
|
||||
|
||||
def test_serialize_response_in_progress_event_dict(self):
|
||||
"""Test serialization of response.in_progress event as dict"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {"type": "response.in_progress", "response_id": "resp_123"}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.in_progress" in serialized
|
||||
|
||||
def test_serialize_output_item_added_event_dict(self):
|
||||
"""Test serialization of response.output_item.added event as dict"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.output_item.added",
|
||||
"response_id": "resp_123",
|
||||
"item_id": "msg_456",
|
||||
"output_index": 0,
|
||||
"item": {"type": "message", "role": "assistant"},
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.output_item.added" in serialized
|
||||
assert "msg_456" in serialized
|
||||
|
||||
def test_serialize_output_text_delta_event_dict(self):
|
||||
"""Test serialization of response.output_text.delta event as dict"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.output_text.delta",
|
||||
"response_id": "resp_123",
|
||||
"item_id": "msg_456",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": "Hello",
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.output_text.delta" in serialized
|
||||
assert "Hello" in serialized
|
||||
|
||||
def test_serialize_output_text_done_event_dict(self):
|
||||
"""Test serialization of response.output_text.done event as dict"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.output_text.done",
|
||||
"response_id": "resp_123",
|
||||
"item_id": "msg_456",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"text": "Hello world",
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.output_text.done" in serialized
|
||||
assert "Hello world" in serialized
|
||||
|
||||
def test_serialize_content_part_done_event_dict(self):
|
||||
"""Test serialization of response.content_part.done event as dict"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.content_part.done",
|
||||
"response_id": "resp_123",
|
||||
"item_id": "msg_456",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"part": {"type": "output_text", "text": "Complete text"},
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.content_part.done" in serialized
|
||||
|
||||
def test_serialize_output_item_done_event_dict(self):
|
||||
"""Test serialization of response.output_item.done event as dict"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.output_item.done",
|
||||
"response_id": "resp_123",
|
||||
"item_id": "msg_456",
|
||||
"output_index": 0,
|
||||
"item": {"type": "message", "role": "assistant", "status": "completed"},
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.output_item.done" in serialized
|
||||
assert "msg_456" in serialized
|
||||
|
||||
def test_serialize_response_completed_event_dict(self):
|
||||
"""Test serialization of response.completed event as dict"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.completed",
|
||||
"response_id": "resp_123",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"status": "completed",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"content": [{"type": "output_text", "text": "Done"}],
|
||||
}
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.completed" in serialized
|
||||
assert "resp_123" in serialized
|
||||
|
||||
def test_serialize_response_failed_event_dict(self):
|
||||
"""Test serialization of response.failed event as dict"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.failed",
|
||||
"response_id": "resp_123",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"status": "failed",
|
||||
"status_details": {"error": {"message": "Rate limit exceeded"}},
|
||||
},
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.failed" in serialized
|
||||
assert "Rate limit exceeded" in serialized
|
||||
|
||||
def test_serialize_response_incomplete_event_dict(self):
|
||||
"""Test serialization of response.incomplete event as dict"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.incomplete",
|
||||
"response_id": "resp_123",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"status": "incomplete",
|
||||
"status_details": {"reason": "max_output_tokens"},
|
||||
},
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.incomplete" in serialized
|
||||
assert "max_output_tokens" in serialized
|
||||
|
||||
|
||||
class TestMultiTurnSessionHistory:
|
||||
"""Test multi-turn conversation handling via session history"""
|
||||
|
||||
def test_extract_output_messages_preserves_multiple_messages(self):
|
||||
"""Test that multiple output messages are all preserved"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "First message"}],
|
||||
},
|
||||
{
|
||||
"type": "function_call",
|
||||
"id": "call_123",
|
||||
"name": "get_weather",
|
||||
"arguments": "{}",
|
||||
},
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "Second message"}],
|
||||
},
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._extract_output_messages(
|
||||
completed_event
|
||||
)
|
||||
assert len(messages) == 3
|
||||
assert messages[0]["content"][0]["text"] == "First message"
|
||||
assert messages[1]["type"] == "function_call"
|
||||
assert messages[2]["content"][0]["text"] == "Second message"
|
||||
|
||||
def test_input_to_messages_with_mixed_content_types(self):
|
||||
"""Test input conversion with mixed content types"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
input_list = [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "input_text", "text": "Question"},
|
||||
{"type": "input_image", "image_url": "https://example.com/img.png"},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._input_to_messages(input_list)
|
||||
assert len(messages) == 1
|
||||
assert len(messages[0]["content"]) == 2
|
||||
assert messages[0]["content"][0]["type"] == "input_text"
|
||||
assert messages[0]["content"][1]["type"] == "input_image"
|
||||
|
||||
def test_extract_output_messages_with_mixed_text_types(self):
|
||||
"""Test that both 'output_text' and 'text' types are extracted"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "output_text", "text": "Part 1"},
|
||||
{"type": "text", "text": "Part 2"},
|
||||
],
|
||||
}
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._extract_output_messages(
|
||||
completed_event
|
||||
)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"][0]["text"] == "Part 1Part 2"
|
||||
|
||||
def test_extract_response_id_from_completed_event(self):
|
||||
"""Test extraction of response ID from completed event"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {
|
||||
"type": "response.completed",
|
||||
"response": {"id": "resp_abc123", "status": "completed"},
|
||||
}
|
||||
|
||||
response_id = ManagedResponsesWebSocketHandler._extract_response_id(
|
||||
completed_event
|
||||
)
|
||||
assert response_id == "resp_abc123"
|
||||
|
||||
def test_extract_response_id_handles_missing_response(self):
|
||||
"""Test that missing response dict returns None"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {"type": "response.completed"}
|
||||
|
||||
response_id = ManagedResponsesWebSocketHandler._extract_response_id(
|
||||
completed_event
|
||||
)
|
||||
assert response_id is None
|
||||
|
||||
|
||||
class TestWebSocketErrorHandling:
|
||||
"""Test error handling in WebSocket mode"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_managed_handler_handles_invalid_json(self):
|
||||
"""Test that invalid JSON in response.create is handled gracefully"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
mock_websocket = MagicMock()
|
||||
mock_websocket.send_text = AsyncMock()
|
||||
mock_websocket.recv = AsyncMock(return_value="invalid json {{{")
|
||||
|
||||
mock_logging_obj = Logging(
|
||||
model="test-model",
|
||||
messages=[],
|
||||
stream=True,
|
||||
call_type="aresponses",
|
||||
start_time=0,
|
||||
litellm_call_id="test-id",
|
||||
function_id="test-func",
|
||||
)
|
||||
|
||||
handler = ManagedResponsesWebSocketHandler(
|
||||
websocket=mock_websocket,
|
||||
model="test-model",
|
||||
logging_obj=mock_logging_obj,
|
||||
)
|
||||
|
||||
# Process invalid JSON
|
||||
await handler._process_response_create("invalid json {{{")
|
||||
|
||||
# Should have sent an error event
|
||||
mock_websocket.send_text.assert_called_once()
|
||||
error_event = mock_websocket.send_text.call_args[0][0]
|
||||
assert "error" in error_event
|
||||
assert "Invalid JSON" in error_event
|
||||
|
||||
|
||||
class TestWebSocketChunkTypes:
|
||||
"""Test handling of different chunk types from streaming responses"""
|
||||
|
||||
def test_serialize_function_call_chunk(self):
|
||||
"""Test serialization of function call chunks"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.function_call.added",
|
||||
"response_id": "resp_123",
|
||||
"item_id": "call_456",
|
||||
"output_index": 0,
|
||||
"call_id": "call_456",
|
||||
"name": "get_weather",
|
||||
"arguments": "",
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.function_call.added" in serialized
|
||||
assert "get_weather" in serialized
|
||||
|
||||
def test_serialize_function_call_arguments_delta(self):
|
||||
"""Test serialization of function call arguments delta"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.function_call_arguments.delta",
|
||||
"response_id": "resp_123",
|
||||
"item_id": "call_456",
|
||||
"output_index": 0,
|
||||
"call_id": "call_456",
|
||||
"delta": '{"location"',
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.function_call_arguments.delta" in serialized
|
||||
assert "location" in serialized
|
||||
|
||||
def test_serialize_function_call_arguments_done(self):
|
||||
"""Test serialization of function call arguments done"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.function_call_arguments.done",
|
||||
"response_id": "resp_123",
|
||||
"item_id": "call_456",
|
||||
"output_index": 0,
|
||||
"call_id": "call_456",
|
||||
"arguments": '{"location": "Paris"}',
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.function_call_arguments.done" in serialized
|
||||
assert "Paris" in serialized
|
||||
|
||||
def test_serialize_reasoning_content_delta(self):
|
||||
"""Test serialization of reasoning content delta"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.reasoning_content.delta",
|
||||
"response_id": "resp_123",
|
||||
"item_id": "msg_456",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": "Thinking step 1...",
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.reasoning_content.delta" in serialized
|
||||
assert "Thinking step 1" in serialized
|
||||
|
||||
def test_serialize_reasoning_content_done(self):
|
||||
"""Test serialization of reasoning content done"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.reasoning_content.done",
|
||||
"response_id": "resp_123",
|
||||
"item_id": "msg_456",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"reasoning_content": "Complete reasoning...",
|
||||
}
|
||||
|
||||
serialized = ManagedResponsesWebSocketHandler._serialize_chunk(chunk)
|
||||
assert serialized is not None
|
||||
assert "response.reasoning_content.done" in serialized
|
||||
assert "Complete reasoning" in serialized
|
||||
|
||||
def test_extract_output_messages_preserves_multiple_messages(self):
|
||||
"""Test that multiple output messages are all preserved"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "First message"}],
|
||||
},
|
||||
{
|
||||
"type": "function_call",
|
||||
"id": "call_123",
|
||||
"name": "get_weather",
|
||||
"arguments": "{}",
|
||||
},
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "Second message"}],
|
||||
},
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._extract_output_messages(
|
||||
completed_event
|
||||
)
|
||||
assert len(messages) == 3
|
||||
assert messages[0]["content"][0]["text"] == "First message"
|
||||
assert messages[1]["type"] == "function_call"
|
||||
assert messages[2]["content"][0]["text"] == "Second message"
|
||||
|
||||
def test_input_to_messages_with_mixed_content_types(self):
|
||||
"""Test input conversion with mixed content types"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
input_list = [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "input_text", "text": "Question"},
|
||||
{"type": "input_image", "image_url": "https://example.com/img.png"},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._input_to_messages(input_list)
|
||||
assert len(messages) == 1
|
||||
assert len(messages[0]["content"]) == 2
|
||||
assert messages[0]["content"][0]["type"] == "input_text"
|
||||
assert messages[0]["content"][1]["type"] == "input_image"
|
||||
|
||||
def test_extract_output_messages_with_mixed_text_types(self):
|
||||
"""Test that both 'output_text' and 'text' types are extracted"""
|
||||
from litellm.responses.streaming_iterator import (
|
||||
ManagedResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
completed_event = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "output_text", "text": "Part 1"},
|
||||
{"type": "text", "text": "Part 2"},
|
||||
],
|
||||
}
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
messages = ManagedResponsesWebSocketHandler._extract_output_messages(
|
||||
completed_event
|
||||
)
|
||||
assert len(messages) == 1
|
||||
assert messages[0]["content"][0]["text"] == "Part 1Part 2"
|
||||
Loading…
Add table
Reference in a new issue