mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(chatgpt): retain realtime call quotas and isolate provider routing
This commit is contained in:
parent
2642d98302
commit
34dcfbf833
17 changed files with 1090 additions and 94 deletions
|
|
@ -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 #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
74
litellm/proxy/hooks/realtime_call_lease.py
Normal file
74
litellm/proxy/hooks/realtime_call_lease.py
Normal 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()
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
85
tests/test_litellm/proxy/hooks/test_realtime_call_lease.py
Normal file
85
tests/test_litellm/proxy/hooks/test_realtime_call_lease.py
Normal 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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue