mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): keep idle responses websockets open until a configurable session limit (#44433)
* add config-driven bound to /responses websocket connection * 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> * chore(ci): sync API schema and Ruff formatting Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * 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> * test(proxy): assert recorded websocket outcomes and simplify session reaping Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * 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> * test(proxy): run provider-close checks off the event loop Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): cover handshake auth, subprotocol, DB override and worker kill for the responses websocket session limit Adds the audit cells for the configurable session limit: a rejected bearer fails the handshake with 403, a requested Sec-WebSocket-Protocol is echoed back, a DB override through /config/field/update caps new sessions on every worker without a restart while sessions accepted before it keep their limit, deleting the override restores the default, out-of-range and wrong-type updates answer 400 and leave the stored row alone, and SIGKILLing the worker holding idle sockets drops only its sockets while the proxy keeps serving * test(integration): open override sockets until both proxy workers hold one uvicorn workers share one listening socket and a worker that wakes first accepts a whole simultaneous burst, so eight sockets opened at once can all land on one worker. The DB override cell now opens its sockets one at a time, alternating the two paths, until at least eight are open and every worker holds one, bounded at forty --------- Co-authored-by: Mrinal Chanshetty <mrinal@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
efe0715831
commit
4f49f88172
5 changed files with 1609 additions and 38 deletions
|
|
@ -2738,6 +2738,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
|
||||
|
|
@ -3048,6 +3051,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=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.",
|
||||
)
|
||||
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",
|
||||
|
|
|
|||
|
|
@ -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 ConfigDict, ValidationError
|
||||
from pydantic import ConfigDict, TypeAdapter, ValidationError
|
||||
from starlette.websockets import WebSocket, WebSocketDisconnect
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
|
|
@ -50,6 +50,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
|
||||
|
||||
|
|
@ -1316,6 +1317,21 @@ 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 = _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 = DEFAULT_RESPONSES_WEBSOCKET_SESSION_LIMIT_SECONDS
|
||||
verbose_proxy_logger.warning("invalid general_settings.%s=%r (%s); using default %ss", field, raw, e, default)
|
||||
return default
|
||||
|
||||
|
||||
async def _read_ws_model_from_first_frame(
|
||||
websocket: WebSocket,
|
||||
query_model: str | None = None,
|
||||
|
|
@ -1323,12 +1339,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:
|
||||
|
|
@ -1474,25 +1488,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,
|
||||
|
|
@ -1507,16 +1507,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
|
||||
|
|
@ -1621,3 +1611,55 @@ 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_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(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")
|
||||
finally:
|
||||
if not session_task.done():
|
||||
session_task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError, Exception):
|
||||
await session_task
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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])
|
||||
|
|
@ -957,6 +963,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(
|
||||
|
|
@ -973,7 +980,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,
|
||||
|
|
@ -1000,7 +1007,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 (
|
||||
|
|
@ -1016,7 +1023,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):
|
||||
|
|
@ -1108,6 +1115,34 @@ class TestReadWSModelFromFirstFrameErrors:
|
|||
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 (
|
||||
|
|
|
|||
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -30330,6 +30330,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;
|
||||
/**
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue