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