mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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 <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
parent
51e5d9d3bc
commit
1e394db9e9
4 changed files with 406 additions and 9 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
||||
######################################################################
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue