From 1e394db9e9ebe188ad87bc7e1c81fb8725bd8e1b Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 25 Feb 2026 19:45:11 +0000 Subject: [PATCH] test: add E2E and live integration tests for Responses WebSocket - test_responses_websocket_e2e.py: 6 tests exercising the full proxy WebSocket route via Starlette TestClient with mocked backends, including complete message flow simulation - test_responses_websocket_live.py: 2 tests for live OpenAI WebSocket (auto-skipped when OPENAI_API_KEY is not set) - Fix websockets v15 deprecation: use InvalidStatus instead of InvalidStatusCode in handler and proxy_server Co-authored-by: Ishaan Jaff --- .../openai/responses/websocket_handler.py | 7 +- litellm/proxy/proxy_server.py | 19 +- .../test_responses_websocket_e2e.py | 242 ++++++++++++++++++ .../test_responses_websocket_live.py | 147 +++++++++++ 4 files changed, 406 insertions(+), 9 deletions(-) create mode 100644 tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_e2e.py create mode 100644 tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_live.py diff --git a/litellm/llms/openai/responses/websocket_handler.py b/litellm/llms/openai/responses/websocket_handler.py index 6e8ab25e3f6..81e61c2e687 100644 --- a/litellm/llms/openai/responses/websocket_handler.py +++ b/litellm/llms/openai/responses/websocket_handler.py @@ -107,8 +107,11 @@ class OpenAIResponsesWebSocket: ) await streaming.bidirectional_forward() - except websockets.exceptions.InvalidStatusCode as e: # type: ignore - await websocket.close(code=e.status_code, reason=str(e)) + 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( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0004f981bc0..99f59e24e30 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7508,14 +7508,19 @@ async def responses_websocket_endpoint( await llm_call except Exception as e: - from websockets.exceptions import InvalidStatusCode + try: + from websockets.exceptions import InvalidStatus - 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") + if isinstance(e, InvalidStatus): + verbose_proxy_logger.exception("Invalid status code") + await websocket.close( + code=e.response.status_code, reason="Invalid status code" + ) + return + except ImportError: + pass + verbose_proxy_logger.exception("Internal server error") + await websocket.close(code=1011, reason="Internal server error") ###################################################################### 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 new file mode 100644 index 00000000000..d26c5fa9c9f --- /dev/null +++ b/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_e2e.py @@ -0,0 +1,242 @@ +""" +End-to-end integration test for the Responses API WebSocket endpoint. + +Uses Starlette's TestClient to exercise the full FastAPI WebSocket route +without requiring an external proxy process or live OpenAI key. +""" + +import asyncio +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from starlette.testclient import TestClient + +from litellm.proxy.proxy_server import app + + +class TestResponsesWebSocketE2E: + """Full round-trip tests through the proxy's WebSocket endpoint.""" + + def test_websocket_accepts_connection_and_routes(self): + """ + Verify the /v1/responses WebSocket: + 1. Accepts the connection + 2. Passes auth + 3. Reaches the routing layer (will fail at model lookup without + a configured router, which is expected) + """ + client = TestClient(app) + + with patch( + "litellm.proxy.proxy_server.user_api_key_auth_websocket" + ) as mock_auth: + mock_auth.return_value = MagicMock( + token="test_token", + user_id="test_user", + team_id=None, + api_key="sk-test", + key_alias=None, + allowed_model_region=None, + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + metadata={}, + ) + + try: + with client.websocket_connect( + "/v1/responses?model=gpt-4o-mini", + headers={"Authorization": "Bearer sk-test"}, + ) as ws: + pass + except Exception as e: + error = str(e) + assert "404" not in error, f"Route not found: {error}" + assert "405" not in error, f"Method not allowed: {error}" + + def test_websocket_openai_path(self): + """Verify /openai/v1/responses WebSocket also works.""" + client = TestClient(app) + + with patch( + "litellm.proxy.proxy_server.user_api_key_auth_websocket" + ) as mock_auth: + mock_auth.return_value = MagicMock( + token="t", user_id="u", team_id=None, api_key="sk-t", + key_alias=None, allowed_model_region=None, + tpm_limit=None, rpm_limit=None, max_budget=None, + spend=0.0, metadata={}, + ) + + try: + with client.websocket_connect( + "/openai/v1/responses?model=gpt-4o-mini", + headers={"Authorization": "Bearer sk-t"}, + ) as ws: + pass + except Exception as e: + assert "404" not in str(e) + + def test_websocket_short_path(self): + """Verify /responses WebSocket also works.""" + client = TestClient(app) + + with patch( + "litellm.proxy.proxy_server.user_api_key_auth_websocket" + ) as mock_auth: + mock_auth.return_value = MagicMock( + token="t", user_id="u", team_id=None, api_key="sk-t", + key_alias=None, allowed_model_region=None, + tpm_limit=None, rpm_limit=None, max_budget=None, + spend=0.0, metadata={}, + ) + + try: + with client.websocket_connect( + "/responses?model=gpt-4o-mini", + headers={"Authorization": "Bearer sk-t"}, + ) as ws: + pass + except Exception as e: + assert "404" not in str(e) + + +class TestResponsesWebSocketHandlerE2E: + """ + Tests that exercise the actual OpenAI handler with a mocked + websockets.connect to simulate the full flow. + """ + + @pytest.mark.asyncio + async def test_full_flow_with_mocked_backend(self): + """ + Simulate a complete WebSocket session: + 1. Client sends response.create + 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 ( + ResponsesWebSocketStreaming, + ) + + backend_events = [ + json.dumps({"type": "response.created", "response": {"id": "resp_test123", "status": "in_progress"}}), + json.dumps({"type": "response.output_text.delta", "delta": "Hello"}), + json.dumps({"type": "response.output_text.delta", "delta": " World"}), + json.dumps({"type": "response.completed", "response": {"id": "resp_test123", "status": "completed", "output": [{"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "Hello World"}]}]}}), + ] + + event_idx = 0 + from websockets.exceptions import ConnectionClosed + + async def mock_recv(decode=False): + nonlocal event_idx + if event_idx < len(backend_events): + msg = backend_events[event_idx] + event_idx += 1 + return msg + raise ConnectionClosed(None, None) + + mock_backend_ws = AsyncMock() + mock_backend_ws.recv = mock_recv + + client_messages_sent = [] + + class MockClientWS: + async def send_text(self, data): + client_messages_sent.append(data) + + async def receive_text(self): + await asyncio.sleep(0.05) + raise Exception("client done") + + mock_client_ws = MockClientWS() + mock_logging = MagicMock() + mock_logging.pre_call = MagicMock() + mock_logging.async_success_handler = AsyncMock() + + streaming = ResponsesWebSocketStreaming( + websocket=mock_client_ws, + backend_ws=mock_backend_ws, + logging_obj=mock_logging, + ) + + await streaming.bidirectional_forward() + + assert len(client_messages_sent) == 4, ( + f"Expected 4 messages forwarded to client, got {len(client_messages_sent)}" + ) + + forwarded_types = [ + json.loads(m).get("type") for m in client_messages_sent + ] + assert forwarded_types == [ + "response.created", + "response.output_text.delta", + "response.output_text.delta", + "response.completed", + ] + + assert len(streaming.messages) == 2 + + @pytest.mark.asyncio + async def test_client_sends_and_backend_receives(self): + """ + Verify client→backend forwarding: client sends response.create + and the backend WS receives it. + """ + from litellm.litellm_core_utils.responses_websocket_streaming import ( + ResponsesWebSocketStreaming, + ) + + request_msg = json.dumps({ + "type": "response.create", + "response": { + "model": "gpt-4o-mini", + "input": "Hello", + }, + }) + + call_count = 0 + + class MockClientWS: + async def receive_text(self): + nonlocal call_count + call_count += 1 + if call_count == 1: + return request_msg + raise Exception("done") + + async def send_text(self, data): + pass + + mock_backend = AsyncMock() + mock_logging = MagicMock() + mock_logging.pre_call = MagicMock() + + streaming = ResponsesWebSocketStreaming( + websocket=MockClientWS(), + backend_ws=mock_backend, + logging_obj=mock_logging, + ) + + await streaming.client_to_backend() + + mock_backend.send.assert_called_once_with(request_msg) + assert len(streaming.input_messages) == 1 + assert streaming.input_messages[0]["type"] == "response.create" + + @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, + ) + + 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" diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_live.py b/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_live.py new file mode 100644 index 00000000000..0599532a50e --- /dev/null +++ b/tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_live.py @@ -0,0 +1,147 @@ +""" +Live integration test for Responses API WebSocket mode through the LiteLLM proxy. + +Requires: + - OPENAI_API_KEY set in environment + - LiteLLM proxy running on localhost:4000 with a model named 'gpt-4o-mini' + +Run standalone: + OPENAI_API_KEY=sk-... python -m pytest tests/test_litellm/proxy/response_api_endpoints/test_responses_websocket_live.py -v -s +""" + +import asyncio +import json +import os + +import pytest + +PROXY_BASE = os.environ.get("LITELLM_PROXY_BASE", "ws://localhost:4000") +PROXY_KEY = os.environ.get("LITELLM_PROXY_KEY", "sk-1234") +MODEL = os.environ.get("LITELLM_WS_TEST_MODEL", "gpt-4o-mini") + + +pytestmark = pytest.mark.skipif( + not os.environ.get("OPENAI_API_KEY"), + reason="OPENAI_API_KEY not set — live OpenAI test skipped", +) + + +@pytest.mark.asyncio +async def test_responses_websocket_live_through_proxy(): + """ + End-to-end: connect to the proxy's /v1/responses WebSocket, send a + response.create event, and collect streamed events until response.completed. + """ + import websockets + + url = f"{PROXY_BASE}/v1/responses?model={MODEL}" + headers = {"Authorization": f"Bearer {PROXY_KEY}"} + + collected_events = [] + + async with websockets.connect(url, additional_headers=headers) as ws: + request_event = { + "type": "response.create", + "response": { + "model": MODEL, + "input": "Say exactly: Hello WebSocket", + }, + } + await ws.send(json.dumps(request_event)) + + async for raw in ws: + event = json.loads(raw) + collected_events.append(event) + event_type = event.get("type", "") + print(f" ← {event_type}") + if event_type in ( + "response.completed", + "response.failed", + "response.incomplete", + "error", + ): + break + + event_types = [e.get("type") for e in collected_events] + print(f"\nAll event types received: {event_types}") + + assert "response.created" in event_types, "Missing response.created event" + assert ( + "response.completed" in event_types + or "response.failed" in event_types + ), "No terminal event received" + + if "response.completed" in event_types: + completed = next(e for e in collected_events if e["type"] == "response.completed") + response_obj = completed.get("response", {}) + assert response_obj.get("id", "").startswith("resp_") + assert response_obj.get("status") == "completed" + print(f"\n✅ Response ID: {response_obj['id']}") + + output = response_obj.get("output", []) + for item in output: + if item.get("type") == "message": + for part in item.get("content", []): + if part.get("type") == "output_text": + print(f" Model said: {part['text']}") + + +@pytest.mark.asyncio +async def test_responses_websocket_continuation(): + """ + Test continuation: send a first response.create, get the response_id, + then send a second response.create with previous_response_id. + """ + import websockets + + url = f"{PROXY_BASE}/v1/responses?model={MODEL}" + headers = {"Authorization": f"Bearer {PROXY_KEY}"} + + async with websockets.connect(url, additional_headers=headers) as ws: + # First turn + await ws.send(json.dumps({ + "type": "response.create", + "response": { + "model": MODEL, + "input": "Remember the number 42.", + }, + })) + + first_response_id = None + async for raw in ws: + event = json.loads(raw) + if event.get("type") == "response.completed": + first_response_id = event.get("response", {}).get("id") + break + if event.get("type") in ("response.failed", "error"): + pytest.skip(f"First turn failed: {event}") + + assert first_response_id, "Did not receive first response ID" + print(f" First response: {first_response_id}") + + # 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, + }, + })) + + second_events = [] + async for raw in ws: + event = json.loads(raw) + second_events.append(event) + if event.get("type") in ( + "response.completed", + "response.failed", + "error", + ): + break + + second_types = [e.get("type") for e in second_events] + print(f" Second turn events: {second_types}") + assert "response.completed" in second_types or "response.failed" in second_types