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