mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
4738911e8f
commit
4f042374a5
7 changed files with 196 additions and 163 deletions
|
|
@ -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 ####################
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue