From 8578a2852c4ee8dac76d78e131a868994dacde9f Mon Sep 17 00:00:00 2001 From: mrinal Date: Sat, 3 Oct 2026 21:46:33 +0000 Subject: [PATCH] fix(proxy): close responses websocket client before upstream cleanup at session limit Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/response_api_endpoints/endpoints.py | 21 ++- .../response_api_endpoints/test_endpoints.py | 136 +++++++++++++++++- 2 files changed, 146 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 0a3cc21d519..320c567be39 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -1644,14 +1644,25 @@ async def responses_websocket_endpoint( await websocket.accept(**accept_kwargs) limit_seconds: Final = _resolve_responses_ws_session_limit_seconds() - session: Final = _responses_websocket_session( - websocket=websocket, - model=model, - user_api_key_dict=user_api_key_dict, + session_task: Final = asyncio.ensure_future( + _responses_websocket_session(websocket=websocket, model=model, user_api_key_dict=user_api_key_dict) ) try: - await asyncio.wait_for(session, timeout=limit_seconds) + await asyncio.wait_for(asyncio.shield(session_task), timeout=limit_seconds) except asyncio.TimeoutError: verbose_proxy_logger.info("Responses WebSocket closed: session duration limit reached") + 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() + with contextlib.suppress(asyncio.CancelledError, Exception): + await session_task diff --git a/tests/unit/proxy/response_api_endpoints/test_endpoints.py b/tests/unit/proxy/response_api_endpoints/test_endpoints.py index 09adda7db7a..0d4db8457a2 100644 --- a/tests/unit/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/unit/proxy/response_api_endpoints/test_endpoints.py @@ -2,6 +2,7 @@ Test for response_api_endpoints/endpoints.py """ +import asyncio import unittest from collections.abc import Mapping from typing import Any, Final, Literal @@ -1121,8 +1122,6 @@ class TestResponsesWSSessionLimit: @pytest.mark.asyncio async def test_idle_connection_is_closed_at_session_limit(self): - import asyncio - from litellm.proxy.response_api_endpoints.endpoints import ( responses_websocket_endpoint, ) @@ -1142,8 +1141,6 @@ class TestResponsesWSSessionLimit: @pytest.mark.asyncio async def test_active_session_is_closed_at_session_limit(self): - import asyncio - from litellm.proxy.response_api_endpoints.endpoints import ( responses_websocket_endpoint, ) @@ -1176,9 +1173,135 @@ class TestResponsesWSSessionLimit: ws.close.assert_awaited_once_with(code=1000, reason="Session duration limit reached") @pytest.mark.asyncio - async def test_delayed_first_frame_is_routed_within_session_limit(self, monkeypatch): - import asyncio + async def test_session_timeout_closes_client_before_slow_session_cleanup(self) -> None: + from litellm.proxy.response_api_endpoints.endpoints import responses_websocket_endpoint + session_started: Final = asyncio.Event() + cleanup_started: Final = asyncio.Event() + cleanup_finished: Final = asyncio.Event() + client_closed: Final = asyncio.Event() + + async def slow_cleanup_session( + *, + websocket: object, + model: str | None, + user_api_key_dict: object, + ) -> None: + session_started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + cleanup_started.set() + await asyncio.sleep(0.5) + cleanup_finished.set() + raise + + async def record_client_close(*, code: int, reason: str) -> None: + assert not cleanup_finished.is_set() + assert (code, reason) == (1000, "Session duration limit reached") + client_closed.set() + + ws: Final = self._ws(lambda: "") + ws.close = AsyncMock(side_effect=record_client_close) + + with ( + patch( + "litellm.proxy.response_api_endpoints.endpoints._resolve_responses_ws_session_limit_seconds", + return_value=0.1, + ), + patch( + "litellm.proxy.response_api_endpoints.endpoints._responses_websocket_session", + new=slow_cleanup_session, + ), + ): + endpoint_task: Final = asyncio.create_task( + responses_websocket_endpoint(websocket=ws, model="gpt-4o-mini", user_api_key_dict=MagicMock()) + ) + await session_started.wait() + await endpoint_task + + assert client_closed.is_set() + assert cleanup_started.is_set() + assert cleanup_finished.is_set() + + @pytest.mark.asyncio + async def test_cancelling_endpoint_cancels_and_reaps_session(self) -> None: + from litellm.proxy.response_api_endpoints.endpoints import responses_websocket_endpoint + + session_started: Final = asyncio.Event() + cancellation_observed: Final = asyncio.Event() + cleanup_finished: Final = asyncio.Event() + + async def waiting_session( + *, + websocket: object, + model: str | None, + user_api_key_dict: object, + ) -> None: + session_started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + cancellation_observed.set() + cleanup_finished.set() + raise + + ws: Final = self._ws(lambda: "") + + with ( + patch( + "litellm.proxy.response_api_endpoints.endpoints._resolve_responses_ws_session_limit_seconds", + return_value=60, + ), + patch( + "litellm.proxy.response_api_endpoints.endpoints._responses_websocket_session", + new=waiting_session, + ), + ): + endpoint_task: Final = asyncio.create_task( + responses_websocket_endpoint(websocket=ws, model="gpt-4o-mini", user_api_key_dict=MagicMock()) + ) + await session_started.wait() + endpoint_task.cancel() + with pytest.raises(asyncio.CancelledError): + await endpoint_task + + assert cancellation_observed.is_set() + assert cleanup_finished.is_set() + + @pytest.mark.asyncio + async def test_session_exception_before_limit_is_propagated(self) -> None: + from litellm.proxy.response_api_endpoints.endpoints import responses_websocket_endpoint + + failure: Final = RuntimeError("session failed") + + async def failed_session( + *, + websocket: object, + model: str | None, + user_api_key_dict: object, + ) -> None: + raise failure + + ws: Final = self._ws(lambda: "") + + with ( + patch( + "litellm.proxy.response_api_endpoints.endpoints._resolve_responses_ws_session_limit_seconds", + return_value=60, + ), + patch( + "litellm.proxy.response_api_endpoints.endpoints._responses_websocket_session", + new=failed_session, + ), + pytest.raises(RuntimeError, match="session failed") as raised, + ): + await responses_websocket_endpoint(websocket=ws, model="gpt-4o-mini", user_api_key_dict=MagicMock()) + + assert raised.value is failure + + @pytest.mark.asyncio + async def test_delayed_first_frame_is_routed_within_session_limit(self, monkeypatch): from litellm.proxy.proxy_server import general_settings from litellm.proxy.response_api_endpoints.endpoints import ( responses_websocket_endpoint, @@ -1214,6 +1337,7 @@ class TestResponsesWSSessionLimit: mock_route_request.assert_awaited_once() ws.close.assert_not_awaited() + @pytest.mark.parametrize( "configured,expected", [