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