From b82b31a44f1b44b0a4e4edb324305ac10ce6f533 Mon Sep 17 00:00:00 2001 From: Young Han Date: Wed, 2 Sep 2026 13:22:47 -0700 Subject: [PATCH 1/7] feat(realtime): add Meta Muse Voice transcription --- .../litellm_core_utils/realtime_streaming.py | 87 ++- litellm/llms/meta/__init__.py | 3 + litellm/llms/meta/realtime/__init__.py | 10 + litellm/llms/meta/realtime/handler.py | 661 ++++++++++++++++++ litellm/llms/meta/realtime/transformation.py | 619 ++++++++++++++++ litellm/llms/openai_like/providers.json | 2 +- ...odel_prices_and_context_window_backup.json | 27 +- litellm/realtime_api/main.py | 34 +- litellm/types/llms/openai.py | 2 + litellm/types/realtime.py | 12 +- model_prices_and_context_window.json | 27 +- .../test_realtime_streaming.py | 191 ++++- .../realtime/test_meta_realtime_handler.py | 449 ++++++++++++ .../test_meta_realtime_transformation.py | 299 ++++++++ tests/test_litellm/realtime_api/test_main.py | 80 ++- 15 files changed, 2446 insertions(+), 57 deletions(-) create mode 100644 litellm/llms/meta/__init__.py create mode 100644 litellm/llms/meta/realtime/__init__.py create mode 100644 litellm/llms/meta/realtime/handler.py create mode 100644 litellm/llms/meta/realtime/transformation.py create mode 100644 tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py create mode 100644 tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 4923bdda305..4be7dd6b4ce 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -19,7 +19,7 @@ from litellm.types.llms.openai import ( OpenAIRealtimeStreamResponseBaseObject, OpenAIRealtimeStreamSessionEvents, ) -from litellm.types.realtime import ALL_DELTA_TYPES +from litellm.types.realtime import ALL_DELTA_TYPES, RealtimeInputAudioTranscriptionUsage from .litellm_logging import Logging as LiteLLMLogging from .realtime_errors import client_close_code, realtime_error_event, websocket_close_reason @@ -116,6 +116,10 @@ class RealtimeEventNormalizer(Protocol): def patch_outgoing_session(self, session: dict) -> dict: ... +class RealtimeUsageProvider(Protocol): + def unbilled_usage_on_session_close(self, model: str) -> RealtimeInputAudioTranscriptionUsage | None: ... + + DefaultLoggedRealTimeEventTypes: Final = [ "session.created", "response.create", @@ -139,6 +143,8 @@ class RealTimeStreaming: force_transcription_model: str | None = None, event_normalizer: RealtimeEventNormalizer | None = None, logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER, + usage_provider: RealtimeUsageProvider | None = None, + exclude_private_content_from_logs: bool = False, ): self.websocket: _ClientWebSocket = websocket self.backend_ws = backend_ws @@ -200,6 +206,10 @@ class RealTimeStreaming: self._is_transcription_session: bool = force_transcription_model is not None # Optional per-provider GA event normalizer (e.g. XAIRealtimeNormalizer). self._event_normalizer = event_normalizer + self._usage_provider: RealtimeUsageProvider | None = ( + usage_provider if usage_provider is not None else provider_config + ) + self._exclude_private_content_from_logs = exclude_private_content_from_logs # Per-connection caps for pre-setup audio frames (message count + total bytes). _MAX_BUFFERED_MESSAGES: int = 200 @@ -237,7 +247,7 @@ class RealTimeStreaming: def _should_store_message( self, - message_obj: dict | OpenAIRealtimeEvents, + message_obj: dict[str, Any] | OpenAIRealtimeEvents, # mutable-ok: existing realtime event contract ) -> bool: _msg_type: Final = message_obj["type"] if "type" in message_obj else None if self.logged_real_time_event_types == "*": @@ -246,16 +256,54 @@ class RealTimeStreaming: return True return False + def _message_for_logging( + self, + message_obj: dict[str, Any], # mutable-ok: existing realtime event contract + ) -> dict[str, Any]: # mutable-ok: logging stores concrete event dictionaries + if not self._exclude_private_content_from_logs: + return message_obj + logged_message: dict[str, Any] = { # mutable-ok: incrementally builds the sanitized event copy + key: message_obj[key] + for key in ( + "type", + "event_id", + "item_id", + "response_id", + "conversation_id", + "session_id", + "content_index", + "output_index", + "model", + "mode", + "usage", + ) + if key in message_obj + } + session: Final = message_obj.get("session") + if isinstance(session, dict): + logged_session: Final[dict[str, Any]] = { # mutable-ok: sanitized JSON session snapshot + key: session[key] for key in ("id", "model", "mode", "type") if key in session + } + if logged_session: + logged_message["session"] = logged_session + return logged_message + def store_message(self, message: str | bytes | dict | OpenAIRealtimeEvents): """Store message in list""" if isinstance(message, bytes): message = message.decode("utf-8") if isinstance(message, dict): # TypedDict union members do not narrow to plain dict for mypy. - message_obj: dict[str, Any] = cast(dict[str, Any], message) + parsed_message_obj: dict[str, Any] = cast( # cast-ok: TypedDict events are JSON dictionaries + dict[str, Any], message + ) else: - message_obj = cast(dict[str, Any], json.loads(cast(str, message))) - self._collect_tool_calls_from_response_done(cast(dict, message_obj)) + parsed_message_obj = cast( # cast-ok: parsed realtime events are JSON dictionaries + dict[str, Any], json.loads(message) + ) + if not self._exclude_private_content_from_logs: + self._collect_tool_calls_from_response_done(parsed_message_obj) + message_obj: Final = self._message_for_logging(parsed_message_obj) if not self._should_store_message(message_obj): return try: @@ -273,6 +321,8 @@ class RealTimeStreaming: def _collect_user_input_from_client_event(self, message: str | dict) -> None: """Extract user text content from client WebSocket events for spend logging.""" + if self._exclude_private_content_from_logs: + return try: if isinstance(message, str): msg_obj = json.loads(message) @@ -309,6 +359,8 @@ class RealTimeStreaming: def _collect_user_input_from_backend_event(self, event_obj: dict | OpenAIRealtimeEvents) -> None: """Extract user voice transcription from backend events for spend logging.""" + if self._exclude_private_content_from_logs: + return try: event_type: Final = event_obj.get("type", "") if event_type == "conversation.item.input_audio_transcription.completed": @@ -364,9 +416,9 @@ class RealTimeStreaming: pass def _flush_unbilled_transcription_usage(self) -> None: - if self.provider_config is None: + if self._usage_provider is None: return - usage: Final = self.provider_config.unbilled_usage_on_session_close(self.model) + usage: Final = self._usage_provider.unbilled_usage_on_session_close(self.model) if usage is None: return flush_event: Final = ( @@ -403,12 +455,27 @@ class RealTimeStreaming: except (AttributeError, TypeError): pass + def _input_for_logging( + self, + message: str | dict, # mutable-ok: existing realtime input contract + ) -> str | dict: # mutable-ok: logging stores concrete event dictionaries + if not self._exclude_private_content_from_logs: + return message + try: + parsed_message: Final[object] = message if isinstance(message, dict) else json.loads(message) + except (json.JSONDecodeError, TypeError): + return {} # mutable-ok: empty JSON logging payload + if not isinstance(parsed_message, dict): + return {} # mutable-ok: empty JSON logging payload + return self._message_for_logging(parsed_message) + def store_input(self, message: str | dict): """Store input message""" - self.input_message = message if isinstance(message, dict) else {} + logged_message: Final[str | dict] = self._input_for_logging(message) # mutable-ok: logging payload + self.input_message = logged_message if isinstance(logged_message, dict) else {} self._collect_user_input_from_client_event(message) if self.logging_obj: - self.logging_obj.pre_call(input=message, api_key="") + self.logging_obj.pre_call(input=logged_message, api_key="") async def log_messages(self): """Log messages in list""" @@ -1009,6 +1076,8 @@ class RealTimeStreaming: self.store_message(event_str) self._capture_transcription_usage(event) await self._send_event_to_client(event, event_str) + if self._is_transcription_session: + continue blocked = await self.run_realtime_guardrails( cast(str, transcript), item_id=cast(str | None, event.get("item_id")), diff --git a/litellm/llms/meta/__init__.py b/litellm/llms/meta/__init__.py new file mode 100644 index 00000000000..7c7d32788a2 --- /dev/null +++ b/litellm/llms/meta/__init__.py @@ -0,0 +1,3 @@ +from .realtime import MetaRealtime, MuseRealtimeAdapter + +__all__ = ("MetaRealtime", "MuseRealtimeAdapter") diff --git a/litellm/llms/meta/realtime/__init__.py b/litellm/llms/meta/realtime/__init__.py new file mode 100644 index 00000000000..6398765da24 --- /dev/null +++ b/litellm/llms/meta/realtime/__init__.py @@ -0,0 +1,10 @@ +from .handler import MetaRealtime, MuseRealtimeAdapter +from .transformation import MuseEventTransformer, MuseProtocolError, MuseSessionConfig + +__all__ = ( + "MetaRealtime", + "MuseEventTransformer", + "MuseProtocolError", + "MuseRealtimeAdapter", + "MuseSessionConfig", +) diff --git a/litellm/llms/meta/realtime/handler.py b/litellm/llms/meta/realtime/handler.py new file mode 100644 index 00000000000..9eb98e4ba0b --- /dev/null +++ b/litellm/llms/meta/realtime/handler.py @@ -0,0 +1,661 @@ +from __future__ import annotations + +import asyncio +import base64 +import binascii +import contextlib +import json +import time +import uuid +from collections.abc import Awaitable, Callable, Mapping +from typing import Final, Protocol +from urllib.parse import urlparse, urlunparse + +from pydantic import JsonValue, TypeAdapter, ValidationError + +from litellm._logging import verbose_proxy_logger +from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging +from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming +from litellm.llms.custom_httpx.http_handler import get_shared_realtime_ssl_context +from litellm.types.realtime import RealtimeInputAudioTranscriptionUsage, RealtimeQueryParams + +from .transformation import ( + MUSE_MODEL, + MuseEventTransformer, + MuseProtocolError, + MuseSessionConfig, + encode_event, + error_event, + parse_session_update, + session_created_event, + session_updated_event, +) + +DEFAULT_MUSE_REALTIME_URL: Final = "wss://api.meta.ai/v1/asr/realtime" +_MAX_AUDIO_BACKLOG_SECONDS: Final = 4 +_MAX_PENDING_PROVIDER_EVENTS: Final = 256 +_JSON_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) + + +class _ProviderWebSocket(Protocol): + async def send(self, message: str | bytes) -> None: ... + + async def recv(self, decode: bool | None = None) -> str | bytes: ... + + async def close(self, code: int = 1000, reason: str = "") -> None: ... + + +class _ClientWebSocketExceptions(Protocol): + ConnectionClosed: type[Exception] + + +class _ClientWebSocket(Protocol): + exceptions: _ClientWebSocketExceptions + + @property + def scope(self) -> Mapping[str, object]: ... + + async def send_text(self, data: str) -> None: ... + + async def receive_text(self) -> str: ... + + async def close(self, code: int = 1000, reason: str | None = None) -> None: ... + + +class WebSocketConnect(Protocol): + def __call__( + self, + url: str, + *, + open_timeout: float, + max_size: int | None, + ssl: object | None, + ) -> Awaitable[_ProviderWebSocket]: ... + + +class MuseAdapterError(RuntimeError): + def __init__(self, message: str, *, close_code: int) -> None: + super().__init__(message) + self.close_code: Final = close_code + + +class MuseRealtimeAdapter: + def __init__( + self, + *, + model: str, + api_key: str, + api_base: str | None = None, + timeout: float | None = None, + websocket_connect: WebSocketConnect | None = None, + monotonic: Callable[[], float] = time.monotonic, + sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, + terminate_client: Callable[[int], Awaitable[None]] | None = None, + ) -> None: + if model.removeprefix("meta/") != MUSE_MODEL: + raise ValueError("unsupported Meta realtime model") + self._model: Final = model.removeprefix("meta/") + self._access_token: Final = normalize_access_token(api_key) + self._url: Final = build_muse_realtime_url(api_base) + self._timeout: Final = timeout or 10.0 + self._websocket_connect = websocket_connect + self._monotonic: Final = monotonic + self._sleep: Final = sleep + self._terminate_client: Final = terminate_client + self._provider_ws: _ProviderWebSocket | None = None + self._config: MuseSessionConfig | None = None + self._session_id: str = f"sess_{uuid.uuid4().hex}" + self._events: Final[asyncio.Queue[str | BaseException]] = asyncio.Queue(maxsize=_MAX_PENDING_PROVIDER_EVENTS) + self._events.put_nowait(encode_event(session_created_event(self._model, self._session_id))) + self._transformer: Final = MuseEventTransformer() + self._audio_condition: Final = asyncio.Condition() + self._pending_audio: bytearray = bytearray() + self._audio_generation: int = 0 + self._flush_requested: bool = False + self._end_requested: bool = False + self._end_stream_sent: bool = False + self._audio_consumed: bool = False + self._closed: bool = False + self._resources_closed: bool = False + self._sender_task: asyncio.Task[None] | None = None + self._receiver_task: asyncio.Task[None] | None = None + self.close_code: int = 1000 + self.close_reason: str = "Session closed" + + async def send(self, message: str | bytes) -> None: + if self._closed: + raise MuseAdapterError("Meta Muse realtime session is closed", close_code=self.close_code) + if isinstance(message, bytes): + await self._reject("invalid_request_error", "invalid_event", "Client events must be JSON text") + return + try: + event: Final = _parse_client_event(message) + event_type: Final = event.get("type") + if event_type in ("session.update", "transcription_session.update"): + await self._handle_session_update(message) + return + if event_type == "input_audio_buffer.append": + await self._handle_audio_append(event) + return + if event_type == "input_audio_buffer.clear": + await self._clear_audio() + return + if event_type == "input_audio_buffer.commit": + await self._commit_audio() + return + if event_type == "input_audio_buffer.end": + await self._end_audio() + return + await self._emit( + error_event( + "invalid_request_error", + "unsupported_event", + f"Event type {event_type!r} is not supported for Meta Muse transcription", + ) + ) + except MuseProtocolError as exc: + await self._reject("invalid_request_error", "invalid_event", str(exc)) + + async def recv(self, decode: bool | None = None) -> str | bytes: + event: Final = await self._events.get() + if isinstance(event, BaseException): + close_code: Final = _exception_close_code(event) if isinstance(event, Exception) else 1011 + if self._terminate_client is not None: + await self._terminate_client(close_code) + raise event + return event.encode("utf-8") if decode is False else event + + async def close(self, code: int = 1000, reason: str = "") -> None: + if self._resources_closed: + return + self._closed = True + self._resources_closed = True + self.close_code = sanitize_close_code(code) + self.close_reason = safe_close_reason(self.close_code) + async with self._audio_condition: + self._end_requested = True + self._audio_condition.notify_all() + tasks: Final = tuple(task for task in (self._sender_task, self._receiver_task) if task is not None) + for task in tasks: + task.cancel() + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + provider_ws: Final = self._provider_ws + if provider_ws is not None: + with contextlib.suppress(Exception): + await provider_ws.close(code=self.close_code, reason=self.close_reason) + + def unbilled_usage_on_session_close(self, model: str) -> RealtimeInputAudioTranscriptionUsage | None: + return self._transformer.take_unbilled_usage() + + async def _handle_session_update(self, message: str) -> None: + config: Final = parse_session_update(message, self._model) + if self._config is not None: + if config != self._config: + await self._reject( + "invalid_request_error", + "session_configuration_locked", + "Meta Muse session configuration cannot change after setup", + ) + return + await self._emit(session_updated_event(config, self._session_id)) + return + await self._connect(config) + + async def _connect(self, config: MuseSessionConfig) -> None: + connector: Final = self._websocket_connect or _default_websocket_connect + last_error: Exception | None = None # rebind-ok: records the latest bounded handshake attempt + for attempt in range(2): + provider_ws: _ProviderWebSocket | None = None + try: + provider_ws = await connector( + self._url, + open_timeout=self._timeout, + max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, + ssl=_ssl_config(self._url), + ) + await provider_ws.send(json.dumps(config.handshake(self._access_token), separators=(",", ":"))) + raw_ack: str | bytes = await asyncio.wait_for( # rebind-ok: one response per handshake attempt + provider_ws.recv(), timeout=self._timeout + ) + session_id: str = _parse_handshake_ack(raw_ack) # rebind-ok: one ID per handshake attempt + self._provider_ws = provider_ws + self._config = config + self._transformer.configure(config) + self._session_id = session_id + self._sender_task = asyncio.create_task(self._send_audio(), name="meta-muse-realtime-send") + self._receiver_task = asyncio.create_task(self._receive_events(), name="meta-muse-realtime-receive") + await self._emit(session_updated_event(config, session_id)) + return + except asyncio.CancelledError: + if provider_ws is not None: + with contextlib.suppress(Exception): + await provider_ws.close() + raise + except Exception as exc: # noqa: BLE001 # connector implementations expose heterogeneous transport errors + last_error = exc + if provider_ws is not None: + with contextlib.suppress(Exception): + await provider_ws.close() + close_code: int = _exception_close_code(exc) # rebind-ok: classified per handshake attempt + retryable_transport_error: bool = not isinstance( # rebind-ok: classified per handshake attempt + exc, (MuseAdapterError, MuseProtocolError) + ) + if attempt == 0 and retryable_transport_error and close_code in (1011, 1013): + continue + self.close_code = close_code + self.close_reason = safe_close_reason(close_code) + await self._emit( + error_event( + "server_error" if close_code != 1008 else "invalid_request_error", + "provider_connection_error", + "Meta Muse realtime handshake failed", + ) + ) + await self._events.put(MuseAdapterError("Meta Muse realtime handshake failed", close_code=close_code)) + await self._mark_terminated(close_code) + return + assert last_error is not None + raise MuseAdapterError("Meta Muse realtime handshake failed", close_code=1011) + + async def _handle_audio_append(self, event: Mapping[str, JsonValue]) -> None: + config: Final = self._require_configured() + if self._end_requested or self._end_stream_sent: + await self._reject("invalid_request_error", "input_ended", "Audio input has already ended") + return + audio_value: Final = event.get("audio") + if not isinstance(audio_value, str): + await self._reject("invalid_request_error", "invalid_audio", "Audio must be a base64 string") + return + try: + audio: Final = base64.b64decode(audio_value, validate=True) + except (binascii.Error, ValueError): + await self._reject("invalid_request_error", "invalid_audio", "Audio must be valid base64") + return + if len(audio) % 2: + await self._reject("invalid_request_error", "invalid_audio", "PCM16 audio must contain complete samples") + return + if not audio: + return + max_backlog_bytes: Final = config.bytes_per_second * _MAX_AUDIO_BACKLOG_SECONDS + if len(audio) > max_backlog_bytes: + await self._reject( + "invalid_request_error", + "audio_backlog_exceeded", + "Audio append exceeds the four-second Muse backlog limit", + ) + return + async with self._audio_condition: + await self._audio_condition.wait_for( + lambda: self._closed or len(self._pending_audio) + len(audio) <= max_backlog_bytes + ) + if self._closed: + raise MuseAdapterError("Meta Muse realtime session is closed", close_code=self.close_code) + self._pending_audio.extend(audio) + self._audio_condition.notify_all() + + async def _clear_audio(self) -> None: + self._require_configured() + async with self._audio_condition: + self._pending_audio.clear() + self._audio_generation += 1 + self._flush_requested = False + self._audio_condition.notify_all() + await self._emit( + { # mutable-ok: OpenAI-compatible JSON event + "type": "input_audio_buffer.cleared", + "event_id": f"event_{uuid.uuid4().hex}", + } + ) + + async def _commit_audio(self) -> None: + config: Final = self._require_configured() + previous_item_id, item_id = self._transformer.commit_item() + async with self._audio_condition: + self._flush_requested = True + if config.mode == "PUSH_TO_TALK": + self._end_requested = True + self._audio_condition.notify_all() + await self._emit( + { # mutable-ok: OpenAI-compatible JSON event + "type": "input_audio_buffer.committed", + "event_id": f"event_{uuid.uuid4().hex}", + "previous_item_id": previous_item_id, + "item_id": item_id, + } + ) + + async def _end_audio(self) -> None: + self._require_configured() + async with self._audio_condition: + self._flush_requested = True + self._end_requested = True + self._audio_condition.notify_all() + + async def _send_audio(self) -> None: + config: Final = self._require_configured() + provider_ws: Final = self._require_provider_ws() + pacing_origin: float | None = None # rebind-ok: initialized when the first packet is ready + sent_duration: float = 0.0 # rebind-ok: absolute pacing clock advances after each packet + try: + while True: + packet, pacing_origin, ended = await self._next_audio_packet( + config, + pacing_origin, + sent_duration, + ) + if ended: + break + if packet is None: + continue + await provider_ws.send(packet) + self._audio_consumed = True + sent_duration += len(packet) / config.bytes_per_second + await self._send_end_stream() + except asyncio.CancelledError: + raise + except Exception as exc: # noqa: BLE001 # WebSocket implementations expose heterogeneous transport errors + await self._fail_provider(exc, phase="audio send") + + async def _next_audio_packet( + self, + config: MuseSessionConfig, + pacing_origin: float | None, + sent_duration: float, + ) -> tuple[bytes | None, float | None, bool]: + async with self._audio_condition: + await self._audio_condition.wait_for( + lambda: ( + self._closed + or len(self._pending_audio) >= config.packet_bytes + or (self._flush_requested and bool(self._pending_audio)) + or (self._end_requested and not self._pending_audio) + ) + ) + if self._closed or (self._end_requested and not self._pending_audio): + return None, pacing_origin, True + packet_size: Final = min(config.packet_bytes, len(self._pending_audio)) + if packet_size < config.packet_bytes and not self._flush_requested: + return None, pacing_origin, False + generation: Final = self._audio_generation + current_time: Final = self._monotonic() + effective_origin: Final = ( + current_time - sent_duration + if pacing_origin is None or current_time > pacing_origin + sent_duration + else pacing_origin + ) + deadline: Final = effective_origin + sent_duration + delay: Final = deadline - self._monotonic() + if delay > 0: + await self._sleep(delay) + async with self._audio_condition: + if generation != self._audio_generation: + return None, effective_origin, False + actual_size: Final = min(packet_size, len(self._pending_audio)) + packet: Final = bytes(self._pending_audio[:actual_size]) + del self._pending_audio[:actual_size] + if not self._pending_audio: + self._flush_requested = False + self._audio_condition.notify_all() + return packet or None, effective_origin, False + + async def _send_end_stream(self) -> None: + if self._end_stream_sent: + return + provider_ws: Final = self._require_provider_ws() + await provider_ws.send('{"type":"endStream"}') + self._end_stream_sent = True + + async def _receive_events(self) -> None: + provider_ws: Final = self._require_provider_ws() + try: + while True: + raw: str | bytes = await provider_ws.recv() # rebind-ok: one provider frame per iteration + if not isinstance(raw, str): + raise MuseProtocolError("provider returned a non-text event") + for event in self._transformer.transform(raw): + await self._emit(event) + except asyncio.CancelledError: + raise + except Exception as exc: # noqa: BLE001 # provider close exceptions vary by WebSocket implementation + close_code: Final = _exception_close_code(exc) + if close_code == 1000 and self._end_stream_sent: + await self._mark_terminated(1000) + await self._events.put(MuseAdapterError("Meta Muse realtime session completed", close_code=1000)) + return + failure: Final = MuseAdapterError( + "Meta Muse realtime closed before input ended", + close_code=1011 if close_code == 1000 else close_code, + ) + await self._fail_provider(failure, phase="receive") + + async def _fail_provider(self, exc: Exception, *, phase: str) -> None: + close_code: Final = _exception_close_code(exc) + self.close_code = close_code + self.close_reason = safe_close_reason(close_code) + await self._emit( + error_event( + "server_error", + "provider_connection_error", + f"Meta Muse realtime {phase} failed", + ) + ) + await self._events.put(MuseAdapterError(f"Meta Muse realtime {phase} failed", close_code=close_code)) + await self._mark_terminated(close_code) + + async def _mark_terminated(self, close_code: int) -> None: + self._closed = True + self.close_code = sanitize_close_code(close_code) + self.close_reason = safe_close_reason(self.close_code) + async with self._audio_condition: + self._audio_condition.notify_all() + + async def _terminate(self, close_code: int) -> None: + await self._mark_terminated(close_code) + if self._terminate_client is not None: + await self._terminate_client(self.close_code) + + async def _reject(self, error_type: str, code: str, message: str) -> None: + self.close_code = 1008 + self.close_reason = safe_close_reason(1008) + await self._emit(error_event(error_type, code, message)) + await self._events.put(MuseAdapterError(message, close_code=1008)) + await self._mark_terminated(1008) + + async def _emit(self, event: Mapping[str, object]) -> None: + await self._events.put(encode_event(event)) + + def _require_configured(self) -> MuseSessionConfig: + if self._config is None: + raise MuseProtocolError("send session.update before audio events") + return self._config + + def _require_provider_ws(self) -> _ProviderWebSocket: + if self._provider_ws is None: + raise MuseProtocolError("Meta Muse provider connection is not ready") + return self._provider_ws + + +class MetaRealtime: + async def async_realtime( + self, + model: str, + websocket: _ClientWebSocket, + logging_obj: LiteLLMLogging, + api_base: str | None = None, + api_key: str | None = None, + client: object | None = None, + timeout: float | None = None, + query_params: RealtimeQueryParams | None = None, + user_api_key_dict: object | None = None, + litellm_metadata: Mapping[str, object] | None = None, + websocket_connect: WebSocketConnect | None = None, + monotonic: Callable[[], float] = time.monotonic, + sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, + **kwargs: object, # kwargs-ok: realtime dispatcher forwards provider-neutral options + ) -> None: + if api_key is None or not api_key.strip(): + await _send_client_error_and_close(websocket, "Meta Model API key is required") + return + try: + adapter: Final = MuseRealtimeAdapter( + model=model, + api_key=api_key, + api_base=api_base, + timeout=timeout, + websocket_connect=websocket_connect, + monotonic=monotonic, + sleep=sleep, + terminate_client=lambda code: _close_client(websocket, code), + ) + except ValueError: + await _send_client_error_and_close(websocket, "Invalid Meta Muse realtime configuration") + return + realtime_streaming: Final = RealTimeStreaming( + websocket, + adapter, # pyright: ignore[reportArgumentType] # raw adapter intentionally matches the websocket surface + logging_obj, + model=model, + user_api_key_dict=user_api_key_dict, + request_data={ # mutable-ok: relay request metadata payload + "litellm_metadata": dict(litellm_metadata or {}) # mutable-ok: relay owns its metadata copy + }, + force_transcription_model=model, + usage_provider=adapter, + exclude_private_content_from_logs=True, + ) + try: + await realtime_streaming.bidirectional_forward() + except MuseAdapterError as exc: + adapter.close_code = exc.close_code + adapter.close_reason = safe_close_reason(exc.close_code) + except Exception: # noqa: BLE001 # relay errors are normalized before closing the accepted client socket + adapter.close_code = 1011 + adapter.close_reason = safe_close_reason(1011) + verbose_proxy_logger.exception("Meta Muse realtime session failed") + finally: + await adapter.close(code=adapter.close_code) + await _close_client(websocket, adapter.close_code) + + +def normalize_access_token(api_key: str) -> str: + stripped: Final = api_key.strip() + if not stripped: + raise ValueError("Meta Model API key is required") + parts: Final = stripped.split(None, 1) + if parts[0].casefold() == "bearer": + if len(parts) != 2 or not parts[1].strip(): + raise ValueError("Meta Model API key must include a token after Bearer") + return f"Bearer {parts[1].strip()}" + return f"Bearer {stripped}" + + +def build_muse_realtime_url(api_base: str | None) -> str: + if api_base is None: + return DEFAULT_MUSE_REALTIME_URL + parsed: Final = urlparse(api_base.strip()) + scheme: Final = "wss" if parsed.scheme == "https" else parsed.scheme + if ( + scheme != "wss" + or not parsed.hostname + or parsed.username is not None + or parsed.password is not None + or parsed.fragment + ): + raise ValueError("Meta api_base must be an absolute wss:// or https:// URL without credentials or a fragment") + netloc: Final = f"{parsed.hostname}:{parsed.port}" if parsed.port is not None else parsed.hostname + return urlunparse((scheme, netloc, "/v1/asr/realtime", "", "", "")) + + +def sanitize_close_code(code: int | None) -> int: + if code is not None and code in (1000, 1008, 1011, 1013): + return code + return 1011 + + +def safe_close_reason(code: int) -> str: + return { # mutable-ok: immutable-by-convention close-reason lookup + 1000: "Session closed", + 1008: "Invalid realtime transcription request", + 1011: "Realtime transcription service error", + 1013: "Realtime transcription service unavailable", + }.get(code, "Realtime transcription service error") + + +def _parse_client_event(payload: str) -> Mapping[str, JsonValue]: + try: + value: Final = _JSON_ADAPTER.validate_json(payload) + except ValidationError: + raise MuseProtocolError("invalid JSON object") from None + if not isinstance(value, dict): + raise MuseProtocolError("message must be a JSON object") + event_type: Final = value.get("type") + if not isinstance(event_type, str) or not event_type: + raise MuseProtocolError("message type must be a non-empty string") + return value + + +def _parse_handshake_ack(raw: str | bytes) -> str: + if not isinstance(raw, str): + raise MuseProtocolError("provider returned a non-text handshake response") + message: Final = _parse_json_object(raw) + if message.get("type") == "error": + raise MuseAdapterError("Meta Muse realtime handshake was rejected", close_code=1008) + session_id: Final = message.get("sessionId") + if not isinstance(session_id, str) or not session_id.strip(): + raise MuseProtocolError("provider returned an invalid handshake response") + return session_id.strip() + + +def _parse_json_object(payload: str) -> Mapping[str, JsonValue]: + try: + value: Final = _JSON_ADAPTER.validate_json(payload) + except ValidationError: + raise MuseProtocolError("invalid provider JSON object") from None + if not isinstance(value, dict): + raise MuseProtocolError("provider message must be a JSON object") + return value + + +def _exception_close_code(exc: Exception) -> int: + code: Final = getattr(exc, "code", None) + if isinstance(exc, MuseAdapterError): + return sanitize_close_code(exc.close_code) + return sanitize_close_code(code if isinstance(code, int) else None) + + +def _ssl_config(url: str) -> object | None: + if not url.startswith("wss://"): + return None + config: Final = get_shared_realtime_ssl_context() + return True if config is False else config + + +async def _default_websocket_connect( + url: str, + *, + open_timeout: float, + max_size: int | None, + ssl: object | None, +) -> _ProviderWebSocket: + import websockets + + connection: Final = await websockets.connect( + url, + open_timeout=open_timeout, + max_size=max_size, + ssl=ssl, # pyright: ignore[reportArgumentType] # shared SSL helper returns the library-supported union + ) + return connection + + +async def _send_client_error_and_close(websocket: _ClientWebSocket, message: str) -> None: + with contextlib.suppress(Exception): + await websocket.send_text(encode_event(error_event("invalid_request_error", "invalid_configuration", message))) + await _close_client(websocket, 1008) + + +async def _close_client(websocket: _ClientWebSocket, code: int) -> None: + with contextlib.suppress(Exception): + await websocket.close(code=sanitize_close_code(code), reason=safe_close_reason(sanitize_close_code(code))) diff --git a/litellm/llms/meta/realtime/transformation.py b/litellm/llms/meta/realtime/transformation.py new file mode 100644 index 00000000000..f37295892f8 --- /dev/null +++ b/litellm/llms/meta/realtime/transformation.py @@ -0,0 +1,619 @@ +from __future__ import annotations + +import json +import math +import uuid +from collections import OrderedDict, deque +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Final, Literal, TypeAlias + +from pydantic import JsonValue, TypeAdapter, ValidationError + +from litellm.types.realtime import RealtimeInputAudioTranscriptionUsage + +MUSE_MODEL: Final = "muse-voice-transcribe-1.0" +SUPPORTED_SAMPLE_RATES: Final = frozenset((16_000, 24_000)) +SUPPORTED_MODES: Final = frozenset(("PUSH_TO_TALK", "ENDPOINTING", "DIARIZATION")) +SUPPORTED_LANGUAGES: Final = ( + "Arabic", + "Bengali", + "Dutch", + "English", + "French", + "German", + "Hebrew", + "Hindi", + "Indonesian", + "Italian", + "Japanese", + "Kannada", + "Korean", + "Malay", + "Mandarin Chinese", + "Marathi", + "Polish", + "Portuguese", + "Spanish", + "Tagalog", + "Tamil", + "Telugu", + "Thai", + "Turkish", + "Vietnamese", +) +_LANGUAGE_NAMES: Final = { # mutable-ok: immutable-by-convention language lookup table + language.casefold(): language for language in SUPPORTED_LANGUAGES +} +_LANGUAGE_CODES: Final = { # mutable-ok: immutable-by-convention language lookup table + "ar": "Arabic", + "bn": "Bengali", + "de": "German", + "en": "English", + "es": "Spanish", + "fil": "Tagalog", + "fr": "French", + "he": "Hebrew", + "hi": "Hindi", + "id": "Indonesian", + "it": "Italian", + "iw": "Hebrew", + "ja": "Japanese", + "kn": "Kannada", + "ko": "Korean", + "ms": "Malay", + "mr": "Marathi", + "nl": "Dutch", + "pl": "Polish", + "pt": "Portuguese", + "ta": "Tamil", + "te": "Telugu", + "th": "Thai", + "tl": "Tagalog", + "tr": "Turkish", + "vi": "Vietnamese", + "zh": "Mandarin Chinese", +} +_JSON_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +OpenAIEvent: TypeAlias = Mapping[str, object] + + +class MuseProtocolError(ValueError): + pass + + +@dataclass(frozen=True, slots=True) +class MuseSessionConfig: + model: str + mode: Literal["PUSH_TO_TALK", "ENDPOINTING", "DIARIZATION"] + sample_rate: Literal[16000, 24000] + keywords: tuple[str, ...] + language_bias: tuple[str, ...] + + @property + def audio_encoding(self) -> Literal["PCM_16KHZ", "PCM_24KHZ"]: + return "PCM_16KHZ" if self.sample_rate == 16_000 else "PCM_24KHZ" + + @property + def bytes_per_second(self) -> int: + return self.sample_rate * 2 + + @property + def packet_bytes(self) -> int: + return self.bytes_per_second * 80 // 1000 + + def handshake(self, access_token: str) -> Mapping[str, object]: + base: Final[Mapping[str, object]] = { # mutable-ok: JSON wire payload + "mode": self.mode, + "authorization": {"accessToken": access_token}, # mutable-ok: JSON wire payload + "audioEncoding": self.audio_encoding, + "model": self.model, + "partialMode": "CUMULATIVE", + "emitAudioProgress": True, + } + payload: dict[str, object] = dict(base) # mutable-ok: incrementally builds JSON wire payload + if self.keywords: + payload["keywords"] = list(self.keywords) # mutable-ok: JSON arrays require concrete lists + if self.language_bias: + payload["languageBias"] = list(self.language_bias) # mutable-ok: JSON arrays require concrete lists + return payload + + def openai_session(self, session_id: str) -> Mapping[str, object]: + turn_detection: Final[Mapping[str, object] | None] = ( + None if self.mode == "PUSH_TO_TALK" else {"type": "server_vad"} # mutable-ok: JSON wire payload + ) + transcription: dict[str, object] = { # mutable-ok: incrementally builds JSON wire payload + "model": self.model, + } + if self.language_bias: + transcription["language"] = self.language_bias[0] + transcription["language_bias"] = list( # mutable-ok: JSON arrays require concrete lists + self.language_bias + ) + if self.keywords: + transcription["keywords"] = list(self.keywords) # mutable-ok: JSON arrays require concrete lists + return { # mutable-ok: JSON wire payload + "id": session_id, + "object": "realtime.transcription_session", + "type": "transcription", + "model": self.model, + "audio": { # mutable-ok: JSON wire payload + "input": { # mutable-ok: JSON wire payload + "format": {"type": "audio/pcm", "rate": self.sample_rate}, # mutable-ok: JSON wire payload + "transcription": transcription, + "turn_detection": turn_detection, + } + }, + } + + +@dataclass(slots=True) +class _TurnState: + item_id: str | None = None + started: bool = False + start_emitted: bool = False + latest_partial: str | None = None + emitted_partial: str = "" + final_text: str | None = None + completed_signal: bool = False + completed_emitted: bool = False + stopped: bool = False + stopped_emitted: bool = False + speaker: str | None = None + + +def _json_object(payload: str) -> Mapping[str, JsonValue]: + try: + value: Final = _JSON_ADAPTER.validate_json(payload) + except ValidationError: + raise MuseProtocolError("invalid JSON object") from None + if not isinstance(value, dict): + raise MuseProtocolError("message must be a JSON object") + return value + + +def _mapping(value: JsonValue | None, name: str) -> Mapping[str, JsonValue]: + if value is None: + return {} # mutable-ok: empty JSON object + if not isinstance(value, dict): + raise MuseProtocolError(f"{name} must be an object") + return value + + +def _string(value: JsonValue | None, name: str) -> str | None: + if value is None: + return None + if not isinstance(value, str): + raise MuseProtocolError(f"{name} must be a string") + return value + + +def _normalize_model(model: str) -> str: + return model.removeprefix("meta/").strip() + + +def normalize_language(language: str) -> str: + value: Final = language.strip() + if not value: + raise MuseProtocolError("language must be non-empty") + documented_name: Final = _LANGUAGE_NAMES.get(value.casefold()) + if documented_name is not None: + return documented_name + primary: Final = value.replace("_", "-").split("-", 1)[0].casefold() + mapped_name: Final = _LANGUAGE_CODES.get(primary) + if mapped_name is None: + raise MuseProtocolError("unsupported Muse Voice language") + return mapped_name + + +def _normalize_string_sequence(value: JsonValue | None, name: str) -> tuple[str, ...]: + if value is None: + return () + if not isinstance(value, list): + raise MuseProtocolError(f"{name} must be an array of strings") + normalized: list[str] = [] # mutable-ok: deduplicates validated language hints before freezing + for entry in value: + if not isinstance(entry, str) or not entry.strip(): + raise MuseProtocolError(f"{name} entries must be non-empty strings") + item: str = entry.strip() # rebind-ok: normalized once for each hint + if item not in normalized: + normalized.append(item) + return tuple(normalized) + + +def _normalize_language_sequence(value: JsonValue | None) -> tuple[str, ...]: + return tuple(dict.fromkeys(normalize_language(item) for item in _normalize_string_sequence(value, "language_bias"))) + + +def _parse_sample_rate(session: Mapping[str, JsonValue]) -> Literal[16000, 24000]: + beta_format: Final = session.get("input_audio_format") + audio: Final = _mapping(session.get("audio"), "session.audio") + audio_input: Final = _mapping(audio.get("input"), "session.audio.input") + ga_format: Final = audio_input.get("format") + if beta_format is not None and ga_format is not None: + raise MuseProtocolError("input audio format must use either beta or GA layout") + if beta_format is not None: + if beta_format != "pcm16": + raise MuseProtocolError("Muse Voice requires pcm16 input audio") + return 24_000 + if ga_format is None: + return 24_000 + if isinstance(ga_format, str): + if ga_format != "pcm16": + raise MuseProtocolError("Muse Voice requires audio/pcm input audio") + return 24_000 + format_mapping: Final = _mapping(ga_format, "session.audio.input.format") + if format_mapping.get("type") != "audio/pcm": + raise MuseProtocolError("Muse Voice requires audio/pcm input audio") + channels: Final = format_mapping.get("channels", 1) + if isinstance(channels, bool) or channels != 1: + raise MuseProtocolError("Muse Voice requires mono input audio") + rate: Final = format_mapping.get("rate", 24_000) + if isinstance(rate, bool) or not isinstance(rate, int) or rate not in SUPPORTED_SAMPLE_RATES: + raise MuseProtocolError("Muse Voice supports PCM16 at 16000 Hz or 24000 Hz") + return rate + + +def _parse_mode( + session: Mapping[str, JsonValue], audio_input: Mapping[str, JsonValue] +) -> Literal["PUSH_TO_TALK", "ENDPOINTING", "DIARIZATION"]: + explicit: Final = session.get("mode") + if explicit is not None: + if not isinstance(explicit, str) or explicit.upper() not in SUPPORTED_MODES: + raise MuseProtocolError("unsupported Muse Voice mode") + normalized_mode: Final = explicit.upper() + if normalized_mode == "PUSH_TO_TALK": + return "PUSH_TO_TALK" + if normalized_mode == "DIARIZATION": + return "DIARIZATION" + return "ENDPOINTING" + turn_detection_present: Final = "turn_detection" in session or "turn_detection" in audio_input + turn_detection: Final = session.get("turn_detection", audio_input.get("turn_detection")) + if turn_detection_present and turn_detection is None: + return "PUSH_TO_TALK" + if turn_detection is None: + return "ENDPOINTING" + turn_detection_mapping: Final = _mapping(turn_detection, "turn_detection") + if turn_detection_mapping.get("type") not in (None, "server_vad"): + raise MuseProtocolError("Muse Voice supports server_vad turn detection or null") + return "ENDPOINTING" + + +def parse_session_update(payload: str, expected_model: str) -> MuseSessionConfig: + message: Final = _json_object(payload) + if message.get("type") not in ("session.update", "transcription_session.update"): + raise MuseProtocolError("expected session.update") + session: Final = _mapping(message.get("session"), "session") + if not session: + raise MuseProtocolError("session.update requires a session object") + session_type: Final = session.get("type") + if session_type not in (None, "transcription", "realtime"): + raise MuseProtocolError("Muse Voice supports transcription sessions only") + audio: Final = _mapping(session.get("audio"), "session.audio") + audio_input: Final = _mapping(audio.get("input"), "session.audio.input") + beta_transcription: Final = session.get("input_audio_transcription") + ga_transcription: Final = audio_input.get("transcription") + if beta_transcription is not None and ga_transcription is not None: + raise MuseProtocolError("input transcription must use either beta or GA layout") + transcription: Final = _mapping( + beta_transcription if beta_transcription is not None else ga_transcription, + "input audio transcription", + ) + requested_model: Final = _string(transcription.get("model"), "transcription model") + normalized_model: Final = _normalize_model(expected_model) + if normalized_model != MUSE_MODEL: + raise MuseProtocolError("unsupported Meta realtime model") + if requested_model is not None and _normalize_model(requested_model) != normalized_model: + raise MuseProtocolError("realtime session model cannot be changed") + language_value: Final = _string(transcription.get("language"), "language") + explicit_bias: Final = _normalize_language_sequence(transcription.get("language_bias")) + language_bias: Final = tuple( + dict.fromkeys((normalize_language(language_value), *explicit_bias)) + if language_value is not None + else explicit_bias + ) + keywords: Final = _normalize_string_sequence(transcription.get("keywords"), "keywords") + return MuseSessionConfig( + model=normalized_model, + mode=_parse_mode(session, audio_input), + sample_rate=_parse_sample_rate(session), + keywords=keywords, + language_bias=language_bias, + ) + + +def session_created_event(model: str, session_id: str) -> OpenAIEvent: + normalized_model: Final = _normalize_model(model) + default_config: Final = MuseSessionConfig( + model=normalized_model, + mode="ENDPOINTING", + sample_rate=24_000, + keywords=(), + language_bias=(), + ) + return { # mutable-ok: OpenAI-compatible JSON event + "type": "session.created", + "event_id": f"event_{uuid.uuid4().hex}", + "session": default_config.openai_session(session_id), + } + + +def session_updated_event(config: MuseSessionConfig, session_id: str) -> OpenAIEvent: + return { # mutable-ok: OpenAI-compatible JSON event + "type": "session.updated", + "event_id": f"event_{uuid.uuid4().hex}", + "session": config.openai_session(session_id), + } + + +def error_event(error_type: str, code: str, message: str) -> OpenAIEvent: + return { # mutable-ok: OpenAI-compatible JSON event + "type": "error", + "event_id": f"event_{uuid.uuid4().hex}", + "error": { # mutable-ok: nested OpenAI-compatible error object + "type": error_type, + "code": code, + "message": message, + }, + } + + +class MuseEventTransformer: + def __init__(self, *, completed_turn_limit: int = 128) -> None: + self._turns: OrderedDict[str, _TurnState] = OrderedDict() # mutable-ok: ordered active-turn state + self._active_turn_id: str | None = None + self._mode: Literal["PUSH_TO_TALK", "ENDPOINTING", "DIARIZATION"] = "ENDPOINTING" + self._completed_turn_ids: set[str] = set() # mutable-ok: bounded completed-turn membership + self._completed_turn_order: deque[str] = deque( # mutable-ok: bounded completion eviction order + maxlen=completed_turn_limit + ) + self._completed_turn_limit: Final = completed_turn_limit + self._pending_item_ids: deque[str] = deque() # mutable-ok: FIFO commit correlation state + self._last_committed_item_id: str | None = None + self._last_audio_processed_ms: float = 0.0 + self._unassigned_usage_seconds: float = 0.0 + + def configure(self, config: MuseSessionConfig) -> None: + self._mode = config.mode + + def transform(self, payload: str) -> tuple[OpenAIEvent, ...]: + message: Final = _json_object(payload) + event_type: Final = message.get("type") + if event_type == "error": + return (error_event("server_error", "provider_error", "Meta Muse realtime transcription failed"),) + if event_type == "audioProgress": + self._update_audio_progress(message) + return () + if event_type == "speechStart": + self._speech_start(message) + elif event_type == "transcript": + self._transcript(message) + elif event_type == "speaker": + self._speaker(message) + elif event_type == "speechEnd": + self._speech_end(message) + elif event_type == "speechComplete": + self._speech_complete(message) + else: + return () + return self._drain() + + def commit_item(self) -> tuple[str | None, str]: + previous_item_id: Final = self._last_committed_item_id + provider_turn_id: Final = self._active_turn_id + active_turn: Final = self._turns.get(provider_turn_id) if provider_turn_id is not None else None + item_id: Final = ( + active_turn.item_id or provider_turn_id + if active_turn is not None and provider_turn_id is not None + else f"item_{uuid.uuid4().hex}" + ) + if active_turn is not None: + active_turn.item_id = item_id + else: + self._pending_item_ids.append(item_id) + self._last_committed_item_id = item_id + return previous_item_id, item_id + + def take_unbilled_usage(self) -> RealtimeInputAudioTranscriptionUsage | None: + seconds: Final = self._unassigned_usage_seconds + if seconds <= 0: + return None + self._unassigned_usage_seconds = 0.0 + return {"type": "duration", "seconds": seconds} # mutable-ok: typed usage wire payload + + def _turn(self, turn_id: str) -> _TurnState: + if turn_id in self._completed_turn_ids: + raise _CompletedTurn + turn: Final = self._turns.get(turn_id) + if turn is not None: + return turn + created: Final = _TurnState(item_id=self._pending_item_ids.popleft() if self._pending_item_ids else turn_id) + self._turns[turn_id] = created + return created + + def _speech_start(self, message: Mapping[str, JsonValue]) -> None: + turn_id: Final = self._required_turn_id(message, "speechStart") + try: + turn: Final = self._turn(turn_id) + except _CompletedTurn: + return + turn.started = True + self._active_turn_id = turn_id + + def _transcript(self, message: Mapping[str, JsonValue]) -> None: + transcript: Final = message.get("transcript") + if not isinstance(transcript, str): + raise MuseProtocolError("transcript event has invalid transcript") + if not transcript and message.get("turnId") is None and self._active_turn_id is None: + return + turn_id: Final = self._transcript_turn_id(message) + try: + turn: Final = self._turn(turn_id) + except _CompletedTurn: + return + final: Final = message.get("final") is True + if final: + turn.final_text = transcript + turn.completed_signal = True + if self._mode == "PUSH_TO_TALK": + turn.stopped = True + if self._active_turn_id == turn_id: + self._active_turn_id = None + return + if turn.final_text is None: + turn.latest_partial = transcript + + def _speaker(self, message: Mapping[str, JsonValue]) -> None: + turn_id: Final = ( + self._required_turn_id(message, "speaker") if message.get("turnId") is not None else self._active_turn_id + ) + if turn_id is None: + raise MuseProtocolError("speaker event arrived outside an active turn") + label: Final = message.get("label") + if not isinstance(label, str) or not label.strip(): + raise MuseProtocolError("speaker event has invalid label") + try: + turn: Final = self._turn(turn_id) + except _CompletedTurn: + return + turn.speaker = label.strip() + + def _speech_end(self, message: Mapping[str, JsonValue]) -> None: + turn_id: Final = self._required_turn_id(message, "speechEnd") + try: + turn: Final = self._turn(turn_id) + except _CompletedTurn: + return + turn.stopped = True + if self._active_turn_id == turn_id: + self._active_turn_id = None + + def _speech_complete(self, message: Mapping[str, JsonValue]) -> None: + turn_id: Final = self._required_turn_id(message, "speechComplete") + transcript: Final = message.get("transcript") + if not isinstance(transcript, str): + raise MuseProtocolError("speechComplete event has invalid transcript") + try: + turn: Final = self._turn(turn_id) + except _CompletedTurn: + return + turn.final_text = transcript + turn.completed_signal = True + + def _update_audio_progress(self, message: Mapping[str, JsonValue]) -> None: + processed_ms: Final = message.get("audioProcessedMs") + if ( + isinstance(processed_ms, bool) + or not isinstance(processed_ms, (int, float)) + or not math.isfinite(processed_ms) + or processed_ms < 0 + ): + raise MuseProtocolError("audioProgress event has invalid audioProcessedMs") + if processed_ms <= self._last_audio_processed_ms: + return + self._unassigned_usage_seconds += (float(processed_ms) - self._last_audio_processed_ms) / 1000 + self._last_audio_processed_ms = float(processed_ms) + + def _drain(self) -> tuple[OpenAIEvent, ...]: + events: list[OpenAIEvent] = [] # mutable-ok: ordered events are frozen to a tuple before return + while self._turns: + turn_id: str = next(iter(self._turns)) # rebind-ok: selects the next ordered turn + turn: _TurnState = self._turns[turn_id] # rebind-ok: state for the selected turn + has_content: bool = ( # rebind-ok: evaluated for the selected turn + turn.latest_partial is not None or turn.final_text is not None + ) + item_id: str = turn.item_id or turn_id # rebind-ok: selected for each ordered turn + if (turn.started or has_content) and not turn.start_emitted: + turn.start_emitted = True + events.append(self._speech_event("input_audio_buffer.speech_started", item_id)) + if turn.latest_partial is not None and turn.final_text is None: + delta: str = self._new_suffix( # rebind-ok: computed for the selected turn + turn.emitted_partial, turn.latest_partial + ) + if delta: + turn.emitted_partial = turn.latest_partial + events.append( + { # mutable-ok: OpenAI-compatible JSON event + "type": "conversation.item.input_audio_transcription.delta", + "event_id": f"event_{uuid.uuid4().hex}", + "item_id": item_id, + "content_index": 0, + "delta": delta, + } + ) + if turn.stopped and not turn.stopped_emitted: + turn.stopped_emitted = True + events.append(self._speech_event("input_audio_buffer.speech_stopped", item_id)) + if turn.final_text is not None and turn.stopped_emitted and not turn.completed_emitted: + turn.completed_emitted = True + usage: RealtimeInputAudioTranscriptionUsage | None = ( # rebind-ok: usage assigned per turn + self.take_unbilled_usage() + ) + completed_event: dict[str, object] = { # mutable-ok: incrementally builds OpenAI JSON event + "type": "conversation.item.input_audio_transcription.completed", + "event_id": f"event_{uuid.uuid4().hex}", + "item_id": item_id, + "content_index": 0, + "transcript": turn.final_text, + } + if turn.speaker is not None: + completed_event["speaker"] = turn.speaker + if usage is not None: + completed_event["usage"] = usage + events.append(completed_event) + if not (turn.completed_emitted and (turn.stopped or turn.completed_signal)): + break + del self._turns[turn_id] + self._remember_completed(turn_id) + return tuple(events) + + def _remember_completed(self, turn_id: str) -> None: + if turn_id in self._completed_turn_ids: + return + if len(self._completed_turn_order) >= self._completed_turn_limit: + self._completed_turn_ids.discard(self._completed_turn_order.popleft()) + self._completed_turn_order.append(turn_id) + self._completed_turn_ids.add(turn_id) + + def _transcript_turn_id(self, message: Mapping[str, JsonValue]) -> str: + if message.get("turnId") is not None: + return self._required_turn_id(message, "transcript") + if self._active_turn_id is not None: + return self._active_turn_id + if self._mode != "PUSH_TO_TALK": + raise MuseProtocolError("transcript event is missing turnId outside an active turn") + turn_id: Final = f"item_{uuid.uuid4().hex}" + self._active_turn_id = turn_id + return turn_id + + @staticmethod + def _required_turn_id(message: Mapping[str, JsonValue], event: str) -> str: + value: Final = message.get("turnId") + if isinstance(value, bool) or not isinstance(value, (str, int)): + raise MuseProtocolError(f"{event} event has invalid turnId") + turn_id: Final = str(value).strip() + if not turn_id: + raise MuseProtocolError(f"{event} event has invalid turnId") + return turn_id + + @staticmethod + def _speech_event(event_type: str, turn_id: str) -> OpenAIEvent: + return { # mutable-ok: OpenAI-compatible JSON event + "type": event_type, + "event_id": f"event_{uuid.uuid4().hex}", + "item_id": turn_id, + } + + @staticmethod + def _new_suffix(previous: str, current: str) -> str: + if current.startswith(previous): + return current[len(previous) :] + return "" + + +class _CompletedTurn(Exception): + pass + + +def encode_event(event: Mapping[str, object]) -> str: + return json.dumps(event, separators=(",", ":")) diff --git a/litellm/llms/openai_like/providers.json b/litellm/llms/openai_like/providers.json index a458a209ea9..fe10293c420 100644 --- a/litellm/llms/openai_like/providers.json +++ b/litellm/llms/openai_like/providers.json @@ -173,7 +173,7 @@ "api_key_env": "META_API_KEY", "api_base_env": "META_API_BASE", "base_class": "openai_gpt", - "supported_endpoints": ["/v1/chat/completions", "/v1/responses", "/v1/messages"] + "supported_endpoints": ["/v1/chat/completions", "/v1/responses", "/v1/messages", "/v1/realtime"] }, "cognition": { "base_url": "https://api.cognition.ai/v1", diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 7fa09951eae..4dc6768e12a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -34717,6 +34717,21 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "meta/muse-voice-transcribe-1.0": { + "litellm_provider": "meta", + "mode": "audio_transcription", + "source": "https://dev.meta.ai/docs/speech-to-text", + "supported_endpoints": [ + "/v1/realtime" + ], + "supported_modalities": [ + "audio" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true + }, "meta_llama/Llama-3.3-70B-Instruct": { "litellm_provider": "meta_llama", "max_input_tokens": 128000, @@ -59098,9 +59113,9 @@ "litellm_provider": "wandb", "mode": "chat", "supports_reasoning": true, - "input_cost_per_token": 0.00000131, - "output_cost_per_token": 0.00000396, - "cache_read_input_token_cost": 0.000000044, + "input_cost_per_token": 1.31e-06, + "output_cost_per_token": 3.96e-06, + "cache_read_input_token_cost": 4.4e-08, "supports_prompt_caching": true, "source": "https://wandb.ai/site/pricing/tokens/" }, @@ -59108,9 +59123,9 @@ "litellm_provider": "wandb", "mode": "chat", "supports_reasoning": true, - "input_cost_per_token": 0.0000001, - "output_cost_per_token": 0.00000015, - "cache_read_input_token_cost": 0.00000005, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 1.5e-07, + "cache_read_input_token_cost": 5e-08, "supports_prompt_caching": true, "source": "https://wandb.ai/site/pricing/tokens/" }, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 44c47af57f4..b0803e44f6b 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -333,7 +333,7 @@ async def _resolve_vertex_access_token_bounded( @wrapper_client -async def _arealtime( +async def _arealtime( # noqa: C901 # central dispatcher branches once per supported realtime provider model: str, websocket: "WebSocket", # fastapi websocket api_base: str | None = None, @@ -391,7 +391,37 @@ async def _arealtime( model=model, provider=LlmProviders(_custom_llm_provider), ) - if provider_config is not None: + if _custom_llm_provider == LlmProviders.META.value: + if model != "muse-voice-transcribe-1.0": + raise ValueError(f"Unsupported Meta realtime model: {model}") + if query_params is None or query_params.get("intent") != "transcription": + raise ValueError("Meta Muse Voice realtime requires intent=transcription") + + from litellm.llms.meta.realtime.handler import MetaRealtime + + meta_api_key: Final = get_secret_str("META_API_KEY") + dynamic_key_override: Final = dynamic_api_key if dynamic_api_key != meta_api_key else None + resolved_meta_api_key: Final = ( + api_key + or litellm_params.api_key + or dynamic_key_override + or get_secret_str("MODEL_API_KEY") + or dynamic_api_key + or meta_api_key + ) + await MetaRealtime().async_realtime( + model=model, + websocket=websocket, + logging_obj=litellm_logging_obj, + api_base=dynamic_api_base or litellm_params.api_base or api_base, + api_key=resolved_meta_api_key, + client=client, + timeout=timeout, + query_params=query_params, + user_api_key_dict=kwargs.get("user_api_key_dict"), + litellm_metadata=_build_litellm_metadata(kwargs), + ) + elif provider_config is not None: await base_llm_http_handler.async_realtime( model=model, websocket=websocket, diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index b7c4371f32f..02a102c9579 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -2202,6 +2202,8 @@ class OpenAIRealtimeInputAudioTranscriptionCompleted(TypedDict): item_id: ReadOnly[str] content_index: ReadOnly[int] transcript: ReadOnly[str] + usage: NotRequired[ReadOnly[Mapping[str, object]]] + speaker: NotRequired[ReadOnly[str]] class OpenAIRealtimeUsageTokenDetails(TypedDict): diff --git a/litellm/types/realtime.py b/litellm/types/realtime.py index 17dc70126f3..30db794c96e 100644 --- a/litellm/types/realtime.py +++ b/litellm/types/realtime.py @@ -169,9 +169,19 @@ class RealtimeInputAudioTranscriptionUsageInputTokenDetails(TypedDict): audio_tokens: ReadOnly[int] -class RealtimeInputAudioTranscriptionUsage(TypedDict): +class RealtimeInputAudioTranscriptionTokenUsage(TypedDict): type: ReadOnly[Literal["tokens"]] input_tokens: ReadOnly[int] output_tokens: ReadOnly[int] total_tokens: ReadOnly[int] input_token_details: ReadOnly[RealtimeInputAudioTranscriptionUsageInputTokenDetails] + + +class RealtimeInputAudioTranscriptionDurationUsage(TypedDict): + type: ReadOnly[Literal["duration"]] + seconds: ReadOnly[float] + + +RealtimeInputAudioTranscriptionUsage = ( + RealtimeInputAudioTranscriptionTokenUsage | RealtimeInputAudioTranscriptionDurationUsage +) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 7fa09951eae..4dc6768e12a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -34717,6 +34717,21 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "meta/muse-voice-transcribe-1.0": { + "litellm_provider": "meta", + "mode": "audio_transcription", + "source": "https://dev.meta.ai/docs/speech-to-text", + "supported_endpoints": [ + "/v1/realtime" + ], + "supported_modalities": [ + "audio" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true + }, "meta_llama/Llama-3.3-70B-Instruct": { "litellm_provider": "meta_llama", "max_input_tokens": 128000, @@ -59098,9 +59113,9 @@ "litellm_provider": "wandb", "mode": "chat", "supports_reasoning": true, - "input_cost_per_token": 0.00000131, - "output_cost_per_token": 0.00000396, - "cache_read_input_token_cost": 0.000000044, + "input_cost_per_token": 1.31e-06, + "output_cost_per_token": 3.96e-06, + "cache_read_input_token_cost": 4.4e-08, "supports_prompt_caching": true, "source": "https://wandb.ai/site/pricing/tokens/" }, @@ -59108,9 +59123,9 @@ "litellm_provider": "wandb", "mode": "chat", "supports_reasoning": true, - "input_cost_per_token": 0.0000001, - "output_cost_per_token": 0.00000015, - "cache_read_input_token_cost": 0.00000005, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 1.5e-07, + "cache_read_input_token_cost": 5e-08, "supports_prompt_caching": true, "source": "https://wandb.ai/site/pricing/tokens/" }, diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 9c0f6f59463..e330cb103b1 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -10,8 +10,6 @@ from websockets.exceptions import ConnectionClosed from websockets.frames import Close import litellm - - from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.realtime_streaming import ( REALTIME_SESSION_SUCCESS_LOGGED_KEY, @@ -20,10 +18,6 @@ from litellm.litellm_core_utils.realtime_streaming import ( ) from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.llms.openai import ( - OpenAIRealtimeStreamResponseBaseObject, - OpenAIRealtimeStreamSessionEvents, -) def _make_transcript_event(text: str, item_id: str = "item_x") -> bytes: @@ -161,6 +155,7 @@ async def test_backend_to_client_send_text_receives_str_not_bytes(): logging_obj = MagicMock() logging_obj.async_success_handler = AsyncMock() logging_obj.success_handler = MagicMock() + logging_obj.dispatch_success_handlers = AsyncMock() streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) await streaming.backend_to_client_send_messages() @@ -812,7 +807,6 @@ async def test_transcription_captured_in_backend_to_client(): Test that conversation.item.input_audio_transcription.completed events from the backend are captured as user input during the WebSocket session. """ - import litellm client_ws = MagicMock() client_ws.send_text = AsyncMock() @@ -838,6 +832,7 @@ async def test_transcription_captured_in_backend_to_client(): logging_obj.model_call_details = {"messages": "default-message-value"} logging_obj.async_success_handler = AsyncMock() logging_obj.success_handler = MagicMock() + logging_obj.dispatch_success_handlers = AsyncMock() streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) await streaming.backend_to_client_send_messages() @@ -883,6 +878,7 @@ async def test_transcription_session_captures_usage_and_skips_response_create(): logging_obj.model_call_details = {} logging_obj.async_success_handler = AsyncMock() logging_obj.success_handler = MagicMock() + logging_obj.dispatch_success_handlers = AsyncMock() streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) await streaming.backend_to_client_send_messages() @@ -1100,7 +1096,6 @@ def test_capture_transcription_usage_deduplicates_when_already_stored(): When the event is already in messages (logged via store_message), it must not be appended a second time by _capture_transcription_usage. """ - import litellm streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) # Add the event type to the default logged list so _should_store_message returns True. @@ -1409,7 +1404,6 @@ async def test_realtime_guardrail_blocks_prompt_injection(monkeypatch: pytest.Mo ) - @pytest.mark.asyncio async def test_realtime_guardrail_allows_clean_transcript(monkeypatch: pytest.MonkeyPatch): """ @@ -1466,7 +1460,6 @@ async def test_realtime_guardrail_allows_clean_transcript(monkeypatch: pytest.Mo assert len(response_creates) == 1, f"Clean transcript should trigger response.create, got: {sent_to_backend}" - @pytest.mark.asyncio async def test_realtime_text_input_guardrail_blocks_and_returns_error(monkeypatch: pytest.MonkeyPatch): """ @@ -1560,7 +1553,6 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error(monkeypatc assert len(original_items) == 0, f"Blocked item should not be forwarded to backend, got: {original_items}" - @pytest.mark.asyncio async def test_realtime_function_call_output_guardrail_blocks_and_returns_error(monkeypatch: pytest.MonkeyPatch): """ @@ -1649,7 +1641,6 @@ async def test_realtime_function_call_output_guardrail_blocks_and_returns_error( assert "test@example.com" not in sanitized_item["output"] - @pytest.mark.asyncio async def test_realtime_function_call_output_guardrail_allows_clean_output(monkeypatch: pytest.MonkeyPatch): """ @@ -1714,7 +1705,6 @@ async def test_realtime_function_call_output_guardrail_allows_clean_output(monke assert len(forwarded) == 1, f"Clean function_call_output should be forwarded, got: {forwarded}" - @pytest.mark.asyncio async def test_realtime_text_input_guardrail_uses_pre_call_mode(monkeypatch: pytest.MonkeyPatch): """ @@ -1750,7 +1740,6 @@ async def test_realtime_text_input_guardrail_uses_pre_call_mode(monkeypatch: pyt ) - @pytest.mark.asyncio async def test_realtime_session_created_injects_session_update_for_audio_guardrail(monkeypatch: pytest.MonkeyPatch): """ @@ -1807,7 +1796,6 @@ async def test_realtime_session_created_injects_session_update_for_audio_guardra ) - @pytest.mark.asyncio async def test_realtime_session_created_does_not_inject_session_update_for_pre_call_only( monkeypatch: pytest.MonkeyPatch, @@ -1852,7 +1840,6 @@ async def test_realtime_session_created_does_not_inject_session_update_for_pre_c assert len(session_updates) == 0, f"pre_call-only guardrail must not inject session.update, got: {sent_to_backend}" - @pytest.mark.asyncio async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad(monkeypatch: pytest.MonkeyPatch): """Model Armor-style pre_call + post_call must not gate audio VAD.""" @@ -1868,17 +1855,17 @@ async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad(monke litellm, "callbacks", [ - ModelArmorStyleGuardrail( - guardrail_name="model_armor_all_pre_call", - event_hook=GuardrailEventHooks.pre_call, - default_on=False, - ), - ModelArmorStyleGuardrail( - guardrail_name="model_armor_all_post_call", - event_hook=GuardrailEventHooks.post_call, - default_on=False, - ), - ], + ModelArmorStyleGuardrail( + guardrail_name="model_armor_all_pre_call", + event_hook=GuardrailEventHooks.pre_call, + default_on=False, + ), + ModelArmorStyleGuardrail( + guardrail_name="model_armor_all_post_call", + event_hook=GuardrailEventHooks.post_call, + default_on=False, + ), + ], ) client_ws = MagicMock() @@ -1902,7 +1889,6 @@ async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad(monke assert streaming._has_audio_transcription_guardrails() is False - @pytest.mark.asyncio async def test_end_session_after_n_fails_closes_connection(monkeypatch: pytest.MonkeyPatch): """ @@ -1949,7 +1935,6 @@ async def test_end_session_after_n_fails_closes_connection(monkeypatch: pytest.M assert streaming._violation_count == 2 - @pytest.mark.asyncio async def test_on_violation_end_session_closes_on_first_fail(monkeypatch: pytest.MonkeyPatch): """ @@ -1995,7 +1980,6 @@ async def test_on_violation_end_session_closes_on_first_fail(monkeypatch: pytest assert streaming._violation_count == 1 - @pytest.mark.asyncio async def test_provider_path_suppresses_duplicate_session_created_after_synthetic(): client_ws = MagicMock() @@ -2956,7 +2940,9 @@ async def test_log_messages_routes_async_logging_through_bounded_worker(): mock_worker.ensure_initialized_and_enqueue.assert_called_once() enqueued = mock_worker.ensure_initialized_and_enqueue.call_args - assert (enqueued.args or tuple(enqueued.kwargs.values()))[0] is logging_obj.dispatch_success_handlers.return_value + assert (enqueued.args or tuple(enqueued.kwargs.values()))[ + 0 + ] is logging_obj.dispatch_success_handlers.return_value logging_obj.dispatch_success_handlers.assert_called_once_with(streaming.messages, prefer_async_handlers=True) logging_obj.success_handler.assert_not_called() # the bare create_task path must no longer be used for success logging @@ -3041,6 +3027,7 @@ async def test_session_close_flushes_unbilled_transcription_usage(): logging_obj: Final = MagicMock() logging_obj.async_success_handler = AsyncMock() logging_obj.success_handler = MagicMock() + logging_obj.dispatch_success_handlers = AsyncMock() usage: Final[RealtimeInputAudioTranscriptionUsage] = { "type": "tokens", @@ -3116,6 +3103,7 @@ async def test_session_close_flush_noop_without_unbilled_usage(): logging_obj: Final = MagicMock() logging_obj.async_success_handler = AsyncMock() logging_obj.success_handler = MagicMock() + logging_obj.dispatch_success_handlers = AsyncMock() provider_config: Final = MagicMock() provider_config.unbilled_usage_on_session_close = MagicMock(return_value=None) @@ -3412,3 +3400,144 @@ async def test_refused_session_does_not_stamp_the_reservation_ownership_marker() assert session.logging.logged_failures == (upstream_close,) assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in session.logging.model_call_details + + +@pytest.mark.asyncio +async def test_separate_usage_provider_flushes_duration_once_without_client_event(): + from typing import Final + + client_ws: Final = MagicMock() + client_ws.send_text = AsyncMock() + backend_ws: Final = MagicMock() + backend_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None)) + logging_obj: Final = MagicMock() + logging_obj.model_call_details = {} + logging_obj.dispatch_success_handlers = AsyncMock() + usage_provider: Final = MagicMock() + usage_provider.unbilled_usage_on_session_close.return_value = { + "type": "duration", + "seconds": 0.75, + } + + streaming: Final = RealTimeStreaming( + client_ws, + backend_ws, + logging_obj, + model="muse-voice-transcribe-1.0", + usage_provider=usage_provider, + ) + + await streaming.backend_to_client_send_messages() + + usage_provider.unbilled_usage_on_session_close.assert_called_once_with("muse-voice-transcribe-1.0") + duration_events: Final = tuple( + message + for message in streaming.messages + if isinstance(message, dict) and message.get("usage") == {"type": "duration", "seconds": 0.75} + ) + assert len(duration_events) == 1 + assert client_ws.send_text.await_count == 0 + logging_obj.dispatch_success_handlers.assert_called_once_with(streaming.messages, prefer_async_handlers=True) + + +@pytest.mark.asyncio +async def test_transformed_transcription_completion_never_sends_response_create(): + from typing import Final + + completed_event: Final = { + "type": "conversation.item.input_audio_transcription.completed", + "event_id": "event_1", + "item_id": "turn_1", + "content_index": 0, + "transcript": "private transcript", + "usage": {"type": "duration", "seconds": 0.5}, + } + provider_config: Final = MagicMock() + provider_config.requires_session_configuration.return_value = True + provider_config.transform_realtime_response.return_value = { + "response": completed_event, + "current_output_item_id": None, + "current_response_id": None, + "current_delta_chunks": None, + "current_conversation_id": None, + "current_item_chunks": None, + "current_delta_type": None, + "session_configuration_request": None, + } + provider_config.transform_realtime_request.return_value = (json.dumps({"type": "response.create"}),) + provider_config.is_setup_message.return_value = False + provider_config.is_content_message.return_value = False + client_ws: Final = MagicMock() + client_ws.send_text = AsyncMock() + backend_ws: Final = MagicMock() + backend_ws.send = AsyncMock() + + streaming: Final = RealTimeStreaming( + client_ws, + backend_ws, + MagicMock(), + provider_config=provider_config, + model="muse-voice-transcribe-1.0", + force_transcription_model="muse-voice-transcribe-1.0", + ) + + await streaming._handle_provider_config_message("{}") + + assert json.loads(client_ws.send_text.await_args.args[0]) == completed_event + backend_ws.send.assert_not_awaited() + + +def test_private_logging_excludes_audio_transcript_hints_and_provider_body(monkeypatch: pytest.MonkeyPatch): + from typing import Final + + monkeypatch.setattr(litellm, "logged_real_time_event_types", "*") + logging_obj: Final = MagicMock() + logging_obj.model_call_details = {} + streaming: Final = RealTimeStreaming( + MagicMock(), + MagicMock(), + logging_obj, + model="muse-voice-transcribe-1.0", + exclude_private_content_from_logs=True, + ) + audio: Final = "cHJpdmF0ZS1hdWRpbw==" + transcript: Final = "private transcript" + keyword: Final = "private keyword" + provider_body: Final = "private provider body" + + streaming.store_input( + json.dumps( + { + "type": "session.update", + "session": { + "type": "transcription", + "model": "muse-voice-transcribe-1.0", + "mode": "ENDPOINTING", + "audio": {"input": {"transcription": {"keywords": [keyword]}}}, + }, + } + ) + ) + streaming.store_input(json.dumps({"type": "input_audio_buffer.append", "audio": audio})) + streaming.store_message( + { + "type": "conversation.item.input_audio_transcription.completed", + "event_id": "event_1", + "item_id": "turn_1", + "transcript": transcript, + "provider_body": provider_body, + "usage": {"type": "duration", "seconds": 1.0}, + } + ) + + logged_inputs: Final = tuple(call.kwargs["input"] for call in logging_obj.pre_call.call_args_list) + serialized: Final = json.dumps({"inputs": logged_inputs, "messages": streaming.messages}) + assert audio not in serialized + assert transcript not in serialized + assert keyword not in serialized + assert provider_body not in serialized + assert "muse-voice-transcribe-1.0" in serialized + assert "ENDPOINTING" in serialized + assert "turn_1" in serialized + assert '"seconds": 1.0' in serialized + assert streaming.input_messages == [] diff --git a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py new file mode 100644 index 00000000000..3d31b3a5534 --- /dev/null +++ b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py @@ -0,0 +1,449 @@ +import asyncio +import base64 +import json +from collections.abc import Callable +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.llms.meta.realtime.handler import ( + DEFAULT_MUSE_REALTIME_URL, + MetaRealtime, + MuseAdapterError, + MuseRealtimeAdapter, + build_muse_realtime_url, + normalize_access_token, + safe_close_reason, + sanitize_close_code, +) +from litellm.llms.meta.realtime.transformation import MUSE_MODEL + + +class FakeProviderWebSocket: + def __init__(self, session_id: str = "provider-session") -> None: + self.sent: list[str | bytes] = [] + self.close_calls: list[tuple[int, str]] = [] + self._session_id: Final = session_id + self._recv_count = 0 + self._closed = asyncio.Event() + + async def send(self, message: str | bytes) -> None: + self.sent.append(message) + + async def recv(self, decode: bool | None = None) -> str | bytes: + self._recv_count += 1 + if self._recv_count == 1: + return json.dumps({"sessionId": self._session_id}) + await self._closed.wait() + raise MuseAdapterError("closed", close_code=1000) + + async def close(self, code: int = 1000, reason: str = "") -> None: + self.close_calls.append((code, reason)) + self._closed.set() + + +class DelayedAckWebSocket(FakeProviderWebSocket): + def __init__(self) -> None: + super().__init__() + self.ack_release = asyncio.Event() + + async def recv(self, decode: bool | None = None) -> str | bytes: + self._recv_count += 1 + if self._recv_count == 1: + await self.ack_release.wait() + return json.dumps({"sessionId": self._session_id}) + await self._closed.wait() + raise MuseAdapterError("closed", close_code=1000) + + +async def _wait_until(predicate: Callable[[], bool]) -> None: + for _ in range(100): + if predicate(): + return + await asyncio.sleep(0) + raise AssertionError("condition did not become true") + + +def _session_update(*, rate: int = 24_000, mode: str = "ENDPOINTING") -> str: + return json.dumps( + { + "type": "session.update", + "session": { + "type": "transcription", + "mode": mode, + "audio": { + "input": { + "format": {"type": "audio/pcm", "rate": rate, "channels": 1}, + "transcription": {"model": MUSE_MODEL}, + } + }, + }, + } + ) + + +async def _configured_adapter( + *, + rate: int = 24_000, + mode: str = "ENDPOINTING", + provider_ws: FakeProviderWebSocket | None = None, + monotonic: Callable[[], float] = lambda: 10.0, + sleep: Callable[[float], object] | None = None, +) -> tuple[MuseRealtimeAdapter, FakeProviderWebSocket, dict[str, object]]: + ws: Final = provider_ws or FakeProviderWebSocket() + connect_call: Final[dict[str, object]] = {} + + async def connect(url: str, **kwargs: object) -> FakeProviderWebSocket: + connect_call.update({"url": url, **kwargs}) + return ws + + async def no_sleep(_: float) -> None: + return None + + adapter: Final = MuseRealtimeAdapter( + model=f"meta/{MUSE_MODEL}", + api_key=" raw-token ", + websocket_connect=connect, + monotonic=monotonic, + sleep=sleep or no_sleep, + ) + created: Final = json.loads(await adapter.recv()) + assert created["type"] == "session.created" + await adapter.send(_session_update(rate=rate, mode=mode)) + updated: Final = json.loads(await adapter.recv()) + assert updated["type"] == "session.updated" + return adapter, ws, connect_call + + +@pytest.mark.parametrize( + ("api_key", "expected"), + [ + ("token", "Bearer token"), + (" Bearer token ", "Bearer token"), + ("bearer token", "Bearer token"), + ], +) +def test_normalize_access_token_emits_exactly_one_bearer_prefix(api_key: str, expected: str): + assert normalize_access_token(api_key) == expected + + +@pytest.mark.parametrize("api_key", ["", " ", "Bearer", " bearer "]) +def test_normalize_access_token_rejects_missing_token(api_key: str): + with pytest.raises(ValueError, match=r"token|key is required"): + normalize_access_token(api_key) + + +def test_build_muse_realtime_url_uses_fixed_secure_path(): + assert build_muse_realtime_url(None) == DEFAULT_MUSE_REALTIME_URL + assert build_muse_realtime_url("https://example.test/custom/path?ignored=yes") == ( + "wss://example.test/v1/asr/realtime" + ) + assert build_muse_realtime_url("wss://example.test:8443/other") == ("wss://example.test:8443/v1/asr/realtime") + + +@pytest.mark.parametrize( + "api_base", + [ + "http://example.test", + "ws://example.test", + "wss://user:pass@example.test", + "wss://example.test/path#fragment", + "not-a-url", + ], +) +def test_build_muse_realtime_url_rejects_insecure_or_ambiguous_overrides(api_base: str): + with pytest.raises(ValueError, match="absolute wss:// or https://"): + build_muse_realtime_url(api_base) + + +@pytest.mark.asyncio +async def test_handshake_contains_bearer_only_in_json_body_and_waits_for_ack(): + provider_ws: Final = DelayedAckWebSocket() + connect_call: Final[dict[str, object]] = {} + + async def connect(url: str, **kwargs: object) -> DelayedAckWebSocket: + connect_call.update({"url": url, **kwargs}) + return provider_ws + + adapter: Final = MuseRealtimeAdapter( + model=MUSE_MODEL, + api_key="Bearer private-token", + websocket_connect=connect, + ) + await adapter.recv() + update_task: Final = asyncio.create_task(adapter.send(_session_update())) + await _wait_until(lambda: len(provider_ws.sent) == 1) + + assert connect_call["url"] == DEFAULT_MUSE_REALTIME_URL + assert "additional_headers" not in connect_call + handshake: Final = json.loads(provider_ws.sent[0]) + assert handshake["authorization"] == {"accessToken": "Bearer private-token"} + assert handshake["audioEncoding"] == "PCM_24KHZ" + assert not update_task.done() + assert not any(isinstance(frame, bytes) for frame in provider_ws.sent) + + provider_ws.ack_release.set() + await update_task + updated: Final = json.loads(await adapter.recv()) + assert updated["type"] == "session.updated" + assert updated["session"]["id"] == "provider-session" + await adapter.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("rate", "packet_bytes"), [(16_000, 2_560), (24_000, 3_840)]) +async def test_audio_is_strictly_decoded_and_packetized_as_raw_pcm(rate: int, packet_bytes: int): + adapter, provider_ws, _ = await _configured_adapter(rate=rate) + pcm: Final = (b"\xff\xfe\x00\x80" * (packet_bytes // 2))[: packet_bytes * 2] + + await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(pcm).decode()})) + await _wait_until(lambda: sum(isinstance(frame, bytes) for frame in provider_ws.sent) == 2) + + binary_frames: Final = tuple(frame for frame in provider_ws.sent if isinstance(frame, bytes)) + assert binary_frames == (pcm[:packet_bytes], pcm[packet_bytes:]) + assert b"\xff\xfe" in pcm + await adapter.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("audio", "expected_message"), + [ + ("not base64!", "valid base64"), + (base64.b64encode(b"\x00").decode(), "complete samples"), + ], +) +async def test_invalid_base64_or_odd_pcm_is_rejected(audio: str, expected_message: str): + adapter, provider_ws, _ = await _configured_adapter() + + await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": audio})) + error: Final = json.loads(await adapter.recv()) + + assert error["type"] == "error" + assert error["error"]["code"] == "invalid_audio" + assert expected_message in error["error"]["message"] + assert adapter.close_code == 1008 + assert not any(isinstance(frame, bytes) for frame in provider_ws.sent) + await adapter.close() + + +@pytest.mark.asyncio +async def test_absolute_pacing_delays_only_audio_ahead_of_wall_time(): + sleeps: Final[list[float]] = [] + + async def record_sleep(delay: float) -> None: + sleeps.append(delay) + + adapter, provider_ws, _ = await _configured_adapter(monotonic=lambda: 10.0, sleep=record_sleep) + pcm: Final = b"\x01\x02" * 3_840 + + await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(pcm).decode()})) + await _wait_until(lambda: sum(isinstance(frame, bytes) for frame in provider_ws.sent) == 2) + + assert sleeps == pytest.approx([0.08]) + await adapter.close() + + +@pytest.mark.asyncio +async def test_append_larger_than_four_seconds_is_rejected_without_dropping_prefix(): + adapter, provider_ws, _ = await _configured_adapter(rate=16_000) + oversized_pcm: Final = b"\x00\x00" * (16_000 * 4 + 1) + + await adapter.send( + json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(oversized_pcm).decode()}) + ) + error: Final = json.loads(await adapter.recv()) + + assert error["error"]["code"] == "audio_backlog_exceeded" + assert adapter.close_code == 1008 + assert not any(isinstance(frame, bytes) for frame in provider_ws.sent) + await adapter.close() + + +@pytest.mark.asyncio +async def test_clear_discards_only_unsent_audio(): + adapter, provider_ws, _ = await _configured_adapter() + old_pcm: Final = b"\x01\x02" * 100 + new_pcm: Final = b"\x03\x04" * 100 + + await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(old_pcm).decode()})) + await adapter.send(_event("input_audio_buffer.clear")) + cleared: Final = json.loads(await adapter.recv()) + await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(new_pcm).decode()})) + await adapter.send(_event("input_audio_buffer.commit")) + await _wait_until(lambda: any(isinstance(frame, bytes) for frame in provider_ws.sent)) + + assert cleared["type"] == "input_audio_buffer.cleared" + assert tuple(frame for frame in provider_ws.sent if isinstance(frame, bytes)) == (new_pcm,) + assert '{"type":"endStream"}' not in provider_ws.sent + await adapter.close() + + +@pytest.mark.asyncio +async def test_endpointing_commit_flushes_partial_packet_without_ending_stream(): + adapter, provider_ws, _ = await _configured_adapter(mode="ENDPOINTING") + pcm: Final = b"\x01\x02" * 100 + + await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(pcm).decode()})) + await adapter.send(_event("input_audio_buffer.commit")) + committed: Final = json.loads(await adapter.recv()) + await _wait_until(lambda: any(isinstance(frame, bytes) for frame in provider_ws.sent)) + + assert committed["type"] == "input_audio_buffer.committed" + assert committed["item_id"].startswith("item_") + assert committed["previous_item_id"] is None + assert tuple(frame for frame in provider_ws.sent if isinstance(frame, bytes)) == (pcm,) + assert '{"type":"endStream"}' not in provider_ws.sent + await adapter.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("mode", "terminal_event"), + [("PUSH_TO_TALK", "input_audio_buffer.commit"), ("ENDPOINTING", "input_audio_buffer.end")], +) +async def test_commit_or_end_sends_end_stream_exactly_once(mode: str, terminal_event: str): + adapter, provider_ws, _ = await _configured_adapter(mode=mode) + pcm: Final = b"\x01\x02" * 100 + + await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(pcm).decode()})) + await adapter.send(_event(terminal_event)) + await adapter.send(_event("input_audio_buffer.end")) + await _wait_until(lambda: '{"type":"endStream"}' in provider_ws.sent) + + assert tuple(frame for frame in provider_ws.sent if isinstance(frame, bytes)) == (pcm,) + assert provider_ws.sent.count('{"type":"endStream"}') == 1 + await adapter.close() + + +@pytest.mark.asyncio +async def test_response_create_is_returned_as_error_and_never_sent_upstream(): + adapter, provider_ws, _ = await _configured_adapter() + + await adapter.send(_event("response.create")) + error: Final = json.loads(await adapter.recv()) + + assert error["type"] == "error" + assert error["error"]["code"] == "unsupported_event" + assert not any(isinstance(frame, str) and "response.create" in frame for frame in provider_ws.sent) + await adapter.close() + + +@pytest.mark.asyncio +async def test_close_codes_and_reasons_are_sanitized_without_secret_leakage(): + adapter, provider_ws, _ = await _configured_adapter() + secret: Final = "Bearer private-token" + + await adapter.close(code=4001, reason=f"provider rejected {secret}") + + assert adapter.close_code == 1011 + assert adapter.close_reason == "Realtime transcription service error" + assert provider_ws.close_calls == [(1011, "Realtime transcription service error")] + assert secret not in json.dumps(provider_ws.close_calls) + assert sanitize_close_code(1013) == 1013 + assert safe_close_reason(1008) == "Invalid realtime transcription request" + + +@pytest.mark.asyncio +async def test_handshake_failure_reports_only_exception_type(): + secret: Final = "private-token" + + async def failing_connect(url: str, **kwargs: object) -> FakeProviderWebSocket: + raise RuntimeError(f"failed with {secret}") + + adapter: Final = MuseRealtimeAdapter( + model=MUSE_MODEL, + api_key=secret, + websocket_connect=failing_connect, + ) + await adapter.recv() + + await adapter.send(_session_update()) + error = json.loads(await adapter.recv()) + + assert error["type"] == "error" + assert error["error"]["message"] == "Meta Muse realtime handshake failed" + assert secret not in json.dumps(error) + with pytest.raises(MuseAdapterError) as exc_info: + await adapter.recv() + assert exc_info.value.close_code == 1011 + assert secret not in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_meta_realtime_missing_credentials_closes_client_with_policy_code(): + client_ws: Final = MagicMock() + client_ws.send_text = AsyncMock() + client_ws.close = AsyncMock() + + await MetaRealtime().async_realtime( + model=MUSE_MODEL, + websocket=client_ws, + logging_obj=MagicMock(), + api_key=None, + ) + + client_ws.close.assert_awaited_once_with( + code=1008, + reason="Invalid realtime transcription request", + ) + sent_error: Final = json.loads(client_ws.send_text.await_args.args[0]) + assert sent_error["type"] == "error" + assert sent_error["error"]["code"] == "invalid_configuration" + + +@pytest.mark.asyncio +async def test_meta_realtime_invalid_constructor_input_sends_error_before_close(): + client_ws: Final = MagicMock() + client_ws.send_text = AsyncMock() + client_ws.close = AsyncMock() + + await MetaRealtime().async_realtime( + model=MUSE_MODEL, + websocket=client_ws, + logging_obj=MagicMock(), + api_key="Bearer", + ) + + sent_error: Final = json.loads(client_ws.send_text.await_args.args[0]) + assert sent_error["error"]["message"] == "Invalid Meta Muse realtime configuration" + client_ws.close.assert_awaited_once_with( + code=1008, + reason="Invalid realtime transcription request", + ) + + +@pytest.mark.asyncio +async def test_meta_realtime_enables_private_logging_usage_and_model_enforcement(monkeypatch: pytest.MonkeyPatch): + captured: Final[dict[str, object]] = {} + + class CapturingStreaming: + def __init__(self, websocket, backend_ws, logging_obj, **kwargs): + captured.update({"websocket": websocket, "backend_ws": backend_ws, "logging_obj": logging_obj, **kwargs}) + + async def bidirectional_forward(self) -> None: + return None + + client_ws: Final = MagicMock() + client_ws.send_text = AsyncMock() + client_ws.close = AsyncMock() + monkeypatch.setattr("litellm.llms.meta.realtime.handler.RealTimeStreaming", CapturingStreaming) + + await MetaRealtime().async_realtime( + model=MUSE_MODEL, + websocket=client_ws, + logging_obj=MagicMock(), + api_key="private-token", + ) + + adapter: Final = captured["backend_ws"] + assert isinstance(adapter, MuseRealtimeAdapter) + assert captured["force_transcription_model"] == MUSE_MODEL + assert captured["usage_provider"] is adapter + assert captured["exclude_private_content_from_logs"] is True + client_ws.close.assert_awaited_once_with(code=1000, reason="Session closed") + + +def _event(event_type: str) -> str: + return json.dumps({"type": event_type}) diff --git a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py new file mode 100644 index 00000000000..8cc4af836a5 --- /dev/null +++ b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py @@ -0,0 +1,299 @@ +import json + +import pytest + +from litellm.llms.meta.realtime.transformation import ( + MUSE_MODEL, + MuseEventTransformer, + MuseProtocolError, + encode_event, + normalize_language, + parse_session_update, + session_created_event, + session_updated_event, +) + + +def _event(event_type: str, **fields: object) -> str: + return json.dumps({"type": event_type, **fields}) + + +def test_beta_session_builds_authenticated_24khz_handshake_with_hints(): + config = parse_session_update( + _event( + "session.update", + session={ + "type": "transcription", + "input_audio_format": "pcm16", + "turn_detection": {"type": "server_vad"}, + "input_audio_transcription": { + "model": "meta/muse-voice-transcribe-1.0", + "language": "en-US", + "language_bias": ["Spanish", "english", "French"], + "keywords": [" Muse ", "LiteLLM", "Muse"], + "prompt": "must not become a keyword", + }, + }, + ), + "meta/muse-voice-transcribe-1.0", + ) + + assert config.sample_rate == 24_000 + assert config.packet_bytes == 3_840 + assert config.mode == "ENDPOINTING" + assert config.language_bias == ("English", "Spanish", "French") + assert config.keywords == ("Muse", "LiteLLM") + assert config.handshake("Bearer token") == { + "mode": "ENDPOINTING", + "authorization": {"accessToken": "Bearer token"}, + "audioEncoding": "PCM_24KHZ", + "model": MUSE_MODEL, + "partialMode": "CUMULATIVE", + "emitAudioProgress": True, + "keywords": ["Muse", "LiteLLM"], + "languageBias": ["English", "Spanish", "French"], + } + assert "must not become a keyword" not in json.dumps(config.handshake("Bearer token")) + + +def test_ga_session_accepts_16khz_mono_push_to_talk(): + config = parse_session_update( + _event( + "session.update", + session={ + "type": "transcription", + "audio": { + "input": { + "format": {"type": "audio/pcm", "rate": 16000, "channels": 1}, + "turn_detection": None, + "transcription": {"model": MUSE_MODEL, "language": "zh-Hans"}, + } + }, + }, + ), + MUSE_MODEL, + ) + + assert config.sample_rate == 16_000 + assert config.packet_bytes == 2_560 + assert config.mode == "PUSH_TO_TALK" + assert config.language_bias == ("Mandarin Chinese",) + assert config.handshake("Bearer token")["audioEncoding"] == "PCM_16KHZ" + + +@pytest.mark.parametrize( + ("source", "expected"), + [ + ("EN_us", "English"), + ("mandarin chinese", "Mandarin Chinese"), + ("fil-PH", "Tagalog"), + ("iw-IL", "Hebrew"), + ("pt-BR", "Portuguese"), + ], +) +def test_language_normalization_uses_official_muse_names(source: str, expected: str): + assert normalize_language(source) == expected + + +@pytest.mark.parametrize( + ("session", "message"), + [ + ({"input_audio_format": "g711_ulaw"}, "requires pcm16"), + ({"audio": {"input": {"format": {"type": "audio/pcm", "rate": 8000}}}}, "16000 Hz or 24000 Hz"), + ( + {"audio": {"input": {"format": {"type": "audio/pcm", "rate": 24000, "channels": 2}}}}, + "requires mono", + ), + ( + {"input_audio_format": "pcm16", "audio": {"input": {"format": {"type": "audio/pcm"}}}}, + "either beta or GA layout", + ), + ({"input_audio_transcription": {"model": "other-model"}}, "cannot be changed"), + ({"input_audio_transcription": {"keywords": ["valid", ""]}}, "non-empty strings"), + ({"input_audio_transcription": {"language": "xx"}}, "unsupported Muse Voice language"), + ], +) +def test_session_rejects_unsupported_audio_model_and_hints(session: dict[str, object], message: str): + with pytest.raises(MuseProtocolError, match=message): + parse_session_update(_event("session.update", session={"type": "transcription", **session}), MUSE_MODEL) + + +def test_session_events_expose_openai_transcription_shapes(): + config = parse_session_update( + _event( + "session.update", + session={ + "mode": "DIARIZATION", + "audio": { + "input": { + "format": {"type": "audio/pcm", "rate": 24000}, + "transcription": {"model": MUSE_MODEL, "language": "ja", "keywords": ["Meta"]}, + } + }, + }, + ), + MUSE_MODEL, + ) + + created = session_created_event(MUSE_MODEL, "session-before-handshake") + updated = session_updated_event(config, "provider-session") + + assert created["type"] == "session.created" + assert created["session"]["type"] == "transcription" + assert updated["type"] == "session.updated" + assert updated["session"]["id"] == "provider-session" + assert updated["session"]["audio"]["input"]["transcription"] == { + "model": MUSE_MODEL, + "language": "Japanese", + "keywords": ["Meta"], + "language_bias": ["Japanese"], + } + + +def test_turnless_empty_silence_transcript_is_ignored(): + transformer = MuseEventTransformer() + + assert transformer.transform(_event("transcript", transcript="", final=True)) == () + + +def test_transcript_without_speech_start_synthesizes_start_before_delta(): + transformer = MuseEventTransformer() + + events = transformer.transform(_event("transcript", turnId="turn-1", transcript="hello", final=False)) + + assert [event["type"] for event in events] == [ + "input_audio_buffer.speech_started", + "conversation.item.input_audio_transcription.delta", + ] + + +def test_cumulative_partials_emit_only_extensions_and_final_is_authoritative(): + transformer = MuseEventTransformer() + + started = transformer.transform(_event("speechStart", turnId="turn-1")) + first = transformer.transform(_event("transcript", turnId="turn-1", transcript="hello", final=False)) + extension = transformer.transform(_event("transcript", turnId="turn-1", transcript="hello world", final=False)) + rewrite = transformer.transform(_event("transcript", turnId="turn-1", transcript="hullo world", final=False)) + assert transformer.transform(_event("speechComplete", turnId="turn-1", transcript="hullo world")) == () + completed = transformer.transform(_event("speechEnd", turnId="turn-1")) + + assert [event["type"] for event in started] == ["input_audio_buffer.speech_started"] + assert first[0]["delta"] == "hello" + assert extension[0]["delta"] == " world" + assert rewrite == () + assert completed[0]["type"] == "input_audio_buffer.speech_stopped" + assert completed[1]["type"] == "conversation.item.input_audio_transcription.completed" + assert completed[1]["item_id"] == "turn-1" + assert completed[1]["transcript"] == "hullo world" + + +def test_completed_transcript_waits_for_speech_stopped(): + transformer = MuseEventTransformer() + + transformer.transform(_event("speechStart", turnId="turn-1")) + assert transformer.transform(_event("speechComplete", turnId="turn-1", transcript="done")) == () + + released = transformer.transform(_event("speechEnd", turnId="turn-1")) + assert [event["type"] for event in released] == [ + "input_audio_buffer.speech_stopped", + "conversation.item.input_audio_transcription.completed", + ] + + +def test_overlapping_turns_are_emitted_in_provider_turn_order(): + transformer = MuseEventTransformer() + + transformer.transform(_event("speechStart", turnId="turn-a")) + transformer.transform(_event("speechStart", turnId="turn-b")) + assert transformer.transform(_event("transcript", turnId="turn-b", transcript="second", final=False)) == () + assert transformer.transform(_event("speechComplete", turnId="turn-a", transcript="first")) == () + released = transformer.transform(_event("speechEnd", turnId="turn-a")) + + assert [(event["type"], event["item_id"]) for event in released] == [ + ("input_audio_buffer.speech_stopped", "turn-a"), + ("conversation.item.input_audio_transcription.completed", "turn-a"), + ("input_audio_buffer.speech_started", "turn-b"), + ("conversation.item.input_audio_transcription.delta", "turn-b"), + ] + assert transformer.transform(_event("speechComplete", turnId="turn-b", transcript="second final")) == () + final_b = transformer.transform(_event("speechEnd", turnId="turn-b")) + assert final_b[0]["type"] == "input_audio_buffer.speech_stopped" + assert final_b[1]["item_id"] == "turn-b" + assert final_b[1]["transcript"] == "second final" + + +def test_committed_item_id_is_used_for_next_provider_turn(): + transformer = MuseEventTransformer() + + previous_item_id, item_id = transformer.commit_item() + started = transformer.transform(_event("speechStart", turnId="provider-turn")) + transformer.transform(_event("speechComplete", turnId="provider-turn", transcript="hello")) + completed = transformer.transform(_event("speechEnd", turnId="provider-turn")) + + assert previous_item_id is None + assert started[0]["item_id"] == item_id + assert completed[-1]["item_id"] == item_id + + +def test_commit_after_speech_start_reuses_active_item_id(): + transformer = MuseEventTransformer() + + started = transformer.transform(_event("speechStart", turnId="provider-turn")) + previous_item_id, item_id = transformer.commit_item() + transformer.transform(_event("speechComplete", turnId="provider-turn", transcript="hello")) + completed = transformer.transform(_event("speechEnd", turnId="provider-turn")) + + assert previous_item_id is None + assert item_id == "provider-turn" + assert started[0]["item_id"] == item_id + assert completed[-1]["item_id"] == item_id + + +def test_speaker_and_positive_audio_progress_deltas_attach_to_next_completion(): + transformer = MuseEventTransformer() + + transformer.transform(_event("audioProgress", audioProcessedMs=1000)) + transformer.transform(_event("audioProgress", audioProcessedMs=750)) + transformer.transform(_event("audioProgress", audioProcessedMs=1600)) + transformer.transform(_event("speaker", turnId=42, label=" Speaker 2 ")) + transformer.transform(_event("speechComplete", turnId=42, transcript="hello")) + completed = transformer.transform(_event("speechEnd", turnId=42)) + + assert completed[-1]["speaker"] == "Speaker 2" + assert completed[-1]["usage"] == {"type": "duration", "seconds": 1.6} + assert transformer.take_unbilled_usage() is None + + +def test_trailing_audio_progress_is_returned_once(): + transformer = MuseEventTransformer() + + transformer.transform(_event("audioProgress", audioProcessedMs=250)) + + assert transformer.take_unbilled_usage() == {"type": "duration", "seconds": 0.25} + assert transformer.take_unbilled_usage() is None + + +def test_completed_turn_tombstone_suppresses_late_duplicates(): + transformer = MuseEventTransformer() + + transformer.transform(_event("speechComplete", turnId="turn-1", transcript="done")) + + assert transformer.transform(_event("speechComplete", turnId="turn-1", transcript="duplicate")) == () + assert transformer.transform(_event("speaker", turnId="turn-1", label="late")) == () + + +def test_provider_error_is_sanitized_and_encodable(): + token = "private-token" + provider_body = f"authorization failed for Bearer {token}" + transformed = MuseEventTransformer().transform( + _event("error", code="AUTH", message=provider_body, request={"accessToken": token}) + ) + + encoded = encode_event(transformed[0]) + assert json.loads(encoded)["error"] == { + "type": "server_error", + "code": "provider_error", + "message": "Meta Muse realtime transcription failed", + } + assert token not in encoded + assert provider_body not in encoded diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index 0827bbcdc38..aa3f7d45d84 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -4,7 +4,6 @@ from types import TracebackType from typing import Final from unittest.mock import MagicMock, patch - import pytest import litellm @@ -152,6 +151,85 @@ async def test_vertex_credential_resolution_bounds_a_thread_offloaded_refresh(): assert time.monotonic() - start < 5 +@pytest.mark.asyncio +async def test_meta_realtime_rejects_missing_transcription_intent(monkeypatch: pytest.MonkeyPatch): + def mock_get_llm_provider(model, api_base, api_key): + return model.removeprefix("meta/"), "meta", api_key, api_base + + monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) + + with pytest.raises(ValueError, match="requires intent=transcription"): + await realtime_main._arealtime.__wrapped__( + model="meta/muse-voice-transcribe-1.0", + websocket=MagicMock(), + litellm_logging_obj=FakeLogging(), + query_params={"model": "meta/muse-voice-transcribe-1.0"}, + ) + + +@pytest.mark.asyncio +async def test_meta_realtime_rejects_unsupported_model_before_connecting(monkeypatch: pytest.MonkeyPatch): + def mock_get_llm_provider(model, api_base, api_key): + return model.removeprefix("meta/"), "meta", api_key, api_base + + monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) + + with pytest.raises(ValueError, match="Unsupported Meta realtime model: other-model"): + await realtime_main._arealtime.__wrapped__( + model="meta/other-model", + websocket=MagicMock(), + litellm_logging_obj=FakeLogging(), + query_params={"model": "meta/other-model", "intent": "transcription"}, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("explicit_key", "model_key", "meta_key", "expected"), + [ + ("explicit", "model-env", "meta-env", "explicit"), + (None, "model-env", "meta-env", "model-env"), + (None, None, "meta-env", "meta-env"), + ], +) +async def test_meta_realtime_credential_precedence_is_forwarded_to_handler( + monkeypatch: pytest.MonkeyPatch, + explicit_key: str | None, + model_key: str | None, + meta_key: str | None, + expected: str, +): + captured: dict[str, object] = {} + + def mock_get_llm_provider(model, api_base, api_key): + return model.removeprefix("meta/"), "meta", meta_key, api_base + + def mock_get_secret_str(name: str): + return {"MODEL_API_KEY": model_key, "META_API_KEY": meta_key}.get(name) + + async def mock_async_realtime(self, **kwargs): + captured.update(kwargs) + + monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) + monkeypatch.setattr(realtime_main, "get_secret_str", mock_get_secret_str) + monkeypatch.setattr( + "litellm.llms.meta.realtime.handler.MetaRealtime.async_realtime", + mock_async_realtime, + ) + + await realtime_main._arealtime.__wrapped__( + model="meta/muse-voice-transcribe-1.0", + websocket=MagicMock(), + litellm_logging_obj=FakeLogging(), + api_key=explicit_key, + query_params={"model": "meta/muse-voice-transcribe-1.0", "intent": "transcription"}, + ) + + assert captured["model"] == "muse-voice-transcribe-1.0" + assert captured["api_key"] == expected + assert captured["query_params"] == {"model": "muse-voice-transcribe-1.0", "intent": "transcription"} + + @pytest.mark.asyncio async def test_arealtime_vertex_branch_resolves_credentials_under_a_bound(monkeypatch): """The wiring half of the regression: the vertex branch of _arealtime must From 1acb994998704e0e1df3f873b52a9226582eb05b Mon Sep 17 00:00:00 2001 From: Young Han Date: Wed, 2 Sep 2026 13:40:35 -0700 Subject: [PATCH 2/7] fix(realtime): bound Muse audio before decoding --- litellm/llms/meta/realtime/handler.py | 10 +++++++++- .../meta/realtime/test_meta_realtime_handler.py | 17 ++++++++++------- 2 files changed, 19 insertions(+), 8 deletions(-) diff --git a/litellm/llms/meta/realtime/handler.py b/litellm/llms/meta/realtime/handler.py index 9eb98e4ba0b..9dbe3ea8b5d 100644 --- a/litellm/llms/meta/realtime/handler.py +++ b/litellm/llms/meta/realtime/handler.py @@ -268,6 +268,15 @@ class MuseRealtimeAdapter: if not isinstance(audio_value, str): await self._reject("invalid_request_error", "invalid_audio", "Audio must be a base64 string") return + max_backlog_bytes: Final = config.bytes_per_second * _MAX_AUDIO_BACKLOG_SECONDS + max_encoded_bytes: Final = 4 * ((max_backlog_bytes + 2) // 3) + if len(audio_value) > max_encoded_bytes: + await self._reject( + "invalid_request_error", + "audio_backlog_exceeded", + "Audio append exceeds the four-second backlog limit", + ) + return try: audio: Final = base64.b64decode(audio_value, validate=True) except (binascii.Error, ValueError): @@ -278,7 +287,6 @@ class MuseRealtimeAdapter: return if not audio: return - max_backlog_bytes: Final = config.bytes_per_second * _MAX_AUDIO_BACKLOG_SECONDS if len(audio) > max_backlog_bytes: await self._reject( "invalid_request_error", diff --git a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py index 3d31b3a5534..14ba972fe14 100644 --- a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py +++ b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py @@ -3,7 +3,7 @@ import base64 import json from collections.abc import Callable from typing import Final -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -246,18 +246,21 @@ async def test_absolute_pacing_delays_only_audio_ahead_of_wall_time(): @pytest.mark.asyncio -async def test_append_larger_than_four_seconds_is_rejected_without_dropping_prefix(): +async def test_append_larger_than_four_seconds_is_rejected_without_decoding(): adapter, provider_ws, _ = await _configured_adapter(rate=16_000) - oversized_pcm: Final = b"\x00\x00" * (16_000 * 4 + 1) + max_pcm_bytes: Final = 16_000 * 2 * 4 + oversized_audio: Final = "A" * (4 * ((max_pcm_bytes + 2) // 3) + 1) + + with patch( # test-quality-ok: proves rejection happens before an attacker-controlled allocation + "litellm.llms.meta.realtime.handler.base64.b64decode" + ) as decode: + await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": oversized_audio})) - await adapter.send( - json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(oversized_pcm).decode()}) - ) error: Final = json.loads(await adapter.recv()) - assert error["error"]["code"] == "audio_backlog_exceeded" assert adapter.close_code == 1008 assert not any(isinstance(frame, bytes) for frame in provider_ws.sent) + decode.assert_not_called() await adapter.close() From 17fde7a261c6c3aff6a2eea72d0a883fd7c17c51 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 11 Sep 2026 19:40:39 -0700 Subject: [PATCH 3/7] refactor(realtime): move Meta Muse Voice onto BaseRealtimeConfig Replace the hand-rolled Meta realtime handler with a MetaRealtimeConfig that plugs into the shared realtime handler and RealTimeStreaming relay. Clients keep speaking the OpenAI realtime wire: session.update, input_audio_buffer.append/commit and the OpenAI transcription events. Unsupported transcription settings are logged and dropped, matching the Gemini realtime precedent, and the Meta-specific session.mode, keywords, language_bias, DIARIZATION and speaker extensions are removed. Drop the MODEL_API_KEY env var in favor of the standard META_API_KEY, remove the private-logging flag so spend logs record the transcript the same way other realtime models do, and add per-second pricing for muse-voice-transcribe-1.0. The relay now sends raw bytes from transform_realtime_request straight to the backend after pace_backend_send, and transcription sessions never trigger response.create. --- .../litellm_core_utils/realtime_streaming.py | 91 +- .../llms/base_llm/realtime/transformation.py | 8 +- litellm/llms/meta/__init__.py | 3 - litellm/llms/meta/realtime/__init__.py | 10 - litellm/llms/meta/realtime/handler.py | 669 --------------- litellm/llms/meta/realtime/transformation.py | 774 ++++++++++-------- ...odel_prices_and_context_window_backup.json | 1 + litellm/realtime_api/main.py | 34 +- litellm/types/llms/meta.py | 58 ++ litellm/types/llms/openai.py | 1 - litellm/utils.py | 4 + model_prices_and_context_window.json | 1 + .../test_realtime_streaming.py | 104 +-- .../realtime/test_meta_realtime_handler.py | 452 ---------- .../test_meta_realtime_transformation.py | 418 ++++++++-- tests/test_litellm/realtime_api/test_main.py | 64 +- 16 files changed, 889 insertions(+), 1803 deletions(-) delete mode 100644 litellm/llms/meta/__init__.py delete mode 100644 litellm/llms/meta/realtime/__init__.py delete mode 100644 litellm/llms/meta/realtime/handler.py create mode 100644 litellm/types/llms/meta.py delete mode 100644 tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 4be7dd6b4ce..06d9241b826 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -19,7 +19,7 @@ from litellm.types.llms.openai import ( OpenAIRealtimeStreamResponseBaseObject, OpenAIRealtimeStreamSessionEvents, ) -from litellm.types.realtime import ALL_DELTA_TYPES, RealtimeInputAudioTranscriptionUsage +from litellm.types.realtime import ALL_DELTA_TYPES from .litellm_logging import Logging as LiteLLMLogging from .realtime_errors import client_close_code, realtime_error_event, websocket_close_reason @@ -116,10 +116,6 @@ class RealtimeEventNormalizer(Protocol): def patch_outgoing_session(self, session: dict) -> dict: ... -class RealtimeUsageProvider(Protocol): - def unbilled_usage_on_session_close(self, model: str) -> RealtimeInputAudioTranscriptionUsage | None: ... - - DefaultLoggedRealTimeEventTypes: Final = [ "session.created", "response.create", @@ -143,8 +139,6 @@ class RealTimeStreaming: force_transcription_model: str | None = None, event_normalizer: RealtimeEventNormalizer | None = None, logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER, - usage_provider: RealtimeUsageProvider | None = None, - exclude_private_content_from_logs: bool = False, ): self.websocket: _ClientWebSocket = websocket self.backend_ws = backend_ws @@ -206,10 +200,6 @@ class RealTimeStreaming: self._is_transcription_session: bool = force_transcription_model is not None # Optional per-provider GA event normalizer (e.g. XAIRealtimeNormalizer). self._event_normalizer = event_normalizer - self._usage_provider: RealtimeUsageProvider | None = ( - usage_provider if usage_provider is not None else provider_config - ) - self._exclude_private_content_from_logs = exclude_private_content_from_logs # Per-connection caps for pre-setup audio frames (message count + total bytes). _MAX_BUFFERED_MESSAGES: int = 200 @@ -247,7 +237,7 @@ class RealTimeStreaming: def _should_store_message( self, - message_obj: dict[str, Any] | OpenAIRealtimeEvents, # mutable-ok: existing realtime event contract + message_obj: dict | OpenAIRealtimeEvents, ) -> bool: _msg_type: Final = message_obj["type"] if "type" in message_obj else None if self.logged_real_time_event_types == "*": @@ -256,54 +246,16 @@ class RealTimeStreaming: return True return False - def _message_for_logging( - self, - message_obj: dict[str, Any], # mutable-ok: existing realtime event contract - ) -> dict[str, Any]: # mutable-ok: logging stores concrete event dictionaries - if not self._exclude_private_content_from_logs: - return message_obj - logged_message: dict[str, Any] = { # mutable-ok: incrementally builds the sanitized event copy - key: message_obj[key] - for key in ( - "type", - "event_id", - "item_id", - "response_id", - "conversation_id", - "session_id", - "content_index", - "output_index", - "model", - "mode", - "usage", - ) - if key in message_obj - } - session: Final = message_obj.get("session") - if isinstance(session, dict): - logged_session: Final[dict[str, Any]] = { # mutable-ok: sanitized JSON session snapshot - key: session[key] for key in ("id", "model", "mode", "type") if key in session - } - if logged_session: - logged_message["session"] = logged_session - return logged_message - def store_message(self, message: str | bytes | dict | OpenAIRealtimeEvents): """Store message in list""" if isinstance(message, bytes): message = message.decode("utf-8") if isinstance(message, dict): # TypedDict union members do not narrow to plain dict for mypy. - parsed_message_obj: dict[str, Any] = cast( # cast-ok: TypedDict events are JSON dictionaries - dict[str, Any], message - ) + message_obj: dict[str, Any] = cast(dict[str, Any], message) else: - parsed_message_obj = cast( # cast-ok: parsed realtime events are JSON dictionaries - dict[str, Any], json.loads(message) - ) - if not self._exclude_private_content_from_logs: - self._collect_tool_calls_from_response_done(parsed_message_obj) - message_obj: Final = self._message_for_logging(parsed_message_obj) + message_obj = cast(dict[str, Any], json.loads(cast(str, message))) + self._collect_tool_calls_from_response_done(cast(dict, message_obj)) if not self._should_store_message(message_obj): return try: @@ -321,8 +273,6 @@ class RealTimeStreaming: def _collect_user_input_from_client_event(self, message: str | dict) -> None: """Extract user text content from client WebSocket events for spend logging.""" - if self._exclude_private_content_from_logs: - return try: if isinstance(message, str): msg_obj = json.loads(message) @@ -359,8 +309,6 @@ class RealTimeStreaming: def _collect_user_input_from_backend_event(self, event_obj: dict | OpenAIRealtimeEvents) -> None: """Extract user voice transcription from backend events for spend logging.""" - if self._exclude_private_content_from_logs: - return try: event_type: Final = event_obj.get("type", "") if event_type == "conversation.item.input_audio_transcription.completed": @@ -416,9 +364,9 @@ class RealTimeStreaming: pass def _flush_unbilled_transcription_usage(self) -> None: - if self._usage_provider is None: + if self.provider_config is None: return - usage: Final = self._usage_provider.unbilled_usage_on_session_close(self.model) + usage: Final = self.provider_config.unbilled_usage_on_session_close(self.model) if usage is None: return flush_event: Final = ( @@ -455,27 +403,12 @@ class RealTimeStreaming: except (AttributeError, TypeError): pass - def _input_for_logging( - self, - message: str | dict, # mutable-ok: existing realtime input contract - ) -> str | dict: # mutable-ok: logging stores concrete event dictionaries - if not self._exclude_private_content_from_logs: - return message - try: - parsed_message: Final[object] = message if isinstance(message, dict) else json.loads(message) - except (json.JSONDecodeError, TypeError): - return {} # mutable-ok: empty JSON logging payload - if not isinstance(parsed_message, dict): - return {} # mutable-ok: empty JSON logging payload - return self._message_for_logging(parsed_message) - def store_input(self, message: str | dict): """Store input message""" - logged_message: Final[str | dict] = self._input_for_logging(message) # mutable-ok: logging payload - self.input_message = logged_message if isinstance(logged_message, dict) else {} + self.input_message = message if isinstance(message, dict) else {} self._collect_user_input_from_client_event(message) if self.logging_obj: - self.logging_obj.pre_call(input=logged_message, api_key="") + self.logging_obj.pre_call(input=message, api_key="") async def log_messages(self): """Log messages in list""" @@ -512,6 +445,12 @@ class RealTimeStreaming: ) sent = False for msg in transformed: + if isinstance(msg, bytes): + await self.provider_config.pace_backend_send(msg) + await self.backend_ws.send(msg) + self._content_sent_after_setup = True + sent = True + continue try: msg_obj = _decode_json_object(msg) except (json.JSONDecodeError, TypeError): diff --git a/litellm/llms/base_llm/realtime/transformation.py b/litellm/llms/base_llm/realtime/transformation.py index cfcde7c6e9e..e44cccc1a62 100644 --- a/litellm/llms/base_llm/realtime/transformation.py +++ b/litellm/llms/base_llm/realtime/transformation.py @@ -1,4 +1,5 @@ from abc import ABC, abstractmethod +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any import httpx @@ -54,9 +55,12 @@ class BaseRealtimeConfig(ABC): message: str, model: str, session_configuration_request: str | None = None, - ) -> list[str]: + ) -> Sequence[str | bytes]: pass + async def pace_backend_send(self, message: bytes) -> None: + return None + def is_setup_message(self, msg_obj: dict) -> bool: return False @@ -79,7 +83,7 @@ class BaseRealtimeConfig(ABC): model: str, logging_session_id: str, session_configuration_request: str | None = None, - ) -> dict | OpenAIRealtimeStreamSessionEvents | None: + ) -> Mapping[str, object] | OpenAIRealtimeStreamSessionEvents | None: """ Optional hook for providers that defer session setup until client `session.update`. diff --git a/litellm/llms/meta/__init__.py b/litellm/llms/meta/__init__.py deleted file mode 100644 index 7c7d32788a2..00000000000 --- a/litellm/llms/meta/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .realtime import MetaRealtime, MuseRealtimeAdapter - -__all__ = ("MetaRealtime", "MuseRealtimeAdapter") diff --git a/litellm/llms/meta/realtime/__init__.py b/litellm/llms/meta/realtime/__init__.py deleted file mode 100644 index 6398765da24..00000000000 --- a/litellm/llms/meta/realtime/__init__.py +++ /dev/null @@ -1,10 +0,0 @@ -from .handler import MetaRealtime, MuseRealtimeAdapter -from .transformation import MuseEventTransformer, MuseProtocolError, MuseSessionConfig - -__all__ = ( - "MetaRealtime", - "MuseEventTransformer", - "MuseProtocolError", - "MuseRealtimeAdapter", - "MuseSessionConfig", -) diff --git a/litellm/llms/meta/realtime/handler.py b/litellm/llms/meta/realtime/handler.py deleted file mode 100644 index 9dbe3ea8b5d..00000000000 --- a/litellm/llms/meta/realtime/handler.py +++ /dev/null @@ -1,669 +0,0 @@ -from __future__ import annotations - -import asyncio -import base64 -import binascii -import contextlib -import json -import time -import uuid -from collections.abc import Awaitable, Callable, Mapping -from typing import Final, Protocol -from urllib.parse import urlparse, urlunparse - -from pydantic import JsonValue, TypeAdapter, ValidationError - -from litellm._logging import verbose_proxy_logger -from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging -from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming -from litellm.llms.custom_httpx.http_handler import get_shared_realtime_ssl_context -from litellm.types.realtime import RealtimeInputAudioTranscriptionUsage, RealtimeQueryParams - -from .transformation import ( - MUSE_MODEL, - MuseEventTransformer, - MuseProtocolError, - MuseSessionConfig, - encode_event, - error_event, - parse_session_update, - session_created_event, - session_updated_event, -) - -DEFAULT_MUSE_REALTIME_URL: Final = "wss://api.meta.ai/v1/asr/realtime" -_MAX_AUDIO_BACKLOG_SECONDS: Final = 4 -_MAX_PENDING_PROVIDER_EVENTS: Final = 256 -_JSON_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) - - -class _ProviderWebSocket(Protocol): - async def send(self, message: str | bytes) -> None: ... - - async def recv(self, decode: bool | None = None) -> str | bytes: ... - - async def close(self, code: int = 1000, reason: str = "") -> None: ... - - -class _ClientWebSocketExceptions(Protocol): - ConnectionClosed: type[Exception] - - -class _ClientWebSocket(Protocol): - exceptions: _ClientWebSocketExceptions - - @property - def scope(self) -> Mapping[str, object]: ... - - async def send_text(self, data: str) -> None: ... - - async def receive_text(self) -> str: ... - - async def close(self, code: int = 1000, reason: str | None = None) -> None: ... - - -class WebSocketConnect(Protocol): - def __call__( - self, - url: str, - *, - open_timeout: float, - max_size: int | None, - ssl: object | None, - ) -> Awaitable[_ProviderWebSocket]: ... - - -class MuseAdapterError(RuntimeError): - def __init__(self, message: str, *, close_code: int) -> None: - super().__init__(message) - self.close_code: Final = close_code - - -class MuseRealtimeAdapter: - def __init__( - self, - *, - model: str, - api_key: str, - api_base: str | None = None, - timeout: float | None = None, - websocket_connect: WebSocketConnect | None = None, - monotonic: Callable[[], float] = time.monotonic, - sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, - terminate_client: Callable[[int], Awaitable[None]] | None = None, - ) -> None: - if model.removeprefix("meta/") != MUSE_MODEL: - raise ValueError("unsupported Meta realtime model") - self._model: Final = model.removeprefix("meta/") - self._access_token: Final = normalize_access_token(api_key) - self._url: Final = build_muse_realtime_url(api_base) - self._timeout: Final = timeout or 10.0 - self._websocket_connect = websocket_connect - self._monotonic: Final = monotonic - self._sleep: Final = sleep - self._terminate_client: Final = terminate_client - self._provider_ws: _ProviderWebSocket | None = None - self._config: MuseSessionConfig | None = None - self._session_id: str = f"sess_{uuid.uuid4().hex}" - self._events: Final[asyncio.Queue[str | BaseException]] = asyncio.Queue(maxsize=_MAX_PENDING_PROVIDER_EVENTS) - self._events.put_nowait(encode_event(session_created_event(self._model, self._session_id))) - self._transformer: Final = MuseEventTransformer() - self._audio_condition: Final = asyncio.Condition() - self._pending_audio: bytearray = bytearray() - self._audio_generation: int = 0 - self._flush_requested: bool = False - self._end_requested: bool = False - self._end_stream_sent: bool = False - self._audio_consumed: bool = False - self._closed: bool = False - self._resources_closed: bool = False - self._sender_task: asyncio.Task[None] | None = None - self._receiver_task: asyncio.Task[None] | None = None - self.close_code: int = 1000 - self.close_reason: str = "Session closed" - - async def send(self, message: str | bytes) -> None: - if self._closed: - raise MuseAdapterError("Meta Muse realtime session is closed", close_code=self.close_code) - if isinstance(message, bytes): - await self._reject("invalid_request_error", "invalid_event", "Client events must be JSON text") - return - try: - event: Final = _parse_client_event(message) - event_type: Final = event.get("type") - if event_type in ("session.update", "transcription_session.update"): - await self._handle_session_update(message) - return - if event_type == "input_audio_buffer.append": - await self._handle_audio_append(event) - return - if event_type == "input_audio_buffer.clear": - await self._clear_audio() - return - if event_type == "input_audio_buffer.commit": - await self._commit_audio() - return - if event_type == "input_audio_buffer.end": - await self._end_audio() - return - await self._emit( - error_event( - "invalid_request_error", - "unsupported_event", - f"Event type {event_type!r} is not supported for Meta Muse transcription", - ) - ) - except MuseProtocolError as exc: - await self._reject("invalid_request_error", "invalid_event", str(exc)) - - async def recv(self, decode: bool | None = None) -> str | bytes: - event: Final = await self._events.get() - if isinstance(event, BaseException): - close_code: Final = _exception_close_code(event) if isinstance(event, Exception) else 1011 - if self._terminate_client is not None: - await self._terminate_client(close_code) - raise event - return event.encode("utf-8") if decode is False else event - - async def close(self, code: int = 1000, reason: str = "") -> None: - if self._resources_closed: - return - self._closed = True - self._resources_closed = True - self.close_code = sanitize_close_code(code) - self.close_reason = safe_close_reason(self.close_code) - async with self._audio_condition: - self._end_requested = True - self._audio_condition.notify_all() - tasks: Final = tuple(task for task in (self._sender_task, self._receiver_task) if task is not None) - for task in tasks: - task.cancel() - if tasks: - await asyncio.gather(*tasks, return_exceptions=True) - provider_ws: Final = self._provider_ws - if provider_ws is not None: - with contextlib.suppress(Exception): - await provider_ws.close(code=self.close_code, reason=self.close_reason) - - def unbilled_usage_on_session_close(self, model: str) -> RealtimeInputAudioTranscriptionUsage | None: - return self._transformer.take_unbilled_usage() - - async def _handle_session_update(self, message: str) -> None: - config: Final = parse_session_update(message, self._model) - if self._config is not None: - if config != self._config: - await self._reject( - "invalid_request_error", - "session_configuration_locked", - "Meta Muse session configuration cannot change after setup", - ) - return - await self._emit(session_updated_event(config, self._session_id)) - return - await self._connect(config) - - async def _connect(self, config: MuseSessionConfig) -> None: - connector: Final = self._websocket_connect or _default_websocket_connect - last_error: Exception | None = None # rebind-ok: records the latest bounded handshake attempt - for attempt in range(2): - provider_ws: _ProviderWebSocket | None = None - try: - provider_ws = await connector( - self._url, - open_timeout=self._timeout, - max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, - ssl=_ssl_config(self._url), - ) - await provider_ws.send(json.dumps(config.handshake(self._access_token), separators=(",", ":"))) - raw_ack: str | bytes = await asyncio.wait_for( # rebind-ok: one response per handshake attempt - provider_ws.recv(), timeout=self._timeout - ) - session_id: str = _parse_handshake_ack(raw_ack) # rebind-ok: one ID per handshake attempt - self._provider_ws = provider_ws - self._config = config - self._transformer.configure(config) - self._session_id = session_id - self._sender_task = asyncio.create_task(self._send_audio(), name="meta-muse-realtime-send") - self._receiver_task = asyncio.create_task(self._receive_events(), name="meta-muse-realtime-receive") - await self._emit(session_updated_event(config, session_id)) - return - except asyncio.CancelledError: - if provider_ws is not None: - with contextlib.suppress(Exception): - await provider_ws.close() - raise - except Exception as exc: # noqa: BLE001 # connector implementations expose heterogeneous transport errors - last_error = exc - if provider_ws is not None: - with contextlib.suppress(Exception): - await provider_ws.close() - close_code: int = _exception_close_code(exc) # rebind-ok: classified per handshake attempt - retryable_transport_error: bool = not isinstance( # rebind-ok: classified per handshake attempt - exc, (MuseAdapterError, MuseProtocolError) - ) - if attempt == 0 and retryable_transport_error and close_code in (1011, 1013): - continue - self.close_code = close_code - self.close_reason = safe_close_reason(close_code) - await self._emit( - error_event( - "server_error" if close_code != 1008 else "invalid_request_error", - "provider_connection_error", - "Meta Muse realtime handshake failed", - ) - ) - await self._events.put(MuseAdapterError("Meta Muse realtime handshake failed", close_code=close_code)) - await self._mark_terminated(close_code) - return - assert last_error is not None - raise MuseAdapterError("Meta Muse realtime handshake failed", close_code=1011) - - async def _handle_audio_append(self, event: Mapping[str, JsonValue]) -> None: - config: Final = self._require_configured() - if self._end_requested or self._end_stream_sent: - await self._reject("invalid_request_error", "input_ended", "Audio input has already ended") - return - audio_value: Final = event.get("audio") - if not isinstance(audio_value, str): - await self._reject("invalid_request_error", "invalid_audio", "Audio must be a base64 string") - return - max_backlog_bytes: Final = config.bytes_per_second * _MAX_AUDIO_BACKLOG_SECONDS - max_encoded_bytes: Final = 4 * ((max_backlog_bytes + 2) // 3) - if len(audio_value) > max_encoded_bytes: - await self._reject( - "invalid_request_error", - "audio_backlog_exceeded", - "Audio append exceeds the four-second backlog limit", - ) - return - try: - audio: Final = base64.b64decode(audio_value, validate=True) - except (binascii.Error, ValueError): - await self._reject("invalid_request_error", "invalid_audio", "Audio must be valid base64") - return - if len(audio) % 2: - await self._reject("invalid_request_error", "invalid_audio", "PCM16 audio must contain complete samples") - return - if not audio: - return - if len(audio) > max_backlog_bytes: - await self._reject( - "invalid_request_error", - "audio_backlog_exceeded", - "Audio append exceeds the four-second Muse backlog limit", - ) - return - async with self._audio_condition: - await self._audio_condition.wait_for( - lambda: self._closed or len(self._pending_audio) + len(audio) <= max_backlog_bytes - ) - if self._closed: - raise MuseAdapterError("Meta Muse realtime session is closed", close_code=self.close_code) - self._pending_audio.extend(audio) - self._audio_condition.notify_all() - - async def _clear_audio(self) -> None: - self._require_configured() - async with self._audio_condition: - self._pending_audio.clear() - self._audio_generation += 1 - self._flush_requested = False - self._audio_condition.notify_all() - await self._emit( - { # mutable-ok: OpenAI-compatible JSON event - "type": "input_audio_buffer.cleared", - "event_id": f"event_{uuid.uuid4().hex}", - } - ) - - async def _commit_audio(self) -> None: - config: Final = self._require_configured() - previous_item_id, item_id = self._transformer.commit_item() - async with self._audio_condition: - self._flush_requested = True - if config.mode == "PUSH_TO_TALK": - self._end_requested = True - self._audio_condition.notify_all() - await self._emit( - { # mutable-ok: OpenAI-compatible JSON event - "type": "input_audio_buffer.committed", - "event_id": f"event_{uuid.uuid4().hex}", - "previous_item_id": previous_item_id, - "item_id": item_id, - } - ) - - async def _end_audio(self) -> None: - self._require_configured() - async with self._audio_condition: - self._flush_requested = True - self._end_requested = True - self._audio_condition.notify_all() - - async def _send_audio(self) -> None: - config: Final = self._require_configured() - provider_ws: Final = self._require_provider_ws() - pacing_origin: float | None = None # rebind-ok: initialized when the first packet is ready - sent_duration: float = 0.0 # rebind-ok: absolute pacing clock advances after each packet - try: - while True: - packet, pacing_origin, ended = await self._next_audio_packet( - config, - pacing_origin, - sent_duration, - ) - if ended: - break - if packet is None: - continue - await provider_ws.send(packet) - self._audio_consumed = True - sent_duration += len(packet) / config.bytes_per_second - await self._send_end_stream() - except asyncio.CancelledError: - raise - except Exception as exc: # noqa: BLE001 # WebSocket implementations expose heterogeneous transport errors - await self._fail_provider(exc, phase="audio send") - - async def _next_audio_packet( - self, - config: MuseSessionConfig, - pacing_origin: float | None, - sent_duration: float, - ) -> tuple[bytes | None, float | None, bool]: - async with self._audio_condition: - await self._audio_condition.wait_for( - lambda: ( - self._closed - or len(self._pending_audio) >= config.packet_bytes - or (self._flush_requested and bool(self._pending_audio)) - or (self._end_requested and not self._pending_audio) - ) - ) - if self._closed or (self._end_requested and not self._pending_audio): - return None, pacing_origin, True - packet_size: Final = min(config.packet_bytes, len(self._pending_audio)) - if packet_size < config.packet_bytes and not self._flush_requested: - return None, pacing_origin, False - generation: Final = self._audio_generation - current_time: Final = self._monotonic() - effective_origin: Final = ( - current_time - sent_duration - if pacing_origin is None or current_time > pacing_origin + sent_duration - else pacing_origin - ) - deadline: Final = effective_origin + sent_duration - delay: Final = deadline - self._monotonic() - if delay > 0: - await self._sleep(delay) - async with self._audio_condition: - if generation != self._audio_generation: - return None, effective_origin, False - actual_size: Final = min(packet_size, len(self._pending_audio)) - packet: Final = bytes(self._pending_audio[:actual_size]) - del self._pending_audio[:actual_size] - if not self._pending_audio: - self._flush_requested = False - self._audio_condition.notify_all() - return packet or None, effective_origin, False - - async def _send_end_stream(self) -> None: - if self._end_stream_sent: - return - provider_ws: Final = self._require_provider_ws() - await provider_ws.send('{"type":"endStream"}') - self._end_stream_sent = True - - async def _receive_events(self) -> None: - provider_ws: Final = self._require_provider_ws() - try: - while True: - raw: str | bytes = await provider_ws.recv() # rebind-ok: one provider frame per iteration - if not isinstance(raw, str): - raise MuseProtocolError("provider returned a non-text event") - for event in self._transformer.transform(raw): - await self._emit(event) - except asyncio.CancelledError: - raise - except Exception as exc: # noqa: BLE001 # provider close exceptions vary by WebSocket implementation - close_code: Final = _exception_close_code(exc) - if close_code == 1000 and self._end_stream_sent: - await self._mark_terminated(1000) - await self._events.put(MuseAdapterError("Meta Muse realtime session completed", close_code=1000)) - return - failure: Final = MuseAdapterError( - "Meta Muse realtime closed before input ended", - close_code=1011 if close_code == 1000 else close_code, - ) - await self._fail_provider(failure, phase="receive") - - async def _fail_provider(self, exc: Exception, *, phase: str) -> None: - close_code: Final = _exception_close_code(exc) - self.close_code = close_code - self.close_reason = safe_close_reason(close_code) - await self._emit( - error_event( - "server_error", - "provider_connection_error", - f"Meta Muse realtime {phase} failed", - ) - ) - await self._events.put(MuseAdapterError(f"Meta Muse realtime {phase} failed", close_code=close_code)) - await self._mark_terminated(close_code) - - async def _mark_terminated(self, close_code: int) -> None: - self._closed = True - self.close_code = sanitize_close_code(close_code) - self.close_reason = safe_close_reason(self.close_code) - async with self._audio_condition: - self._audio_condition.notify_all() - - async def _terminate(self, close_code: int) -> None: - await self._mark_terminated(close_code) - if self._terminate_client is not None: - await self._terminate_client(self.close_code) - - async def _reject(self, error_type: str, code: str, message: str) -> None: - self.close_code = 1008 - self.close_reason = safe_close_reason(1008) - await self._emit(error_event(error_type, code, message)) - await self._events.put(MuseAdapterError(message, close_code=1008)) - await self._mark_terminated(1008) - - async def _emit(self, event: Mapping[str, object]) -> None: - await self._events.put(encode_event(event)) - - def _require_configured(self) -> MuseSessionConfig: - if self._config is None: - raise MuseProtocolError("send session.update before audio events") - return self._config - - def _require_provider_ws(self) -> _ProviderWebSocket: - if self._provider_ws is None: - raise MuseProtocolError("Meta Muse provider connection is not ready") - return self._provider_ws - - -class MetaRealtime: - async def async_realtime( - self, - model: str, - websocket: _ClientWebSocket, - logging_obj: LiteLLMLogging, - api_base: str | None = None, - api_key: str | None = None, - client: object | None = None, - timeout: float | None = None, - query_params: RealtimeQueryParams | None = None, - user_api_key_dict: object | None = None, - litellm_metadata: Mapping[str, object] | None = None, - websocket_connect: WebSocketConnect | None = None, - monotonic: Callable[[], float] = time.monotonic, - sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, - **kwargs: object, # kwargs-ok: realtime dispatcher forwards provider-neutral options - ) -> None: - if api_key is None or not api_key.strip(): - await _send_client_error_and_close(websocket, "Meta Model API key is required") - return - try: - adapter: Final = MuseRealtimeAdapter( - model=model, - api_key=api_key, - api_base=api_base, - timeout=timeout, - websocket_connect=websocket_connect, - monotonic=monotonic, - sleep=sleep, - terminate_client=lambda code: _close_client(websocket, code), - ) - except ValueError: - await _send_client_error_and_close(websocket, "Invalid Meta Muse realtime configuration") - return - realtime_streaming: Final = RealTimeStreaming( - websocket, - adapter, # pyright: ignore[reportArgumentType] # raw adapter intentionally matches the websocket surface - logging_obj, - model=model, - user_api_key_dict=user_api_key_dict, - request_data={ # mutable-ok: relay request metadata payload - "litellm_metadata": dict(litellm_metadata or {}) # mutable-ok: relay owns its metadata copy - }, - force_transcription_model=model, - usage_provider=adapter, - exclude_private_content_from_logs=True, - ) - try: - await realtime_streaming.bidirectional_forward() - except MuseAdapterError as exc: - adapter.close_code = exc.close_code - adapter.close_reason = safe_close_reason(exc.close_code) - except Exception: # noqa: BLE001 # relay errors are normalized before closing the accepted client socket - adapter.close_code = 1011 - adapter.close_reason = safe_close_reason(1011) - verbose_proxy_logger.exception("Meta Muse realtime session failed") - finally: - await adapter.close(code=adapter.close_code) - await _close_client(websocket, adapter.close_code) - - -def normalize_access_token(api_key: str) -> str: - stripped: Final = api_key.strip() - if not stripped: - raise ValueError("Meta Model API key is required") - parts: Final = stripped.split(None, 1) - if parts[0].casefold() == "bearer": - if len(parts) != 2 or not parts[1].strip(): - raise ValueError("Meta Model API key must include a token after Bearer") - return f"Bearer {parts[1].strip()}" - return f"Bearer {stripped}" - - -def build_muse_realtime_url(api_base: str | None) -> str: - if api_base is None: - return DEFAULT_MUSE_REALTIME_URL - parsed: Final = urlparse(api_base.strip()) - scheme: Final = "wss" if parsed.scheme == "https" else parsed.scheme - if ( - scheme != "wss" - or not parsed.hostname - or parsed.username is not None - or parsed.password is not None - or parsed.fragment - ): - raise ValueError("Meta api_base must be an absolute wss:// or https:// URL without credentials or a fragment") - netloc: Final = f"{parsed.hostname}:{parsed.port}" if parsed.port is not None else parsed.hostname - return urlunparse((scheme, netloc, "/v1/asr/realtime", "", "", "")) - - -def sanitize_close_code(code: int | None) -> int: - if code is not None and code in (1000, 1008, 1011, 1013): - return code - return 1011 - - -def safe_close_reason(code: int) -> str: - return { # mutable-ok: immutable-by-convention close-reason lookup - 1000: "Session closed", - 1008: "Invalid realtime transcription request", - 1011: "Realtime transcription service error", - 1013: "Realtime transcription service unavailable", - }.get(code, "Realtime transcription service error") - - -def _parse_client_event(payload: str) -> Mapping[str, JsonValue]: - try: - value: Final = _JSON_ADAPTER.validate_json(payload) - except ValidationError: - raise MuseProtocolError("invalid JSON object") from None - if not isinstance(value, dict): - raise MuseProtocolError("message must be a JSON object") - event_type: Final = value.get("type") - if not isinstance(event_type, str) or not event_type: - raise MuseProtocolError("message type must be a non-empty string") - return value - - -def _parse_handshake_ack(raw: str | bytes) -> str: - if not isinstance(raw, str): - raise MuseProtocolError("provider returned a non-text handshake response") - message: Final = _parse_json_object(raw) - if message.get("type") == "error": - raise MuseAdapterError("Meta Muse realtime handshake was rejected", close_code=1008) - session_id: Final = message.get("sessionId") - if not isinstance(session_id, str) or not session_id.strip(): - raise MuseProtocolError("provider returned an invalid handshake response") - return session_id.strip() - - -def _parse_json_object(payload: str) -> Mapping[str, JsonValue]: - try: - value: Final = _JSON_ADAPTER.validate_json(payload) - except ValidationError: - raise MuseProtocolError("invalid provider JSON object") from None - if not isinstance(value, dict): - raise MuseProtocolError("provider message must be a JSON object") - return value - - -def _exception_close_code(exc: Exception) -> int: - code: Final = getattr(exc, "code", None) - if isinstance(exc, MuseAdapterError): - return sanitize_close_code(exc.close_code) - return sanitize_close_code(code if isinstance(code, int) else None) - - -def _ssl_config(url: str) -> object | None: - if not url.startswith("wss://"): - return None - config: Final = get_shared_realtime_ssl_context() - return True if config is False else config - - -async def _default_websocket_connect( - url: str, - *, - open_timeout: float, - max_size: int | None, - ssl: object | None, -) -> _ProviderWebSocket: - import websockets - - connection: Final = await websockets.connect( - url, - open_timeout=open_timeout, - max_size=max_size, - ssl=ssl, # pyright: ignore[reportArgumentType] # shared SSL helper returns the library-supported union - ) - return connection - - -async def _send_client_error_and_close(websocket: _ClientWebSocket, message: str) -> None: - with contextlib.suppress(Exception): - await websocket.send_text(encode_event(error_event("invalid_request_error", "invalid_configuration", message))) - await _close_client(websocket, 1008) - - -async def _close_client(websocket: _ClientWebSocket, code: int) -> None: - with contextlib.suppress(Exception): - await websocket.close(code=sanitize_close_code(code), reason=safe_close_reason(sanitize_close_code(code))) diff --git a/litellm/llms/meta/realtime/transformation.py b/litellm/llms/meta/realtime/transformation.py index f37295892f8..096f7c8f1fe 100644 --- a/litellm/llms/meta/realtime/transformation.py +++ b/litellm/llms/meta/realtime/transformation.py @@ -1,20 +1,51 @@ -from __future__ import annotations - +import asyncio +import base64 +import binascii import json import math -import uuid -from collections import OrderedDict, deque -from collections.abc import Mapping +import time +from collections import deque +from collections.abc import Awaitable, Callable, Iterator, Mapping from dataclasses import dataclass -from typing import Final, Literal, TypeAlias +from types import MappingProxyType +from typing import Final, Literal, cast +from urllib.parse import urlparse, urlunparse from pydantic import JsonValue, TypeAdapter, ValidationError -from litellm.types.realtime import RealtimeInputAudioTranscriptionUsage +from litellm import verbose_logger +from litellm._uuid import uuid +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.meta import ( + MuseAudioEncoding, + MuseHandshake, + MuseMode, + MuseSampleRate, + MuseSessionCreatedEvent, + MuseTranscriptionSession, + MuseTranscriptionSettings, + MuseTurnDetection, +) +from litellm.types.llms.openai import ( + OpenAIRealtimeEvents, + OpenAIRealtimeInputAudioBufferSpeechEvent, + OpenAIRealtimeInputAudioTranscriptionCompleted, + OpenAIRealtimeInputAudioTranscriptionDelta, +) +from litellm.types.realtime import ( + RealtimeErrorDetail, + RealtimeErrorEvent, + RealtimeInputAudioTranscriptionDurationUsage, + RealtimeInputAudioTranscriptionUsage, + RealtimeResponseTransformInput, + RealtimeResponseTypedDict, +) MUSE_MODEL: Final = "muse-voice-transcribe-1.0" +DEFAULT_MUSE_REALTIME_URL: Final = "wss://api.meta.ai/v1/asr/realtime" SUPPORTED_SAMPLE_RATES: Final = frozenset((16_000, 24_000)) -SUPPORTED_MODES: Final = frozenset(("PUSH_TO_TALK", "ENDPOINTING", "DIARIZATION")) SUPPORTED_LANGUAGES: Final = ( "Arabic", "Bengali", @@ -42,40 +73,46 @@ SUPPORTED_LANGUAGES: Final = ( "Turkish", "Vietnamese", ) -_LANGUAGE_NAMES: Final = { # mutable-ok: immutable-by-convention language lookup table - language.casefold(): language for language in SUPPORTED_LANGUAGES -} -_LANGUAGE_CODES: Final = { # mutable-ok: immutable-by-convention language lookup table - "ar": "Arabic", - "bn": "Bengali", - "de": "German", - "en": "English", - "es": "Spanish", - "fil": "Tagalog", - "fr": "French", - "he": "Hebrew", - "hi": "Hindi", - "id": "Indonesian", - "it": "Italian", - "iw": "Hebrew", - "ja": "Japanese", - "kn": "Kannada", - "ko": "Korean", - "ms": "Malay", - "mr": "Marathi", - "nl": "Dutch", - "pl": "Polish", - "pt": "Portuguese", - "ta": "Tamil", - "te": "Telugu", - "th": "Thai", - "tl": "Tagalog", - "tr": "Turkish", - "vi": "Vietnamese", - "zh": "Mandarin Chinese", -} +_LANGUAGE_NAMES: Final = MappingProxyType({language.casefold(): language for language in SUPPORTED_LANGUAGES}) +_LANGUAGE_CODES: Final = MappingProxyType( + { + "ar": "Arabic", + "bn": "Bengali", + "de": "German", + "en": "English", + "es": "Spanish", + "fil": "Tagalog", + "fr": "French", + "he": "Hebrew", + "hi": "Hindi", + "id": "Indonesian", + "it": "Italian", + "iw": "Hebrew", + "ja": "Japanese", + "kn": "Kannada", + "ko": "Korean", + "ms": "Malay", + "mr": "Marathi", + "nl": "Dutch", + "pl": "Polish", + "pt": "Portuguese", + "ta": "Tamil", + "te": "Telugu", + "th": "Thai", + "tl": "Tagalog", + "tr": "Turkish", + "vi": "Vietnamese", + "zh": "Mandarin Chinese", + } +) +_SUPPORTED_TRANSCRIPTION_KEYS: Final = frozenset(("model", "language")) +_MAX_AUDIO_BACKLOG_SECONDS: Final = 4 +_PACKET_MS: Final = 80 +_END_STREAM: Final = '{"type":"endStream"}' +_PROVIDER_ERROR_MESSAGE: Final = "Meta Muse realtime transcription failed" _JSON_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) -OpenAIEvent: TypeAlias = Mapping[str, object] +_EMPTY_OBJECT: Final[Mapping[str, JsonValue]] = MappingProxyType({}) +_SERVER_VAD: Final[MuseTurnDetection] = {"type": "server_vad"} class MuseProtocolError(ValueError): @@ -85,13 +122,12 @@ class MuseProtocolError(ValueError): @dataclass(frozen=True, slots=True) class MuseSessionConfig: model: str - mode: Literal["PUSH_TO_TALK", "ENDPOINTING", "DIARIZATION"] - sample_rate: Literal[16000, 24000] - keywords: tuple[str, ...] + mode: MuseMode + sample_rate: MuseSampleRate language_bias: tuple[str, ...] @property - def audio_encoding(self) -> Literal["PCM_16KHZ", "PCM_24KHZ"]: + def audio_encoding(self) -> MuseAudioEncoding: return "PCM_16KHZ" if self.sample_rate == 16_000 else "PCM_24KHZ" @property @@ -100,66 +136,52 @@ class MuseSessionConfig: @property def packet_bytes(self) -> int: - return self.bytes_per_second * 80 // 1000 + return self.bytes_per_second * _PACKET_MS // 1000 - def handshake(self, access_token: str) -> Mapping[str, object]: - base: Final[Mapping[str, object]] = { # mutable-ok: JSON wire payload - "mode": self.mode, - "authorization": {"accessToken": access_token}, # mutable-ok: JSON wire payload + @property + def max_encoded_append_bytes(self) -> int: + return 4 * ((self.bytes_per_second * _MAX_AUDIO_BACKLOG_SECONDS + 2) // 3) + + def handshake(self, access_token: str) -> MuseHandshake: + base: Final[MuseHandshake] = { + "authorization": {"accessToken": access_token}, "audioEncoding": self.audio_encoding, "model": self.model, + "mode": self.mode, "partialMode": "CUMULATIVE", "emitAudioProgress": True, } - payload: dict[str, object] = dict(base) # mutable-ok: incrementally builds JSON wire payload - if self.keywords: - payload["keywords"] = list(self.keywords) # mutable-ok: JSON arrays require concrete lists - if self.language_bias: - payload["languageBias"] = list(self.language_bias) # mutable-ok: JSON arrays require concrete lists - return payload + if not self.language_bias: + return base + biased: Final[MuseHandshake] = {**base, "languageBias": self.language_bias} + return biased - def openai_session(self, session_id: str) -> Mapping[str, object]: - turn_detection: Final[Mapping[str, object] | None] = ( - None if self.mode == "PUSH_TO_TALK" else {"type": "server_vad"} # mutable-ok: JSON wire payload - ) - transcription: dict[str, object] = { # mutable-ok: incrementally builds JSON wire payload - "model": self.model, - } - if self.language_bias: - transcription["language"] = self.language_bias[0] - transcription["language_bias"] = list( # mutable-ok: JSON arrays require concrete lists - self.language_bias - ) - if self.keywords: - transcription["keywords"] = list(self.keywords) # mutable-ok: JSON arrays require concrete lists - return { # mutable-ok: JSON wire payload + def openai_session(self, session_id: str) -> MuseTranscriptionSession: + session: Final[MuseTranscriptionSession] = { "id": session_id, "object": "realtime.transcription_session", "type": "transcription", - "model": self.model, - "audio": { # mutable-ok: JSON wire payload - "input": { # mutable-ok: JSON wire payload - "format": {"type": "audio/pcm", "rate": self.sample_rate}, # mutable-ok: JSON wire payload - "transcription": transcription, - "turn_detection": turn_detection, + "audio": { + "input": { + "format": {"type": "audio/pcm", "rate": self.sample_rate}, + "transcription": self._transcription_settings(), + "turn_detection": None if self.mode == "PUSH_TO_TALK" else _SERVER_VAD, } }, } + return session + + def _transcription_settings(self) -> MuseTranscriptionSettings: + base: Final[MuseTranscriptionSettings] = {"model": self.model} + if not self.language_bias: + return base + localized: Final[MuseTranscriptionSettings] = {**base, "language": self.language_bias[0]} + return localized -@dataclass(slots=True) -class _TurnState: - item_id: str | None = None - started: bool = False - start_emitted: bool = False - latest_partial: str | None = None - emitted_partial: str = "" - final_text: str | None = None - completed_signal: bool = False - completed_emitted: bool = False - stopped: bool = False - stopped_emitted: bool = False - speaker: str | None = None +_DEFAULT_SESSION_CONFIG: Final = MuseSessionConfig( + model=MUSE_MODEL, mode="ENDPOINTING", sample_rate=24_000, language_bias=() +) def _json_object(payload: str) -> Mapping[str, JsonValue]: @@ -174,7 +196,7 @@ def _json_object(payload: str) -> Mapping[str, JsonValue]: def _mapping(value: JsonValue | None, name: str) -> Mapping[str, JsonValue]: if value is None: - return {} # mutable-ok: empty JSON object + return _EMPTY_OBJECT if not isinstance(value, dict): raise MuseProtocolError(f"{name} must be an object") return value @@ -192,6 +214,10 @@ def _normalize_model(model: str) -> str: return model.removeprefix("meta/").strip() +def _event_id() -> str: + return f"event_{uuid.uuid4().hex}" + + def normalize_language(language: str) -> str: value: Final = language.strip() if not value: @@ -206,26 +232,36 @@ def normalize_language(language: str) -> str: return mapped_name -def _normalize_string_sequence(value: JsonValue | None, name: str) -> tuple[str, ...]: - if value is None: - return () - if not isinstance(value, list): - raise MuseProtocolError(f"{name} must be an array of strings") - normalized: list[str] = [] # mutable-ok: deduplicates validated language hints before freezing - for entry in value: - if not isinstance(entry, str) or not entry.strip(): - raise MuseProtocolError(f"{name} entries must be non-empty strings") - item: str = entry.strip() # rebind-ok: normalized once for each hint - if item not in normalized: - normalized.append(item) - return tuple(normalized) +def normalize_access_token(api_key: str) -> str: + stripped: Final = api_key.strip() + if not stripped: + raise ValueError("Meta API key is required") + parts: Final = stripped.split(None, 1) + if parts[0].casefold() != "bearer": + return f"Bearer {stripped}" + if len(parts) != 2 or not parts[1].strip(): + raise ValueError("Meta API key must include a token after Bearer") + return f"Bearer {parts[1].strip()}" -def _normalize_language_sequence(value: JsonValue | None) -> tuple[str, ...]: - return tuple(dict.fromkeys(normalize_language(item) for item in _normalize_string_sequence(value, "language_bias"))) +def build_muse_realtime_url(api_base: str | None) -> str: + if api_base is None: + return DEFAULT_MUSE_REALTIME_URL + parsed: Final = urlparse(api_base.strip()) + scheme: Final = "wss" if parsed.scheme == "https" else parsed.scheme + if ( + scheme != "wss" + or not parsed.hostname + or parsed.username is not None + or parsed.password is not None + or parsed.fragment + ): + raise ValueError("Meta api_base must be an absolute wss:// or https:// URL without credentials or a fragment") + netloc: Final = f"{parsed.hostname}:{parsed.port}" if parsed.port is not None else parsed.hostname + return urlunparse((scheme, netloc, "/v1/asr/realtime", "", "", "")) -def _parse_sample_rate(session: Mapping[str, JsonValue]) -> Literal[16000, 24000]: +def _parse_sample_rate(session: Mapping[str, JsonValue]) -> MuseSampleRate: beta_format: Final = session.get("input_audio_format") audio: Final = _mapping(session.get("audio"), "session.audio") audio_input: Final = _mapping(audio.get("input"), "session.audio.input") @@ -251,22 +287,10 @@ def _parse_sample_rate(session: Mapping[str, JsonValue]) -> Literal[16000, 24000 rate: Final = format_mapping.get("rate", 24_000) if isinstance(rate, bool) or not isinstance(rate, int) or rate not in SUPPORTED_SAMPLE_RATES: raise MuseProtocolError("Muse Voice supports PCM16 at 16000 Hz or 24000 Hz") - return rate + return 16_000 if rate == 16_000 else 24_000 -def _parse_mode( - session: Mapping[str, JsonValue], audio_input: Mapping[str, JsonValue] -) -> Literal["PUSH_TO_TALK", "ENDPOINTING", "DIARIZATION"]: - explicit: Final = session.get("mode") - if explicit is not None: - if not isinstance(explicit, str) or explicit.upper() not in SUPPORTED_MODES: - raise MuseProtocolError("unsupported Muse Voice mode") - normalized_mode: Final = explicit.upper() - if normalized_mode == "PUSH_TO_TALK": - return "PUSH_TO_TALK" - if normalized_mode == "DIARIZATION": - return "DIARIZATION" - return "ENDPOINTING" +def _parse_mode(session: Mapping[str, JsonValue], audio_input: Mapping[str, JsonValue]) -> MuseMode: turn_detection_present: Final = "turn_detection" in session or "turn_detection" in audio_input turn_detection: Final = session.get("turn_detection", audio_input.get("turn_detection")) if turn_detection_present and turn_detection is None: @@ -286,8 +310,7 @@ def parse_session_update(payload: str, expected_model: str) -> MuseSessionConfig session: Final = _mapping(message.get("session"), "session") if not session: raise MuseProtocolError("session.update requires a session object") - session_type: Final = session.get("type") - if session_type not in (None, "transcription", "realtime"): + if session.get("type") not in (None, "transcription", "realtime"): raise MuseProtocolError("Muse Voice supports transcription sessions only") audio: Final = _mapping(session.get("audio"), "session.audio") audio_input: Final = _mapping(audio.get("input"), "session.audio.input") @@ -299,88 +322,144 @@ def parse_session_update(payload: str, expected_model: str) -> MuseSessionConfig beta_transcription if beta_transcription is not None else ga_transcription, "input audio transcription", ) + unsupported: Final = tuple(sorted(key for key in transcription if key not in _SUPPORTED_TRANSCRIPTION_KEYS)) + if unsupported: + verbose_logger.warning("Meta realtime: dropping unsupported transcription settings %s", unsupported) requested_model: Final = _string(transcription.get("model"), "transcription model") normalized_model: Final = _normalize_model(expected_model) if normalized_model != MUSE_MODEL: raise MuseProtocolError("unsupported Meta realtime model") if requested_model is not None and _normalize_model(requested_model) != normalized_model: raise MuseProtocolError("realtime session model cannot be changed") - language_value: Final = _string(transcription.get("language"), "language") - explicit_bias: Final = _normalize_language_sequence(transcription.get("language_bias")) - language_bias: Final = tuple( - dict.fromkeys((normalize_language(language_value), *explicit_bias)) - if language_value is not None - else explicit_bias - ) - keywords: Final = _normalize_string_sequence(transcription.get("keywords"), "keywords") + language: Final = _string(transcription.get("language"), "language") return MuseSessionConfig( model=normalized_model, mode=_parse_mode(session, audio_input), sample_rate=_parse_sample_rate(session), - keywords=keywords, - language_bias=language_bias, + language_bias=() if language is None else (normalize_language(language),), ) -def session_created_event(model: str, session_id: str) -> OpenAIEvent: - normalized_model: Final = _normalize_model(model) - default_config: Final = MuseSessionConfig( - model=normalized_model, - mode="ENDPOINTING", - sample_rate=24_000, - keywords=(), - language_bias=(), - ) - return { # mutable-ok: OpenAI-compatible JSON event +def session_created_event(config: MuseSessionConfig, session_id: str) -> MuseSessionCreatedEvent: + event: Final[MuseSessionCreatedEvent] = { "type": "session.created", - "event_id": f"event_{uuid.uuid4().hex}", - "session": default_config.openai_session(session_id), - } - - -def session_updated_event(config: MuseSessionConfig, session_id: str) -> OpenAIEvent: - return { # mutable-ok: OpenAI-compatible JSON event - "type": "session.updated", - "event_id": f"event_{uuid.uuid4().hex}", + "event_id": _event_id(), "session": config.openai_session(session_id), } + return event -def error_event(error_type: str, code: str, message: str) -> OpenAIEvent: - return { # mutable-ok: OpenAI-compatible JSON event - "type": "error", - "event_id": f"event_{uuid.uuid4().hex}", - "error": { # mutable-ok: nested OpenAI-compatible error object - "type": error_type, - "code": code, - "message": message, - }, +def error_event(message: str) -> OpenAIRealtimeEvents: + detail: Final[RealtimeErrorDetail] = {"type": "server_error", "message": message} + event: Final[RealtimeErrorEvent] = {"type": "error", "error": detail} + return cast(OpenAIRealtimeEvents, event) # cast-ok: the union has no error member; the relay only serializes it + + +def _speech_event( + event_type: Literal["input_audio_buffer.speech_started", "input_audio_buffer.speech_stopped"], item_id: str +) -> OpenAIRealtimeInputAudioBufferSpeechEvent: + event: Final[OpenAIRealtimeInputAudioBufferSpeechEvent] = { + "type": event_type, + "event_id": _event_id(), + "item_id": item_id, } + return event + + +def _delta_event(item_id: str, delta: str) -> OpenAIRealtimeInputAudioTranscriptionDelta: + event: Final[OpenAIRealtimeInputAudioTranscriptionDelta] = { + "type": "conversation.item.input_audio_transcription.delta", + "event_id": _event_id(), + "item_id": item_id, + "content_index": 0, + "delta": delta, + } + return event + + +def _completed_event( + item_id: str, transcript: str, usage: RealtimeInputAudioTranscriptionUsage | None +) -> OpenAIRealtimeInputAudioTranscriptionCompleted: + event: Final[OpenAIRealtimeInputAudioTranscriptionCompleted] = { + "type": "conversation.item.input_audio_transcription.completed", + "event_id": _event_id(), + "item_id": item_id, + "content_index": 0, + "transcript": transcript, + } + if usage is None: + return event + billed: Final[OpenAIRealtimeInputAudioTranscriptionCompleted] = {**event, "usage": usage} + return billed + + +def _required_turn_id(message: Mapping[str, JsonValue], event: str) -> str: + value: Final = message.get("turnId") + if isinstance(value, bool) or not isinstance(value, (str, int)): + raise MuseProtocolError(f"{event} event has invalid turnId") + turn_id: Final = str(value).strip() + if not turn_id: + raise MuseProtocolError(f"{event} event has invalid turnId") + return turn_id + + +def _new_suffix(previous: str, current: str) -> str: + return current[len(previous) :] if current.startswith(previous) else "" + + +@dataclass(slots=True) +class _TurnState: + item_id: str + started: bool = False + start_emitted: bool = False + latest_partial: str | None = None + emitted_partial: str = "" + final_text: str | None = None + completed_signal: bool = False + completed_emitted: bool = False + stopped: bool = False + stopped_emitted: bool = False + + @property + def settled(self) -> bool: + return self.completed_emitted and (self.stopped or self.completed_signal) + + def drain( + self, take_usage: Callable[[], RealtimeInputAudioTranscriptionUsage | None] + ) -> Iterator[OpenAIRealtimeEvents]: + has_content: Final = self.latest_partial is not None or self.final_text is not None + if (self.started or has_content) and not self.start_emitted: + self.start_emitted = True + yield _speech_event("input_audio_buffer.speech_started", self.item_id) + if self.latest_partial is not None and self.final_text is None: + delta: Final = _new_suffix(self.emitted_partial, self.latest_partial) + if delta: + self.emitted_partial = self.latest_partial + yield _delta_event(self.item_id, delta) + if self.stopped and not self.stopped_emitted: + self.stopped_emitted = True + yield _speech_event("input_audio_buffer.speech_stopped", self.item_id) + if self.final_text is not None and self.stopped_emitted and not self.completed_emitted: + self.completed_emitted = True + yield _completed_event(self.item_id, self.final_text, take_usage()) class MuseEventTransformer: def __init__(self, *, completed_turn_limit: int = 128) -> None: - self._turns: OrderedDict[str, _TurnState] = OrderedDict() # mutable-ok: ordered active-turn state + self._turns: dict[str, _TurnState] = {} # mutable-ok: insertion-ordered live turn state machine + self._completed_turns: deque[str] = deque(maxlen=completed_turn_limit) # mutable-ok: bounded tombstones self._active_turn_id: str | None = None - self._mode: Literal["PUSH_TO_TALK", "ENDPOINTING", "DIARIZATION"] = "ENDPOINTING" - self._completed_turn_ids: set[str] = set() # mutable-ok: bounded completed-turn membership - self._completed_turn_order: deque[str] = deque( # mutable-ok: bounded completion eviction order - maxlen=completed_turn_limit - ) - self._completed_turn_limit: Final = completed_turn_limit - self._pending_item_ids: deque[str] = deque() # mutable-ok: FIFO commit correlation state - self._last_committed_item_id: str | None = None + self._mode: MuseMode = "ENDPOINTING" self._last_audio_processed_ms: float = 0.0 - self._unassigned_usage_seconds: float = 0.0 + self._unbilled_seconds: float = 0.0 def configure(self, config: MuseSessionConfig) -> None: self._mode = config.mode - def transform(self, payload: str) -> tuple[OpenAIEvent, ...]: - message: Final = _json_object(payload) + def transform(self, message: Mapping[str, JsonValue]) -> tuple[OpenAIRealtimeEvents, ...]: event_type: Final = message.get("type") if event_type == "error": - return (error_event("server_error", "provider_error", "Meta Muse realtime transcription failed"),) + return (error_event(_PROVIDER_ERROR_MESSAGE),) if event_type == "audioProgress": self._update_audio_progress(message) return () @@ -388,57 +467,38 @@ class MuseEventTransformer: self._speech_start(message) elif event_type == "transcript": self._transcript(message) - elif event_type == "speaker": - self._speaker(message) elif event_type == "speechEnd": self._speech_end(message) elif event_type == "speechComplete": self._speech_complete(message) else: return () - return self._drain() - - def commit_item(self) -> tuple[str | None, str]: - previous_item_id: Final = self._last_committed_item_id - provider_turn_id: Final = self._active_turn_id - active_turn: Final = self._turns.get(provider_turn_id) if provider_turn_id is not None else None - item_id: Final = ( - active_turn.item_id or provider_turn_id - if active_turn is not None and provider_turn_id is not None - else f"item_{uuid.uuid4().hex}" - ) - if active_turn is not None: - active_turn.item_id = item_id - else: - self._pending_item_ids.append(item_id) - self._last_committed_item_id = item_id - return previous_item_id, item_id + return tuple(self._drained_events()) def take_unbilled_usage(self) -> RealtimeInputAudioTranscriptionUsage | None: - seconds: Final = self._unassigned_usage_seconds + seconds: Final = self._unbilled_seconds if seconds <= 0: return None - self._unassigned_usage_seconds = 0.0 - return {"type": "duration", "seconds": seconds} # mutable-ok: typed usage wire payload + self._unbilled_seconds = 0.0 + usage: Final[RealtimeInputAudioTranscriptionDurationUsage] = {"type": "duration", "seconds": seconds} + return usage - def _turn(self, turn_id: str) -> _TurnState: - if turn_id in self._completed_turn_ids: - raise _CompletedTurn - turn: Final = self._turns.get(turn_id) - if turn is not None: - return turn - created: Final = _TurnState(item_id=self._pending_item_ids.popleft() if self._pending_item_ids else turn_id) + def _turn(self, turn_id: str) -> _TurnState | None: + if turn_id in self._completed_turns: + return None + existing: Final = self._turns.get(turn_id) + if existing is not None: + return existing + created: Final = _TurnState(item_id=turn_id) self._turns[turn_id] = created return created def _speech_start(self, message: Mapping[str, JsonValue]) -> None: - turn_id: Final = self._required_turn_id(message, "speechStart") - try: - turn: Final = self._turn(turn_id) - except _CompletedTurn: + turn: Final = self._turn(_required_turn_id(message, "speechStart")) + if turn is None: return turn.started = True - self._active_turn_id = turn_id + self._active_turn_id = turn.item_id def _transcript(self, message: Mapping[str, JsonValue]) -> None: transcript: Final = message.get("transcript") @@ -446,56 +506,34 @@ class MuseEventTransformer: raise MuseProtocolError("transcript event has invalid transcript") if not transcript and message.get("turnId") is None and self._active_turn_id is None: return - turn_id: Final = self._transcript_turn_id(message) - try: - turn: Final = self._turn(turn_id) - except _CompletedTurn: + turn: Final = self._turn(self._transcript_turn_id(message)) + if turn is None: return - final: Final = message.get("final") is True - if final: - turn.final_text = transcript - turn.completed_signal = True - if self._mode == "PUSH_TO_TALK": - turn.stopped = True - if self._active_turn_id == turn_id: - self._active_turn_id = None + if message.get("final") is not True: + if turn.final_text is None: + turn.latest_partial = transcript return - if turn.final_text is None: - turn.latest_partial = transcript - - def _speaker(self, message: Mapping[str, JsonValue]) -> None: - turn_id: Final = ( - self._required_turn_id(message, "speaker") if message.get("turnId") is not None else self._active_turn_id - ) - if turn_id is None: - raise MuseProtocolError("speaker event arrived outside an active turn") - label: Final = message.get("label") - if not isinstance(label, str) or not label.strip(): - raise MuseProtocolError("speaker event has invalid label") - try: - turn: Final = self._turn(turn_id) - except _CompletedTurn: - return - turn.speaker = label.strip() + turn.final_text = transcript + turn.completed_signal = True + if self._mode == "PUSH_TO_TALK": + turn.stopped = True + if self._active_turn_id == turn.item_id: + self._active_turn_id = None def _speech_end(self, message: Mapping[str, JsonValue]) -> None: - turn_id: Final = self._required_turn_id(message, "speechEnd") - try: - turn: Final = self._turn(turn_id) - except _CompletedTurn: + turn: Final = self._turn(_required_turn_id(message, "speechEnd")) + if turn is None: return turn.stopped = True - if self._active_turn_id == turn_id: + if self._active_turn_id == turn.item_id: self._active_turn_id = None def _speech_complete(self, message: Mapping[str, JsonValue]) -> None: - turn_id: Final = self._required_turn_id(message, "speechComplete") transcript: Final = message.get("transcript") if not isinstance(transcript, str): raise MuseProtocolError("speechComplete event has invalid transcript") - try: - turn: Final = self._turn(turn_id) - except _CompletedTurn: + turn: Final = self._turn(_required_turn_id(message, "speechComplete")) + if turn is None: return turn.final_text = transcript turn.completed_signal = True @@ -511,73 +549,21 @@ class MuseEventTransformer: raise MuseProtocolError("audioProgress event has invalid audioProcessedMs") if processed_ms <= self._last_audio_processed_ms: return - self._unassigned_usage_seconds += (float(processed_ms) - self._last_audio_processed_ms) / 1000 + self._unbilled_seconds += (float(processed_ms) - self._last_audio_processed_ms) / 1000 self._last_audio_processed_ms = float(processed_ms) - def _drain(self) -> tuple[OpenAIEvent, ...]: - events: list[OpenAIEvent] = [] # mutable-ok: ordered events are frozen to a tuple before return + def _drained_events(self) -> Iterator[OpenAIRealtimeEvents]: while self._turns: - turn_id: str = next(iter(self._turns)) # rebind-ok: selects the next ordered turn - turn: _TurnState = self._turns[turn_id] # rebind-ok: state for the selected turn - has_content: bool = ( # rebind-ok: evaluated for the selected turn - turn.latest_partial is not None or turn.final_text is not None - ) - item_id: str = turn.item_id or turn_id # rebind-ok: selected for each ordered turn - if (turn.started or has_content) and not turn.start_emitted: - turn.start_emitted = True - events.append(self._speech_event("input_audio_buffer.speech_started", item_id)) - if turn.latest_partial is not None and turn.final_text is None: - delta: str = self._new_suffix( # rebind-ok: computed for the selected turn - turn.emitted_partial, turn.latest_partial - ) - if delta: - turn.emitted_partial = turn.latest_partial - events.append( - { # mutable-ok: OpenAI-compatible JSON event - "type": "conversation.item.input_audio_transcription.delta", - "event_id": f"event_{uuid.uuid4().hex}", - "item_id": item_id, - "content_index": 0, - "delta": delta, - } - ) - if turn.stopped and not turn.stopped_emitted: - turn.stopped_emitted = True - events.append(self._speech_event("input_audio_buffer.speech_stopped", item_id)) - if turn.final_text is not None and turn.stopped_emitted and not turn.completed_emitted: - turn.completed_emitted = True - usage: RealtimeInputAudioTranscriptionUsage | None = ( # rebind-ok: usage assigned per turn - self.take_unbilled_usage() - ) - completed_event: dict[str, object] = { # mutable-ok: incrementally builds OpenAI JSON event - "type": "conversation.item.input_audio_transcription.completed", - "event_id": f"event_{uuid.uuid4().hex}", - "item_id": item_id, - "content_index": 0, - "transcript": turn.final_text, - } - if turn.speaker is not None: - completed_event["speaker"] = turn.speaker - if usage is not None: - completed_event["usage"] = usage - events.append(completed_event) - if not (turn.completed_emitted and (turn.stopped or turn.completed_signal)): - break + turn_id, turn = next(iter(self._turns.items())) + yield from turn.drain(self.take_unbilled_usage) + if not turn.settled: + return del self._turns[turn_id] - self._remember_completed(turn_id) - return tuple(events) - - def _remember_completed(self, turn_id: str) -> None: - if turn_id in self._completed_turn_ids: - return - if len(self._completed_turn_order) >= self._completed_turn_limit: - self._completed_turn_ids.discard(self._completed_turn_order.popleft()) - self._completed_turn_order.append(turn_id) - self._completed_turn_ids.add(turn_id) + self._completed_turns.append(turn_id) def _transcript_turn_id(self, message: Mapping[str, JsonValue]) -> str: if message.get("turnId") is not None: - return self._required_turn_id(message, "transcript") + return _required_turn_id(message, "transcript") if self._active_turn_id is not None: return self._active_turn_id if self._mode != "PUSH_TO_TALK": @@ -586,34 +572,162 @@ class MuseEventTransformer: self._active_turn_id = turn_id return turn_id - @staticmethod - def _required_turn_id(message: Mapping[str, JsonValue], event: str) -> str: - value: Final = message.get("turnId") - if isinstance(value, bool) or not isinstance(value, (str, int)): - raise MuseProtocolError(f"{event} event has invalid turnId") - turn_id: Final = str(value).strip() - if not turn_id: - raise MuseProtocolError(f"{event} event has invalid turnId") - return turn_id - @staticmethod - def _speech_event(event_type: str, turn_id: str) -> OpenAIEvent: - return { # mutable-ok: OpenAI-compatible JSON event - "type": event_type, - "event_id": f"event_{uuid.uuid4().hex}", - "item_id": turn_id, +class MetaRealtimeConfig(BaseRealtimeConfig): + def __init__( + self, + *, + monotonic: Callable[[], float] = time.monotonic, + sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, + ) -> None: + self._monotonic: Final = monotonic + self._sleep: Final = sleep + self._transformer: Final = MuseEventTransformer() + self._access_token: str | None = None + self._config: MuseSessionConfig | None = None + self._pending_audio: bytes = b"" + self._end_stream_sent: bool = False + self._pacing_origin: float | None = None + self._sent_duration: float = 0.0 + + def validate_environment( + self, + headers: dict[str, str], # mutable-ok: BaseRealtimeConfig contract + model: str, + api_key: str | None = None, + ) -> dict[str, str]: # mutable-ok: BaseRealtimeConfig contract + token: Final = api_key or get_secret_str("META_API_KEY") + if token is None: + raise ValueError("api_key is required for Meta API calls") + self._access_token = normalize_access_token(token) + return headers + + def get_complete_url(self, api_base: str | None, model: str, api_key: str | None = None) -> str: + if _normalize_model(model) != MUSE_MODEL: + raise ValueError(f"Unsupported Meta realtime model: {model}") + return build_muse_realtime_url(api_base) + + def is_setup_message(self, msg_obj: Mapping[str, object]) -> bool: + return "authorization" in msg_obj + + def transform_session_created_event( + self, + model: str, + logging_session_id: str, + session_configuration_request: str | None = None, + ) -> MuseSessionCreatedEvent: + return session_created_event(_DEFAULT_SESSION_CONFIG, logging_session_id) + + def transform_realtime_request( + self, + message: str, + model: str, + session_configuration_request: str | None = None, + ) -> tuple[str | bytes, ...]: + request: Final = _json_object(message) + event_type: Final = request.get("type") + if event_type in ("session.update", "transcription_session.update"): + return self._configure(message, model) + if event_type == "input_audio_buffer.append": + return self._append_audio(request) + if event_type == "input_audio_buffer.commit": + return self._flush_audio(end_stream=self._require_config().mode == "PUSH_TO_TALK") + if event_type == "input_audio_buffer.end": + return self._flush_audio(end_stream=True) + if event_type == "input_audio_buffer.clear": + self._pending_audio = b"" + return () + verbose_logger.debug("Meta realtime: dropping unsupported client event %s", event_type) + return () + + async def pace_backend_send(self, message: bytes) -> None: + now: Final = self._monotonic() + origin: Final = self._pacing_origin + effective_origin: Final = ( + now - self._sent_duration if origin is None or now > origin + self._sent_duration else origin + ) + delay: Final = effective_origin + self._sent_duration - now + if delay > 0: + await self._sleep(delay) + self._pacing_origin = effective_origin + self._sent_duration += len(message) / self._require_config().bytes_per_second + + def unbilled_usage_on_session_close(self, model: str) -> RealtimeInputAudioTranscriptionUsage | None: + return self._transformer.take_unbilled_usage() + + def transform_realtime_response( + self, + message: str | bytes, + model: str, + logging_obj: LiteLLMLoggingObj, + realtime_response_transform_input: RealtimeResponseTransformInput, + ) -> RealtimeResponseTypedDict: + payload: Final = message.decode("utf-8") if isinstance(message, bytes) else message + result: Final[RealtimeResponseTypedDict] = { + "response": list(self._backend_events(payload)), # mutable-ok: RealtimeResponseTypedDict.response is a list + "current_output_item_id": realtime_response_transform_input.get("current_output_item_id"), + "current_response_id": realtime_response_transform_input.get("current_response_id"), + "current_delta_chunks": realtime_response_transform_input.get("current_delta_chunks"), + "current_conversation_id": realtime_response_transform_input.get("current_conversation_id"), + "current_item_chunks": realtime_response_transform_input.get("current_item_chunks"), + "current_delta_type": realtime_response_transform_input.get("current_delta_type"), + "session_configuration_request": realtime_response_transform_input.get("session_configuration_request"), } + return result - @staticmethod - def _new_suffix(previous: str, current: str) -> str: - if current.startswith(previous): - return current[len(previous) :] - return "" + def _backend_events(self, payload: str) -> tuple[OpenAIRealtimeEvents, ...]: + frame: Final = _json_object(payload) + session_id: Final = frame.get("sessionId") + if session_id is None: + return self._transformer.transform(frame) + if not isinstance(session_id, str) or not session_id.strip(): + raise MuseProtocolError("provider returned an invalid handshake response") + created: Final = session_created_event(self._require_config(), session_id.strip()) + event: Final = cast(OpenAIRealtimeEvents, created) # cast-ok: ReadOnly Muse session vs writable OpenAI fields + return (event,) + def _configure(self, message: str, model: str) -> tuple[str, ...]: + if self._config is not None: + verbose_logger.debug("Meta realtime: ignoring session.update after the Muse handshake was sent") + return () + access_token: Final = self._access_token + if access_token is None: + raise MuseProtocolError("Meta API key was not validated before the session was configured") + config: Final = parse_session_update(message, model) + self._config = config + self._transformer.configure(config) + return (json.dumps(config.handshake(access_token), separators=(",", ":")),) -class _CompletedTurn(Exception): - pass + def _append_audio(self, request: Mapping[str, JsonValue]) -> tuple[bytes, ...]: + config: Final = self._require_config() + encoded: Final = request.get("audio") + if not isinstance(encoded, str): + raise MuseProtocolError("Audio must be a base64 string") + if len(encoded) > config.max_encoded_append_bytes: + raise MuseProtocolError("Audio append exceeds the four-second backlog limit") + try: + audio: Final = base64.b64decode(encoded, validate=True) + except (binascii.Error, ValueError): + raise MuseProtocolError("Audio must be valid base64") from None + if len(audio) % 2: + raise MuseProtocolError("PCM16 audio must contain complete samples") + buffered: Final = self._pending_audio + audio + packet_end: Final = len(buffered) - len(buffered) % config.packet_bytes + self._pending_audio = buffered[packet_end:] + return tuple( + buffered[start : start + config.packet_bytes] for start in range(0, packet_end, config.packet_bytes) + ) + def _flush_audio(self, *, end_stream: bool) -> tuple[str | bytes, ...]: + remainder: Final = self._pending_audio + self._pending_audio = b"" + frames: Final[tuple[bytes, ...]] = (remainder,) if remainder else () + if not end_stream or self._end_stream_sent: + return frames + self._end_stream_sent = True + return (*frames, _END_STREAM) -def encode_event(event: Mapping[str, object]) -> str: - return json.dumps(event, separators=(",", ":")) + def _require_config(self) -> MuseSessionConfig: + if self._config is None: + raise MuseProtocolError("session.update must configure the Muse session before audio is sent") + return self._config diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 4dc6768e12a..2ae66254b68 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -34718,6 +34718,7 @@ "supports_xhigh_reasoning_effort": true }, "meta/muse-voice-transcribe-1.0": { + "input_cost_per_second": 0.00005, "litellm_provider": "meta", "mode": "audio_transcription", "source": "https://dev.meta.ai/docs/speech-to-text", diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index b0803e44f6b..44c47af57f4 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -333,7 +333,7 @@ async def _resolve_vertex_access_token_bounded( @wrapper_client -async def _arealtime( # noqa: C901 # central dispatcher branches once per supported realtime provider +async def _arealtime( model: str, websocket: "WebSocket", # fastapi websocket api_base: str | None = None, @@ -391,37 +391,7 @@ async def _arealtime( # noqa: C901 # central dispatcher branches once per supp model=model, provider=LlmProviders(_custom_llm_provider), ) - if _custom_llm_provider == LlmProviders.META.value: - if model != "muse-voice-transcribe-1.0": - raise ValueError(f"Unsupported Meta realtime model: {model}") - if query_params is None or query_params.get("intent") != "transcription": - raise ValueError("Meta Muse Voice realtime requires intent=transcription") - - from litellm.llms.meta.realtime.handler import MetaRealtime - - meta_api_key: Final = get_secret_str("META_API_KEY") - dynamic_key_override: Final = dynamic_api_key if dynamic_api_key != meta_api_key else None - resolved_meta_api_key: Final = ( - api_key - or litellm_params.api_key - or dynamic_key_override - or get_secret_str("MODEL_API_KEY") - or dynamic_api_key - or meta_api_key - ) - await MetaRealtime().async_realtime( - model=model, - websocket=websocket, - logging_obj=litellm_logging_obj, - api_base=dynamic_api_base or litellm_params.api_base or api_base, - api_key=resolved_meta_api_key, - client=client, - timeout=timeout, - query_params=query_params, - user_api_key_dict=kwargs.get("user_api_key_dict"), - litellm_metadata=_build_litellm_metadata(kwargs), - ) - elif provider_config is not None: + if provider_config is not None: await base_llm_http_handler.async_realtime( model=model, websocket=websocket, diff --git a/litellm/types/llms/meta.py b/litellm/types/llms/meta.py new file mode 100644 index 00000000000..d7487bb09c8 --- /dev/null +++ b/litellm/types/llms/meta.py @@ -0,0 +1,58 @@ +from typing import Literal, TypeAlias + +from typing_extensions import NotRequired, ReadOnly, TypedDict + +MuseMode: TypeAlias = Literal["PUSH_TO_TALK", "ENDPOINTING"] +MuseAudioEncoding: TypeAlias = Literal["PCM_16KHZ", "PCM_24KHZ"] +MuseSampleRate: TypeAlias = Literal[16000, 24000] + + +class MuseAuthorization(TypedDict): + accessToken: ReadOnly[str] + + +class MuseHandshake(TypedDict): + authorization: ReadOnly[MuseAuthorization] + audioEncoding: ReadOnly[MuseAudioEncoding] + model: ReadOnly[str] + mode: ReadOnly[MuseMode] + partialMode: ReadOnly[Literal["CUMULATIVE"]] + emitAudioProgress: ReadOnly[bool] + languageBias: NotRequired[ReadOnly[tuple[str, ...]]] + + +class MuseTranscriptionAudioFormat(TypedDict): + type: ReadOnly[Literal["audio/pcm"]] + rate: ReadOnly[MuseSampleRate] + + +class MuseTranscriptionSettings(TypedDict): + model: ReadOnly[str] + language: NotRequired[ReadOnly[str]] + + +class MuseTurnDetection(TypedDict): + type: ReadOnly[Literal["server_vad"]] + + +class MuseTranscriptionAudioInput(TypedDict): + format: ReadOnly[MuseTranscriptionAudioFormat] + transcription: ReadOnly[MuseTranscriptionSettings] + turn_detection: ReadOnly[MuseTurnDetection | None] + + +class MuseTranscriptionAudio(TypedDict): + input: ReadOnly[MuseTranscriptionAudioInput] + + +class MuseTranscriptionSession(TypedDict): + id: ReadOnly[str] + object: ReadOnly[Literal["realtime.transcription_session"]] + type: ReadOnly[Literal["transcription"]] + audio: ReadOnly[MuseTranscriptionAudio] + + +class MuseSessionCreatedEvent(TypedDict): + type: ReadOnly[Literal["session.created"]] + event_id: ReadOnly[str] + session: ReadOnly[MuseTranscriptionSession] diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 02a102c9579..765ecaffdae 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -2203,7 +2203,6 @@ class OpenAIRealtimeInputAudioTranscriptionCompleted(TypedDict): content_index: ReadOnly[int] transcript: ReadOnly[str] usage: NotRequired[ReadOnly[Mapping[str, object]]] - speaker: NotRequired[ReadOnly[str]] class OpenAIRealtimeUsageTokenDetails(TypedDict): diff --git a/litellm/utils.py b/litellm/utils.py index 1a77655a5a4..394ab4b4094 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9285,6 +9285,10 @@ class ProviderConfigManager: from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig return GeminiRealtimeConfig() + if LlmProviders.META == provider: + from litellm.llms.meta.realtime.transformation import MetaRealtimeConfig + + return MetaRealtimeConfig() return None @staticmethod diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 4dc6768e12a..2ae66254b68 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -34718,6 +34718,7 @@ "supports_xhigh_reasoning_effort": true }, "meta/muse-voice-transcribe-1.0": { + "input_cost_per_second": 0.00005, "litellm_provider": "meta", "mode": "audio_transcription", "source": "https://dev.meta.ai/docs/speech-to-text", diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index e330cb103b1..295110c6bce 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -3124,7 +3124,6 @@ async def test_session_close_flush_noop_without_unbilled_usage(): ) - _UPSTREAM_REFUSAL: Final = "Publisher model `publishers/google/models/gemini-live-2.5-flash` was not found" @@ -3192,9 +3191,7 @@ def _backend_ws_closing_with(*frames: bytes | Exception) -> MagicMock: def _relay_session(client_ws: MagicMock, backend_ws: MagicMock) -> _RelaySession: logging: Final = _RecordingLogging() worker: Final = _InlineLoggingWorker() - streaming: Final = RealTimeStreaming( - client_ws, backend_ws, logging, model="gpt-realtime", logging_worker=worker - ) + streaming: Final = RealTimeStreaming(client_ws, backend_ws, logging, model="gpt-realtime", logging_worker=worker) return _RelaySession(streaming=streaming, logging=logging, worker=worker) @@ -3402,44 +3399,6 @@ async def test_refused_session_does_not_stamp_the_reservation_ownership_marker() assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in session.logging.model_call_details -@pytest.mark.asyncio -async def test_separate_usage_provider_flushes_duration_once_without_client_event(): - from typing import Final - - client_ws: Final = MagicMock() - client_ws.send_text = AsyncMock() - backend_ws: Final = MagicMock() - backend_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None)) - logging_obj: Final = MagicMock() - logging_obj.model_call_details = {} - logging_obj.dispatch_success_handlers = AsyncMock() - usage_provider: Final = MagicMock() - usage_provider.unbilled_usage_on_session_close.return_value = { - "type": "duration", - "seconds": 0.75, - } - - streaming: Final = RealTimeStreaming( - client_ws, - backend_ws, - logging_obj, - model="muse-voice-transcribe-1.0", - usage_provider=usage_provider, - ) - - await streaming.backend_to_client_send_messages() - - usage_provider.unbilled_usage_on_session_close.assert_called_once_with("muse-voice-transcribe-1.0") - duration_events: Final = tuple( - message - for message in streaming.messages - if isinstance(message, dict) and message.get("usage") == {"type": "duration", "seconds": 0.75} - ) - assert len(duration_events) == 1 - assert client_ws.send_text.await_count == 0 - logging_obj.dispatch_success_handlers.assert_called_once_with(streaming.messages, prefer_async_handlers=True) - - @pytest.mark.asyncio async def test_transformed_transcription_completion_never_sends_response_create(): from typing import Final @@ -3487,57 +3446,26 @@ async def test_transformed_transcription_completion_never_sends_response_create( backend_ws.send.assert_not_awaited() -def test_private_logging_excludes_audio_transcript_hints_and_provider_body(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.asyncio +async def test_provider_bytes_are_sent_raw_after_pacing(): from typing import Final - monkeypatch.setattr(litellm, "logged_real_time_event_types", "*") - logging_obj: Final = MagicMock() - logging_obj.model_call_details = {} + backend_ws: Final = MagicMock() + backend_ws.send = AsyncMock() + provider_config: Final = MagicMock() + provider_config.requires_session_configuration.return_value = True + provider_config.transform_realtime_request.return_value = (b"\x00\x01", '{"type":"endStream"}') + provider_config.pace_backend_send = AsyncMock() + provider_config.is_setup_message.return_value = False streaming: Final = RealTimeStreaming( MagicMock(), + backend_ws, MagicMock(), - logging_obj, + provider_config=provider_config, model="muse-voice-transcribe-1.0", - exclude_private_content_from_logs=True, - ) - audio: Final = "cHJpdmF0ZS1hdWRpbw==" - transcript: Final = "private transcript" - keyword: Final = "private keyword" - provider_body: Final = "private provider body" - - streaming.store_input( - json.dumps( - { - "type": "session.update", - "session": { - "type": "transcription", - "model": "muse-voice-transcribe-1.0", - "mode": "ENDPOINTING", - "audio": {"input": {"transcription": {"keywords": [keyword]}}}, - }, - } - ) - ) - streaming.store_input(json.dumps({"type": "input_audio_buffer.append", "audio": audio})) - streaming.store_message( - { - "type": "conversation.item.input_audio_transcription.completed", - "event_id": "event_1", - "item_id": "turn_1", - "transcript": transcript, - "provider_body": provider_body, - "usage": {"type": "duration", "seconds": 1.0}, - } ) - logged_inputs: Final = tuple(call.kwargs["input"] for call in logging_obj.pre_call.call_args_list) - serialized: Final = json.dumps({"inputs": logged_inputs, "messages": streaming.messages}) - assert audio not in serialized - assert transcript not in serialized - assert keyword not in serialized - assert provider_body not in serialized - assert "muse-voice-transcribe-1.0" in serialized - assert "ENDPOINTING" in serialized - assert "turn_1" in serialized - assert '"seconds": 1.0' in serialized - assert streaming.input_messages == [] + assert await streaming._send_to_backend(json.dumps({"type": "input_audio_buffer.commit"})) is True + + assert [call.args[0] for call in backend_ws.send.await_args_list] == [b"\x00\x01", '{"type":"endStream"}'] + provider_config.pace_backend_send.assert_awaited_once_with(b"\x00\x01") diff --git a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py deleted file mode 100644 index 14ba972fe14..00000000000 --- a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_handler.py +++ /dev/null @@ -1,452 +0,0 @@ -import asyncio -import base64 -import json -from collections.abc import Callable -from typing import Final -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest - -from litellm.llms.meta.realtime.handler import ( - DEFAULT_MUSE_REALTIME_URL, - MetaRealtime, - MuseAdapterError, - MuseRealtimeAdapter, - build_muse_realtime_url, - normalize_access_token, - safe_close_reason, - sanitize_close_code, -) -from litellm.llms.meta.realtime.transformation import MUSE_MODEL - - -class FakeProviderWebSocket: - def __init__(self, session_id: str = "provider-session") -> None: - self.sent: list[str | bytes] = [] - self.close_calls: list[tuple[int, str]] = [] - self._session_id: Final = session_id - self._recv_count = 0 - self._closed = asyncio.Event() - - async def send(self, message: str | bytes) -> None: - self.sent.append(message) - - async def recv(self, decode: bool | None = None) -> str | bytes: - self._recv_count += 1 - if self._recv_count == 1: - return json.dumps({"sessionId": self._session_id}) - await self._closed.wait() - raise MuseAdapterError("closed", close_code=1000) - - async def close(self, code: int = 1000, reason: str = "") -> None: - self.close_calls.append((code, reason)) - self._closed.set() - - -class DelayedAckWebSocket(FakeProviderWebSocket): - def __init__(self) -> None: - super().__init__() - self.ack_release = asyncio.Event() - - async def recv(self, decode: bool | None = None) -> str | bytes: - self._recv_count += 1 - if self._recv_count == 1: - await self.ack_release.wait() - return json.dumps({"sessionId": self._session_id}) - await self._closed.wait() - raise MuseAdapterError("closed", close_code=1000) - - -async def _wait_until(predicate: Callable[[], bool]) -> None: - for _ in range(100): - if predicate(): - return - await asyncio.sleep(0) - raise AssertionError("condition did not become true") - - -def _session_update(*, rate: int = 24_000, mode: str = "ENDPOINTING") -> str: - return json.dumps( - { - "type": "session.update", - "session": { - "type": "transcription", - "mode": mode, - "audio": { - "input": { - "format": {"type": "audio/pcm", "rate": rate, "channels": 1}, - "transcription": {"model": MUSE_MODEL}, - } - }, - }, - } - ) - - -async def _configured_adapter( - *, - rate: int = 24_000, - mode: str = "ENDPOINTING", - provider_ws: FakeProviderWebSocket | None = None, - monotonic: Callable[[], float] = lambda: 10.0, - sleep: Callable[[float], object] | None = None, -) -> tuple[MuseRealtimeAdapter, FakeProviderWebSocket, dict[str, object]]: - ws: Final = provider_ws or FakeProviderWebSocket() - connect_call: Final[dict[str, object]] = {} - - async def connect(url: str, **kwargs: object) -> FakeProviderWebSocket: - connect_call.update({"url": url, **kwargs}) - return ws - - async def no_sleep(_: float) -> None: - return None - - adapter: Final = MuseRealtimeAdapter( - model=f"meta/{MUSE_MODEL}", - api_key=" raw-token ", - websocket_connect=connect, - monotonic=monotonic, - sleep=sleep or no_sleep, - ) - created: Final = json.loads(await adapter.recv()) - assert created["type"] == "session.created" - await adapter.send(_session_update(rate=rate, mode=mode)) - updated: Final = json.loads(await adapter.recv()) - assert updated["type"] == "session.updated" - return adapter, ws, connect_call - - -@pytest.mark.parametrize( - ("api_key", "expected"), - [ - ("token", "Bearer token"), - (" Bearer token ", "Bearer token"), - ("bearer token", "Bearer token"), - ], -) -def test_normalize_access_token_emits_exactly_one_bearer_prefix(api_key: str, expected: str): - assert normalize_access_token(api_key) == expected - - -@pytest.mark.parametrize("api_key", ["", " ", "Bearer", " bearer "]) -def test_normalize_access_token_rejects_missing_token(api_key: str): - with pytest.raises(ValueError, match=r"token|key is required"): - normalize_access_token(api_key) - - -def test_build_muse_realtime_url_uses_fixed_secure_path(): - assert build_muse_realtime_url(None) == DEFAULT_MUSE_REALTIME_URL - assert build_muse_realtime_url("https://example.test/custom/path?ignored=yes") == ( - "wss://example.test/v1/asr/realtime" - ) - assert build_muse_realtime_url("wss://example.test:8443/other") == ("wss://example.test:8443/v1/asr/realtime") - - -@pytest.mark.parametrize( - "api_base", - [ - "http://example.test", - "ws://example.test", - "wss://user:pass@example.test", - "wss://example.test/path#fragment", - "not-a-url", - ], -) -def test_build_muse_realtime_url_rejects_insecure_or_ambiguous_overrides(api_base: str): - with pytest.raises(ValueError, match="absolute wss:// or https://"): - build_muse_realtime_url(api_base) - - -@pytest.mark.asyncio -async def test_handshake_contains_bearer_only_in_json_body_and_waits_for_ack(): - provider_ws: Final = DelayedAckWebSocket() - connect_call: Final[dict[str, object]] = {} - - async def connect(url: str, **kwargs: object) -> DelayedAckWebSocket: - connect_call.update({"url": url, **kwargs}) - return provider_ws - - adapter: Final = MuseRealtimeAdapter( - model=MUSE_MODEL, - api_key="Bearer private-token", - websocket_connect=connect, - ) - await adapter.recv() - update_task: Final = asyncio.create_task(adapter.send(_session_update())) - await _wait_until(lambda: len(provider_ws.sent) == 1) - - assert connect_call["url"] == DEFAULT_MUSE_REALTIME_URL - assert "additional_headers" not in connect_call - handshake: Final = json.loads(provider_ws.sent[0]) - assert handshake["authorization"] == {"accessToken": "Bearer private-token"} - assert handshake["audioEncoding"] == "PCM_24KHZ" - assert not update_task.done() - assert not any(isinstance(frame, bytes) for frame in provider_ws.sent) - - provider_ws.ack_release.set() - await update_task - updated: Final = json.loads(await adapter.recv()) - assert updated["type"] == "session.updated" - assert updated["session"]["id"] == "provider-session" - await adapter.close() - - -@pytest.mark.asyncio -@pytest.mark.parametrize(("rate", "packet_bytes"), [(16_000, 2_560), (24_000, 3_840)]) -async def test_audio_is_strictly_decoded_and_packetized_as_raw_pcm(rate: int, packet_bytes: int): - adapter, provider_ws, _ = await _configured_adapter(rate=rate) - pcm: Final = (b"\xff\xfe\x00\x80" * (packet_bytes // 2))[: packet_bytes * 2] - - await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(pcm).decode()})) - await _wait_until(lambda: sum(isinstance(frame, bytes) for frame in provider_ws.sent) == 2) - - binary_frames: Final = tuple(frame for frame in provider_ws.sent if isinstance(frame, bytes)) - assert binary_frames == (pcm[:packet_bytes], pcm[packet_bytes:]) - assert b"\xff\xfe" in pcm - await adapter.close() - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("audio", "expected_message"), - [ - ("not base64!", "valid base64"), - (base64.b64encode(b"\x00").decode(), "complete samples"), - ], -) -async def test_invalid_base64_or_odd_pcm_is_rejected(audio: str, expected_message: str): - adapter, provider_ws, _ = await _configured_adapter() - - await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": audio})) - error: Final = json.loads(await adapter.recv()) - - assert error["type"] == "error" - assert error["error"]["code"] == "invalid_audio" - assert expected_message in error["error"]["message"] - assert adapter.close_code == 1008 - assert not any(isinstance(frame, bytes) for frame in provider_ws.sent) - await adapter.close() - - -@pytest.mark.asyncio -async def test_absolute_pacing_delays_only_audio_ahead_of_wall_time(): - sleeps: Final[list[float]] = [] - - async def record_sleep(delay: float) -> None: - sleeps.append(delay) - - adapter, provider_ws, _ = await _configured_adapter(monotonic=lambda: 10.0, sleep=record_sleep) - pcm: Final = b"\x01\x02" * 3_840 - - await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(pcm).decode()})) - await _wait_until(lambda: sum(isinstance(frame, bytes) for frame in provider_ws.sent) == 2) - - assert sleeps == pytest.approx([0.08]) - await adapter.close() - - -@pytest.mark.asyncio -async def test_append_larger_than_four_seconds_is_rejected_without_decoding(): - adapter, provider_ws, _ = await _configured_adapter(rate=16_000) - max_pcm_bytes: Final = 16_000 * 2 * 4 - oversized_audio: Final = "A" * (4 * ((max_pcm_bytes + 2) // 3) + 1) - - with patch( # test-quality-ok: proves rejection happens before an attacker-controlled allocation - "litellm.llms.meta.realtime.handler.base64.b64decode" - ) as decode: - await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": oversized_audio})) - - error: Final = json.loads(await adapter.recv()) - assert error["error"]["code"] == "audio_backlog_exceeded" - assert adapter.close_code == 1008 - assert not any(isinstance(frame, bytes) for frame in provider_ws.sent) - decode.assert_not_called() - await adapter.close() - - -@pytest.mark.asyncio -async def test_clear_discards_only_unsent_audio(): - adapter, provider_ws, _ = await _configured_adapter() - old_pcm: Final = b"\x01\x02" * 100 - new_pcm: Final = b"\x03\x04" * 100 - - await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(old_pcm).decode()})) - await adapter.send(_event("input_audio_buffer.clear")) - cleared: Final = json.loads(await adapter.recv()) - await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(new_pcm).decode()})) - await adapter.send(_event("input_audio_buffer.commit")) - await _wait_until(lambda: any(isinstance(frame, bytes) for frame in provider_ws.sent)) - - assert cleared["type"] == "input_audio_buffer.cleared" - assert tuple(frame for frame in provider_ws.sent if isinstance(frame, bytes)) == (new_pcm,) - assert '{"type":"endStream"}' not in provider_ws.sent - await adapter.close() - - -@pytest.mark.asyncio -async def test_endpointing_commit_flushes_partial_packet_without_ending_stream(): - adapter, provider_ws, _ = await _configured_adapter(mode="ENDPOINTING") - pcm: Final = b"\x01\x02" * 100 - - await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(pcm).decode()})) - await adapter.send(_event("input_audio_buffer.commit")) - committed: Final = json.loads(await adapter.recv()) - await _wait_until(lambda: any(isinstance(frame, bytes) for frame in provider_ws.sent)) - - assert committed["type"] == "input_audio_buffer.committed" - assert committed["item_id"].startswith("item_") - assert committed["previous_item_id"] is None - assert tuple(frame for frame in provider_ws.sent if isinstance(frame, bytes)) == (pcm,) - assert '{"type":"endStream"}' not in provider_ws.sent - await adapter.close() - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("mode", "terminal_event"), - [("PUSH_TO_TALK", "input_audio_buffer.commit"), ("ENDPOINTING", "input_audio_buffer.end")], -) -async def test_commit_or_end_sends_end_stream_exactly_once(mode: str, terminal_event: str): - adapter, provider_ws, _ = await _configured_adapter(mode=mode) - pcm: Final = b"\x01\x02" * 100 - - await adapter.send(json.dumps({"type": "input_audio_buffer.append", "audio": base64.b64encode(pcm).decode()})) - await adapter.send(_event(terminal_event)) - await adapter.send(_event("input_audio_buffer.end")) - await _wait_until(lambda: '{"type":"endStream"}' in provider_ws.sent) - - assert tuple(frame for frame in provider_ws.sent if isinstance(frame, bytes)) == (pcm,) - assert provider_ws.sent.count('{"type":"endStream"}') == 1 - await adapter.close() - - -@pytest.mark.asyncio -async def test_response_create_is_returned_as_error_and_never_sent_upstream(): - adapter, provider_ws, _ = await _configured_adapter() - - await adapter.send(_event("response.create")) - error: Final = json.loads(await adapter.recv()) - - assert error["type"] == "error" - assert error["error"]["code"] == "unsupported_event" - assert not any(isinstance(frame, str) and "response.create" in frame for frame in provider_ws.sent) - await adapter.close() - - -@pytest.mark.asyncio -async def test_close_codes_and_reasons_are_sanitized_without_secret_leakage(): - adapter, provider_ws, _ = await _configured_adapter() - secret: Final = "Bearer private-token" - - await adapter.close(code=4001, reason=f"provider rejected {secret}") - - assert adapter.close_code == 1011 - assert adapter.close_reason == "Realtime transcription service error" - assert provider_ws.close_calls == [(1011, "Realtime transcription service error")] - assert secret not in json.dumps(provider_ws.close_calls) - assert sanitize_close_code(1013) == 1013 - assert safe_close_reason(1008) == "Invalid realtime transcription request" - - -@pytest.mark.asyncio -async def test_handshake_failure_reports_only_exception_type(): - secret: Final = "private-token" - - async def failing_connect(url: str, **kwargs: object) -> FakeProviderWebSocket: - raise RuntimeError(f"failed with {secret}") - - adapter: Final = MuseRealtimeAdapter( - model=MUSE_MODEL, - api_key=secret, - websocket_connect=failing_connect, - ) - await adapter.recv() - - await adapter.send(_session_update()) - error = json.loads(await adapter.recv()) - - assert error["type"] == "error" - assert error["error"]["message"] == "Meta Muse realtime handshake failed" - assert secret not in json.dumps(error) - with pytest.raises(MuseAdapterError) as exc_info: - await adapter.recv() - assert exc_info.value.close_code == 1011 - assert secret not in str(exc_info.value) - - -@pytest.mark.asyncio -async def test_meta_realtime_missing_credentials_closes_client_with_policy_code(): - client_ws: Final = MagicMock() - client_ws.send_text = AsyncMock() - client_ws.close = AsyncMock() - - await MetaRealtime().async_realtime( - model=MUSE_MODEL, - websocket=client_ws, - logging_obj=MagicMock(), - api_key=None, - ) - - client_ws.close.assert_awaited_once_with( - code=1008, - reason="Invalid realtime transcription request", - ) - sent_error: Final = json.loads(client_ws.send_text.await_args.args[0]) - assert sent_error["type"] == "error" - assert sent_error["error"]["code"] == "invalid_configuration" - - -@pytest.mark.asyncio -async def test_meta_realtime_invalid_constructor_input_sends_error_before_close(): - client_ws: Final = MagicMock() - client_ws.send_text = AsyncMock() - client_ws.close = AsyncMock() - - await MetaRealtime().async_realtime( - model=MUSE_MODEL, - websocket=client_ws, - logging_obj=MagicMock(), - api_key="Bearer", - ) - - sent_error: Final = json.loads(client_ws.send_text.await_args.args[0]) - assert sent_error["error"]["message"] == "Invalid Meta Muse realtime configuration" - client_ws.close.assert_awaited_once_with( - code=1008, - reason="Invalid realtime transcription request", - ) - - -@pytest.mark.asyncio -async def test_meta_realtime_enables_private_logging_usage_and_model_enforcement(monkeypatch: pytest.MonkeyPatch): - captured: Final[dict[str, object]] = {} - - class CapturingStreaming: - def __init__(self, websocket, backend_ws, logging_obj, **kwargs): - captured.update({"websocket": websocket, "backend_ws": backend_ws, "logging_obj": logging_obj, **kwargs}) - - async def bidirectional_forward(self) -> None: - return None - - client_ws: Final = MagicMock() - client_ws.send_text = AsyncMock() - client_ws.close = AsyncMock() - monkeypatch.setattr("litellm.llms.meta.realtime.handler.RealTimeStreaming", CapturingStreaming) - - await MetaRealtime().async_realtime( - model=MUSE_MODEL, - websocket=client_ws, - logging_obj=MagicMock(), - api_key="private-token", - ) - - adapter: Final = captured["backend_ws"] - assert isinstance(adapter, MuseRealtimeAdapter) - assert captured["force_transcription_model"] == MUSE_MODEL - assert captured["usage_provider"] is adapter - assert captured["exclude_private_content_from_logs"] is True - client_ws.close.assert_awaited_once_with(code=1000, reason="Session closed") - - -def _event(event_type: str) -> str: - return json.dumps({"type": event_type}) diff --git a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py index 8cc4af836a5..058e075860b 100644 --- a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py +++ b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py @@ -1,24 +1,70 @@ +import base64 import json +from typing import Final +from unittest.mock import MagicMock import pytest from litellm.llms.meta.realtime.transformation import ( + DEFAULT_MUSE_REALTIME_URL, MUSE_MODEL, + MetaRealtimeConfig, MuseEventTransformer, MuseProtocolError, - encode_event, + MuseSessionConfig, + build_muse_realtime_url, + normalize_access_token, normalize_language, parse_session_update, session_created_event, - session_updated_event, ) +from litellm.types.realtime import RealtimeResponseTransformInput + +EMPTY_TRANSFORM_INPUT: Final[RealtimeResponseTransformInput] = { + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_delta_chunks": None, + "current_item_chunks": None, + "current_conversation_id": None, + "current_delta_type": None, +} def _event(event_type: str, **fields: object) -> str: return json.dumps({"type": event_type, **fields}) -def test_beta_session_builds_authenticated_24khz_handshake_with_hints(): +def _ga_session_update(rate: int = 24_000, turn_detection: object = "server_vad") -> str: + return _event( + "session.update", + session={ + "type": "transcription", + "audio": { + "input": { + "format": {"type": "audio/pcm", "rate": rate}, + "turn_detection": None if turn_detection is None else {"type": turn_detection}, + "transcription": {"model": f"meta/{MUSE_MODEL}"}, + } + }, + }, + ) + + +def _configured(rate: int = 24_000, turn_detection: object = "server_vad", **kwargs: object) -> MetaRealtimeConfig: + config = MetaRealtimeConfig(**kwargs) + config.validate_environment({}, MUSE_MODEL, api_key="secret-token") + config.transform_realtime_request(_ga_session_update(rate, turn_detection), MUSE_MODEL) + return config + + +def _backend_events(config: MetaRealtimeConfig, payload: str) -> list[dict[str, object]]: + response = config.transform_realtime_response(payload, MUSE_MODEL, MagicMock(), EMPTY_TRANSFORM_INPUT)["response"] + assert isinstance(response, list) + return response + + +def test_beta_session_translates_language_and_drops_non_openai_hints(): config = parse_session_update( _event( "session.update", @@ -29,8 +75,6 @@ def test_beta_session_builds_authenticated_24khz_handshake_with_hints(): "input_audio_transcription": { "model": "meta/muse-voice-transcribe-1.0", "language": "en-US", - "language_bias": ["Spanish", "english", "French"], - "keywords": [" Muse ", "LiteLLM", "Muse"], "prompt": "must not become a keyword", }, }, @@ -41,8 +85,7 @@ def test_beta_session_builds_authenticated_24khz_handshake_with_hints(): assert config.sample_rate == 24_000 assert config.packet_bytes == 3_840 assert config.mode == "ENDPOINTING" - assert config.language_bias == ("English", "Spanish", "French") - assert config.keywords == ("Muse", "LiteLLM") + assert config.language_bias == ("English",) assert config.handshake("Bearer token") == { "mode": "ENDPOINTING", "authorization": {"accessToken": "Bearer token"}, @@ -50,8 +93,7 @@ def test_beta_session_builds_authenticated_24khz_handshake_with_hints(): "model": MUSE_MODEL, "partialMode": "CUMULATIVE", "emitAudioProgress": True, - "keywords": ["Muse", "LiteLLM"], - "languageBias": ["English", "Spanish", "French"], + "languageBias": ("English",), } assert "must not become a keyword" not in json.dumps(config.handshake("Bearer token")) @@ -79,6 +121,7 @@ def test_ga_session_accepts_16khz_mono_push_to_talk(): assert config.mode == "PUSH_TO_TALK" assert config.language_bias == ("Mandarin Chinese",) assert config.handshake("Bearer token")["audioEncoding"] == "PCM_16KHZ" + assert "languageBias" not in MuseSessionConfig(MUSE_MODEL, "ENDPOINTING", 24_000, ()).handshake("Bearer token") @pytest.mark.parametrize( @@ -109,8 +152,9 @@ def test_language_normalization_uses_official_muse_names(source: str, expected: "either beta or GA layout", ), ({"input_audio_transcription": {"model": "other-model"}}, "cannot be changed"), - ({"input_audio_transcription": {"keywords": ["valid", ""]}}, "non-empty strings"), ({"input_audio_transcription": {"language": "xx"}}, "unsupported Muse Voice language"), + ({"turn_detection": {"type": "semantic_vad"}}, "server_vad turn detection or null"), + ({"type": "realtime", "audio": {"input": {"turn_detection": {"type": "semantic_vad"}}}}, "server_vad"), ], ) def test_session_rejects_unsupported_audio_model_and_hints(session: dict[str, object], message: str): @@ -118,16 +162,15 @@ def test_session_rejects_unsupported_audio_model_and_hints(session: dict[str, ob parse_session_update(_event("session.update", session={"type": "transcription", **session}), MUSE_MODEL) -def test_session_events_expose_openai_transcription_shapes(): +def test_session_created_event_exposes_openai_transcription_shape(): config = parse_session_update( _event( "session.update", session={ - "mode": "DIARIZATION", "audio": { "input": { "format": {"type": "audio/pcm", "rate": 24000}, - "transcription": {"model": MUSE_MODEL, "language": "ja", "keywords": ["Meta"]}, + "transcription": {"model": MUSE_MODEL, "language": "ja"}, } }, }, @@ -135,31 +178,25 @@ def test_session_events_expose_openai_transcription_shapes(): MUSE_MODEL, ) - created = session_created_event(MUSE_MODEL, "session-before-handshake") - updated = session_updated_event(config, "provider-session") + created = session_created_event(config, "provider-session") assert created["type"] == "session.created" + assert created["session"]["id"] == "provider-session" assert created["session"]["type"] == "transcription" - assert updated["type"] == "session.updated" - assert updated["session"]["id"] == "provider-session" - assert updated["session"]["audio"]["input"]["transcription"] == { - "model": MUSE_MODEL, - "language": "Japanese", - "keywords": ["Meta"], - "language_bias": ["Japanese"], - } + assert created["session"]["audio"]["input"]["turn_detection"] == {"type": "server_vad"} + assert created["session"]["audio"]["input"]["transcription"] == {"model": MUSE_MODEL, "language": "Japanese"} def test_turnless_empty_silence_transcript_is_ignored(): transformer = MuseEventTransformer() - assert transformer.transform(_event("transcript", transcript="", final=True)) == () + assert transformer.transform(json.loads(_event("transcript", transcript="", final=True))) == () def test_transcript_without_speech_start_synthesizes_start_before_delta(): transformer = MuseEventTransformer() - events = transformer.transform(_event("transcript", turnId="turn-1", transcript="hello", final=False)) + events = transformer.transform(json.loads(_event("transcript", turnId="turn-1", transcript="hello", final=False))) assert [event["type"] for event in events] == [ "input_audio_buffer.speech_started", @@ -170,12 +207,15 @@ def test_transcript_without_speech_start_synthesizes_start_before_delta(): def test_cumulative_partials_emit_only_extensions_and_final_is_authoritative(): transformer = MuseEventTransformer() - started = transformer.transform(_event("speechStart", turnId="turn-1")) - first = transformer.transform(_event("transcript", turnId="turn-1", transcript="hello", final=False)) - extension = transformer.transform(_event("transcript", turnId="turn-1", transcript="hello world", final=False)) - rewrite = transformer.transform(_event("transcript", turnId="turn-1", transcript="hullo world", final=False)) - assert transformer.transform(_event("speechComplete", turnId="turn-1", transcript="hullo world")) == () - completed = transformer.transform(_event("speechEnd", turnId="turn-1")) + def send(payload: str) -> tuple[dict[str, object], ...]: + return transformer.transform(json.loads(payload)) + + started = send(_event("speechStart", turnId="turn-1")) + first = send(_event("transcript", turnId="turn-1", transcript="hello", final=False)) + extension = send(_event("transcript", turnId="turn-1", transcript="hello world", final=False)) + rewrite = send(_event("transcript", turnId="turn-1", transcript="hullo world", final=False)) + assert send(_event("speechComplete", turnId="turn-1", transcript="hullo world")) == () + completed = send(_event("speechEnd", turnId="turn-1")) assert [event["type"] for event in started] == ["input_audio_buffer.speech_started"] assert first[0]["delta"] == "hello" @@ -190,10 +230,10 @@ def test_cumulative_partials_emit_only_extensions_and_final_is_authoritative(): def test_completed_transcript_waits_for_speech_stopped(): transformer = MuseEventTransformer() - transformer.transform(_event("speechStart", turnId="turn-1")) - assert transformer.transform(_event("speechComplete", turnId="turn-1", transcript="done")) == () + transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) + assert transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="done"))) == () - released = transformer.transform(_event("speechEnd", turnId="turn-1")) + released = transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) assert [event["type"] for event in released] == [ "input_audio_buffer.speech_stopped", "conversation.item.input_audio_transcription.completed", @@ -203,11 +243,14 @@ def test_completed_transcript_waits_for_speech_stopped(): def test_overlapping_turns_are_emitted_in_provider_turn_order(): transformer = MuseEventTransformer() - transformer.transform(_event("speechStart", turnId="turn-a")) - transformer.transform(_event("speechStart", turnId="turn-b")) - assert transformer.transform(_event("transcript", turnId="turn-b", transcript="second", final=False)) == () - assert transformer.transform(_event("speechComplete", turnId="turn-a", transcript="first")) == () - released = transformer.transform(_event("speechEnd", turnId="turn-a")) + def send(payload: str) -> tuple[dict[str, object], ...]: + return transformer.transform(json.loads(payload)) + + send(_event("speechStart", turnId="turn-a")) + send(_event("speechStart", turnId="turn-b")) + assert send(_event("transcript", turnId="turn-b", transcript="second", final=False)) == () + assert send(_event("speechComplete", turnId="turn-a", transcript="first")) == () + released = send(_event("speechEnd", turnId="turn-a")) assert [(event["type"], event["item_id"]) for event in released] == [ ("input_audio_buffer.speech_stopped", "turn-a"), @@ -215,51 +258,42 @@ def test_overlapping_turns_are_emitted_in_provider_turn_order(): ("input_audio_buffer.speech_started", "turn-b"), ("conversation.item.input_audio_transcription.delta", "turn-b"), ] - assert transformer.transform(_event("speechComplete", turnId="turn-b", transcript="second final")) == () - final_b = transformer.transform(_event("speechEnd", turnId="turn-b")) + assert send(_event("speechComplete", turnId="turn-b", transcript="second final")) == () + final_b = send(_event("speechEnd", turnId="turn-b")) assert final_b[0]["type"] == "input_audio_buffer.speech_stopped" assert final_b[1]["item_id"] == "turn-b" assert final_b[1]["transcript"] == "second final" -def test_committed_item_id_is_used_for_next_provider_turn(): +def test_push_to_talk_final_transcript_completes_without_speech_end(): + transformer = MuseEventTransformer() + transformer.configure(MuseSessionConfig(MUSE_MODEL, "PUSH_TO_TALK", 24_000, ())) + + events = transformer.transform(json.loads(_event("transcript", transcript="hello there", final=True))) + + assert [event["type"] for event in events] == [ + "input_audio_buffer.speech_started", + "input_audio_buffer.speech_stopped", + "conversation.item.input_audio_transcription.completed", + ] + assert events[2]["transcript"] == "hello there" + assert str(events[0]["item_id"]).startswith("item_") + + +def test_positive_audio_progress_deltas_attach_to_next_completion_and_speaker_is_ignored(): transformer = MuseEventTransformer() - previous_item_id, item_id = transformer.commit_item() - started = transformer.transform(_event("speechStart", turnId="provider-turn")) - transformer.transform(_event("speechComplete", turnId="provider-turn", transcript="hello")) - completed = transformer.transform(_event("speechEnd", turnId="provider-turn")) + def send(payload: str) -> tuple[dict[str, object], ...]: + return transformer.transform(json.loads(payload)) - assert previous_item_id is None - assert started[0]["item_id"] == item_id - assert completed[-1]["item_id"] == item_id + send(_event("audioProgress", audioProcessedMs=1000)) + send(_event("audioProgress", audioProcessedMs=750)) + send(_event("audioProgress", audioProcessedMs=1600)) + assert send(_event("speaker", turnId=42, label=" Speaker 2 ")) == () + send(_event("speechComplete", turnId=42, transcript="hello")) + completed = send(_event("speechEnd", turnId=42)) - -def test_commit_after_speech_start_reuses_active_item_id(): - transformer = MuseEventTransformer() - - started = transformer.transform(_event("speechStart", turnId="provider-turn")) - previous_item_id, item_id = transformer.commit_item() - transformer.transform(_event("speechComplete", turnId="provider-turn", transcript="hello")) - completed = transformer.transform(_event("speechEnd", turnId="provider-turn")) - - assert previous_item_id is None - assert item_id == "provider-turn" - assert started[0]["item_id"] == item_id - assert completed[-1]["item_id"] == item_id - - -def test_speaker_and_positive_audio_progress_deltas_attach_to_next_completion(): - transformer = MuseEventTransformer() - - transformer.transform(_event("audioProgress", audioProcessedMs=1000)) - transformer.transform(_event("audioProgress", audioProcessedMs=750)) - transformer.transform(_event("audioProgress", audioProcessedMs=1600)) - transformer.transform(_event("speaker", turnId=42, label=" Speaker 2 ")) - transformer.transform(_event("speechComplete", turnId=42, transcript="hello")) - completed = transformer.transform(_event("speechEnd", turnId=42)) - - assert completed[-1]["speaker"] == "Speaker 2" + assert "speaker" not in completed[-1] assert completed[-1]["usage"] == {"type": "duration", "seconds": 1.6} assert transformer.take_unbilled_usage() is None @@ -267,7 +301,7 @@ def test_speaker_and_positive_audio_progress_deltas_attach_to_next_completion(): def test_trailing_audio_progress_is_returned_once(): transformer = MuseEventTransformer() - transformer.transform(_event("audioProgress", audioProcessedMs=250)) + transformer.transform(json.loads(_event("audioProgress", audioProcessedMs=250))) assert transformer.take_unbilled_usage() == {"type": "duration", "seconds": 0.25} assert transformer.take_unbilled_usage() is None @@ -276,24 +310,244 @@ def test_trailing_audio_progress_is_returned_once(): def test_completed_turn_tombstone_suppresses_late_duplicates(): transformer = MuseEventTransformer() - transformer.transform(_event("speechComplete", turnId="turn-1", transcript="done")) + transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="done"))) + released = transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) - assert transformer.transform(_event("speechComplete", turnId="turn-1", transcript="duplicate")) == () - assert transformer.transform(_event("speaker", turnId="turn-1", label="late")) == () + assert [event["type"] for event in released] == [ + "input_audio_buffer.speech_stopped", + "conversation.item.input_audio_transcription.completed", + ] + assert transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="duplicate"))) == () + assert transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) == () + assert ( + transformer.transform(json.loads(_event("transcript", turnId="turn-1", transcript="late", final=False))) == () + ) def test_provider_error_is_sanitized_and_encodable(): token = "private-token" provider_body = f"authorization failed for Bearer {token}" transformed = MuseEventTransformer().transform( - _event("error", code="AUTH", message=provider_body, request={"accessToken": token}) + json.loads(_event("error", code="AUTH", message=provider_body, request={"accessToken": token})) ) - encoded = encode_event(transformed[0]) + encoded = json.dumps(transformed[0]) assert json.loads(encoded)["error"] == { "type": "server_error", - "code": "provider_error", "message": "Meta Muse realtime transcription failed", } assert token not in encoded assert provider_body not in encoded + + +@pytest.mark.parametrize( + ("raw", "expected"), + [("token", "Bearer token"), (" Bearer token ", "Bearer token"), ("bearer token", "Bearer token")], +) +def test_access_token_normalization_adds_single_bearer_prefix(raw: str, expected: str): + assert normalize_access_token(raw) == expected + + +@pytest.mark.parametrize("raw", ["", " ", "Bearer", " bearer "]) +def test_access_token_normalization_rejects_empty_tokens(raw: str): + with pytest.raises(ValueError, match=r"token|key is required"): + normalize_access_token(raw) + + +@pytest.mark.parametrize( + ("api_base", "expected"), + [ + (None, DEFAULT_MUSE_REALTIME_URL), + ("https://example.test/custom/path?ignored=yes", "wss://example.test/v1/asr/realtime"), + ("wss://example.test:8443/other", "wss://example.test:8443/v1/asr/realtime"), + ], +) +def test_realtime_url_pins_muse_path(api_base: str | None, expected: str): + assert build_muse_realtime_url(api_base) == expected + assert MetaRealtimeConfig().get_complete_url(api_base, f"meta/{MUSE_MODEL}") == expected + + +@pytest.mark.parametrize( + "api_base", + [ + "http://example.test", + "ws://example.test", + "wss://user:pass@example.test", + "wss://example.test/path#fragment", + "not-a-url", + ], +) +def test_realtime_url_rejects_insecure_or_ambiguous_bases(api_base: str): + with pytest.raises(ValueError, match="absolute wss:// or https://"): + build_muse_realtime_url(api_base) + + +def test_unsupported_model_is_rejected_before_connecting(): + with pytest.raises(ValueError, match="Unsupported Meta realtime model: meta/other-model"): + MetaRealtimeConfig().get_complete_url(None, "meta/other-model") + + +def test_missing_api_key_is_rejected(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("META_API_KEY", raising=False) + + with pytest.raises(ValueError, match="api_key is required for Meta API calls"): + MetaRealtimeConfig().validate_environment({}, MUSE_MODEL) + + +def test_bearer_token_travels_only_in_the_json_handshake(): + config = MetaRealtimeConfig() + headers = {"x-existing": "kept"} + + assert config.validate_environment(headers, MUSE_MODEL, api_key="secret-token") == {"x-existing": "kept"} + (handshake,) = config.transform_realtime_request(_ga_session_update(), MUSE_MODEL) + + assert isinstance(handshake, str) + assert json.loads(handshake)["authorization"] == {"accessToken": "Bearer secret-token"} + assert config.is_setup_message(json.loads(handshake)) is True + assert config.is_setup_message({"type": "input_audio_buffer.append"}) is False + assert config.transform_realtime_request(_ga_session_update(), MUSE_MODEL) == () + + +def test_synthetic_session_created_uses_default_transcription_shape(): + created = MetaRealtimeConfig().transform_session_created_event(f"meta/{MUSE_MODEL}", "trace-1") + + assert created["type"] == "session.created" + assert created["session"]["id"] == "trace-1" + assert created["session"]["audio"]["input"]["format"] == {"type": "audio/pcm", "rate": 24000} + assert created["session"]["audio"]["input"]["transcription"] == {"model": MUSE_MODEL} + + +def test_audio_before_session_update_is_rejected(): + config = MetaRealtimeConfig() + config.validate_environment({}, MUSE_MODEL, api_key="secret-token") + + with pytest.raises(MuseProtocolError, match=r"session\.update must configure"): + config.transform_realtime_request(_event("input_audio_buffer.append", audio="AAAA"), MUSE_MODEL) + + +@pytest.mark.parametrize(("rate", "packet_bytes"), [(16_000, 2_560), (24_000, 3_840)]) +def test_pcm_is_packetized_into_raw_binary_frames(rate: int, packet_bytes: int): + config = _configured(rate=rate) + pcm = b"\xff\xfe\x00\x80" * (packet_bytes // 2) + b"\x01\x02\x03\x04" + + frames = config.transform_realtime_request( + _event("input_audio_buffer.append", audio=base64.b64encode(pcm).decode()), MUSE_MODEL + ) + remainder = config.transform_realtime_request(_event("input_audio_buffer.commit"), MUSE_MODEL) + + assert frames == (pcm[:packet_bytes], pcm[packet_bytes : packet_bytes * 2]) + assert remainder == (pcm[packet_bytes * 2 :],) + + +@pytest.mark.parametrize( + ("audio", "message"), + [ + ("not base64!", "valid base64"), + (base64.b64encode(b"\x00").decode(), "complete samples"), + (12, "base64 string"), + ("A" * (4 * ((24_000 * 2 * 4 + 2) // 3) + 4), "four-second backlog"), + ], +) +def test_invalid_audio_appends_are_rejected(audio: object, message: str): + config = _configured() + + with pytest.raises(MuseProtocolError, match=message): + config.transform_realtime_request(_event("input_audio_buffer.append", audio=audio), MUSE_MODEL) + + +@pytest.mark.asyncio +async def test_backend_sends_are_paced_to_real_time(): + sleeps: list[float] = [] + + async def record_sleep(delay: float) -> None: + sleeps.append(delay) + + config = _configured(monotonic=lambda: 10.0, sleep=record_sleep) + packet = b"\x01\x02" * 1_920 + + await config.pace_backend_send(packet) + await config.pace_backend_send(packet) + await config.pace_backend_send(packet) + + assert sleeps == pytest.approx([0.08, 0.16]) + + +def test_endpointing_commit_flushes_without_end_stream_but_end_sends_it_once(): + config = _configured(turn_detection="server_vad") + + assert config.transform_realtime_request(_event("input_audio_buffer.commit"), MUSE_MODEL) == () + assert config.transform_realtime_request(_event("input_audio_buffer.end"), MUSE_MODEL) == ('{"type":"endStream"}',) + assert config.transform_realtime_request(_event("input_audio_buffer.end"), MUSE_MODEL) == () + + +def test_push_to_talk_commit_ends_the_stream_once(): + config = _configured(turn_detection=None) + config.transform_realtime_request( + _event("input_audio_buffer.append", audio=base64.b64encode(b"\x01\x02").decode()), MUSE_MODEL + ) + + assert config.transform_realtime_request(_event("input_audio_buffer.commit"), MUSE_MODEL) == ( + b"\x01\x02", + '{"type":"endStream"}', + ) + assert config.transform_realtime_request(_event("input_audio_buffer.end"), MUSE_MODEL) == () + + +def test_clear_drops_buffered_remainder_and_unknown_events_are_ignored(): + config = _configured() + config.transform_realtime_request( + _event("input_audio_buffer.append", audio=base64.b64encode(b"\x01\x02").decode()), MUSE_MODEL + ) + + assert config.transform_realtime_request(_event("input_audio_buffer.clear"), MUSE_MODEL) == () + assert config.transform_realtime_request(_event("response.create"), MUSE_MODEL) == () + assert config.transform_realtime_request(_event("input_audio_buffer.commit"), MUSE_MODEL) == () + + +def test_provider_ack_becomes_session_created_with_provider_id(): + config = _configured(rate=16_000, turn_detection=None) + + (created,) = _backend_events(config, json.dumps({"sessionId": " provider-session "})) + + assert created["type"] == "session.created" + assert created["session"]["id"] == "provider-session" + assert created["session"]["audio"]["input"]["format"]["rate"] == 16000 + assert created["session"]["audio"]["input"]["turn_detection"] is None + + +def test_provider_turn_events_and_close_usage_flow_through_config(): + config = _configured() + + assert _backend_events(config, json.dumps({"type": "audioProgress", "audioProcessedMs": 1349})) == [] + assert _backend_events(config, _event("speechStart", turnId="t1"))[0]["type"] == "input_audio_buffer.speech_started" + assert _backend_events(config, _event("speechComplete", turnId="t1", transcript="what is the weather")) == [] + completed = _backend_events(config, _event("speechEnd", turnId="t1")) + + assert [event["type"] for event in completed] == [ + "input_audio_buffer.speech_stopped", + "conversation.item.input_audio_transcription.completed", + ] + assert completed[1]["usage"] == {"type": "duration", "seconds": 1.349} + assert config.unbilled_usage_on_session_close(MUSE_MODEL) is None + + assert _backend_events(config, json.dumps({"type": "audioProgress", "audioProcessedMs": 2349})) == [] + assert config.unbilled_usage_on_session_close(MUSE_MODEL) == {"type": "duration", "seconds": 1.0} + + +def test_provider_error_frame_becomes_openai_error_without_leaking_token(): + config = _configured() + + (error,) = _backend_events(config, _event("error", message="bad token secret-token")) + + assert error == { + "type": "error", + "error": {"type": "server_error", "message": "Meta Muse realtime transcription failed"}, + } + assert "secret-token" not in json.dumps(error) + + +def test_invalid_provider_ack_is_rejected(): + config = _configured() + + with pytest.raises(MuseProtocolError, match="invalid handshake response"): + _backend_events(config, json.dumps({"sessionId": ""})) diff --git a/tests/test_litellm/realtime_api/test_main.py b/tests/test_litellm/realtime_api/test_main.py index aa3f7d45d84..d3d41c5b54b 100644 --- a/tests/test_litellm/realtime_api/test_main.py +++ b/tests/test_litellm/realtime_api/test_main.py @@ -152,81 +152,29 @@ async def test_vertex_credential_resolution_bounds_a_thread_offloaded_refresh(): @pytest.mark.asyncio -async def test_meta_realtime_rejects_missing_transcription_intent(monkeypatch: pytest.MonkeyPatch): - def mock_get_llm_provider(model, api_base, api_key): - return model.removeprefix("meta/"), "meta", api_key, api_base +async def test_meta_realtime_dispatches_to_base_handler_with_meta_config(monkeypatch: pytest.MonkeyPatch): + from litellm.llms.meta.realtime.transformation import MetaRealtimeConfig - monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) - - with pytest.raises(ValueError, match="requires intent=transcription"): - await realtime_main._arealtime.__wrapped__( - model="meta/muse-voice-transcribe-1.0", - websocket=MagicMock(), - litellm_logging_obj=FakeLogging(), - query_params={"model": "meta/muse-voice-transcribe-1.0"}, - ) - - -@pytest.mark.asyncio -async def test_meta_realtime_rejects_unsupported_model_before_connecting(monkeypatch: pytest.MonkeyPatch): - def mock_get_llm_provider(model, api_base, api_key): - return model.removeprefix("meta/"), "meta", api_key, api_base - - monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) - - with pytest.raises(ValueError, match="Unsupported Meta realtime model: other-model"): - await realtime_main._arealtime.__wrapped__( - model="meta/other-model", - websocket=MagicMock(), - litellm_logging_obj=FakeLogging(), - query_params={"model": "meta/other-model", "intent": "transcription"}, - ) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("explicit_key", "model_key", "meta_key", "expected"), - [ - ("explicit", "model-env", "meta-env", "explicit"), - (None, "model-env", "meta-env", "model-env"), - (None, None, "meta-env", "meta-env"), - ], -) -async def test_meta_realtime_credential_precedence_is_forwarded_to_handler( - monkeypatch: pytest.MonkeyPatch, - explicit_key: str | None, - model_key: str | None, - meta_key: str | None, - expected: str, -): captured: dict[str, object] = {} def mock_get_llm_provider(model, api_base, api_key): - return model.removeprefix("meta/"), "meta", meta_key, api_base + return model.removeprefix("meta/"), "meta", None, api_base - def mock_get_secret_str(name: str): - return {"MODEL_API_KEY": model_key, "META_API_KEY": meta_key}.get(name) - - async def mock_async_realtime(self, **kwargs): + async def mock_async_realtime(**kwargs): captured.update(kwargs) monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider) - monkeypatch.setattr(realtime_main, "get_secret_str", mock_get_secret_str) - monkeypatch.setattr( - "litellm.llms.meta.realtime.handler.MetaRealtime.async_realtime", - mock_async_realtime, - ) + monkeypatch.setattr(realtime_main.base_llm_http_handler, "async_realtime", mock_async_realtime) await realtime_main._arealtime.__wrapped__( model="meta/muse-voice-transcribe-1.0", websocket=MagicMock(), litellm_logging_obj=FakeLogging(), - api_key=explicit_key, query_params={"model": "meta/muse-voice-transcribe-1.0", "intent": "transcription"}, ) + assert isinstance(captured["provider_config"], MetaRealtimeConfig) assert captured["model"] == "muse-voice-transcribe-1.0" - assert captured["api_key"] == expected assert captured["query_params"] == {"model": "muse-voice-transcribe-1.0", "intent": "transcription"} From 6b78438c9986286d9ba9e34c6a2e360a9a1f131b Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 12 Sep 2026 12:39:59 -0700 Subject: [PATCH 4/7] fix(realtime): close every Muse turn on its own terminal signal Turns no longer wait behind each other in a FIFO queue, so an empty server_vad turn (speechStart then speechEnd with no transcript) cannot stall every later turn, and a PUSH_TO_TALK speechComplete now closes its turn without waiting for a speechEnd that never arrives. Each turn keeps its own idempotent emit state, so late or duplicate speechEnd, speechComplete and transcript frames are no-ops, and finished turns are remembered in a bounded map instead of a separate tombstone deque. The session.created ack and the sanitized error frame are now typed as members of OpenAIRealtimeEvents, which removes the typing.cast calls that the strict ruff budget flagged. --- litellm/llms/meta/realtime/transformation.py | 157 ++++++++-------- litellm/types/llms/meta.py | 37 ---- litellm/types/llms/openai.py | 49 +++++ .../test_meta_realtime_transformation.py | 169 ++++++++++++++---- 4 files changed, 251 insertions(+), 161 deletions(-) diff --git a/litellm/llms/meta/realtime/transformation.py b/litellm/llms/meta/realtime/transformation.py index 096f7c8f1fe..c43e1897fbc 100644 --- a/litellm/llms/meta/realtime/transformation.py +++ b/litellm/llms/meta/realtime/transformation.py @@ -4,11 +4,10 @@ import binascii import json import math import time -from collections import deque from collections.abc import Awaitable, Callable, Iterator, Mapping from dataclasses import dataclass from types import MappingProxyType -from typing import Final, Literal, cast +from typing import Final, Literal from urllib.parse import urlparse, urlunparse from pydantic import JsonValue, TypeAdapter, ValidationError @@ -18,25 +17,19 @@ from litellm._uuid import uuid from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.secret_managers.main import get_secret_str -from litellm.types.llms.meta import ( - MuseAudioEncoding, - MuseHandshake, - MuseMode, - MuseSampleRate, - MuseSessionCreatedEvent, - MuseTranscriptionSession, - MuseTranscriptionSettings, - MuseTurnDetection, -) +from litellm.types.llms.meta import MuseAudioEncoding, MuseHandshake, MuseMode, MuseSampleRate from litellm.types.llms.openai import ( + OpenAIRealtimeErrorEvent, OpenAIRealtimeEvents, OpenAIRealtimeInputAudioBufferSpeechEvent, OpenAIRealtimeInputAudioTranscriptionCompleted, OpenAIRealtimeInputAudioTranscriptionDelta, + OpenAIRealtimeServerVadTurnDetection, + OpenAIRealtimeTranscriptionSession, + OpenAIRealtimeTranscriptionSessionCreated, + OpenAIRealtimeTranscriptionSettings, ) from litellm.types.realtime import ( - RealtimeErrorDetail, - RealtimeErrorEvent, RealtimeInputAudioTranscriptionDurationUsage, RealtimeInputAudioTranscriptionUsage, RealtimeResponseTransformInput, @@ -112,7 +105,7 @@ _END_STREAM: Final = '{"type":"endStream"}' _PROVIDER_ERROR_MESSAGE: Final = "Meta Muse realtime transcription failed" _JSON_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) _EMPTY_OBJECT: Final[Mapping[str, JsonValue]] = MappingProxyType({}) -_SERVER_VAD: Final[MuseTurnDetection] = {"type": "server_vad"} +_SERVER_VAD: Final[OpenAIRealtimeServerVadTurnDetection] = {"type": "server_vad"} class MuseProtocolError(ValueError): @@ -156,8 +149,8 @@ class MuseSessionConfig: biased: Final[MuseHandshake] = {**base, "languageBias": self.language_bias} return biased - def openai_session(self, session_id: str) -> MuseTranscriptionSession: - session: Final[MuseTranscriptionSession] = { + def openai_session(self, session_id: str) -> OpenAIRealtimeTranscriptionSession: + session: Final[OpenAIRealtimeTranscriptionSession] = { "id": session_id, "object": "realtime.transcription_session", "type": "transcription", @@ -171,11 +164,11 @@ class MuseSessionConfig: } return session - def _transcription_settings(self) -> MuseTranscriptionSettings: - base: Final[MuseTranscriptionSettings] = {"model": self.model} + def _transcription_settings(self) -> OpenAIRealtimeTranscriptionSettings: + base: Final[OpenAIRealtimeTranscriptionSettings] = {"model": self.model} if not self.language_bias: return base - localized: Final[MuseTranscriptionSettings] = {**base, "language": self.language_bias[0]} + localized: Final[OpenAIRealtimeTranscriptionSettings] = {**base, "language": self.language_bias[0]} return localized @@ -340,8 +333,8 @@ def parse_session_update(payload: str, expected_model: str) -> MuseSessionConfig ) -def session_created_event(config: MuseSessionConfig, session_id: str) -> MuseSessionCreatedEvent: - event: Final[MuseSessionCreatedEvent] = { +def session_created_event(config: MuseSessionConfig, session_id: str) -> OpenAIRealtimeTranscriptionSessionCreated: + event: Final[OpenAIRealtimeTranscriptionSessionCreated] = { "type": "session.created", "event_id": _event_id(), "session": config.openai_session(session_id), @@ -349,10 +342,12 @@ def session_created_event(config: MuseSessionConfig, session_id: str) -> MuseSes return event -def error_event(message: str) -> OpenAIRealtimeEvents: - detail: Final[RealtimeErrorDetail] = {"type": "server_error", "message": message} - event: Final[RealtimeErrorEvent] = {"type": "error", "error": detail} - return cast(OpenAIRealtimeEvents, event) # cast-ok: the union has no error member; the relay only serializes it +def error_event(message: str) -> OpenAIRealtimeErrorEvent: + event: Final[OpenAIRealtimeErrorEvent] = { + "type": "error", + "error": {"type": "server_error", "message": message}, + } + return event def _speech_event( @@ -415,14 +410,13 @@ class _TurnState: latest_partial: str | None = None emitted_partial: str = "" final_text: str | None = None - completed_signal: bool = False completed_emitted: bool = False stopped: bool = False stopped_emitted: bool = False - @property - def settled(self) -> bool: - return self.completed_emitted and (self.stopped or self.completed_signal) + def finish(self, transcript: str) -> None: + self.final_text = transcript + self.stopped = True def drain( self, take_usage: Callable[[], RealtimeInputAudioTranscriptionUsage | None] @@ -445,9 +439,9 @@ class _TurnState: class MuseEventTransformer: - def __init__(self, *, completed_turn_limit: int = 128) -> None: - self._turns: dict[str, _TurnState] = {} # mutable-ok: insertion-ordered live turn state machine - self._completed_turns: deque[str] = deque(maxlen=completed_turn_limit) # mutable-ok: bounded tombstones + def __init__(self, *, turn_limit: int = 128) -> None: + self._turns: dict[str, _TurnState] = {} # mutable-ok: bounded, insertion-ordered per-turn emit state + self._turn_limit: Final = turn_limit self._active_turn_id: str | None = None self._mode: MuseMode = "ENDPOINTING" self._last_audio_processed_ms: float = 0.0 @@ -463,17 +457,10 @@ class MuseEventTransformer: if event_type == "audioProgress": self._update_audio_progress(message) return () - if event_type == "speechStart": - self._speech_start(message) - elif event_type == "transcript": - self._transcript(message) - elif event_type == "speechEnd": - self._speech_end(message) - elif event_type == "speechComplete": - self._speech_complete(message) - else: + turn: Final = self._apply_turn_event(event_type, message) + if turn is None: return () - return tuple(self._drained_events()) + return tuple(turn.drain(self.take_unbilled_usage)) def take_unbilled_usage(self) -> RealtimeInputAudioTranscriptionUsage | None: seconds: Final = self._unbilled_seconds @@ -483,60 +470,69 @@ class MuseEventTransformer: usage: Final[RealtimeInputAudioTranscriptionDurationUsage] = {"type": "duration", "seconds": seconds} return usage - def _turn(self, turn_id: str) -> _TurnState | None: - if turn_id in self._completed_turns: - return None + def _apply_turn_event(self, event_type: JsonValue | None, message: Mapping[str, JsonValue]) -> _TurnState | None: + match event_type: + case "speechStart": + return self._speech_start(message) + case "transcript": + return self._transcript(message) + case "speechEnd": + return self._speech_end(message) + case "speechComplete": + return self._speech_complete(message) + case _: + return None + + def _turn(self, turn_id: str) -> _TurnState: existing: Final = self._turns.get(turn_id) if existing is not None: return existing created: Final = _TurnState(item_id=turn_id) self._turns[turn_id] = created + if len(self._turns) > self._turn_limit: + del self._turns[next(iter(self._turns))] return created - def _speech_start(self, message: Mapping[str, JsonValue]) -> None: + def _speech_start(self, message: Mapping[str, JsonValue]) -> _TurnState: turn: Final = self._turn(_required_turn_id(message, "speechStart")) - if turn is None: - return turn.started = True self._active_turn_id = turn.item_id + return turn - def _transcript(self, message: Mapping[str, JsonValue]) -> None: + def _transcript(self, message: Mapping[str, JsonValue]) -> _TurnState | None: transcript: Final = message.get("transcript") if not isinstance(transcript, str): raise MuseProtocolError("transcript event has invalid transcript") if not transcript and message.get("turnId") is None and self._active_turn_id is None: - return + return None turn: Final = self._turn(self._transcript_turn_id(message)) - if turn is None: - return - if message.get("final") is not True: - if turn.final_text is None: - turn.latest_partial = transcript - return - turn.final_text = transcript - turn.completed_signal = True - if self._mode == "PUSH_TO_TALK": - turn.stopped = True - if self._active_turn_id == turn.item_id: - self._active_turn_id = None + if message.get("final") is True: + self._finish(turn, transcript) + elif turn.final_text is None: + turn.latest_partial = transcript + return turn - def _speech_end(self, message: Mapping[str, JsonValue]) -> None: + def _speech_end(self, message: Mapping[str, JsonValue]) -> _TurnState: turn: Final = self._turn(_required_turn_id(message, "speechEnd")) - if turn is None: - return turn.stopped = True - if self._active_turn_id == turn.item_id: - self._active_turn_id = None + self._release_active(turn) + return turn - def _speech_complete(self, message: Mapping[str, JsonValue]) -> None: + def _speech_complete(self, message: Mapping[str, JsonValue]) -> _TurnState: transcript: Final = message.get("transcript") if not isinstance(transcript, str): raise MuseProtocolError("speechComplete event has invalid transcript") turn: Final = self._turn(_required_turn_id(message, "speechComplete")) - if turn is None: - return - turn.final_text = transcript - turn.completed_signal = True + self._finish(turn, transcript) + return turn + + def _finish(self, turn: _TurnState, transcript: str) -> None: + turn.finish(transcript) + self._release_active(turn) + + def _release_active(self, turn: _TurnState) -> None: + if self._active_turn_id == turn.item_id: + self._active_turn_id = None def _update_audio_progress(self, message: Mapping[str, JsonValue]) -> None: processed_ms: Final = message.get("audioProcessedMs") @@ -552,15 +548,6 @@ class MuseEventTransformer: self._unbilled_seconds += (float(processed_ms) - self._last_audio_processed_ms) / 1000 self._last_audio_processed_ms = float(processed_ms) - def _drained_events(self) -> Iterator[OpenAIRealtimeEvents]: - while self._turns: - turn_id, turn = next(iter(self._turns.items())) - yield from turn.drain(self.take_unbilled_usage) - if not turn.settled: - return - del self._turns[turn_id] - self._completed_turns.append(turn_id) - def _transcript_turn_id(self, message: Mapping[str, JsonValue]) -> str: if message.get("turnId") is not None: return _required_turn_id(message, "transcript") @@ -615,7 +602,7 @@ class MetaRealtimeConfig(BaseRealtimeConfig): model: str, logging_session_id: str, session_configuration_request: str | None = None, - ) -> MuseSessionCreatedEvent: + ) -> OpenAIRealtimeTranscriptionSessionCreated: return session_created_event(_DEFAULT_SESSION_CONFIG, logging_session_id) def transform_realtime_request( @@ -682,9 +669,7 @@ class MetaRealtimeConfig(BaseRealtimeConfig): return self._transformer.transform(frame) if not isinstance(session_id, str) or not session_id.strip(): raise MuseProtocolError("provider returned an invalid handshake response") - created: Final = session_created_event(self._require_config(), session_id.strip()) - event: Final = cast(OpenAIRealtimeEvents, created) # cast-ok: ReadOnly Muse session vs writable OpenAI fields - return (event,) + return (session_created_event(self._require_config(), session_id.strip()),) def _configure(self, message: str, model: str) -> tuple[str, ...]: if self._config is not None: diff --git a/litellm/types/llms/meta.py b/litellm/types/llms/meta.py index d7487bb09c8..40ecd6f7b67 100644 --- a/litellm/types/llms/meta.py +++ b/litellm/types/llms/meta.py @@ -19,40 +19,3 @@ class MuseHandshake(TypedDict): partialMode: ReadOnly[Literal["CUMULATIVE"]] emitAudioProgress: ReadOnly[bool] languageBias: NotRequired[ReadOnly[tuple[str, ...]]] - - -class MuseTranscriptionAudioFormat(TypedDict): - type: ReadOnly[Literal["audio/pcm"]] - rate: ReadOnly[MuseSampleRate] - - -class MuseTranscriptionSettings(TypedDict): - model: ReadOnly[str] - language: NotRequired[ReadOnly[str]] - - -class MuseTurnDetection(TypedDict): - type: ReadOnly[Literal["server_vad"]] - - -class MuseTranscriptionAudioInput(TypedDict): - format: ReadOnly[MuseTranscriptionAudioFormat] - transcription: ReadOnly[MuseTranscriptionSettings] - turn_detection: ReadOnly[MuseTurnDetection | None] - - -class MuseTranscriptionAudio(TypedDict): - input: ReadOnly[MuseTranscriptionAudioInput] - - -class MuseTranscriptionSession(TypedDict): - id: ReadOnly[str] - object: ReadOnly[Literal["realtime.transcription_session"]] - type: ReadOnly[Literal["transcription"]] - audio: ReadOnly[MuseTranscriptionAudio] - - -class MuseSessionCreatedEvent(TypedDict): - type: ReadOnly[Literal["session.created"]] - event_id: ReadOnly[str] - session: ReadOnly[MuseTranscriptionSession] diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 765ecaffdae..f852e125968 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -2188,6 +2188,53 @@ class OpenAIRealtimeInputAudioBufferSpeechEvent(TypedDict): item_id: ReadOnly[str] +class OpenAIRealtimeErrorDetail(TypedDict): + type: ReadOnly[str] + message: ReadOnly[str] + + +class OpenAIRealtimeErrorEvent(TypedDict): + type: ReadOnly[Literal["error"]] + error: ReadOnly[OpenAIRealtimeErrorDetail] + + +class OpenAIRealtimeTranscriptionAudioFormat(TypedDict): + type: ReadOnly[Literal["audio/pcm"]] + rate: ReadOnly[int] + + +class OpenAIRealtimeTranscriptionSettings(TypedDict): + model: ReadOnly[str] + language: NotRequired[ReadOnly[str]] + + +class OpenAIRealtimeServerVadTurnDetection(TypedDict): + type: ReadOnly[Literal["server_vad"]] + + +class OpenAIRealtimeTranscriptionAudioInput(TypedDict): + format: ReadOnly[OpenAIRealtimeTranscriptionAudioFormat] + transcription: ReadOnly[OpenAIRealtimeTranscriptionSettings] + turn_detection: ReadOnly[OpenAIRealtimeServerVadTurnDetection | None] + + +class OpenAIRealtimeTranscriptionAudio(TypedDict): + input: ReadOnly[OpenAIRealtimeTranscriptionAudioInput] + + +class OpenAIRealtimeTranscriptionSession(TypedDict): + id: ReadOnly[str] + object: ReadOnly[Literal["realtime.transcription_session"]] + type: ReadOnly[Literal["transcription"]] + audio: ReadOnly[OpenAIRealtimeTranscriptionAudio] + + +class OpenAIRealtimeTranscriptionSessionCreated(TypedDict): + type: ReadOnly[Literal["session.created"]] + event_id: ReadOnly[str] + session: ReadOnly[OpenAIRealtimeTranscriptionSession] + + class OpenAIRealtimeInputAudioTranscriptionDelta(TypedDict): type: ReadOnly[Literal["conversation.item.input_audio_transcription.delta"]] event_id: ReadOnly[str] @@ -2259,6 +2306,8 @@ OpenAIRealtimeEvents = ( | OpenAIRealtimeInputAudioBufferSpeechEvent | OpenAIRealtimeInputAudioTranscriptionDelta | OpenAIRealtimeInputAudioTranscriptionCompleted + | OpenAIRealtimeTranscriptionSessionCreated + | OpenAIRealtimeErrorEvent ) OpenAIRealtimeStreamList = list[OpenAIRealtimeEvents] diff --git a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py index 058e075860b..8b9eaf12dc2 100644 --- a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py +++ b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py @@ -1,4 +1,5 @@ import base64 +import itertools import json from typing import Final from unittest.mock import MagicMock @@ -18,6 +19,7 @@ from litellm.llms.meta.realtime.transformation import ( parse_session_update, session_created_event, ) +from litellm.types.llms.meta import MuseMode from litellm.types.realtime import RealtimeResponseTransformInput EMPTY_TRANSFORM_INPUT: Final[RealtimeResponseTransformInput] = { @@ -214,8 +216,7 @@ def test_cumulative_partials_emit_only_extensions_and_final_is_authoritative(): first = send(_event("transcript", turnId="turn-1", transcript="hello", final=False)) extension = send(_event("transcript", turnId="turn-1", transcript="hello world", final=False)) rewrite = send(_event("transcript", turnId="turn-1", transcript="hullo world", final=False)) - assert send(_event("speechComplete", turnId="turn-1", transcript="hullo world")) == () - completed = send(_event("speechEnd", turnId="turn-1")) + completed = send(_event("speechComplete", turnId="turn-1", transcript="hullo world")) assert [event["type"] for event in started] == ["input_audio_buffer.speech_started"] assert first[0]["delta"] == "hello" @@ -225,44 +226,125 @@ def test_cumulative_partials_emit_only_extensions_and_final_is_authoritative(): assert completed[1]["type"] == "conversation.item.input_audio_transcription.completed" assert completed[1]["item_id"] == "turn-1" assert completed[1]["transcript"] == "hullo world" + assert send(_event("speechEnd", turnId="turn-1")) == () -def test_completed_transcript_waits_for_speech_stopped(): +def test_speech_end_then_speech_complete_emits_stopped_then_completed(): transformer = MuseEventTransformer() transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) - assert transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="done"))) == () + stopped = transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) + completed = transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="done"))) - released = transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) - assert [event["type"] for event in released] == [ + assert [event["type"] for event in stopped] == ["input_audio_buffer.speech_stopped"] + assert [event["type"] for event in completed] == ["conversation.item.input_audio_transcription.completed"] + assert completed[0]["transcript"] == "done" + + +def _typed(events: tuple[dict[str, object], ...]) -> list[tuple[object, object]]: + return [(event["type"], event["item_id"]) for event in events] + + +def test_overlapping_turns_emit_independently_and_correlate_by_item_id(): + transformer = MuseEventTransformer() + + def send(payload: str) -> list[tuple[object, object]]: + return _typed(transformer.transform(json.loads(payload))) + + assert send(_event("speechStart", turnId="turn-a")) == [("input_audio_buffer.speech_started", "turn-a")] + assert send(_event("speechStart", turnId="turn-b")) == [("input_audio_buffer.speech_started", "turn-b")] + assert send(_event("transcript", turnId="turn-b", transcript="second", final=False)) == [ + ("conversation.item.input_audio_transcription.delta", "turn-b") + ] + assert send(_event("speechComplete", turnId="turn-a", transcript="first")) == [ + ("input_audio_buffer.speech_stopped", "turn-a"), + ("conversation.item.input_audio_transcription.completed", "turn-a"), + ] + assert send(_event("speechEnd", turnId="turn-a")) == [] + assert send(_event("speechEnd", turnId="turn-b")) == [("input_audio_buffer.speech_stopped", "turn-b")] + assert send(_event("speechComplete", turnId="turn-b", transcript="second final")) == [ + ("conversation.item.input_audio_transcription.completed", "turn-b") + ] + + +def test_empty_vad_turn_is_closed_and_does_not_block_the_next_turn(): + transformer = MuseEventTransformer() + + def send(payload: str) -> list[tuple[object, object]]: + return _typed(transformer.transform(json.loads(payload))) + + assert send(_event("speechStart", turnId="noise")) == [("input_audio_buffer.speech_started", "noise")] + assert send(_event("speechEnd", turnId="noise")) == [("input_audio_buffer.speech_stopped", "noise")] + assert send(_event("speechStart", turnId="speech")) == [("input_audio_buffer.speech_started", "speech")] + assert send(_event("transcript", turnId="speech", transcript="hello", final=False)) == [ + ("conversation.item.input_audio_transcription.delta", "speech") + ] + assert send(_event("speechEnd", turnId="speech")) == [("input_audio_buffer.speech_stopped", "speech")] + assert send(_event("speechComplete", turnId="speech", transcript="hello world")) == [ + ("conversation.item.input_audio_transcription.completed", "speech") + ] + + +@pytest.mark.parametrize("transcript", ["", "late words"]) +def test_late_speech_complete_after_an_empty_speech_end_completes_that_item(transcript: str): + transformer = MuseEventTransformer() + transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) + transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) + transformer.transform(json.loads(_event("speechStart", turnId="turn-2"))) + + (completed,) = transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript=transcript))) + + assert completed["type"] == "conversation.item.input_audio_transcription.completed" + assert completed["item_id"] == "turn-1" + assert completed["transcript"] == transcript + + +def test_push_to_talk_speech_complete_closes_the_turn_without_speech_end(): + transformer = MuseEventTransformer() + transformer.configure(MuseSessionConfig(MUSE_MODEL, "PUSH_TO_TALK", 24_000, ())) + + transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) + transformer.transform(json.loads(_event("transcript", turnId="turn-1", transcript="hel", final=False))) + events = transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="hello"))) + + assert [event["type"] for event in events] == [ "input_audio_buffer.speech_stopped", "conversation.item.input_audio_transcription.completed", ] + assert events[1]["transcript"] == "hello" -def test_overlapping_turns_are_emitted_in_provider_turn_order(): +_TERMINAL_SIGNALS: Final = { + "speechEnd": _event("speechEnd", turnId="turn-1"), + "speechComplete": _event("speechComplete", turnId="turn-1", transcript="final words"), + "final": _event("transcript", turnId="turn-1", transcript="final words", final=True), +} +_TERMINAL_ORDERINGS: Final = tuple( + ordering for size in (1, 2, 3) for ordering in itertools.permutations(_TERMINAL_SIGNALS, size) +) + + +@pytest.mark.parametrize("mode", ["ENDPOINTING", "PUSH_TO_TALK"]) +@pytest.mark.parametrize("ordering", _TERMINAL_ORDERINGS, ids="-".join) +def test_every_terminal_signal_order_closes_the_turn_exactly_once(mode: MuseMode, ordering: tuple[str, ...]): transformer = MuseEventTransformer() + transformer.configure(MuseSessionConfig(MUSE_MODEL, mode, 24_000, ())) + transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) + transformer.transform(json.loads(_event("transcript", turnId="turn-1", transcript="fin", final=False))) - def send(payload: str) -> tuple[dict[str, object], ...]: - return transformer.transform(json.loads(payload)) - - send(_event("speechStart", turnId="turn-a")) - send(_event("speechStart", turnId="turn-b")) - assert send(_event("transcript", turnId="turn-b", transcript="second", final=False)) == () - assert send(_event("speechComplete", turnId="turn-a", transcript="first")) == () - released = send(_event("speechEnd", turnId="turn-a")) - - assert [(event["type"], event["item_id"]) for event in released] == [ - ("input_audio_buffer.speech_stopped", "turn-a"), - ("conversation.item.input_audio_transcription.completed", "turn-a"), - ("input_audio_buffer.speech_started", "turn-b"), - ("conversation.item.input_audio_transcription.delta", "turn-b"), + emitted = [ + event["type"] for signal in ordering for event in transformer.transform(json.loads(_TERMINAL_SIGNALS[signal])) ] - assert send(_event("speechComplete", turnId="turn-b", transcript="second final")) == () - final_b = send(_event("speechEnd", turnId="turn-b")) - assert final_b[0]["type"] == "input_audio_buffer.speech_stopped" - assert final_b[1]["item_id"] == "turn-b" - assert final_b[1]["transcript"] == "second final" + replayed = [ + event["type"] for signal in ordering for event in transformer.transform(json.loads(_TERMINAL_SIGNALS[signal])) + ] + + has_text = bool(set(ordering) & {"speechComplete", "final"}) + assert emitted == [ + "input_audio_buffer.speech_stopped", + *(["conversation.item.input_audio_transcription.completed"] if has_text else []), + ] + assert replayed == [] def test_push_to_talk_final_transcript_completes_without_speech_end(): @@ -290,12 +372,12 @@ def test_positive_audio_progress_deltas_attach_to_next_completion_and_speaker_is send(_event("audioProgress", audioProcessedMs=750)) send(_event("audioProgress", audioProcessedMs=1600)) assert send(_event("speaker", turnId=42, label=" Speaker 2 ")) == () - send(_event("speechComplete", turnId=42, transcript="hello")) - completed = send(_event("speechEnd", turnId=42)) + completed = send(_event("speechComplete", turnId=42, transcript="hello")) assert "speaker" not in completed[-1] assert completed[-1]["usage"] == {"type": "duration", "seconds": 1.6} assert transformer.take_unbilled_usage() is None + assert send(_event("speechEnd", turnId=42)) == () def test_trailing_audio_progress_is_returned_once(): @@ -307,23 +389,37 @@ def test_trailing_audio_progress_is_returned_once(): assert transformer.take_unbilled_usage() is None -def test_completed_turn_tombstone_suppresses_late_duplicates(): +def test_finished_turn_ignores_late_duplicates(): transformer = MuseEventTransformer() - transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="done"))) - released = transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) + released = transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="done"))) assert [event["type"] for event in released] == [ + "input_audio_buffer.speech_started", "input_audio_buffer.speech_stopped", "conversation.item.input_audio_transcription.completed", ] assert transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="duplicate"))) == () assert transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) == () + assert transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) == () assert ( transformer.transform(json.loads(_event("transcript", turnId="turn-1", transcript="late", final=False))) == () ) +def test_turn_memory_is_bounded_by_turn_limit(): + transformer = MuseEventTransformer(turn_limit=2) + + transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="one"))) + transformer.transform(json.loads(_event("speechComplete", turnId="turn-2", transcript="two"))) + assert transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) == () + transformer.transform(json.loads(_event("speechComplete", turnId="turn-3", transcript="three"))) + + forgotten = transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) + + assert [event["type"] for event in forgotten] == ["input_audio_buffer.speech_stopped"] + + def test_provider_error_is_sanitized_and_encodable(): token = "private-token" provider_body = f"authorization failed for Bearer {token}" @@ -520,14 +616,11 @@ def test_provider_turn_events_and_close_usage_flow_through_config(): assert _backend_events(config, json.dumps({"type": "audioProgress", "audioProcessedMs": 1349})) == [] assert _backend_events(config, _event("speechStart", turnId="t1"))[0]["type"] == "input_audio_buffer.speech_started" - assert _backend_events(config, _event("speechComplete", turnId="t1", transcript="what is the weather")) == [] - completed = _backend_events(config, _event("speechEnd", turnId="t1")) + assert _backend_events(config, _event("speechEnd", turnId="t1"))[0]["type"] == "input_audio_buffer.speech_stopped" + completed = _backend_events(config, _event("speechComplete", turnId="t1", transcript="what is the weather")) - assert [event["type"] for event in completed] == [ - "input_audio_buffer.speech_stopped", - "conversation.item.input_audio_transcription.completed", - ] - assert completed[1]["usage"] == {"type": "duration", "seconds": 1.349} + assert [event["type"] for event in completed] == ["conversation.item.input_audio_transcription.completed"] + assert completed[0]["usage"] == {"type": "duration", "seconds": 1.349} assert config.unbilled_usage_on_session_close(MUSE_MODEL) is None assert _backend_events(config, json.dumps({"type": "audioProgress", "audioProcessedMs": 2349})) == [] From a2e383a1a582c8d05170915ad63828046dd9f823 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 12 Sep 2026 12:53:03 -0700 Subject: [PATCH 5/7] fix(realtime): ignore a late speechStart for a finished Muse turn A duplicate speechStart for a turn that already stopped used to make that closed turn active again, so the next turnless PUSH_TO_TALK transcript was routed to the finished item and dropped. --- litellm/llms/meta/realtime/transformation.py | 2 ++ .../test_meta_realtime_transformation.py | 18 ++++++++++++++++++ 2 files changed, 20 insertions(+) diff --git a/litellm/llms/meta/realtime/transformation.py b/litellm/llms/meta/realtime/transformation.py index c43e1897fbc..549db60be85 100644 --- a/litellm/llms/meta/realtime/transformation.py +++ b/litellm/llms/meta/realtime/transformation.py @@ -495,6 +495,8 @@ class MuseEventTransformer: def _speech_start(self, message: Mapping[str, JsonValue]) -> _TurnState: turn: Final = self._turn(_required_turn_id(message, "speechStart")) + if turn.stopped: + return turn turn.started = True self._active_turn_id = turn.item_id return turn diff --git a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py index 8b9eaf12dc2..3a4c0591fb5 100644 --- a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py +++ b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py @@ -407,6 +407,24 @@ def test_finished_turn_ignores_late_duplicates(): ) +def test_late_duplicate_speech_start_does_not_capture_the_next_turnless_transcript(): + transformer = MuseEventTransformer() + transformer.configure(MuseSessionConfig(MUSE_MODEL, "PUSH_TO_TALK", 24_000, ())) + transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) + transformer.transform(json.loads(_event("speechComplete", turnId="turn-1", transcript="first"))) + + assert transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) == () + events = transformer.transform(json.loads(_event("transcript", transcript="second", final=True))) + + assert [event["type"] for event in events] == [ + "input_audio_buffer.speech_started", + "input_audio_buffer.speech_stopped", + "conversation.item.input_audio_transcription.completed", + ] + assert events[2]["transcript"] == "second" + assert events[2]["item_id"] != "turn-1" + + def test_turn_memory_is_bounded_by_turn_limit(): transformer = MuseEventTransformer(turn_limit=2) From 4647cd121593b46c36d3502d36ee7be3e1070358 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 12 Sep 2026 15:13:44 -0700 Subject: [PATCH 6/7] fix(realtime): keep a Muse turn active for turnless partials after speechEnd Muse partials carry no turnId and belong to the most recent speechStart, and the docs say the model may keep post processing a turn after speechEnd until speechComplete. Releasing the active turn on speechEnd made any partial arriving in that window raise and get dropped in ENDPOINTING mode. The turn now stays active until its speechComplete or final transcript. --- litellm/llms/meta/realtime/transformation.py | 1 - .../test_meta_realtime_transformation.py | 19 +++++++++++++++++++ 2 files changed, 19 insertions(+), 1 deletion(-) diff --git a/litellm/llms/meta/realtime/transformation.py b/litellm/llms/meta/realtime/transformation.py index 549db60be85..1b8943f0cee 100644 --- a/litellm/llms/meta/realtime/transformation.py +++ b/litellm/llms/meta/realtime/transformation.py @@ -517,7 +517,6 @@ class MuseEventTransformer: def _speech_end(self, message: Mapping[str, JsonValue]) -> _TurnState: turn: Final = self._turn(_required_turn_id(message, "speechEnd")) turn.stopped = True - self._release_active(turn) return turn def _speech_complete(self, message: Mapping[str, JsonValue]) -> _TurnState: diff --git a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py index 3a4c0591fb5..a5d7e47fb65 100644 --- a/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py +++ b/tests/test_litellm/llms/meta/realtime/test_meta_realtime_transformation.py @@ -241,6 +241,25 @@ def test_speech_end_then_speech_complete_emits_stopped_then_completed(): assert completed[0]["transcript"] == "done" +def test_turnless_partial_between_speech_end_and_speech_complete_stays_on_that_turn(): + transformer = MuseEventTransformer() + transformer.transform(json.loads(_event("speechStart", turnId="turn-1"))) + transformer.transform(json.loads(_event("transcript", transcript="what is", final=False))) + transformer.transform(json.loads(_event("speechEnd", turnId="turn-1"))) + + post_processed = transformer.transform( + json.loads(_event("transcript", transcript="what is the weather", final=False)) + ) + completed = transformer.transform( + json.loads(_event("speechComplete", turnId="turn-1", transcript="What is the weather?")) + ) + + assert _typed(post_processed) == [("conversation.item.input_audio_transcription.delta", "turn-1")] + assert post_processed[0]["delta"] == " the weather" + assert _typed(completed) == [("conversation.item.input_audio_transcription.completed", "turn-1")] + assert completed[0]["transcript"] == "What is the weather?" + + def _typed(events: tuple[dict[str, object], ...]) -> list[tuple[object, object]]: return [(event["type"], event["item_id"]) for event in events] From 0e435e41486e860a5c8a998dcfd54ce3e23e84ab Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 12 Sep 2026 15:32:45 -0700 Subject: [PATCH 7/7] fix(realtime): run transcription guardrails on transcription-only sessions The provider_config path skipped run_realtime_guardrails for transcription sessions to avoid sending response.create, which also dropped every realtime_input_transcription guardrail: no violation error reached the client and on_violation / end_session_after_n_fails never fired. Run the guardrail for every completed transcript and only suppress response.create when the session has no assistant turn. --- .../litellm_core_utils/realtime_streaming.py | 4 +- .../test_realtime_streaming.py | 61 +++++++++++++++++++ 2 files changed, 62 insertions(+), 3 deletions(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 06d9241b826..e3f8786a39a 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1015,13 +1015,11 @@ class RealTimeStreaming: self.store_message(event_str) self._capture_transcription_usage(event) await self._send_event_to_client(event, event_str) - if self._is_transcription_session: - continue blocked = await self.run_realtime_guardrails( cast(str, transcript), item_id=cast(str | None, event.get("item_id")), ) - if not blocked: + if not blocked and not self._is_transcription_session: await self._send_to_backend(json.dumps({"type": "response.create"})) continue ## LOGGING diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 295110c6bce..2a33d84ec78 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -3446,6 +3446,67 @@ async def test_transformed_transcription_completion_never_sends_response_create( backend_ws.send.assert_not_awaited() +@pytest.mark.asyncio +async def test_transcription_session_still_runs_transcription_guardrail(monkeypatch: pytest.MonkeyPatch): + class BlockingGuardrail(CustomGuardrail): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + raise ValueError("blocked transcript") + + guardrail: Final = BlockingGuardrail( + guardrail_name="transcription-blocker", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + ) + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + + completed_event: Final = { + "type": "conversation.item.input_audio_transcription.completed", + "event_id": "event_1", + "item_id": "turn_1", + "content_index": 0, + "transcript": "blocked transcript", + "usage": {"type": "duration", "seconds": 0.5}, + } + provider_config: Final = MagicMock() + provider_config.requires_session_configuration.return_value = True + provider_config.transform_realtime_response.return_value = { + "response": completed_event, + "current_output_item_id": None, + "current_response_id": None, + "current_delta_chunks": None, + "current_conversation_id": None, + "current_item_chunks": None, + "current_delta_type": None, + "session_configuration_request": None, + } + provider_config.transform_realtime_request.return_value = () + provider_config.is_setup_message.return_value = False + provider_config.is_content_message.return_value = False + client_ws: Final = MagicMock() + client_ws.send_text = AsyncMock() + backend_ws: Final = MagicMock() + backend_ws.send = AsyncMock() + + streaming: Final = RealTimeStreaming( + client_ws, + backend_ws, + MagicMock(), + provider_config=provider_config, + model="muse-voice-transcribe-1.0", + force_transcription_model="muse-voice-transcribe-1.0", + ) + + await streaming._handle_provider_config_message("{}") + + sent_to_client: Final = [json.loads(call.args[0]) for call in client_ws.send_text.await_args_list] + assert completed_event in sent_to_client + error_events: Final = [event for event in sent_to_client if event.get("type") == "error"] + assert len(error_events) == 1 + assert error_events[0]["error"]["type"] == "guardrail_violation" + backend_ws.send.assert_not_awaited() + assert streaming._violation_count == 1 + + @pytest.mark.asyncio async def test_provider_bytes_are_sent_raw_after_pacing(): from typing import Final