Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 78 additions & 9 deletions litellm/litellm_core_utils/realtime_streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
OpenAIRealtimeStreamResponseBaseObject,
OpenAIRealtimeStreamSessionEvents,
)
from litellm.types.realtime import ALL_DELTA_TYPES
from litellm.types.realtime import ALL_DELTA_TYPES, RealtimeInputAudioTranscriptionUsage

from .litellm_logging import Logging as LiteLLMLogging

Expand Down Expand Up @@ -86,6 +86,10 @@ def normalize(self, event: dict) -> dict: ...
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",
Expand All @@ -108,6 +112,8 @@ def __init__(
backend_uses_beta_protocol: bool | None = None,
force_transcription_model: str | None = None,
event_normalizer: RealtimeEventNormalizer | None = None,
usage_provider: RealtimeUsageProvider | None = None,
exclude_private_content_from_logs: bool = False,
):
self.websocket: _ClientWebSocket = websocket
self.backend_ws = backend_ws
Expand Down Expand Up @@ -167,6 +173,10 @@ def __init__(
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
Expand Down Expand Up @@ -204,7 +214,7 @@ def __init__(

def _should_store_message(
self,
message_obj: dict | OpenAIRealtimeEvents,
message_obj: dict[str, Any] | OpenAIRealtimeEvents, # mutable-ok: existing realtime event contract
) -> bool:
_msg_type: Final = message_obj["type"] if "type" in message_obj else None
if self.logged_real_time_event_types == "*":
Expand All @@ -213,16 +223,54 @@ def _should_store_message(
return True
return False

def _message_for_logging(
self,
message_obj: dict[str, Any], # mutable-ok: existing realtime event contract
) -> dict[str, Any]: # mutable-ok: logging stores concrete event dictionaries
if not self._exclude_private_content_from_logs:
return message_obj
logged_message: dict[str, Any] = { # mutable-ok: incrementally builds the sanitized event copy
key: message_obj[key]
for key in (
"type",
"event_id",
"item_id",
"response_id",
"conversation_id",
"session_id",
"content_index",
"output_index",
"model",
"mode",
"usage",
)
if key in message_obj
}
session: Final = message_obj.get("session")
if isinstance(session, dict):
logged_session: Final[dict[str, Any]] = { # mutable-ok: sanitized JSON session snapshot
key: session[key] for key in ("id", "model", "mode", "type") if key in session
}
if logged_session:
logged_message["session"] = logged_session
return logged_message

def store_message(self, message: str | bytes | dict | OpenAIRealtimeEvents):
"""Store message in list"""
if isinstance(message, bytes):
message = message.decode("utf-8")
if isinstance(message, dict):
# TypedDict union members do not narrow to plain dict for mypy.
message_obj: dict[str, Any] = cast(dict[str, Any], message)
parsed_message_obj: dict[str, Any] = cast( # cast-ok: TypedDict events are JSON dictionaries
dict[str, Any], message
)
else:
message_obj = cast(dict[str, Any], json.loads(cast(str, message)))
self._collect_tool_calls_from_response_done(cast(dict, message_obj))
parsed_message_obj = cast( # cast-ok: parsed realtime events are JSON dictionaries
dict[str, Any], json.loads(message)
)
if not self._exclude_private_content_from_logs:
self._collect_tool_calls_from_response_done(parsed_message_obj)
message_obj: Final = self._message_for_logging(parsed_message_obj)
if not self._should_store_message(message_obj):
return
try:
Expand All @@ -240,6 +288,8 @@ def store_message(self, message: str | bytes | dict | OpenAIRealtimeEvents):

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)
Expand Down Expand Up @@ -276,6 +326,8 @@ def _collect_user_input_from_client_event(self, message: str | dict) -> None:

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":
Expand Down Expand Up @@ -331,9 +383,9 @@ def _capture_transcription_usage(self, event_obj: dict | OpenAIRealtimeEvents) -
pass

def _flush_unbilled_transcription_usage(self) -> None:
if self.provider_config is None:
if self._usage_provider is None:
return
usage: Final = self.provider_config.unbilled_usage_on_session_close(self.model)
usage: Final = self._usage_provider.unbilled_usage_on_session_close(self.model)
if usage is None:
return
flush_event: Final = (
Expand Down Expand Up @@ -370,12 +422,27 @@ def _collect_tool_calls_from_response_done(self, event_obj: dict | OpenAIRealtim
except (AttributeError, TypeError):
pass

def _input_for_logging(
self,
message: str | dict, # mutable-ok: existing realtime input contract
) -> str | dict: # mutable-ok: logging stores concrete event dictionaries
if not self._exclude_private_content_from_logs:
return message
try:
parsed_message: Final[object] = message if isinstance(message, dict) else json.loads(message)
except (json.JSONDecodeError, TypeError):
return {} # mutable-ok: empty JSON logging payload
if not isinstance(parsed_message, dict):
return {} # mutable-ok: empty JSON logging payload
return self._message_for_logging(parsed_message)

def store_input(self, message: str | dict):
"""Store input message"""
self.input_message = message if isinstance(message, dict) else {}
logged_message: Final[str | dict] = self._input_for_logging(message) # mutable-ok: logging payload
self.input_message = logged_message if isinstance(logged_message, dict) else {}
self._collect_user_input_from_client_event(message)
if self.logging_obj:
self.logging_obj.pre_call(input=message, api_key="")
self.logging_obj.pre_call(input=logged_message, api_key="")

async def log_messages(self):
"""Log messages in list"""
Expand Down Expand Up @@ -975,6 +1042,8 @@ async def _handle_provider_config_message(self, raw_response: str) -> None:
self.store_message(event_str)
self._capture_transcription_usage(event)
await self._send_event_to_client(event, event_str)
if self._is_transcription_session:
continue
blocked = await self.run_realtime_guardrails(
cast(str, transcript),
item_id=cast(str | None, event.get("item_id")),
Expand Down
3 changes: 3 additions & 0 deletions litellm/llms/meta/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
from .realtime import MetaRealtime, MuseRealtimeAdapter

__all__ = ("MetaRealtime", "MuseRealtimeAdapter")
10 changes: 10 additions & 0 deletions litellm/llms/meta/realtime/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
from .handler import MetaRealtime, MuseRealtimeAdapter
from .transformation import MuseEventTransformer, MuseProtocolError, MuseSessionConfig

__all__ = (
"MetaRealtime",
"MuseEventTransformer",
"MuseProtocolError",
"MuseRealtimeAdapter",
"MuseSessionConfig",
)
Loading
Loading