mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
57630f44ac
commit
8578a2852c
2 changed files with 146 additions and 11 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue