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