fix(chatgpt): preserve observer quotas and bound call cleanup

This commit is contained in:
jibanez-staticduo 2026-09-10 15:54:39 +02:00
parent bac2ad4abb
commit c360d187f5
No known key found for this signature in database
11 changed files with 464 additions and 46 deletions

View file

@ -3,8 +3,9 @@ from enum import Enum, auto
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from httpx import URL
from httpx import URL, QueryParams
from pydantic import TypeAdapter
from websockets.exceptions import ConnectionClosed
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.llms.openai.realtime.handler import OpenAIRealtime
@ -40,11 +41,14 @@ def configured_realtime_query(params: GenericLiteLLMParams) -> Mapping[str, str]
inbound: Final = TypeAdapter(Mapping[str, str]).validate_python(
getattr(params, "chatgpt_realtime_client_query", None) or MappingProxyType({})
)
configured: Final = TypeAdapter(Mapping[str, str]).validate_python(
configured: Final = TypeAdapter(Mapping[str, str | int | float | bool | None]).validate_python(
getattr(params, "extra_query", None) or MappingProxyType({})
)
return MappingProxyType(
{**{key: value for key, value in inbound.items() if key in ("intent", "architecture")}, **configured}
{
**{key: value for key, value in inbound.items() if key in ("intent", "architecture")},
**QueryParams(configured),
}
)
@ -111,8 +115,12 @@ class ChatGPTRealtime(OpenAIRealtime):
async def close_call(self, connection: "ClientConnection", model: str, api_base: str) -> None:
if realtime_endpoint(model) == "live":
await connection.send('{"type":"session.close"}')
return
try:
await connection.send('{"type":"session.close"}')
return
except (ConnectionClosed, OSError):
await self.hangup_call(api_base)
return
await self.hangup_call(api_base)
async def hangup_call(self, api_base: str) -> None:

View file

@ -97,6 +97,10 @@ class ReconcileOutcome(NamedTuple):
live_after: frozenset[str] | None
class InternalRequestOrigin(enum.Enum):
REALTIME_OBSERVER = enum.auto()
class SupportedDBObjectType(str, enum.Enum):
"""
Supported database object types for fine-grained DB storage control.

View file

@ -1832,6 +1832,8 @@ class ProxyBaseLLMRequestProcessing:
user_api_base: str | None = None,
model: str | None = None,
llm_router: Router | None = None,
*,
internal_realtime_observer: bool = False,
) -> tuple[dict, LiteLLMLoggingObj]:
start_time: Final = datetime.now() # start before calling guardrail hooks
@ -1996,6 +1998,11 @@ class ProxyBaseLLMRequestProcessing:
user_api_key_dict=user_api_key_dict,
data=self.data,
call_type=route_type,
**(
MappingProxyType({"internal_realtime_observer": True})
if internal_realtime_observer
else MappingProxyType({})
),
)
if route_type == "aget_responses":
attach_post_call_pipelines_to_retrieval(

View file

@ -12,7 +12,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.exceptions import RateLimitType
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth
from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, InternalRequestOrigin, UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
get_key_model_rpm_limit,
get_key_model_tpm_limit,
@ -489,6 +489,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
)
async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time):
releases_slot: Final = kwargs.get("internal_request_origin") is not InternalRequestOrigin.REALTIME_OBSERVER
from litellm.proxy.common_utils.callback_utils import (
get_model_group_from_litellm_kwargs,
)
@ -521,7 +522,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
# Setup values
# ------------
if global_max_parallel_requests is not None:
if releases_slot and global_max_parallel_requests is not None:
# get value from cache
_key: Final = "global_max_parallel_requests"
# decrement
@ -552,13 +553,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
key=request_count_api_key,
litellm_parent_otel_span=litellm_parent_otel_span,
) or {
"current_requests": 1,
"current_requests": int(releases_slot),
"current_tpm": 0,
"current_rpm": 0,
}
new_val = {
"current_requests": max(current["current_requests"] - 1, 0),
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
"current_tpm": current["current_tpm"] + total_tokens,
"current_rpm": current["current_rpm"],
}
@ -593,13 +594,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
key=request_count_api_key,
litellm_parent_otel_span=litellm_parent_otel_span,
) or {
"current_requests": 1,
"current_requests": int(releases_slot),
"current_tpm": 0,
"current_rpm": 0,
}
new_val = {
"current_requests": max(current["current_requests"] - 1, 0),
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
"current_tpm": current["current_tpm"] + total_tokens,
"current_rpm": current["current_rpm"],
}
@ -619,13 +620,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
key=request_count_api_key,
litellm_parent_otel_span=litellm_parent_otel_span,
) or {
"current_requests": 1,
"current_tpm": total_tokens,
"current_rpm": 1,
"current_requests": int(releases_slot),
"current_tpm": total_tokens if releases_slot else 0,
"current_rpm": int(releases_slot),
}
new_val = {
"current_requests": max(current["current_requests"] - 1, 0),
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
"current_tpm": current["current_tpm"] + total_tokens,
"current_rpm": current["current_rpm"],
}
@ -645,13 +646,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
key=request_count_api_key,
litellm_parent_otel_span=litellm_parent_otel_span,
) or {
"current_requests": 1,
"current_tpm": total_tokens,
"current_rpm": 1,
"current_requests": int(releases_slot),
"current_tpm": total_tokens if releases_slot else 0,
"current_rpm": int(releases_slot),
}
new_val = {
"current_requests": max(current["current_requests"] - 1, 0),
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
"current_tpm": current["current_tpm"] + total_tokens,
"current_rpm": current["current_rpm"],
}
@ -671,13 +672,13 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
key=request_count_api_key,
litellm_parent_otel_span=litellm_parent_otel_span,
) or {
"current_requests": 1,
"current_tpm": total_tokens,
"current_rpm": 1,
"current_requests": int(releases_slot),
"current_tpm": total_tokens if releases_slot else 0,
"current_rpm": int(releases_slot),
}
new_val = {
"current_requests": max(current["current_requests"] - 1, 0),
"current_requests": max(current["current_requests"] - int(releases_slot), 0),
"current_tpm": current["current_tpm"] + total_tokens,
"current_rpm": current["current_rpm"],
}
@ -694,6 +695,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
self.print_verbose(e)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
if kwargs.get("internal_request_origin") is InternalRequestOrigin.REALTIME_OBSERVER:
return
try:
self.print_verbose("Inside Max Parallel Request Failure Hook")
litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs)

View file

@ -29,7 +29,7 @@ from litellm.llms.chatgpt.realtime import (
configured_realtime_headers,
realtime_endpoint,
)
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy._types import InternalRequestOrigin, ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
from litellm.proxy.auth.user_api_key_auth import (
get_api_key,
@ -73,6 +73,7 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth:
auth,
call.alias,
"_arealtime",
internal_realtime_observer=True,
)
pinned: Final = { # mutable-ok: logging and provider parameter contract
**processed,
@ -116,6 +117,9 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth:
async def close_call() -> None:
await handler.close_call(connection, call.model, api_base)
async def force_close_call() -> None:
await handler.hangup_call(api_base)
frontend: Final = WebSocket(
{**request.scope, "type": "websocket"}, receive=receive, send=send
) # mutable-ok: ASGI scope
@ -126,6 +130,7 @@ async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth:
logger,
auth,
close_call,
force_close_call=force_close_call,
terminal_usage_required=realtime_endpoint(call.model) == "live",
)
supervision_owned = True
@ -191,6 +196,8 @@ async def process_codex_request(
auth: UserAPIKeyAuth,
model: str,
route_type: Literal["arealtime_calls", "_arealtime"],
*,
internal_realtime_observer: bool = False,
) -> tuple[dict[str, object], Logging]: # mutable-ok: common request processor returns enriched routing arguments
from litellm.proxy import proxy_server as server
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
@ -211,7 +218,14 @@ async def process_codex_request(
user_api_base=server.user_api_base,
model=model,
route_type=route_type,
**(
MappingProxyType({"internal_realtime_observer": True})
if internal_realtime_observer
else MappingProxyType({})
),
)
if internal_realtime_observer:
logging_obj.model_call_details["internal_request_origin"] = InternalRequestOrigin.REALTIME_OBSERVER
return processed, logging_obj

View file

@ -7,6 +7,7 @@ from pydantic import BaseModel
from websockets.exceptions import ConnectionClosedOK
from litellm._logging import verbose_proxy_logger
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
@ -45,23 +46,28 @@ class CallSupervisor:
lifetime: float = 3600,
drain_timeout: float = 5,
termination_timeout: float = 60,
logging_timeout: float = LOGGING_WORKER_MAX_TIME_PER_COROUTINE,
terminal_usage_required: bool = True,
force_close_call: Callable[[], Awaitable[None]] | None = None,
) -> None:
self._upstream = upstream
self._stream = stream
self._logging = logging_obj
self._auth = auth
self._close_call = close_call
self._force_close_call = force_close_call
self._ready_timeout = ready_timeout
self._lifetime = lifetime
self._drain_timeout = drain_timeout
self._termination_timeout = termination_timeout
self._logging_timeout = logging_timeout
self._terminal_usage_required = terminal_usage_required
self._ready = asyncio.Event()
self._stop = asyncio.Event()
self._started = False
self._terminal = False
self._close_confirmed = False
self._accounting_complete = False
self._task: asyncio.Task[None] | None = None
async def start(self) -> None:
@ -70,7 +76,7 @@ class CallSupervisor:
self._task = asyncio.create_task(self._run())
try:
await asyncio.wait_for(self._ready.wait(), timeout=self._ready_timeout)
if not self._started or self._task.done():
if not self._started or self._terminal or self._task.done():
raise RuntimeError("Call observer ended before session became available")
except BaseException:
await self.close()
@ -113,12 +119,21 @@ class CallSupervisor:
finally:
try:
if not self._terminal:
deadline: Final = asyncio.get_running_loop().time() + self._termination_timeout
try:
await asyncio.wait_for(self._close_call(), timeout=self._termination_timeout)
self._close_confirmed = True
except Exception: # noqa: BLE001 # provider exceptions can contain credentials
verbose_proxy_logger.error("Realtime observer could not terminate upstream call")
await self._drain(reader)
await self._drain(reader, timeout=max(0.0, deadline - asyncio.get_running_loop().time()))
if self._terminal_usage_required and not self._terminal and self._force_close_call is not None:
remaining: Final = max(0.0, deadline - asyncio.get_running_loop().time())
try:
await asyncio.wait_for(self._force_close_call(), timeout=remaining)
self._close_confirmed = True
except Exception: # noqa: BLE001 # provider exceptions can contain credentials
verbose_proxy_logger.error("Realtime observer independent hangup failed")
await self._drain(reader, timeout=max(0.0, deadline - asyncio.get_running_loop().time()))
finally:
stopped.cancel()
reader.cancel()
@ -132,9 +147,16 @@ class CallSupervisor:
)
try:
try:
await self._stream.log_messages(wait_for_dispatch=True)
await asyncio.wait_for(
self._stream.log_messages(wait_for_dispatch=True), timeout=self._logging_timeout
)
self._accounting_complete = True
except asyncio.TimeoutError:
verbose_proxy_logger.error("Realtime observer timed out dispatching usage accounting")
finally:
if self._started and not self._usage_complete():
if not self._accounting_complete:
self._logging.model_call_details["realtime_accounting_incomplete"] = True
if self._started and (not self._usage_complete() or not self._accounting_complete):
await invalidate_budget_reservation_counters(
budget_reservation=self._auth.budget_reservation
)
@ -145,9 +167,12 @@ class CallSupervisor:
finally:
self._ready.set()
async def _drain(self, reader: asyncio.Task[None]) -> None:
async def _drain(self, reader: asyncio.Task[None], *, timeout: float | None = None) -> None:
try:
await asyncio.wait_for(asyncio.shield(reader), timeout=self._drain_timeout)
await asyncio.wait_for(
asyncio.shield(reader),
timeout=self._drain_timeout if timeout is None else min(self._drain_timeout, timeout),
)
except asyncio.TimeoutError:
if not self._usage_complete():
verbose_proxy_logger.error("Realtime observer timed out draining terminal usage")

View file

@ -2046,6 +2046,8 @@ class ProxyLogging:
data: None,
call_type: CallTypesLiteral,
guardrails_only: bool = False,
*,
internal_realtime_observer: bool = False,
) -> None:
pass
@ -2056,6 +2058,8 @@ class ProxyLogging:
data: dict,
call_type: CallTypesLiteral,
guardrails_only: bool = False,
*,
internal_realtime_observer: bool = False,
) -> dict:
pass
@ -2065,6 +2069,8 @@ class ProxyLogging:
data: dict | None,
call_type: CallTypesLiteral,
guardrails_only: bool = False,
*,
internal_realtime_observer: bool = False,
) -> dict | None:
"""
Allows users to modify/reject the incoming request to the proxy, without having to deal with parsing Request body.
@ -2163,6 +2169,10 @@ class ProxyLogging:
deferred_route_exc: SensitiveDataRouteException | None = None
for _callback in caps.resolved_callbacks:
if internal_realtime_observer and isinstance(
_callback, (_PROXY_MaxParallelRequestsHandler, _PROXY_MaxParallelRequestsHandler_v3)
):
continue
start_time = time.time()
try:
if isinstance(_callback, CustomGuardrail) and data is not None:

View file

@ -1,4 +1,5 @@
import json
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
@ -11,6 +12,51 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.types.router import GenericLiteLLMParams
@pytest.mark.asyncio
@pytest.mark.parametrize("failure", ["closed", "network"])
@pytest.mark.parametrize(
"hangup_status, expectation", [(200, nullcontext()), (503, pytest.raises(httpx.HTTPStatusError))]
)
async def test_live_closed_observer_uses_independent_hangup(failure, hangup_status, expectation, chatgpt_tokens):
from websockets.exceptions import ConnectionClosedOK
from websockets.frames import Close
handler = ChatGPTRealtime(
GenericLiteLLMParams(
chatgpt_realtime_call_id="rtc_live_closed",
chatgpt_token_dir=chatgpt_tokens,
extra_query={"gateway": "tenant"},
),
{},
{"x-gateway-token": "test-only"},
)
connection = SimpleNamespace(
send=AsyncMock(
side_effect=(
ConnectionClosedOK(Close(1000, ""), Close(1000, ""), True)
if failure == "closed"
else OSError("socket unavailable")
)
)
)
requests = []
def respond(request):
requests.append(request)
return httpx.Response(hangup_status)
client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
with patch("httpx.AsyncClient", return_value=client):
with expectation:
await handler.close_call(connection, "gpt-live-1-codex", "https://gateway.example/v1")
assert len(requests) == 1
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].headers["x-gateway-token"] == "test-only"
assert requests[0].headers["Authorization"] == "Bearer test-token-default"
assert client.is_closed
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["client_secrets", "transcription_sessions"])
@pytest.mark.parametrize("source", ["default", "explicit", "CHATGPT_API_BASE", "OPENAI_CHATGPT_API_BASE"])
@ -47,8 +93,17 @@ async def test_realtime_session_urls_honor_gateway(endpoint, source, chatgpt_tok
@pytest.mark.asyncio
@pytest.mark.parametrize("inbound_headers", [{}, {"openai-alpha": "quicksilver=v2"}])
async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, chatgpt_tokens, monkeypatch):
from litellm.llms.chatgpt.codex import CodexRealtimeOffer, build_call_request
@pytest.mark.parametrize("model, endpoint", [("gpt-live-1-codex", "live"), ("gpt-realtime-1.5", "realtime")])
async def test_routed_call_preserves_deployment_gateway_headers(
inbound_headers, model, endpoint, chatgpt_tokens, monkeypatch
):
from litellm.llms.chatgpt.codex import (
CodexRealtimeCall,
CodexRealtimeOffer,
build_call_request,
build_sideband_request,
parse_call_response,
)
monkeypatch.setenv("CHATGPT_TOKEN_DIR", chatgpt_tokens)
requests = []
@ -64,11 +119,22 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers,
{
"model_name": "voice-gateway",
"litellm_params": {
"model": "chatgpt/gpt-live-1-codex",
"model": f"chatgpt/{model}",
"api_base": "https://voice.example/backend-api/codex",
"extra_headers": {"x-gateway-route": "configured"},
"extra_query": {"gateway_token": "configured", "intent": "pinned-intent"},
"extra_query": {
"gateway_token": "configured",
"intent": "pinned-intent",
"count": 7,
"fraction": 1.5,
"enabled": True,
"disabled": False,
"blank": None,
"model": "other-model",
"call_id": "rtc_wrong",
},
},
"model_info": {"id": "selected-gateway-deployment"},
}
],
num_retries=0,
@ -84,11 +150,29 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers,
"gateway_token": "configured",
"intent": "pinned-intent",
"architecture": "avas",
"count": "7",
"fraction": "1.5",
"enabled": "true",
"disabled": "false",
"blank": "",
"model": "other-model",
"call_id": "rtc_wrong",
}
assert response.extensions["chatgpt_realtime"]["extra_query"] == dict(requests[0].url.params)
assert response.extensions["chatgpt_realtime"]["extra_headers"]["x-gateway-route"] == "configured"
for name, value in inbound_headers.items():
assert requests[0].headers[name] == value
call = parse_call_response(response, alias="voice-gateway", owner="test-owner", expires_at=1)
restored = CodexRealtimeCall.model_validate_json(call.model_dump_json())
assert restored.model_id == "selected-gateway-deployment"
assert restored.model == model
handler = ChatGPTRealtime(GenericLiteLLMParams.model_validate(build_sideband_request(restored)), {})
sideband_url = httpx.URL(handler._construct_url(restored.api_base, {"model": restored.model}))
assert {key: value for key, value in sideband_url.params.items() if key != "call_id"} == {
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.path.endswith("/realtime" if endpoint == "realtime" else "/live/rtc_test")
finally:
await client.client.aclose()

View file

@ -13,7 +13,8 @@ from litellm.proxy.realtime_endpoints.call_sessions import decode_call, encode_c
@pytest.mark.asyncio
@pytest.mark.parametrize("route_type", ["arealtime_calls", "_arealtime"])
async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type):
@pytest.mark.parametrize("observer", [False, True])
async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type, observer):
from fastapi import Request
from litellm import Router
from litellm.proxy import proxy_server as server
@ -21,7 +22,8 @@ async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type)
from litellm.proxy.realtime_endpoints.call_sessions import process_codex_request
class PolicyHook:
async def pre_call_hook(self, user_api_key_dict, data, call_type):
async def pre_call_hook(self, user_api_key_dict, data, call_type, *, internal_realtime_observer=False):
assert internal_realtime_observer is observer
if "model-policy" in data.get("metadata", {}).get("guardrails", []):
raise HTTPException(403, "Model policy rejected request")
return data
@ -34,7 +36,14 @@ async def test_codex_processing_merges_model_guardrails(monkeypatch, route_type)
monkeypatch.setattr(server, "proxy_logging_obj", PolicyHook())
request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", "headers": [], "query_string": b"", "scheme": "http", "server": ("localhost", 80)})
with pytest.raises(HTTPException) as error:
await process_codex_request(request, {"model": "voice-policy"}, UserAPIKeyAuth(), "voice-policy", route_type)
await process_codex_request(
request,
{"model": "voice-policy"},
UserAPIKeyAuth(),
"voice-policy",
route_type,
internal_realtime_observer=observer,
)
assert error.value.status_code == 403
assert error.value.detail == "Model policy rejected request"

View file

@ -30,6 +30,67 @@ class Socket:
self.closed = True
@pytest.mark.asyncio
@pytest.mark.parametrize("fallback", ["terminal", "no_terminal", "timeout"])
async def test_live_unacknowledged_close_uses_bounded_independent_hangup(monkeypatch, fallback):
from litellm.proxy.realtime_endpoints import call_supervision
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
sink = Sink(logger)
invalidate = AsyncMock()
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
async def force_close():
if fallback == "terminal":
await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}})
elif fallback == "timeout":
await asyncio.Event().wait()
force = AsyncMock(side_effect=force_close)
close = AsyncMock()
supervisor = CallSupervisor(
socket,
sink,
logger,
UserAPIKeyAuth(),
close,
force_close_call=force,
drain_timeout=0.01,
termination_timeout=0.08,
)
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await asyncio.wait_for(supervisor.close(), timeout=0.5)
close.assert_awaited_once()
force.assert_awaited_once()
assert socket.closed
if fallback == "terminal":
invalidate.assert_not_awaited()
assert not logger.model_call_details.get("realtime_usage_incomplete")
else:
invalidate.assert_awaited_once()
assert logger.model_call_details["realtime_usage_incomplete"] is True
@pytest.mark.asyncio
async def test_live_confirmed_terminal_does_not_force_hangup():
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
async def close():
await socket.messages.put({"type": "session.closed"})
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 supervisor.close()
force.assert_not_awaited()
class Sink:
def __init__(self, logger):
self.logger = logger
@ -178,7 +239,7 @@ async def test_observer_error_rejects_start(caplog):
@pytest.mark.asyncio
async def test_failed_logging_releases_reservation(monkeypatch):
async def test_failed_logging_invalidates_reservation_without_zeroing_spend(monkeypatch):
from litellm.proxy.realtime_endpoints import call_supervision
socket = Socket()
@ -187,18 +248,89 @@ async def test_failed_logging_releases_reservation(monkeypatch):
sink = MagicMock()
sink.log_messages = AsyncMock(side_effect=RuntimeError("logging unavailable"))
release = AsyncMock()
invalidate = AsyncMock()
monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release)
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock())
await socket.messages.put({"type": "session.started"})
await supervisor.start()
await socket.messages.put({"type": "session.closed"})
with pytest.raises(RuntimeError, match="logging unavailable"):
await supervisor.wait()
release.assert_awaited_once_with(budget_reservation=None)
release.assert_not_awaited()
invalidate.assert_awaited_once_with(budget_reservation=None)
assert logger.model_call_details["realtime_accounting_incomplete"] is True
assert socket.closed
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
@pytest.mark.asyncio
async def test_start_rejects_terminal_session_while_accounting_is_pending():
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
dispatch_started = asyncio.Event()
allow_dispatch = asyncio.Event()
async def log_messages(*, wait_for_dispatch=False):
dispatch_started.set()
await allow_dispatch.wait()
logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
sink = MagicMock()
sink.log_messages = AsyncMock(side_effect=log_messages)
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), AsyncMock())
await socket.messages.put({"type": "session.created"})
await socket.messages.put({"type": "session.closed"})
startup = asyncio.create_task(supervisor.start())
try:
await asyncio.wait_for(dispatch_started.wait(), timeout=1)
finally:
allow_dispatch.set()
with pytest.raises(RuntimeError, match="ended before"):
await asyncio.wait_for(startup, timeout=1)
assert socket.closed
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
@pytest.mark.asyncio
async def test_shutdown_bounds_accounting_and_invalidates_partial_dispatch(monkeypatch):
from litellm.proxy.realtime_endpoints import call_supervision
socket = Socket()
logger = MagicMock(spec=Logging)
logger.model_call_details = {}
dispatch_cancelled = asyncio.Event()
invalidate = AsyncMock()
release = AsyncMock()
monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate)
monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release)
async def log_messages(*, wait_for_dispatch=False):
try:
await asyncio.Event().wait()
finally:
dispatch_cancelled.set()
async def hangup():
await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}})
sink = MagicMock()
sink.log_messages = AsyncMock(side_effect=log_messages)
supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), hangup, logging_timeout=0.01)
registry = CallSupervisors()
await socket.messages.put({"type": "session.created"})
await registry.start(supervisor)
await asyncio.wait_for(registry.shutdown(), timeout=1)
assert dispatch_cancelled.is_set()
assert socket.closed
assert logger.model_call_details["realtime_accounting_incomplete"] is True
assert not logger.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY)
invalidate.assert_awaited_once_with(budget_reservation=None)
release.assert_not_awaited()
sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True)
@pytest.mark.asyncio
async def test_shutdown_waits_for_usage_dispatch_completion():
socket = Socket()

View file

@ -1,5 +1,6 @@
import datetime as real_datetime
import smtplib
from unittest.mock import MagicMock, patch
import pytest
from fastapi import HTTPException
@ -8,15 +9,10 @@ from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import ProxyErrorTypes, UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.proxy.utils import ProxyLogging, get_custom_url, join_paths
from litellm.types.guardrails import GuardrailEventHooks
from unittest.mock import MagicMock, patch
from litellm.proxy.utils import get_custom_url, join_paths
def test_get_custom_url(monkeypatch):
monkeypatch.setenv("SERVER_ROOT_PATH", "/litellm")
custom_url = get_custom_url(request_base_url="http://0.0.0.0:4000", route="ui/")
@ -2030,7 +2026,9 @@ async def test_post_call_failure_hook_redacts_traceback_before_callbacks(monkeyp
with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()):
await proxy_logging_obj.post_call_failure_hook(
request_data={"metadata": {}},
original_exception=HTTPException(status_code=400, detail="Upstream passthrough request failed with status 400"),
original_exception=HTTPException(
status_code=400, detail="Upstream passthrough request failed with status 400"
),
user_api_key_dict=UserAPIKeyAuth(),
traceback_str=upstream_traceback,
)
@ -2038,3 +2036,127 @@ async def test_post_call_failure_hook_redacts_traceback_before_callbacks(monkeyp
assert recorder.received_traceback is not None
assert provider_key not in recorder.received_traceback
assert "REDACTED" in recorder.received_traceback
@pytest.mark.asyncio
@pytest.mark.parametrize("limiter_version", [1, 3])
@pytest.mark.parametrize("limit", ["rpm_limit", "max_parallel_requests"])
async def test_internal_realtime_observer_preserves_quota_and_custom_hooks(monkeypatch, limiter_version, limit):
import asyncio
from datetime import datetime
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3,
_request_stash,
get_request_stash,
)
from litellm.proxy.utils import InternalUsageCache, ProxyLogging
observed = []
class Hook(CustomLogger):
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
observed.append(call_type)
return {**data, "extra_headers": {"x-hook": "required"}}
cache = DualCache()
limiter_type = _PROXY_MaxParallelRequestsHandler if limiter_version == 1 else _PROXY_MaxParallelRequestsHandler_v3
limiter = limiter_type(InternalUsageCache(dual_cache=cache))
proxy = ProxyLogging(UserApiKeyCache())
monkeypatch.setattr(litellm, "callbacks", [limiter, Hook()])
token = _request_stash.set(None)
try:
auth = UserAPIKeyAuth(api_key="observer-quota-test", **{limit: 1})
await proxy.pre_call_hook(
auth, {"model": "voice", "litellm_call_id": "signaling", "metadata": {}}, "arealtime_calls"
)
await asyncio.sleep(0)
initial_stash = get_request_stash()
result = await proxy.pre_call_hook(
auth,
{"model": "voice", "litellm_call_id": "observer", "metadata": {}},
"_arealtime",
internal_realtime_observer=True,
)
assert result["extra_headers"] == {"x-hook": "required"}
assert observed == ["arealtime_calls", "_arealtime"]
if limiter_version == 3:
assert get_request_stash() is initial_stash
assert initial_stash.owner_litellm_call_id == "signaling"
if limit == "max_parallel_requests":
await limiter.async_log_success_event(
{
"litellm_call_id": "signaling",
"litellm_params": {"metadata": {"user_api_key": auth.api_key, "user_api_key_model_max_budget": {}}},
},
litellm.ModelResponse(usage=litellm.Usage(total_tokens=0)),
datetime.now(),
datetime.now(),
)
if limiter_version == 3:
assert initial_stash.parallel_slot is None
await proxy.pre_call_hook(
auth, {"model": "voice", "litellm_call_id": "next", "metadata": {}}, "arealtime_calls"
)
if limiter_version == 1 and limit == "max_parallel_requests":
from litellm.proxy._types import InternalRequestOrigin
await asyncio.sleep(0)
observer_kwargs = {
"internal_request_origin": InternalRequestOrigin.REALTIME_OBSERVER,
"litellm_call_id": "observer",
"litellm_params": {"metadata": {"user_api_key": auth.api_key, "user_api_key_model_max_budget": {}}},
}
await limiter.async_log_success_event(
observer_kwargs,
litellm.ModelResponse(usage=litellm.Usage(total_tokens=17)),
datetime.now(),
datetime.now(),
)
current = await limiter.internal_usage_cache.async_get_cache(
key=f"{auth.api_key}::{datetime.now():%Y-%m-%d-%H-%M}::request_count", litellm_parent_otel_span=None
)
assert current["current_requests"] == 1
assert current["current_tpm"] == 17
with pytest.raises(HTTPException) as error:
await proxy.pre_call_hook(
auth,
{"model": "voice", "litellm_call_id": "forged", "metadata": {}, "internal_realtime_observer": True},
"_arealtime",
)
assert error.value.status_code == 429
finally:
_request_stash.reset(token)
@pytest.mark.asyncio
@pytest.mark.parametrize("scope", ["key", "user", "team", "end_user"])
async def test_internal_observer_missing_legacy_counter_only_adds_usage(scope):
from datetime import datetime
from litellm.proxy._types import InternalRequestOrigin
from litellm.proxy.hooks.parallel_request_limiter import _PROXY_MaxParallelRequestsHandler
from litellm.proxy.utils import InternalUsageCache
limiter = _PROXY_MaxParallelRequestsHandler(InternalUsageCache(dual_cache=DualCache()))
metadata = {"user_api_key": "expired-key", "user_api_key_model_max_budget": {}}
if scope in ("user", "team"):
metadata[f"user_api_key_{scope}_id"] = "expired-scope"
kwargs = {
"internal_request_origin": InternalRequestOrigin.REALTIME_OBSERVER,
"litellm_params": {"metadata": metadata},
**({"user": "expired-scope"} if scope == "end_user" else {}),
}
await limiter.async_log_success_event(
kwargs, litellm.ModelResponse(usage=litellm.Usage(total_tokens=23)), datetime.now(), datetime.now()
)
identity = "expired-key" if scope == "key" else "expired-scope"
current = await limiter.internal_usage_cache.async_get_cache(
key=f"{identity}::{datetime.now():%Y-%m-%d-%H-%M}::request_count", litellm_parent_otel_span=None
)
assert current == {"current_requests": 0, "current_tpm": 23, "current_rpm": 0}