From b5f8f2b8b8f50b852f4a7475db3f8cd5833c8add Mon Sep 17 00:00:00 2001 From: mrinal Date: Sat, 3 Oct 2026 22:03:48 +0000 Subject: [PATCH] test(proxy): assert recorded websocket outcomes and simplify session reaping Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/response_api_endpoints/endpoints.py | 7 -- .../response_api_endpoints/test_endpoints.py | 68 ++++++++++++++++--- 2 files changed, 57 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 320c567be39..188055b74bd 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -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() diff --git a/tests/unit/proxy/response_api_endpoints/test_endpoints.py b/tests/unit/proxy/response_api_endpoints/test_endpoints.py index 0d4db8457a2..a06499d41d4 100644 --- a/tests/unit/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/unit/proxy/response_api_endpoints/test_endpoints.py @@ -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(