From 642a2f196d1fce9e6d0c6788e53d0926661df61e Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 30 Sep 2026 23:51:42 +0000 Subject: [PATCH] feat(realtime): proxy OpenAI Live sessions on /v1/live/sessions Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- gateway/routes/allowlist.py | 2 + litellm/anthropic_beta_headers_config.json | 2 +- litellm/constants.py | 2 + litellm/cost_calculator.py | 214 ++++- litellm/llms/openai/live/__init__.py | 3 + litellm/llms/openai/live/handler.py | 478 +++++++++++ litellm/proxy/_types.py | 3 + .../auth/managed_authorization.py | 11 +- litellm/proxy/auth/auth_checks.py | 49 +- litellm/proxy/proxy_server.py | 241 ++++-- litellm/realtime_api/main.py | 98 ++- litellm/types/utils.py | 3 + tests/unit/llms/openai/live/__init__.py | 0 tests/unit/llms/openai/live/test_handler.py | 762 ++++++++++++++++++ tests/unit/proxy/test_proxy_server.py | 249 +++++- tests/unit/realtime_api/test_main.py | 151 ++++ .../test_anthropic_beta_headers_filtering.py | 2 +- tests/unit/test_cost_calculator.py | 261 ++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 114 +++ 19 files changed, 2563 insertions(+), 82 deletions(-) create mode 100644 litellm/llms/openai/live/__init__.py create mode 100644 litellm/llms/openai/live/handler.py create mode 100644 tests/unit/llms/openai/live/__init__.py create mode 100644 tests/unit/llms/openai/live/test_handler.py diff --git a/gateway/routes/allowlist.py b/gateway/routes/allowlist.py index 6e91f5486d0..a76d9a12b81 100644 --- a/gateway/routes/allowlist.py +++ b/gateway/routes/allowlist.py @@ -116,6 +116,8 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = ( # Realtime / streaming "/v1/realtime", "/realtime", + "/v1/live/sessions", + "/live/sessions", # Health & ops "/health", "/metrics", diff --git a/litellm/anthropic_beta_headers_config.json b/litellm/anthropic_beta_headers_config.json index b57239f8699..7bb4c6e58df 100644 --- a/litellm/anthropic_beta_headers_config.json +++ b/litellm/anthropic_beta_headers_config.json @@ -49,7 +49,7 @@ "computer-use-2025-11-24": "computer-use-2025-11-24", "context-1m-2025-08-07": "context-1m-2025-08-07", "context-management-2025-06-27": "context-management-2025-06-27", - "dangerous-tool-use-2026-09-03": "dangerous-tool-use-2026-09-03", + "dangerous-tool-use-2026-09-03": null, "effort-2025-11-24": "effort-2025-11-24", "fast-mode-2026-02-01": null, "files-api-2025-04-14": "files-api-2025-04-14", diff --git a/litellm/constants.py b/litellm/constants.py index 9af40744896..68761914754 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -355,6 +355,8 @@ BEDROCK_REALTIME_SDK_SUPPORTED_RANGE: Final = ">=0.10.0,<0.12.0" CLIENT_REQUESTED_MODEL_SCOPE_KEY: Final = "litellm.client_requested_model" MODEL_GROUP_ALIAS_RESOLVED_SCOPE_KEY: Final = "litellm.model_group_alias_resolved" REALTIME_SESSION_SUCCESS_LOGGED_KEY: Final = "realtime_session_success_logged" +OPENAI_LIVE_SESSION_START_TIMEOUT_SECONDS: Final = 30 +OPENAI_LIVE_SESSION_CLOSE_TIMEOUT_SECONDS: Final = 5 REALTIME_SESSION_FAILURE_LOGGED_KEY: Final = "realtime_session_failure_logged" # SSL/TLS cipher configuration for faster handshakes diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 238b7cc3fdd..3d9ae26646e 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -1,8 +1,10 @@ # What is this? ## File for 'response_cost' calculation in Logging import logging +import math import time from collections.abc import Mapping, Sequence +from datetime import datetime from functools import lru_cache from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, cast @@ -106,7 +108,6 @@ from litellm.types.llms.openai import ( OpenAIModerationResponse, OpenAIRealtimeStreamList, OpenAIRealtimeStreamResponseBaseObject, - OpenAIRealtimeStreamSessionEvents, ResponseAPIUsage, ResponsesAPIResponse, ) @@ -2880,6 +2881,145 @@ class ResponsesWebSocketTokenUsageProcessor(BaseTokenUsageProcessor): _TRANSCRIPTION_COMPLETED_EVENT_TYPE: Final = "conversation.item.input_audio_transcription.completed" +_LIVE_TERMINAL_RESPONSE_EVENT_TYPES: Final = frozenset({"response.completed", "response.incomplete", "response.failed"}) + + +class _LiveResponsePayload(BaseModel): + id: str | None = None + model: str | None = None + usage: Mapping[str, object] | None = None + + +class _LiveResponseEvent(BaseModel): + type: str = "" + response: _LiveResponsePayload | None = None + + +class _LiveResponseEnvelope(BaseModel): + type: str = "" + event: _LiveResponseEvent | None = None + + +def _live_response_payload(result: Mapping[str, object]) -> _LiveResponsePayload | None: + try: + envelope: Final = _LiveResponseEnvelope.model_validate(result) + except Exception: + return None + if ( + envelope.type != "response.event" + or envelope.event is None + or envelope.event.type not in _LIVE_TERMINAL_RESPONSE_EVENT_TYPES + or envelope.event.response is None + or envelope.event.response.usage is None + ): + return None + return envelope.event.response + + +def _unique_live_response_payloads(results: Sequence[Mapping[str, object]]) -> tuple[_LiveResponsePayload, ...]: + response_payloads: Final = tuple( + response for result in results if (response := _live_response_payload(result)) is not None + ) + response_ids: Final = tuple(dict.fromkeys(response.id for response in response_payloads if response.id is not None)) + return tuple( + next(response for response in response_payloads if response.id == response_id) for response_id in response_ids + ) + + +def _live_session_model(result: Mapping[str, object]) -> str | None: + session: Final = result.get("session") + if not isinstance(session, Mapping): + return None + model: Final = session.get("model") + return model if isinstance(model, str) else None + + +def _live_usage_seconds(result: Mapping[str, object]) -> float | None: + usage: Final = result.get("usage") + if not isinstance(usage, Mapping): + return None + seconds: Final = usage.get("seconds") + if isinstance(seconds, (int, float)) and not isinstance(seconds, bool) and math.isfinite(seconds): + return float(seconds) + return None + + +def _last_live_usage_seconds(results: Sequence[Mapping[str, object]], event_type: str) -> float | None: + return next( + ( + seconds + for result in reversed(results) + if result.get("type") == event_type and (seconds := _live_usage_seconds(result)) is not None + ), + None, + ) + + +def _live_session_wall_clock_seconds(litellm_logging_obj: LitellmLoggingObject | None) -> float: + if litellm_logging_obj is None: + return 0.0 + model_call_details: Final = litellm_logging_obj.model_call_details + start_time: Final = model_call_details.get("start_time") + end_time: Final = model_call_details.get("end_time") + if isinstance(start_time, datetime) and isinstance(end_time, datetime): + return max((end_time - start_time).total_seconds(), 0.0) + return 0.0 + + +def _live_input_cost_per_second(model_name: str, custom_llm_provider: str) -> float | None: + try: + model_info: Final = _cached_get_model_info_helper( + model=model_name, + custom_llm_provider=custom_llm_provider, + ) + except Exception: + return None + rate: Final = model_info.get("input_cost_per_second") + if isinstance(rate, (int, float)) and not isinstance(rate, bool) and math.isfinite(rate): + return float(rate) + return None + + +def _first_live_input_cost_per_second( + potential_model_names: Sequence[str | None], + custom_llm_provider: str, +) -> float: + return next( + ( + rate + for model_name in potential_model_names + if model_name is not None + and (rate := _live_input_cost_per_second(model_name, custom_llm_provider)) is not None + ), + 0.0, + ) + + +def _live_response_token_costs( + response: _LiveResponsePayload, + custom_llm_provider: str, + data_residency: str | None, +) -> tuple[float, float]: + if response.id is None or response.model is None or response.usage is None: + verbose_logger.debug("Skipping delegated Live response without an id, model, or usage") + return 0.0, 0.0 + try: + usage: Final = get_usage_object(response.model_dump()) + except Exception as error: + verbose_logger.debug("Could not transform delegated Live response usage: %s", error) + return 0.0, 0.0 + if usage is None: + return 0.0, 0.0 + costs: Final = _candidate_realtime_token_costs( + model_name=response.model, + combined_usage_object=usage, + custom_llm_provider=custom_llm_provider, + data_residency=data_residency, + ) + if costs is None: + verbose_logger.debug("Skipping unpriced delegated Live response model=%s", response.model) + return 0.0, 0.0 + return costs def _candidate_realtime_token_costs( @@ -2968,14 +3108,21 @@ def handle_realtime_stream_cost_calculation( base_pricing_model: the deployment's resolved base_model, tried ahead of the session-reported model but after custom rates """ - received_model = None - potential_model_names: Final = [custom_pricing_model, base_pricing_model] - for result in results: - if result["type"] == "session.created": - received_model = cast(OpenAIRealtimeStreamSessionEvents, result)["session"].get("model", None) - potential_model_names.append(received_model) - - potential_model_names.append(litellm_model_name) + received_model: Final = next( + ( + model_name + for result in results + if result.get("type") in ("session.created", "session.started") + and (model_name := _live_session_model(result)) is not None + ), + None, + ) + potential_model_names: Final = ( + custom_pricing_model, + base_pricing_model, + received_model, + litellm_model_name, + ) input_cost_per_token, output_cost_per_token = _first_priced_realtime_token_costs( potential_model_names=potential_model_names, combined_usage_object=combined_usage_object, @@ -2992,15 +3139,56 @@ 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_lifecycle_event_types: Final = frozenset({"session.started", "session.usage.updated", "session.closed"}) + live_session_started: Final = any(result.get("type") == "session.started" for result in results) + live_session_closed: Final = any(result.get("type") == "session.closed" for result in results) + live_event_seen: Final = any(result.get("type") in live_lifecycle_event_types for result in results) + closed_seconds: Final = _last_live_usage_seconds(results, "session.closed") + reported_seconds: Final = _last_live_usage_seconds(results, "session.usage.updated") + latest_reported_seconds: Final = 0.0 if reported_seconds is None else reported_seconds + voice_seconds: Final = ( + closed_seconds + if closed_seconds is not None + else max( + latest_reported_seconds, + _live_session_wall_clock_seconds(litellm_logging_obj), + ) + if live_session_started and not live_session_closed + else latest_reported_seconds + ) + live_model_names: Final = ( + (custom_pricing_model, base_pricing_model, received_model, litellm_model_name) if live_event_seen else () + ) + live_voice_cost: Final = voice_seconds * _first_live_input_cost_per_second( + live_model_names, + custom_llm_provider, + ) + live_response_costs: Final = tuple( + _live_response_token_costs( + response=response, + custom_llm_provider=custom_llm_provider, + data_residency=data_residency, + ) + for response in _unique_live_response_payloads(results) + ) + live_response_input_cost: Final = sum(costs[0] for costs in live_response_costs) + live_response_output_cost: Final = sum(costs[1] for costs in live_response_costs) + input_cost: Final = input_cost_per_token + live_response_input_cost + output_cost: Final = output_cost_per_token + live_response_output_cost + total_cost: Final = input_cost + output_cost + transcription_cost + live_voice_cost + additional_costs: Final = tuple( + (cost_name, cost) + for cost_name, cost in (("transcription_cost", transcription_cost), ("live_voice_cost", live_voice_cost)) + if cost > 0 + ) _store_cost_breakdown_in_logging_obj( litellm_logging_obj=litellm_logging_obj, - prompt_tokens_cost_usd_dollar=input_cost_per_token, - completion_tokens_cost_usd_dollar=output_cost_per_token, + prompt_tokens_cost_usd_dollar=input_cost, + completion_tokens_cost_usd_dollar=output_cost, 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=dict(additional_costs) if additional_costs else None, data_residency=data_residency, ) diff --git a/litellm/llms/openai/live/__init__.py b/litellm/llms/openai/live/__init__.py new file mode 100644 index 00000000000..1aa84ccd6cb --- /dev/null +++ b/litellm/llms/openai/live/__init__.py @@ -0,0 +1,3 @@ +from .handler import OpenAILiveSessions + +__all__ = ("OpenAILiveSessions",) diff --git a/litellm/llms/openai/live/handler.py b/litellm/llms/openai/live/handler.py new file mode 100644 index 00000000000..6196605683a --- /dev/null +++ b/litellm/llms/openai/live/handler.py @@ -0,0 +1,478 @@ +import asyncio +import json +import math +import ssl as ssl_module +from collections.abc import AsyncIterator, Mapping +from contextlib import AbstractAsyncContextManager, suppress +from itertools import count +from types import MappingProxyType +from typing import Final, Protocol, cast +from urllib.parse import urlsplit, urlunsplit + +from pydantic import TypeAdapter, ValidationError +from typing_extensions import NotRequired, ReadOnly, TypedDict + +from litellm._logging import redact_secrets +from litellm.constants import ( + OPENAI_LIVE_SESSION_CLOSE_TIMEOUT_SECONDS, + REALTIME_SESSION_SUCCESS_LOGGED_KEY, + REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, +) +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER, LoggingWorker +from litellm.litellm_core_utils.realtime_errors import ( + close_after_upstream_handshake_refusal, + realtime_error_event, + websocket_close_reason, +) +from litellm.litellm_core_utils.realtime_streaming import backend_close_from +from litellm.llms.openai.realtime.handler import OpenAIRealtime +from litellm.types.llms.openai import OpenAIRealtimeStreamList +from litellm.types.utils import LiteLLMRealtimeStreamLoggingObject, Usage + +_LIVE_EVENT_FRAME_ADAPTER: Final = TypeAdapter(dict[str, object]) +_LIVE_TERMINAL_RESPONSE_EVENT_TYPES: Final = frozenset({"response.completed", "response.incomplete", "response.failed"}) +_LIVE_USAGE_EVENT_TYPES: Final = frozenset({"session.usage.updated", "session.closed"}) +_LIVE_WEBSOCKET_SCHEMES: Final = MappingProxyType({"https": "wss", "http": "ws"}) +_LIVE_SESSION_CLOSE_FRAME: Final = '{"type": "session.close"}' + + +class _LiveSessionModel(TypedDict): + model: ReadOnly[str] + + +class _LiveSessionStartedEvent(TypedDict): + type: ReadOnly[str] + session: NotRequired[ReadOnly[_LiveSessionModel]] + + +class _LiveUsageSeconds(TypedDict): + seconds: ReadOnly[object] + + +class _LiveUsageEvent(TypedDict): + type: ReadOnly[str] + usage: NotRequired[ReadOnly[_LiveUsageSeconds]] + + +class _LiveResponseSummary(TypedDict): + id: ReadOnly[str] + model: ReadOnly[object] + usage: ReadOnly[object] + + +class _LiveTerminalResponseEvent(TypedDict): + type: ReadOnly[str] + response: ReadOnly[_LiveResponseSummary] + + +class _LiveResponseEvent(TypedDict): + type: ReadOnly[str] + event: ReadOnly[_LiveTerminalResponseEvent] + + +class _LivePreCallInput(TypedDict): + session_start: ReadOnly[Mapping[str, object]] + + +class _LivePreCallArgs(TypedDict): + api_base: ReadOnly[str] + headers: ReadOnly[Mapping[str, str]] + complete_input_dict: ReadOnly[_LivePreCallInput] + + +class LiveClientWebSocket(Protocol): + async def receive_text(self) -> str: ... + + async def send_text(self, data: str) -> None: ... + + async def close(self, code: int = 1000, reason: str | None = None) -> None: ... + + +class LiveBackendWebSocket(Protocol): + async def send(self, data: str) -> None: ... + + async def recv(self) -> str | bytes: ... + + async def close(self) -> None: ... + + +class LiveWebSocketConnector(Protocol): + def __call__( + self, + url: str, + *, + additional_headers: Mapping[str, str], + max_size: int | None, + ssl: bool | str | ssl_module.SSLContext | None, + ) -> AbstractAsyncContextManager[LiveBackendWebSocket]: ... + + +def _openai_live_websocket_connect( + url: str, + *, + additional_headers: Mapping[str, str], + max_size: int | None, + ssl: bool | str | ssl_module.SSLContext | None, +) -> AbstractAsyncContextManager[LiveBackendWebSocket]: + import websockets + + return websockets.connect( + url, + additional_headers=additional_headers, + max_size=max_size, + ssl=ssl, + ) + + +def _is_finite_usage_seconds(value: object) -> bool: + return ( + isinstance(value, (int, float)) + and not isinstance(value, bool) + and (not isinstance(value, float) or math.isfinite(value)) + ) + + +def _session_started_event(event: Mapping[str, object]) -> _LiveSessionStartedEvent: + session: Final = event.get("session") + session_model: Final = session.get("model") if isinstance(session, Mapping) else None + if isinstance(session_model, str): + started_with_model: Final[_LiveSessionStartedEvent] = { + "type": "session.started", + "session": {"model": session_model}, + } + return started_with_model + started: Final[_LiveSessionStartedEvent] = {"type": "session.started"} + return started + + +def _usage_event(event: Mapping[str, object], event_type: str) -> _LiveUsageEvent: + usage: Final = event.get("usage") + seconds: Final = usage.get("seconds") if isinstance(usage, Mapping) else None + if _is_finite_usage_seconds(seconds): + usage_with_seconds: Final[_LiveUsageEvent] = {"type": event_type, "usage": {"seconds": seconds}} + return usage_with_seconds + usage_without_seconds: Final[_LiveUsageEvent] = {"type": event_type} + return usage_without_seconds + + +def _terminal_response_event(event: Mapping[str, object]) -> _LiveResponseEvent | None: + nested_event: Final = event.get("event") + if not isinstance(nested_event, Mapping): + return None + nested_event_type: Final = nested_event.get("type") + if not isinstance(nested_event_type, str) or nested_event_type not in _LIVE_TERMINAL_RESPONSE_EVENT_TYPES: + return None + response: Final = nested_event.get("response") + response_id: Final = response.get("id") if isinstance(response, Mapping) else None + if not isinstance(response_id, str) or not isinstance(response, Mapping): + return None + response_event: Final[_LiveResponseEvent] = { + "type": "response.event", + "event": { + "type": nested_event_type, + "response": { + "id": response_id, + "model": response.get("model"), + "usage": response.get("usage"), + }, + }, + } + return response_event + + +def _metering_event(message: str) -> Mapping[str, object] | None: + try: + event: Final = _LIVE_EVENT_FRAME_ADAPTER.validate_json(message) + except ValidationError: + return None + event_type: Final = event.get("type") + if event_type == "session.started": + return _session_started_event(event) + if isinstance(event_type, str) and event_type in _LIVE_USAGE_EVENT_TYPES: + return _usage_event(event, event_type) + if event_type == "response.event": + return _terminal_response_event(event) + return None + + +def _is_session_close_message(message: str) -> bool: + try: + event: Final = _LIVE_EVENT_FRAME_ADAPTER.validate_json(message) + except ValidationError: + return False + return event.get("type") == "session.close" + + +_RetainedEvent = tuple[int, Mapping[str, object]] + + +class _LiveEventRetention: + def __init__(self) -> None: + self._event_sequence = count() + self._session_started: _RetainedEvent | None = None + self._latest_usage: _RetainedEvent | None = None + self._session_closed: _RetainedEvent | None = None + self._response_events: dict[str, _RetainedEvent] = {} # mutable-ok: one insert per response id, no copy + + def retain(self, event: Mapping[str, object]) -> None: + retained_item: Final = (next(self._event_sequence), event) + event_type: Final = event.get("type") + if event_type == "session.started": + if self._session_started is None: + self._session_started = retained_item + return + if event_type == "session.usage.updated": + usage: Final = event.get("usage") + seconds: Final = usage.get("seconds") if isinstance(usage, Mapping) else None + if _is_finite_usage_seconds(seconds): + self._latest_usage = retained_item + return + if event_type == "session.closed": + self._session_closed = retained_item + return + nested_event: Final = event.get("event") + response: Final = nested_event.get("response") if isinstance(nested_event, Mapping) else None + response_id: Final = response.get("id") if isinstance(response, Mapping) else None + if isinstance(response_id, str): + self._response_events.setdefault(response_id, retained_item) + + def results(self) -> tuple[Mapping[str, object], ...]: + retained_items: Final = tuple( + item for item in (self._session_started, self._latest_usage, self._session_closed) if item is not None + ) + tuple(self._response_events.values()) + return tuple(event for _, event in sorted(retained_items, key=lambda item: item[0])) + + +async def _forward_live_client_frames( + *, + websocket: LiveClientWebSocket, + backend_websocket: LiveBackendWebSocket, + client_gone: asyncio.Event, + client_close_sent: asyncio.Event, +) -> None: + async for message in _live_client_messages(websocket, client_gone): + await backend_websocket.send(message) + if _is_session_close_message(message): + client_close_sent.set() + + +async def _live_client_messages( + websocket: LiveClientWebSocket, + client_gone: asyncio.Event, +) -> AsyncIterator[str]: + from starlette.websockets import WebSocketDisconnect + + try: + while True: + yield await websocket.receive_text() + except WebSocketDisconnect: + client_gone.set() + + +async def _forward_live_backend_frames( + *, + websocket: LiveClientWebSocket, + backend_websocket: LiveBackendWebSocket, + client_gone: asyncio.Event, + event_retention: _LiveEventRetention, +) -> None: + from starlette.websockets import WebSocketDisconnect + from websockets.exceptions import ConnectionClosed + + async for message in _live_backend_messages(websocket, backend_websocket): + if (event := _metering_event(message)) is not None: + event_retention.retain(event) + if client_gone.is_set(): + continue + try: + await websocket.send_text(message) + except (WebSocketDisconnect, RuntimeError, ConnectionClosed, OSError): + client_gone.set() + + +async def _live_backend_messages( + websocket: LiveClientWebSocket, + backend_websocket: LiveBackendWebSocket, +) -> AsyncIterator[str]: + from starlette.websockets import WebSocketDisconnect + from websockets.exceptions import ConnectionClosed + + try: + while True: + yield _decode_live_websocket_frame(await backend_websocket.recv()) + except ConnectionClosed as error: + upstream_close: Final = backend_close_from(error) + with suppress(WebSocketDisconnect, RuntimeError, ConnectionClosed, OSError): + await websocket.close(code=upstream_close.code, reason=upstream_close.reason) + + +def _decode_live_websocket_frame(frame: str | bytes) -> str: + return frame.decode("utf-8") if isinstance(frame, bytes) else frame + + +class OpenAILiveSessions(OpenAIRealtime): + def __init__( + self, + logging_worker: LoggingWorker = GLOBAL_LOGGING_WORKER, + websocket_connector: LiveWebSocketConnector | None = None, + ) -> None: + super().__init__() + self._logging_worker = logging_worker + self._websocket_connector = ( + websocket_connector if websocket_connector is not None else _openai_live_websocket_connect + ) + + @staticmethod + def _construct_live_url(api_base: str) -> str: + parsed_url: Final = urlsplit(api_base) + websocket_scheme: Final = _LIVE_WEBSOCKET_SCHEMES.get(parsed_url.scheme, parsed_url.scheme) + return urlunsplit((websocket_scheme, parsed_url.netloc, "/v1/live/sessions", "", "")) + + async def async_live_session( + self, + *, + model: str, + websocket: LiveClientWebSocket, + logging_obj: LiteLLMLogging, + session_start: Mapping[str, object], + api_base: str | None, + api_key: str | None, + ) -> None: + if api_key is None: + raise ValueError("api_key is required for OpenAI Live session calls") + from websockets.exceptions import InvalidStatus + + resolved_api_base: Final = api_base or self._get_default_api_base() + url: Final = self._construct_live_url(resolved_api_base) + headers: Final = self._get_additional_headers(api_key) + ssl_config: Final = self._get_ssl_config(url) + session: Final = session_start.get("session") + if not isinstance(session, Mapping): + raise ValueError("session_start.session must be a mapping") + rewritten_start: Final = { # mutable-ok: json.dumps frame + **session_start, + "session": {**session, "model": model}, # mutable-ok: json.dumps frame + } + upstream_start: Final = json.dumps(rewritten_start) + pre_call_args: Final[_LivePreCallArgs] = { + "api_base": url, + "headers": headers, + "complete_input_dict": {"session_start": rewritten_start}, + } + event_retention: Final = _LiveEventRetention() + try: + logging_obj.pre_call(input=None, api_key=api_key, additional_args=pre_call_args) + async with self._websocket_connector( + url, + additional_headers=headers, + max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, + ssl=ssl_config, + ) as backend_websocket: + upstream_connected: Final = True + try: + await backend_websocket.send(upstream_start) + await self._bidirectional_forward( + websocket=websocket, + backend_websocket=backend_websocket, + event_retention=event_retention, + ) + finally: + if upstream_connected: + retained_events: Final = event_retention.results() + logging_result: Final = LiteLLMRealtimeStreamLoggingObject( + usage=Usage(), + results=cast(OpenAIRealtimeStreamList, list(retained_events)), # mutable-ok: list field + ) + try: + self._logging_worker.ensure_initialized_and_enqueue( + logging_obj.dispatch_success_handlers(logging_result, prefer_async_handlers=True) + ) + finally: + logging_obj.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True + except InvalidStatus as error: + await close_after_upstream_handshake_refusal(websocket, error.response.status_code) + except Exception as error: + redacted_error: Final = redact_secrets(str(error)) + with suppress(Exception): + await websocket.send_text(realtime_error_event(redacted_error, error_type="server_error")) + try: + await websocket.close( + code=1011, + reason=websocket_close_reason(redacted_error, fallback="Internal server error"), + ) + except RuntimeError as close_error: + if "already completed" not in str(close_error) and "websocket.close" not in str(close_error): + raise Exception(f"Unexpected error while closing WebSocket: {close_error}") from close_error + + async def _bidirectional_forward( + self, + *, + websocket: LiveClientWebSocket, + backend_websocket: LiveBackendWebSocket, + event_retention: _LiveEventRetention, + ) -> None: + from websockets.exceptions import ConnectionClosed + + client_gone: Final = asyncio.Event() + client_close_sent: Final = asyncio.Event() + client_forward_task: Final = asyncio.create_task( + _forward_live_client_frames( + websocket=websocket, + backend_websocket=backend_websocket, + client_gone=client_gone, + client_close_sent=client_close_sent, + ) + ) + backend_forward_task: Final = asyncio.create_task( + _forward_live_backend_frames( + websocket=websocket, + backend_websocket=backend_websocket, + client_gone=client_gone, + event_retention=event_retention, + ) + ) + client_gone_task: Final = asyncio.create_task(client_gone.wait()) + forwarding_tasks: Final = (client_forward_task, backend_forward_task, client_gone_task) + try: + completed_tasks: Final = await asyncio.wait( + forwarding_tasks, + return_when=asyncio.FIRST_COMPLETED, + ) + done_tasks: Final = completed_tasks[0] + if backend_forward_task in done_tasks: + backend_forward_task.result() + return + if client_forward_task in done_tasks: + client_forward_task.result() + elif client_gone_task in done_tasks: + client_forward_task.cancel() + event_loop: Final = asyncio.get_running_loop() + close_deadline: Final = event_loop.time() + OPENAI_LIVE_SESSION_CLOSE_TIMEOUT_SECONDS + if not client_close_sent.is_set(): + remaining_send_timeout: Final = max(close_deadline - event_loop.time(), 0) + with suppress(ConnectionClosed, TimeoutError, asyncio.TimeoutError): + await asyncio.wait_for( + backend_websocket.send(_LIVE_SESSION_CLOSE_FRAME), + timeout=remaining_send_timeout, + ) + remaining_wait_timeout: Final = max(close_deadline - event_loop.time(), 0) + close_wait_result: Final = await asyncio.wait( + (backend_forward_task,), + timeout=remaining_wait_timeout, + ) + remaining_close_timeout: Final = max(close_deadline - event_loop.time(), 0) + with suppress(TimeoutError, asyncio.TimeoutError): + await asyncio.wait_for( + backend_websocket.close(), + timeout=remaining_close_timeout, + ) + if backend_forward_task in close_wait_result[0]: + backend_forward_task.result() + elif not backend_forward_task.done(): + backend_forward_task.cancel() + finally: + for task in forwarding_tasks: + if not task.done(): + task.cancel() + await asyncio.gather(*forwarding_tasks, return_exceptions=True) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d2fad212dd9..3bf859c722e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -419,6 +419,9 @@ class LiteLLMRoutes(enum.Enum): "/realtime", "/v1/realtime", "/openai/v1/realtime", + "/live/sessions", + "/v1/live/sessions", + "/openai/v1/live/sessions", "/realtime?{model}", "/v1/realtime?{model}", "/openai/v1/realtime?{model}", diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index 17d988127ec..03cf901b2da 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -11,7 +11,16 @@ from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failur from litellm.types.agents import AgentResponse from litellm.types.proxy.agent_identity import AgentIdentityFailure, ManagedAgentContext -_MANAGED_REALTIME_ROUTES: Final = frozenset(("/realtime", "/v1/realtime", "/openai/v1/realtime")) +_MANAGED_REALTIME_ROUTES: Final = frozenset( + ( + "/realtime", + "/v1/realtime", + "/openai/v1/realtime", + "/live/sessions", + "/v1/live/sessions", + "/openai/v1/live/sessions", + ) +) _MANAGED_MODEL_ROUTES: Final = frozenset( f"{prefix}/{operation}" for prefix, operation in product( diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3ec430332ee..d99d8b48999 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1059,22 +1059,7 @@ async def common_checks( code=status.HTTP_400_BAD_REQUEST, ) - managed_policy: Final = managed_agent_policy(valid_token) - if _model and valid_token is not None and managed_policy is not None: - managed_models: Final = (managed_policy.object_permission or MappingProxyType({})).get("models", ()) - if not isinstance(managed_models, (list, tuple)) or not managed_models: - raise HTTPException(403, "This agent has no model grants") - _can_object_call_model( - model=_resolve_team_alias(_model, valid_token.team_model_aliases, valid_token.team_id, llm_router), - llm_router=llm_router, - models=list(managed_models), - team_id=valid_token.team_id, - object_type="agent", - key_model_aliases=key_model_aliases_for_auth_check(valid_token), - ) - - await _check_agent_access_group_model_access(model=_model, valid_token=valid_token, llm_router=llm_router) - await _check_agent_caller_model_access( + await can_agent_call_model( model=_model, valid_token=valid_token, llm_router=llm_router, @@ -4545,6 +4530,38 @@ def _live_team_alias_target( return model if deleted_team_deployment else target +async def can_agent_call_model( + model: str | list[str] | None, # mutable-ok: the model checks it delegates to take list[str] + valid_token: UserAPIKeyAuth | None, + llm_router: Router | None, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + proxy_logging_obj: ProxyLogging, +) -> None: + managed_policy: Final = managed_agent_policy(valid_token) + if model and valid_token is not None and managed_policy is not None: + managed_models: Final = (managed_policy.object_permission or MappingProxyType({})).get("models", ()) + if not isinstance(managed_models, (list, tuple)) or not managed_models: + raise HTTPException(403, "This agent has no model grants") + _can_object_call_model( + model=_resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router), + llm_router=llm_router, + models=list(managed_models), + team_id=valid_token.team_id, + object_type="agent", + key_model_aliases=key_model_aliases_for_auth_check(valid_token), + ) + await _check_agent_access_group_model_access(model=model, valid_token=valid_token, llm_router=llm_router) + await _check_agent_caller_model_access( + model=model, + valid_token=valid_token, + llm_router=llm_router, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + async def _check_agent_access_group_model_access( model: str | list[str] | None, # mutable-ok: _can_object_call_model and the client message helper take list[str] valid_token: UserAPIKeyAuth | None, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0e151199f41..38ec6e1fb88 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -32,6 +32,7 @@ from itertools import chain from types import MappingProxyType, UnionType from typing import ( TYPE_CHECKING, + Annotated, Any, Final, Literal, @@ -74,6 +75,7 @@ from litellm.constants import ( LITELLM_SETTINGS_SAFE_DB_OVERRIDES, LITELLM_UI_ALLOW_HEADERS, LITELLM_UI_SESSION_DURATION, + OPENAI_LIVE_SESSION_START_TIMEOUT_SECONDS, RUNTIME_UPDATABLE_ROUTER_SETTINGS, ) from litellm.litellm_core_utils.asyncify import asyncify @@ -342,6 +344,7 @@ from litellm.proxy.analytics_endpoints.analytics_endpoints import ( from litellm.proxy.auth.auth_checks import ( ROLE_BASED_PERMISSIONS_ADAPTER, ExperimentalUIJWTToken, + can_agent_call_model, can_key_call_resolved_model, get_team_object, log_db_metrics, @@ -12839,6 +12842,17 @@ async def _release_realtime_max_parallel_slot(user_api_key_dict: UserAPIKeyAuth) await release_like_http_disconnect(user_api_key_dict) +class _RealtimeErrorBody(TypedDict): + type: ReadOnly[str] + message: ReadOnly[str] + code: NotRequired[ReadOnly[str]] + + +class _RealtimeErrorEvent(TypedDict): + type: ReadOnly[str] + error: ReadOnly[_RealtimeErrorBody] + + async def _reject_realtime_session( websocket: WebSocket, user_api_key_dict: UserAPIKeyAuth, @@ -12846,13 +12860,19 @@ async def _reject_realtime_session( code: int, reason: str, error_message: str | None = None, + error_type: str = "guardrail_error", + error_code: str | None = None, ) -> None: try: if error_message is not None: try: - await websocket.send_text( - json.dumps({"type": "error", "error": {"type": "guardrail_error", "message": error_message}}) - ) + error_event: Final[_RealtimeErrorEvent] = { + "type": "error", + "error": {"type": error_type, "message": error_message} + if error_code is None + else {"type": error_type, "message": error_message, "code": error_code}, + } + await websocket.send_text(json.dumps(error_event)) except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below verbose_proxy_logger.debug("Could not send realtime pre-call error event to client; closing anyway") await websocket.close(code=code, reason=reason) @@ -12861,67 +12881,36 @@ async def _reject_realtime_session( await _release_realtime_max_parallel_slot(user_api_key_dict) -@app.websocket("/openai/v1/realtime") -@app.websocket("/v1/realtime") -@app.websocket("/realtime") -async def realtime_websocket_endpoint( +async def _route_realtime_websocket_session( websocket: WebSocket, - model: str | None = fastapi.Query(None, description="The model to use for the websocket connection."), - intent: str | None = fastapi.Query(None, description="The intent of the websocket connection."), - guardrails: str | None = fastapi.Query( - None, - description="Comma-separated list of guardrail names to apply to this request.", - ), - user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), -): - requested_protocols: Final = [ - p.strip() for p in (websocket.headers.get("sec-websocket-protocol") or "").split(",") if p.strip() - ] - accept_kwargs: Final[dict] = {} - if requested_protocols: - accept_kwargs["subprotocol"] = requested_protocols[0] - - route_model = model - if route_model is None: - if intent == "transcription": - route_model = "gpt-realtime-whisper" - else: - await _reject_realtime_session( - websocket, user_api_key_dict, code=1008, reason="model query parameter is required" - ) - return - assert route_model is not None - try: - await can_key_call_resolved_model( - model=route_model, - llm_model_list=llm_model_list, - valid_token=user_api_key_dict, - llm_router=llm_router, - ) - except ProxyException as e: - _log_model_access_denial(e) - await _reject_realtime_session(websocket, user_api_key_dict, code=1008, reason=e.message[:120]) - return - await websocket.accept(**accept_kwargs) - + user_api_key_dict: UserAPIKeyAuth, + *, + route_model: str, + model: str | None, + intent: str | None, + guardrails: str | None, + request_path: str | None = None, + live_session_start: Mapping[str, object] | None = None, +) -> None: # Only use explicit parameters, not all query params query_params: Final = cast(RealtimeQueryParams, dict(_realtime_query_params_template(model, intent))) - data: dict[str, object] = { "model": route_model, "websocket": websocket, "query_params": query_params, # Only explicit params } - # Pass guardrails into data so pre-call guardrail processing picks them up if guardrails: data["guardrails"] = [g.strip() for g in guardrails.split(",") if g.strip()] + if live_session_start is not None: + data["live_session_start"] = live_session_start # Use raw ASGI headers (already lowercase bytes) to avoid extra work headers_list: Final = list(websocket.scope.get("headers") or []) - scope: Final = REALTIME_REQUEST_SCOPE_TEMPLATE.copy() scope["headers"] = headers_list + if request_path is not None: + scope["path"] = request_path request: Final = Request(scope=scope) @@ -13001,6 +12990,162 @@ async def realtime_websocket_endpoint( await _release_realtime_max_parallel_slot(user_api_key_dict) +@app.websocket("/openai/v1/realtime") +@app.websocket("/v1/realtime") +@app.websocket("/realtime") +async def realtime_websocket_endpoint( + websocket: WebSocket, + model: str | None = fastapi.Query(None, description="The model to use for the websocket connection."), + intent: str | None = fastapi.Query(None, description="The intent of the websocket connection."), + guardrails: str | None = fastapi.Query( + None, + description="Comma-separated list of guardrail names to apply to this request.", + ), + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), +): + requested_protocols: Final = tuple( + protocol.strip() + for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",") + if protocol.strip() + ) + + route_model = model + if route_model is None: + if intent == "transcription": + route_model = "gpt-realtime-whisper" + else: + await _reject_realtime_session( + websocket, user_api_key_dict, code=1008, reason="model query parameter is required" + ) + return + assert route_model is not None + try: + await can_key_call_resolved_model( + model=route_model, + llm_model_list=llm_model_list, + valid_token=user_api_key_dict, + llm_router=llm_router, + ) + except ProxyException as e: + _log_model_access_denial(e) + await _reject_realtime_session(websocket, user_api_key_dict, code=1008, reason=e.message[:120]) + return + await websocket.accept(subprotocol=requested_protocols[0] if requested_protocols else None) + await _route_realtime_websocket_session( + websocket, + user_api_key_dict, + route_model=route_model, + model=model, + intent=intent, + guardrails=guardrails, + ) + + +async def _reject_invalid_live_session_start( + websocket: WebSocket, + user_api_key_dict: UserAPIKeyAuth, +) -> None: + await _reject_realtime_session( + websocket, + user_api_key_dict, + code=1008, + reason="Invalid session.start frame", + error_message="First message must be a session.start JSON object with a non-empty session.model.", + error_type="invalid_request_error", + error_code="invalid_session_start", + ) + + +@app.websocket("/openai/v1/live/sessions") +@app.websocket("/v1/live/sessions") +@app.websocket("/live/sessions") +async def live_sessions_websocket_endpoint( + websocket: WebSocket, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)], +) -> None: + requested_protocols: Final = tuple( + protocol.strip() + for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",") + if protocol.strip() + ) + await websocket.accept(subprotocol=requested_protocols[0] if requested_protocols else None) + + try: + first_message: Final = await asyncio.wait_for( + websocket.receive_text(), + timeout=OPENAI_LIVE_SESSION_START_TIMEOUT_SECONDS, + ) + except asyncio.TimeoutError: + await _reject_invalid_live_session_start(websocket, user_api_key_dict) + return + except WebSocketDisconnect: + await _release_realtime_budget_reservation(user_api_key_dict) + await _release_realtime_max_parallel_slot(user_api_key_dict) + return + except RuntimeError: + await _reject_invalid_live_session_start(websocket, user_api_key_dict) + return + except BaseException: + await _release_realtime_budget_reservation(user_api_key_dict) + await _release_realtime_max_parallel_slot(user_api_key_dict) + raise + + try: + session_start: Final = TypeAdapter(dict[str, object]).validate_json( + first_message.replace("\r", "").replace("\n", "") + ) + session: Final = TypeAdapter(dict[str, object]).validate_python(session_start.get("session")) + except ValidationError: + await _reject_invalid_live_session_start(websocket, user_api_key_dict) + return + + model: Final = session.get("model") + if session_start.get("type") != "session.start" or not isinstance(model, str) or not model.strip(): + await _reject_invalid_live_session_start(websocket, user_api_key_dict) + return + + try: + await can_key_call_resolved_model( + model=model, + llm_model_list=llm_model_list, + valid_token=user_api_key_dict, + llm_router=llm_router, + ) + await can_agent_call_model( + model=model, + valid_token=user_api_key_dict, + llm_router=llm_router, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + except (ProxyException, HTTPException) as error: + denial_message: Final = error.message if isinstance(error, ProxyException) else str(error.detail) + if isinstance(error, ProxyException): + _log_model_access_denial(error) + await _reject_realtime_session( + websocket, + user_api_key_dict, + code=1008, + reason=websocket_close_reason(denial_message, fallback="Model access denied"), + error_message=denial_message, + error_type="invalid_request_error", + error_code="model_not_allowed", + ) + return + + await _route_realtime_websocket_session( + websocket, + user_api_key_dict, + route_model=model, + model=None, + intent=None, + guardrails=None, + request_path=websocket.url.path, + live_session_start=session_start, + ) + + ###################################################################### # /v1/assistant Endpoints diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 28814741852..5a6cbd06aa9 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -3,8 +3,11 @@ import asyncio import os from collections.abc import Mapping +from dataclasses import dataclass from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, cast +from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, cast + +from typing_extensions import ReadOnly, TypedDict, assert_never import litellm from litellm.constants import ( @@ -37,6 +40,7 @@ from ..llms.azure.common_utils import get_azure_ad_token from ..llms.azure.realtime.handler import AzureOpenAIRealtime, azure_realtime_protocol_for_client from ..llms.bedrock.realtime.handler import BedrockRealtime from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context +from ..llms.openai.live.handler import OpenAILiveSessions from ..llms.openai.realtime.handler import OpenAIRealtime from ..llms.vertex_ai.audio_transcription.realtime_transformation import is_vertex_speech_to_text_model from ..llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig, vertex_realtime_config @@ -51,6 +55,7 @@ if TYPE_CHECKING: azure_realtime: Final = AzureOpenAIRealtime() openai_realtime: Final = OpenAIRealtime() +openai_live_sessions: Final = OpenAILiveSessions() bedrock_realtime: Final = BedrockRealtime() xai_realtime: Final = XAIRealtime() vertex_llm_base: Final = VertexBase() @@ -345,6 +350,83 @@ async def _resolve_vertex_access_token_bounded( ) from e +class _RealtimeGuardrailRequestData(TypedDict): + litellm_metadata: ReadOnly[Mapping[str, object]] + + +def _realtime_guardrail_configured(litellm_metadata: Mapping[str, object]) -> bool: + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.types.guardrails import GuardrailEventHooks + + event_hooks: Final = ( + GuardrailEventHooks.realtime_input_transcription, + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + ) + request_data: Final[_RealtimeGuardrailRequestData] = {"litellm_metadata": litellm_metadata} + return any( + isinstance(callback, CustomGuardrail) + and any(callback.should_run_guardrail(data=request_data, event_type=hook) for hook in event_hooks) + for callback in litellm.callbacks + ) + + +@dataclass(frozen=True, slots=True) +class _LiveProviderUnsupported: + provider: str + + +@dataclass(frozen=True, slots=True) +class _LiveGuardrailsUnsupported: + pass + + +_LiveSessionRejection = _LiveProviderUnsupported | _LiveGuardrailsUnsupported + + +def _live_session_rejection( + custom_llm_provider: str, litellm_metadata: Mapping[str, object] +) -> _LiveSessionRejection | None: + if custom_llm_provider != "openai": + return _LiveProviderUnsupported(provider=custom_llm_provider) + if _realtime_guardrail_configured(litellm_metadata): + return _LiveGuardrailsUnsupported() + return None + + +def _raise_live_session_rejection(rejection: _LiveSessionRejection) -> NoReturn: + match rejection: + case _LiveProviderUnsupported(provider=provider): + raise ValueError(f"OpenAI Live sessions require the openai provider, got {provider}") + case _LiveGuardrailsUnsupported(): + raise ValueError("Guardrails are not supported on OpenAI Live sessions") + case _: + assert_never(rejection) + + +async def _alive_session( + model: str, + websocket: "WebSocket", + custom_llm_provider: str, + litellm_logging_obj: LiteLLMLogging, + session_start: Mapping[str, object], + api_base: str | None, + api_key: str | None, + litellm_metadata: Mapping[str, object], +) -> None: + rejection: Final = _live_session_rejection(custom_llm_provider, litellm_metadata) + if rejection is not None: + _raise_live_session_rejection(rejection) + await openai_live_sessions.async_live_session( + model=model, + websocket=websocket, + logging_obj=litellm_logging_obj, + session_start=session_start, + api_base=api_base or litellm.api_base or "https://api.openai.com/", + api_key=api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY"), + ) + + @wrapper_client async def _arealtime( model: str, @@ -356,6 +438,7 @@ async def _arealtime( client: object | None = None, timeout: float | None = None, query_params: RealtimeQueryParams | None = None, + live_session_start: Mapping[str, object] | None = None, **kwargs, ): """ @@ -398,6 +481,19 @@ async def _arealtime( custom_llm_provider=_custom_llm_provider, ) + if live_session_start is not None: + await _alive_session( + model=model, + websocket=websocket, + custom_llm_provider=_custom_llm_provider, + litellm_logging_obj=litellm_logging_obj, + session_start=live_session_start, + api_base=dynamic_api_base or litellm_params.api_base, + api_key=dynamic_api_key, + litellm_metadata=_build_litellm_metadata(kwargs), + ) + return + provider_config: BaseRealtimeConfig | None = None if _custom_llm_provider in LlmProviders._member_map_.values(): provider_config = ProviderConfigManager.get_provider_realtime_config( diff --git a/litellm/types/utils.py b/litellm/types/utils.py index c12a4def69a..b0c9dcf4342 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1020,6 +1020,9 @@ API_ROUTE_TO_CALL_TYPES: Final[Mapping[str, Sequence[CallTypes]]] = { "/realtime": [CallTypes.arealtime], "/v1/realtime": [CallTypes.arealtime], "/openai/v1/realtime": [CallTypes.arealtime], + "/live/sessions": (CallTypes.arealtime,), + "/v1/live/sessions": (CallTypes.arealtime,), + "/openai/v1/live/sessions": (CallTypes.arealtime,), # Provider-specific routes "/anthropic/v1/messages": [CallTypes.anthropic_messages], # Google GenAI routes diff --git a/tests/unit/llms/openai/live/__init__.py b/tests/unit/llms/openai/live/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/openai/live/test_handler.py b/tests/unit/llms/openai/live/test_handler.py new file mode 100644 index 00000000000..34e5622ca12 --- /dev/null +++ b/tests/unit/llms/openai/live/test_handler.py @@ -0,0 +1,762 @@ +import asyncio +import json +from collections.abc import Coroutine +from typing import Final, NoReturn, cast + +import pytest +from pydantic import TypeAdapter +from starlette.websockets import WebSocketDisconnect +from websockets.datastructures import Headers +from websockets.exceptions import ConnectionClosedOK, InvalidStatus +from websockets.frames import Close +from websockets.http11 import Response + +import litellm +from litellm.constants import REALTIME_SESSION_SUCCESS_LOGGED_KEY, REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES +from litellm.cost_calculator import handle_realtime_stream_cost_calculation +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.llms.openai.live import handler as live_handler +from litellm.llms.openai.live.handler import OpenAILiveSessions +from litellm.types.utils import LiteLLMRealtimeStreamLoggingObject, Usage + +_LIVE_EVENT_FRAME_ADAPTER: Final = TypeAdapter(dict[str, object]) + + +def _event_type(message: str) -> str | None: + event_type: Final = _LIVE_EVENT_FRAME_ADAPTER.validate_json(message).get("type") + return event_type if isinstance(event_type, str) else None + + +class _FakeClientWebSocket: + def __init__( + self, + messages: tuple[str, ...], + *, + disconnect_after_send_type: str | None = None, + disconnect_after_receive_type: str | None = None, + fail_on_send_type: str | None = None, + ) -> None: + self._messages = iter(messages) + self._disconnect_after_send_type = disconnect_after_send_type + self._disconnect_after_receive_type = disconnect_after_receive_type + self._fail_on_send_type = fail_on_send_type + self.sent_text: asyncio.Queue[str] = asyncio.Queue() + self.closed = asyncio.Event() + self.close_info: tuple[int, str | None] | None = None + + async def receive_text(self) -> str: + try: + message: Final = next(self._messages) + if _event_type(message) == self._disconnect_after_receive_type: + self.closed.set() + return message + except StopIteration: + await self.closed.wait() + raise WebSocketDisconnect(code=1000) + + async def send_text(self, data: str) -> None: + event_type: Final = _event_type(data) + if event_type == self._fail_on_send_type: + raise RuntimeError("client is disconnected") + await self.sent_text.put(data) + if event_type == self._disconnect_after_send_type: + self.closed.set() + + async def close(self, code: int = 1000, reason: str | None = None) -> None: + self.close_info = (code, reason) + self.closed.set() + + +class _FakeUpstream: + def __init__( + self, + client_frames: tuple[str, ...], + upstream_frames: tuple[str, ...], + *, + close_response: str | None = None, + never_respond_after_close: bool = False, + block_session_close_send: bool = False, + block_close: bool = False, + ) -> None: + self.client_frames = client_frames + self._upstream_frames = iter(upstream_frames) + self._close_response = close_response + self._close_response_sent = False + self._never_respond_after_close = never_respond_after_close + self._block_session_close_send = block_session_close_send + self._block_close = block_close + self.sent_text: asyncio.Queue[str] = asyncio.Queue() + self.client_frames_sent = asyncio.Event() + self.session_close_sent = asyncio.Event() + self.closed_event = asyncio.Event() + self.closed = False + + async def send(self, data: str) -> None: + event_type: Final = _event_type(data) + if event_type == "session.close" and self._block_session_close_send: + await asyncio.Event().wait() + await self.sent_text.put(data) + if event_type == "session.close": + self.session_close_sent.set() + if self.sent_text.qsize() == len(self.client_frames) + 1: + self.client_frames_sent.set() + + async def recv(self) -> str: + await self.client_frames_sent.wait() + try: + return next(self._upstream_frames) + except StopIteration as error: + if self._close_response is not None and not self._close_response_sent: + await self.session_close_sent.wait() + self._close_response_sent = True + return self._close_response + if self._never_respond_after_close: + await self.session_close_sent.wait() + await self.closed_event.wait() + raise ConnectionClosedOK( + Close(code=1000, reason="upstream closed"), + Close(code=1000, reason="upstream closed"), + True, + ) from error + + async def close(self) -> None: + if self._block_close: + await asyncio.Event().wait() + self.closed = True + self.closed_event.set() + + async def __aenter__(self) -> "_FakeUpstream": + return self + + async def __aexit__(self, exc_type: object, exc: object, traceback: object) -> None: + self.closed = True + self.closed_event.set() + + +class _FailingConnection: + def __init__(self, error: Exception) -> None: + self.error = error + + async def __aenter__(self) -> NoReturn: + raise self.error + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + traceback: object | None, + ) -> None: + return None + + +class _FakeLogging: + def __init__(self, pre_call_error: Exception | None = None) -> None: + self.model_call_details: dict[str, object] = {} + self.pre_call_kwargs: dict[str, object] = {} + self.dispatch_count = 0 + self.dispatched_result: object | None = None + self.prefer_async_handlers = False + self._pre_call_error = pre_call_error + + def pre_call(self, **kwargs: object) -> None: + if self._pre_call_error is not None: + raise self._pre_call_error + self.pre_call_kwargs = kwargs + + async def dispatch_success_handlers( + self, + result: object, + *, + prefer_async_handlers: bool = False, + ) -> None: + self.dispatch_count += 1 + self.dispatched_result = result + self.prefer_async_handlers = prefer_async_handlers + + +class _FakeLoggingWorker: + def __init__(self) -> None: + self.coroutine: Coroutine[object, object, None] | None = None + + def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine[object, object, None]) -> None: + self.coroutine = async_coroutine + + +class _FakeConnector: + def __init__(self, upstream: _FakeUpstream) -> None: + self.upstream = upstream + self.url: str | None = None + self.kwargs: dict[str, object] = {} + + def __call__(self, url: str, **kwargs: object) -> _FakeUpstream: + self.url = url + self.kwargs = kwargs + return self.upstream + + +@pytest.mark.parametrize( + ("api_base", "expected_url"), + [ + ("https://api.openai.com/", "wss://api.openai.com/v1/live/sessions"), + ("https://api.openai.com/v1", "wss://api.openai.com/v1/live/sessions"), + ("http://localhost:8080", "ws://localhost:8080/v1/live/sessions"), + ("https://api.openai.com/v1?api-version=1", "wss://api.openai.com/v1/live/sessions"), + ], +) +def test_construct_live_url_uses_exact_path_and_drops_query(api_base: str, expected_url: str) -> None: + assert OpenAILiveSessions._construct_live_url(api_base) == expected_url + + +@pytest.mark.asyncio +async def test_async_live_session_rewrites_model_and_relays_frames_without_mutating_input() -> None: + client_frames: Final = ('{"type":"input_audio.append","audio":"client-a"}', '{"type":"session.close"}') + upstream_frames: Final = ( + '{"type":"session.started","session":{"model":"gpt-live-1"}}', + '{"type":"response.audio.delta","delta":"server-a"}', + '{"type":"session.closed","usage":{"seconds":16}}', + ) + upstream: Final = _FakeUpstream(client_frames=client_frames, upstream_frames=upstream_frames) + connector: Final = _FakeConnector(upstream) + client: Final = _FakeClientWebSocket(messages=client_frames) + logger: Final = _FakeLogging() + worker: Final = _FakeLoggingWorker() + session_start: Final[dict[str, object]] = { + "type": "session.start", + "session": { + "model": "proxy-live-alias", + "instructions": "Be concise", + "audio": {"input": {"format": "pcm16"}}, + }, + } + original_session_start: Final = json.loads(json.dumps(session_start)) + + await OpenAILiveSessions(logging_worker=worker, websocket_connector=connector).async_live_session( + model="gpt-live-1", + websocket=client, + logging_obj=cast(Logging, logger), + session_start=session_start, + api_base="https://api.openai.com/", + api_key="live-test-key", + ) + + sent_frames: Final = tuple( + await asyncio.gather(*(upstream.sent_text.get() for _ in range(upstream.sent_text.qsize()))) + ) + forwarded_frames: Final = tuple( + await asyncio.gather(*(client.sent_text.get() for _ in range(client.sent_text.qsize()))) + ) + rewritten_start: Final = json.loads(sent_frames[0]) + assert connector.url == "wss://api.openai.com/v1/live/sessions" + assert connector.kwargs["additional_headers"] == {"Authorization": "Bearer live-test-key"} + assert connector.kwargs["max_size"] == REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES + assert "ssl" in connector.kwargs + assert rewritten_start == { + "type": "session.start", + "session": { + "model": "gpt-live-1", + "instructions": "Be concise", + "audio": {"input": {"format": "pcm16"}}, + }, + } + assert sent_frames[1:] == client_frames + assert forwarded_frames == upstream_frames + assert session_start == original_session_start + assert client.close_info == (1000, "upstream closed") + assert logger.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] is True + assert worker.coroutine is not None + await worker.coroutine + assert logger.dispatch_count == 1 + assert logger.prefer_async_handlers is True + logged_result: Final = cast(LiteLLMRealtimeStreamLoggingObject, logger.dispatched_result) + assert tuple(event["type"] for event in logged_result.results) == ( + "session.started", + "session.closed", + ) + + +@pytest.mark.asyncio +async def test_async_live_session_coalesces_cumulative_usage_updates() -> None: + client_frames: Final = ('{"type":"session.close"}',) + upstream_frames: Final = ( + '{"type":"session.started","session":{"model":"gpt-live-1"}}', + '{"type":"response.event","event":{"type":"response.completed","response":{"id":"resp-1"}}}', + '{"type":"response.event","event":{"type":"response.failed","response":{"id":"resp-2"}}}', + '{"type":"session.usage.updated","usage":{"seconds":1}}', + '{"type":"session.usage.updated","usage":{"seconds":2}}', + '{"type":"session.usage.updated","usage":{"seconds":3}}', + '{"type":"session.closed","usage":{"seconds":3}}', + ) + upstream: Final = _FakeUpstream(client_frames=client_frames, upstream_frames=upstream_frames) + connector: Final = _FakeConnector(upstream) + worker: Final = _FakeLoggingWorker() + logger: Final = _FakeLogging() + + await OpenAILiveSessions(logging_worker=worker, websocket_connector=connector).async_live_session( + model="gpt-live-1", + websocket=_FakeClientWebSocket(messages=client_frames), + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ) + + assert worker.coroutine is not None + await worker.coroutine + logged_result: Final = cast(LiteLLMRealtimeStreamLoggingObject, logger.dispatched_result) + assert tuple((event["type"], event.get("usage")) for event in logged_result.results) == ( + ("session.started", None), + ("response.event", None), + ("response.event", None), + ("session.usage.updated", {"seconds": 3}), + ("session.closed", {"seconds": 3}), + ) + assert tuple(logged_result.results[1:3]) == ( + { + "type": "response.event", + "event": { + "type": "response.completed", + "response": {"id": "resp-1", "model": None, "usage": None}, + }, + }, + { + "type": "response.event", + "event": { + "type": "response.failed", + "response": {"id": "resp-2", "model": None, "usage": None}, + }, + }, + ) + + +@pytest.mark.asyncio +async def test_async_live_session_requires_api_key() -> None: + with pytest.raises(ValueError, match="api_key"): + await OpenAILiveSessions().async_live_session( + model="gpt-live-1", + websocket=_FakeClientWebSocket(messages=()), + logging_obj=cast(Logging, _FakeLogging()), + session_start={"type": "session.start", "session": {"model": "proxy-live-alias"}}, + api_base="https://api.openai.com/", + api_key=None, + ) + + +@pytest.mark.asyncio +async def test_async_live_session_does_not_log_success_on_handshake_refusal() -> None: + response: Final = Response(401, "Unauthorized", Headers()) + invalid_status: Final = InvalidStatus(response) + + def refuse_connection(url: str, **kwargs: object) -> _FailingConnection: + return _FailingConnection(invalid_status) + + client: Final = _FakeClientWebSocket(messages=()) + logger: Final = _FakeLogging() + worker: Final = _FakeLoggingWorker() + + await OpenAILiveSessions(logging_worker=worker, websocket_connector=refuse_connection).async_live_session( + model="gpt-live-1", + websocket=client, + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ) + + try: + assert worker.coroutine is None + assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in logger.model_call_details + assert logger.dispatch_count == 0 + finally: + if worker.coroutine is not None: + await worker.coroutine + + +@pytest.mark.asyncio +async def test_async_live_session_does_not_log_success_when_connection_fails_before_connect() -> None: + def fail_connection(url: str, **kwargs: object) -> _FailingConnection: + return _FailingConnection(RuntimeError("connection failed")) + + logger: Final = _FakeLogging() + worker: Final = _FakeLoggingWorker() + + await OpenAILiveSessions(logging_worker=worker, websocket_connector=fail_connection).async_live_session( + model="gpt-live-1", + websocket=_FakeClientWebSocket(messages=()), + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ) + + assert worker.coroutine is None + assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in logger.model_call_details + assert logger.dispatch_count == 0 + + +@pytest.mark.asyncio +async def test_async_live_session_retains_final_usage_after_client_disconnect_and_bills_voice() -> None: + session_closed: Final = json.dumps({"type": "session.closed", "usage": {"seconds": 7}}) + upstream: Final = _FakeUpstream( + client_frames=(), + upstream_frames=('{"type":"session.started","session":{"model":"gpt-live-1"}}',), + close_response=session_closed, + ) + connector: Final = _FakeConnector(upstream) + client: Final = _FakeClientWebSocket(messages=(), disconnect_after_send_type="session.started") + logger: Final = _FakeLogging() + worker: Final = _FakeLoggingWorker() + + await OpenAILiveSessions(logging_worker=worker, websocket_connector=connector).async_live_session( + model="gpt-live-1", + websocket=client, + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ) + + assert worker.coroutine is not None + await worker.coroutine + logged_result: Final = cast(LiteLLMRealtimeStreamLoggingObject, logger.dispatched_result) + assert logged_result.results == [ + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.closed", "usage": {"seconds": 7}}, + ] + sent_frames: Final = tuple( + await asyncio.gather(*(upstream.sent_text.get() for _ in range(upstream.sent_text.qsize()))) + ) + session_close: Final = json.dumps({"type": "session.close"}) + assert tuple(frame for frame in sent_frames if _event_type(frame) == "session.close") == (session_close,) + + rate: Final = cast(float, litellm.model_cost["gpt-live-1"]["input_cost_per_second"]) + cost: Final = handle_realtime_stream_cost_calculation( + results=logged_result.results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + ) + assert cost == pytest.approx(7 * rate) + + +@pytest.mark.asyncio +async def test_async_live_session_keeps_finite_usage_after_invalid_snapshot() -> None: + upstream: Final = _FakeUpstream( + client_frames=(), + upstream_frames=( + '{"type":"session.started","session":{"model":"gpt-live-1"}}', + '{"type":"session.usage.updated","usage":{"seconds":7}}', + '{"type":"session.usage.updated","usage":{"seconds":"nan"}}', + ), + never_respond_after_close=True, + ) + connector: Final = _FakeConnector(upstream) + client: Final = _FakeClientWebSocket(messages=(), disconnect_after_send_type="session.started") + logger: Final = _FakeLogging() + worker: Final = _FakeLoggingWorker() + + await OpenAILiveSessions(logging_worker=worker, websocket_connector=connector).async_live_session( + model="gpt-live-1", + websocket=client, + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ) + + assert worker.coroutine is not None + await worker.coroutine + logged_result: Final = cast(LiteLLMRealtimeStreamLoggingObject, logger.dispatched_result) + assert logged_result.results == [ + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.usage.updated", "usage": {"seconds": 7}}, + ] + rate: Final = cast(float, litellm.model_cost["gpt-live-1"]["input_cost_per_second"]) + cost: Final = handle_realtime_stream_cost_calculation( + results=logged_result.results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + ) + assert cost == pytest.approx(7 * rate) + + +@pytest.mark.asyncio +async def test_async_live_session_does_not_send_a_second_close_after_client_close( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(live_handler, "OPENAI_LIVE_SESSION_CLOSE_TIMEOUT_SECONDS", 0.01) + client_frames: Final = (json.dumps({"type": "session.close"}),) + upstream: Final = _FakeUpstream( + client_frames=client_frames, + upstream_frames=(), + never_respond_after_close=True, + ) + connector: Final = _FakeConnector(upstream) + worker: Final = _FakeLoggingWorker() + logger: Final = _FakeLogging() + + await OpenAILiveSessions(logging_worker=worker, websocket_connector=connector).async_live_session( + model="gpt-live-1", + websocket=_FakeClientWebSocket( + messages=client_frames, + disconnect_after_receive_type="session.close", + ), + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ) + + sent_frames: Final = tuple( + await asyncio.gather(*(upstream.sent_text.get() for _ in range(upstream.sent_text.qsize()))) + ) + assert tuple(frame for frame in sent_frames if _event_type(frame) == "session.close") == client_frames + assert worker.coroutine is not None + await worker.coroutine + + +@pytest.mark.asyncio +async def test_async_live_session_bounds_a_blocking_upstream_close_send( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(live_handler, "OPENAI_LIVE_SESSION_CLOSE_TIMEOUT_SECONDS", 0.05) + upstream: Final = _FakeUpstream( + client_frames=(), + upstream_frames=('{"type":"session.started","session":{"model":"gpt-live-1"}}',), + never_respond_after_close=True, + block_session_close_send=True, + ) + connector: Final = _FakeConnector(upstream) + client: Final = _FakeClientWebSocket(messages=(), disconnect_after_send_type="session.started") + logger: Final = _FakeLogging() + worker: Final = _FakeLoggingWorker() + + await asyncio.wait_for( + OpenAILiveSessions(logging_worker=worker, websocket_connector=connector).async_live_session( + model="gpt-live-1", + websocket=client, + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ), + timeout=1, + ) + + assert worker.coroutine is not None + await worker.coroutine + assert logger.dispatch_count == 1 + + +@pytest.mark.asyncio +async def test_async_live_session_bounds_a_blocking_upstream_close( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(live_handler, "OPENAI_LIVE_SESSION_CLOSE_TIMEOUT_SECONDS", 0.05) + upstream: Final = _FakeUpstream( + client_frames=(), + upstream_frames=('{"type":"session.started","session":{"model":"gpt-live-1"}}',), + close_response='{"type":"session.closed","usage":{"seconds":7}}', + block_close=True, + ) + connector: Final = _FakeConnector(upstream) + client: Final = _FakeClientWebSocket(messages=(), disconnect_after_send_type="session.started") + logger: Final = _FakeLogging() + worker: Final = _FakeLoggingWorker() + + await asyncio.wait_for( + OpenAILiveSessions(logging_worker=worker, websocket_connector=connector).async_live_session( + model="gpt-live-1", + websocket=client, + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ), + timeout=1, + ) + + assert worker.coroutine is not None + await worker.coroutine + assert logger.dispatch_count == 1 + + +@pytest.mark.asyncio +async def test_async_live_session_retains_closed_event_when_client_send_fails() -> None: + upstream: Final = _FakeUpstream( + client_frames=(), + upstream_frames=( + '{"type":"session.started","session":{"model":"gpt-live-1"}}', + '{"type":"session.closed","usage":{"seconds":7}}', + ), + ) + connector: Final = _FakeConnector(upstream) + client: Final = _FakeClientWebSocket(messages=(), fail_on_send_type="session.closed") + logger: Final = _FakeLogging() + worker: Final = _FakeLoggingWorker() + + await OpenAILiveSessions(logging_worker=worker, websocket_connector=connector).async_live_session( + model="gpt-live-1", + websocket=client, + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ) + + assert worker.coroutine is not None + await worker.coroutine + logged_result: Final = cast(LiteLLMRealtimeStreamLoggingObject, logger.dispatched_result) + assert {"type": "session.closed", "usage": {"seconds": 7}} in logged_result.results + + +@pytest.mark.asyncio +async def test_async_live_session_bounds_wait_for_usage_after_client_disconnect( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(live_handler, "OPENAI_LIVE_SESSION_CLOSE_TIMEOUT_SECONDS", 0.01) + upstream: Final = _FakeUpstream( + client_frames=(), + upstream_frames=('{"type":"session.started","session":{"model":"gpt-live-1"}}',), + never_respond_after_close=True, + ) + connector: Final = _FakeConnector(upstream) + client: Final = _FakeClientWebSocket(messages=(), disconnect_after_send_type="session.started") + logger: Final = _FakeLogging() + worker: Final = _FakeLoggingWorker() + + await asyncio.wait_for( + OpenAILiveSessions(logging_worker=worker, websocket_connector=connector).async_live_session( + model="gpt-live-1", + websocket=client, + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ), + timeout=1, + ) + + assert upstream.closed + assert worker.coroutine is not None + await worker.coroutine + logged_result: Final = cast(LiteLLMRealtimeStreamLoggingObject, logger.dispatched_result) + assert all(event["type"] != "session.closed" for event in logged_result.results) + + +@pytest.mark.asyncio +async def test_async_live_session_compacts_and_deduplicates_terminal_response_events() -> None: + large_output: Final = "x" * 4096 + large_instructions: Final = "y" * 4096 + response_usage: Final[dict[str, object]] = {"input_tokens": 10, "output_tokens": 5} + duplicate_response_usage: Final[dict[str, object]] = {"input_tokens": 100, "output_tokens": 50} + first_response: Final = json.dumps( + { + "type": "response.event", + "event": { + "type": "response.completed", + "response": { + "id": "resp-1", + "model": "gpt-live-1", + "usage": response_usage, + "output": large_output, + "instructions": large_instructions, + }, + }, + } + ) + duplicate_response: Final = json.dumps( + { + "type": "response.event", + "event": { + "type": "response.completed", + "response": { + "id": "resp-1", + "model": "duplicate-model", + "usage": duplicate_response_usage, + "output": large_output, + "instructions": large_instructions, + }, + }, + } + ) + response_without_id: Final = json.dumps( + { + "type": "response.event", + "event": { + "type": "response.completed", + "response": {"model": "gpt-live-1", "usage": response_usage}, + }, + } + ) + upstream: Final = _FakeUpstream( + client_frames=(), + upstream_frames=( + '{"type":"session.started","session":{"model":"gpt-live-1"}}', + first_response, + duplicate_response, + response_without_id, + ), + ) + connector: Final = _FakeConnector(upstream) + logger: Final = _FakeLogging() + worker: Final = _FakeLoggingWorker() + + await OpenAILiveSessions(logging_worker=worker, websocket_connector=connector).async_live_session( + model="gpt-live-1", + websocket=_FakeClientWebSocket(messages=()), + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "gpt-live-1"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ) + + assert worker.coroutine is not None + await worker.coroutine + logged_result: Final = cast(LiteLLMRealtimeStreamLoggingObject, logger.dispatched_result) + response_events: Final = tuple(event for event in logged_result.results if event["type"] == "response.event") + assert response_events == ( + { + "type": "response.event", + "event": { + "type": "response.completed", + "response": { + "id": "resp-1", + "model": "gpt-live-1", + "usage": response_usage, + }, + }, + }, + ) + + +@pytest.mark.asyncio +async def test_async_live_session_reports_redacted_server_error_and_bounds_close_reason() -> None: + error_message: Final = "x" * 300 + client: Final = _FakeClientWebSocket(messages=()) + logger: Final = _FakeLogging(pre_call_error=RuntimeError(error_message)) + worker: Final = _FakeLoggingWorker() + + await OpenAILiveSessions(logging_worker=worker).async_live_session( + model="gpt-live-1", + websocket=client, + logging_obj=cast(Logging, logger), + session_start={"type": "session.start", "session": {"model": "proxy-live-alias"}}, + api_base="https://api.openai.com/", + api_key="live-test-key", + ) + + error_event: Final = json.loads(client.sent_text.get_nowait()) + assert error_event == {"type": "error", "error": {"type": "server_error", "message": error_message}} + assert client.close_info is not None + assert client.close_info[0] == 1011 + assert client.close_info[1] is not None + assert len(client.close_info[1].encode("utf-8")) <= 123 + assert client.close_info[1] == "x" * 123 + try: + assert worker.coroutine is None + assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in logger.model_call_details + assert logger.dispatch_count == 0 + finally: + if worker.coroutine is not None: + await worker.coroutine diff --git a/tests/unit/proxy/test_proxy_server.py b/tests/unit/proxy/test_proxy_server.py index 8947da4d9fc..16147daa5f4 100644 --- a/tests/unit/proxy/test_proxy_server.py +++ b/tests/unit/proxy/test_proxy_server.py @@ -1,6 +1,7 @@ import os import traceback -from typing import Final +from collections.abc import Coroutine +from typing import Final, cast from unittest import mock from dotenv import load_dotenv @@ -3108,3 +3109,249 @@ def test_get_litellm_model_info(data): ): get_litellm_model_info(model=model) get_info_mock.assert_called_once_with(data["expected"]) + + +@pytest.mark.parametrize("path", ["/openai/v1/live/sessions", "/v1/live/sessions", "/live/sessions"]) +def test_live_session_invalid_start_is_rejected_without_routing( + path: str, + client_no_auth: TestClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy import proxy_server + from starlette.websockets import WebSocketDisconnect + + async def no_op(user_api_key_dict: object) -> None: + return None + + route_call: Final = AsyncMock() + monkeypatch.setattr(proxy_server, "route_request", route_call) + monkeypatch.setattr(proxy_server, "_release_realtime_budget_reservation", no_op) + monkeypatch.setattr(proxy_server, "_release_realtime_max_parallel_slot", no_op) + + with client_no_auth.websocket_connect(path, headers={"Authorization": "Bearer test-api-key"}) as websocket: + websocket.send_text('{"type":"session.start","session":{"model":""}}') + assert websocket.receive_json() == { + "type": "error", + "error": { + "type": "invalid_request_error", + "code": "invalid_session_start", + "message": "First message must be a session.start JSON object with a non-empty session.model.", + }, + } + with pytest.raises(WebSocketDisconnect) as error: + websocket.receive_json() + assert error.value.code == 1008 + + route_call.assert_not_awaited() + + +@pytest.mark.parametrize(("frame_model", "granted"), (("gpt-3.5-turbo", True), ("gpt-4o", False))) +def test_live_session_frame_model_is_checked_against_managed_agent_grants( + client_no_auth: TestClient, + monkeypatch: pytest.MonkeyPatch, + frame_model: str, + granted: bool, +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_websocket + from litellm.types.agents import AgentResponse + from starlette.websockets import WebSocket, WebSocketDisconnect + + async def managed_agent_auth() -> UserAPIKeyAuth: + auth: Final = UserAPIKeyAuth(token="managed-token", agent_id="managed") + auth.managed_agent_policy = AgentResponse( + agent_id="managed", + agent_name="Managed", + agent_card_params={}, + object_permission={"models": ["gpt-3.5-turbo"]}, + ) + return auth + + async def no_op(user_api_key_dict: object) -> None: + return None + + async def complete_call() -> None: + return None + + async def route( + *, + data: dict[str, object], + route_type: str, + llm_router: object, + user_model: str | None, + ) -> Coroutine[object, object, None]: + await cast(WebSocket, data["websocket"]).send_text("routed") + return complete_call() + + route_call: Final = AsyncMock(side_effect=route) + monkeypatch.setitem(proxy_server.app.dependency_overrides, user_api_key_auth_websocket, managed_agent_auth) + monkeypatch.setattr(proxy_server, "route_request", route_call) + monkeypatch.setattr(proxy_server, "_release_realtime_budget_reservation", no_op) + monkeypatch.setattr(proxy_server, "_release_realtime_max_parallel_slot", no_op) + + with client_no_auth.websocket_connect( + "/v1/live/sessions", + headers={"Authorization": "Bearer managed-token"}, + ) as websocket: + websocket.send_json({"type": "session.start", "session": {"model": frame_model}}) + if granted: + assert websocket.receive_text() == "routed" + else: + assert websocket.receive_json()["type"] == "error" + with pytest.raises(WebSocketDisconnect) as error: + websocket.receive_json() + assert error.value.code == 1008 + + assert route_call.await_count == (1 if granted else 0) + + +def test_live_session_valid_start_routes_with_model_and_frame( + client_no_auth: TestClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + from types import SimpleNamespace + + from litellm.proxy import proxy_server + from starlette.websockets import WebSocket + + async def no_op(user_api_key_dict: object) -> None: + return None + + async def complete_call() -> None: + return None + + captured: Final = SimpleNamespace(data=None, route_type=None) + + async def route( + *, + data: dict[str, object], + route_type: str, + llm_router: object, + user_model: str | None, + ) -> Coroutine[object, object, None]: + captured.data = data + captured.route_type = route_type + await cast(WebSocket, data["websocket"]).send_text("routed") + return complete_call() + + async def allow_model(**kwargs: object) -> None: + return None + + monkeypatch.setattr(proxy_server, "route_request", route) + monkeypatch.setattr(proxy_server, "can_key_call_resolved_model", allow_model) + monkeypatch.setattr(proxy_server, "_release_realtime_budget_reservation", no_op) + monkeypatch.setattr(proxy_server, "_release_realtime_max_parallel_slot", no_op) + start_frame: Final[dict[str, object]] = { + "type": "session.start", + "session": {"model": "gpt-3.5-turbo", "instructions": "Stay brief"}, + } + + with client_no_auth.websocket_connect( + "/v1/live/sessions", + headers={"Authorization": "Bearer test-api-key"}, + ) as websocket: + websocket.send_json(start_frame) + assert websocket.receive_text() == "routed" + + routed_data: Final = cast(dict[str, object], captured.data) + assert captured.route_type == "_arealtime" + assert routed_data["model"] == "gpt-3.5-turbo" + assert routed_data["live_session_start"] == start_frame + + +def test_live_session_preserves_pretty_printed_start_frame( + client_no_auth: TestClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + from types import SimpleNamespace + + from starlette.websockets import WebSocket + + from litellm.proxy import proxy_server + + async def no_op(user_api_key_dict: object) -> None: + return None + + async def complete_call() -> None: + return None + + captured: Final = SimpleNamespace(data=None, route_type=None) + + async def route( + *, + data: dict[str, object], + route_type: str, + llm_router: object, + user_model: str | None, + ) -> Coroutine[object, object, None]: + captured.data = data + captured.route_type = route_type + await cast(WebSocket, data["websocket"]).send_text("routed") + return complete_call() + + async def allow_model(**kwargs: object) -> None: + return None + + monkeypatch.setattr(proxy_server, "route_request", route) + monkeypatch.setattr(proxy_server, "can_key_call_resolved_model", allow_model) + monkeypatch.setattr(proxy_server, "_release_realtime_budget_reservation", no_op) + monkeypatch.setattr(proxy_server, "_release_realtime_max_parallel_slot", no_op) + + start_frame: Final = "\r\n".join( + [ + "{", + ' "type": "session.start",', + ' "session": {"model": "gpt-live-1", "instructions": "a\\nb"}', + "}", + ] + ) + expected_start: Final[dict[str, object]] = { + "type": "session.start", + "session": {"model": "gpt-live-1", "instructions": "a\nb"}, + } + + with client_no_auth.websocket_connect( + "/v1/live/sessions", + headers={"Authorization": "Bearer test-api-key"}, + ) as websocket: + websocket.send_text(start_frame) + assert websocket.receive_text() == "routed" + + routed_data: Final = cast(dict[str, object], captured.data) + assert captured.route_type == "_arealtime" + assert routed_data["model"] == "gpt-live-1" + assert routed_data["live_session_start"] == expected_start + + +def test_live_session_denied_model_closes_with_policy_code( + client_no_auth: TestClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy._types import ProxyException + from starlette.websockets import WebSocketDisconnect + + async def no_op(user_api_key_dict: object) -> None: + return None + + async def deny_model(**kwargs: object) -> None: + raise ProxyException("model access denied", "invalid_request_error", None, 403) + + route_call: Final = AsyncMock() + monkeypatch.setattr(proxy_server, "can_key_call_resolved_model", deny_model) + monkeypatch.setattr(proxy_server, "route_request", route_call) + monkeypatch.setattr(proxy_server, "_release_realtime_budget_reservation", no_op) + monkeypatch.setattr(proxy_server, "_release_realtime_max_parallel_slot", no_op) + + with client_no_auth.websocket_connect( + "/v1/live/sessions", + headers={"Authorization": "Bearer test-api-key"}, + ) as websocket: + websocket.send_json({"type": "session.start", "session": {"model": "gpt-3.5-turbo"}}) + assert websocket.receive_json()["type"] == "error" + with pytest.raises(WebSocketDisconnect) as error: + websocket.receive_json() + assert error.value.code == 1008 + + route_call.assert_not_awaited() diff --git a/tests/unit/realtime_api/test_main.py b/tests/unit/realtime_api/test_main.py index 5d3276dfae1..614acd01a20 100644 --- a/tests/unit/realtime_api/test_main.py +++ b/tests/unit/realtime_api/test_main.py @@ -1,5 +1,6 @@ import asyncio import time +from collections.abc import Mapping from types import TracebackType from typing import Final from unittest.mock import MagicMock, patch @@ -7,9 +8,11 @@ from unittest.mock import MagicMock, patch import pytest import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.models.credentials import CredentialItem from litellm.realtime_api import main as realtime_main from litellm.realtime_api.main import _with_resolved_session_model +from litellm.types.guardrails import GuardrailEventHooks @pytest.fixture @@ -26,6 +29,11 @@ class FakeLogging: pass +class _CallCapture: + def __init__(self) -> None: + self.values: Mapping[str, object] | None = None + + def test_resolves_top_level_session_model(): resolved = _with_resolved_session_model({"model": "alias/gpt-realtime"}, "gpt-realtime") assert resolved == {"model": "gpt-realtime"} @@ -231,6 +239,149 @@ def test_client_secret_forwards_nested_transcription_model_untouched(monkeypatch assert session["input_audio_transcription"]["model"] == "whisper-1" +@pytest.mark.asyncio +async def test_arealtime_live_session_dispatches_to_openai_live_handler(monkeypatch: pytest.MonkeyPatch) -> None: + captured: Final = _CallCapture() + + def mock_get_llm_provider( + model: str, + api_base: str | None, + api_key: str | None, + ) -> tuple[str, str, str | None, str | None]: + return "gpt-live-1", "openai", "resolved-key", "https://live-gateway.example/v1" + + class FakeLiveHandler: + async def async_live_session(self, **kwargs: object) -> None: + captured.values = kwargs + + monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) + monkeypatch.setattr(realtime_main, "openai_live_sessions", FakeLiveHandler()) + start_frame: Final[dict[str, object]] = { + "type": "session.start", + "session": {"model": "live-alias", "instructions": "Stay brief"}, + } + websocket: Final = MagicMock() + logging_obj: Final = FakeLogging() + + await realtime_main._arealtime.__wrapped__( + model="live-alias", + websocket=websocket, + litellm_logging_obj=logging_obj, + live_session_start=start_frame, + ) + + assert captured.values is not None + assert captured.values["model"] == "gpt-live-1" + assert captured.values["websocket"] is websocket + assert captured.values["logging_obj"] is logging_obj + assert captured.values["session_start"] == start_frame + assert captured.values["api_base"] == "https://live-gateway.example/v1" + assert captured.values["api_key"] == "resolved-key" + + +@pytest.mark.asyncio +async def test_arealtime_live_session_rejects_matching_guardrail_before_dispatch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class RecordingGuardrail(CustomGuardrail): + def __init__(self) -> None: + super().__init__( + guardrail_name="live-blocker", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + ) + self.request_data: dict[str, object] | None = None + + def should_run_guardrail(self, data: dict[str, object], event_type: GuardrailEventHooks) -> bool: + self.request_data = data + return super().should_run_guardrail(data, event_type) + + class FakeLiveHandler: + def __init__(self) -> None: + self.called = False + + async def async_live_session(self, **kwargs: object) -> None: + self.called = True + + guardrail: Final = RecordingGuardrail() + live_handler: Final = FakeLiveHandler() + + def mock_get_llm_provider( + model: str, + api_base: str | None, + api_key: str | None, + ) -> tuple[str, str, str | None, str | None]: + return "gpt-live-1", "openai", "resolved-key", "https://live-gateway.example/v1" + + monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) + monkeypatch.setattr(realtime_main, "openai_live_sessions", live_handler) + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + + with pytest.raises(ValueError, match="Guardrails are not supported on OpenAI Live sessions"): + await realtime_main._arealtime.__wrapped__( + model="live-alias", + websocket=MagicMock(), + litellm_logging_obj=FakeLogging(), + litellm_metadata={"session_id": "live-session"}, + live_session_start={"type": "session.start", "session": {"model": "live-alias"}}, + ) + + assert guardrail.request_data == {"litellm_metadata": {"session_id": "live-session"}} + assert live_handler.called is False + + +@pytest.mark.asyncio +async def test_arealtime_live_session_rejects_non_openai_provider(monkeypatch: pytest.MonkeyPatch) -> None: + def mock_get_llm_provider( + model: str, + api_base: str | None, + api_key: str | None, + ) -> tuple[str, str, str | None, str | None]: + return "claude-live", "anthropic", None, None + + monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) + + with pytest.raises(ValueError, match="anthropic"): + await realtime_main._arealtime.__wrapped__( + model="anthropic/claude-live", + websocket=MagicMock(), + litellm_logging_obj=FakeLogging(), + live_session_start={"type": "session.start", "session": {"model": "claude-live"}}, + ) + + +@pytest.mark.asyncio +async def test_arealtime_without_live_start_preserves_openai_realtime_dispatch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + captured: Final = _CallCapture() + + def mock_get_llm_provider( + model: str, + api_base: str | None, + api_key: str | None, + ) -> tuple[str, str, str | None, str | None]: + return "gpt-realtime", "openai", "resolved-key", "https://realtime-gateway.example/v1" + + class FakeRealtimeHandler: + async def async_realtime(self, **kwargs: object) -> None: + captured.values = kwargs + + monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) + monkeypatch.setattr(realtime_main, "openai_realtime", FakeRealtimeHandler()) + + await realtime_main._arealtime.__wrapped__( + model="gpt-realtime", + websocket=MagicMock(), + litellm_logging_obj=FakeLogging(), + ) + + assert captured.values is not None + assert captured.values["model"] == "gpt-realtime" + assert captured.values["api_key"] == "resolved-key" + assert captured.values["api_base"] == "https://realtime-gateway.example/v1" + + class _CapturingConnect: def __init__(self) -> None: self.url: str | None = None diff --git a/tests/unit/test_anthropic_beta_headers_filtering.py b/tests/unit/test_anthropic_beta_headers_filtering.py index 1a6899f16ba..8656a7564d2 100644 --- a/tests/unit/test_anthropic_beta_headers_filtering.py +++ b/tests/unit/test_anthropic_beta_headers_filtering.py @@ -444,7 +444,7 @@ class TestAnthropicBetaHeadersFiltering: assert filtered == ["thinking-binding-controls-2026-08-01"] - @pytest.mark.parametrize("provider", ["anthropic", "azure_ai", "bedrock", "bedrock_mantle", "vertex_ai"]) + @pytest.mark.parametrize("provider", ["anthropic", "bedrock", "bedrock_mantle", "vertex_ai"]) def test_dangerous_tool_use_forwarded(self, provider): """Claude Code's server-side auto-mode classifier sends `safeguards` together with dangerous-tool-use-2026-09-03. Bedrock Invoke, Bedrock Mantle, and Vertex rawPredict diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index 36e188e82d6..5c8a70cf8d9 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -5639,3 +5639,264 @@ def test_completion_cost_bills_base_when_gemini_serves_on_demand( ) assert cost == pytest.approx(100 * 0.001 + 50 * 0.002) + + +def _live_session_results(*events: dict[str, object]) -> OpenAIRealtimeStreamList: + return cast(OpenAIRealtimeStreamList, list(events)) + + +def _live_logging_object( + start_time: datetime.datetime, + end_time: datetime.datetime, +) -> Logging: + return cast( + Logging, + SimpleNamespace(model_call_details={"start_time": start_time, "end_time": end_time}), + ) + + +def test_realtime_keeps_first_reported_session_model_for_pricing(_local_model_cost_map: None) -> None: + first_model: Final = "gpt-realtime-2" + later_model: Final = "gpt-realtime-mini" + combined_usage: Final = Usage(prompt_tokens=120, completion_tokens=60, total_tokens=180) + first_model_cost: Final = handle_realtime_stream_cost_calculation( + results=_live_session_results({"type": "session.created", "session": {"model": first_model}}), + combined_usage_object=combined_usage, + custom_llm_provider="openai", + litellm_model_name="unmapped-realtime-alias", + ) + first_model_with_later_update_cost: Final = handle_realtime_stream_cost_calculation( + results=_live_session_results( + {"type": "session.created", "session": {"model": first_model}}, + {"type": "session.started", "session": {"model": later_model}}, + ), + combined_usage_object=combined_usage, + custom_llm_provider="openai", + litellm_model_name="unmapped-realtime-alias", + ) + later_model_cost: Final = handle_realtime_stream_cost_calculation( + results=_live_session_results({"type": "session.started", "session": {"model": later_model}}), + combined_usage_object=combined_usage, + custom_llm_provider="openai", + litellm_model_name="unmapped-realtime-alias", + ) + + assert first_model_with_later_update_cost == pytest.approx(first_model_cost) + assert first_model_cost != pytest.approx(later_model_cost) + + +def test_live_session_charges_latest_closed_voice_usage(_local_model_cost_map: None) -> None: + results: Final = _live_session_results( + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.usage.updated", "usage": {"seconds": 12}}, + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + {"type": "session.closed", "usage": {"seconds": 16}}, + ) + rate: Final = cast(float, litellm.model_cost["gpt-live-1"]["input_cost_per_second"]) + + cost: Final = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + ) + + assert rate > 0 + assert cost == pytest.approx(16 * rate) + + +def test_live_session_uses_last_reported_duration_when_wall_clock_is_short(_local_model_cost_map: None) -> None: + start_time: Final = datetime.datetime(2025, 1, 1) + results: Final = _live_session_results( + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + ) + rate: Final = cast(float, litellm.model_cost["gpt-live-1"]["input_cost_per_second"]) + + cost: Final = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + litellm_logging_obj=_live_logging_object(start_time, start_time + datetime.timedelta(seconds=5)), + ) + + assert cost == pytest.approx(15 * rate) + + +def test_live_session_uses_latest_cumulative_usage_snapshot_not_sum(_local_model_cost_map: None) -> None: + results: Final = _live_session_results( + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.usage.updated", "usage": {"seconds": 12}}, + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + ) + rate: Final = cast(float, litellm.model_cost["gpt-live-1"]["input_cost_per_second"]) + + cost: Final = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + ) + + assert cost == pytest.approx(15 * rate) + + +def test_live_session_ignores_non_finite_usage_snapshot_after_finite_one(_local_model_cost_map: None) -> None: + results: Final = _live_session_results( + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + {"type": "session.usage.updated", "usage": {"seconds": float("nan")}}, + ) + rate: Final = cast(float, litellm.model_cost["gpt-live-1"]["input_cost_per_second"]) + + cost: Final = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + ) + + assert cost == pytest.approx(15 * rate) + + +def test_live_session_ignores_non_finite_closed_usage_snapshot(_local_model_cost_map: None) -> None: + results: Final = _live_session_results( + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + {"type": "session.closed", "usage": {"seconds": float("inf")}}, + ) + rate: Final = cast(float, litellm.model_cost["gpt-live-1"]["input_cost_per_second"]) + + cost: Final = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + ) + + assert cost == pytest.approx(15 * rate) + + +def test_live_session_skips_non_finite_rate_for_next_candidate( + _local_model_cost_map: None, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setitem( + litellm.model_cost, + "live-non-finite-rate", + {"input_cost_per_second": float("nan"), "litellm_provider": "openai", "mode": "realtime"}, + ) + monkeypatch.setitem( + litellm.model_cost, + "live-fallback-rate", + {"input_cost_per_second": 0.25, "litellm_provider": "openai", "mode": "realtime"}, + ) + litellm.get_model_info.cache_clear() + results: Final = _live_session_results( + {"type": "session.started", "session": {"model": "live-non-finite-rate"}}, + {"type": "session.usage.updated", "usage": {"seconds": 4}}, + ) + + cost: Final = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="live-fallback-rate", + ) + + assert cost == pytest.approx(4 * 0.25) + + +def test_live_session_uses_wall_clock_duration_when_greater_than_reported(_local_model_cost_map: None) -> None: + start_time: Final = datetime.datetime(2025, 1, 1) + results: Final = _live_session_results( + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + ) + rate: Final = cast(float, litellm.model_cost["gpt-live-1"]["input_cost_per_second"]) + + cost: Final = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + litellm_logging_obj=_live_logging_object(start_time, start_time + datetime.timedelta(seconds=20)), + ) + + assert cost == pytest.approx(20 * rate) + + +def test_live_session_deduplicates_delegated_responses_and_charges_voice(_local_model_cost_map: None) -> None: + delegated_response: Final[dict[str, object]] = { + "type": "response.event", + "delegation_id": "delegation-1", + "event": { + "type": "response.completed", + "response": { + "id": "resp-1", + "model": "gpt-4o-mini", + "usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, + }, + }, + } + results: Final = _live_session_results( + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.closed", "usage": {"seconds": 12}}, + delegated_response, + delegated_response, + ) + live_rate: Final = cast(float, litellm.model_cost["gpt-live-1"]["input_cost_per_second"]) + backend_model: Final = litellm.model_cost["gpt-4o-mini"] + expected_backend_cost: Final = ( + 10 * cast(float, backend_model["input_cost_per_token"]) + + 5 * cast(float, backend_model["output_cost_per_token"]) + ) + + cost: Final = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + ) + + assert cost == pytest.approx(12 * live_rate + expected_backend_cost) + + +def test_live_session_zero_closed_seconds_do_not_use_usage_fallback(_local_model_cost_map: None) -> None: + start_time: Final = datetime.datetime(2025, 1, 1) + results: Final = _live_session_results( + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + {"type": "session.closed", "usage": {"seconds": 0}}, + ) + + cost: Final = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + litellm_logging_obj=_live_logging_object(start_time, start_time + datetime.timedelta(seconds=20)), + ) + + assert cost == 0.0 + + +def test_live_session_uses_latest_snapshot_when_closed_usage_is_missing(_local_model_cost_map: None) -> None: + start_time: Final = datetime.datetime(2025, 1, 1) + results: Final = _live_session_results( + {"type": "session.started", "session": {"model": "gpt-live-1"}}, + {"type": "session.usage.updated", "usage": {"seconds": 15}}, + {"type": "session.closed", "reason": "server_end"}, + ) + rate: Final = cast(float, litellm.model_cost["gpt-live-1"]["input_cost_per_second"]) + + cost: Final = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-live-1", + litellm_logging_obj=_live_logging_object(start_time, start_time + datetime.timedelta(seconds=20)), + ) + + assert cost == pytest.approx(15 * rate) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 8ffd5c96ab0..7fdc3cf73fa 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -8882,6 +8882,26 @@ export interface paths { patch?: never; trace?: never; }; + "/live/sessions": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * WebSocket: live_sessions_websocket_endpoint + * @description WebSocket connection endpoint + */ + get: operations["websocket_live_sessions_websocket_endpoint_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/login": { parameters: { query?: never; @@ -10523,6 +10543,26 @@ export interface paths { patch?: never; trace?: never; }; + "/openai/v1/live/sessions": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * WebSocket: live_sessions_websocket_endpoint + * @description WebSocket connection endpoint + */ + get: operations["websocket_live_sessions_websocket_endpoint_get_3"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/openai/v1/realtime": { parameters: { query?: never; @@ -19684,6 +19724,26 @@ export interface paths { patch?: never; trace?: never; }; + "/v1/live/sessions": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * WebSocket: live_sessions_websocket_endpoint + * @description WebSocket connection endpoint + */ + get: operations["websocket_live_sessions_websocket_endpoint_get_2"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/v1/mcp/access_groups": { parameters: { query?: never; @@ -60491,6 +60551,24 @@ export interface operations { }; }; }; + websocket_live_sessions_websocket_endpoint_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description WebSocket Protocol Switched */ + 101: { + headers: { + [name: string]: unknown; + }; + content?: never; + }; + }; + }; login_login_post: { parameters: { query?: never; @@ -63035,6 +63113,24 @@ export interface operations { }; }; }; + websocket_live_sessions_websocket_endpoint_get_3: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description WebSocket Protocol Switched */ + 101: { + headers: { + [name: string]: unknown; + }; + content?: never; + }; + }; + }; websocket_realtime_websocket_endpoint_get_3: { parameters: { query?: never; @@ -74276,6 +74372,24 @@ export interface operations { }; }; }; + websocket_live_sessions_websocket_endpoint_get_2: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description WebSocket Protocol Switched */ + 101: { + headers: { + [name: string]: unknown; + }; + content?: never; + }; + }; + }; get_mcp_access_groups_v1_mcp_access_groups_get: { parameters: { query?: never;