test(proxy): move responses websocket session limit timing coverage to integration

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mrinal 2026-10-03 23:02:34 +00:00
parent b5f8f2b8b8
commit c90e63c456
4 changed files with 91 additions and 293 deletions

View file

@ -2719,6 +2719,9 @@ class ScheduledJobStaggerSettings(LiteLLMPydanticObjectBase):
)
DEFAULT_RESPONSES_WEBSOCKET_SESSION_LIMIT_SECONDS: Final[float] = 3600.0
class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"""
Documents all the fields supported by `general_settings` in config.yaml
@ -3007,7 +3010,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
description="Default upstream request timeout in seconds for native and custom pass-through endpoints that use pass_through_request. Defaults to 600 when unset.",
)
responses_websocket_session_limit_seconds: float = Field(
default=3600.0,
default=DEFAULT_RESPONSES_WEBSOCKET_SESSION_LIMIT_SECONDS,
ge=60,
le=7200,
description="Maximum lifetime in seconds of a Responses API WebSocket session, measured from connection accept and covering the idle wait for the first response.create frame. Defaults to 3600, matching OpenAI's documented 60-minute WebSocket connection limit. Must be between 60 and 7200 seconds.",

View file

@ -12,7 +12,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi.responses import JSONResponse
from openai.types.responses import ResponseItemList
from openai.types.responses.response_create_params import ResponseInputParam
from pydantic import BaseModel, ConfigDict, ValidationError
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from starlette.websockets import WebSocket, WebSocketDisconnect
from typing_extensions import ReadOnly, TypedDict
@ -47,6 +47,7 @@ if TYPE_CHECKING:
from litellm.router import Router
router: Final = APIRouter()
_RESPONSES_WS_CONFIG_VALUE_ADAPTER: Final[TypeAdapter[object | None]] = TypeAdapter(object | None)
_ResponseDocSchemas: TypeAlias = dict[int | str, dict[str, object]] # fastapi's responses kwarg
@ -1317,15 +1318,15 @@ def _resolve_responses_ws_session_limit_seconds() -> float:
from litellm.proxy.proxy_server import general_settings
field: Final = "responses_websocket_session_limit_seconds"
raw: Final = general_settings.get(field)
raw: Final = _RESPONSES_WS_CONFIG_VALUE_ADAPTER.validate_python(general_settings.get(field))
try:
return ConfigGeneralSettings.model_validate(
{} if raw is None else {field: raw}
).responses_websocket_session_limit_seconds
except ValidationError as e:
default: Final = ConfigGeneralSettings.model_fields[field].default
default: Final = DEFAULT_RESPONSES_WEBSOCKET_SESSION_LIMIT_SECONDS
verbose_proxy_logger.warning("invalid general_settings.%s=%r (%s); using default %ss", field, raw, e, default)
return float(default)
return default
async def _read_ws_model_from_first_frame(

View file

@ -7,12 +7,14 @@ import ssl
import threading
import time
import uuid
from collections.abc import Generator, Iterator
from collections.abc import Generator, Iterator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from queue import SimpleQueue
from queue import Empty, SimpleQueue
from types import MappingProxyType
from typing import Final, TypeVar
from urllib.parse import parse_qs, urlsplit
import httpx
import pytest
@ -34,6 +36,9 @@ pytestmark: Final = pytest.mark.timeout(360)
PROVIDER_MODEL: Final = "ws-peer-model"
STALL_PROVIDER_MODEL: Final = "ws-stall-peer-model"
DEAF_PROVIDER_MODEL: Final = "ws-deaf-peer-model"
DEAF_RESPONSE_DELAY_SECONDS: Final = 10
DEAF_READ_PAUSE_SECONDS: Final = 20
PEER_TEXT: Final = "responses websocket peer"
TERMINAL: Final = frozenset({"response.completed", "response.failed", "error"})
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
@ -51,7 +56,7 @@ class PeerConnection:
class ResponsesPeer:
url: str
connections: SimpleQueue[PeerConnection]
closed: SimpleQueue[PeerConnection]
connections_by_model: Mapping[str, SimpleQueue[PeerConnection]]
@dataclass(frozen=True, slots=True)
@ -130,6 +135,7 @@ class CapResults:
idle: CloseResult
active: ActiveResult
mid: MidResult
deaf: MidResult
provider: tuple[PeerSnapshot, ...]
@ -209,17 +215,24 @@ def _events(response_id: str, model: str, text: str, *, stall: bool = False) ->
)
def _path_model(path: str) -> str | None:
values: Final = parse_qs(urlsplit(path).query).get("model")
return values[0] if values else None
async def _peer_handler(connection: ServerConnection, peer: ResponsesPeer) -> None:
path: Final = connection.request.path if connection.request is not None else ""
record: Final = PeerConnection(path, SimpleQueue(), SimpleQueue())
peer.connections.put(record)
model: Final = _path_model(path)
if model is not None and model in peer.connections_by_model:
peer.connections_by_model[model].put(record)
turns: Final = itertools.count(1)
try:
async for raw in connection:
await _peer_frame(raw, connection, record, turns)
finally:
record.closed.put(time.monotonic())
peer.closed.put(record)
async def _peer_frame(
@ -233,9 +246,19 @@ async def _peer_frame(
if frame.get("type") != "response.create":
return
model: Final = _string(frame.get("model", ""))
stall: Final = model == STALL_PROVIDER_MODEL
stall: Final = model in (STALL_PROVIDER_MODEL, DEAF_PROVIDER_MODEL)
if model == DEAF_PROVIDER_MODEL:
await asyncio.sleep(DEAF_RESPONSE_DELAY_SECONDS)
for event in _events(f"resp_peer_{next(turns)}", model, PEER_TEXT, stall=stall):
await connection.send(json.dumps(event))
if model == DEAF_PROVIDER_MODEL:
transport: Final = connection.transport
assert transport is not None
transport.pause_reading()
try:
await asyncio.sleep(DEAF_READ_PAUSE_SECONDS)
finally:
transport.resume_reading()
async def _serve_peer(
@ -256,7 +279,10 @@ def responses_peer(cert: tuple[Path, Path]) -> Generator[ResponsesPeer, None, No
loop: Final = asyncio.new_event_loop()
stop: Final = asyncio.Event()
ports: Final = SimpleQueue[int]()
peer: Final = ResponsesPeer("", SimpleQueue(), SimpleQueue())
connections_by_model: Final[Mapping[str, SimpleQueue[PeerConnection]]] = MappingProxyType(
{model: SimpleQueue[PeerConnection]() for model in (PROVIDER_MODEL, STALL_PROVIDER_MODEL, DEAF_PROVIDER_MODEL)}
)
peer: Final = ResponsesPeer("", SimpleQueue(), connections_by_model)
thread: Final = threading.Thread(
target=loop.run_until_complete,
args=(_serve_peer(server_context(*cert), peer, ports, stop),),
@ -265,7 +291,7 @@ def responses_peer(cert: tuple[Path, Path]) -> Generator[ResponsesPeer, None, No
thread.start()
try:
port: Final = ports.get(timeout=10)
yield ResponsesPeer(f"https://127.0.0.1:{port}/v1", peer.connections, peer.closed)
yield ResponsesPeer(f"https://127.0.0.1:{port}/v1", peer.connections, peer.connections_by_model)
finally:
loop.call_soon_threadsafe(stop.set)
thread.join(timeout=10)
@ -523,7 +549,7 @@ async def _active_session(candidate: Gateway, key: str, model: str, peer: Respon
turn, turn_error = await _turn_result(connection, _create(model, f"active-{uuid.uuid4().hex}"))
remaining: Final = max(0, 75 - (time.monotonic() - started))
close: Final = CloseResult(await _wait_for_close(connection, remaining), time.monotonic() - started)
provider_closed: Final = _provider_closed_within(peer, 5)
provider_closed: Final = _provider_closed_within(peer, PROVIDER_MODEL, 5)
return ActiveResult(turn, turn_error, close, provider_closed)
except (ConnectionClosed, asyncio.TimeoutError) as error:
code, reason = _auth_close(error) if isinstance(error, ConnectionClosed) else (None, None)
@ -535,9 +561,18 @@ async def _active_session(candidate: Gateway, key: str, model: str, peer: Respon
)
def _provider_closed_within(peer: ResponsesPeer, seconds: float) -> bool:
def _provider_closed_within(peer: ResponsesPeer, model: str, seconds: float) -> bool:
deadline: Final = time.monotonic() + seconds
try:
eventually(lambda: peer.closed.qsize(), lambda count: count >= 1, seconds=seconds)
record: Final = peer.connections_by_model[model].get(timeout=seconds)
except Empty:
return False
try:
eventually(
lambda: record.closed.qsize(),
lambda count: count >= 1,
seconds=max(0, deadline - time.monotonic()),
)
except AssertionError:
return False
return True
@ -602,7 +637,20 @@ async def _mid_response(
peer: ResponsesPeer,
) -> MidResult:
created, turn_error, close = await _mid_session(candidate, key, stall_model)
provider_closed: Final = _provider_closed_within(peer, 5)
provider_closed: Final = _provider_closed_within(peer, STALL_PROVIDER_MODEL, 5)
fresh_completed, fresh_error = await _fresh_session(candidate, key, normal_model)
return MidResult(created, turn_error, close, provider_closed, fresh_completed, fresh_error)
async def _deaf_response(
candidate: Gateway,
key: str,
normal_model: str,
deaf_model: str,
peer: ResponsesPeer,
) -> MidResult:
created, turn_error, close = await _mid_session(candidate, key, deaf_model)
provider_closed: Final = _provider_closed_within(peer, DEAF_PROVIDER_MODEL, 30)
fresh_completed, fresh_error = await _fresh_session(candidate, key, normal_model)
return MidResult(created, turn_error, close, provider_closed, fresh_completed, fresh_error)
@ -612,14 +660,16 @@ async def _cap_workload(
key: str,
normal_model: str,
stall_model: str,
deaf_model: str,
peer: ResponsesPeer,
) -> tuple[CloseResult, ActiveResult, MidResult]:
idle, active, mid = await asyncio.gather(
) -> tuple[CloseResult, ActiveResult, MidResult, MidResult]:
idle, active, mid, deaf = await asyncio.gather(
_close_at_limit(candidate, key, normal_model),
_active_session(candidate, key, normal_model, peer),
_mid_response(candidate, key, normal_model, stall_model, peer),
_deaf_response(candidate, key, normal_model, deaf_model, peer),
)
return idle, active, mid
return idle, active, mid, deaf
def _session_config(path: Path, seconds: int) -> Path:
@ -788,10 +838,11 @@ def cap_results(
):
normal: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url)
stall: Final = scenario.model(model=f"openai/{STALL_PROVIDER_MODEL}", api_base=peer.url)
key: Final = scenario.key(models=[normal, stall])
idle, active, mid = asyncio.run(_cap_workload(candidate, key, normal, stall, peer))
deaf: Final = scenario.model(model=f"openai/{DEAF_PROVIDER_MODEL}", api_base=peer.url)
key: Final = scenario.key(models=[normal, stall, deaf])
idle, active, mid, deaf_result = asyncio.run(_cap_workload(candidate, key, normal, stall, deaf, peer))
provider: Final = _available_snapshots(peer)
yield CapResults(idle, active, mid, provider)
yield CapResults(idle, active, mid, deaf_result, provider)
@pytest.fixture(scope="module")
@ -947,6 +998,18 @@ def test_session_cap_closes_mid_response_and_allows_new_session(cap_results: Cap
assert cap_results.mid.fresh_completed, cap_results.mid
def test_session_cap_closes_client_promptly_when_provider_ignores_close(cap_results: CapResults) -> None:
assert cap_results.deaf.created, cap_results.deaf
assert cap_results.deaf.turn_error is None, cap_results.deaf
assert cap_results.deaf.close.outcome.closed, cap_results.deaf
assert cap_results.deaf.close.outcome.close_code == 1000, cap_results.deaf
assert cap_results.deaf.close.outcome.close_reason == "Session duration limit reached", cap_results.deaf
assert 59 <= cap_results.deaf.close.elapsed <= 63, cap_results.deaf
assert cap_results.deaf.provider_closed, cap_results.deaf
assert cap_results.deaf.fresh_error is None, cap_results.deaf
assert cap_results.deaf.fresh_completed, cap_results.deaf
def test_invalid_session_cap_falls_back_to_default(invalid_results: InvalidResult) -> None:
assert invalid_results.session.idle.timed_out, invalid_results.session
assert invalid_results.session.error is None, invalid_results.session

View file

@ -2,7 +2,6 @@
Test for response_api_endpoints/endpoints.py
"""
import asyncio
import unittest
from collections.abc import Mapping
from typing import Any, Final, Literal
@ -963,6 +962,7 @@ class TestResponsesWSFirstFrameModelAuth:
request = Request({"type": "http", "method": "POST", "path": "/v1/responses", "headers": []})
user_api_key_dict = MagicMock()
llm_router = MagicMock()
empty_settings: Final[dict[str, object]] = {}
with (
patch(
@ -979,7 +979,7 @@ class TestResponsesWSFirstFrameModelAuth:
),
patch("litellm.proxy.proxy_server.master_key", "sk-test"),
patch("litellm.proxy.proxy_server.user_custom_auth", None),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.general_settings", empty_settings),
):
await _enforce_responses_ws_first_frame_model_auth(
request=request,
@ -1114,275 +1114,6 @@ class TestReadWSModelFromFirstFrameErrors:
ws.close.assert_not_awaited()
class TestResponsesWSSessionLimit:
def _ws(self, receive_text):
ws = MagicMock()
ws.headers = {}
ws.scope = {"headers": []}
ws.url = "ws://testserver/v1/responses"
ws.accept = AsyncMock()
ws.receive_text = AsyncMock(side_effect=receive_text)
ws.send_text = AsyncMock()
ws.close = AsyncMock()
return ws
@pytest.mark.asyncio
async def test_idle_connection_is_closed_at_session_limit(self):
from litellm.proxy.response_api_endpoints.endpoints import (
responses_websocket_endpoint,
)
async def silent_socket():
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",
return_value=0.05,
):
await responses_websocket_endpoint(websocket=ws, model="gpt-4o-mini", user_api_key_dict=MagicMock())
assert close_calls == [(1000, "Session duration limit reached")]
@pytest.mark.asyncio
async def test_active_session_is_closed_at_session_limit(self):
from litellm.proxy.response_api_endpoints.endpoints import (
responses_websocket_endpoint,
)
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",
return_value=0.05,
),
patch(
"litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing",
return_value=processor,
),
patch(
"litellm.proxy.route_llm_request.route_request",
new_callable=AsyncMock,
side_effect=record_route_request,
),
):
await responses_websocket_endpoint(websocket=ws, model="gpt-4o-mini", user_api_key_dict=MagicMock())
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:
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,
)
monkeypatch.setitem(general_settings, "responses_websocket_session_limit_seconds", 60)
async def delayed_frame():
await asyncio.sleep(0.2)
return json.dumps({"type": "response.create", "model": "gpt-4o-mini", "input": []})
ws = self._ws(delayed_frame)
processor = MagicMock()
processor.common_processing_pre_call_logic = AsyncMock(return_value=({"model": "gpt-4o-mini"}, MagicMock()))
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",
return_value=processor,
),
patch(
"litellm.proxy.route_llm_request.route_request",
new_callable=AsyncMock,
side_effect=record_route_request,
),
):
await responses_websocket_endpoint(websocket=ws, model="gpt-4o-mini", user_api_key_dict=MagicMock())
assert route_calls == [("_aresponses_websocket", "gpt-4o-mini")]
assert close_calls == []
@pytest.mark.parametrize(
"configured,expected",