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>
This commit is contained in:
mrinal 2026-10-03 21:46:33 +00:00
parent 57630f44ac
commit 8578a2852c
2 changed files with 146 additions and 11 deletions

View file

@ -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

View file

@ -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",
[