fix(chatgpt): release attachment quota before upstream websocket close

This commit is contained in:
jibanez-staticduo 2026-09-10 22:08:10 +02:00
parent 09e6b3a89d
commit 2642d98302
No known key found for this signature in database
4 changed files with 154 additions and 4 deletions

View file

@ -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)

View file

@ -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)

View file

@ -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

View file

@ -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