From 18193632db9ac3e2a571a3d270b7dbe15d6a3b81 Mon Sep 17 00:00:00 2001 From: Mrinal Chanshetty Date: Sat, 3 Oct 2026 11:51:47 -0700 Subject: [PATCH 1/7] add config-driven bound to /responses websocket connection --- litellm/proxy/_types.py | 6 + .../proxy/response_api_endpoints/endpoints.py | 103 +++++++++---- .../response_api_endpoints/test_endpoints.py | 138 +++++++++++++++++- 3 files changed, 213 insertions(+), 34 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0abec51cc49..4cf3ff89bdc 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3006,6 +3006,12 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): default=None, 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, + 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.", + ) pass_through_endpoints: list[PassThroughGenericEndpoint] | None = Field( default=None, description="Set-up pass-through endpoints for provider-specific endpoints. Docs - https://docs.litellm.ai/docs/proxy/pass_through", diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 882aa240fad..86581d90039 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -1313,6 +1313,23 @@ async def cancel_response( ) +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) + 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 + verbose_proxy_logger.warning( + "invalid general_settings.%s=%r (%s); using default %ss", field, raw, e, default + ) + return float(default) + + async def _read_ws_model_from_first_frame( websocket: WebSocket, query_model: str | None = None, @@ -1320,12 +1337,10 @@ async def _read_ws_model_from_first_frame( """Read the first WS frame and return (model, raw_message), or None on error. Sends an appropriate error frame and closes the socket before returning None. + The session-duration deadline is enforced by the caller, not here. """ try: - first_message: Final = await asyncio.wait_for(websocket.receive_text(), timeout=30) - except asyncio.TimeoutError: - await websocket.close(code=1008, reason="Timed out waiting for first message") - return None + first_message: Final = await websocket.receive_text() except WebSocketDisconnect: return None except Exception: @@ -1471,25 +1486,11 @@ async def _enforce_responses_ws_first_frame_model_auth( ) -@router.websocket("/v1/responses") -@router.websocket("/responses") -async def responses_websocket_endpoint( +async def _responses_websocket_session( websocket: WebSocket, - model: str | None = fastapi.Query(None, description="The model to use for the responses WebSocket session."), - user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), -): - """ - Responses API WebSocket mode endpoint. - - Keeps a persistent WebSocket connection for response.create events, - enabling lower-latency agentic workflows with many tool-call round trips. - - Follows the OpenAI split: the bearer token is validated at connection time - (before accept); the model is resolved either from the ?model= query param - or from the first response.create frame, whichever is present. - - See: https://developers.openai.com/api/docs/guides/websocket-mode/ - """ + model: str | None, + user_api_key_dict: UserAPIKeyAuth, +) -> None: from litellm.proxy.proxy_server import ( general_settings, llm_router, @@ -1504,16 +1505,6 @@ async def responses_websocket_endpoint( ) from litellm.proxy.route_llm_request import route_request - # Accept the WebSocket handshake. Key was already validated by the Depends - # above; we can safely accept regardless of whether ?model= was supplied. - requested_protocols: Final = [ - p.strip() for p in (websocket.headers.get("sec-websocket-protocol") or "").split(",") if p.strip() - ] - accept_kwargs: Final[dict] = {} - if requested_protocols: - accept_kwargs["subprotocol"] = requested_protocols[0] - await websocket.accept(**accept_kwargs) - result: Final = await _read_ws_model_from_first_frame(websocket, query_model=model) if result is None: return @@ -1618,3 +1609,51 @@ async def responses_websocket_endpoint( request_data=routed_data, ) await websocket.close(code=1011, reason="Internal server error") + + +@router.websocket("/v1/responses") +@router.websocket("/responses") +async def responses_websocket_endpoint( + websocket: WebSocket, + model: str | None = fastapi.Query(None, description="The model to use for the responses WebSocket session."), + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), +): + """ + Responses API WebSocket mode endpoint. + + Keeps a persistent WebSocket connection for response.create events, + enabling lower-latency agentic workflows with many tool-call round trips. + + Follows the OpenAI split: the bearer token is validated at connection time + (before accept); the model is resolved either from the ?model= query param + or from the first response.create frame, whichever is present. + + The session is bounded by a lifetime measured from accept, configured via + general_settings.responses_websocket_session_limit_seconds (60-7200, + default 3600). There is no separate first-frame deadline, so + pre-established connections may sit idle until their first response.create. + + See: https://developers.openai.com/api/docs/guides/websocket-mode/ + """ + # Accept the WebSocket handshake. Key was already validated by the Depends + # above; we can safely accept regardless of whether ?model= was supplied. + requested_protocols: Final = [ + p.strip() for p in (websocket.headers.get("sec-websocket-protocol") or "").split(",") if p.strip() + ] + accept_kwargs: Final[dict] = {} + if requested_protocols: + accept_kwargs["subprotocol"] = requested_protocols[0] + 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, + ) + try: + await asyncio.wait_for(session, timeout=limit_seconds) + except asyncio.TimeoutError: + verbose_proxy_logger.info("Responses WebSocket closed: session duration limit reached") + with contextlib.suppress(Exception): + await websocket.close(code=1000, reason="Session duration limit reached") diff --git a/tests/unit/proxy/response_api_endpoints/test_endpoints.py b/tests/unit/proxy/response_api_endpoints/test_endpoints.py index 656dc33e88c..09adda7db7a 100644 --- a/tests/unit/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/unit/proxy/response_api_endpoints/test_endpoints.py @@ -999,7 +999,7 @@ class TestResponsesWSFirstFrameModelAuth: class TestReadWSModelFromFirstFrameErrors: @pytest.mark.asyncio - async def test_timeout_closes_without_error_frame(self): + async def test_transport_error_first_frame_closes_with_internal_error(self): import asyncio from litellm.proxy.response_api_endpoints.endpoints import ( @@ -1015,7 +1015,7 @@ class TestReadWSModelFromFirstFrameErrors: assert result is None ws.send_text.assert_not_awaited() - ws.close.assert_awaited_once_with(code=1008, reason="Timed out waiting for first message") + ws.close.assert_awaited_once_with(code=1011, reason="Internal server error") @pytest.mark.asyncio async def test_invalid_json_sends_error_and_closes(self): @@ -1107,6 +1107,140 @@ 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): + import asyncio + + from litellm.proxy.response_api_endpoints.endpoints import ( + responses_websocket_endpoint, + ) + + async def silent_socket(): + await asyncio.sleep(60) + + ws = self._ws(silent_socket) + + 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()) + + ws.close.assert_awaited_once_with(code=1000, reason="Session duration limit reached") + + @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, + ) + + 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) + + 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, + return_value=hanging_relay(), + ), + ): + 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") + + @pytest.mark.asyncio + async def test_delayed_first_frame_is_routed_within_session_limit(self, monkeypatch): + import asyncio + + 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 + + with ( + patch( + "litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing", + return_value=processor, + ), + patch( + "litellm.proxy.route_llm_request.route_request", + new_callable=AsyncMock, + return_value=fake_llm_call(), + ) as mock_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() + +@pytest.mark.parametrize( + "configured,expected", + [ + (None, 3600.0), + (60, 60.0), + (1200, 1200.0), + (7200, 7200.0), + (59, 3600.0), + (0, 3600.0), + (9000, 3600.0), + ("not-a-number", 3600.0), + ], +) +def test_responses_ws_session_limit_resolution(monkeypatch, configured, expected): + from litellm.proxy.proxy_server import general_settings + from litellm.proxy.response_api_endpoints.endpoints import ( + _resolve_responses_ws_session_limit_seconds, + ) + + if configured is None: + monkeypatch.delitem(general_settings, "responses_websocket_session_limit_seconds", raising=False) + else: + monkeypatch.setitem(general_settings, "responses_websocket_session_limit_seconds", configured) + + assert _resolve_responses_ws_session_limit_seconds() == expected + + class TestManagedResponsesSameProvider: def _handler(self, model, custom_llm_provider=None): from litellm.responses.streaming_iterator import ( From 2320d13ee995746fbb890a54e244d8d8a4a5f8d6 Mon Sep 17 00:00:00 2001 From: mrinal Date: Sat, 3 Oct 2026 20:59:07 +0000 Subject: [PATCH 2/7] test(proxy): cover responses websocket session limit and idle first frame Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_responses_websocket_session_limit.py | 982 ++++++++++++++++++ 1 file changed, 982 insertions(+) create mode 100644 tests/integration/providers/test_responses_websocket_session_limit.py diff --git a/tests/integration/providers/test_responses_websocket_session_limit.py b/tests/integration/providers/test_responses_websocket_session_limit.py new file mode 100644 index 00000000000..fc71aea8f2c --- /dev/null +++ b/tests/integration/providers/test_responses_websocket_session_limit.py @@ -0,0 +1,982 @@ +from __future__ import annotations + +import asyncio +import itertools +import json +import ssl +import threading +import time +import uuid +from collections.abc import Generator, Iterator +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from typing import Final, TypeVar + +import httpx +import pytest +import websockets +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.process import owned_proxy, owned_proxy_process +from integration._support.tls import server_context, write_self_signed_cert +from integration._support.upstream import ( + JsonResponse, + RoutedResponse, + delete_scenario, + register_scenario, +) +from pydantic import JsonValue, TypeAdapter, ValidationError +from websockets.asyncio.server import ServerConnection, serve +from websockets.exceptions import ConnectionClosed + +pytestmark: Final = pytest.mark.timeout(360) + +PROVIDER_MODEL: Final = "ws-peer-model" +STALL_PROVIDER_MODEL: Final = "ws-stall-peer-model" +PEER_TEXT: Final = "responses websocket peer" +TERMINAL: Final = frozenset({"response.completed", "response.failed", "error"}) +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +T = TypeVar("T") + + +@dataclass(frozen=True, slots=True) +class PeerConnection: + path: str + frames: SimpleQueue[dict[str, JsonValue]] + closed: SimpleQueue[float] + + +@dataclass(frozen=True, slots=True) +class ResponsesPeer: + url: str + connections: SimpleQueue[PeerConnection] + closed: SimpleQueue[PeerConnection] + + +@dataclass(frozen=True, slots=True) +class PeerSnapshot: + path: str + frames: tuple[dict[str, JsonValue], ...] + closed: bool + + +@dataclass(frozen=True, slots=True) +class ReceiveResult: + timed_out: bool + closed: bool + frame: dict[str, JsonValue] | None + close_code: int | None + close_reason: str | None + + +@dataclass(frozen=True, slots=True) +class SessionResult: + text: str + idle: ReceiveResult + events: tuple[dict[str, JsonValue], ...] + error: str | None + + +@dataclass(frozen=True, slots=True) +class DefaultResults: + pool: tuple[SessionResult, ...] + query: SessionResult + provider: tuple[PeerSnapshot, ...] + + +@dataclass(frozen=True, slots=True) +class AuthResult: + idle: ReceiveResult + frame: dict[str, JsonValue] | None + close_code: int | None + close_reason: str | None + error: str | None + + +@dataclass(frozen=True, slots=True) +class AuthResults: + immediate: AuthResult + delayed: AuthResult + provider_connections: int + + +@dataclass(frozen=True, slots=True) +class CloseResult: + outcome: ReceiveResult + elapsed: float + + +@dataclass(frozen=True, slots=True) +class ActiveResult: + turn: tuple[dict[str, JsonValue], ...] + turn_error: str | None + close: CloseResult + provider_closed: bool + + +@dataclass(frozen=True, slots=True) +class MidResult: + created: bool + turn_error: str | None + close: CloseResult + provider_closed: bool + fresh_completed: bool + fresh_error: str | None + + +@dataclass(frozen=True, slots=True) +class CapResults: + idle: CloseResult + active: ActiveResult + mid: MidResult + provider: tuple[PeerSnapshot, ...] + + +@dataclass(frozen=True, slots=True) +class BurstResult: + path: str + status_code: int | None + body: dict[str, JsonValue] | None + error: str | None + + +@dataclass(frozen=True, slots=True) +class InvalidResult: + session: SessionResult + warning_found: bool + + +@dataclass(frozen=True, slots=True) +class ChaosResults: + burst: tuple[BurstResult, ...] + health_status: int + websocket_completed: bool + + +def _object(value: JsonValue) -> dict[str, JsonValue]: + assert isinstance(value, dict), value + return value + + +def _list(value: JsonValue) -> list[JsonValue]: + assert isinstance(value, list), value + return value + + +def _string(value: JsonValue) -> str: + assert isinstance(value, str), value + return value + + +def _drain(queue: SimpleQueue[T]) -> tuple[T, ...]: + return tuple(queue.get_nowait() for _ in range(queue.qsize())) + + +def _events(response_id: str, model: str, text: str, *, stall: bool = False) -> tuple[dict[str, JsonValue], ...]: + message: Final[dict[str, JsonValue]] = { + "type": "message", + "id": f"msg_{response_id}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": text, "annotations": []}], + } + response: Final[dict[str, JsonValue]] = { + "id": response_id, + "object": "response", + "created_at": 1700000000, + "model": model, + } + created: Final[dict[str, JsonValue]] = { + "type": "response.created", + "response": {**response, "status": "in_progress", "output": []}, + } + if stall: + return (created,) + return ( + created, + { + "type": "response.output_text.delta", + "item_id": f"msg_{response_id}", + "output_index": 0, + "content_index": 0, + "delta": text, + }, + { + "type": "response.completed", + "response": {**response, "status": "completed", "output": [message]}, + }, + ) + + +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) + 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( + raw: str | bytes, + connection: ServerConnection, + record: PeerConnection, + turns: itertools.count[int], +) -> None: + frame: Final = JSON_OBJECT.validate_json(raw) + record.frames.put(frame) + if frame.get("type") != "response.create": + return + model: Final = _string(frame.get("model", "")) + stall: Final = model == STALL_PROVIDER_MODEL + for event in _events(f"resp_peer_{next(turns)}", model, PEER_TEXT, stall=stall): + await connection.send(json.dumps(event)) + + +async def _serve_peer( + tls: ssl.SSLContext, + peer: ResponsesPeer, + ports: SimpleQueue[int], + stop: asyncio.Event, +) -> None: + async with serve(lambda connection: _peer_handler(connection, peer), "127.0.0.1", 0, ssl=tls) as server: + address: object = next(iter(server.sockets)).getsockname() # pyright: ignore[reportAny] # socket.getsockname is typed Any + port: Final = TypeAdapter(tuple[str, int]).validate_python(address)[1] + ports.put(port) + await stop.wait() + + +@contextmanager +def responses_peer(cert: tuple[Path, Path]) -> Generator[ResponsesPeer, None, None]: + loop: Final = asyncio.new_event_loop() + stop: Final = asyncio.Event() + ports: Final = SimpleQueue[int]() + peer: Final = ResponsesPeer("", SimpleQueue(), SimpleQueue()) + thread: Final = threading.Thread( + target=loop.run_until_complete, + args=(_serve_peer(server_context(*cert), peer, ports, stop),), + daemon=True, + ) + thread.start() + try: + port: Final = ports.get(timeout=10) + yield ResponsesPeer(f"https://127.0.0.1:{port}/v1", peer.connections, peer.closed) + finally: + loop.call_soon_threadsafe(stop.set) + thread.join(timeout=10) + loop.close() + + +def _proxy_ws_url(candidate: Gateway, path: str) -> str: + base: Final = str(candidate.client.base_url).rstrip("/").replace("http://", "ws://", 1) + return f"{base}{path}" + + +def _create(model: str, text: str, *, include_model: bool = True) -> str: + body: Final[dict[str, JsonValue]] = { + "type": "response.create", + "input": [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": text}]}], + **({"model": model} if include_model else {}), + } + return json.dumps(body) + + +async def _receive(connection: websockets.ClientConnection, timeout: float) -> ReceiveResult: + try: + raw: Final = await asyncio.wait_for(connection.recv(), timeout=timeout) + except asyncio.TimeoutError: + return ReceiveResult(True, False, None, None, None) + except ConnectionClosed as error: + received: Final = error.rcvd + return ReceiveResult( + False, + True, + None, + received.code if received is not None else None, + received.reason if received is not None else None, + ) + return ReceiveResult(False, False, JSON_OBJECT.validate_json(raw), None, None) + + +async def _wait_for_close(connection: websockets.ClientConnection, timeout: float) -> ReceiveResult: + started: Final = time.monotonic() + result: Final = await _receive(connection, timeout) + if result.closed or result.timed_out: + return result + remaining: Final = timeout - (time.monotonic() - started) + if remaining <= 0: + return ReceiveResult(True, False, None, None, None) + return await _wait_for_close(connection, remaining) + + +async def _until_terminal( + connection: websockets.ClientConnection, + received: tuple[dict[str, JsonValue], ...] = (), +) -> tuple[dict[str, JsonValue], ...]: + event: Final = JSON_OBJECT.validate_json(await asyncio.wait_for(connection.recv(), timeout=20)) + collected: Final = (*received, event) + if event.get("type") in TERMINAL or len(collected) >= 50: + return collected + return await _until_terminal(connection, collected) + + +async def _turn(connection: websockets.ClientConnection, frame: str) -> tuple[dict[str, JsonValue], ...]: + await connection.send(frame) + return await _until_terminal(connection) + + +async def _idle_turn( + proxy: str, + key: str, + model: str, + path: str, + text: str, + *, + query_model: bool = False, +) -> SessionResult: + query: Final = f"?model={model}" if query_model else "" + try: + async with websockets.connect( + f"{proxy}{path}{query}", + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + idle: Final = await _receive(connection, 35) + try: + events: Final = await _turn(connection, _create(model, text, include_model=not query_model)) + except (ConnectionClosed, asyncio.TimeoutError) as error: + return SessionResult(text, idle, (), f"{type(error).__name__}: {error}") + return SessionResult(text, idle, events, None) + except (ConnectionClosed, asyncio.TimeoutError) as error: + return SessionResult(text, ReceiveResult(False, True, None, None, None), (), f"{type(error).__name__}: {error}") + + +async def _pool_workload(candidate: Gateway, key: str, model: str) -> tuple[SessionResult, ...]: + proxy: Final = _proxy_ws_url(candidate, "") + paths: Final = ("/v1/responses",) * 4 + ("/responses",) * 4 + return tuple( + await asyncio.gather(*tuple(_idle_turn(proxy, key, model, path, f"pool-{uuid.uuid4().hex}") for path in paths)) + ) + + +async def _default_workload( + candidate: Gateway, key: str, model: str +) -> tuple[tuple[SessionResult, ...], SessionResult]: + proxy: Final = _proxy_ws_url(candidate, "") + pool, query = await asyncio.gather( + _pool_workload(candidate, key, model), + _idle_turn(proxy, key, model, "/v1/responses", f"query-{uuid.uuid4().hex}", query_model=True), + ) + return pool, query + + +def _snapshot(record: PeerConnection) -> PeerSnapshot: + return PeerSnapshot(record.path, _drain(record.frames), record.closed.qsize() > 0) + + +def _snapshots(peer: ResponsesPeer, count: int) -> tuple[PeerSnapshot, ...]: + records: Final = tuple(peer.connections.get(timeout=15) for _ in range(count)) + eventually( + lambda: tuple(record.closed.qsize() for record in records), + lambda counts: all(count > 0 for count in counts), + seconds=10, + ) + return tuple(_snapshot(record) for record in records) + + +def _available_snapshots(peer: ResponsesPeer) -> tuple[PeerSnapshot, ...]: + eventually(lambda: peer.connections.qsize(), lambda count: count >= 1, seconds=10) + records: Final = _drain(peer.connections) + eventually( + lambda: tuple(record.closed.qsize() for record in records), + lambda counts: all(count > 0 for count in counts), + seconds=10, + ) + return tuple(_snapshot(record) for record in records) + + +def _frame_text(frame: dict[str, JsonValue]) -> str: + input_value: Final = _list(frame["input"]) + message: Final = _object(input_value[0]) + content: Final = _list(message["content"]) + item: Final = _object(content[0]) + return _string(item["text"]) + + +def _completed_text(events: tuple[dict[str, JsonValue], ...]) -> str: + assert events, events + completed: Final = events[-1] + assert completed.get("type") == "response.completed", completed + response: Final = _object(completed["response"]) + output: Final = _list(response["output"]) + message: Final = _object(output[0]) + content: Final = _list(message["content"]) + item: Final = _object(content[0]) + return _string(item["text"]) + + +def _auth_close(error: ConnectionClosed) -> tuple[int | None, str | None]: + received: Final = error.rcvd + return ( + received.code if received is not None else None, + received.reason if received is not None else None, + ) + + +async def _auth_attempt(proxy: str, key: str, model: str, *, delay: bool) -> AuthResult: + try: + async with websockets.connect( + f"{proxy}/v1/responses", + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + idle: Final = await _receive(connection, 35) if delay else ReceiveResult(False, False, None, None, None) + try: + await connection.send(_create(model, f"auth-{uuid.uuid4().hex}")) + frame: Final = JSON_OBJECT.validate_json(await asyncio.wait_for(connection.recv(), timeout=20)) + try: + await asyncio.wait_for(connection.recv(), timeout=20) + except ConnectionClosed as error: + code, reason = _auth_close(error) + return AuthResult(idle, frame, code, reason, None) + return AuthResult(idle, frame, None, None, "connection remained open after rejection") + except (ConnectionClosed, asyncio.TimeoutError) as error: + code, reason = _auth_close(error) if isinstance(error, ConnectionClosed) else (None, None) + return AuthResult(idle, None, code, reason, f"{type(error).__name__}: {error}") + except (ConnectionClosed, asyncio.TimeoutError) as error: + return AuthResult( + ReceiveResult(False, True, None, None, None), + None, + None, + None, + f"{type(error).__name__}: {error}", + ) + + +async def _auth_workload(candidate: Gateway, key: str, model: str, peer: ResponsesPeer) -> AuthResults: + proxy: Final = _proxy_ws_url(candidate, "") + immediate, delayed = await asyncio.gather( + _auth_attempt(proxy, key, model, delay=False), + _auth_attempt(proxy, key, model, delay=True), + ) + return AuthResults(immediate, delayed, peer.connections.qsize()) + + +async def _close_at_limit( + candidate: Gateway, + key: str, + model: str, + *, + delay: float | None = None, +) -> CloseResult: + started: Final = time.monotonic() + outcome: Final = await _close_session(candidate, key, model, delay=delay) + return CloseResult(outcome, time.monotonic() - started) + + +async def _close_session( + candidate: Gateway, + key: str, + model: str, + *, + delay: float | None, +) -> ReceiveResult: + started: Final = time.monotonic() + try: + async with websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + if delay is not None: + await asyncio.sleep(max(0, started + delay - time.monotonic())) + return await _wait_for_close(connection, 75) + except (ConnectionClosed, asyncio.TimeoutError) as error: + code, reason = _auth_close(error) if isinstance(error, ConnectionClosed) else (None, None) + return ReceiveResult(False, True, None, code, reason) + + +async def _turn_result( + connection: websockets.ClientConnection, + frame: str, +) -> tuple[tuple[dict[str, JsonValue], ...], str | None]: + try: + return await _turn(connection, frame), None + except (ConnectionClosed, asyncio.TimeoutError) as error: + return (), f"{type(error).__name__}: {error}" + + +async def _active_session(candidate: Gateway, key: str, model: str, peer: ResponsesPeer) -> ActiveResult: + started: Final = time.monotonic() + try: + async with websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + await asyncio.sleep(max(0, started + 5 - time.monotonic())) + 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) + 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) + return ActiveResult( + (), + f"{type(error).__name__}: {error}", + CloseResult(ReceiveResult(False, True, None, code, reason), time.monotonic() - started), + False, + ) + + +def _provider_closed_within(peer: ResponsesPeer, seconds: float) -> bool: + try: + eventually(lambda: peer.closed.qsize(), lambda count: count >= 1, seconds=seconds) + except AssertionError: + return False + return True + + +async def _stall_turn(connection: websockets.ClientConnection, model: str) -> tuple[bool, str | None]: + try: + await connection.send(_create(model, f"stall-{uuid.uuid4().hex}")) + first: Final = JSON_OBJECT.validate_json(await asyncio.wait_for(connection.recv(), timeout=20)) + return first.get("type") == "response.created", None + except (ConnectionClosed, asyncio.TimeoutError) as error: + return False, f"{type(error).__name__}: {error}" + + +async def _mid_session( + candidate: Gateway, + key: str, + stall_model: str, +) -> tuple[bool, str | None, CloseResult]: + started: Final = time.monotonic() + try: + async with websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + await asyncio.sleep(max(0, started + 40 - time.monotonic())) + created, turn_error = await _stall_turn(connection, stall_model) + remaining: Final = max(0, 75 - (time.monotonic() - started)) + close: Final = CloseResult(await _wait_for_close(connection, remaining), time.monotonic() - started) + return created, turn_error, close + except (ConnectionClosed, asyncio.TimeoutError) as error: + code, reason = _auth_close(error) if isinstance(error, ConnectionClosed) else (None, None) + return ( + False, + f"{type(error).__name__}: {error}", + CloseResult( + ReceiveResult(False, True, None, code, reason), + time.monotonic() - started, + ), + ) + + +async def _fresh_session(candidate: Gateway, key: str, normal_model: str) -> tuple[bool, str | None]: + try: + async with websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + events: Final = await _turn(connection, _create(normal_model, f"fresh-{uuid.uuid4().hex}")) + return _completed_text(events) == PEER_TEXT, None + except (ConnectionClosed, asyncio.TimeoutError) as error: + return False, f"{type(error).__name__}: {error}" + + +async def _mid_response( + candidate: Gateway, + key: str, + normal_model: str, + stall_model: str, + peer: ResponsesPeer, +) -> MidResult: + created, turn_error, close = await _mid_session(candidate, key, stall_model) + provider_closed: Final = _provider_closed_within(peer, 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 _cap_workload( + candidate: Gateway, + key: str, + normal_model: str, + stall_model: str, + peer: ResponsesPeer, +) -> tuple[CloseResult, ActiveResult, MidResult]: + idle, active, mid = 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), + ) + return idle, active, mid + + +def _session_config(path: Path, seconds: int) -> Path: + source: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" + content: Final = source.read_text() + updated: Final = content.replace( + "general_settings:\n", + f"general_settings:\n responses_websocket_session_limit_seconds: {seconds}\n", + 1, + ) + path.write_text(updated) + return path + + +def _responses_body(model: str) -> JsonResponse: + message: Final[dict[str, JsonValue]] = { + "type": "message", + "id": "msg_$UNIQUE_ID", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "responses-$UNIQUE_ID", "annotations": []}], + } + return JsonResponse( + content_type="application/json", + body={ + "id": "$UNIQUE_ID", + "object": "response", + "created_at": 1700000000, + "status": "completed", + "model": model, + "output": [message], + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + }, + ) + + +def _chat_body(model: str) -> JsonResponse: + return JsonResponse( + content_type="application/json", + body={ + "id": "$UNIQUE_ID", + "object": "chat.completion", + "created": 1700000000, + "model": model, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "chat-$UNIQUE_ID"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + ) + + +def _responses_spec(model: str) -> tuple[str, dict[str, JsonValue]]: + return "/v1/responses", {"model": model, "input": f"responses-{uuid.uuid4().hex}"} + + +def _chat_spec(model: str) -> tuple[str, dict[str, JsonValue]]: + return "/v1/chat/completions", { + "model": model, + "messages": [{"role": "user", "content": f"chat-{uuid.uuid4().hex}"}], + } + + +async def _burst_request( + client: httpx.AsyncClient, + key: str, + path: str, + body: dict[str, JsonValue], +) -> BurstResult: + try: + response: Final = await client.post(path, json=body, headers={"Authorization": f"Bearer {key}"}) + try: + parsed: Final = JSON_OBJECT.validate_json(response.content) + except ValidationError as error: + return BurstResult(path, response.status_code, None, str(error)) + return BurstResult(path, response.status_code, parsed, None) + except httpx.HTTPError as error: + return BurstResult(path, None, None, str(error)) + + +async def _chaos_workload( + candidate: Gateway, + key: str, + scripted_model: str, + peer_model: str, +) -> ChaosResults: + base: Final = str(candidate.client.base_url) + specs: Final = tuple( + _responses_spec(scripted_model) if index % 2 == 0 else _chat_spec(scripted_model) for index in range(20) + ) + connections: Final = await asyncio.gather( + *tuple( + websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) + for _ in range(30) + ) + ) + async with httpx.AsyncClient(base_url=base, timeout=30, trust_env=False) as client: + results: Final = await asyncio.gather(*tuple(_burst_request(client, key, path, body) for path, body in specs)) + tuple(connection.transport.abort() for connection in connections) + health: Final = candidate.request("GET", "/health/liveliness").status_code + websocket_completed: Final = await _chaos_turn(candidate, key, peer_model) + return ChaosResults(tuple(results), health, websocket_completed) + + +async def _chaos_turn(candidate: Gateway, key: str, model: str) -> bool: + async with websockets.connect( + _proxy_ws_url(candidate, "/v1/responses"), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as connection: + events: Final = await _turn(connection, _create(model, f"chaos-peer-{uuid.uuid4().hex}")) + return _completed_text(events) == PEER_TEXT + + +@pytest.fixture(scope="module") +def cert(tmp_path_factory: pytest.TempPathFactory) -> tuple[Path, Path]: + return write_self_signed_cert(tmp_path_factory.mktemp("responses-ws-session-limit")) + + +@pytest.fixture(scope="module") +def default_results( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[DefaultResults]: + with responses_peer(cert) as peer, gateway_from_environment() as base: + with ( + owned_proxy( + base, + tmp_path_factory.mktemp("responses-ws-default"), + {"SSL_CERT_FILE": str(cert[0])}, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url) + key: Final = scenario.key(models=[model]) + pool, query = asyncio.run(_default_workload(candidate, key, model)) + provider: Final = _snapshots(peer, 9) + yield DefaultResults(pool, query, provider) + + +@pytest.fixture(scope="module") +def cap_results( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[CapResults]: + with responses_peer(cert) as peer, gateway_from_environment() as base: + config: Final = _session_config(tmp_path_factory.mktemp("responses-ws-cap") / "config.yaml", 60) + with ( + owned_proxy( + base, + tmp_path_factory.mktemp("responses-ws-cap-proxy"), + {"SSL_CERT_FILE": str(cert[0])}, + config=config, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + 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)) + provider: Final = _available_snapshots(peer) + yield CapResults(idle, active, mid, provider) + + +@pytest.fixture(scope="module") +def invalid_results( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[InvalidResult]: + with responses_peer(cert) as peer, gateway_from_environment() as base: + config: Final = _session_config(tmp_path_factory.mktemp("responses-ws-invalid") / "config.yaml", 30) + with owned_proxy_process( + base, + tmp_path_factory.mktemp("responses-ws-invalid-proxy"), + {"SSL_CERT_FILE": str(cert[0])}, + config=config, + workers=2, + ) as owned: + with owned.gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url) + key: Final = scenario.key(models=[model]) + session: Final = asyncio.run( + _idle_turn( + _proxy_ws_url(owned.gateway, ""), + key, + model, + "/v1/responses", + f"invalid-{uuid.uuid4().hex}", + ) + ) + warning: Final = "invalid general_settings.responses_websocket_session_limit_seconds=30" + yield InvalidResult(session, warning in owned.log.read_text()) + + +@pytest.fixture(scope="module") +def chaos_results( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> Iterator[ChaosResults]: + with responses_peer(cert) as peer, gateway_from_environment() as base: + with ( + owned_proxy( + base, + tmp_path_factory.mktemp("responses-ws-chaos"), + {"SSL_CERT_FILE": str(cert[0])}, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + scenario_id: Final = f"responses-chaos-{uuid.uuid4().hex}" + scripted: Final = register_scenario( + scenario_id, + RoutedResponse( + content_type="application/x-routed", + routes={ + "POST /responses": _responses_body("scripted-responses-model"), + "POST /chat/completions": _chat_body("scripted-chat-model"), + }, + ), + control_url=candidate.upstream_url, + ) + scenario.cleanups.callback(delete_scenario, scripted) + scripted_model: Final = scenario.model( + model="openai/scripted-responses-model", + api_base=scripted.api_base(), + ) + peer_model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url) + key: Final = scenario.key(models=[scripted_model, peer_model]) + yield asyncio.run(_chaos_workload(candidate, key, scripted_model, peer_model)) + + +def test_default_config_idle_pool_stays_open_and_completes(default_results: DefaultResults) -> None: + assert all(result.idle.timed_out for result in default_results.pool), default_results.pool + assert all(result.error is None for result in default_results.pool), default_results.pool + assert all(_completed_text(result.events) == PEER_TEXT for result in default_results.pool) + texts: Final = frozenset(result.text for result in default_results.pool) + provider: Final = tuple(snapshot for snapshot in default_results.provider if snapshot.frames) + assert len(provider) == 9, default_results.provider + assert all(snapshot.path == f"/v1/responses?model={PROVIDER_MODEL}" for snapshot in provider) + frames: Final = tuple(snapshot.frames[0] for snapshot in provider) + assert frozenset(_frame_text(frame) for frame in frames) == texts | {default_results.query.text} + assert all(frame.get("model") == PROVIDER_MODEL for frame in frames) + + +def test_default_config_query_model_socket_stays_open_and_completes(default_results: DefaultResults) -> None: + assert default_results.query.idle.timed_out, default_results.query + assert default_results.query.error is None, default_results.query + assert _completed_text(default_results.query.events) == PEER_TEXT + assert any( + default_results.query.text == _frame_text(snapshot.frames[0]) + for snapshot in default_results.provider + if snapshot.frames + ) + + +def test_delayed_first_frame_model_auth_matches_immediate_rejection( + cert: tuple[Path, Path], + tmp_path_factory: pytest.TempPathFactory, +) -> None: + with responses_peer(cert) as peer, gateway_from_environment() as base: + with ( + owned_proxy( + base, + tmp_path_factory.mktemp("responses-ws-auth"), + {"SSL_CERT_FILE": str(cert[0])}, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url) + key: Final = scenario.key(models=["not-authorized-model"]) + results: Final = asyncio.run(_auth_workload(candidate, key, model, peer)) + assert results.immediate.error is None, results.immediate + assert results.delayed.idle.timed_out, results.delayed + assert results.delayed.error is None, results.delayed + assert results.immediate.frame is not None, results.immediate + assert results.immediate.frame.get("type") == "error", results.immediate + rejection: Final = _object(results.immediate.frame["error"]) + assert rejection.get("type") == "invalid_request_error", results.immediate + assert results.immediate.close_code == 1008, results.immediate + assert results.immediate.close_reason == "Pre-call error", results.immediate + assert results.immediate.frame == results.delayed.frame + assert results.immediate.close_code == results.delayed.close_code + assert results.immediate.close_reason == results.delayed.close_reason + assert results.provider_connections == 0, results + + +def test_session_cap_closes_never_started_socket(cap_results: CapResults) -> None: + assert cap_results.idle.outcome.closed, cap_results.idle + assert cap_results.idle.outcome.close_code == 1000, cap_results.idle + assert cap_results.idle.outcome.close_reason == "Session duration limit reached", cap_results.idle + assert 59 <= cap_results.idle.elapsed <= 75, cap_results.idle + + +def test_session_cap_closes_active_socket_and_provider(cap_results: CapResults) -> None: + assert cap_results.active.turn_error is None, cap_results.active + assert _completed_text(cap_results.active.turn) == PEER_TEXT + assert cap_results.active.close.outcome.closed, cap_results.active + assert cap_results.active.close.outcome.close_code == 1000, cap_results.active + assert cap_results.active.close.outcome.close_reason == "Session duration limit reached", cap_results.active + assert 59 <= cap_results.active.close.elapsed <= 75, cap_results.active + assert cap_results.active.provider_closed, cap_results.active + assert all(snapshot.closed for snapshot in cap_results.provider), cap_results.provider + + +def test_session_cap_closes_mid_response_and_allows_new_session(cap_results: CapResults) -> None: + assert cap_results.mid.created, cap_results.mid + assert cap_results.mid.turn_error is None, cap_results.mid + assert cap_results.mid.close.outcome.closed, cap_results.mid + assert cap_results.mid.close.outcome.close_code == 1000, cap_results.mid + assert cap_results.mid.close.outcome.close_reason == "Session duration limit reached", cap_results.mid + assert 59 <= cap_results.mid.close.elapsed <= 75, cap_results.mid + assert cap_results.mid.provider_closed, cap_results.mid + assert cap_results.mid.fresh_error is None, cap_results.mid + assert cap_results.mid.fresh_completed, cap_results.mid + + +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 + assert _completed_text(invalid_results.session.events) == PEER_TEXT + assert invalid_results.warning_found + + +def _matches_scripted_body(result: BurstResult) -> bool: + if result.body is None: + return False + if result.path == "/v1/responses": + output: Final = _list(result.body["output"]) + responses_message: Final = _object(output[0]) + message_id: Final = _string(responses_message["id"]) + if not message_id.startswith("msg_"): + return False + content: Final = _list(responses_message["content"]) + return _string(_object(content[0])["text"]) == f"responses-{message_id.removeprefix('msg_')}" + identity: Final = _string(result.body["id"]) + choices: Final = _list(result.body["choices"]) + chat_message: Final = _object(_object(choices[0])["message"]) + return _string(chat_message["content"]) == f"chat-{identity}" + + +def test_aborted_idle_pool_preserves_http_health_and_new_websocket(chaos_results: ChaosResults) -> None: + results: Final = chaos_results.burst + assert len(results) == 20 + assert all(result.status_code == 200 and result.error is None for result in results), results + identities: Final = tuple(_string(_object(result.body)["id"]) for result in results if result.body is not None) + assert len(set(identities)) == 20, identities + assert all(_matches_scripted_body(result) for result in results), results + assert chaos_results.health_status == 200 + assert chaos_results.websocket_completed From 57630f44aca17287ecd48f27f60e2424cd4407f7 Mon Sep 17 00:00:00 2001 From: mrinal Date: Sat, 3 Oct 2026 21:11:18 +0000 Subject: [PATCH 3/7] chore(ci): sync API schema and Ruff formatting Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/response_api_endpoints/endpoints.py | 4 +--- ui/litellm-dashboard/src/lib/http/schema.d.ts | 6 ++++++ 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 86581d90039..0a3cc21d519 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -1324,9 +1324,7 @@ def _resolve_responses_ws_session_limit_seconds() -> float: ).responses_websocket_session_limit_seconds except ValidationError as e: default: Final = ConfigGeneralSettings.model_fields[field].default - verbose_proxy_logger.warning( - "invalid general_settings.%s=%r (%s); using default %ss", field, raw, e, default - ) + verbose_proxy_logger.warning("invalid general_settings.%s=%r (%s); using default %ss", field, raw, e, default) return float(default) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index ba7dd127aae..e07dc68c562 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -29841,6 +29841,12 @@ export interface components { * @description When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata. */ reject_clientside_metadata_tags?: boolean | null; + /** + * Responses Websocket Session Limit Seconds + * @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. + * @default 3600 + */ + responses_websocket_session_limit_seconds: number; /** @description Spreads the proxy's scheduled background jobs (spend flushes, budget resets, config reloads, exports) across a window instead of firing them together on every replica. On by default; set to tune the window, pin a job, or turn it off. */ scheduled_job_stagger?: components["schemas"]["ScheduledJobStaggerSettings"] | null; /** @description Daily check of the spend LiteLLM captured against the provider's own bill (OpenAI via OPENAI_ADMIN_KEY). Publishes litellm_spend_capture_rate per provider and alerts when the ratio over the lookback window falls under the threshold (default 0.9). Off unless set. */ From 8578a2852c4ee8dac76d78e131a868994dacde9f Mon Sep 17 00:00:00 2001 From: mrinal Date: Sat, 3 Oct 2026 21:46:33 +0000 Subject: [PATCH 4/7] 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", [ From b5f8f2b8b8f50b852f4a7475db3f8cd5833c8add Mon Sep 17 00:00:00 2001 From: mrinal Date: Sat, 3 Oct 2026 22:03:48 +0000 Subject: [PATCH 5/7] test(proxy): assert recorded websocket outcomes and simplify session reaping Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/response_api_endpoints/endpoints.py | 7 -- .../response_api_endpoints/test_endpoints.py | 68 ++++++++++++++++--- 2 files changed, 57 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 320c567be39..188055b74bd 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -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() diff --git a/tests/unit/proxy/response_api_endpoints/test_endpoints.py b/tests/unit/proxy/response_api_endpoints/test_endpoints.py index 0d4db8457a2..a06499d41d4 100644 --- a/tests/unit/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/unit/proxy/response_api_endpoints/test_endpoints.py @@ -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( From c90e63c456139b1e534cb8ba54ff93038f145479 Mon Sep 17 00:00:00 2001 From: mrinal Date: Sat, 3 Oct 2026 23:02:34 +0000 Subject: [PATCH 6/7] 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> --- litellm/proxy/_types.py | 5 +- .../proxy/response_api_endpoints/endpoints.py | 9 +- .../test_responses_websocket_session_limit.py | 97 +++++-- .../response_api_endpoints/test_endpoints.py | 273 +----------------- 4 files changed, 91 insertions(+), 293 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 4cf3ff89bdc..24dfba8e55e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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.", diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 188055b74bd..b19d3d79348 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -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( diff --git a/tests/integration/providers/test_responses_websocket_session_limit.py b/tests/integration/providers/test_responses_websocket_session_limit.py index fc71aea8f2c..57698e1c248 100644 --- a/tests/integration/providers/test_responses_websocket_session_limit.py +++ b/tests/integration/providers/test_responses_websocket_session_limit.py @@ -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 diff --git a/tests/unit/proxy/response_api_endpoints/test_endpoints.py b/tests/unit/proxy/response_api_endpoints/test_endpoints.py index a06499d41d4..07735303027 100644 --- a/tests/unit/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/unit/proxy/response_api_endpoints/test_endpoints.py @@ -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", From b1efbc72ab8af75f97d96bc98123f2772ed6eaf7 Mon Sep 17 00:00:00 2001 From: mrinal Date: Sun, 4 Oct 2026 07:24:42 +0000 Subject: [PATCH 7/7] test(proxy): run provider-close checks off the event loop Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../providers/test_responses_websocket_session_limit.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/integration/providers/test_responses_websocket_session_limit.py b/tests/integration/providers/test_responses_websocket_session_limit.py index 57698e1c248..03e834c222b 100644 --- a/tests/integration/providers/test_responses_websocket_session_limit.py +++ b/tests/integration/providers/test_responses_websocket_session_limit.py @@ -549,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, PROVIDER_MODEL, 5) + provider_closed: Final = await asyncio.to_thread(_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) @@ -637,7 +637,7 @@ 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, STALL_PROVIDER_MODEL, 5) + provider_closed: Final = await asyncio.to_thread(_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) @@ -650,7 +650,7 @@ async def _deaf_response( 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) + provider_closed: Final = await asyncio.to_thread(_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)