fix(chatgpt): supervise call usage and preserve gateway query routing

This commit is contained in:
jibanez-staticduo 2026-09-10 12:52:04 +02:00
parent a9cd1dc56d
commit d6b8360ac3
No known key found for this signature in database
16 changed files with 1195 additions and 48 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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