fix(chatgpt): release sideband quotas and preserve repeated queries

This commit is contained in:
jibanez-staticduo 2026-09-10 19:11:03 +02:00
parent 8ec259bdbf
commit 09e6b3a89d
No known key found for this signature in database
11 changed files with 389 additions and 30 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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