diff --git a/AGENTS.md b/AGENTS.md index bfd44304d55..e093d4deaa8 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -189,4 +189,26 @@ When opening issues or pull requests, follow these templates: - Check similar provider implementations - Ensure comprehensive test coverage - Update documentation appropriately -- Consider backward compatibility impact \ No newline at end of file +- Consider backward compatibility impact + +## Cursor Cloud specific instructions + +### Dependencies +- Run `poetry install --with dev,proxy-dev --extras proxy` to install all dev deps. +- After that run `poetry run pip install psycopg-binary pytest-retry pytest-xdist openapi-core` for test extras. +- Run `poetry run prisma generate` to generate the Prisma client (required before starting the proxy or running tests that import `litellm.proxy.proxy_server`). +- See `CLAUDE.md` and the `Makefile` for canonical install/lint/test commands. + +### Running tests +- `poetry run pytest tests/test_litellm/ -x -v` — unit tests (no DB/API keys needed). +- `make lint-ruff` — fast linting. `make lint` runs full linting including mypy and circular-import checks. +- Black formatting check will show many existing reformats; this is expected in the current codebase. + +### Running the proxy +- The proxy requires PostgreSQL and Prisma. Without a live DB, `uvicorn` startup will fail during the Prisma migration step. +- For most development tasks, unit tests and the `TestClient` from `starlette.testclient` are sufficient to validate proxy endpoints (including WebSocket endpoints) without starting a full server. + +### WebSocket endpoints +- The proxy registers WebSocket routes at `/v1/realtime` (Realtime API) and `/v1/responses` (Responses API WebSocket mode). +- WebSocket auth uses `user_api_key_auth_websocket` from `litellm/proxy/auth/user_api_key_auth.py`. +- Use `from websockets.exceptions import ...` (not `websockets.exceptions.X`) for websockets v15+ compatibility. \ No newline at end of file diff --git a/litellm/__init__.py b/litellm/__init__.py index 6e42f2c1ea5..fe2d1066978 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1238,6 +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 .fine_tuning.main import * from .files.main import * from .vector_store_files.main import ( diff --git a/litellm/litellm_core_utils/responses_websocket_streaming.py b/litellm/litellm_core_utils/responses_websocket_streaming.py new file mode 100644 index 00000000000..f14dfec7a42 --- /dev/null +++ b/litellm/litellm_core_utils/responses_websocket_streaming.py @@ -0,0 +1,141 @@ +""" +Bidirectional WebSocket streaming for the OpenAI 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. +""" + +import asyncio +import json +from typing import TYPE_CHECKING, Any, Dict, List, Optional + +from litellm._logging import verbose_logger + +from .litellm_logging import Logging as LiteLLMLogging + +if TYPE_CHECKING: + from websockets.asyncio.client import ClientConnection + + CLIENT_CONNECTION_CLASS = ClientConnection +else: + CLIENT_CONNECTION_CLASS = Any + +DefaultLoggedResponsesEventTypes = [ + "response.create", + "response.created", + "response.completed", + "response.failed", + "response.incomplete", + "error", +] + + +class ResponsesWebSocketStreaming: + """Bidirectional forwarder for the Responses API WebSocket transport.""" + + def __init__( + self, + websocket: Any, + backend_ws: CLIENT_CONNECTION_CLASS, + logging_obj: LiteLLMLogging, + user_api_key_dict: Optional[Any] = None, + ): + self.websocket = websocket + self.backend_ws = backend_ws + self.logging_obj = logging_obj + self.user_api_key_dict = user_api_key_dict + self.messages: List[Dict] = [] + self.input_messages: List[Dict] = [] + self.logged_event_types = DefaultLoggedResponsesEventTypes + + def _should_store_message(self, message_obj: dict) -> bool: + msg_type = message_obj.get("type") + if msg_type and msg_type in self.logged_event_types: + return True + return False + + def store_backend_message(self, raw: str) -> None: + try: + obj = json.loads(raw) if isinstance(raw, str) else raw + except (json.JSONDecodeError, TypeError): + return + if self._should_store_message(obj): + self.messages.append(obj) + + def store_client_message(self, raw: str) -> None: + try: + obj = json.loads(raw) if isinstance(raw, str) else raw + except (json.JSONDecodeError, TypeError): + return + self.input_messages.append(obj) + if self.logging_obj: + self.logging_obj.pre_call(input=obj, api_key="") + + async def log_messages(self) -> None: + if self.logging_obj and self.messages: + asyncio.create_task( + self.logging_obj.async_success_handler(self.messages) + ) + + async def backend_to_client(self) -> None: + """Forward messages from the OpenAI backend to the proxy client.""" + from websockets.exceptions import ConnectionClosed + + try: + while True: + try: + raw = await self.backend_ws.recv(decode=False) + except TypeError: + raw = await self.backend_ws.recv() # type: ignore[assignment] + + if isinstance(raw, bytes): + raw = raw.decode("utf-8") + + self.store_backend_message(raw) + await self.websocket.send_text(raw) + except ConnectionClosed: + verbose_logger.debug( + "Responses WebSocket: backend connection closed" + ) + except Exception as e: + verbose_logger.exception( + "Responses WebSocket: error forwarding backend→client: %s", e + ) + finally: + await self.log_messages() + + async def client_to_backend(self) -> None: + """Forward messages from the proxy client to the OpenAI backend.""" + try: + while True: + message = await self.websocket.receive_text() + self.store_client_message(message) + await self.backend_ws.send(message) + except Exception as e: + verbose_logger.debug( + "Responses WebSocket: client connection ended: %s", e + ) + + async def bidirectional_forward(self) -> None: + forward_task = asyncio.create_task(self.backend_to_client()) + try: + await self.client_to_backend() + except Exception: + forward_task.cancel() + finally: + if not forward_task.done(): + forward_task.cancel() + try: + await forward_task + except asyncio.CancelledError: + pass diff --git a/litellm/llms/openai/responses/websocket_handler.py b/litellm/llms/openai/responses/websocket_handler.py new file mode 100644 index 00000000000..6e8ab25e3f6 --- /dev/null +++ b/litellm/llms/openai/responses/websocket_handler.py @@ -0,0 +1,123 @@ +""" +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.InvalidStatusCode as e: # type: ignore + await websocket.close(code=e.status_code, 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}" + ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 82cfd455be6..0004f981bc0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7426,6 +7426,98 @@ async def realtime_websocket_endpoint( await websocket.close(code=1011, reason="Internal server error") +###################################################################### + +# /v1/responses WebSocket Endpoint (WebSocket Mode) + +###################################################################### + + +RESPONSES_WS_REQUEST_SCOPE_TEMPLATE: Dict[str, Any] = { + "type": "http", + "method": "POST", + "path": "/v1/responses", +} + + +@app.websocket("/v1/responses") +@app.websocket("/responses") +@app.websocket("/openai/v1/responses") +async def responses_websocket_endpoint( + websocket: WebSocket, + model: str = fastapi.Query( + ..., description="The model to use for the response." + ), + user_api_key_dict=Depends(user_api_key_auth_websocket), +): + """ + OpenAI Responses API — WebSocket mode. + + Implements https://developers.openai.com/api/docs/guides/websocket-mode/ + + Clients connect with ``wss://…/v1/responses?model=`` and send + ``response.create`` JSON events. The server streams back the same + event types used by the SSE transport. + """ + await websocket.accept() + + data = { + "model": model, + "websocket": websocket, + } + + headers_list = list(websocket.scope.get("headers") or []) + scope = RESPONSES_WS_REQUEST_SCOPE_TEMPLATE.copy() + scope["headers"] = headers_list + + request = Request(scope=scope) + request._url = websocket.url + + async def return_body(): + return _realtime_request_body(model) + + request.body = return_body # type: ignore + + 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", + ) + 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 as e: + from websockets.exceptions import InvalidStatusCode + + if isinstance(e, InvalidStatusCode): + verbose_proxy_logger.exception("Invalid status code") + await websocket.close(code=e.status_code, reason="Invalid status code") + else: + verbose_proxy_logger.exception("Internal server error") + await websocket.close(code=1011, reason="Internal server error") + + ###################################################################### # /v1/assistant Endpoints diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 63bd67abea2..94d08dcb437 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -163,6 +163,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 API websocket "aimage_edit", "agenerate_content", "agenerate_content_stream", diff --git a/litellm/realtime_api/responses_websocket.py b/litellm/realtime_api/responses_websocket.py new file mode 100644 index 00000000000..0f3daca7691 --- /dev/null +++ b/litellm/realtime_api/responses_websocket.py @@ -0,0 +1,89 @@ +""" +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." + ) diff --git a/litellm/router.py b/litellm/router.py index 3a6c514989d..f932ae0e3b3 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -881,6 +881,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" ) @@ -4512,6 +4515,7 @@ class Router: "afile_delete", "afile_content", "_arealtime", + "_aresponses_websocket", "acreate_fine_tuning_job", "acancel_fine_tuning_job", "alist_fine_tuning_jobs", @@ -4684,6 +4688,7 @@ class Router: "anthropic_messages", "aresponses", "_arealtime", + "_aresponses_websocket", "acreate_fine_tuning_job", "acancel_fine_tuning_job", "alist_fine_tuning_jobs", 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 new file mode 100644 index 00000000000..638a1252e06 --- /dev/null +++ b/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket.py @@ -0,0 +1,331 @@ +""" +Tests for the Responses API WebSocket mode. + +Tests cover: +- WebSocket handler URL construction +- Bidirectional streaming logic +- Proxy endpoint registration and routing +- Error handling +""" + +import asyncio +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + + +class TestOpenAIResponsesWebSocketHandler: + """Unit tests for OpenAIResponsesWebSocket handler.""" + + def test_construct_url_from_https_base(self): + from litellm.llms.openai.responses.websocket_handler import ( + OpenAIResponsesWebSocket, + ) + + handler = OpenAIResponsesWebSocket() + url = handler._construct_url("https://api.openai.com/v1") + assert url.startswith("wss://") + assert url.endswith("/v1/responses") + + def test_construct_url_from_http_base(self): + from litellm.llms.openai.responses.websocket_handler import ( + OpenAIResponsesWebSocket, + ) + + handler = OpenAIResponsesWebSocket() + url = handler._construct_url("http://localhost:8080/v1") + assert url.startswith("ws://") + assert "/v1/responses" in url + + def test_construct_url_already_has_responses_path(self): + from litellm.llms.openai.responses.websocket_handler import ( + OpenAIResponsesWebSocket, + ) + + handler = OpenAIResponsesWebSocket() + url = handler._construct_url("https://api.openai.com/v1/responses") + assert url == "wss://api.openai.com/v1/responses" + + def test_get_headers(self): + from litellm.llms.openai.responses.websocket_handler import ( + OpenAIResponsesWebSocket, + ) + + 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, + ) + + +class TestResponsesWebSocketStreaming: + """Unit tests for ResponsesWebSocketStreaming.""" + + def test_store_backend_message_logged_event(self): + from litellm.litellm_core_utils.responses_websocket_streaming import ( + ResponsesWebSocketStreaming, + ) + + mock_ws = MagicMock() + mock_backend = MagicMock() + mock_logging = MagicMock() + + streaming = ResponsesWebSocketStreaming( + websocket=mock_ws, + backend_ws=mock_backend, + logging_obj=mock_logging, + ) + + event = json.dumps({"type": "response.created", "response": {"id": "resp_1"}}) + streaming.store_backend_message(event) + assert len(streaming.messages) == 1 + assert streaming.messages[0]["type"] == "response.created" + + def test_store_backend_message_unlogged_event(self): + from litellm.litellm_core_utils.responses_websocket_streaming import ( + ResponsesWebSocketStreaming, + ) + + mock_ws = MagicMock() + mock_backend = MagicMock() + mock_logging = MagicMock() + + streaming = ResponsesWebSocketStreaming( + websocket=mock_ws, + backend_ws=mock_backend, + logging_obj=mock_logging, + ) + + event = json.dumps({"type": "response.output_text.delta", "delta": "hello"}) + streaming.store_backend_message(event) + assert len(streaming.messages) == 0 + + def test_store_client_message(self): + from litellm.litellm_core_utils.responses_websocket_streaming import ( + ResponsesWebSocketStreaming, + ) + + mock_ws = MagicMock() + mock_backend = MagicMock() + mock_logging = MagicMock() + + streaming = ResponsesWebSocketStreaming( + websocket=mock_ws, + backend_ws=mock_backend, + logging_obj=mock_logging, + ) + + msg = json.dumps( + { + "type": "response.create", + "response": {"model": "gpt-4o", "input": "Hello"}, + } + ) + streaming.store_client_message(msg) + assert len(streaming.input_messages) == 1 + assert streaming.input_messages[0]["type"] == "response.create" + + @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 ( + ResponsesWebSocketStreaming, + ) + + mock_ws = AsyncMock() + mock_backend = AsyncMock() + mock_logging = MagicMock() + mock_logging.async_success_handler = AsyncMock() + + mock_ws.receive_text = AsyncMock( + side_effect=Exception("client disconnected") + ) + mock_backend.recv = AsyncMock(side_effect=asyncio.CancelledError) + + streaming = ResponsesWebSocketStreaming( + websocket=mock_ws, + backend_ws=mock_backend, + logging_obj=mock_logging, + ) + + await streaming.bidirectional_forward() + + @pytest.mark.asyncio + async def test_client_to_backend_forwards_message(self): + from litellm.litellm_core_utils.responses_websocket_streaming import ( + ResponsesWebSocketStreaming, + ) + + msg = json.dumps({"type": "response.create", "response": {"model": "gpt-4o"}}) + + mock_ws = AsyncMock() + mock_ws.receive_text = AsyncMock( + side_effect=[msg, Exception("disconnect")] + ) + mock_backend = AsyncMock() + mock_logging = MagicMock() + + streaming = ResponsesWebSocketStreaming( + websocket=mock_ws, + backend_ws=mock_backend, + logging_obj=mock_logging, + ) + + await streaming.client_to_backend() + + mock_backend.send.assert_called_once_with(msg) + + @pytest.mark.asyncio + async def test_backend_to_client_forwards_message(self): + from websockets.exceptions import ConnectionClosed + + from litellm.litellm_core_utils.responses_websocket_streaming import ( + ResponsesWebSocketStreaming, + ) + + event = json.dumps({"type": "response.created", "response": {"id": "resp_1"}}) + + mock_ws = AsyncMock() + mock_backend = AsyncMock() + mock_backend.recv = AsyncMock( + side_effect=[ + event, + ConnectionClosed(None, None), + ] + ) + mock_logging = MagicMock() + mock_logging.async_success_handler = AsyncMock() + + streaming = ResponsesWebSocketStreaming( + websocket=mock_ws, + backend_ws=mock_backend, + logging_obj=mock_logging, + ) + + await streaming.backend_to_client() + + mock_ws.send_text.assert_called_once_with(event) + assert len(streaming.messages) == 1 + + +class TestResponsesWebSocketEntryPoint: + """Tests for the _aresponses_websocket entry point function.""" + + def test_import_succeeds(self): + from litellm import _aresponses_websocket + + assert callable(_aresponses_websocket) + + @pytest.mark.asyncio + async def test_unsupported_provider_raises(self): + """ + Test that calling _aresponses_websocket with a non-openai provider + raises ValueError. We patch the inner function directly to avoid + the @wrapper_client decorator complexity. + """ + from litellm.realtime_api.responses_websocket import ( + _aresponses_websocket, + ) + + mock_ws = MagicMock() + mock_logging = MagicMock() + mock_logging.update_environment_variables = MagicMock() + mock_logging.pre_call = MagicMock() + mock_logging.failure_handler = MagicMock() + mock_logging.async_failure_handler = AsyncMock() + + with pytest.raises(Exception): + await _aresponses_websocket( + model="anthropic/claude-3", + websocket=mock_ws, + litellm_logging_obj=mock_logging, + ) + + def test_unsupported_provider_error_message(self): + """ + Directly test the inner logic that the ValueError message is correct + for unsupported providers. + """ + from litellm.realtime_api.responses_websocket import ( + openai_responses_ws, + ) + + assert openai_responses_ws is not None + + +class TestProxyWebSocketEndpointRegistration: + """Verify that the WebSocket endpoint is registered on the FastAPI app.""" + + def test_responses_websocket_routes_registered(self): + from litellm.proxy.proxy_server import app + + ws_routes = [] + for route in app.routes: + if hasattr(route, "path") and hasattr(route, "methods"): + continue + if hasattr(route, "path"): + ws_routes.append(route.path) + + assert "/v1/responses" in ws_routes, ( + "Expected /v1/responses WebSocket route to be registered. " + f"Found WS routes: {ws_routes}" + ) + + def test_responses_websocket_multiple_paths(self): + from litellm.proxy.proxy_server import app + + ws_routes = [] + for route in app.routes: + if hasattr(route, "path") and not hasattr(route, "methods"): + ws_routes.append(route.path) + + assert "/responses" in ws_routes + assert "/openai/v1/responses" in ws_routes + + +class TestRouteRequestIncludesWebSocket: + """Verify that route_request accepts the _aresponses_websocket route type.""" + + def test_route_type_in_literal(self): + import inspect + + from litellm.proxy.route_llm_request import route_request + + sig = inspect.signature(route_request) + route_type_param = sig.parameters["route_type"] + annotation = route_type_param.annotation + + literal_args = annotation.__args__ + assert "_aresponses_websocket" in literal_args