mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
test(proxy): move responses websocket session limit timing coverage to integration
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
b5f8f2b8b8
commit
c90e63c456
4 changed files with 91 additions and 293 deletions
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue