test(proxy): assert recorded websocket outcomes and simplify session reaping

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mrinal 2026-10-03 22:03:48 +00:00
parent 8578a2852c
commit b5f8f2b8b8
2 changed files with 57 additions and 18 deletions

View file

@ -1654,13 +1654,6 @@ async def responses_websocket_endpoint(
session_task.cancel()
with contextlib.suppress(Exception):
await websocket.close(code=1000, reason="Session duration limit reached")
with contextlib.suppress(asyncio.CancelledError, Exception):
await session_task
except asyncio.CancelledError:
session_task.cancel()
with contextlib.suppress(asyncio.CancelledError, Exception):
await session_task
raise
finally:
if not session_task.done():
session_task.cancel()

View file

@ -728,11 +728,17 @@ class TestResponsesWSFirstFrameModelAuth:
async def fake_llm_call():
return None
authenticated_models: Final[list[str]] = []
async def record_model_auth(*, model: str, **_kwargs: object) -> None:
authenticated_models.append(model)
with (
patch(
"litellm.proxy.response_api_endpoints.endpoints._enforce_responses_ws_first_frame_model_auth",
new_callable=AsyncMock,
) as mock_model_auth,
side_effect=record_model_auth,
),
patch(
"litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing",
return_value=processor,
@ -749,7 +755,7 @@ class TestResponsesWSFirstFrameModelAuth:
user_api_key_dict=MagicMock(),
)
mock_model_auth.assert_awaited_once()
assert authenticated_models == ["gpt-4o-mini"]
@pytest.mark.asyncio
@pytest.mark.parametrize("nested", [False, True])
@ -1130,6 +1136,12 @@ class TestResponsesWSSessionLimit:
await asyncio.sleep(60)
ws = self._ws(silent_socket)
close_calls: Final[list[tuple[int, str]]] = []
async def record_client_close(*, code: int, reason: str) -> None:
close_calls.append((code, reason))
ws.close = AsyncMock(side_effect=record_client_close)
with patch(
"litellm.proxy.response_api_endpoints.endpoints._resolve_responses_ws_session_limit_seconds",
@ -1137,7 +1149,7 @@ class TestResponsesWSSessionLimit:
):
await responses_websocket_endpoint(websocket=ws, model="gpt-4o-mini", user_api_key_dict=MagicMock())
ws.close.assert_awaited_once_with(code=1000, reason="Session duration limit reached")
assert close_calls == [(1000, "Session duration limit reached")]
@pytest.mark.asyncio
async def test_active_session_is_closed_at_session_limit(self):
@ -1145,14 +1157,30 @@ class TestResponsesWSSessionLimit:
responses_websocket_endpoint,
)
ws = self._ws(lambda: json.dumps({"type": "response.create", "model": "gpt-4o-mini", "input": []}))
processor = MagicMock()
processor.common_processing_pre_call_logic = AsyncMock(return_value=({"model": "gpt-4o-mini"}, MagicMock()))
async def hanging_relay():
await asyncio.sleep(60)
close_calls: Final[list[tuple[int, str]]] = []
route_calls: Final[list[tuple[str, object | None]]] = []
async def record_client_close(*, code: int, reason: str) -> None:
close_calls.append((code, reason))
async def record_route_request(
*,
data: dict[str, object],
route_type: str,
**_kwargs: object,
) -> object:
route_calls.append((route_type, data.get("model")))
return hanging_relay()
ws = self._ws(lambda: json.dumps({"type": "response.create", "model": "gpt-4o-mini", "input": []}))
ws.close = AsyncMock(side_effect=record_client_close)
with (
patch(
"litellm.proxy.response_api_endpoints.endpoints._resolve_responses_ws_session_limit_seconds",
@ -1165,12 +1193,13 @@ class TestResponsesWSSessionLimit:
patch(
"litellm.proxy.route_llm_request.route_request",
new_callable=AsyncMock,
return_value=hanging_relay(),
side_effect=record_route_request,
),
):
await responses_websocket_endpoint(websocket=ws, model="gpt-4o-mini", user_api_key_dict=MagicMock())
ws.close.assert_awaited_once_with(code=1000, reason="Session duration limit reached")
assert route_calls == [("_aresponses_websocket", "gpt-4o-mini")]
assert close_calls == [(1000, "Session duration limit reached")]
@pytest.mark.asyncio
async def test_session_timeout_closes_client_before_slow_session_cleanup(self) -> None:
@ -1321,6 +1350,23 @@ class TestResponsesWSSessionLimit:
async def fake_llm_call():
return None
route_calls: Final[list[tuple[str, object | None]]] = []
close_calls: Final[list[tuple[int, str]]] = []
async def record_route_request(
*,
data: dict[str, object],
route_type: str,
**_kwargs: object,
) -> object:
route_calls.append((route_type, data.get("model")))
return fake_llm_call()
async def record_client_close(*, code: int, reason: str) -> None:
close_calls.append((code, reason))
ws.close = AsyncMock(side_effect=record_client_close)
with (
patch(
"litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing",
@ -1329,13 +1375,13 @@ class TestResponsesWSSessionLimit:
patch(
"litellm.proxy.route_llm_request.route_request",
new_callable=AsyncMock,
return_value=fake_llm_call(),
) as mock_route_request,
side_effect=record_route_request,
),
):
await responses_websocket_endpoint(websocket=ws, model="gpt-4o-mini", user_api_key_dict=MagicMock())
mock_route_request.assert_awaited_once()
ws.close.assert_not_awaited()
assert route_calls == [("_aresponses_websocket", "gpt-4o-mini")]
assert close_calls == []
@pytest.mark.parametrize(