mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
refactor(realtime): move Meta Muse Voice onto BaseRealtimeConfig
Replace the hand-rolled Meta realtime handler with a MetaRealtimeConfig that plugs into the shared realtime handler and RealTimeStreaming relay. Clients keep speaking the OpenAI realtime wire: session.update, input_audio_buffer.append/commit and the OpenAI transcription events. Unsupported transcription settings are logged and dropped, matching the Gemini realtime precedent, and the Meta-specific session.mode, keywords, language_bias, DIARIZATION and speaker extensions are removed. Drop the MODEL_API_KEY env var in favor of the standard META_API_KEY, remove the private-logging flag so spend logs record the transcript the same way other realtime models do, and add per-second pricing for muse-voice-transcribe-1.0. The relay now sends raw bytes from transform_realtime_request straight to the backend after pace_backend_send, and transcription sessions never trigger response.create.
This commit is contained in:
parent
1acb994998
commit
17fde7a261
16 changed files with 889 additions and 1803 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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`.
|
||||
|
||||
|
|
|
|||
|
|
@ -1,3 +0,0 @@
|
|||
from .realtime import MetaRealtime, MuseRealtimeAdapter
|
||||
|
||||
__all__ = ("MetaRealtime", "MuseRealtimeAdapter")
|
||||
|
|
@ -1,10 +0,0 @@
|
|||
from .handler import MetaRealtime, MuseRealtimeAdapter
|
||||
from .transformation import MuseEventTransformer, MuseProtocolError, MuseSessionConfig
|
||||
|
||||
__all__ = (
|
||||
"MetaRealtime",
|
||||
"MuseEventTransformer",
|
||||
"MuseProtocolError",
|
||||
"MuseRealtimeAdapter",
|
||||
"MuseSessionConfig",
|
||||
)
|
||||
|
|
@ -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)))
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
58
litellm/types/llms/meta.py
Normal file
58
litellm/types/llms/meta.py
Normal file
|
|
@ -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]
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
@ -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": ""}))
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue