refactor: make WebSocket responses extensible via BaseResponsesAPIConfig transforms

Architecture now matches the realtime API pattern (BaseLLMHTTPHandler +
BaseRealtimeConfig transforms):

BaseResponsesAPIConfig (new non-abstract hooks with passthrough defaults):
  - get_websocket_url()         — HTTP→WSS URL, override for custom paths
  - transform_websocket_client_message()  — client→backend message transform
  - transform_websocket_backend_message() — backend→client message transform

BaseLLMHTTPHandler.async_responses_websocket():
  - Generic handler that accepts any BaseResponsesAPIConfig
  - Calls config.transform_* hooks on every message
  - Same pattern as async_realtime() with BaseRealtimeConfig

ResponsesWebSocketStreaming:
  - Now accepts optional provider_config, calls transforms in the loop

When Azure (or any provider) ships WebSocket mode on /responses, they
just override the 3 hooks on their config — zero handler code changes.

Deleted: litellm/responses/websocket_handler.py (logic moved to
BaseLLMHTTPHandler)

Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-02-25 22:51:49 +00:00
parent 4738911e8f
commit 4f042374a5
7 changed files with 196 additions and 163 deletions

View file

@ -269,3 +269,56 @@ class BaseResponsesAPIConfig(ABC):
#########################################################
########## END COMPACT RESPONSE API TRANSFORMATION ######
#########################################################
#########################################################
########## WEBSOCKET MODE HOOKS ########################
#########################################################
def get_websocket_url(
self,
api_base: Optional[str],
litellm_params: dict,
) -> str:
"""
Return the WSS URL for the Responses API WebSocket mode.
Default: take ``get_complete_url`` (HTTP) and swap the scheme.
Override for providers that use a different WS path or query params.
"""
http_url = self.get_complete_url(api_base=api_base, litellm_params=litellm_params)
return (
http_url
.replace("https://", "wss://")
.replace("http://", "ws://")
)
def transform_websocket_client_message(
self,
message: str,
model: str,
) -> str:
"""
Transform a client→backend message before forwarding.
Default: pass through unchanged.
Override for providers that need different field names, extra wrapping,
or model-name rewriting.
"""
return message
def transform_websocket_backend_message(
self,
message: str,
model: str,
) -> str:
"""
Transform a backend→client message before forwarding.
Default: pass through unchanged.
Override for providers that return non-OpenAI event shapes.
"""
return message
#########################################################
########## END WEBSOCKET MODE HOOKS ####################
#########################################################

View file

@ -4708,6 +4708,80 @@ 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",
ws_url: str,
auth_headers: dict,
user_api_key_dict: Optional[Any] = None,
):
"""
Generic WebSocket handler for the Responses API WebSocket mode.
Follows the same pattern as ``async_realtime`` but uses
``BaseResponsesAPIConfig`` transform hooks:
- ``transform_websocket_client_message``
- ``transform_websocket_backend_message``
Provider-specific behaviour (URL construction, auth, message
transforms) is delegated to ``responses_api_provider_config``.
"""
import websockets
from litellm.responses.websocket_streaming import (
ResponsesWebSocketStreaming,
)
ssl_config: Any = None
if ws_url.startswith("wss://"):
ssl_config = get_shared_realtime_ssl_context()
if ssl_config is False:
ssl_config = True
logging_obj.pre_call(
input=None,
api_key="",
additional_args={
"api_base": ws_url,
"headers": auth_headers,
"complete_input_dict": {},
},
)
try:
async with websockets.connect( # type: ignore
ws_url,
additional_headers=auth_headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_config,
) as backend_ws:
streaming = ResponsesWebSocketStreaming(
websocket=websocket,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=responses_api_provider_config,
model=model,
user_api_key_dict=user_api_key_dict,
)
await streaming.bidirectional_forward()
except Exception as e:
verbose_logger.exception("Responses WebSocket error: %s", e)
try:
await websocket.close(
code=1011, reason=f"Internal server error: {str(e)}"
)
except RuntimeError as close_error:
if "already completed" not in str(
close_error
) and "websocket.close" not in str(close_error):
raise Exception(
f"Unexpected error while closing WebSocket: {close_error}"
)
def image_edit_handler(
self,
model: str,

View file

@ -1,13 +1,13 @@
"""
Entry point for the Responses API WebSocket mode.
Lives alongside ``litellm.responses.main`` and reuses the same
``ProviderConfigManager`` / ``get_llm_provider`` patterns for
credential resolution.
Currently supports OpenAI. Other providers can be added by
implementing ``get_websocket_url`` / ``get_websocket_headers``
on their ``BaseResponsesAPIConfig`` subclass.
Follows the same pattern as ``litellm.realtime_api.main._arealtime``:
1. Resolve provider via ``get_llm_provider``
2. Get ``BaseResponsesAPIConfig`` via ``ProviderConfigManager``
3. Call ``config.validate_environment`` for auth headers
4. Call ``config.get_websocket_url`` for the WSS URL
5. Delegate to ``BaseLLMHTTPHandler.async_responses_websocket``
which runs the generic WS loop with config-driven transforms.
"""
from typing import Any, Optional, cast
@ -16,13 +16,12 @@ from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager, client as wrapper_client
from .websocket_handler import OpenAIResponsesWebSocketHandler
openai_responses_ws_handler = OpenAIResponsesWebSocketHandler()
base_llm_http_handler = BaseLLMHTTPHandler()
@wrapper_client
@ -35,12 +34,6 @@ async def _aresponses_websocket(
) -> None:
"""
Responses API WebSocket transport. For proxy use only.
Follows the same provider-resolution flow as ``litellm.responses.main.responses``:
1. ``get_llm_provider`` → model, provider, dynamic key/base
2. ``ProviderConfigManager.get_provider_responses_api_config`` → config
3. ``config.validate_environment`` → auth headers
4. ``config.get_complete_url`` → HTTP URL (converted to WSS)
"""
headers = cast(Optional[dict], kwargs.get("headers"))
extra_headers = cast(Optional[dict], kwargs.get("extra_headers"))
@ -90,15 +83,17 @@ async def _aresponses_websocket(
headers={}, model=model, litellm_params=litellm_params
)
http_url = responses_config.get_complete_url(
ws_url = responses_config.get_websocket_url(
api_base=litellm_params.api_base,
litellm_params=litellm_params_dict,
)
await openai_responses_ws_handler.async_responses_websocket(
await base_llm_http_handler.async_responses_websocket(
model=model,
websocket=websocket,
logging_obj=litellm_logging_obj,
http_url=http_url,
responses_api_provider_config=responses_config,
ws_url=ws_url,
auth_headers=auth_headers,
user_api_key_dict=kwargs.get("user_api_key_dict"),
)

View file

@ -1,95 +0,0 @@
"""
WebSocket handler for the OpenAI Responses API WebSocket mode.
Receives a fully-resolved HTTP URL and auth headers from the entry
point in ``litellm.responses.websocket`` (which uses the same
``BaseResponsesAPIConfig`` credential-resolution as the HTTP path).
This handler only owns the WebSocket-specific concerns:
- Converting the HTTP URL to a WSS URL
- Opening the ``websockets`` connection
- Bidirectional message forwarding via ``ResponsesWebSocketStreaming``
"""
from typing import Any, Optional
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
from litellm.responses.websocket_streaming import ResponsesWebSocketStreaming
class OpenAIResponsesWebSocketHandler:
"""
Handles the WebSocket connection lifecycle for the Responses API.
Analogous to ``OpenAIRealtime`` but for ``/v1/responses``.
"""
@staticmethod
def _http_url_to_ws(http_url: str) -> str:
return (
http_url
.replace("https://", "wss://")
.replace("http://", "ws://")
)
@staticmethod
def _get_ssl_config(url: str) -> Any:
if url.startswith("ws://"):
return None
ssl_config = get_shared_realtime_ssl_context()
if ssl_config is False:
return True
return ssl_config
async def async_responses_websocket(
self,
websocket: Any,
logging_obj: LiteLLMLogging,
http_url: str,
auth_headers: dict,
user_api_key_dict: Optional[Any] = None,
) -> None:
import websockets
ws_url = self._http_url_to_ws(http_url)
ssl_config = self._get_ssl_config(ws_url)
logging_obj.pre_call(
input=None,
api_key="",
additional_args={
"api_base": ws_url,
"headers": auth_headers,
"complete_input_dict": {},
},
)
try:
async with websockets.connect( # type: ignore
ws_url,
additional_headers=auth_headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_config,
) as backend_ws:
streaming = ResponsesWebSocketStreaming(
websocket=websocket,
backend_ws=backend_ws,
logging_obj=logging_obj,
user_api_key_dict=user_api_key_dict,
)
await streaming.bidirectional_forward()
except Exception as e:
try:
await websocket.close(
code=1011, reason=f"Internal server error: {str(e)}"
)
except RuntimeError as close_error:
if "already completed" not in str(
close_error
) and "websocket.close" not in str(close_error):
raise Exception(
f"Unexpected error while closing WebSocket: {close_error}"
)

View file

@ -1,18 +1,12 @@
"""
Bidirectional WebSocket streaming for the OpenAI Responses API WebSocket mode.
Bidirectional WebSocket streaming for the Responses API WebSocket mode.
Unlike the Realtime API streaming (which handles audio sessions, VAD, and
guardrail interception on transcription events), the Responses WebSocket
protocol is simpler:
Client ──response.create──▸ Backend
Client ◂──streaming events── Backend
The client sends ``response.create`` JSON messages. The backend sends back
streaming response events (the same events used by the SSE transport, but
delivered as individual WebSocket text frames).
This module handles the bidirectional forwarding and logging.
Follows the same pattern as ``RealTimeStreaming`` in
``litellm.litellm_core_utils.realtime_streaming``:
- Accepts an optional ``provider_config`` (``BaseResponsesAPIConfig``)
- Calls ``transform_websocket_client_message`` on outbound messages
- Calls ``transform_websocket_backend_message`` on inbound messages
- If no config is supplied, messages pass through unchanged (OpenAI)
"""
import asyncio
@ -20,10 +14,10 @@ import json
from typing import TYPE_CHECKING, Any, Dict, List, Optional
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
if TYPE_CHECKING:
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from websockets.asyncio.client import ClientConnection
CLIENT_CONNECTION_CLASS = ClientConnection
@ -48,11 +42,15 @@ class ResponsesWebSocketStreaming:
websocket: Any,
backend_ws: CLIENT_CONNECTION_CLASS,
logging_obj: LiteLLMLogging,
provider_config: Optional["BaseResponsesAPIConfig"] = None,
model: str = "",
user_api_key_dict: Optional[Any] = None,
):
self.websocket = websocket
self.backend_ws = backend_ws
self.logging_obj = logging_obj
self.provider_config = provider_config
self.model = model
self.user_api_key_dict = user_api_key_dict
self.messages: List[Dict] = []
self.input_messages: List[Dict] = []
@ -101,6 +99,11 @@ class ResponsesWebSocketStreaming:
if isinstance(raw, bytes):
raw = raw.decode("utf-8")
if self.provider_config:
raw = self.provider_config.transform_websocket_backend_message(
raw, self.model
)
self.store_backend_message(raw)
await self.websocket.send_text(raw)
except ConnectionClosed:
@ -120,6 +123,12 @@ class ResponsesWebSocketStreaming:
while True:
message = await self.websocket.receive_text()
self.store_client_message(message)
if self.provider_config:
message = self.provider_config.transform_websocket_client_message(
message, self.model
)
await self.backend_ws.send(message)
except Exception as e:
verbose_logger.debug(

View file

@ -15,42 +15,40 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
class TestOpenAIResponsesWebSocketHandler:
"""Unit tests for OpenAIResponsesWebSocketHandler."""
class TestBaseResponsesAPIConfigWebSocket:
"""Test the WebSocket hooks on BaseResponsesAPIConfig."""
def test_http_url_to_ws_https(self):
from litellm.responses.websocket_handler import (
OpenAIResponsesWebSocketHandler,
def test_get_websocket_url_converts_https(self):
from litellm.llms.openai.responses.transformation import (
OpenAIResponsesAPIConfig,
)
assert OpenAIResponsesWebSocketHandler._http_url_to_ws(
"https://api.openai.com/v1/responses"
) == "wss://api.openai.com/v1/responses"
config = OpenAIResponsesAPIConfig()
url = config.get_websocket_url(api_base=None, litellm_params={})
assert url.startswith("wss://")
assert url.endswith("/responses")
def test_http_url_to_ws_http(self):
from litellm.responses.websocket_handler import (
OpenAIResponsesWebSocketHandler,
def test_get_websocket_url_converts_http(self):
from litellm.llms.openai.responses.transformation import (
OpenAIResponsesAPIConfig,
)
result = OpenAIResponsesWebSocketHandler._http_url_to_ws(
"http://localhost:4000/v1/responses"
config = OpenAIResponsesAPIConfig()
url = config.get_websocket_url(
api_base="http://localhost:4000/v1", litellm_params={}
)
assert result == "ws://localhost:4000/v1/responses"
assert url.startswith("ws://")
assert "/v1/responses" in url
def test_ssl_config_ws(self):
from litellm.responses.websocket_handler import (
OpenAIResponsesWebSocketHandler,
def test_default_transforms_are_passthrough(self):
from litellm.llms.openai.responses.transformation import (
OpenAIResponsesAPIConfig,
)
assert OpenAIResponsesWebSocketHandler._get_ssl_config("ws://localhost:8080") is None
def test_ssl_config_wss(self):
from litellm.responses.websocket_handler import (
OpenAIResponsesWebSocketHandler,
)
result = OpenAIResponsesWebSocketHandler._get_ssl_config("wss://api.openai.com/v1/responses")
assert result is not None
config = OpenAIResponsesAPIConfig()
msg = '{"type":"response.create","model":"gpt-4o"}'
assert config.transform_websocket_client_message(msg, "gpt-4o") == msg
assert config.transform_websocket_backend_message(msg, "gpt-4o") == msg
class TestResponsesWebSocketStreaming:
@ -241,11 +239,9 @@ class TestResponsesWebSocketEntryPoint:
Directly test the inner logic that the ValueError message is correct
for unsupported providers.
"""
from litellm.responses.websocket import (
openai_responses_ws_handler,
)
from litellm.responses.websocket import base_llm_http_handler
assert openai_responses_ws_handler is not None
assert hasattr(base_llm_http_handler, "async_responses_websocket")
class TestProxyWebSocketEndpointRegistration:

View file

@ -231,10 +231,11 @@ class TestResponsesWebSocketHandlerE2E:
@pytest.mark.asyncio
async def test_handler_constructs_correct_wss_url(self):
"""Verify OpenAIResponsesWebSocket builds the correct WSS URL."""
from litellm.responses.websocket_handler import (
OpenAIResponsesWebSocketHandler,
from litellm.llms.openai.responses.transformation import (
OpenAIResponsesAPIConfig,
)
assert OpenAIResponsesWebSocketHandler._http_url_to_ws("https://api.openai.com/v1/responses") == "wss://api.openai.com/v1/responses"
assert OpenAIResponsesWebSocketHandler._http_url_to_ws("http://localhost:4000/v1/responses") == "ws://localhost:4000/v1/responses"
assert OpenAIResponsesWebSocketHandler._http_url_to_ws("https://custom.endpoint.com/v1/responses") == "wss://custom.endpoint.com/v1/responses"
config = OpenAIResponsesAPIConfig()
assert config.get_websocket_url(api_base="https://api.openai.com/v1", litellm_params={}) == "wss://api.openai.com/v1/responses"
assert config.get_websocket_url(api_base="http://localhost:4000/v1", litellm_params={}) == "ws://localhost:4000/v1/responses"
assert config.get_websocket_url(api_base="https://custom.endpoint.com/v1", litellm_params={}) == "wss://custom.endpoint.com/v1/responses"