From 4f042374a585bec7f1a07e98d2d1c009bc25fd11 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 25 Feb 2026 22:51:49 +0000 Subject: [PATCH] refactor: make WebSocket responses extensible via BaseResponsesAPIConfig transforms MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 --- .../llms/base_llm/responses/transformation.py | 53 +++++++++++ litellm/llms/custom_httpx/llm_http_handler.py | 74 +++++++++++++++ litellm/responses/websocket.py | 33 +++---- litellm/responses/websocket_handler.py | 95 ------------------- litellm/responses/websocket_streaming.py | 37 +++++--- .../test_responses_websocket.py | 56 +++++------ .../test_responses_websocket_e2e.py | 11 ++- 7 files changed, 196 insertions(+), 163 deletions(-) delete mode 100644 litellm/responses/websocket_handler.py diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py index 7a4da985528..f48478cfb75 100644 --- a/litellm/llms/base_llm/responses/transformation.py +++ b/litellm/llms/base_llm/responses/transformation.py @@ -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 #################### + ######################################################### diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 7267532933d..8616412f0e0 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, diff --git a/litellm/responses/websocket.py b/litellm/responses/websocket.py index 714bce6a26f..48b3145c774 100644 --- a/litellm/responses/websocket.py +++ b/litellm/responses/websocket.py @@ -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"), ) diff --git a/litellm/responses/websocket_handler.py b/litellm/responses/websocket_handler.py deleted file mode 100644 index d7234052be8..00000000000 --- a/litellm/responses/websocket_handler.py +++ /dev/null @@ -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}" - ) diff --git a/litellm/responses/websocket_streaming.py b/litellm/responses/websocket_streaming.py index 0c461fcfef7..78fed254cff 100644 --- a/litellm/responses/websocket_streaming.py +++ b/litellm/responses/websocket_streaming.py @@ -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( diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket.py b/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket.py index 917b9d0dd07..7587b60805d 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket.py @@ -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: diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_e2e.py b/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_e2e.py index 49b7d460854..c7edfb19340 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_e2e.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_e2e.py @@ -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"