From 34dcfbf833373c7fa47ae56600041d6708f68d5b Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Fri, 11 Sep 2026 02:25:21 +0200 Subject: [PATCH] fix(chatgpt): retain realtime call quotas and isolate provider routing --- .../base_llm/realtime/http_transformation.py | 14 + litellm/llms/chatgpt/codex.py | 1 + litellm/llms/chatgpt/realtime.py | 34 ++- litellm/llms/openai/realtime/handler.py | 17 +- .../proxy/hooks/parallel_request_limiter.py | 40 ++- .../hooks/parallel_request_limiter_v3.py | 118 +++++++- litellm/proxy/hooks/realtime_call_lease.py | 74 +++++ .../proxy/realtime_endpoints/call_sessions.py | 97 +++++-- .../realtime_endpoints/call_supervision.py | 29 +- litellm/realtime_api/main.py | 82 +++--- litellm/utils.py | 23 +- .../hooks/test_parallel_request_limiter.py | 105 +++++++ .../hooks/test_parallel_request_limiter_v3.py | 129 +++++++++ .../proxy/hooks/test_realtime_call_lease.py | 85 ++++++ .../realtime_endpoints/test_call_sessions.py | 274 ++++++++++++++++++ .../test_call_supervision.py | 36 +++ tests/test_litellm/realtime_api/test_main.py | 26 ++ 17 files changed, 1090 insertions(+), 94 deletions(-) create mode 100644 litellm/proxy/hooks/realtime_call_lease.py create mode 100644 tests/test_litellm/proxy/hooks/test_realtime_call_lease.py diff --git a/litellm/llms/base_llm/realtime/http_transformation.py b/litellm/llms/base_llm/realtime/http_transformation.py index 43a80edb493..80daebd88b2 100644 --- a/litellm/llms/base_llm/realtime/http_transformation.py +++ b/litellm/llms/base_llm/realtime/http_transformation.py @@ -7,6 +7,7 @@ These are HTTP (not WebSocket) endpoints used by the WebRTC flow: """ from abc import ABC, abstractmethod +from collections.abc import Mapping from typing import Final import httpx @@ -36,6 +37,14 @@ class BaseRealtimeHTTPConfig(ABC): explicit api_base → litellm.api_base → env var → hard-coded default """ + def resolve_api_base(self, api_base: str | None, dynamic_api_base: str | None) -> str: + return self.get_api_base(dynamic_api_base or api_base) + + def get_realtime_calls_extra_headers( + self, headers: dict[str, object] | None + ) -> dict[str, object] | None: # mutable-ok: shared HTTP handler accepts a mutable header dictionary + return headers + @abstractmethod def get_api_key( self, @@ -97,6 +106,11 @@ class BaseRealtimeHTTPConfig(ABC): "Authorization": f"Bearer {ephemeral_key}", } + def transform_realtime_calls_response( + self, response: httpx.Response, model: str, model_id: str | None, headers: Mapping[str, object] | None + ) -> httpx.Response: + return response + # ------------------------------------------------------------------ # # Error handling # # ------------------------------------------------------------------ # diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index 1a66f30008c..ea736c97a0d 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -23,6 +23,7 @@ class CodexRealtimeCall(BaseModel): extra_headers: Mapping[str, str] | None = None extra_query: Mapping[str, str | tuple[str, ...]] | None = None usage_supervised: bool = False + parallel_reserved: bool = False owner: str expires_at: float diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index cc5773db66c..f6aa04fc29d 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -3,7 +3,7 @@ from enum import Enum, auto from types import MappingProxyType from typing import TYPE_CHECKING, Final -from httpx import URL, QueryParams +from httpx import URL, QueryParams, Response from pydantic import TypeAdapter from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES @@ -156,6 +156,16 @@ class ChatGPTRealtime(OpenAIRealtime): self._profile_headers = realtime_headers(params, headers, extra_headers) self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None)) self._extra_query = configured_realtime_query(params) + self._account_usage = accounts_for_call_usage(params) + + def _get_default_api_base(self) -> str: + return self.get_api_base() + + def _resolve_api_key(self, api_key: str | None) -> str: + return "chatgpt-oauth" + + def _accounts_for_call_usage(self) -> bool: + return self._account_usage def _get_additional_headers( self, api_key: str, *, openai_beta_realtime: bool = False @@ -212,6 +222,14 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): ) -> str: return api_base or (Authenticator.get_api_base() if self._use_codex_backend else ChatGPTRealtime.get_api_base()) + def resolve_api_base(self, api_base: str | None, dynamic_api_base: str | None) -> str: + return self.get_api_base(api_base) + + def get_realtime_calls_extra_headers( + self, headers: dict[str, object] | None + ) -> dict[str, object]: # mutable-ok: shared HTTP handler accepts a mutable header dictionary + return {**realtime_call_headers(self._params)} # mutable-ok: shared HTTP header contract + def get_api_key( self, api_key: str | None, @@ -223,6 +241,20 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): query: Final = configured_realtime_query(self._params) return str(URL(f"{self.get_api_base(api_base).rstrip('/')}/realtime/calls", params=query)) + def transform_realtime_calls_response( + self, response: Response, model: str, model_id: str | None, headers: Mapping[str, object] | None + ) -> Response: + response.extensions["chatgpt_realtime"] = MappingProxyType( + { + "model": model, + "model_id": model_id, + "api_base": ChatGPTRealtime.get_api_base(self._params.api_base), + "extra_headers": configured_realtime_headers(headers), + "extra_query": configured_realtime_query(self._params), + } + ) + return response + def get_realtime_calls_headers( self, ephemeral_key: str ) -> dict[str, str]: # mutable-ok: HTTP handler header contract diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index ca141c0b958..a896594db41 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -38,6 +38,14 @@ class OpenAIRealtime(OpenAIChatCompletion): """ return "https://api.openai.com/" + def _resolve_api_key(self, api_key: str | None) -> str: + if api_key is None: + raise ValueError("api_key is required for OpenAI realtime calls") + return api_key + + def _accounts_for_call_usage(self) -> bool: + return True + def _get_additional_headers( self, api_key: str, @@ -125,8 +133,7 @@ class OpenAIRealtime(OpenAIChatCompletion): if api_base is None: api_base = self._get_default_api_base() - if api_key is None: - raise ValueError("api_key is required for OpenAI realtime calls") + resolved_api_key: Final = self._resolve_api_key(api_key) # Use all query params if provided, else fallback to just model if query_params is None: @@ -144,12 +151,12 @@ class OpenAIRealtime(OpenAIChatCompletion): "If your client expects beta event names, add 'OpenAI-Beta: realtime=v1' " "to the WebSocket headers sent to the LiteLLM proxy." ) - headers: Final = self._get_additional_headers(api_key, openai_beta_realtime=openai_beta_realtime) + headers: Final = self._get_additional_headers(resolved_api_key, openai_beta_realtime=openai_beta_realtime) # Log a masked request preview consistent with other endpoints. logging_obj.pre_call( input=None, - api_key=api_key, + api_key=resolved_api_key, additional_args={ "api_base": url, "headers": headers, @@ -173,7 +180,7 @@ class OpenAIRealtime(OpenAIChatCompletion): model if (query_params or {}).get("intent") == "transcription" else None ), event_normalizer=self._make_event_normalizer(), - account_usage=account_usage, + account_usage=account_usage and self._accounts_for_call_usage(), ) await realtime_streaming.bidirectional_forward() diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 819f0b8324a..5bf9ab1b8f3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -68,6 +68,16 @@ class _RealtimeAttachmentReservations(BaseModel): return owned +_RELEASE_REALTIME_COUNTER_LUA: Final = """ +local raw = redis.call('GET', KEYS[1]) +if not raw then return 0 end +local value = cjson.decode(raw) +value.current_requests = math.max(value.current_requests - 1, 0) +redis.call('SET', KEYS[1], cjson.encode(value), 'KEEPTTL') +return 1 +""" + + class _PROXY_MaxParallelRequestsHandler(CustomLogger): # Class variables or attributes def __init__(self, internal_usage_cache: InternalUsageCache): @@ -91,23 +101,23 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): 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) + await self._release_realtime_counter(key) - 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, + async def _release_realtime_counter(self, key: str) -> None: + local: Final = self.internal_usage_cache.dual_cache.in_memory_cache + remote: Final = self.internal_usage_cache.dual_cache.redis_cache + raw: Final[object] = local.get_cache(key) + current: Final = TypeAdapter(Mapping[str, int] | None).validate_python(raw) + updated: Final = ( + {**current, "current_requests": max(current["current_requests"] - 1, 0)} if current is not None else None ) + if updated is not None: + local.set_cache(key, updated, ttl=60) + if remote is not None: + release: Final = remote.async_register_script(_RELEASE_REALTIME_COUNTER_LUA) + await release(keys=(key,), args=()) + if local.get_cache(key) is updated: + local.delete_cache(key) def print_verbose(self, print_statement): try: diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 905a2f02fb5..67412ccfb96 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -8,10 +8,12 @@ import asyncio import binascii import os import uuid -from collections.abc import Awaitable, Callable, Mapping, Sequence, Set +from collections.abc import Awaitable, Callable, Generator, Mapping, Sequence, Set +from contextlib import contextmanager from contextvars import ContextVar from dataclasses import dataclass, field from datetime import datetime +from types import MappingProxyType from typing import ( TYPE_CHECKING, Any, @@ -53,6 +55,7 @@ from litellm.proxy.hooks.batch_enqueued_tokens import ( canonical_provider_batch_id, ) from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit +from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease, is_realtime_call_attachment from litellm.router_utils.add_retry_fallback_headers import ( ensure_response_additional_headers, response_has_hidden_params, @@ -304,6 +307,23 @@ end return results """ +PARALLEL_RENEW_SCRIPT: Final = """ +local clock = redis.call('TIME') +local now = tonumber(clock[1]) +local ttl = tonumber(ARGV[2]) +for i = 1, #KEYS do + local score = redis.call('ZSCORE', KEYS[i], ARGV[1]) + if not score or tonumber(score) <= now - ttl then + return {0} + end +end +for i = 1, #KEYS do + redis.call('ZADD', KEYS[i], 'XX', now, ARGV[1]) + redis.call('EXPIRE', KEYS[i], ttl) +end +return {1} +""" + TOKEN_INCREMENT_SCRIPT: Final = """ local results = {} @@ -416,6 +436,14 @@ class ParallelRequestGauge(TypedDict): descriptor_key: str +def _without_parallel_limit(descriptor: RateLimitDescriptor) -> RateLimitDescriptor: + rate_limit: Final[RateLimitDescriptorRateLimitObject] = { + **(descriptor.get("rate_limit") or MappingProxyType({})), + "max_parallel_requests": None, + } + return RateLimitDescriptor(key=descriptor["key"], value=descriptor["value"], rate_limit=rate_limit) + + class ParallelSlotAcquisition(TypedDict): slot_id: str counter_keys: list[str] @@ -546,6 +574,15 @@ def get_request_stash() -> RequestRateLimiterStash | None: return _request_stash.get() +@contextmanager +def isolated_request_stash() -> Generator[None]: + token: Final = _request_stash.set(None) + try: + yield + finally: + _request_stash.reset(token) + + def get_or_create_request_stash() -> RequestRateLimiterStash: stash = _request_stash.get() if stash is None: @@ -595,6 +632,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): parallel_acquire_script: _AsyncLuaScript | None parallel_release_script: _AsyncLuaScript | None parallel_count_script: _AsyncLuaScript | None + parallel_renew_script: _AsyncLuaScript | None def __init__( self, @@ -627,6 +665,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.parallel_count_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( PARALLEL_COUNT_SCRIPT ) + self.parallel_renew_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( + PARALLEL_RENEW_SCRIPT + ) else: self.batch_rate_limiter_script = None self.token_increment_script = None @@ -635,6 +676,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.parallel_acquire_script = None self.parallel_release_script = None self.parallel_count_script = None + self.parallel_renew_script = None self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60)) @@ -1598,6 +1640,66 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): statuses.append(self._gauge_status(gauge, in_flight + 1, "OK")) return RateLimitResponse(overall_code="OK", statuses=statuses) + def transfer_realtime_call_slot(self, request_data: Mapping[str, object]) -> RealtimeCallLease | None: + call_id: Final = request_data.get("litellm_call_id") + if not isinstance(call_id, str): + return None + stash: Final = get_request_stash_for_call(call_id) + if stash is None or stash.parallel_slot is None: + return None + slot_id: Final = stash.parallel_slot["slot_id"] + counter_keys: Final = tuple(stash.parallel_slot["counter_keys"]) + stash.parallel_slot = None + + async def renew() -> bool: + return await self._renew_realtime_call_slot(slot_id, counter_keys) + + async def release() -> None: + await self._release_parallel_request_slots( + ParallelSlotAcquisition(slot_id=slot_id, counter_keys=list(counter_keys)) + ) + + return RealtimeCallLease(renew=renew, release=release) + + async def _renew_realtime_call_slot(self, slot_id: str, counter_keys: tuple[str, ...]) -> bool: + if self.parallel_renew_script is not None: + try: + result: Final = await self.parallel_renew_script( + keys=counter_keys, args=(slot_id, PARALLEL_REQUEST_SLOT_TTL_SECONDS) + ) + return tuple(result) == (1,) + except Exception: # noqa: BLE001 # Redis ownership cannot be established by a local count mirror + return False + async with self._check_and_increment_lock: + now: Final = self._get_current_time().timestamp() + cutoff: Final = now - PARALLEL_REQUEST_SLOT_TTL_SECONDS + values: Final[tuple[ParallelGaugeCacheValue | None, ...]] = tuple( + [ + await self.internal_usage_cache.async_get_cache( + key=counter_key, local_only=True, litellm_parent_otel_span=None + ) + for counter_key in counter_keys + ] + ) + if any( + not isinstance(value, dict) + or not isinstance(score := value.get(slot_id), (int, float)) + or score <= cutoff + for value in values + ): + return False + for counter_key, value in zip(counter_keys, values): + if not isinstance(value, dict): + return False + await self.internal_usage_cache.async_set_cache( + key=counter_key, + value={**value, slot_id: now}, + ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, + local_only=True, + litellm_parent_otel_span=None, + ) + return True + async def _release_parallel_request_slots( self, acquisition: ParallelSlotAcquisition, @@ -3477,6 +3579,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Org Level Rate Limits descriptors.extend(self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model)) + effective_descriptors: Final = ( + tuple(_without_parallel_limit(descriptor) for descriptor in descriptors) + if call_type == "_arealtime" and is_realtime_call_attachment(data.get("websocket")) + else descriptors + ) + # Only check rate limits if we have descriptors with actual limits if descriptors: # First pass: RPM and max_parallel_requests sliding-window check. @@ -3495,16 +3603,18 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # double-charge every request. parallel_counter_keys: Final = [ self.create_rate_limit_keys(d["key"], d["value"], "max_parallel_requests") - for d in descriptors + for d in effective_descriptors if (d.get("rate_limit") or {}).get("max_parallel_requests") is not None ] parallel_slot_id: Final = uuid.uuid4().hex if parallel_counter_keys else None first_pass_descriptors: Final = ( - descriptors + effective_descriptors if self.tpm_reservation_enabled else tuple( - d for d in descriptors if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) + d + for d in effective_descriptors + if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) ) ) response: Final = await self.should_rate_limit( diff --git a/litellm/proxy/hooks/realtime_call_lease.py b/litellm/proxy/hooks/realtime_call_lease.py new file mode 100644 index 00000000000..c33d81f92d3 --- /dev/null +++ b/litellm/proxy/hooks/realtime_call_lease.py @@ -0,0 +1,74 @@ +import asyncio +from collections.abc import Awaitable, Callable, Generator +from contextlib import contextmanager +from contextvars import ContextVar +from typing import Final + +_realtime_call_attachment: Final[ContextVar[object | None]] = ContextVar("realtime_call_attachment", default=None) + + +@contextmanager +def realtime_call_attachment(websocket: object) -> Generator[None]: + token: Final = _realtime_call_attachment.set(websocket) + try: + yield + finally: + _realtime_call_attachment.reset(token) + + +def is_realtime_call_attachment(websocket: object) -> bool: + bound: Final = _realtime_call_attachment.get() + return bound is not None and bound is websocket + + +class RealtimeCallLease: + def __init__( + self, + *, + renew: Callable[[], Awaitable[bool]], + release: Callable[[], Awaitable[None]], + interval: float = 300, + renewal_timeout: float = 10, + ) -> None: + self._renew = renew + self._release = release + self._interval = interval + self._renewal_timeout = renewal_timeout + self._failed = asyncio.Event() + self._heartbeat: asyncio.Task[None] | None = None + self._closing: asyncio.Task[None] | None = None + + def start(self) -> None: + if self._heartbeat is None and self._closing is None: + self._heartbeat = asyncio.create_task(self._run()) + + async def renew(self) -> bool: + if self._closing is not None or self._failed.is_set(): + return False + try: + renewed: Final = await asyncio.wait_for(self._renew(), timeout=self._renewal_timeout) + except Exception: # noqa: BLE001 # fail closed without exposing cache credentials + self._failed.set() + return False + if not renewed: + self._failed.set() + return renewed and not self._failed.is_set() and self._closing is None + + async def wait_failed(self) -> None: + await self._failed.wait() + + async def _run(self) -> None: + while await self.renew(): + await asyncio.sleep(self._interval) + self._failed.set() + + async def close(self) -> None: + if self._closing is None: + self._closing = asyncio.create_task(self._close()) + await asyncio.shield(self._closing) + + async def _close(self) -> None: + if self._heartbeat is not None: + self._heartbeat.cancel() + await asyncio.gather(self._heartbeat, return_exceptions=True) + await self._release() diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index 6ed2e83ab1c..8ea329945c5 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -1,9 +1,10 @@ +import asyncio import base64 import hashlib import json import time from collections.abc import Awaitable, Callable, Mapping -from contextlib import AsyncExitStack +from contextlib import AsyncExitStack, nullcontext from contextvars import Token from types import MappingProxyType from typing import Final, Literal @@ -48,24 +49,38 @@ from litellm.proxy.hooks.parallel_request_limiter import ( ) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # existing built-in limiter has no public alias + isolated_request_stash, ) +from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease, realtime_call_attachment from litellm.proxy.spend_tracking.budget_reservation import ( invalidate_budget_reservation_counters, release_or_invalidate_budget_reservation, ) +from litellm.types.realtime import RealtimeQueryParams from litellm.types.router import GenericLiteLLMParams -async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth) -> None: +async def supervise_codex_call( + request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth, lease: RealtimeCallLease | None = None +) -> None: + with isolated_request_stash(): + await _start_codex_supervisor(request, call, auth, lease) + + +async def _start_codex_supervisor( + request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth, lease: RealtimeCallLease | None +) -> None: import litellm from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, CallSupervisor async def receive() -> Message: - return { + body: Final[RealtimeQueryParams] = {"model": call.alias} + message: Final[Message] = { "type": "http.request", - "body": json.dumps({"model": call.alias}).encode(), + "body": json.dumps(body).encode(), "more_body": False, - } # mutable-ok: ASGI message + } + return message async def send(_message: Message) -> None: return None @@ -143,6 +158,7 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth: close_call, force_close_call=force_close_call, terminal_usage_required=realtime_endpoint(call.model) == "live", + lease=lease, ) supervision_owned = True sockets.pop_all() @@ -241,6 +257,11 @@ async def process_codex_request( async def create_codex_realtime_call(request: Request) -> Response: + with isolated_request_stash(): + return await _create_codex_realtime_call(request) + + +async def _create_codex_realtime_call(request: Request) -> Response: from litellm.proxy import proxy_server as server try: @@ -275,6 +296,10 @@ async def create_codex_realtime_call(request: Request) -> Response: get_api_key_from_custom_header(request, custom_header) if isinstance(custom_header, str) else selected_key ) supervision_started = False # rebind-ok: transfer reservation ownership only after supervision is established + call_lease: RealtimeCallLease | None = None + lease_transferred = False # rebind-ok: failed startup leaves the signaling task responsible for its lease + preprocessing_started = False # rebind-ok: only refund reservations belonging to this signaling request + limiter: Final = server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter") try: await can_key_call_resolved_model( model=model, @@ -286,17 +311,30 @@ async def create_codex_realtime_call(request: Request) -> Response: signaling_auth: Final = auth.model_copy( update={"budget_reservation": None} ) # mutable-ok: Pydantic update contract + if isinstance(limiter, _PROXY_MaxParallelRequestsHandler) and ( + auth.max_parallel_requests is not None + or server.general_settings.get("global_max_parallel_requests") is not None + ): + raise HTTPException(400, "Realtime calls with parallel limits require the V3 rate limiter") + preprocessing_started = True processed, _ = await process_codex_request(request, data, signaling_auth, model, "arealtime_calls") - result: Final = await server.route_request( - data=processed, - route_type="arealtime_calls", - llm_router=server.llm_router, - user_model=server.user_model, - ) - try: - response: Final = await result - except BaseLLMException as exc: - raise HTTPException(exc.status_code, str(exc)) from exc + if isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + call_lease = limiter.transfer_realtime_call_slot(processed) + if call_lease is not None: + call_lease.start() + if not await call_lease.renew(): + raise HTTPException(503, "Realtime call quota reservation was lost") + with isolated_request_stash(): + result: Final = await server.route_request( + data=processed, + route_type="arealtime_calls", + llm_router=server.llm_router, + user_model=server.user_model, + ) + try: + response: Final = await result + except BaseLLMException as exc: + raise HTTPException(exc.status_code, str(exc)) from exc if not isinstance(response, httpx.Response): raise HTTPException(502, "Invalid realtime signaling response") if response.is_error: @@ -311,11 +349,15 @@ async def create_codex_realtime_call(request: Request) -> Response: except ValueError as exc: raise HTTPException(400, str(exc)) from exc supervised_call: Final = call.model_copy( - update={"usage_supervised": True} + update={"usage_supervised": True, "parallel_reserved": call_lease is not None} ) # mutable-ok: Pydantic update contract token: Final = encode_call(supervised_call) supervision_started = True - await supervise_codex_call(request, supervised_call, auth) + if call_lease is None: + await supervise_codex_call(request, supervised_call, auth) + else: + await supervise_codex_call(request, supervised_call, auth, call_lease) + lease_transferred = True return Response( response.content, status_code=response.status_code, @@ -323,8 +365,22 @@ async def create_codex_realtime_call(request: Request) -> Response: headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}), ) finally: - if not supervision_started: - await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) + try: + if call_lease is not None and not lease_transferred: + await call_lease.close() + finally: + try: + if preprocessing_started and isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + await asyncio.shield( + limiter.async_post_call_failure_hook( + request_data={}, # mutable-ok: existing failure-hook contract + original_exception=Exception("Realtime signaling completed without token usage"), + user_api_key_dict=auth, + ) + ) + finally: + if not supervision_started: + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAPIKeyAuth) -> None: @@ -385,7 +441,8 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP if isinstance(limiter, _PROXY_MaxParallelRequestsHandler): limiter.begin_realtime_attachment(data) try: - processed, logging_obj = await process_codex_request(request, data, auth, call.alias, "_arealtime") + with realtime_call_attachment(websocket) if call.parallel_reserved else nullcontext(): + processed, logging_obj = await process_codex_request(request, data, auth, call.alias, "_arealtime") except Exception: # noqa: BLE001 # custom hook exceptions must reject the connection verbose_proxy_logger.exception("Realtime sideband pre-call rejected") await websocket.close(code=1008, reason="Realtime pre-call rejected") diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index ddf1e8f3f49..8947edb8758 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -11,6 +11,7 @@ from litellm.constants import LOGGING_WORKER_MAX_TIME_PER_COROUTINE from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease from litellm.proxy.spend_tracking.budget_reservation import ( invalidate_budget_reservation_counters, release_or_invalidate_budget_reservation, @@ -57,6 +58,7 @@ class CallSupervisor: logging_timeout: float = LOGGING_WORKER_MAX_TIME_PER_COROUTINE, terminal_usage_required: bool = True, force_close_call: Callable[[], Awaitable[None]] | None = None, + lease: RealtimeCallLease | None = None, ) -> None: self._upstream = upstream self._stream = stream @@ -64,6 +66,7 @@ class CallSupervisor: self._auth = auth self._close_call = close_call self._force_close_call = force_close_call + self._lease = lease self._ready_timeout = ready_timeout self._lifetime = lifetime self._drain_timeout = drain_timeout @@ -85,6 +88,8 @@ class CallSupervisor: self._task = asyncio.create_task(self._run()) try: await asyncio.wait_for(self._ready.wait(), timeout=self._ready_timeout) + if self._lease is not None and not await self._lease.renew(): + raise RuntimeError("Call observer lost its quota reservation during startup") if not self._started or self._terminal or self._task.done(): raise RuntimeError("Call observer ended before session became available") except BaseException: @@ -129,10 +134,25 @@ class CallSupervisor: return self._terminal or self._close_confirmed async def _run(self) -> None: + try: + await self._observe() + finally: + try: + if self._lease is not None: + await self._lease.close() + finally: + self._ready.set() + + async def _observe(self) -> None: reader: Final = asyncio.create_task(self._read()) stopped: Final = asyncio.create_task(self._stop.wait()) + lease_failed: Final = asyncio.create_task(self._lease.wait_failed()) if self._lease is not None else None try: - await asyncio.wait((reader, stopped), timeout=self._lifetime, return_when=asyncio.FIRST_COMPLETED) + await asyncio.wait( + (reader, stopped, lease_failed) if lease_failed is not None else (reader, stopped), + timeout=self._lifetime, + return_when=asyncio.FIRST_COMPLETED, + ) finally: try: if not self._terminal: @@ -161,7 +181,12 @@ class CallSupervisor: finally: stopped.cancel() reader.cancel() - await asyncio.gather(reader, stopped, return_exceptions=True) + if lease_failed is not None: + lease_failed.cancel() + await asyncio.gather( + *((reader, stopped, lease_failed) if lease_failed is not None else (reader, stopped)), + return_exceptions=True, + ) with suppress(Exception): await self._upstream.close() if not self._usage_complete(): diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 23963475ed0..0b8a32dcfc1 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -76,7 +76,7 @@ def _get_realtime_http_provider_config( dynamic_api_base: str | None, dynamic_api_key: str | None, litellm_params: GenericLiteLLMParams, - use_codex_backend: bool = False, + is_call: bool = False, ) -> tuple["BaseRealtimeHTTPConfig | None", str, str]: """ Return (provider_config, resolved_api_base, resolved_api_key) for the @@ -90,23 +90,19 @@ def _get_realtime_http_provider_config( ) provider_config: BaseRealtimeHTTPConfig | None = None - if custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig - - provider_config = ChatGPTRealtimeHTTPConfig(litellm_params, use_codex_backend=use_codex_backend) - elif custom_llm_provider in LlmProviders._member_map_.values(): + if custom_llm_provider in LlmProviders._member_map_.values(): provider_config = ProviderConfigManager.get_provider_realtime_http_config( model="", provider=LlmProviders(custom_llm_provider), + params=litellm_params, + is_call=is_call, ) - raw_api_base: Final = ( - litellm_params.api_base if custom_llm_provider == "chatgpt" else dynamic_api_base or litellm_params.api_base - ) + raw_api_base: Final = dynamic_api_base or litellm_params.api_base raw_api_key: Final = dynamic_api_key or litellm_params.api_key if provider_config is not None: - resolved_api_base = provider_config.get_api_base(api_base=raw_api_base) + resolved_api_base = provider_config.resolve_api_base(litellm_params.api_base, dynamic_api_base) resolved_api_key = provider_config.get_api_key(api_key=raw_api_key) else: # Fallback for providers without a dedicated HTTP config (treated as OpenAI-compatible). @@ -260,8 +256,6 @@ async def arealtime_calls( timeout: float | None = None, **kwargs, ): - from litellm.llms.chatgpt.realtime import realtime_call_headers - model_name = model or "gpt-4o-realtime-preview" litellm_logging_obj: Final[LiteLLMLogging] = kwargs.get("litellm_logging_obj") litellm_params: Final = GenericLiteLLMParams(**kwargs) @@ -281,12 +275,14 @@ async def arealtime_calls( dynamic_api_base=dynamic_api_base, dynamic_api_key=dynamic_api_key, litellm_params=litellm_params, - use_codex_backend=True, + is_call=True, ) if session is not None: session = _with_resolved_session_model(session, model_name) call_headers: Final = ( - realtime_call_headers(litellm_params) if custom_llm_provider == "chatgpt" else kwargs.get("extra_headers") + provider_config.get_realtime_calls_extra_headers(kwargs.get("extra_headers")) + if provider_config is not None + else kwargs.get("extra_headers") ) litellm_logging_obj.update_from_kwargs( kwargs=kwargs, @@ -308,23 +304,13 @@ async def arealtime_calls( client=kwargs.get("client"), api_version=litellm_params.api_version, ) - if custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ( - ChatGPTRealtime, - configured_realtime_headers, - configured_realtime_query, + return ( + provider_config.transform_realtime_calls_response( + response, model_name, litellm_logging_obj.get_router_model_id(), call_headers ) - - response.extensions["chatgpt_realtime"] = MappingProxyType( - { - "model": model_name, - "model_id": litellm_logging_obj.get_router_model_id(), - "api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base), - "extra_headers": configured_realtime_headers(call_headers), - "extra_query": configured_realtime_query(litellm_params), - } - ) - return response + if provider_config is not None + else response + ) async def vertex_access_token_resolver( @@ -421,7 +407,26 @@ async def _arealtime( model=model, provider=LlmProviders(_custom_llm_provider), ) - if provider_config is not None: + provider_handler: Final = ( + ProviderConfigManager.get_provider_realtime_handler( + LlmProviders(_custom_llm_provider), litellm_params, websocket.headers, headers + ) + if _custom_llm_provider in LlmProviders._member_map_.values() + else None + ) + if provider_handler is not None: + await provider_handler.async_realtime( + model=model, + websocket=websocket, + logging_obj=litellm_logging_obj, + api_base=api_base or None, + api_key=api_key, + timeout=timeout, + query_params=query_params, + user_api_key_dict=kwargs.get("user_api_key_dict"), + litellm_metadata=_build_litellm_metadata(kwargs), + ) + elif provider_config is not None: await base_llm_http_handler.async_realtime( model=model, websocket=websocket, @@ -469,21 +474,6 @@ async def _arealtime( user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata(kwargs), ) - elif _custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ChatGPTRealtime, accounts_for_call_usage - - await ChatGPTRealtime(litellm_params, websocket.headers, headers).async_realtime( - model=model, - websocket=websocket, - logging_obj=litellm_logging_obj, - api_base=ChatGPTRealtime.get_api_base(api_base), - api_key="chatgpt-oauth", - timeout=timeout, - query_params=query_params, - user_api_key_dict=kwargs.get("user_api_key_dict"), - litellm_metadata=_build_litellm_metadata(kwargs), - account_usage=accounts_for_call_usage(litellm_params), - ) elif _custom_llm_provider == "openai": api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or "https://api.openai.com/" # set API KEY diff --git a/litellm/utils.py b/litellm/utils.py index 5aab0210d4a..b4f1ce8afc2 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -419,6 +419,7 @@ if TYPE_CHECKING: from litellm.llms.cohere.common_utils import CohereModelInfo from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.mistral.ocr.transformation import MistralOCRConfig + from litellm.llms.openai.realtime.handler import OpenAIRealtime from litellm.proxy._types import AllowedModelRegion from litellm.router_utils.get_retry_from_policy import ( get_num_retries_from_retry_policy, @@ -434,7 +435,7 @@ if TYPE_CHECKING: ChatCompletionToolCallFunctionChunk, ) from litellm.types.rerank import RerankResponse - from litellm.types.router import LiteLLM_Params + from litellm.types.router import GenericLiteLLMParams, LiteLLM_Params from litellm.llms.base_llm.chat.transformation import BaseConfig from litellm.llms.base_llm.completion.transformation import BaseTextCompletionConfig @@ -9211,16 +9212,36 @@ class ProviderConfigManager: return GeminiRealtimeConfig() return None + @staticmethod + def get_provider_realtime_handler( + provider: LlmProviders, + params: GenericLiteLLMParams, + headers: Mapping[str, str], + extra_headers: Mapping[str, object] | None = None, + ) -> OpenAIRealtime | None: + if provider == LlmProviders.CHATGPT: + from litellm.llms.chatgpt.realtime import ChatGPTRealtime + + return ChatGPTRealtime(params, headers, extra_headers) + return None + @staticmethod def get_provider_realtime_http_config( model: str, provider: LlmProviders, + params: GenericLiteLLMParams | None = None, + is_call: bool = False, ) -> BaseRealtimeHTTPConfig | None: """ Return the HTTP transformation config for realtime HTTP endpoints (POST /realtime/client_secrets and POST /realtime/calls). """ + if LlmProviders.CHATGPT == provider: + from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig + from litellm.types.router import GenericLiteLLMParams + + return ChatGPTRealtimeHTTPConfig(params or GenericLiteLLMParams(), use_codex_backend=is_call) if LlmProviders.OPENAI == provider: from litellm.llms.openai.realtime.http_transformation import ( OpenAIRealtimeHTTPConfig, 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 0bf488016c2..3ee364a57a4 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -2,11 +2,17 @@ Unit Tests for the max parallel request limiter v1 for the proxy """ +import asyncio +import shutil +import socket +import subprocess from datetime import datetime +from unittest.mock import AsyncMock, MagicMock import pytest from litellm.caching.caching import DualCache +from litellm.caching.redis_cache import RedisCache 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 ( @@ -16,6 +22,105 @@ from litellm.proxy.utils import InternalUsageCache, hash_token from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage +@pytest.fixture +def isolated_legacy_redis(tmp_path): + executable = shutil.which("redis-server") + if executable is None: + pytest.skip("redis-server is required to exercise atomic Lua updates") + with socket.socket() as listener: + listener.bind(("127.0.0.1", 0)) + port = listener.getsockname()[1] + process = subprocess.Popen( + [ + executable, + "--bind", + "127.0.0.1", + "--port", + str(port), + "--save", + "", + "--appendonly", + "no", + "--dir", + str(tmp_path), + ], + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + try: + import redis + + client = redis.Redis(host="127.0.0.1", port=port) + for _ in range(100): + try: + client.ping() + break + except redis.ConnectionError: + import time + + time.sleep(0.01) + else: + pytest.fail("isolated Redis did not start") + yield port + client.close() + finally: + process.terminate() + process.wait(timeout=5) + + +@pytest.mark.asyncio +async def test_concurrent_realtime_releases_update_redis_without_lost_decrement(isolated_legacy_redis): + remote = RedisCache(host="127.0.0.1", port=isolated_legacy_redis, namespace="legacy-test") + first_cache, second_cache = DualCache(redis_cache=remote), DualCache(redis_cache=remote) + first, second = (_PROXY_MaxParallelRequestsHandler(InternalUsageCache(c)) for c in (first_cache, second_cache)) + auth = UserAPIKeyAuth(api_key="concurrent-key", max_parallel_requests=2) + first_data, second_data = {"model": "test"}, {"model": "test"} + first.begin_realtime_attachment(first_data) + second.begin_realtime_attachment(second_data) + await first.async_pre_call_hook(auth, first_cache, first_data, "_arealtime") + await second.async_pre_call_hook(auth, second_cache, second_data, "_arealtime") + key = f"concurrent-key::{datetime.now().strftime('%Y-%m-%d-%H-%M')}::request_count" + counter = {"current_requests": 2, "current_rpm": 2, "current_tpm": 17} + await first_cache.async_set_cache(key, counter) + await second_cache.async_set_cache(key, counter, local_only=True) + remote.redis_client.pexpire(remote.check_and_fix_namespace(key), 15000) + await asyncio.gather( + first.async_release_realtime_attachment(first_data, auth), + second.async_release_realtime_attachment(second_data, auth), + ) + expected = {"current_requests": 0, "current_rpm": 2, "current_tpm": 17} + assert await remote.async_get_cache(key) == expected + assert 0 < remote.redis_client.pttl(remote.check_and_fix_namespace(key)) <= 15000 + assert await first_cache.async_get_cache(key) == expected + assert await second_cache.async_get_cache(key) == expected + await first_cache.async_set_cache("missing", counter, local_only=True) + await first._release_realtime_counter("missing") + assert await remote.async_get_cache("missing") is None + assert await first_cache.async_get_cache("missing", local_only=True) is None + + +@pytest.mark.asyncio +async def test_realtime_release_preserves_newer_local_admission_while_redis_finishes(): + started, finish = asyncio.Event(), asyncio.Event() + + async def release(**kwargs): + started.set() + await finish.wait() + + remote = MagicMock(spec=RedisCache) + remote.async_register_script.return_value = AsyncMock(side_effect=release) + cache = DualCache(redis_cache=remote) + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache)) + await cache.async_set_cache("key", {"current_requests": 1, "current_rpm": 1, "current_tpm": 7}, local_only=True) + task = asyncio.create_task(handler._release_realtime_counter("key")) + await started.wait() + next_admission = {"current_requests": 1, "current_rpm": 2, "current_tpm": 7} + await cache.async_set_cache("key", next_admission, local_only=True) + finish.set() + await task + assert await cache.async_get_cache("key", local_only=True) == next_admission + + @pytest.mark.asyncio @pytest.mark.parametrize("reject_team", [False, True]) async def test_realtime_attachment_releases_only_acquired_legacy_slots(reject_team): diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 6d382370f5f..181b3d8901e 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -31,6 +31,8 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, ) from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token +from litellm.proxy.hooks.parallel_request_limiter_v3 import isolated_request_stash +from litellm.proxy.hooks.realtime_call_lease import realtime_call_attachment from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.utils import ( EmbeddingResponse, @@ -51,6 +53,133 @@ class TimeController: self._current += timedelta(seconds=seconds) +@pytest.mark.asyncio +async def test_realtime_lease_retains_quota_across_signaling_and_three_attachments(): + cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache)) + auth = UserAPIKeyAuth(api_key="logical-owner", max_parallel_requests=1, rpm_limit=4, tpm_limit=100000) + data = {"model": "gpt-3.5-turbo", "litellm_call_id": "signaling"} + await handler.async_pre_call_hook(auth, cache, data, "arealtime_calls") + stash = get_request_stash() + assert stash.reserved_tokens > 0 + assert handler.transfer_realtime_call_slot({"litellm_call_id": "other-call"}) is None + assert stash.parallel_slot is not None + lease = handler.transfer_realtime_call_slot(data) + assert lease is not None + assert stash.parallel_slot is None + assert stash.reserved_tokens > 0 + assert not stash.reservation_released + assert handler.transfer_realtime_call_slot(data) is None + await handler.async_log_success_event( + kwargs={"litellm_call_id": "signaling", "standard_logging_object": {"metadata": {"user_api_key_hash": auth.api_key}}}, + response_obj=ModelResponse(usage=Usage()), start_time=datetime.now(), end_time=datetime.now(), + ) + assert await cache.async_get_cache("{api_key:logical-owner}:tokens") == 0 + socket = object() + for attachment in range(3): + with isolated_request_stash(), realtime_call_attachment(socket): + attachment_data = {"model": "gpt-3.5-turbo", "litellm_call_id": f"attachment-{attachment}", "websocket": socket} + await handler.async_pre_call_hook(auth, cache, attachment_data, "_arealtime") + assert get_request_stash().parallel_slot is None + await handler.async_release_realtime_attachment(attachment_data, auth) + with isolated_request_stash(), realtime_call_attachment(socket), pytest.raises(HTTPException) as error: + await handler.async_pre_call_hook(auth, cache, {"model": "gpt-3.5-turbo", "websocket": socket}, "_arealtime") + assert error.value.status_code == 429 + assert "requests" in str(error.value.detail) + quota_only = auth.model_copy(update={"rpm_limit": None}) + with isolated_request_stash(), realtime_call_attachment(object()), pytest.raises(HTTPException) as error: + await handler.async_pre_call_hook( + quota_only, cache, {"model": "gpt-3.5-turbo", "websocket": socket}, "_arealtime" + ) + assert "max_parallel_requests" in str(error.value.detail) + with isolated_request_stash(), realtime_call_attachment(socket), pytest.raises(HTTPException) as error: + await handler.async_pre_call_hook( + quota_only, cache, {"model": "gpt-3.5-turbo", "websocket": socket}, "acompletion" + ) + assert "max_parallel_requests" in str(error.value.detail) + with isolated_request_stash(), pytest.raises(HTTPException) as error: + await handler.async_pre_call_hook(quota_only, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") + assert "max_parallel_requests" in str(error.value.detail) + assert get_request_stash() is stash + await lease.close() + with isolated_request_stash(): + await handler.async_pre_call_hook(quota_only, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") + + +@pytest.mark.asyncio +async def test_realtime_lease_renewal_preserves_quota_past_ttl_and_does_not_resurrect_expiry(): + cache = DualCache() + clock = TimeController() + handler = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(cache), time_provider=clock.now) + auth = UserAPIKeyAuth(api_key="long-call", max_parallel_requests=1) + data = {"model": "gpt-3.5-turbo", "litellm_call_id": "long-call"} + await handler.async_pre_call_hook(auth, cache, data, "arealtime_calls") + lease = handler.transfer_realtime_call_slot(data) + assert lease is not None + clock.advance(PARALLEL_REQUEST_SLOT_TTL_SECONDS - 1) + assert await lease.renew() + clock.advance(2) + with isolated_request_stash(), pytest.raises(HTTPException): + await handler.async_pre_call_hook(auth, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") + + clock.advance(PARALLEL_REQUEST_SLOT_TTL_SECONDS) + assert not await lease.renew() + await asyncio.wait_for(lease.wait_failed(), 1) + with isolated_request_stash(): + await handler.async_pre_call_hook(auth, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") + await lease.close() + with isolated_request_stash(), pytest.raises(HTTPException): + await handler.async_pre_call_hook(auth, cache, {"model": "gpt-3.5-turbo"}, "arealtime_calls") + + +@pytest.mark.asyncio +async def test_realtime_lease_redis_renewal_is_atomic_and_does_not_resurrect(): + import shutil + import subprocess + import tempfile + from redis.asyncio import Redis + from redis.exceptions import ConnectionError as RedisConnectionError + from litellm.proxy.hooks.parallel_request_limiter_v3 import PARALLEL_RENEW_SCRIPT + + executable = shutil.which("redis-server") + if executable is None: + pytest.skip("redis-server is required for the Lua regression") + with tempfile.TemporaryDirectory(prefix="rtc-") as temporary: + socket = f"{temporary}/redis.sock" + process = subprocess.Popen( + [executable, "--port", "0", "--unixsocket", socket, "--save", "", "--appendonly", "no"], + stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) + client = Redis(unix_socket_path=socket) + try: + for attempt in range(100): + try: + await client.ping() + break + except RedisConnectionError: + await asyncio.sleep(0.01) + else: + pytest.fail("isolated Redis did not start") + now = (await client.time())[0] + await client.zadd("first", {"owner": now - 10, "other": now}) + await client.zadd("second", {"owner": now - PARALLEL_REQUEST_SLOT_TTL_SECONDS}) + renew = client.register_script(PARALLEL_RENEW_SCRIPT) + assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] + assert await client.zscore("first", "owner") == now - 10 + await client.zadd("second", {"owner": now - 10}) + assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [1] + assert await client.zscore("first", "owner") >= now + assert await client.ttl("first") > PARALLEL_REQUEST_SLOT_TTL_SECONDS - 10 + await client.zrem("second", "owner") + assert await renew(keys=["first", "second"], args=["owner", PARALLEL_REQUEST_SLOT_TTL_SECONDS]) == [0] + assert await client.zscore("second", "owner") is None + assert await client.zscore("first", "other") == now + finally: + await client.aclose() + process.terminate() + process.wait(timeout=5) + + @pytest.fixture def time_controller(monkeypatch): controller = TimeController() diff --git a/tests/test_litellm/proxy/hooks/test_realtime_call_lease.py b/tests/test_litellm/proxy/hooks/test_realtime_call_lease.py new file mode 100644 index 00000000000..cb8ee7fc6b9 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_realtime_call_lease.py @@ -0,0 +1,85 @@ +import asyncio +from unittest.mock import AsyncMock + +import pytest + +from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease + + +@pytest.mark.asyncio +async def test_failed_renewal_signals_owner_and_close_releases_once(): + renew = AsyncMock(side_effect=[True, False, True]) + release = AsyncMock() + lease = RealtimeCallLease(renew=renew, release=release, interval=0.001) + lease.start() + await asyncio.wait_for(lease.wait_failed(), timeout=1) + assert renew.await_count == 2 + assert not await lease.renew() + assert renew.await_count == 2 + await asyncio.gather(lease.close(), lease.close()) + assert release.await_count == 1 + + +@pytest.mark.asyncio +async def test_renewal_exception_and_close_before_start(): + release = AsyncMock() + lease = RealtimeCallLease(renew=AsyncMock(side_effect=RuntimeError("backend")), release=release, interval=0.001) + lease.start() + await asyncio.wait_for(lease.wait_failed(), timeout=1) + await lease.close() + assert release.await_count == 1 + unused = RealtimeCallLease(renew=AsyncMock(), release=release) + await unused.close() + assert release.await_count == 2 + + +@pytest.mark.asyncio +async def test_renewal_timeout_signals_failure_without_start(): + lease = RealtimeCallLease(renew=asyncio.Event().wait, release=AsyncMock(), renewal_timeout=0.001) + assert not await lease.renew() + await asyncio.wait_for(lease.wait_failed(), timeout=1) + await lease.close() + + +@pytest.mark.asyncio +async def test_cancelled_close_still_releases_exactly_once(): + entered = asyncio.Event() + finish = asyncio.Event() + + async def release(): + entered.set() + await finish.wait() + + cleanup = AsyncMock(side_effect=release) + lease = RealtimeCallLease(renew=AsyncMock(return_value=True), release=cleanup) + lease.start() + closing = asyncio.create_task(lease.close()) + await asyncio.wait_for(entered.wait(), timeout=1) + closing.cancel() + with pytest.raises(asyncio.CancelledError): + await closing + finish.set() + await lease.close() + assert cleanup.await_count == 1 + + +@pytest.mark.asyncio +async def test_concurrent_renewal_cannot_restore_a_failed_lease(): + pending = asyncio.Event() + entered = asyncio.Event() + + async def delayed_success(): + entered.set() + await pending.wait() + return True + + renew = AsyncMock(side_effect=delayed_success) + lease = RealtimeCallLease(renew=renew, release=AsyncMock()) + first = asyncio.create_task(lease.renew()) + await asyncio.wait_for(entered.wait(), timeout=1) + renew.side_effect = None + renew.return_value = False + assert not await lease.renew() + pending.set() + assert not await first + await lease.close() 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 2591d45f9a5..75e1f64a65b 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -180,6 +180,7 @@ async def test_supervised_attachments_release_real_limiter_before_reconnect(monk 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( @@ -755,3 +756,276 @@ async def test_attachment_releases_quota_before_upstream_close_handshake(monkeyp "{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 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", [None, "provider", "observer", "renewal", "legacy_key", "legacy_global"]) +async def test_signaling_keeps_or_releases_owned_call_lease(monkeypatch, failure): + import json + from unittest.mock import AsyncMock, MagicMock + + import httpx + from fastapi import Request + from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease + + from litellm.proxy import proxy_server as server + from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + + auth = UserAPIKeyAuth(max_parallel_requests=None if failure == "legacy_global" else 1) + lease = MagicMock(spec=RealtimeCallLease) + lease.renew = AsyncMock(return_value=failure != "renewal") + lease.close = AsyncMock() + legacy = failure in ("legacy_key", "legacy_global") + limiter = MagicMock(spec=_PROXY_MaxParallelRequestsHandler if legacy else _PROXY_MaxParallelRequestsHandler_v3) + if not legacy: + limiter.transfer_realtime_call_slot.return_value = lease + proxy = MagicMock() + proxy.get_proxy_hook.return_value = limiter + monkeypatch.setattr(server, "proxy_logging_obj", proxy) + monkeypatch.setattr( + server, "general_settings", {"global_max_parallel_requests": 1} if failure == "legacy_global" else {} + ) + monkeypatch.setattr(codex, "user_api_key_auth", AsyncMock(return_value=auth)) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + process = AsyncMock(return_value=({}, None)) + monkeypatch.setattr(codex, "process_codex_request", process) + monkeypatch.setenv("LITELLM_SALT_KEY", "lease-transfer-test") + + async def route(**kwargs): + lease.start.assert_called_once() + if failure == "provider": + raise RuntimeError("Provider unavailable") + + async def respond(): + return httpx.Response( + 201, + text="v=0\r\n", + headers={"Location": "/v1/realtime/calls/rtc_lease"}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}, + ) + + return respond() + + async def supervise(request, call, owner, selected_lease): + assert owner is auth + assert selected_lease is lease + assert call.parallel_reserved + if failure == "observer": + raise RuntimeError("Observer unavailable") + + monkeypatch.setattr(server, "route_request", route) + monkeypatch.setattr(codex, "supervise_codex_call", supervise) + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/realtime/calls", + "query_string": b"", + "headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer owner")], + }, + AsyncMock( + return_value={ + "type": "http.request", + "body": json.dumps({"sdp": "v=0", "session": {"model": "voice"}}).encode(), + } + ), + ) + if failure is None: + response = await codex.create_codex_realtime_call(request) + token = response.headers["location"].rsplit("/", 1)[-1] + assert codex.decode_call(token, "Bearer owner").parallel_reserved + lease.close.assert_not_awaited() + else: + with pytest.raises((RuntimeError, HTTPException)) as raised: + await codex.create_codex_realtime_call(request) + if legacy: + assert raised.value.status_code == 400 + assert "V3 rate limiter" in raised.value.detail + process.assert_not_awaited() + lease.close.assert_not_awaited() + else: + if failure == "renewal": + assert raised.value.status_code == 503 + lease.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_signaling_rejection_after_admission_refunds_parallel_slot(monkeypatch): + import json + from unittest.mock import AsyncMock + + from fastapi import Request + + import litellm + from litellm.integrations.custom_logger import CustomLogger + from litellm.proxy import proxy_server as server + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.utils import ProxyLogging + + proxy = ProxyLogging(UserApiKeyCache()) + monkeypatch.setattr(litellm, "callbacks", []) + proxy._add_proxy_hooks() + limiter = proxy.get_proxy_hook("parallel_request_limiter") + key = "{api_key:rejected-signaling-owner}:max_parallel_requests" + + class Reject(CustomLogger): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + current = await proxy.internal_usage_cache.async_get_cache(key, litellm_parent_otel_span=None, local_only=True) + assert limiter._gauge_in_flight_from_cache_value(current) == 1 + raise RuntimeError("Policy rejected after admission") + + litellm.callbacks.append(Reject()) + monkeypatch.setattr(server, "proxy_logging_obj", proxy) + monkeypatch.setattr(server, "general_settings", {}) + monkeypatch.setattr( + server, + "llm_router", + litellm.Router( + model_list=[ + {"model_name": "voice", "litellm_params": {"model": "openai/gpt-realtime-1.5", "api_key": "test"}} + ] + ), + ) + auth = UserAPIKeyAuth(api_key="rejected-signaling-owner", max_parallel_requests=1) + monkeypatch.setattr(codex, "user_api_key_auth", AsyncMock(return_value=auth)) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + route = AsyncMock() + monkeypatch.setattr(server, "route_request", route) + for _ in range(2): + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/realtime/calls", + "query_string": b"", + "headers": [(b"content-type", b"application/json")], + }, + AsyncMock( + return_value={ + "type": "http.request", + "body": json.dumps({"sdp": "v=0", "session": {"model": "voice"}}).encode(), + } + ), + ) + with pytest.raises(RuntimeError, match="Policy rejected after admission"): + await codex.create_codex_realtime_call(request) + current = await proxy.internal_usage_cache.async_get_cache(key, litellm_parent_otel_span=None, local_only=True) + assert limiter._gauge_in_flight_from_cache_value(current) == 0 + route.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("callback_order", ["before", "after", "cancel"]) +async def test_signaling_settles_tokens_once_with_isolated_sdk_callbacks(monkeypatch, callback_order): + import asyncio + import json + from datetime import datetime + from unittest.mock import AsyncMock + + import httpx + import litellm + from fastapi import Request + 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 get_request_stash, isolated_request_stash + from litellm.proxy.utils import ProxyLogging + + proxy = ProxyLogging(UserApiKeyCache()) + monkeypatch.setattr(litellm, "callbacks", []) + proxy._add_proxy_hooks() + limiter = proxy.get_proxy_hook("parallel_request_limiter") + monkeypatch.setattr(server, "proxy_logging_obj", proxy) + monkeypatch.setattr(server, "general_settings", {}) + monkeypatch.setattr( + server, + "llm_router", + litellm.Router( + model_list=[ + {"model_name": "voice", "litellm_params": {"model": "openai/gpt-realtime-1.5", "api_key": "test"}} + ] + ), + ) + auth = UserAPIKeyAuth(api_key="signaling-settlement-owner", max_parallel_requests=1, tpm_limit=10000) + monkeypatch.setattr(codex, "user_api_key_auth", AsyncMock(return_value=auth)) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + supervisor = AsyncMock() + monkeypatch.setattr(codex, "supervise_codex_call", supervisor) + monkeypatch.setenv("LITELLM_SALT_KEY", "signaling-settlement-test") + ready, release_callback = asyncio.Event(), asyncio.Event() + callbacks = [] + + async def counter(kind): + return await proxy.internal_usage_cache.async_get_cache( + f"{{api_key:{auth.api_key}}}:{kind}", litellm_parent_otel_span=None, local_only=True + ) + + async def route(**kwargs): + assert get_request_stash() is None + assert await counter("tokens") > 0 + + async def callback(): + await release_callback.wait() + assert get_request_stash() is None + await limiter.async_log_success_event( + kwargs={ + "litellm_call_id": kwargs["data"]["litellm_call_id"], + "standard_logging_object": {"metadata": {"user_api_key_hash": auth.api_key}}, + }, + response_obj=litellm.ModelResponse(usage=litellm.Usage()), + start_time=datetime.now(), + end_time=datetime.now(), + ) + + async def respond(): + assert get_request_stash() is None + ready.set() + if callback_order == "cancel": + await asyncio.Event().wait() + callbacks.append(asyncio.create_task(callback())) + if callback_order == "before": + release_callback.set() + await callbacks[0] + assert await counter("tokens") > 0 + return httpx.Response( + 201, + text="v=0\r\n", + headers={"Location": "/v1/realtime/calls/rtc_settlement"}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}, + ) + + return respond() + + monkeypatch.setattr(server, "route_request", route) + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/realtime/calls", + "query_string": b"", + "headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer owner")], + }, + AsyncMock( + return_value={ + "type": "http.request", + "body": json.dumps({"sdp": "v=0", "session": {"model": "voice"}}).encode(), + } + ), + ) + with isolated_request_stash(): + signaling = asyncio.create_task(codex.create_codex_realtime_call(request)) + await asyncio.wait_for(ready.wait(), timeout=2) + if callback_order == "cancel": + signaling.cancel() + with pytest.raises(asyncio.CancelledError): + await signaling + supervisor.assert_not_awaited() + else: + assert (await signaling).status_code == 201 + assert limiter._gauge_in_flight_from_cache_value(await counter("max_parallel_requests")) == 1 + await supervisor.call_args.args[3].close() + assert await counter("tokens") == 0 + release_callback.set() + await asyncio.gather(*callbacks) + assert await counter("tokens") == 0 + assert limiter._gauge_in_flight_from_cache_value(await counter("max_parallel_requests")) == 0 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 93b97a28fe6..3377f63433f 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -30,6 +30,42 @@ class Socket: self.closed = True +@pytest.mark.asyncio +@pytest.mark.parametrize("lease_lost", [False, True]) +async def test_supervisor_holds_call_lease_until_terminal_accounting(lease_lost): + from litellm.proxy.hooks.realtime_call_lease import RealtimeCallLease + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + lost = asyncio.Event() + lease = MagicMock(spec=RealtimeCallLease) + lease.wait_failed = lost.wait + + async def release(): + assert socket.closed + assert sink.logs == 1 + + lease.close = AsyncMock(side_effect=release) + + async def close(): + await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}}) + + terminate = AsyncMock(side_effect=close) + supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), terminate, lease=lease) + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + lease.close.assert_not_awaited() + if lease_lost: + lost.set() + else: + await close() + await asyncio.wait_for(supervisor.wait(), 1) + assert terminate.await_count == int(lease_lost) + lease.close.assert_awaited_once() + + @pytest.mark.asyncio @pytest.mark.parametrize("stalled_step", ["close", "drain"]) async def test_live_initial_close_reserves_time_for_independent_hangup(stalled_step): diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index 761e87ac764..3ab843dcea7 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -331,3 +331,29 @@ async def test_arealtime_azure_ai_on_a_foundry_host_connects_to_the_azure_openai "wss://my-project.services.ai.azure.com/openai/realtime" "?api-version=2024-10-01-preview&deployment=gpt-realtime-mini" ) + + +@pytest.mark.parametrize("is_call", [False, True]) +@pytest.mark.parametrize("provider", ["chatgpt", "openai", "azure"]) +def test_realtime_http_provider_controls_dynamic_base_precedence(provider, is_call, monkeypatch): + from litellm.types.router import GenericLiteLLMParams + + monkeypatch.delenv("CHATGPT_API_BASE", raising=False) + monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False) + config, base, key = realtime_main._get_realtime_http_provider_config( + custom_llm_provider=provider, + dynamic_api_base="https://dynamic.example/v1", + dynamic_api_key="dynamic-key", + litellm_params=GenericLiteLLMParams(api_base="https://configured.example/v1"), + is_call=is_call, + ) + expected_base = "https://configured.example/v1" if provider == "chatgpt" else "https://dynamic.example/v1" + assert base == expected_base + assert key == ("chatgpt-oauth" if provider == "chatgpt" else "dynamic-key") + assert config is not None + if provider == "chatgpt": + assert config.get_realtime_calls_url(base, "gpt-realtime-1.5") == expected_base + "/realtime/calls" + else: + assert config.get_realtime_calls_extra_headers({"x-gateway-route": "required"}) == { + "x-gateway-route": "required" + }