mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): release attachment quota before upstream websocket close
This commit is contained in:
parent
09e6b3a89d
commit
2642d98302
4 changed files with 154 additions and 4 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue