From c360d187f5a78d43a3b8a9bd5771ddd4abe26d17 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 15:54:39 +0200 Subject: [PATCH] fix(chatgpt): preserve observer quotas and bound call cleanup --- litellm/llms/chatgpt/realtime.py | 18 ++- litellm/proxy/_types.py | 4 + litellm/proxy/common_request_processing.py | 7 + .../proxy/hooks/parallel_request_limiter.py | 39 ++--- .../proxy/realtime_endpoints/call_sessions.py | 16 ++- .../realtime_endpoints/call_supervision.py | 37 ++++- litellm/proxy/utils.py | 10 ++ .../llms/chatgpt/test_realtime.py | 92 +++++++++++- .../realtime_endpoints/test_call_sessions.py | 15 +- .../test_call_supervision.py | 136 +++++++++++++++++- tests/test_litellm/proxy/test_proxy_utils.py | 136 +++++++++++++++++- 11 files changed, 464 insertions(+), 46 deletions(-) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index f79713fe03f..27ce6cb3966 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -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: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 5a3633afa3e..78bfa2bf668 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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. diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e3a2b892721..476f80b4158 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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( diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index b313cb64c3f..2ec576b97b3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -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) diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index b85df30bb31..8af489e8eb1 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -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 diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index b574a229bbe..e750ce61d6a 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -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") diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 00ccad33b6d..b43835fd7d9 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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: diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 8b29d440b4c..6db166e956b 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -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() diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py index df1c53ec675..ef85e324680 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -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" diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index 264eccc787e..7d0d614f31d 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -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() diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index b78ec7dcff6..2f7cc16a70d 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -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}