Merge pull request #22771 from BerriAI/litellm_responses_websocket_2

Add support for responses websocket for all providers
This commit is contained in:
Sameer Kankute 2026-03-04 22:12:12 +05:30 • committed by GitHub
commit 23d312dbd2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 1509 additions and 202 deletions

View file

@ -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.

View file

@ -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 ##########
#########################################################

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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 ##############
#########################################################

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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,
)

View file

@ -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

View file

@ -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"