Merge pull request #22559 from BerriAI/litellm_responses_websocket

[Feat] Add support for Responses Websocket
This commit is contained in:
Sameer Kankute 2026-03-04 17:57:34 +05:30 • committed by GitHub
commit 8764e5da8c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 1162 additions and 6 deletions

View file

@ -1246,6 +1246,7 @@ from .ocr.main import *
from .rag.main import *
from .search.main import *
from .realtime_api.main import _arealtime
from .responses.main import _aresponses_websocket
from .fine_tuning.main import *
from .files.main import *
from .vector_store_files.main import (

View file

@ -69,6 +69,7 @@ from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
MockResponsesAPIStreamingIterator,
ResponsesAPIStreamingIterator,
ResponsesWebSocketStreaming,
SyncResponsesAPIStreamingIterator,
)
from litellm.types.containers.main import (
@ -4731,6 +4732,98 @@ class BaseLLMHTTPHandler:
f"Unexpected error while closing WebSocket: {close_error}"
)
async def async_responses_websocket(
self,
model: str,
websocket: Any,
logging_obj: LiteLLMLoggingObj,
responses_api_provider_config: 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,
):
"""
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.
"""
import websockets
from websockets.asyncio.client import ClientConnection
litellm_params = GenericLiteLLMParams()
headers = responses_api_provider_config.validate_environment(
headers={},
model=model,
litellm_params=litellm_params,
)
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
http_url = responses_api_provider_config.get_complete_url(
api_base=api_base,
litellm_params={},
)
# /responses -> wss:// URL
ws_url = http_url.replace("https://", "wss://").replace("http://", "ws://")
try:
ssl_context = get_shared_realtime_ssl_context()
if ws_url.startswith("wss://") and ssl_context is False:
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
ssl_context.check_hostname = False
ssl_context.verify_mode = ssl.CERT_NONE
logging_obj.pre_call(
input=None,
api_key=api_key or "",
additional_args={
"api_base": ws_url,
"headers": headers,
"complete_input_dict": {"mode": "responses_websocket"},
},
)
async with websockets.connect( # type: ignore
ws_url,
additional_headers=headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_context,
) as backend_ws:
_request_data: Dict[str, Any] = {}
if litellm_metadata:
_request_data["litellm_metadata"] = litellm_metadata
streaming = ResponsesWebSocketStreaming(
websocket=websocket,
backend_ws=cast(ClientConnection, backend_ws),
logging_obj=logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=_request_data,
)
await streaming.bidirectional_forward()
except websockets.exceptions.InvalidStatusCode as e: # type: ignore
verbose_logger.exception(f"Error connecting to responses WS backend: {e}")
await websocket.close(code=e.status_code, reason=str(e))
except Exception as e:
verbose_logger.exception(f"Error in responses WS: {e}")
try:
await websocket.close(
code=1011, reason=f"Internal server error: {str(e)}"
)
except RuntimeError as close_error:
if "already completed" in str(close_error) or "websocket.close" in str(
close_error
):
pass
else:
raise Exception(
f"Unexpected error while closing WebSocket: {close_error}"
)
def image_edit_handler(
self,
model: str,

View file

@ -499,6 +499,7 @@ class ProxyBaseLLMRequestProcessing:
"aembedding",
"aresponses",
"_arealtime",
"_aresponses_websocket",
"aget_responses",
"adelete_responses",
"acancel_responses",

View file

@ -1,14 +1,21 @@
import asyncio
import json
import time
from typing import Any, AsyncIterator, Optional, cast
from typing import Any, AsyncIterator, Dict, Optional, cast
from uuid import uuid4
import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from starlette.websockets import WebSocket
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.auth.user_api_key_auth import (
UserAPIKeyAuth,
user_api_key_auth,
user_api_key_auth_websocket,
)
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse
from litellm.types.responses.main import DeleteResponseResult
@ -904,3 +911,121 @@ async def cancel_response(
proxy_logging_obj=proxy_logging_obj,
version=version,
)
@router.websocket("/v1/responses")
@router.websocket("/responses")
async def responses_websocket_endpoint(
websocket: WebSocket,
model: str = fastapi.Query(
..., description="The model to use for the responses WebSocket session."
),
user_api_key_dict=Depends(user_api_key_auth_websocket),
):
"""
Responses API WebSocket mode endpoint.
Keeps a persistent WebSocket connection for response.create events,
enabling lower-latency agentic workflows with many tool-call round trips.
See: https://developers.openai.com/api/docs/guides/websocket-mode/
"""
from litellm.proxy.proxy_server import (
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
user_api_base,
user_max_tokens,
user_model,
user_request_timeout,
user_temperature,
version,
)
from litellm.proxy.route_llm_request import route_request
# Accept the WebSocket handshake
requested_protocols = [
p.strip()
for p in (websocket.headers.get("sec-websocket-protocol") or "").split(",")
if p.strip()
]
accept_kwargs: dict = {}
if requested_protocols:
accept_kwargs["subprotocol"] = requested_protocols[0]
await websocket.accept(**accept_kwargs)
data: Dict[str, Any] = {
"model": model,
"websocket": websocket,
}
# Construct a synthetic Request for pre-call processing
headers_list = list(websocket.scope.get("headers") or [])
scope: Dict[str, Any] = {
"type": "http",
"method": "POST",
"path": "/v1/responses",
"headers": headers_list,
}
request = Request(scope=scope)
request._url = websocket.url
async def return_body():
return f'{{"model": "{model}"}}'.encode()
request.body = return_body # type: ignore
# Phase 1: pre-call processing (auth, guardrails, rate limits)
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
try:
(
data,
litellm_logging_obj,
) = await base_llm_response_processor.common_processing_pre_call_logic(
request=request,
general_settings=general_settings,
user_api_key_dict=user_api_key_dict,
version=version,
proxy_logging_obj=proxy_logging_obj,
proxy_config=proxy_config,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
model=model,
route_type="_aresponses_websocket",
)
except Exception as e:
verbose_proxy_logger.exception("Responses WebSocket pre-call error")
try:
await websocket.send_text(
json.dumps(
{
"type": "error",
"error": {
"type": "pre_call_error",
"message": str(e),
},
}
)
)
except Exception:
pass
await websocket.close(code=1011, reason="Pre-call error")
return
# Phase 2: route to upstream provider
try:
data["user_api_key_dict"] = user_api_key_dict
llm_call = await route_request(
data=data,
route_type="_aresponses_websocket",
llm_router=llm_router,
user_model=user_model,
)
await llm_call
except Exception:
verbose_proxy_logger.exception("Responses WebSocket error")
await websocket.close(code=1011, reason="Internal server error")

View file

@ -42,6 +42,7 @@ ROUTE_ENDPOINT_MAPPING = {
"amoderation": "/moderations",
"arerank": "/rerank",
"aresponses": "/responses",
"_aresponses_websocket": "/responses",
"alist_input_items": "/responses/{response_id}/input_items",
"aimage_edit": "/images/edits",
"acancel_responses": "/responses/{response_id}/cancel",
@ -163,6 +164,7 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin
"acreate_response_reply",
"alist_input_items",
"_arealtime", # private function for realtime API
"_aresponses_websocket", # private function for responses WebSocket mode
"aimage_edit",
"agenerate_content",
"agenerate_content_stream",

View file

@ -51,6 +51,8 @@ if TYPE_CHECKING:
from litellm.types.llms.openai import ResponseText # type: ignore
else:
ResponseText = str # Fallback for ResponseText import
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.secret_managers.main import get_secret_str
from litellm.types.responses.main import *
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager, client
@ -182,8 +184,6 @@ async def aresponses_api_with_mcp(
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None
secret_fields = kwargs.get("secret_fields")
if secret_fields and isinstance(secret_fields, dict):
from litellm.responses.utils import ResponsesAPIRequestUtils
mcp_auth_header, mcp_server_auth_headers, _, _ = (
ResponsesAPIRequestUtils.extract_mcp_headers_from_request(
secret_fields=secret_fields, tools=tools
@ -1662,3 +1662,100 @@ def compact_responses(
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
# ---------------------------------------------------------------------------
# Responses API WebSocket mode
# ---------------------------------------------------------------------------
def _build_litellm_metadata_for_ws(kwargs: dict) -> dict:
metadata: dict = {**(kwargs.get("litellm_metadata") or {})}
guardrails = (
(kwargs.get("metadata") or {}).get("guardrails")
or kwargs.get("guardrails")
or []
)
if guardrails:
metadata["guardrails"] = guardrails
return metadata
@client
async def _aresponses_websocket(
model: str,
websocket: Any,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
timeout: Optional[float] = None,
**kwargs,
):
"""
Private function to handle the Responses API WebSocket mode.
For PROXY use only.
Resolves the LLM provider from ``model``, looks up the matching
``BaseResponsesAPIConfig``, and hands off to
``BaseLLMHTTPHandler.async_responses_websocket``.
"""
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
user = kwargs.get("user", None)
litellm_params = GenericLiteLLMParams(**kwargs)
litellm_params_dict = get_litellm_params(**kwargs)
model, _custom_llm_provider, dynamic_api_key, dynamic_api_base = (
litellm.get_llm_provider(
model=model,
api_base=api_base,
api_key=api_key,
)
)
litellm_logging_obj.update_environment_variables(
model=model,
user=user,
optional_params={},
litellm_params=litellm_params_dict,
custom_llm_provider=_custom_llm_provider,
)
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = None
if _custom_llm_provider is not None:
responses_api_provider_config = (
ProviderConfigManager.get_provider_responses_api_config(
model=model,
provider=litellm.LlmProviders(_custom_llm_provider),
)
)
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
or litellm.api_base
or None
)
resolved_api_key = (
dynamic_api_key
or litellm_params.api_key
or litellm.api_key
or litellm.openai_key
or get_secret_str("OPENAI_API_KEY")
)
await base_llm_http_handler.async_responses_websocket(
model=model,
websocket=websocket,
logging_obj=litellm_logging_obj,
responses_api_provider_config=responses_api_provider_config,
api_base=resolved_api_base,
api_key=resolved_api_key,
timeout=timeout,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata_for_ws(kwargs),
)

View file

@ -3,7 +3,7 @@ import json
import time
import traceback
from datetime import datetime
from typing import Any, Dict, Optional
from typing import TYPE_CHECKING, Any, Dict, List, Optional
import httpx
@ -682,3 +682,594 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
for c in getattr(out_item, "content", []):
out += c.text
return out
# ---------------------------------------------------------------------------
# WebSocket mode streaming (bidirectional forwarding)
# ---------------------------------------------------------------------------
if TYPE_CHECKING:
from websockets.asyncio.client import ClientConnection as _WsClientConnection
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.thread_pool_executor import executor as _ws_executor
RESPONSES_WS_LOGGED_EVENT_TYPES = [
"response.created",
"response.completed",
"response.failed",
"response.incomplete",
"error",
]
class ResponsesWebSocketStreaming:
"""
Manages bidirectional WebSocket forwarding for the Responses API
WebSocket mode (wss://.../v1/responses).
Unlike the Realtime API, the Responses API WebSocket mode:
- Uses response.create as the client-to-server event
- Streams back the same events as the HTTP streaming Responses API
- Supports previous_response_id for incremental continuation
- Supports generate: false for warmup
- One response at a time per connection (sequential, no multiplexing)
"""
def __init__(
self,
websocket: Any,
backend_ws: Any,
logging_obj: LiteLLMLoggingObj,
user_api_key_dict: Optional[Any] = None,
request_data: Optional[Dict] = None,
):
self.websocket = websocket
self.backend_ws = backend_ws
self.logging_obj = logging_obj
self.user_api_key_dict = user_api_key_dict
self.request_data: Dict = request_data or {}
self.messages: list[Dict] = []
self.input_messages: list[Dict[str, str]] = []
def _should_store_event(self, event_obj: dict) -> bool:
return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES
def _store_event(self, event: Any) -> None:
if isinstance(event, bytes):
event = event.decode("utf-8")
if isinstance(event, str):
try:
event_obj = json.loads(event)
except (json.JSONDecodeError, TypeError):
return
else:
event_obj = event
if self._should_store_event(event_obj):
self.messages.append(event_obj)
def _collect_input_from_client_event(self, message: Any) -> None:
"""Extract user input content from response.create for logging."""
try:
if isinstance(message, str):
msg_obj = json.loads(message)
elif isinstance(message, dict):
msg_obj = message
else:
return
if msg_obj.get("type") != "response.create":
return
input_items = msg_obj.get("input", [])
if isinstance(input_items, str):
self.input_messages.append({"role": "user", "content": input_items})
return
if isinstance(input_items, list):
for item in input_items:
if not isinstance(item, dict):
continue
if item.get("type") == "message" and item.get("role") == "user":
content = item.get("content", [])
if isinstance(content, str):
self.input_messages.append(
{"role": "user", "content": content}
)
elif isinstance(content, list):
for c in content:
if (
isinstance(c, dict)
and c.get("type") == "input_text"
):
text = c.get("text", "")
if text:
self.input_messages.append(
{"role": "user", "content": text}
)
except (json.JSONDecodeError, AttributeError, TypeError):
pass
def _store_input(self, message: Any) -> None:
self._collect_input_from_client_event(message)
if self.logging_obj:
self.logging_obj.pre_call(input=message, api_key="")
async def _log_messages(self) -> None:
if not self.logging_obj:
return
if self.input_messages:
self.logging_obj.model_call_details["messages"] = self.input_messages
if self.messages:
asyncio.create_task(
self.logging_obj.async_success_handler(self.messages)
)
_ws_executor.submit(self.logging_obj.success_handler, self.messages)
async def backend_to_client(self) -> None:
"""Forward events from backend WebSocket to the client."""
import websockets
try:
while True:
try:
raw_response = await self.backend_ws.recv(decode=False) # type: ignore[union-attr]
except TypeError:
raw_response = await self.backend_ws.recv() # type: ignore[union-attr, assignment]
if isinstance(raw_response, bytes):
response_str = raw_response.decode("utf-8")
else:
response_str = raw_response
self._store_event(response_str)
await self.websocket.send_text(response_str)
except websockets.exceptions.ConnectionClosed as e: # type: ignore
verbose_logger.debug(
"Responses WS backend connection closed: %s", e
)
except Exception as e:
verbose_logger.exception(
"Error in responses WS backend_to_client: %s", e
)
finally:
await self._log_messages()
async def client_to_backend(self) -> None:
"""Forward response.create events from client to backend."""
try:
while True:
message = await self.websocket.receive_text()
self._store_input(message)
self._store_event(message)
await self.backend_ws.send(message) # type: ignore[union-attr]
except Exception as e:
verbose_logger.debug("Responses WS client_to_backend ended: %s", e)
async def bidirectional_forward(self) -> None:
"""Run both forwarding directions concurrently."""
forward_task = asyncio.create_task(self.backend_to_client())
try:
await self.client_to_backend()
except Exception:
pass
finally:
if not forward_task.done():
forward_task.cancel()
try:
await forward_task
except asyncio.CancelledError:
pass
try:
await self.backend_ws.close()
except Exception:
pass
# ---------------------------------------------------------------------------
# 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",
)
_MANAGED_WS_SKIP_KWARGS = frozenset(
{
"litellm_logging_obj",
"litellm_call_id",
"aresponses",
"_aresponses_websocket",
"user_api_key_dict",
}
)
class ManagedResponsesWebSocketHandler:
"""
Handles Responses API WebSocket mode for providers that do not expose a
native ``wss://`` responses endpoint.
Instead of proxying to a provider WebSocket, this handler:
- Listens for ``response.create`` events from the client
- Makes HTTP streaming calls via ``litellm.aresponses(stream=True)``
- Serialises and forwards every streaming event back over the WebSocket
- Supports ``previous_response_id`` for multi-turn conversations via
in-memory session tracking (avoids async DB-write timing issues)
- Supports sequential requests over a single persistent connection
This makes every provider that LiteLLM can reach over HTTP available on
the WebSocket transport without any provider-specific changes.
"""
def __init__(
self,
websocket: Any,
model: str,
logging_obj: "LiteLLMLoggingObj",
user_api_key_dict: Optional[Any] = None,
litellm_metadata: Optional[Dict[str, Any]] = None,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
timeout: Optional[float] = None,
custom_llm_provider: Optional[str] = None,
**kwargs: Any,
) -> None:
self.websocket = websocket
self.model = model
self.logging_obj = logging_obj
self.user_api_key_dict = user_api_key_dict
self.litellm_metadata: Dict[str, Any] = litellm_metadata or {}
self.api_key = api_key
self.api_base = api_base
self.timeout = timeout
self.custom_llm_provider = custom_llm_provider
# Carry through safe pass-through kwargs (e.g. extra_headers)
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.
# 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.
self._session_history: Dict[str, List[Dict[str, Any]]] = {}
# ------------------------------------------------------------------
# Internal helpers
# ------------------------------------------------------------------
@staticmethod
def _serialize_chunk(chunk: Any) -> Optional[str]:
"""Serialize a streaming chunk to a JSON string for WebSocket transmission."""
try:
if hasattr(chunk, "model_dump_json"):
return chunk.model_dump_json(exclude_none=True)
if hasattr(chunk, "model_dump"):
return json.dumps(chunk.model_dump(exclude_none=True), default=str)
if isinstance(chunk, dict):
return json.dumps(chunk, default=str)
return json.dumps(str(chunk))
except Exception as exc:
verbose_logger.debug("ManagedResponsesWS: failed to serialize chunk: %s", exc)
return None
async def _send_error(self, message: str, error_type: str = "server_error") -> None:
try:
await self.websocket.send_text(
json.dumps({"type": "error", "error": {"type": error_type, "message": message}})
)
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:
"""
Persist a turn's messages in the in-memory session store.
*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`.
"""
prior: List[Dict[str, Any]] = self._session_history.get(response_id, [])
self._session_history[response_id] = prior + input_messages + output_messages
@staticmethod
def _extract_response_id(completed_event: Dict[str, Any]) -> Optional[str]:
"""
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:
return None
decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(encoded_id)
return decoded.get("response_id", encoded_id)
@staticmethod
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``.
"""
resp_obj = completed_event.get("response", {})
if not isinstance(resp_obj, dict):
return []
messages: List[Dict[str, Any]] = []
for item in resp_obj.get("output", []) or []:
if not isinstance(item, dict):
continue
item_type = item.get("type")
role = item.get("role", "assistant")
if item_type == "message":
content_parts = item.get("content") or []
text_parts = [
p.get("text", "")
for p in content_parts
if isinstance(p, dict) and p.get("type") in ("output_text", "text")
]
text = "".join(text_parts)
if text:
messages.append({"type": "message", "role": role, "content": [{"type": "output_text", "text": text}]})
elif item_type == "function_call":
messages.append(item)
return messages
@staticmethod
def _input_to_messages(input_val: Any) -> List[Dict[str, Any]]:
"""
Normalise the ``input`` field of a ``response.create`` event to a list
of Responses API message dicts.
"""
if isinstance(input_val, str):
return [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": input_val}]}]
if isinstance(input_val, list):
return [item for item in input_val if isinstance(item, dict)]
return []
async def _process_response_create(self, raw_message: str) -> None:
"""
Parse one ``response.create`` event, call ``litellm.aresponses(stream=True)``,
and forward every streaming event to the client.
Multi-turn support via in-memory session history
------------------------------------------------
When ``previous_response_id`` is present in the event:
1. Look up the accumulated message history in ``self._session_history``
(keyed by the decoded provider response ID).
2. Prepend those messages to the current ``input`` so the model has full
conversation context.
3. After the stream completes, extract the new response ID and output
messages from ``response.completed`` and store them in
``self._session_history`` for the next turn.
This in-memory approach avoids the async DB-write race condition that
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")
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["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)
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
# ---------------------------------------------------------------------------
# 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)
# 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.)
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
except Exception as 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,
)
# ---------------------------------------------------------------------------
# ------------------------------------------------------------------
# Main entry point
# ------------------------------------------------------------------
async def run(self) -> None:
"""
Main loop: accept ``response.create`` events sequentially and handle
each one before waiting for the next message.
"""
try:
while True:
try:
message = await self.websocket.receive_text()
except Exception as exc:
verbose_logger.debug(
"ManagedResponsesWS: client disconnected: %s", exc
)
break
await self._process_response_create(message)
except Exception as exc:
verbose_logger.exception("ManagedResponsesWS: unexpected error: %s", exc)
await self._send_error(f"Internal server error: {exc}")

View file

@ -885,6 +885,9 @@ class Router:
self._arealtime = self.factory_function(
litellm._arealtime, call_type="_arealtime"
)
self._aresponses_websocket = self.factory_function(
litellm._aresponses_websocket, call_type="_aresponses_websocket"
)
self.acreate_fine_tuning_job = self.factory_function(
litellm.acreate_fine_tuning_job, call_type="acreate_fine_tuning_job"
)
@ -1848,7 +1851,7 @@ class Router:
finally:
if hasattr(model_response, "close"):
try:
model_response.close()
model_response.close() # type: ignore[reportAttributeAccessIssue]
except BaseException as close_err:
verbose_router_logger.debug(
"stream_with_fallbacks: error closing model_response: %s",
@ -4683,6 +4686,7 @@ class Router:
"afile_delete",
"afile_content",
"_arealtime",
"_aresponses_websocket",
"acreate_fine_tuning_job",
"acancel_fine_tuning_job",
"alist_fine_tuning_jobs",
@ -4855,6 +4859,7 @@ class Router:
"anthropic_messages",
"aresponses",
"_arealtime",
"_aresponses_websocket",
"acreate_fine_tuning_job",
"acancel_fine_tuning_job",
"alist_fine_tuning_jobs",

View file

@ -291,6 +291,7 @@ class CallTypes(str, Enum):
search = "search"
asearch = "asearch"
arealtime = "_arealtime"
aresponses_websocket = "_aresponses_websocket"
create_batch = "create_batch"
acreate_batch = "acreate_batch"
aretrieve_batch = "aretrieve_batch"

View file

@ -16,3 +16,4 @@ exclude = ["litellm/types/*", "litellm/__init__.py", "litellm/proxy/example_conf
"litellm/proxy/utils.py" = ["F401", "PLR0915"]
"litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py" = ["PLR0915"]
"litellm/proxy/guardrails/guardrail_hooks/guardrail_benchmarks/test_eval.py" = ["PLR0915"]
"litellm/responses/streaming_iterator.py" = ["PLR0915"]

View file

@ -0,0 +1,239 @@
"""
E2E tests for OpenAI Responses API WebSocket mode through the LiteLLM proxy.
Connects to ws://0.0.0.0:4000/v1/responses, sends response.create events,
and validates the streamed response events.
Requires:
- Proxy running: python -m litellm.proxy.proxy_cli --config <config> --port 4000
- Model configured in proxy (e.g. gpt-4o-mini)
See: https://developers.openai.com/api/docs/guides/websocket-mode/
"""
import asyncio
import json
import os
import httpx
import pytest
# ── Configuration ─────────────────────────────────────────────────────────────
PROXY_BASE_URL = os.environ.get("LITELLM_PROXY_BASE_URL", "ws://0.0.0.0:4000")
PROXY_MASTER_KEY = os.environ.get("LITELLM_PROXY_KEY", "sk-1234")
PROXY_MODEL = os.environ.get("LITELLM_PROXY_RESPONSES_MODEL", "gpt-4o-mini")
# ──────────────────────────────────────────────────────────────────────────────
def _generate_key() -> str:
"""Generate a key for testing via proxy key/generate endpoint."""
url = "http://0.0.0.0:4000/key/generate"
headers = {
"Authorization": f"Bearer {PROXY_MASTER_KEY}",
"Content-Type": "application/json",
}
response = httpx.post(url, headers=headers, json={}, timeout=10)
if response.status_code != 200:
raise Exception(
f"Key generation failed with status: {response.status_code}. "
"Is the proxy running?"
)
return response.json()["key"]
def _assert_basic_response(events: list[dict], label: str = "") -> None:
"""Assert that events contain response.created, response.completed, and usage."""
prefix = f"[{label}] " if label else ""
types = [e.get("type") for e in events]
assert len(events) > 0, f"{prefix}no events received"
assert "response.created" in types, f"{prefix}missing response.created, got: {types}"
assert "response.completed" in types, (
f"{prefix}missing response.completed, got: {types}"
)
completed = next(e for e in events if e.get("type") == "response.completed")
resp = completed.get("response", {})
assert resp.get("status") == "completed", (
f"{prefix}status != completed: {resp.get('status')}"
)
usage = resp.get("usage", {})
assert usage.get("input_tokens", 0) > 0, f"{prefix}input_tokens=0"
assert usage.get("output_tokens", 0) > 0, f"{prefix}output_tokens=0"
streaming_types = {
"response.output_item.added",
"response.content_part.added",
"response.output_text.delta",
"response.output_item.done",
}
found = streaming_types & set(types)
assert found, f"{prefix}no streaming delta events found, got: {types}"
@pytest.mark.asyncio
async def test_responses_websocket_proxy_basic():
"""
Sends a simple response.create event to the proxy WebSocket endpoint
and validates response.created, response.completed, and streaming events.
"""
try:
import websockets
except ImportError:
pytest.skip("websockets not installed")
try:
key = _generate_key()
except Exception as e:
pytest.skip(
f"Proxy not available or key generation failed: {e}. "
"Start proxy: python -m litellm.proxy.proxy_cli --config <config> --port 4000"
)
url = f"{PROXY_BASE_URL}/v1/responses?model={PROXY_MODEL}"
headers = {"Authorization": f"Bearer {key}"}
events: list[dict] = []
try:
async with websockets.connect(
url, additional_headers=headers, open_timeout=5
) as ws:
payload = {
"type": "response.create",
"model": PROXY_MODEL,
"store": False,
"input": [
{
"type": "message",
"role": "user",
"content": [
{"type": "input_text", "text": "Say hello in one word."}
],
}
],
"tools": [],
}
await ws.send(json.dumps(payload))
for _ in range(50):
msg = await asyncio.wait_for(ws.recv(), timeout=15)
event = json.loads(msg)
events.append(event)
if event.get("type") in (
"response.completed",
"response.failed",
"error",
):
break
except Exception as e:
pytest.fail(
f"WebSocket connection failed: {e}. "
"Ensure proxy is running and model is configured."
)
_assert_basic_response(events, "proxy-basic")
@pytest.mark.asyncio
async def test_responses_websocket_proxy_multi_turn():
"""
Sends two sequential response.create events with previous_response_id
to validate multi-turn conversation over a single WebSocket.
"""
try:
import websockets
except ImportError:
pytest.skip("websockets not installed")
try:
key = _generate_key()
except Exception as e:
pytest.skip(
f"Proxy not available or key generation failed: {e}. "
"Start proxy: python -m litellm.proxy.proxy_cli --config <config> --port 4000"
)
url = f"{PROXY_BASE_URL}/v1/responses?model={PROXY_MODEL}"
headers = {"Authorization": f"Bearer {key}"}
all_events: list[dict] = []
completed: list[dict] = []
first_id = None
try:
async with websockets.connect(
url, additional_headers=headers, open_timeout=5
) as ws:
# Turn 1
await ws.send(
json.dumps(
{
"type": "response.create",
"model": PROXY_MODEL,
"store": True,
"input": [
{
"type": "message",
"role": "user",
"content": [
{
"type": "input_text",
"text": "Remember the number 7. Just say OK.",
}
],
}
],
}
)
)
for _ in range(50):
msg = await asyncio.wait_for(ws.recv(), timeout=15)
event = json.loads(msg)
all_events.append(event)
if event.get("type") == "response.completed":
completed.append(event)
first_id = event.get("response", {}).get("id")
break
if event.get("type") in ("response.failed", "error"):
break
assert first_id, "Turn 1 never completed"
# Turn 2
await ws.send(
json.dumps(
{
"type": "response.create",
"model": PROXY_MODEL,
"store": True,
"previous_response_id": first_id,
"input": [
{
"type": "message",
"role": "user",
"content": [
{
"type": "input_text",
"text": "What number did I tell you to remember?",
}
],
}
],
}
)
)
for _ in range(50):
msg = await asyncio.wait_for(ws.recv(), timeout=15)
event = json.loads(msg)
all_events.append(event)
if event.get("type") == "response.completed":
completed.append(event)
break
if event.get("type") in ("response.failed", "error"):
break
except Exception as e:
pytest.fail(
f"WebSocket multi-turn failed: {e}. "
"Ensure proxy is running and model is configured."
)
assert len(completed) >= 2, (
f"Expected 2 response.completed events, got {len(completed)}"
)
assert completed[1].get("response", {}).get("status") == "completed"