diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index f71a4d5bbad..1a66f30008c 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -21,7 +21,7 @@ class CodexRealtimeCall(BaseModel): alias: str api_base: str | None = None extra_headers: Mapping[str, str] | None = None - extra_query: Mapping[str, str] | None = None + extra_query: Mapping[str, str | tuple[str, ...]] | None = None usage_supervised: bool = False owner: str expires_at: float @@ -32,7 +32,7 @@ class ChatGPTCallRouting(BaseModel): model_id: str | None = None api_base: str | None = None extra_headers: Mapping[str, str] | None = None - extra_query: Mapping[str, str] | None = None + extra_query: Mapping[str, str | tuple[str, ...]] | None = None class CodexSidebandRequest(TypedDict): @@ -41,7 +41,7 @@ class CodexSidebandRequest(TypedDict): chatgpt_realtime_call_id: ReadOnly[str] query_params: ReadOnly[RealtimeQueryParams] extra_headers: ReadOnly[Mapping[str, str] | None] - extra_query: ReadOnly[Mapping[str, str] | None] + extra_query: ReadOnly[Mapping[str, str | tuple[str, ...]] | None] def build_call_request( diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 559a6ef18f6..cc5773db66c 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -36,18 +36,18 @@ def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping return MappingProxyType({key.lower(): value for key, value in validated.items()}) -def configured_realtime_query(params: GenericLiteLLMParams) -> Mapping[str, str]: +def configured_realtime_query(params: GenericLiteLLMParams) -> Mapping[str, str | tuple[str, ...]]: inbound: Final = TypeAdapter(Mapping[str, str]).validate_python( getattr(params, "chatgpt_realtime_client_query", None) or MappingProxyType({}) ) - configured: Final = TypeAdapter(Mapping[str, str | int | float | bool | None]).validate_python( - getattr(params, "extra_query", None) or MappingProxyType({}) - ) + configured: Final = TypeAdapter( + Mapping[str, str | int | float | bool | None | tuple[str | int | float | bool | None, ...]] + ).validate_python(getattr(params, "extra_query", None) or MappingProxyType({})) + merged: Final = QueryParams( + tuple((key, value) for key, value in inbound.items() if key in ("intent", "architecture")) + ).merge(configured) return MappingProxyType( - { - **{key: value for key, value in inbound.items() if key in ("intent", "architecture")}, - **QueryParams(configured), - } + {key: merged[key] if len(merged.get_list(key)) == 1 else tuple(merged.get_list(key)) for key in merged} ) @@ -132,7 +132,11 @@ class ChatGPTRealtime(OpenAIRealtime): url: Final = base.copy_with( scheme="https" if base.scheme in ("https", "wss") else "http", path=f"{base.path.rstrip('/')}/realtime/calls/{self._call_id}/hangup", - params=tuple((key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id")), + params=tuple( + (key, value) + for key, value in QueryParams(self._extra_query).multi_items() + if key not in ("model", "call_id") + ), ) client: Final = get_async_httpx_client(llm_provider=LlmProviders.CHATGPT) response: Final = await client.post(str(url), headers=self._profile_headers, data=b"", timeout=10) @@ -166,7 +170,9 @@ class ChatGPTRealtime(OpenAIRealtime): endpoint: Final = realtime_endpoint(query_params.get("model", "")) if self._call_id: gateway_query: Final = tuple( - (key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id") + (key, value) + for key, value in QueryParams(self._extra_query).multi_items() + if key not in ("model", "call_id") ) return str( base.copy_with( @@ -182,7 +188,11 @@ class ChatGPTRealtime(OpenAIRealtime): scheme="wss" if base.scheme in ("https", "wss") else "ws", path=f"{base.path.rstrip('/')}/{endpoint}", params=QueryParams(TypeAdapter(Mapping[str, str | None]).validate_python(query_params)).merge( - tuple((key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id")) + tuple( + (key, value) + for key, value in QueryParams(self._extra_query).multi_items() + if key not in ("model", "call_id") + ) ), ) ) diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 2ec576b97b3..819f0b8324a 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -1,9 +1,10 @@ import asyncio import sys +from collections.abc import Mapping from datetime import datetime, timedelta from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn -from pydantic import BaseModel +from pydantic import BaseModel, TypeAdapter from typing_extensions import TypedDict import litellm @@ -50,11 +51,64 @@ class CacheObject(TypedDict): request_count_end_user_id: dict | None +class _RealtimeAttachmentReservations(BaseModel): + cache_keys: tuple[str, ...] = () + global_acquired: bool = False + + def acquire(self, key: str) -> None: + self.cache_keys = tuple(dict.fromkeys((*self.cache_keys, key))) + + def acquire_global(self) -> None: + self.global_acquired = True + + def take(self) -> tuple[tuple[str, ...], bool]: + owned: Final = (self.cache_keys, self.global_acquired) + self.cache_keys = () + self.global_acquired = False + return owned + + class _PROXY_MaxParallelRequestsHandler(CustomLogger): # Class variables or attributes def __init__(self, internal_usage_cache: InternalUsageCache): self.internal_usage_cache = internal_usage_cache + def begin_realtime_attachment(self, request_data: dict[str, object]) -> None: + request_data["_legacy_realtime_attachment_reservations"] = _RealtimeAttachmentReservations() + + async def async_release_realtime_attachment( + self, request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth + ) -> None: + receipt: Final = request_data.get("_legacy_realtime_attachment_reservations") + if not isinstance(receipt, _RealtimeAttachmentReservations): + return + keys, global_acquired = receipt.take() + if global_acquired: + await self.internal_usage_cache.async_increment_cache( + key="global_max_parallel_requests", + value=-1, + local_only=True, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + for key in keys: + await self._release_realtime_counter(key, user_api_key_dict) + + async def _release_realtime_counter(self, key: str, user_api_key_dict: UserAPIKeyAuth) -> None: + raw: Final[object] = await self.internal_usage_cache.async_get_cache( + key=key, + local_only=True, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + if raw is None: + return + current: Final = TypeAdapter(Mapping[str, int]).validate_python(raw) + await self.internal_usage_cache.async_set_cache( + key=key, + value={**current, "current_requests": max(current["current_requests"] - 1, 0)}, + ttl=60, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + def print_verbose(self, print_statement): try: verbose_proxy_logger.debug(print_statement) @@ -142,6 +196,9 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): litellm_parent_otel_span=user_api_key_dict.parent_otel_span, local_only=True, ) + receipt: Final = data.get("_legacy_realtime_attachment_reservations") + if isinstance(receipt, _RealtimeAttachmentReservations): + receipt.acquire(request_count_api_key) return new_val def time_to_next_minute(self) -> float: @@ -299,6 +356,9 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): local_only=True, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, ) + receipt: Final = data.get("_legacy_realtime_attachment_reservations") + if isinstance(receipt, _RealtimeAttachmentReservations): + receipt.acquire_global() _model = data.get("model", None) current_date: Final = datetime.now().strftime("%Y-%m-%d") @@ -480,6 +540,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): values_to_update_in_cache=values_to_update_in_cache, ) + if isinstance(data.get("_legacy_realtime_attachment_reservations"), _RealtimeAttachmentReservations): + await self.internal_usage_cache.async_batch_set_cache( + cache_list=values_to_update_in_cache, + ttl=60, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + return asyncio.create_task( self.internal_usage_cache.async_batch_set_cache( cache_list=values_to_update_in_cache, diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index c6c3dde4b6e..905a2f02fb5 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -4677,6 +4677,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) stash.parallel_slot = None + async def async_release_realtime_attachment( + self, request_data: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth + ) -> None: + await self.async_post_call_failure_hook( + request_data={}, # mutable-ok: existing failure hook requires dict; attachment has no billable usage + original_exception=Exception("Realtime attachment completed"), + user_api_key_dict=user_api_key_dict, + ) + async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): """ Post-call hook to update rate limit headers in the response. diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 8af489e8eb1..1cec36876b2 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -38,6 +38,12 @@ from litellm.proxy.auth.user_api_key_auth import ( user_api_key_auth, ) from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # existing built-in limiter has no public alias +) +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # existing built-in limiter has no public alias +) from litellm.proxy.spend_tracking.budget_reservation import ( invalidate_budget_reservation_counters, release_or_invalidate_budget_reservation, @@ -324,6 +330,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP p.strip() for p in websocket.headers.get("sec-websocket-protocol", "").split(",") if p.strip() ) 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 try: try: api_key: Final = get_websocket_api_key(websocket) @@ -364,6 +371,13 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP name.strip() for name in websocket.query_params.get("guardrails", "").split(",") if name.strip() ], } + limiter: Final = server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter") + if call.usage_supervised and isinstance( + limiter, (_PROXY_MaxParallelRequestsHandler, _PROXY_MaxParallelRequestsHandler_v3) + ): + attachment_limiter = limiter + if isinstance(limiter, _PROXY_MaxParallelRequestsHandler): + limiter.begin_realtime_attachment(data) try: processed, logging_obj = await process_codex_request(request, data, auth, call.alias, "_arealtime") except Exception: # noqa: BLE001 # custom hook exceptions must reject the connection @@ -397,5 +411,9 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP }, ) finally: - 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) + try: + if attachment_limiter is not None: + await attachment_limiter.async_release_realtime_attachment(data, auth) + finally: + 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/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index ae04582a36b..ddf1e8f3f49 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -3,7 +3,7 @@ from collections.abc import AsyncIterator, Awaitable, Callable from contextlib import suppress from typing import Final, Protocol -from pydantic import BaseModel +from pydantic import BaseModel, Field, ValidationError from websockets.exceptions import ConnectionClosedOK from litellm._logging import verbose_proxy_logger @@ -33,6 +33,14 @@ class _ObserverEvent(BaseModel): type: str +class _LiveDurationUsage(BaseModel): + audio_duration_ms: float = Field(strict=True, ge=0, allow_inf_nan=False) + + +class _LiveTerminalEvent(BaseModel): + usage: _LiveDurationUsage + + class CallSupervisor: def __init__( self, @@ -66,6 +74,7 @@ class CallSupervisor: self._stop = asyncio.Event() self._started = False self._terminal = False + self._terminal_usage_valid = False self._close_confirmed = False self._accounting_complete = False self._task: asyncio.Task[None] | None = None @@ -106,10 +115,18 @@ class CallSupervisor: self._ready.set() if event.type == "session.closed": self._terminal = True + try: + _LiveTerminalEvent.model_validate_json(message) + except ValidationError: + self._terminal_usage_valid = False + else: + self._terminal_usage_valid = True return def _usage_complete(self) -> bool: - return self._terminal or (not self._terminal_usage_required and self._close_confirmed) + if self._terminal_usage_required: + return self._terminal and self._terminal_usage_valid + return self._terminal or self._close_confirmed async def _run(self) -> None: reader: Final = asyncio.create_task(self._read()) diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 402afaaa58c..35606931e24 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -1,9 +1,30 @@ +import hashlib +import time + import httpx import pytest from litellm.llms.chatgpt.codex import CodexRealtimeCall, build_sideband_request, parse_call_response +def test_encrypted_call_preserves_repeated_gateway_query(monkeypatch): + from litellm.proxy.realtime_endpoints.call_sessions import decode_call, encode_call + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-repeated-query") + authorization = "Bearer test-owner" + call = CodexRealtimeCall( + call_id="rtc_repeated", + model="gpt-live-1-codex", + alias="voice", + owner=hashlib.sha256(authorization.encode()).hexdigest(), + expires_at=time.time() + 60, + extra_query={"tag": ["alpha +/&", "beta"], "gateway": "tenant"}, + ) + restored = decode_call(encode_call(call), authorization) + assert restored.extra_query == {"tag": ("alpha +/&", "beta"), "gateway": "tenant"} + assert build_sideband_request(restored)["extra_query"] == restored.extra_query + + @pytest.mark.parametrize("location", ["", "/v1/realtime/calls/foreign-id"]) def test_signaling_rejects_invalid_upstream_call_id(location): response = httpx.Response(201, headers={"Location": location}, diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index b02b3a551f2..f3258fd20b4 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -27,7 +27,7 @@ async def test_live_closed_observer_uses_independent_hangup(failure, hangup_stat GenericLiteLLMParams( chatgpt_realtime_call_id="rtc_live_closed", chatgpt_token_dir=chatgpt_tokens, - extra_query={"gateway": "tenant"}, + extra_query={"gateway": "tenant", "tag": ["alpha +/&", "beta"]}, ), {}, {"x-gateway-token": "test-only"}, @@ -59,7 +59,9 @@ async def test_live_closed_observer_uses_independent_hangup(failure, hangup_stat await client.aclose() assert len(requests) == 2 assert requests[0].method == "POST" - assert str(requests[0].url) == "https://gateway.example/v1/realtime/calls/rtc_live_closed/hangup?gateway=tenant" + assert requests[0].url.path == "/v1/realtime/calls/rtc_live_closed/hangup" + assert requests[0].url.params.get_list("tag") == ["alpha +/&", "beta"] + assert requests[0].url.params["gateway"] == "tenant" assert requests[0].headers["x-gateway-token"] == "test-only" assert requests[0].headers["Authorization"] == "Bearer test-token-default" assert requests[0].extensions["timeout"]["read"] == 10 @@ -138,6 +140,8 @@ async def test_routed_call_preserves_deployment_gateway_headers( "enabled": True, "disabled": False, "blank": None, + "tag": ["alpha +/&", "beta"], + "empty": [], "model": "other-model", "call_id": "rtc_wrong", }, @@ -163,10 +167,16 @@ async def test_routed_call_preserves_deployment_gateway_headers( "enabled": "true", "disabled": "false", "blank": "", + "tag": "alpha +/&", "model": "other-model", "call_id": "rtc_wrong", } - assert response.extensions["chatgpt_realtime"]["extra_query"] == dict(requests[0].url.params) + assert requests[0].url.params.get_list("tag") == ["alpha +/&", "beta"] + assert response.extensions["chatgpt_realtime"]["extra_query"] == { + **dict(requests[0].url.params), + "tag": ("alpha +/&", "beta"), + "empty": (), + } assert response.extensions["chatgpt_realtime"]["extra_headers"]["x-gateway-route"] == "configured" for name, value in inbound_headers.items(): assert requests[0].headers[name] == value @@ -180,6 +190,7 @@ async def test_routed_call_preserves_deployment_gateway_headers( key: value for key, value in requests[0].url.params.items() if key not in ("model", "call_id") } assert sideband_url.params.get("call_id") == ("rtc_test" if endpoint == "realtime" else None) + assert sideband_url.params.get_list("tag") == ["alpha +/&", "beta"] assert sideband_url.path.endswith("/realtime" if endpoint == "realtime" else "/live/rtc_test") finally: await client.client.aclose() @@ -203,6 +214,8 @@ async def test_websocket_forwards_configured_headers_without_client_identity(mod websocket=websocket, api_base="https://voice.example/codex", chatgpt_realtime_call_id=call_id, + query_params={"model": model, "intent": "client-intent"}, + extra_query={"intent": "configured-intent", "tag": ["alpha +/&", "beta"]}, headers={"x-deployment-header": "configured"}, extra_headers={ "X-Gateway-Route": "voice", @@ -213,6 +226,9 @@ async def test_websocket_forwards_configured_headers_without_client_identity(mod ) connect.assert_called_once() headers = httpx.Headers(connect.call_args.kwargs["additional_headers"]) + upstream_url = httpx.URL(connect.call_args.args[0]) + assert upstream_url.params.get_list("intent") == ["configured-intent"] + assert upstream_url.params.get_list("tag") == ["alpha +/&", "beta"] assert headers["x-deployment-header"] == "configured" assert headers["x-gateway-route"] == "voice" assert headers["openai-alpha"] == "configured-value" diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py index 0e2683dcbfd..0bf488016c2 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -7,6 +7,8 @@ from datetime import datetime import pytest from litellm.caching.caching import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.parallel_request_limiter import ( _PROXY_MaxParallelRequestsHandler, ) @@ -14,6 +16,78 @@ from litellm.proxy.utils import InternalUsageCache, hash_token from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage +@pytest.mark.asyncio +@pytest.mark.parametrize("reject_team", [False, True]) +async def test_realtime_attachment_releases_only_acquired_legacy_slots(reject_team): + cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache)) + auth = UserAPIKeyAuth( + api_key="attachment-key", + user_id="attachment-user", + team_id="attachment-team", + team_rpm_limit=0 if reject_team else 100, + max_parallel_requests=1, + end_user_id="attachment-end-user", + metadata={"model_rpm_limit": {"test-model": 100}}, + ) + data = {"model": "test-model", "metadata": {"global_max_parallel_requests": 10}} + minute = datetime.now().strftime("%Y-%m-%d-%H-%M") + team_key = f"attachment-team::{minute}::request_count" + await cache.async_set_cache(team_key, {"current_requests": 3, "current_tpm": 7, "current_rpm": 4}) + handler.begin_realtime_attachment(data) + if reject_team: + with pytest.raises(ProxyRateLimitError, match="Rate Limit Handler"): + await handler.async_pre_call_hook(auth, cache, data, "_arealtime") + else: + await handler.async_pre_call_hook(auth, cache, data, "_arealtime") + await handler.async_release_realtime_attachment(data, auth) + await handler.async_release_realtime_attachment(data, auth) + assert await cache.async_get_cache("global_max_parallel_requests") == 0 + assert await cache.async_get_cache(f"attachment-key::{minute}::request_count") == { + "current_requests": 0, + "current_tpm": 0, + "current_rpm": 1, + } + assert await cache.async_get_cache(f"attachment-user::{minute}::request_count") == { + "current_requests": 0, + "current_tpm": 0, + "current_rpm": 1, + } + assert await cache.async_get_cache(team_key) == { + "current_requests": 3, + "current_tpm": 7, + "current_rpm": 4 if reject_team else 5, + } + assert await cache.async_get_cache(f"attachment-key::test-model::{minute}::request_count") == { + "current_requests": 0, + "current_tpm": 0, + "current_rpm": 1, + } + end_user = await cache.async_get_cache(f"attachment-end-user::{minute}::request_count") + assert end_user == (None if reject_team else {"current_requests": 0, "current_tpm": 0, "current_rpm": 1}) + if not reject_team: + handler.begin_realtime_attachment(data) + await handler.async_pre_call_hook(auth, cache, data, "_arealtime") + await handler.async_release_realtime_attachment(data, auth) + + +@pytest.mark.asyncio +async def test_realtime_attachment_rejected_before_acquisition_preserves_other_slot(): + cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache)) + auth = UserAPIKeyAuth(api_key="busy-key", max_parallel_requests=1) + minute = datetime.now().strftime("%Y-%m-%d-%H-%M") + key = f"busy-key::{minute}::request_count" + current = {"current_requests": 1, "current_tpm": 13, "current_rpm": 2} + await cache.async_set_cache(key, current) + data = {"model": "test-model"} + handler.begin_realtime_attachment(data) + with pytest.raises(ProxyRateLimitError, match="Rate Limit Handler"): + await handler.async_pre_call_hook(auth, cache, data, "_arealtime") + await handler.async_release_realtime_attachment(data, auth) + assert await cache.async_get_cache(key) == current + + @pytest.mark.parametrize( "response_obj", [ @@ -39,9 +113,7 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ team_id = "litellm-team" end_user_id = "customer-1" - parallel_request_handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(DualCache()) - ) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) current_date = datetime.now().strftime("%Y-%m-%d") current_hour = datetime.now().strftime("%H") @@ -80,7 +152,4 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ key=f"{scope_id}::{precise_minute}::request_count", litellm_parent_otel_span=None, ) - assert current["current_tpm"] == 50, ( - f"expected 50 tokens counted for {scope_id}, " - f"got {current['current_tpm']}" - ) + assert current["current_tpm"] == 50, f"expected 50 tokens counted for {scope_id}, got {current['current_tpm']}" 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 ef85e324680..50e3aebfbe2 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -83,6 +83,103 @@ async def test_sideband_preserves_pending_cost_reconciliation(monkeypatch, logge assert auth.budget_reservation["finalized"] is not logged_success +@pytest.mark.asyncio +@pytest.mark.parametrize("ending", ["normal", "disconnect", "pre_call", "admission"]) +async def test_supervised_attachments_release_real_limiter_before_reconnect(monkeypatch, ending): + import asyncio + from unittest.mock import AsyncMock + + import litellm + from litellm.caching.caching import DualCache + from litellm.proxy import proxy_server as server + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + _request_stash, + get_request_stash, + ) + from litellm.proxy.utils import InternalUsageCache + + cache = DualCache() + limiter = _PROXY_MaxParallelRequestsHandler_v3(InternalUsageCache(cache)) + auth = UserAPIKeyAuth(api_key="attachment-owner", max_parallel_requests=1, tpm_limit=10000) + token_key = limiter.create_rate_limit_keys(key="api_key", value=auth.api_key, rate_limit_type="tokens") + parallel_key = f"{{api_key:{auth.api_key}}}:max_parallel_requests" + 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, + ) + monkeypatch.setenv("LITELLM_SALT_KEY", "attachment-cleanup-test") + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(server, "proxy_logging_obj", SimpleNamespace(get_proxy_hook=lambda name: limiter)) + + async def process(request, data, selected_auth, model, call_type): + await limiter.async_pre_call_hook( + user_api_key_dict=selected_auth, + cache=cache, + data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hello"}], "max_tokens": 50}, + call_type="completion", + ) + assert get_request_stash().reserved_tokens > 0 + if ending == "pre_call": + raise RuntimeError("Later policy rejected attachment") + return data, SimpleNamespace(model_call_details={}) + + async def forward(**kwargs): + if ending == "disconnect": + raise asyncio.CancelledError() + + monkeypatch.setattr(codex, "process_codex_request", process) + monkeypatch.setattr(litellm, "_arealtime", forward) + blocker_stash = None + blocker_reserved = 0 + if ending == "admission": + setup_token = _request_stash.set(None) + try: + await limiter.async_pre_call_hook( + user_api_key_dict=auth, + cache=cache, + data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hello"}], "max_tokens": 50}, + call_type="completion", + ) + blocker_stash = get_request_stash() + blocker_reserved = blocker_stash.reserved_tokens + finally: + _request_stash.reset(setup_token) + for _ in range(3): + stash_token = _request_stash.set(None) + try: + websocket = WebSocket( + { + "type": "websocket", + "path": "/v1/live/opaque", + "query_string": b"", + "headers": [(b"authorization", b"Bearer owner")], + }, + AsyncMock(return_value={"type": "websocket.connect"}), + AsyncMock(), + ) + if ending == "disconnect": + with pytest.raises(asyncio.CancelledError): + await codex.codex_realtime_sideband(websocket, encode_call(call), auth) + else: + await codex.codex_realtime_sideband(websocket, encode_call(call), auth) + assert limiter._gauge_in_flight_from_cache_value(await cache.async_get_cache(parallel_key)) == int( + ending == "admission" + ) + assert int(await cache.async_get_cache(token_key) or 0) == blocker_reserved + finally: + _request_stash.reset(stash_token) + if blocker_stash is not None: + cleanup_token = _request_stash.set(blocker_stash) + try: + await limiter.async_release_realtime_attachment({}, auth) + finally: + _request_stash.reset(cleanup_token) + def test_sideband_token_binds_owner_and_model(monkeypatch): monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") call = CodexRealtimeCall( diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index c14597a012a..93b97a28fe6 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -147,6 +147,34 @@ class Sink: self.logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True +@pytest.mark.asyncio +@pytest.mark.parametrize( + "duration,valid", [(0, True), (1000, True), (None, False), (-1, False), (True, False), ("1000", False)] +) +async def test_live_terminal_requires_valid_duration_for_accounting(monkeypatch, duration, valid): + from litellm.proxy.realtime_endpoints import call_supervision + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + invalidate = AsyncMock() + monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate) + close = AsyncMock() + force = AsyncMock() + supervisor = CallSupervisor(socket, Sink(logger), logger, UserAPIKeyAuth(), close, force_close_call=force) + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await socket.messages.put( + {"type": "session.closed", **({"usage": {"audio_duration_ms": duration}} if duration is not None else {})} + ) + await supervisor.wait() + close.assert_not_awaited() + force.assert_not_awaited() + assert socket.closed + assert bool(logger.model_call_details.get("realtime_usage_incomplete")) is not valid + assert invalidate.await_count == (0 if valid else 1) + + def fixture(*, ready_timeout=1, lifetime=1): socket = Socket() logger = MagicMock(spec=Logging) @@ -454,7 +482,14 @@ async def test_shutdown_allows_hangup_longer_than_usage_drain_timeout(): hangup_finished.set() supervisor = CallSupervisor( - socket, sink, logger, UserAPIKeyAuth(), hangup, drain_timeout=0.01, termination_timeout=1 + socket, + sink, + logger, + UserAPIKeyAuth(), + hangup, + drain_timeout=0.01, + termination_timeout=1, + terminal_usage_required=False, ) registry = CallSupervisors() await socket.messages.put({"type": "session.created"})