mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(realtime): proxy OpenAI Live sessions on /v1/live/sessions
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
ed4caebb65
commit
642a2f196d
19 changed files with 2563 additions and 82 deletions
|
|
@ -116,6 +116,8 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
# Realtime / streaming
|
||||
"/v1/realtime",
|
||||
"/realtime",
|
||||
"/v1/live/sessions",
|
||||
"/live/sessions",
|
||||
# Health & ops
|
||||
"/health",
|
||||
"/metrics",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
3
litellm/llms/openai/live/__init__.py
Normal file
3
litellm/llms/openai/live/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .handler import OpenAILiveSessions
|
||||
|
||||
__all__ = ("OpenAILiveSessions",)
|
||||
478
litellm/llms/openai/live/handler.py
Normal file
478
litellm/llms/openai/live/handler.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
0
tests/unit/llms/openai/live/__init__.py
Normal file
0
tests/unit/llms/openai/live/__init__.py
Normal file
762
tests/unit/llms/openai/live/test_handler.py
Normal file
762
tests/unit/llms/openai/live/test_handler.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
114
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
114
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue