From d6b8360ac3f3ca393a895e0f21e2a8f5bf2d17b8 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Thu, 10 Sep 2026 12:52:04 +0200 Subject: [PATCH] fix(chatgpt): supervise call usage and preserve gateway query routing --- litellm/cost_calculator.py | 45 ++- .../litellm_core_utils/realtime_streaming.py | 21 +- litellm/llms/chatgpt/codex.py | 11 +- litellm/llms/chatgpt/realtime.py | 71 ++++- litellm/llms/openai/realtime/handler.py | 2 + litellm/proxy/proxy_server.py | 3 + .../proxy/realtime_endpoints/call_sessions.py | 148 ++++++++- .../realtime_endpoints/call_supervision.py | 187 +++++++++++ litellm/realtime_api/main.py | 11 +- litellm/types/llms/openai.py | 6 + .../test_realtime_streaming.py | 26 ++ tests/test_litellm/llms/chatgpt/test_codex.py | 23 +- .../llms/chatgpt/test_realtime.py | 68 +++- .../realtime_endpoints/test_call_sessions.py | 252 +++++++++++++-- .../test_call_supervision.py | 301 ++++++++++++++++++ tests/test_litellm/test_cost_calculator.py | 68 ++++ 16 files changed, 1195 insertions(+), 48 deletions(-) create mode 100644 litellm/proxy/realtime_endpoints/call_supervision.py create mode 100644 tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 814eaaf76f7..24e048eafd1 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -8,7 +8,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, cast from httpx import Response -from pydantic import BaseModel +from pydantic import BaseModel, Field, ValidationError import litellm import litellm._logging @@ -2563,7 +2563,12 @@ def handle_realtime_stream_cost_calculation( if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results) else 0.0 ) - total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost + live_audio_cost: Final = handle_live_session_duration_cost( + results=results, + custom_llm_provider=custom_llm_provider, + litellm_model_name=litellm_model_name, + ) + total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost + live_audio_cost _store_cost_breakdown_in_logging_obj( litellm_logging_obj=litellm_logging_obj, @@ -2571,13 +2576,47 @@ def handle_realtime_stream_cost_calculation( completion_tokens_cost_usd_dollar=output_cost_per_token, cost_for_built_in_tools_cost_usd_dollar=0.0, total_cost_usd_dollar=total_cost, - additional_costs={"transcription_cost": transcription_cost} if transcription_cost > 0 else None, + additional_costs={ + name: cost + for name, cost in (("transcription_cost", transcription_cost), ("live_audio_cost", live_audio_cost)) + if cost > 0 + } + or None, data_residency=data_residency, ) return total_cost +class _LiveSessionDurationUsage(BaseModel): + audio_duration_ms: float = Field(strict=True, ge=0, allow_inf_nan=False) + + +class _LiveSessionClosedEvent(BaseModel): + usage: _LiveSessionDurationUsage + + +def handle_live_session_duration_cost( + results: OpenAIRealtimeStreamList, + custom_llm_provider: str, + litellm_model_name: str, +) -> float: + if any(event.get("type") == "response.done" for event in results): + return 0.0 + terminal: Final = next((event for event in reversed(results) if event.get("type") == "session.closed"), None) + if terminal is None: + return 0.0 + try: + usage: Final = _LiveSessionClosedEvent.model_validate(terminal).usage + except ValidationError: + return 0.0 + try: + model_info: Final = litellm.get_model_info(model=litellm_model_name, custom_llm_provider=custom_llm_provider) + except Exception: + return 0.0 + return usage.audio_duration_ms / 1000 * (model_info.get("input_cost_per_second") or 0.0) + + def handle_realtime_transcription_cost_calculation( results: OpenAIRealtimeStreamList, custom_llm_provider: str, diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 75046f2cf87..c5e6b0fafe3 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -6,6 +6,7 @@ from dataclasses import dataclass from enum import Enum, auto from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast +from pydantic import TypeAdapter from typing_extensions import ReadOnly import litellm @@ -16,6 +17,7 @@ from litellm.types.llms.openai import ( OpenAIRealtimeEvents, OpenAIRealtimeOutputItemDone, OpenAIRealtimeResponseDelta, + OpenAIRealtimeSessionClosed, OpenAIRealtimeStreamResponseBaseObject, OpenAIRealtimeStreamSessionEvents, ) @@ -139,11 +141,14 @@ class RealTimeStreaming: force_transcription_model: str | None = None, event_normalizer: RealtimeEventNormalizer | None = None, logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER, + *, + account_usage: bool = True, ): self.websocket: _ClientWebSocket = websocket self.backend_ws = backend_ws self.logging_obj = logging_obj self._logging_worker = logging_worker + self._account_usage = account_usage self.messages: list[OpenAIRealtimeEvents] = [] self._backend_sent_frames: bool = False self.input_message: dict = {} @@ -256,6 +261,9 @@ class RealTimeStreaming: else: message_obj = cast(dict[str, Any], json.loads(cast(str, message))) self._collect_tool_calls_from_response_done(cast(dict, message_obj)) + if message_obj.get("type") == "session.closed" and isinstance(message_obj.get("usage"), dict): + self.messages.append(TypeAdapter(OpenAIRealtimeSessionClosed).validate_python(message_obj)) + return if not self._should_store_message(message_obj): return try: @@ -410,8 +418,10 @@ class RealTimeStreaming: if self.logging_obj: self.logging_obj.pre_call(input=message, api_key="") - async def log_messages(self): + async def log_messages(self, *, wait_for_dispatch: bool = False): """Log messages in list""" + if not self._account_usage: + return if self.logging_obj: if self.input_messages: self.logging_obj.model_call_details["messages"] = self.input_messages @@ -421,9 +431,12 @@ class RealTimeStreaming: # Route through the bounded logging worker (per-coroutine timeout + # concurrency cap) instead of a bare create_task, so a slow callback # can't leave suspended tasks pinning each call's response in memory. - self._logging_worker.ensure_initialized_and_enqueue( - self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True) - ) + if wait_for_dispatch: + await self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True) + else: + self._logging_worker.ensure_initialized_and_enqueue( + self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True) + ) self.logging_obj.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True async def _send_to_backend(self, message: str) -> bool: diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index 78e71963b7f..f71a4d5bbad 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -17,17 +17,22 @@ class CodexRealtimeOffer(BaseModel): class CodexRealtimeCall(BaseModel): call_id: str = Field(pattern=r"^rtc_[A-Za-z0-9_-]+$") model: str + model_id: str | None = None alias: str api_base: str | None = None extra_headers: Mapping[str, str] | None = None + extra_query: Mapping[str, str] | None = None + usage_supervised: bool = False owner: str expires_at: float class ChatGPTCallRouting(BaseModel): model: str + model_id: str | None = None api_base: str | None = None extra_headers: Mapping[str, str] | None = None + extra_query: Mapping[str, str] | None = None class CodexSidebandRequest(TypedDict): @@ -36,6 +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] def build_call_request( @@ -46,7 +52,7 @@ def build_call_request( "sdp_body": offer.sdp.encode(), "session": offer.session.model_dump(exclude_none=True), "openai_ephemeral_key": "", - "extra_query": { # mutable-ok: router request parameters + "chatgpt_realtime_client_query": { # mutable-ok: router request parameters key: value for key, value in query.items() if key in ("intent", "architecture") }, "chatgpt_realtime_client_headers": { # mutable-ok: router request headers @@ -66,11 +72,13 @@ def parse_call_response(response: httpx.Response, alias: str, owner: str, expire return CodexRealtimeCall( call_id=call_id, model=routing.model, + model_id=routing.model_id, alias=alias, owner=owner, expires_at=expires_at, api_base=routing.api_base, extra_headers=routing.extra_headers, + extra_query=routing.extra_query, ) @@ -81,4 +89,5 @@ def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest: chatgpt_realtime_call_id=call.call_id, query_params=RealtimeQueryParams(model=call.model), extra_headers=call.extra_headers, + extra_query=call.extra_query, ) diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index c44a1e4f85e..f79713fe03f 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -1,10 +1,12 @@ from collections.abc import Mapping +from enum import Enum, auto from types import MappingProxyType -from typing import Final +from typing import TYPE_CHECKING, Final from httpx import URL from pydantic import TypeAdapter +from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES from litellm.llms.openai.realtime.handler import OpenAIRealtime from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig from litellm.types.realtime import RealtimeQueryParams @@ -15,6 +17,17 @@ from .authenticator import Authenticator from .common_utils import without_oauth_identity_headers from .responses.transformation import ChatGPTResponsesAPIConfig +if TYPE_CHECKING: + from websockets.asyncio.client import ClientConnection + + +class CallAccounting(Enum): + SUPERVISED = auto() + + +def accounts_for_call_usage(params: GenericLiteLLMParams) -> bool: + return getattr(params, "chatgpt_call_accounting", None) is not CallAccounting.SUPERVISED + def configured_realtime_headers(headers: Mapping[str, object] | None) -> Mapping[str, str]: validated: Final = TypeAdapter(Mapping[str, str]).validate_python( @@ -23,6 +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]: + 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( + getattr(params, "extra_query", None) or MappingProxyType({}) + ) + return MappingProxyType( + {**{key: value for key, value in inbound.items() if key in ("intent", "architecture")}, **configured} + ) + + def realtime_call_headers(params: GenericLiteLLMParams) -> dict[str, str]: # mutable-ok: HTTP handler header contract inbound: Final = TypeAdapter(Mapping[str, str]).validate_python( getattr(params, "chatgpt_realtime_client_headers", None) or MappingProxyType({}) @@ -72,6 +97,40 @@ def realtime_endpoint(model: str) -> str: class ChatGPTRealtime(OpenAIRealtime): + async def open_call_connection(self, model: str, api_base: str) -> "ClientConnection": + import websockets + + url: Final = self._construct_url(api_base, RealtimeQueryParams(model=model)) + return await websockets.connect( + url, + additional_headers=self._profile_headers, + max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, + ssl=self._get_ssl_config(url), + open_timeout=20, + ) + + 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 + await self.hangup_call(api_base) + + async def hangup_call(self, api_base: str) -> None: + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + base: Final = URL(api_base) + 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")), + ) + client: Final = AsyncHTTPHandler() + try: + response: Final = await client.post(str(url), headers=self._profile_headers, data=b"", timeout=10) + response.raise_for_status() + finally: + await client.close() + @staticmethod def get_api_base(api_base: str | None = None) -> str: return api_base or Authenticator.get_api_base(default_base="https://api.openai.com/v1") @@ -85,6 +144,7 @@ class ChatGPTRealtime(OpenAIRealtime): super().__init__() self._profile_headers = realtime_headers(params, headers, extra_headers) self._call_id = TypeAdapter(str | None).validate_python(getattr(params, "chatgpt_realtime_call_id", None)) + self._extra_query = configured_realtime_query(params) def _get_additional_headers( self, api_key: str, *, openai_beta_realtime: bool = False @@ -98,13 +158,16 @@ class ChatGPTRealtime(OpenAIRealtime): base: Final = URL(api_base) 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") + ) return str( base.copy_with( scheme="wss" if base.scheme in ("https", "wss") else "ws", path=f"{base.path.rstrip('/')}/{endpoint}/{self._call_id}" if endpoint == "live" else f"{base.path.rstrip('/')}/realtime", - params=() if endpoint == "live" else (("call_id", self._call_id),), + params=gateway_query + (() if endpoint == "live" else (("call_id", self._call_id),)), ) ) return str( @@ -138,9 +201,7 @@ class ChatGPTRealtimeHTTPConfig(OpenAIRealtimeHTTPConfig): return "chatgpt-oauth" def get_realtime_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str: - query: Final = TypeAdapter(Mapping[str, str]).validate_python( - getattr(self._params, "extra_query", None) or MappingProxyType({}) - ) + query: Final = configured_realtime_query(self._params) return str(URL(f"{self.get_api_base(api_base).rstrip('/')}/realtime/calls", params=query)) def get_realtime_calls_headers( diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index e3ecbac1a53..ca141c0b958 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -117,6 +117,7 @@ class OpenAIRealtime(OpenAIChatCompletion): query_params: RealtimeQueryParams | None = None, user_api_key_dict: object | None = None, litellm_metadata: dict | None = None, + account_usage: bool = True, **kwargs: object, ): import websockets @@ -172,6 +173,7 @@ class OpenAIRealtime(OpenAIChatCompletion): model if (query_params or {}).get("intent") == "transcription" else None ), event_normalizer=self._make_event_normalizer(), + account_usage=account_usage, ) await realtime_streaming.bidirectional_forward() diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8f631ccde52..a9ed397845b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1362,6 +1362,9 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: except Exception as e: verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e) + from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS + + await CALL_SUPERVISORS.shutdown() await _flush_spend_logs_queue_on_shutdown() await proxy_config.stop_config_sync_subscriber() diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py index e02a66193b8..b0833b19414 100644 --- a/litellm/proxy/realtime_endpoints/call_sessions.py +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -2,6 +2,7 @@ import base64 import hashlib import json import time +from contextlib import AsyncExitStack from types import MappingProxyType from typing import Final, Literal @@ -11,7 +12,7 @@ from starlette.types import Message from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import Logging -from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY +from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY, RealTimeStreaming from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.chatgpt.codex import ( CodexRealtimeCall, @@ -20,7 +21,12 @@ from litellm.llms.chatgpt.codex import ( build_sideband_request, parse_call_response, ) -from litellm.llms.chatgpt.realtime import configured_realtime_headers +from litellm.llms.chatgpt.realtime import ( + CallAccounting, + ChatGPTRealtime, + configured_realtime_headers, + realtime_endpoint, +) from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.auth.auth_checks import can_key_call_resolved_model from litellm.proxy.auth.user_api_key_auth import ( @@ -30,7 +36,126 @@ 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.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation +from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + release_or_invalidate_budget_reservation, +) +from litellm.types.router import GenericLiteLLMParams + + +async def supervise_codex_call(request: Request, call: CodexRealtimeCall, auth: UserAPIKeyAuth) -> None: + from collections.abc import Mapping + + from pydantic import TypeAdapter + + import litellm + from litellm.proxy.realtime_endpoints.call_supervision import CALL_SUPERVISORS, CallSupervisor + + async def receive() -> Message: + return { + "type": "http.request", + "body": json.dumps({"model": call.alias}).encode(), + "more_body": False, + } # mutable-ok: ASGI message + + async def send(_message: Message) -> None: + return None + + supervision_owned = False # rebind-ok: supervisor owns cleanup after construction + effective_handler: ChatGPTRealtime | None = None # rebind-ok: reuse hook-enriched credentials for cleanup + sockets: Final = AsyncExitStack() + try: + observer_request: Final = Request({**request.scope}, receive=receive) # mutable-ok: ASGI request scope + processed, logger = await process_codex_request( + observer_request, + { + **build_sideband_request(call), + "model": call.alias, + }, # mutable-ok: common request processing enriches metadata + auth, + call.alias, + "_arealtime", + ) + pinned: Final = { # mutable-ok: logging and provider parameter contract + **processed, + **build_sideband_request(call), + "extra_headers": { + **configured_realtime_headers( + TypeAdapter(Mapping[str, object] | None).validate_python(processed.get("extra_headers")) + ), + **configured_realtime_headers(call.extra_headers), + }, + "litellm_metadata": { + **TypeAdapter(Mapping[str, object]).validate_python(processed.get("litellm_metadata") or {}), + **( + {"model_info": {**litellm.get_model_info(model=call.model_id), "id": call.model_id}} + if call.model_id is not None + else {} + ), + }, + } + logger.update_from_kwargs( + kwargs=pinned, + model=call.model, + user=None, + optional_params={}, # mutable-ok: logging contract + litellm_params={ + **logger.litellm_params, + "litellm_metadata": pinned["litellm_metadata"], + "arealtime": True, + }, # mutable-ok: logging contract + custom_llm_provider="chatgpt", + ) + params: Final = GenericLiteLLMParams.model_validate(pinned) + handler: Final = ChatGPTRealtime( + params, request.headers, TypeAdapter(Mapping[str, object]).validate_python(pinned["extra_headers"]) + ) + effective_handler = handler + api_base: Final = ChatGPTRealtime.get_api_base(call.api_base) + connection: Final = await handler.open_call_connection(call.model, api_base) + sockets.push_async_callback(connection.close) + + async def close_call() -> None: + await handler.close_call(connection, call.model, api_base) + + frontend: Final = WebSocket( + {**request.scope, "type": "websocket"}, receive=receive, send=send + ) # mutable-ok: ASGI scope + stream: Final = RealTimeStreaming(frontend, connection, logger, model=call.model, user_api_key_dict=auth) + supervisor: Final = CallSupervisor( + connection, + stream, + logger, + auth, + close_call, + terminal_usage_required=realtime_endpoint(call.model) == "live", + ) + supervision_owned = True + sockets.pop_all() + await CALL_SUPERVISORS.start(supervisor) + except BaseException: + if not supervision_owned: + try: + fallback_handler: Final = effective_handler or ChatGPTRealtime( + GenericLiteLLMParams.model_validate(build_sideband_request(call)), + request.headers, + call.extra_headers, + ) + await fallback_handler.hangup_call(ChatGPTRealtime.get_api_base(call.api_base)) + except Exception: # noqa: BLE001 # preserve original failure without logging provider credentials + verbose_proxy_logger.error("Realtime startup cleanup could not confirm upstream termination") + try: + await invalidate_budget_reservation_counters(budget_reservation=auth.budget_reservation) + except Exception: # noqa: BLE001 # cleanup errors must not replace the original startup failure + verbose_proxy_logger.error("Realtime startup cleanup could not invalidate budget counters") + else: + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) + finally: + try: + await sockets.aclose() + except Exception: # noqa: BLE001 # socket cleanup must preserve the original startup failure + verbose_proxy_logger.error("Realtime startup cleanup could not close observer socket") + raise def encode_call(call: CodexRealtimeCall) -> str: @@ -126,6 +251,7 @@ async def create_codex_realtime_call(request: Request) -> Response: owner_key: Final = ( get_api_key_from_custom_header(request, custom_header) if isinstance(custom_header, str) else selected_key ) + supervision_started = False # rebind-ok: transfer reservation ownership only after supervision is established try: await can_key_call_resolved_model( model=model, @@ -134,7 +260,10 @@ async def create_codex_realtime_call(request: Request) -> Response: llm_router=server.llm_router, ) data: Final = build_call_request(offer, request.query_params, request.headers) - processed, _ = await process_codex_request(request, data, auth, model, "arealtime_calls") + signaling_auth: Final = auth.model_copy( + update={"budget_reservation": None} + ) # mutable-ok: Pydantic update contract + processed, _ = await process_codex_request(request, data, signaling_auth, model, "arealtime_calls") result: Final = await server.route_request( data=processed, route_type="arealtime_calls", @@ -158,7 +287,12 @@ async def create_codex_realtime_call(request: Request) -> Response: ) except ValueError as exc: raise HTTPException(400, str(exc)) from exc - token: Final = encode_call(call) + supervised_call: Final = call.model_copy( + update={"usage_supervised": True} + ) # mutable-ok: Pydantic update contract + token: Final = encode_call(supervised_call) + supervision_started = True + await supervise_codex_call(request, supervised_call, auth) return Response( response.content, status_code=response.status_code, @@ -166,7 +300,8 @@ async def create_codex_realtime_call(request: Request) -> Response: headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}), ) finally: - await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) + if not supervision_started: + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAPIKeyAuth) -> None: @@ -238,6 +373,7 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP ), "websocket": websocket, "user_api_key_dict": auth, + "chatgpt_call_accounting": CallAccounting.SUPERVISED if call.usage_supervised else None, } ) finally: diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py new file mode 100644 index 00000000000..98f018353ba --- /dev/null +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -0,0 +1,187 @@ +import asyncio +from collections.abc import AsyncIterator, Awaitable, Callable +from contextlib import suppress +from typing import Final, Protocol + +from pydantic import BaseModel +from websockets.exceptions import ConnectionClosedOK + +from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + release_or_invalidate_budget_reservation, +) + + +class ObserverSocket(Protocol): + def __aiter__(self) -> AsyncIterator[str | bytes]: ... + + async def close(self) -> None: ... + + +class UsageSink(Protocol): + def store_message(self, message: str) -> None: ... + + async def log_messages(self, *, wait_for_dispatch: bool = False) -> None: ... + + +class _ObserverEvent(BaseModel): + type: str + + +class CallSupervisor: + def __init__( + self, + upstream: ObserverSocket, + stream: UsageSink, + logging_obj: Logging, + auth: UserAPIKeyAuth, + close_call: Callable[[], Awaitable[None]], + *, + ready_timeout: float = 20, + lifetime: float = 3600, + drain_timeout: float = 5, + terminal_usage_required: bool = True, + ) -> None: + self._upstream = upstream + self._stream = stream + self._logging = logging_obj + self._auth = auth + self._close_call = close_call + self._ready_timeout = ready_timeout + self._lifetime = lifetime + self._drain_timeout = drain_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._task: asyncio.Task[None] | None = None + + async def start(self) -> None: + if self._task is not None: + raise RuntimeError("Call observer already started") + 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(): + raise RuntimeError("Call observer ended before session became available") + except BaseException: + await self.close() + raise + + async def close(self) -> None: + self._stop.set() + await self.wait() + + async def wait(self) -> None: + if self._task is not None: + await asyncio.shield(self._task) + + async def _read(self) -> None: + try: + await self._read_events() + except ConnectionClosedOK: + return + + async def _read_events(self) -> None: + event: _ObserverEvent + async for message in self._upstream: + self._stream.store_message(message.decode("utf-8") if isinstance(message, bytes) else message) + event = _ObserverEvent.model_validate_json(message) + if event.type in ("session.started", "session.created"): + self._started = True + self._ready.set() + if event.type == "session.closed": + self._terminal = True + return + + def _usage_complete(self) -> bool: + return self._terminal or (not self._terminal_usage_required and self._close_confirmed) + + async def _run(self) -> None: + reader: Final = asyncio.create_task(self._read()) + stopped: Final = asyncio.create_task(self._stop.wait()) + try: + await asyncio.wait((reader, stopped), timeout=self._lifetime, return_when=asyncio.FIRST_COMPLETED) + finally: + try: + if not self._terminal: + try: + await asyncio.wait_for(self._close_call(), timeout=self._drain_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) + finally: + stopped.cancel() + reader.cancel() + await asyncio.gather(reader, stopped, return_exceptions=True) + with suppress(Exception): + await self._upstream.close() + if not self._usage_complete(): + self._logging.model_call_details["realtime_usage_incomplete"] = True + verbose_proxy_logger.error( + "Realtime observer ended without terminal usage; recorded usage is partial" + ) + try: + try: + await self._stream.log_messages(wait_for_dispatch=True) + finally: + if self._started and not self._usage_complete(): + await invalidate_budget_reservation_counters( + budget_reservation=self._auth.budget_reservation + ) + elif not self._logging.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY): + await release_or_invalidate_budget_reservation( + budget_reservation=self._auth.budget_reservation + ) + finally: + self._ready.set() + + async def _drain(self, reader: asyncio.Task[None]) -> None: + try: + await asyncio.wait_for(asyncio.shield(reader), timeout=self._drain_timeout) + except asyncio.TimeoutError: + if not self._usage_complete(): + verbose_proxy_logger.error("Realtime observer timed out draining terminal usage") + except Exception: # noqa: BLE001 # cleanup must settle the socket even when reading or closing fails + verbose_proxy_logger.error("Realtime observer could not drain terminal usage") + return + + +class CallSupervisors: + def __init__(self) -> None: + self._tasks: tuple[asyncio.Task[None], ...] = () + self._calls: tuple[CallSupervisor, ...] = () + + async def start(self, supervisor: CallSupervisor) -> None: + self._calls = (*self._calls, supervisor) + try: + await supervisor.start() + except BaseException: + self._calls = tuple(call for call in self._calls if call is not supervisor) + raise + task: Final = asyncio.create_task(self._watch(supervisor)) + self._tasks = (*self._tasks, task) + + async def _watch(self, supervisor: CallSupervisor) -> None: + try: + try: + await supervisor.wait() + except Exception: # noqa: BLE001 # task must be consumed without exposing provider exception payloads + verbose_proxy_logger.error("Realtime observer accounting failed") + finally: + self._calls = tuple(call for call in self._calls if call is not supervisor) + self._tasks = tuple(task for task in self._tasks if task is not asyncio.current_task()) + + async def shutdown(self) -> None: + await asyncio.gather(*(call.close() for call in self._calls), return_exceptions=True) + await asyncio.gather(*self._tasks, return_exceptions=True) + + +CALL_SUPERVISORS: Final = CallSupervisors() diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 849934b59cd..23963475ed0 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -309,13 +309,19 @@ async def arealtime_calls( api_version=litellm_params.api_version, ) if custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ChatGPTRealtime, configured_realtime_headers + from litellm.llms.chatgpt.realtime import ( + ChatGPTRealtime, + configured_realtime_headers, + configured_realtime_query, + ) response.extensions["chatgpt_realtime"] = MappingProxyType( { "model": model_name, + "model_id": litellm_logging_obj.get_router_model_id(), "api_base": ChatGPTRealtime.get_api_base(litellm_params.api_base), "extra_headers": configured_realtime_headers(call_headers), + "extra_query": configured_realtime_query(litellm_params), } ) return response @@ -464,7 +470,7 @@ async def _arealtime( litellm_metadata=_build_litellm_metadata(kwargs), ) elif _custom_llm_provider == "chatgpt": - from litellm.llms.chatgpt.realtime import ChatGPTRealtime + from litellm.llms.chatgpt.realtime import ChatGPTRealtime, accounts_for_call_usage await ChatGPTRealtime(litellm_params, websocket.headers, headers).async_realtime( model=model, @@ -476,6 +482,7 @@ async def _arealtime( query_params=query_params, user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata(kwargs), + account_usage=accounts_for_call_usage(litellm_params), ) elif _custom_llm_provider == "openai": api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or "https://api.openai.com/" diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index b6da9490e01..1c87e355382 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -2031,6 +2031,11 @@ class OpenAIRealtimeStreamResponseBaseObject(TypedDict): type: str +class OpenAIRealtimeSessionClosed(TypedDict): + type: ReadOnly[Literal["session.closed"]] + usage: ReadOnly[Mapping[str, object]] + + class OpenAIRealtimeConversationObject(TypedDict, total=False): id: str object: Required[Literal["realtime.conversation"]] @@ -2236,6 +2241,7 @@ class OpenAIRealtimeEventTypes(Enum): OpenAIRealtimeEvents = ( OpenAIRealtimeStreamResponseBaseObject + | OpenAIRealtimeSessionClosed | OpenAIRealtimeStreamSessionEvents | OpenAIRealtimeStreamResponseOutputItemAdded | OpenAIRealtimeResponseContentPartAdded diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 9c0f6f59463..cba46374837 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -3412,3 +3412,29 @@ async def test_refused_session_does_not_stamp_the_reservation_ownership_marker() assert session.logging.logged_failures == (upstream_close,) assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in session.logging.model_call_details + + +def test_live_terminal_usage_survives_filtered_event_logging(monkeypatch): + from litellm.cost_calculator import RealtimeAPITokenUsageProcessor + + def terminal(): + return {"type": "session.closed", "usage": {"audio_duration_ms": 4000, "backend_model_usage": []}} + + monkeypatch.setattr(litellm, "logged_real_time_event_types", []) + stream = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + event = {**terminal(), "private_transcript": "Do not retain this text"} + stream.store_message(event) + assert stream.messages == [terminal()] + usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(stream.messages) + assert usage.total_tokens == 0 + + +@pytest.mark.asyncio +async def test_live_attachment_does_not_dispatch_duplicate_usage(): + worker = MagicMock() + logger = MagicMock() + stream = RealTimeStreaming(MagicMock(), MagicMock(), logger, logging_worker=worker, account_usage=False) + stream.store_message({"type": "session.closed", "usage": {"audio_duration_ms": 4000}}) + await stream.log_messages() + worker.ensure_initialized_and_enqueue.assert_not_called() + logger.dispatch_success_handlers.assert_not_called() diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 06f04b9b81c..402afaaa58c 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -1,7 +1,7 @@ import httpx import pytest -from litellm.llms.chatgpt.codex import build_sideband_request, parse_call_response +from litellm.llms.chatgpt.codex import CodexRealtimeCall, build_sideband_request, parse_call_response @pytest.mark.parametrize("location", ["", "/v1/realtime/calls/foreign-id"]) @@ -12,14 +12,25 @@ def test_signaling_rejects_invalid_upstream_call_id(location): parse_call_response(response, "voice", "owner", 1000) -def test_signaling_preserves_selected_model_for_sideband(): - response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"}, - extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex", - "extra_headers": {"x-gateway-route": "voice"}}}) +@pytest.mark.parametrize("extra_query", [None, {"gateway_token": "opaque +/& value"}]) +def test_signaling_preserves_selected_model_for_sideband(extra_query): + response = httpx.Response( + 201, + headers={"Location": "/v1/realtime/calls/rtc_provider"}, + extensions={ + "chatgpt_realtime": { + "model": "gpt-live-1-codex", + "api_base": "https://voice.example/codex", + "extra_headers": {"x-gateway-route": "voice"}, + **({"extra_query": extra_query} if extra_query is not None else {}), + } + }, + ) call = parse_call_response(response, "voice", "owner", 1000) - request = build_sideband_request(call) + request = build_sideband_request(CodexRealtimeCall.model_validate_json(call.model_dump_json(exclude_none=True))) assert request["api_base"] == "https://voice.example/codex" assert request["model"] == "chatgpt/gpt-live-1-codex" assert request["chatgpt_realtime_call_id"] == "rtc_provider" assert request["query_params"] == {"model": "gpt-live-1-codex"} assert request["extra_headers"] == {"x-gateway-route": "voice"} + assert request["extra_query"] == extra_query diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index b299e2a6842..8b29d440b4c 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -67,6 +67,7 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, "model": "chatgpt/gpt-live-1-codex", "api_base": "https://voice.example/backend-api/codex", "extra_headers": {"x-gateway-route": "configured"}, + "extra_query": {"gateway_token": "configured", "intent": "pinned-intent"}, }, } ], @@ -74,8 +75,17 @@ async def test_routed_call_preserves_deployment_gateway_headers(inbound_headers, ) offer = CodexRealtimeOffer(sdp="v=0\r\n", session={"model": "voice-gateway"}) try: - response = await router.arealtime_calls(**build_call_request(offer, {}, inbound_headers), client=client) + response = await router.arealtime_calls( + **build_call_request(offer, {"intent": "quicksilver", "architecture": "avas"}, inbound_headers), + client=client, + ) assert requests[0].headers.get("x-gateway-route") == "configured" + assert dict(requests[0].url.params) == { + "gateway_token": "configured", + "intent": "pinned-intent", + "architecture": "avas", + } + 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 @@ -137,6 +147,7 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap sdp_body=b"v=0\r\n", session={"model": "chatgpt/gpt-live-1-codex", "audio": {"output": {"voice": "sol"}}}, extra_query={"intent": "quicksilver", "architecture": "avas"}, + chatgpt_realtime_client_query={"intent": "untrusted-override", "architecture": "avas", "untrusted": "bad"}, extra_headers={ "openai-alpha": "quicksilver=v2", "x-gateway-route": "voice", @@ -146,9 +157,13 @@ async def test_chatgpt_call_keeps_oauth_and_frameless_session(chatgpt_tokens, ap client=client, ) assert response.extensions["chatgpt_realtime"]["api_base"] == (api_base or "https://api.openai.com/v1") - assert response.extensions["chatgpt_realtime"]["extra_headers"] == {"openai-alpha": "quicksilver=v2", "x-gateway-route": "voice"} + assert response.extensions["chatgpt_realtime"]["extra_headers"] == { + "openai-alpha": "quicksilver=v2", + "x-gateway-route": "voice", + } assert requests[0].url.host == ("voice.example" if api_base else "chatgpt.com") assert response.status_code == 201 + assert response.extensions["chatgpt_realtime"]["extra_query"] == {"intent": "quicksilver", "architecture": "avas"} assert requests[0].url.path == "/backend-api/codex/realtime/calls" assert requests[0].url.params["architecture"] == "avas" assert requests[0].headers["authorization"] == "Bearer test-token-" + "default" @@ -244,3 +259,52 @@ def test_realtime_routes_use_configured_gateway(monkeypatch, env_name, api_base, assert handler._construct_url(handler.get_api_base(api_base), {"model": "gpt-realtime-1.5"}) == ( expected.replace("https://", "wss://") + "/realtime?model=gpt-realtime-1.5" ) + + +@pytest.mark.parametrize("model,endpoint", [("gpt-live-1-codex", "live"), ("gpt-realtime-1.5", "realtime")]) +def test_sideband_restores_gateway_query_without_overriding_call(model, endpoint, chatgpt_tokens): + handler = ChatGPTRealtime( + GenericLiteLLMParams( + chatgpt_realtime_call_id="rtc_selected", + extra_query={"gateway_token": "opaque +/& value", "model": "other", "call_id": "rtc_other"}, + ), + {}, + ) + url = httpx.URL(handler._construct_url("https://gateway.example/v1", {"model": model})) + assert url.params["gateway_token"] == "opaque +/& value" + assert "model" not in url.params + if endpoint == "live": + assert url.path == "/v1/live/rtc_selected" + assert "call_id" not in url.params + else: + assert url.path == "/v1/realtime" + assert url.params["call_id"] == "rtc_selected" + + +def test_client_cannot_forge_supervised_call_accounting(chatgpt_tokens): + from litellm.llms.chatgpt.realtime import CallAccounting, accounts_for_call_usage + + assert accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting={"supervised": True})) + assert accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting="supervised")) + assert not accounts_for_call_usage(GenericLiteLLMParams(chatgpt_call_accounting=CallAccounting.SUPERVISED)) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ["gpt-live-1-codex", "gpt-realtime-1.5"]) +async def test_supervisor_connection_preserves_call_routing(model, chatgpt_tokens): + handler = ChatGPTRealtime( + GenericLiteLLMParams( + chatgpt_token_dir=chatgpt_tokens, + chatgpt_realtime_call_id="rtc_owner", + extra_query={"gateway_token": "a+b&c"}, + ), + {"openai-alpha": "quicksilver=v2"}, + {"x-gateway-token": "configured"}, + ) + connection = AsyncMock() + with patch("websockets.connect", AsyncMock(return_value=connection)) as connect: + assert await handler.open_call_connection(model, "https://gateway.example/v1") is connection + url = httpx.URL(connect.call_args.args[0]) + assert url.params["gateway_token"] == "a+b&c" + assert connect.call_args.kwargs["additional_headers"]["x-gateway-token"] == "configured" + assert url.path.endswith("/rtc_owner") if model == "gpt-live-1-codex" else url.params["call_id"] == "rtc_owner" 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 32b13ad678a..df1c53ec675 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -174,7 +174,9 @@ async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id @pytest.mark.parametrize("multipart", [False, True]) @pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol", "x-litellm-api-key", "custom"]) @pytest.mark.parametrize("signaling_credential", ["authorization", "api-key", "x-litellm-api-key", "mixed"]) -async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential, signaling_credential): +async def test_offer_exchange_wraps_call_and_filters_client_headers( + monkeypatch, multipart, credential, signaling_credential +): import json from unittest.mock import AsyncMock @@ -188,11 +190,15 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") session = {"model": "voice-alias", "audio": {"output": {"voice": "sol"}}} if multipart: - body_request = httpx.Request("POST", "http://test/v1/realtime/calls", files={ - "sdp": (None, "v=0\r\n"), "session": (None, json.dumps(session)) - }) + body_request = httpx.Request( + "POST", + "http://test/v1/realtime/calls", + files={"sdp": (None, "v=0\r\n"), "session": (None, json.dumps(session))}, + ) else: - body_request = httpx.Request("POST", "http://test/v1/realtime/calls", json={"sdp": "v=0\r\n", "session": session}) + body_request = httpx.Request( + "POST", "http://test/v1/realtime/calls", json={"sdp": "v=0\r\n", "session": session} + ) body = body_request.read() async def receive(): @@ -203,16 +209,30 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, if signaling_credential == "mixed" else [(signaling_credential.encode(), b"Bearer owner" if signaling_credential == "authorization" else b"owner")] ) - request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", - "scheme": "http", "server": ("localhost", 80), - "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", - "headers": [(b"content-type", body_request.headers["content-type"].encode()), - *signaling_headers, *([(b"x-proxy-key", b"Bearer owner")] if credential == "custom" else []), (b"openai-alpha", b"quicksilver=v2"), - (b"x-untrusted", b"bad")]}, receive) + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/realtime/calls", + "scheme": "http", + "server": ("localhost", 80), + "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", + "headers": [ + (b"content-type", body_request.headers["content-type"].encode()), + *signaling_headers, + *([(b"x-proxy-key", b"Bearer owner")] if credential == "custom" else []), + (b"openai-alpha", b"quicksilver=v2"), + (b"x-untrusted", b"bad"), + ], + }, + receive, + ) auth = UserAPIKeyAuth() authorize = AsyncMock() monkeypatch.setattr(proxy_server, "master_key", "owner") - monkeypatch.setattr(proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"} if credential == "custom" else {}) + monkeypatch.setattr( + proxy_server, "general_settings", {"litellm_key_header_name": "x-proxy-key"} if credential == "custom" else {} + ) monkeypatch.setattr(codex, "can_key_call_resolved_model", authorize) class Processor: @@ -225,7 +245,16 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, assert self.data["model"] == "voice-alias" assert self.data["guardrails"] == ["query-guardrail"] assert await kwargs["request"].json() == {"model": "voice-alias"} - return {**self.data, "extra_headers": {"X-Hook-Required": "policy-value", "x-gateway-token": "untrusted-override", "Authorization": "Bearer untrusted"}, "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}}, None + return { + **self.data, + "extra_headers": { + "X-Hook-Required": "policy-value", + "x-gateway-token": "untrusted-override", + "Authorization": "Bearer untrusted", + }, + "extra_query": {"gateway_token": "untrusted-override"}, + "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}, + }, None return self.data, None monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", Processor) @@ -236,14 +265,28 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, assert data["session"] == session assert data["chatgpt_realtime_client_headers"] == {"openai-alpha": "quicksilver=v2"} assert "extra_headers" not in data - assert data["extra_query"] == {"intent": "quicksilver", "architecture": "avas"} + assert data["chatgpt_realtime_client_query"] == {"intent": "quicksilver", "architecture": "avas"} async def respond(): - return httpx.Response(201, content=b"v=0\r\nanswer", headers={"Location": "/v1/realtime/calls/rtc_private"}, - extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex", "api_base": "https://voice.example/codex", "extra_headers": {"X-Gateway-Token": "pinned-value"}}}) + return httpx.Response( + 201, + content=b"v=0\r\nanswer", + headers={"Location": "/v1/realtime/calls/rtc_private"}, + extensions={ + "chatgpt_realtime": { + "model": "gpt-live-1-codex", + "api_base": "https://voice.example/codex", + "extra_headers": {"X-Gateway-Token": "pinned-value"}, + "extra_query": {"gateway_token": "pinned-query-value"}, + } + }, + ) + return respond() monkeypatch.setattr(proxy_server, "route_request", route) + supervise = AsyncMock() + monkeypatch.setattr(codex, "supervise_codex_call", supervise) response = await codex.create_codex_realtime_call(request) assert response.status_code == 201 assert response.body == b"v=0\r\nanswer" @@ -252,7 +295,11 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, assert call.call_id == "rtc_private" assert call.alias == "voice-alias" assert call.model == "gpt-live-1-codex" + assert call.usage_supervised + supervise.assert_awaited_once() assert "rtc_private" not in token + assert "pinned-query-value" not in token + assert call.extra_query == {"gateway_token": "pinned-query-value"} assert time.time() < call.expires_at < time.time() + 3601 authorize.assert_awaited_once() @@ -271,16 +318,28 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, "custom": [(b"x-proxy-key", b"Bearer owner")], "subprotocol": [(b"sec-websocket-protocol", b"realtime, openai-insecure-api-key.owner")], } - websocket = WebSocket({"type": "websocket", "path": "/v1/live/opaque", - "query_string": b"guardrails=query-guardrail", "headers": credential_headers[credential]}, receive_ws, send) + websocket = WebSocket( + { + "type": "websocket", + "path": "/v1/live/opaque", + "query_string": b"guardrails=query-guardrail", + "headers": credential_headers[credential], + }, + receive_ws, + send, + ) forward = AsyncMock() monkeypatch.setattr(litellm, "_arealtime", forward) await codex.codex_realtime_sideband(websocket, token, auth) assert sent[0]["type"] == "websocket.accept" if credential == "subprotocol": assert sent[0]["subprotocol"] == "realtime" - assert forward.await_args.kwargs["extra_headers"] == {"x-hook-required": "policy-value", "x-gateway-token": "pinned-value"} + assert forward.await_args.kwargs["extra_headers"] == { + "x-hook-required": "policy-value", + "x-gateway-token": "pinned-value", + } assert forward.await_args.kwargs["metadata"] == {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"} + assert forward.await_args.kwargs["extra_query"] == {"gateway_token": "pinned-query-value"} assert forward.await_args.kwargs["chatgpt_realtime_call_id"] == "rtc_private" assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-codex" assert forward.await_args.kwargs["api_base"] == "https://voice.example/codex" @@ -344,3 +403,158 @@ async def test_sideband_pre_call_block_prevents_upstream_connection(monkeypatch) await codex.codex_realtime_sideband(websocket, token, UserAPIKeyAuth()) forward.assert_not_called() assert sent == [{"type": "websocket.close", "code": 1008, "reason": "Realtime pre-call rejected"}] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("observer_fails", [False, True]) +async def test_signaling_transfers_reservation_only_to_ready_observer(monkeypatch, observer_fails): + import json + from unittest.mock import AsyncMock + + import httpx + from fastapi import Request + + from litellm.proxy import proxy_server + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-reservation-transfer") + reservation = {"reserved_cost": 0.55, "input_cost": 0.0, "finalized": False, "entries": []} + auth = UserAPIKeyAuth(budget_reservation=reservation) + monkeypatch.setattr(codex, "user_api_key_auth", AsyncMock(return_value=auth)) + monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock()) + monkeypatch.setattr(proxy_server, "general_settings", {}) + process = AsyncMock(return_value=({}, None)) + monkeypatch.setattr(codex, "process_codex_request", process) + + async def response(): + return httpx.Response( + 201, + text="v=0\r\n", + headers={"Location": "/v1/realtime/calls/rtc_ready"}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}, + ) + + async def route(**kwargs): + return response() + + monkeypatch.setattr(proxy_server, "route_request", route) + + async def supervise(request, call, owner): + assert owner is auth + assert not owner.budget_reservation["finalized"] + assert call.usage_supervised + if observer_fails: + await codex.release_or_invalidate_budget_reservation(budget_reservation=owner.budget_reservation) + raise RuntimeError("Observer unavailable") + + monkeypatch.setattr(codex, "supervise_codex_call", supervise) + + async def receive(): + return {"type": "http.request", "body": json.dumps({"sdp": "v=0", "session": {"model": "voice"}}).encode()} + + request = Request( + { + "type": "http", + "method": "POST", + "path": "/v1/realtime/calls", + "query_string": b"", + "headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer owner")], + }, + receive, + ) + if observer_fails: + with pytest.raises(RuntimeError, match="Observer unavailable"): + await codex.create_codex_realtime_call(request) + else: + assert (await codex.create_codex_realtime_call(request)).status_code == 201 + assert process.await_args.args[2].budget_reservation is None + assert auth.budget_reservation["finalized"] is observer_fails + + +@pytest.mark.asyncio +async def test_supervisor_policy_failure_hangs_up_before_releasing(monkeypatch): + from unittest.mock import AsyncMock + + from fastapi import Request + + call = CodexRealtimeCall( + call_id="rtc_open", model="gpt-live-1-codex", alias="voice", owner="owner", expires_at=time.time() + 60 + ) + auth = UserAPIKeyAuth(budget_reservation={"reserved_cost": 0.5, "finalized": False, "entries": []}) + monkeypatch.setattr(codex, "process_codex_request", AsyncMock(side_effect=HTTPException(403, "Policy rejected"))) + closed = [] + + class Handler: + def __init__(self, *args): + pass + + @staticmethod + def get_api_base(base): + return "https://gateway.test/v1" + + async def hangup_call(self, base): + assert not auth.budget_reservation["finalized"] + closed.append(base) + + monkeypatch.setattr(codex, "ChatGPTRealtime", Handler) + request = Request({"type": "http", "headers": [], "method": "POST", "path": "/v1/realtime/calls"}) + with pytest.raises(HTTPException) as error: + await codex.supervise_codex_call(request, call, auth) + assert error.value.status_code == 403 + assert closed == ["https://gateway.test/v1"] + assert auth.budget_reservation["finalized"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("hangup_fails", [False, True]) +async def test_supervisor_constructor_failure_closes_effective_connection(monkeypatch, hangup_fails, caplog): + from unittest.mock import AsyncMock, MagicMock + + from fastapi import Request + + call = CodexRealtimeCall( + call_id="rtc_open", model="gpt-live-1-codex", alias="voice", owner="owner", expires_at=time.time() + 60 + ) + auth = UserAPIKeyAuth(budget_reservation={"reserved_cost": 0.5, "finalized": False, "entries": []}) + logger = MagicMock() + logger.litellm_params = {} + connection = AsyncMock() + handlers = [] + invalidate = AsyncMock() + release = AsyncMock() + monkeypatch.setattr(codex, "invalidate_budget_reservation_counters", invalidate, raising=False) + monkeypatch.setattr(codex, "release_or_invalidate_budget_reservation", release) + monkeypatch.setattr( + codex, "process_codex_request", AsyncMock(return_value=({"extra_headers": {"x-hook": "effective"}}, logger)) + ) + + class Handler: + def __init__(self, params, headers, extra_headers): + self.headers = extra_headers + handlers.append(self) + + @staticmethod + def get_api_base(base): + return "https://gateway.test/v1" + + async def open_call_connection(self, model, base): + return connection + + async def hangup_call(self, base): + assert self.headers["x-hook"] == "effective" + if hangup_fails: + raise RuntimeError("private-cleanup-credential") + + monkeypatch.setattr(codex, "ChatGPTRealtime", Handler) + monkeypatch.setattr(codex, "RealTimeStreaming", MagicMock(side_effect=ValueError("original constructor failure"))) + request = Request({"type": "http", "headers": [], "method": "POST", "path": "/v1/realtime/calls"}) + with pytest.raises(ValueError, match="original constructor failure"): + await codex.supervise_codex_call(request, call, auth) + connection.close.assert_awaited_once() + assert len(handlers) == 1 + if hangup_fails: + invalidate.assert_awaited_once_with(budget_reservation=auth.budget_reservation) + release.assert_not_awaited() + else: + release.assert_awaited_once_with(budget_reservation=auth.budget_reservation) + invalidate.assert_not_awaited() + assert "private-cleanup-credential" not in caplog.text diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py new file mode 100644 index 00000000000..642bfd4982b --- /dev/null +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -0,0 +1,301 @@ +import asyncio +import json +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.litellm_core_utils.realtime_streaming import REALTIME_SESSION_SUCCESS_LOGGED_KEY +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.realtime_endpoints.call_supervision import CallSupervisor, CallSupervisors + + +class Socket: + def __init__(self): + self.messages = asyncio.Queue() + self.closed = False + + def __aiter__(self): + return self + + async def __anext__(self): + message = await self.messages.get() + if message is None: + raise StopAsyncIteration + if isinstance(message, Exception): + raise message + return json.dumps(message) + + async def close(self): + self.closed = True + + +class Sink: + def __init__(self, logger): + self.logger = logger + self.events = [] + self.logs = 0 + + def store_message(self, message): + self.events.append(json.loads(message)) + + async def log_messages(self, *, wait_for_dispatch=False): + assert wait_for_dispatch + self.logs += 1 + self.logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True + + +def fixture(*, ready_timeout=1, lifetime=1): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + + async def hangup(): + assert not socket.closed + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 42}}) + + close_call = AsyncMock(side_effect=hangup) + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close_call, + ready_timeout=ready_timeout, + lifetime=lifetime, + drain_timeout=0.05, + ) + return socket, sink, close_call, supervisor + + +@pytest.mark.asyncio +async def test_observer_logs_webrtc_usage_without_client_sideband(): + socket, sink, close_call, supervisor = fixture() + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await socket.messages.put({"type": "response.done", "response": {"usage": {"total_tokens": 15}}}) + await socket.messages.put({"type": "session.closed", "usage": {"total_tokens": 19}}) + await supervisor.wait() + await supervisor.close() + assert sink.logs == 1 + assert sink.events[-1]["usage"]["total_tokens"] == 19 + assert sink.events[1]["response"]["usage"]["total_tokens"] == 15 + assert socket.closed + close_call.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_early_upstream_eof_rejects_start(): + socket, sink, close_call, supervisor = fixture() + await socket.messages.put(None) + with pytest.raises(RuntimeError, match="ended before"): + await supervisor.start() + assert socket.closed + assert sink.logs == 1 + close_call.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_cancelled_start_hangs_up_and_drains_terminal_usage(): + socket, sink, close_call, supervisor = fixture() + started = asyncio.create_task(supervisor.start()) + await asyncio.sleep(0) + started.cancel() + with pytest.raises(asyncio.CancelledError): + await started + close_call.assert_awaited_once() + assert socket.closed + assert sink.logs == 1 + assert sink.events[-1]["usage"]["total_tokens"] == 42 + + +@pytest.mark.asyncio +async def test_worker_shutdown_drains_all_calls(): + registry = CallSupervisors() + socket, sink, close_call, supervisor = fixture() + await socket.messages.put({"type": "session.created"}) + await registry.start(supervisor) + await registry.shutdown() + await registry.shutdown() + close_call.assert_awaited_once() + assert socket.closed + assert sink.logs == 1 + assert sink.events[-1]["usage"]["total_tokens"] == 42 + + +@pytest.mark.asyncio +async def test_ready_timeout_hangs_up_before_returning_error(): + socket, sink, close_call, supervisor = fixture(ready_timeout=0.01) + with pytest.raises(asyncio.TimeoutError): + await supervisor.start() + close_call.assert_awaited_once() + assert socket.closed + assert sink.logs == 1 + + +@pytest.mark.asyncio +async def test_lifetime_limit_closes_call_and_collects_final_usage(): + socket, sink, close_call, supervisor = fixture(lifetime=0.01) + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await supervisor.wait() + close_call.assert_awaited_once() + assert socket.closed + assert sink.events[-1]["usage"]["total_tokens"] == 42 + + +@pytest.mark.asyncio +async def test_socket_eof_after_ready_still_hangs_up_provider_call(monkeypatch): + from litellm.proxy.realtime_endpoints import call_supervision + + invalidate = AsyncMock() + release = AsyncMock() + monkeypatch.setattr(call_supervision, "invalidate_budget_reservation_counters", invalidate) + monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release) + socket, sink, close_call, supervisor = fixture() + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await socket.messages.put(None) + await supervisor.wait() + close_call.assert_awaited_once() + assert socket.closed + assert sink.logs == 1 + assert sink.logger.model_call_details["realtime_usage_incomplete"] is True + invalidate.assert_awaited_once() + release.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_observer_error_rejects_start(caplog): + socket, sink, close_call, supervisor = fixture() + await socket.messages.put(RuntimeError("private-provider-credential")) + with pytest.raises(RuntimeError, match="ended before"): + await supervisor.start() + assert socket.closed + close_call.assert_awaited_once() + assert "private-provider-credential" not in caplog.text + + +@pytest.mark.asyncio +async def test_failed_logging_releases_reservation(monkeypatch): + from litellm.proxy.realtime_endpoints import call_supervision + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = MagicMock() + sink.log_messages = AsyncMock(side_effect=RuntimeError("logging unavailable")) + release = AsyncMock() + monkeypatch.setattr(call_supervision, "release_or_invalidate_budget_reservation", release) + 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) + assert socket.closed + 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() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + dispatch_started = asyncio.Event() + dispatch_complete = asyncio.Event() + dispatch_finished = asyncio.Event() + + async def log_messages(*, wait_for_dispatch=False): + assert wait_for_dispatch + dispatch_started.set() + await dispatch_complete.wait() + logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True + dispatch_finished.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) + registry = CallSupervisors() + supervisor = CallSupervisor(socket, sink, logger, UserAPIKeyAuth(), hangup) + await socket.messages.put({"type": "session.created"}) + await registry.start(supervisor) + shutdown = asyncio.create_task(registry.shutdown()) + try: + await asyncio.wait_for(dispatch_started.wait(), timeout=1) + assert not shutdown.done() + assert not dispatch_finished.is_set() + finally: + dispatch_complete.set() + await asyncio.wait_for(shutdown, timeout=1) + assert dispatch_finished.is_set() + sink.log_messages.assert_awaited_once_with(wait_for_dispatch=True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("terminal_usage_required", [True, False]) +async def test_confirmed_hangup_without_terminal_usage_matches_protocol(terminal_usage_required): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + close_call = AsyncMock() + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close_call, + terminal_usage_required=terminal_usage_required, + drain_timeout=0.01, + ) + await socket.messages.put({"type": "session.created"}) + await supervisor.start() + await socket.messages.put({"type": "response.done", "response": {"usage": {"total_tokens": 17}}}) + await supervisor.close() + assert bool(logger.model_call_details.get("realtime_usage_incomplete")) == terminal_usage_required + assert sink.events[-1]["response"]["usage"]["total_tokens"] == 17 + assert sink.logs == 1 + assert socket.closed + close_call.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("closure", ["eof", "normal_close", "error"]) +@pytest.mark.parametrize("hangup_succeeds", [True, False]) +async def test_ga_observer_disconnect_requires_confirmed_hangup(closure, hangup_succeeds): + from websockets.exceptions import ConnectionClosedOK + from websockets.frames import Close + + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + sink = Sink(logger) + close_call = AsyncMock(side_effect=None if hangup_succeeds else RuntimeError("unconfirmed hangup")) + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close_call, + terminal_usage_required=False, + drain_timeout=0.01, + ) + await socket.messages.put({"type": "session.created"}) + await supervisor.start() + await socket.messages.put( + None + if closure == "eof" + else ConnectionClosedOK(Close(1000, ""), Close(1000, ""), True) + if closure == "normal_close" + else RuntimeError("observer failed") + ) + await supervisor.wait() + close_call.assert_awaited_once() + assert bool(logger.model_call_details.get("realtime_usage_incomplete")) == (not hangup_succeeds) + assert sink.logs == 1 + assert socket.closed diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 8f8a7640c08..692939d3ed0 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -4751,3 +4751,71 @@ def test_collect_and_combine_realtime_usage_stores_partitioned_text_tokens() -> assert combined.completion_tokens_details.reasoning_tokens == 95 assert combined.completion_tokens_details.text_tokens == 38 assert combined.completion_tokens_details.audio_tokens == 0 + + +def _live_terminal_event(duration=4000): + return {"type": "session.closed", "usage": {"audio_duration_ms": duration, "backend_model_usage": []}} + + +@pytest.mark.parametrize("rate,expected", [(0.025, 0.1), (0, 0), (None, 0)]) +def test_live_terminal_duration_uses_configured_second_price(monkeypatch, rate, expected): + monkeypatch.setitem( + litellm.model_cost, + "live-priced-test", + {"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": rate}, + ) + assert handle_realtime_stream_cost_calculation( + [_live_terminal_event()], Usage(), "chatgpt", "live-priced-test" + ) == pytest.approx(expected) + + +def test_live_terminal_duration_honors_deployment_override(monkeypatch): + monkeypatch.setitem( + litellm.model_cost, + "live-deployment-test", + {"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025}, + ) + result = RealtimeAPITokenUsageProcessor.create_logging_realtime_object(Usage(), [_live_terminal_event()]) + assert completion_cost( + completion_response=result, + model="gpt-live-1", + custom_llm_provider="chatgpt", + call_type="_arealtime", + custom_pricing=True, + router_model_id="live-deployment-test", + ) == pytest.approx(0.1) + + +@pytest.mark.parametrize("duration", [-1, True, "4000", float("inf"), float("nan"), None]) +def test_live_terminal_invalid_duration_does_not_create_spend(monkeypatch, duration): + monkeypatch.setitem( + litellm.model_cost, + "live-priced-test", + {"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025}, + ) + assert ( + handle_realtime_stream_cost_calculation( + [_live_terminal_event(duration)], Usage(), "chatgpt", "live-priced-test" + ) + == 0 + ) + + +def test_live_terminal_is_not_counted_twice(monkeypatch): + monkeypatch.setitem( + litellm.model_cost, + "live-priced-test", + {"litellm_provider": "chatgpt", "mode": "realtime", "input_cost_per_second": 0.025}, + ) + assert handle_realtime_stream_cost_calculation( + [_live_terminal_event(), _live_terminal_event()], Usage(), "chatgpt", "live-priced-test" + ) == pytest.approx(0.1) + assert ( + handle_realtime_stream_cost_calculation( + [{"type": "response.done", "response": {"usage": {}}}, _live_terminal_event()], + Usage(), + "chatgpt", + "live-priced-test", + ) + == 0 + )