refactor: move websocket code into litellm/responses/ to match existing organization

- litellm/responses/websocket.py: entry point reusing ProviderConfigManager
  and BaseResponsesAPIConfig.validate_environment/get_complete_url for
  credential resolution — same pattern as responses/main.py
- litellm/responses/websocket_handler.py: handler receives resolved URL +
  auth headers, only owns WS-specific concerns (protocol upgrade, SSL,
  bidirectional forwarding)
- litellm/responses/websocket_streaming.py: moved from litellm_core_utils/

Removed:
- litellm/realtime_api/responses_websocket.py (duplicated provider logic)
- litellm/llms/openai/responses/websocket_handler.py (duplicated credential
  resolution)

Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-02-25 22:35:41 +00:00
parent 1e394db9e9
commit 4738911e8f
9 changed files with 250 additions and 311 deletions

View file

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

View file

@ -1,126 +0,0 @@
"""
OpenAI Responses API WebSocket Mode handler.
Implements the WebSocket transport for OpenAI's Responses API
(wss://api.openai.com/v1/responses).
Protocol summary (from https://developers.openai.com/api/docs/guides/websocket-mode/):
- Client connects via WebSocket to /v1/responses
- Client sends `response.create` events; payload mirrors the HTTP Responses
create body but omits transport-specific fields (`stream`, `background`).
- Server streams back the same SSE event types used by the HTTP streaming
endpoint, wrapped in JSON-framed WebSocket messages.
- Client may continue a conversation by sending another `response.create`
with `previous_response_id` and incremental input.
- A warmup request (`generate: false`) can be sent to pre-populate
connection state without triggering generation.
"""
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.litellm_core_utils.responses_websocket_streaming import (
ResponsesWebSocketStreaming,
)
from litellm.llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
class OpenAIResponsesWebSocket:
"""
Handler for OpenAI Responses API WebSocket connections.
Mirrors the structure of ``OpenAIRealtime`` but targets the
``/v1/responses`` WebSocket endpoint instead of ``/v1/realtime``.
"""
def _get_default_api_base(self) -> str:
return "https://api.openai.com/v1"
def _get_headers(self, api_key: str) -> dict:
return {
"Authorization": f"Bearer {api_key}",
}
def _get_ssl_config(self, 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
def _construct_url(self, api_base: str) -> str:
from httpx import URL
api_base = api_base.replace("https://", "wss://").replace(
"http://", "ws://"
)
url = URL(api_base)
if not url.raw_path.endswith(b"/responses"):
url = url.copy_with(path="/v1/responses")
return str(url)
async def async_responses_websocket(
self,
websocket: Any,
logging_obj: LiteLLMLogging,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
timeout: Optional[float] = None,
user_api_key_dict: Optional[Any] = None,
**kwargs: Any,
) -> None:
import websockets
if api_base is None:
api_base = self._get_default_api_base()
if api_key is None:
raise ValueError("api_key is required for OpenAI Responses WebSocket calls")
url = self._construct_url(api_base)
headers = self._get_headers(api_key)
ssl_config = self._get_ssl_config(url)
logging_obj.pre_call(
input=None,
api_key=api_key,
additional_args={
"api_base": url,
"headers": headers,
"complete_input_dict": {},
},
)
try:
async with websockets.connect( # type: ignore
url,
additional_headers=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 websockets.exceptions.InvalidStatus as e: # type: ignore
status = getattr(
getattr(e, "response", None), "status_code", 1011
)
await websocket.close(code=status, reason=str(e))
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,89 +0,0 @@
"""
Entry point for the Responses API WebSocket mode.
Analogous to ``litellm.realtime_api.main._arealtime`` but for the
``/v1/responses`` WebSocket transport.
Currently supports OpenAI only. Other providers can be added as they ship
their own WebSocket modes.
"""
from typing import Any, Optional, cast
import litellm
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.openai.responses.websocket_handler import OpenAIResponsesWebSocket
from litellm.secret_managers.main import get_secret_str
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import client as wrapper_client
openai_responses_ws = OpenAIResponsesWebSocket()
@wrapper_client
async def _aresponses_websocket(
model: str,
websocket: Any,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
**kwargs: Any,
) -> None:
"""
Private function for the Responses API WebSocket transport.
For PROXY use only.
"""
headers = cast(Optional[dict], kwargs.get("headers"))
extra_headers = cast(Optional[dict], kwargs.get("extra_headers"))
if headers is None:
headers = {}
if extra_headers is not None:
headers.update(extra_headers)
litellm_logging_obj: LiteLLMLogging = 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 = 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,
)
if _custom_llm_provider == "openai":
resolved_api_base = (
dynamic_api_base
or litellm_params.api_base
or litellm.api_base
or "https://api.openai.com/v1"
)
resolved_api_key = (
dynamic_api_key
or litellm.api_key
or litellm.openai_key
or get_secret_str("OPENAI_API_KEY")
)
await openai_responses_ws.async_responses_websocket(
websocket=websocket,
logging_obj=litellm_logging_obj,
api_base=resolved_api_base,
api_key=resolved_api_key,
user_api_key_dict=kwargs.get("user_api_key_dict"),
)
else:
raise ValueError(
f"Responses API WebSocket mode is not supported for provider: {_custom_llm_provider}. "
"Currently only 'openai' is supported."
)

View file

@ -0,0 +1,104 @@
"""
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.
"""
from typing import Any, Optional, cast
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.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()
@wrapper_client
async def _aresponses_websocket(
model: str,
websocket: Any,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
**kwargs: Any,
) -> 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"))
if headers is None:
headers = {}
if extra_headers is not None:
headers.update(extra_headers)
litellm_logging_obj: LiteLLMLogging = 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 = get_llm_provider(
model=model,
api_base=api_base,
api_key=api_key,
)
if dynamic_api_key is not None:
litellm_params.api_key = dynamic_api_key
if dynamic_api_base is not None:
litellm_params.api_base = dynamic_api_base
litellm_logging_obj.update_environment_variables(
model=model,
user=user,
optional_params={},
litellm_params=litellm_params_dict,
custom_llm_provider=custom_llm_provider,
)
responses_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=model,
provider=LlmProviders(custom_llm_provider),
)
)
if responses_config is None:
raise ValueError(
f"Responses API WebSocket mode is not supported for provider: "
f"{custom_llm_provider}. No responses config found."
)
auth_headers = responses_config.validate_environment(
headers={}, model=model, litellm_params=litellm_params
)
http_url = responses_config.get_complete_url(
api_base=litellm_params.api_base,
litellm_params=litellm_params_dict,
)
await openai_responses_ws_handler.async_responses_websocket(
websocket=websocket,
logging_obj=litellm_logging_obj,
http_url=http_url,
auth_headers=auth_headers,
user_api_key_dict=kwargs.get("user_api_key_dict"),
)

View file

@ -0,0 +1,95 @@
"""
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

@ -21,7 +21,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional
from litellm._logging import verbose_logger
from .litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
if TYPE_CHECKING:
from websockets.asyncio.client import ClientConnection

View file

@ -16,85 +16,48 @@ import pytest
class TestOpenAIResponsesWebSocketHandler:
"""Unit tests for OpenAIResponsesWebSocket handler."""
"""Unit tests for OpenAIResponsesWebSocketHandler."""
def test_construct_url_from_https_base(self):
from litellm.llms.openai.responses.websocket_handler import (
OpenAIResponsesWebSocket,
def test_http_url_to_ws_https(self):
from litellm.responses.websocket_handler import (
OpenAIResponsesWebSocketHandler,
)
handler = OpenAIResponsesWebSocket()
url = handler._construct_url("https://api.openai.com/v1")
assert url.startswith("wss://")
assert url.endswith("/v1/responses")
assert OpenAIResponsesWebSocketHandler._http_url_to_ws(
"https://api.openai.com/v1/responses"
) == "wss://api.openai.com/v1/responses"
def test_construct_url_from_http_base(self):
from litellm.llms.openai.responses.websocket_handler import (
OpenAIResponsesWebSocket,
def test_http_url_to_ws_http(self):
from litellm.responses.websocket_handler import (
OpenAIResponsesWebSocketHandler,
)
handler = OpenAIResponsesWebSocket()
url = handler._construct_url("http://localhost:8080/v1")
assert url.startswith("ws://")
assert "/v1/responses" in url
result = OpenAIResponsesWebSocketHandler._http_url_to_ws(
"http://localhost:4000/v1/responses"
)
assert result == "ws://localhost:4000/v1/responses"
def test_construct_url_already_has_responses_path(self):
from litellm.llms.openai.responses.websocket_handler import (
OpenAIResponsesWebSocket,
def test_ssl_config_ws(self):
from litellm.responses.websocket_handler import (
OpenAIResponsesWebSocketHandler,
)
handler = OpenAIResponsesWebSocket()
url = handler._construct_url("https://api.openai.com/v1/responses")
assert url == "wss://api.openai.com/v1/responses"
assert OpenAIResponsesWebSocketHandler._get_ssl_config("ws://localhost:8080") is None
def test_get_headers(self):
from litellm.llms.openai.responses.websocket_handler import (
OpenAIResponsesWebSocket,
def test_ssl_config_wss(self):
from litellm.responses.websocket_handler import (
OpenAIResponsesWebSocketHandler,
)
handler = OpenAIResponsesWebSocket()
headers = handler._get_headers("sk-test-key")
assert headers == {"Authorization": "Bearer sk-test-key"}
def test_get_default_api_base(self):
from litellm.llms.openai.responses.websocket_handler import (
OpenAIResponsesWebSocket,
)
handler = OpenAIResponsesWebSocket()
assert handler._get_default_api_base() == "https://api.openai.com/v1"
def test_get_ssl_config_ws(self):
from litellm.llms.openai.responses.websocket_handler import (
OpenAIResponsesWebSocket,
)
handler = OpenAIResponsesWebSocket()
assert handler._get_ssl_config("ws://localhost:8080") is None
@pytest.mark.asyncio
async def test_missing_api_key_raises(self):
from litellm.llms.openai.responses.websocket_handler import (
OpenAIResponsesWebSocket,
)
handler = OpenAIResponsesWebSocket()
mock_ws = MagicMock()
mock_logging = MagicMock()
with pytest.raises(ValueError, match="api_key is required"):
await handler.async_responses_websocket(
websocket=mock_ws,
logging_obj=mock_logging,
api_key=None,
)
result = OpenAIResponsesWebSocketHandler._get_ssl_config("wss://api.openai.com/v1/responses")
assert result is not None
class TestResponsesWebSocketStreaming:
"""Unit tests for ResponsesWebSocketStreaming."""
def test_store_backend_message_logged_event(self):
from litellm.litellm_core_utils.responses_websocket_streaming import (
from litellm.responses.websocket_streaming import (
ResponsesWebSocketStreaming,
)
@ -114,7 +77,7 @@ class TestResponsesWebSocketStreaming:
assert streaming.messages[0]["type"] == "response.created"
def test_store_backend_message_unlogged_event(self):
from litellm.litellm_core_utils.responses_websocket_streaming import (
from litellm.responses.websocket_streaming import (
ResponsesWebSocketStreaming,
)
@ -133,7 +96,7 @@ class TestResponsesWebSocketStreaming:
assert len(streaming.messages) == 0
def test_store_client_message(self):
from litellm.litellm_core_utils.responses_websocket_streaming import (
from litellm.responses.websocket_streaming import (
ResponsesWebSocketStreaming,
)
@ -160,7 +123,7 @@ class TestResponsesWebSocketStreaming:
@pytest.mark.asyncio
async def test_bidirectional_forward_client_disconnect(self):
"""When the client disconnects, the forward task should be cancelled."""
from litellm.litellm_core_utils.responses_websocket_streaming import (
from litellm.responses.websocket_streaming import (
ResponsesWebSocketStreaming,
)
@ -184,7 +147,7 @@ class TestResponsesWebSocketStreaming:
@pytest.mark.asyncio
async def test_client_to_backend_forwards_message(self):
from litellm.litellm_core_utils.responses_websocket_streaming import (
from litellm.responses.websocket_streaming import (
ResponsesWebSocketStreaming,
)
@ -211,7 +174,7 @@ class TestResponsesWebSocketStreaming:
async def test_backend_to_client_forwards_message(self):
from websockets.exceptions import ConnectionClosed
from litellm.litellm_core_utils.responses_websocket_streaming import (
from litellm.responses.websocket_streaming import (
ResponsesWebSocketStreaming,
)
@ -255,7 +218,7 @@ class TestResponsesWebSocketEntryPoint:
raises ValueError. We patch the inner function directly to avoid
the @wrapper_client decorator complexity.
"""
from litellm.realtime_api.responses_websocket import (
from litellm.responses.websocket import (
_aresponses_websocket,
)
@ -278,11 +241,11 @@ class TestResponsesWebSocketEntryPoint:
Directly test the inner logic that the ValueError message is correct
for unsupported providers.
"""
from litellm.realtime_api.responses_websocket import (
openai_responses_ws,
from litellm.responses.websocket import (
openai_responses_ws_handler,
)
assert openai_responses_ws is not None
assert openai_responses_ws_handler is not None
class TestProxyWebSocketEndpointRegistration:

View file

@ -117,7 +117,7 @@ class TestResponsesWebSocketHandlerE2E:
2. Backend sends back response.created, output_text.delta, response.completed
3. Verify all messages are forwarded correctly
"""
from litellm.litellm_core_utils.responses_websocket_streaming import (
from litellm.responses.websocket_streaming import (
ResponsesWebSocketStreaming,
)
@ -187,7 +187,7 @@ class TestResponsesWebSocketHandlerE2E:
Verify client→backend forwarding: client sends response.create
and the backend WS receives it.
"""
from litellm.litellm_core_utils.responses_websocket_streaming import (
from litellm.responses.websocket_streaming import (
ResponsesWebSocketStreaming,
)
@ -231,12 +231,10 @@ class TestResponsesWebSocketHandlerE2E:
@pytest.mark.asyncio
async def test_handler_constructs_correct_wss_url(self):
"""Verify OpenAIResponsesWebSocket builds the correct WSS URL."""
from litellm.llms.openai.responses.websocket_handler import (
OpenAIResponsesWebSocket,
from litellm.responses.websocket_handler import (
OpenAIResponsesWebSocketHandler,
)
handler = OpenAIResponsesWebSocket()
assert handler._construct_url("https://api.openai.com/v1") == "wss://api.openai.com/v1/responses"
assert handler._construct_url("http://localhost:4000/v1") == "ws://localhost:4000/v1/responses"
assert handler._construct_url("https://custom.endpoint.com/v1/responses") == "wss://custom.endpoint.com/v1/responses"
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"

View file

@ -42,10 +42,8 @@ async def test_responses_websocket_live_through_proxy():
async with websockets.connect(url, additional_headers=headers) as ws:
request_event = {
"type": "response.create",
"response": {
"model": MODEL,
"input": "Say exactly: Hello WebSocket",
},
"model": MODEL,
"input": "Say exactly: Hello WebSocket",
}
await ws.send(json.dumps(request_event))
@ -101,10 +99,8 @@ async def test_responses_websocket_continuation():
# First turn
await ws.send(json.dumps({
"type": "response.create",
"response": {
"model": MODEL,
"input": "Remember the number 42.",
},
"model": MODEL,
"input": "Remember the number 42.",
}))
first_response_id = None
@ -122,13 +118,11 @@ async def test_responses_websocket_continuation():
# Second turn — continuation
await ws.send(json.dumps({
"type": "response.create",
"response": {
"model": MODEL,
"input": [
{"type": "message", "role": "user", "content": "What number did I mention?"},
],
"previous_response_id": first_response_id,
},
"model": MODEL,
"input": [
{"type": "message", "role": "user", "content": "What number did I mention?"},
],
"previous_response_id": first_response_id,
}))
second_events = []