Add support for responses websocket for all providers

This commit is contained in:
Sameer Kankute
2026-03-04 18:24:50 +05:30
parent 8764e5da8c
commit eec6f6ee69
3 changed files with 233 additions and 196 deletions
+30 -5
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:
+5 -5
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,
)
+198 -186
View File
@@ -26,6 +26,7 @@ from litellm.types.llms.openai import (
OutputTextDeltaEvent,
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIRequestParams,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
ResponsesAPIStreamingResponse,
@@ -874,37 +875,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",
}
)
@@ -952,7 +933,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.
@@ -985,41 +966,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]:
@@ -1027,8 +994,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:
@@ -1040,7 +1005,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):
@@ -1077,6 +1042,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)``,
@@ -1097,157 +1225,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