diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index c5e6b0fafe3..18f7eaa0e32 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1,7 +1,8 @@ import asyncio import json import traceback -from collections.abc import Coroutine, Mapping, Sequence +from collections.abc import Awaitable, Callable, Coroutine, Mapping, Sequence +from contextvars import ContextVar from dataclasses import dataclass from enum import Enum, auto from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast @@ -26,6 +27,10 @@ from litellm.types.realtime import ALL_DELTA_TYPES from .litellm_logging import Logging as LiteLLMLogging from .realtime_errors import client_close_code, realtime_error_event, websocket_close_reason +realtime_attachment_cleanup: Final[ContextVar[Callable[[], Awaitable[None]] | None]] = ContextVar( + "realtime_attachment_cleanup", default=None +) + if TYPE_CHECKING: from websockets.asyncio.client import ClientConnection from websockets.exceptions import ConnectionClosed @@ -1581,7 +1586,12 @@ class RealTimeStreaming: finally: forward_task.cancel() client_task.cancel() - await asyncio.gather(forward_task, client_task, return_exceptions=True) + try: + await asyncio.gather(forward_task, client_task, return_exceptions=True) + finally: + cleanup: Final = realtime_attachment_cleanup.get() + if not self._account_usage and cleanup is not None: + await cleanup() async def _close_client(self, close: BackendClose) -> None: redacted_message: Final = redact_internal_details_from_client_message(close.message) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 1cec36876b2..6ed2e83ab1c 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -2,8 +2,9 @@ import base64 import hashlib import json import time -from collections.abc import Mapping +from collections.abc import Awaitable, Callable, Mapping from contextlib import AsyncExitStack +from contextvars import Token from types import MappingProxyType from typing import Final, Literal @@ -14,7 +15,11 @@ from starlette.types import Message from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import Logging -from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY, RealTimeStreaming +from litellm.litellm_core_utils.realtime_streaming import ( + REALTIME_SESSION_SUCCESS_LOGGED_KEY, + RealTimeStreaming, + realtime_attachment_cleanup, +) from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.chatgpt.codex import ( CodexRealtimeCall, @@ -331,6 +336,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP ) logging_obj: Logging | None = None # rebind-ok: cleanup needs the logger only after pre-call succeeds attachment_limiter: _PROXY_MaxParallelRequestsHandler | _PROXY_MaxParallelRequestsHandler_v3 | None = None + cleanup_token: Token[Callable[[], Awaitable[None]] | None] | None = None try: try: api_key: Final = get_websocket_api_key(websocket) @@ -387,6 +393,13 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP await websocket.accept( subprotocol=next((p for p in protocols if not p.startswith("openai-insecure-api-key.")), None) ) + if attachment_limiter is not None: + selected_limiter: Final = attachment_limiter + + async def release_attachment() -> None: + await selected_limiter.async_release_realtime_attachment(data, auth) + + cleanup_token = realtime_attachment_cleanup.set(release_attachment) await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # dispatch for an already authorized call model=f"chatgpt/{call.model}", websocket=websocket, @@ -415,5 +428,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP if attachment_limiter is not None: await attachment_limiter.async_release_realtime_attachment(data, auth) finally: + if cleanup_token is not None: + realtime_attachment_cleanup.reset(cleanup_token) if logging_obj is None or not logging_obj.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY): await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index cba46374837..5b263f9d557 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -3438,3 +3438,37 @@ async def test_live_attachment_does_not_dispatch_duplicate_usage(): await stream.log_messages() worker.ensure_initialized_and_enqueue.assert_not_called() logger.dispatch_success_handlers.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("account_usage", [False, True]) +async def test_attachment_cleanup_runs_in_owning_context_only(account_usage): + from litellm.litellm_core_utils.realtime_streaming import realtime_attachment_cleanup + + contexts = [] + + async def one(name): + task = asyncio.current_task() + callback = AsyncMock(side_effect=lambda: contexts.append((name, asyncio.current_task() is task))) + token = realtime_attachment_cleanup.set(callback) + try: + websocket = MagicMock() + websocket.receive_text = AsyncMock(side_effect=RuntimeError("disconnected")) + backend = MagicMock() + + async def recv(**kwargs): + await asyncio.Event().wait() + + backend.recv = recv + stream = RealTimeStreaming(websocket, backend, MagicMock(), account_usage=account_usage) + await stream.bidirectional_forward() + if account_usage: + callback.assert_not_awaited() + else: + callback.assert_awaited_once() + finally: + realtime_attachment_cleanup.reset(token) + + await asyncio.gather(one("first"), one("second")) + assert sorted(contexts) == ([] if account_usage else [("first", True), ("second", True)]) + assert realtime_attachment_cleanup.get() is None diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index 50e3aebfbe2..2591d45f9a5 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -664,3 +664,94 @@ async def test_supervisor_constructor_failure_closes_effective_connection(monkey release.assert_awaited_once_with(budget_reservation=auth.budget_reservation) invalidate.assert_not_awaited() assert "private-cleanup-credential" not in caplog.text + + +@pytest.mark.asyncio +async def test_attachment_releases_quota_before_upstream_close_handshake(monkeypatch): + import asyncio + from unittest.mock import AsyncMock + + import websockets + + import litellm + from litellm.proxy import proxy_server as server + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.hooks.parallel_request_limiter_v3 import _request_stash + from litellm.proxy.utils import ProxyLogging + + proxy = ProxyLogging(UserApiKeyCache()) + monkeypatch.setattr(litellm, "callbacks", []) + proxy._add_proxy_hooks() + monkeypatch.setattr(server, "proxy_logging_obj", proxy) + monkeypatch.setattr( + server, + "llm_router", + litellm.Router( + model_list=[ + {"model_name": "voice", "litellm_params": {"model": "openai/gpt-realtime-1.5", "api_key": "test"}} + ] + ), + ) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setenv("LITELLM_SALT_KEY", "attachment-close-order-test") + from litellm.llms.chatgpt.authenticator import Authenticator + + monkeypatch.setattr(Authenticator, "get_access_token", lambda self: "test-token") + monkeypatch.setattr(Authenticator, "get_account_id", lambda self: "test-account") + closing = asyncio.Event() + finish_close = asyncio.Event() + + class Backend: + async def recv(self, **kwargs): + await asyncio.Event().wait() + + async def send(self, value): + return None + + class Connection: + async def __aenter__(self): + return Backend() + + async def __aexit__(self, *args): + closing.set() + await finish_close.wait() + + monkeypatch.setattr(websockets, "connect", lambda *args, **kwargs: Connection()) + auth = UserAPIKeyAuth(api_key="close-order-owner", max_parallel_requests=1) + call = CodexRealtimeCall( + call_id="rtc_test", + model="gpt-live-1-codex", + alias="voice", + usage_supervised=True, + owner=hashlib.sha256(b"Bearer owner").hexdigest(), + expires_at=time.time() + 300, + ) + ws = WebSocket( + { + "type": "websocket", + "path": "/v1/live/test", + "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")], + "scheme": "ws", + "server": ("localhost", 80), + }, + AsyncMock(side_effect=[{"type": "websocket.connect"}, {"type": "websocket.disconnect", "code": 1000}]), + AsyncMock(), + ) + token = _request_stash.set(None) + request = asyncio.create_task(codex.codex_realtime_sideband(ws, encode_call(call), auth)) + try: + await asyncio.wait_for(closing.wait(), timeout=5) + limiter = proxy.get_proxy_hook("parallel_request_limiter") + value = await proxy.internal_usage_cache.async_get_cache( + "{api_key:close-order-owner}:max_parallel_requests", litellm_parent_otel_span=None, local_only=True + ) + assert limiter._gauge_in_flight_from_cache_value(value) == 0 + finally: + finish_close.set() + await asyncio.wait_for(request, timeout=5) + _request_stash.reset(token) + value = await proxy.internal_usage_cache.async_get_cache( + "{api_key:close-order-owner}:max_parallel_requests", litellm_parent_otel_span=None, local_only=True + ) + assert limiter._gauge_in_flight_from_cache_value(value) == 0