fix(chatgpt): retain realtime call quotas and isolate provider routing

This commit is contained in:
jibanez-staticduo 2026-09-11 02:25:21 +02:00
parent 2642d98302
commit 34dcfbf833
No known key found for this signature in database
17 changed files with 1090 additions and 94 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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