diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index ca7208458d8..e2ad3c13337 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1,7 +1,7 @@ import asyncio import concurrent.futures import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Protocol, Union, cast import litellm from litellm._logging import verbose_logger @@ -27,6 +27,13 @@ else: # Create a thread pool with a maximum of 10 threads executor = concurrent.futures.ThreadPoolExecutor(max_workers=10) + +class RealtimeEventNormalizer(Protocol): + def should_drop(self, event: object) -> bool: ... + def normalize(self, event: dict) -> dict: ... + def patch_outgoing_session(self, session: dict) -> dict: ... + + DefaultLoggedRealTimeEventTypes = [ "session.created", "response.create", @@ -48,6 +55,7 @@ class RealTimeStreaming: request_data: Optional[Dict] = None, backend_uses_beta_protocol: Optional[bool] = None, force_transcription_model: Optional[str] = None, + event_normalizer: Optional[RealtimeEventNormalizer] = None, ): self.websocket = websocket self.backend_ws = backend_ws @@ -101,11 +109,16 @@ class RealTimeStreaming: self._flushing_pending_messages_until_setup: bool = False self._pending_messages_until_setup: List[str] = [] self._pending_messages_byte_total: int = 0 + # Gemini Live rejects a follow-up BidiGenerateContentSetup once any + # content (realtimeInput / clientContent / toolResponse) has been sent. + self._content_sent_after_setup: bool = False # Whether this is a transcription-only session (session.type == "transcription", # e.g. gpt-realtime-whisper). Such sessions must not be sent response.create and # their input_audio_transcription.completed usage drives duration-based cost. self._force_transcription_model = force_transcription_model 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 # Per-connection caps for pre-setup audio frames (message count + total bytes). _MAX_BUFFERED_MESSAGES: int = 200 @@ -353,15 +366,36 @@ class RealTimeStreaming: ) sent = False for msg in transformed: - # Send first; only cache the setup payload once the backend - # has actually accepted it. Caching before send would leave - # ``session_configuration_request`` populated after a failed - # send, causing subsequent client session.update messages to - # be treated as "subsequent" and dropped even though the - # backend never received the original setup. - await self.backend_ws.send(msg) # type: ignore[union-attr, attr-defined] - self._cache_session_configuration_request(msg) - sent = True + try: + msg_obj = json.loads(msg) + except (json.JSONDecodeError, TypeError): + msg_obj = None + if isinstance(msg_obj, dict) and self.provider_config.is_setup_message( + msg_obj + ): + if self._content_sent_after_setup: + verbose_logger.debug( + "Dropping follow-up setup after content was already sent to backend" + ) + continue + await self.backend_ws.send(msg) # type: ignore[union-attr, attr-defined] + self._cache_session_configuration_request(msg) + sent = True + else: + is_content_message = isinstance( + msg_obj, dict + ) and self.provider_config.is_content_message(msg_obj) + # Send first, then mutate state, so a failed send leaves both + # ``session_configuration_request`` and + # ``_content_sent_after_setup`` untouched. Caching or marking + # content before send would leave the session believing the + # backend received a setup/content frame it never got, causing + # subsequent client session.update messages to be dropped. + await self.backend_ws.send(msg) # type: ignore[union-attr, attr-defined] + self._cache_session_configuration_request(msg) + if is_content_message: + self._content_sent_after_setup = True + sent = True return sent await self.backend_ws.send(message) # type: ignore[union-attr, attr-defined] return True @@ -557,7 +591,27 @@ class RealTimeStreaming: return False return True + def _should_drop_event_from_client(self, event: object) -> bool: + """Return True for provider-specific events that must not reach GA clients.""" + if self._event_normalizer is not None: + return self._event_normalizer.should_drop(event) + return False + + def _normalize_event_for_ga_client(self, event: dict) -> dict: + """Apply per-provider GA normalization before forwarding to clients.""" + if self._event_normalizer is not None: + return self._event_normalizer.normalize(event) + return event + + def _event_to_client_json(self, event: dict) -> str: + return json.dumps(self._normalize_event_for_ga_client(event)) + async def _send_event_to_client(self, event: Any, event_str: str) -> bool: + if self._should_drop_event_from_client(event): + return False + if isinstance(event, dict): + event = self._normalize_event_for_ga_client(event) + event_str = json.dumps(event) if self._client_wants_beta and isinstance(event, dict): try: translated = self._translate_event_to_beta(event) @@ -845,6 +899,8 @@ class RealTimeStreaming: else [transformed_response] ) for event in events: + if self._should_drop_event_from_client(event): + continue is_session_created_event = ( isinstance(event, dict) and event.get("type") == "session.created" ) @@ -936,7 +992,7 @@ class RealTimeStreaming: and self._has_audio_transcription_guardrails() ): self.store_message(event_obj) - await self.websocket.send_text(raw_response) + await self.websocket.send_text(self._event_to_client_json(event_obj)) await self._send_to_backend(self._make_disable_auto_response_message()) return True @@ -944,7 +1000,7 @@ class RealTimeStreaming: transcript = event_obj.get("transcript", "") self._collect_user_input_from_backend_event(event_obj) self.store_message(event_obj) - await self.websocket.send_text(raw_response) + await self.websocket.send_text(self._event_to_client_json(event_obj)) # Transcription-only sessions (e.g. gpt-realtime-whisper) have no # assistant turn: capture audio-duration usage for cost and never @@ -997,20 +1053,23 @@ class RealTimeStreaming: await self.websocket.send_text(raw_response) continue + if self._should_drop_event_from_client(event): + continue + if await self._handle_raw_backend_message(event, raw_response): continue + + event = self._normalize_event_for_ga_client(event) self.store_message(event) if not self._client_wants_beta: - await self.websocket.send_text(raw_response) + await self.websocket.send_text(json.dumps(event)) continue translated = self._translate_event_to_beta(event) if translated is None: continue - await self.websocket.send_text( - raw_response if translated is event else json.dumps(translated) - ) + await self.websocket.send_text(json.dumps(translated)) except websockets.exceptions.ConnectionClosed as e: # type: ignore verbose_logger.exception( @@ -1142,8 +1201,7 @@ class RealTimeStreaming: Returns None when the event must be dropped (the GA-only conversation.item.done has no beta counterpart). Returns the original - event object unchanged when no translation applies, so the caller can - forward the raw frame without re-serializing; otherwise returns a + event object unchanged when no translation applies; otherwise returns a translated copy. """ event_type = event.get("type", "") @@ -1404,6 +1462,14 @@ class RealTimeStreaming: msg_obj["session"] = session message = json.dumps(msg_obj) + if msg_type == "session.update" and self._event_normalizer: + session = msg_obj.get("session") + if isinstance(session, dict): + msg_obj["session"] = ( + self._event_normalizer.patch_outgoing_session(session) + ) + message = json.dumps(msg_obj) + except (json.JSONDecodeError, AttributeError): pass diff --git a/litellm/llms/base_llm/realtime/transformation.py b/litellm/llms/base_llm/realtime/transformation.py index 0f239b4ad45..b66bdbd2dd9 100644 --- a/litellm/llms/base_llm/realtime/transformation.py +++ b/litellm/llms/base_llm/realtime/transformation.py @@ -60,6 +60,12 @@ class BaseRealtimeConfig(ABC): ) -> List[str]: pass + def is_setup_message(self, msg_obj: dict) -> bool: + return False + + def is_content_message(self, msg_obj: dict) -> bool: + return False + def requires_session_configuration( self, ) -> bool: # initial configuration message sent to setup the realtime session diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index a8cb40ac6db..9c0a4a30efb 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -70,39 +70,30 @@ MAP_GEMINI_FIELD_TO_OPENAI_EVENT: Dict[ "toolCall": ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, } -# Top-level keys in a Gemini realtime message that map_openai_event knows how -# to handle. Other keys (e.g. ``usageMetadata``) can appear alongside these as -# siblings and must be skipped by the main transform loop — otherwise -# map_openai_event raises ``ValueError`` and the WebSocket session terminates. +# Keys the main transform loop handles; siblings like ``usageMetadata`` are skipped. _KNOWN_GEMINI_TOP_LEVEL_KEYS: set = { map_key.split(".", 1)[0] for map_key in MAP_GEMINI_FIELD_TO_OPENAI_EVENT } -# Gemini Live native-audio model ids carry this marker (e.g. -# ``gemini-2.5-flash-native-audio-preview-09-2025``). These models reject a -# ``speechConfig`` on ``setup`` with a 1007 invalid-argument error, so it is -# stripped in ``_finalize_gemini_live_setup``. -_GEMINI_NATIVE_AUDIO_MODEL_MARKER = "native-audio" - class GeminiRealtimeConfig(BaseRealtimeConfig): - # Cap the LRU of in-flight tool calls so long sessions with many tool - # calls don't grow the dict without bound. Sized large enough to cover - # bursts of pending tool responses; the oldest entry is evicted when a - # new call beyond the cap arrives. - _TOOL_CALL_ID_TO_NAME_MAX = 256 + _TOOL_CALL_ID_TO_NAME_MAX = 256 # LRU cap for call_id→name mapping def __init__(self): super().__init__() - # Store call_id → function_name mapping for tool call round-trip self._tool_call_id_to_name: "OrderedDict[str, str]" = OrderedDict() - # Buffer ``usageMetadata`` that Gemini Live emits as a standalone - # frame (between turns) so the next ``response.done`` attributes the - # tokens consumed. Without this an authenticated client can drive - # tool-call or normal turns whose token usage is recorded as zero, - # bypassing spend and budget accounting. + # Gemini Live sometimes emits usageMetadata in a standalone frame between + # turns; buffer it here so the next response.done carries the token counts. self._pending_usage_metadata: Optional[dict] = None + def is_setup_message(self, msg_obj: dict) -> bool: + return "setup" in msg_obj + + def is_content_message(self, msg_obj: dict) -> bool: + return any( + k in msg_obj for k in ("realtimeInput", "clientContent", "toolResponse") + ) + def _include_function_response_id(self) -> bool: """Google AI Studio Gemini 3.5+ accepts ``id`` on functionResponses; Vertex AI rejects it.""" return True @@ -415,16 +406,57 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return normalized + @staticmethod + def _model_cost_entry(model: str) -> dict: + entry = litellm.model_cost.get(model) + if entry is None: + stripped = model.split("/", 1)[-1] + entry = litellm.model_cost.get(stripped) or litellm.model_cost.get( + f"gemini/{stripped}" + ) + return entry or {} + + @staticmethod + def _is_audio_only_live_model(model: str) -> bool: + entry = GeminiRealtimeConfig._model_cost_entry(model) + return bool( + entry.get("gemini_native_audio") or entry.get("gemini_audio_only_live") + ) + + @staticmethod + def _is_native_audio_model(model: str) -> bool: + return bool( + GeminiRealtimeConfig._model_cost_entry(model).get("gemini_native_audio") + ) + + @staticmethod + def _coerce_response_modalities(model: str, modalities: list[Any]) -> list[str]: + """Map unsupported TEXT responseModalities to AUDIO for audio-only Live models.""" + normalized = [ + modality.upper() if isinstance(modality, str) else str(modality).upper() + for modality in modalities + ] + if not GeminiRealtimeConfig._is_audio_only_live_model(model): + return normalized + if "TEXT" not in normalized: + return normalized + without_text = [modality for modality in normalized if modality != "TEXT"] + return without_text if without_text else ["AUDIO"] + @staticmethod def _finalize_gemini_live_setup( model: str, setup: Dict[str, Any] ) -> Dict[str, Any]: """Drop fields Gemini Live native-audio rejects on ``setup``.""" - if _GEMINI_NATIVE_AUDIO_MODEL_MARKER not in model.lower(): - return setup generation_config = setup.get("generationConfig") if isinstance(generation_config, dict): - generation_config.pop("speechConfig", None) + modalities = generation_config.get("responseModalities") + if isinstance(modalities, list): + generation_config["responseModalities"] = ( + GeminiRealtimeConfig._coerce_response_modalities(model, modalities) + ) + if GeminiRealtimeConfig._is_native_audio_model(model): + generation_config.pop("speechConfig", None) return setup def _handle_session_update( @@ -540,18 +572,22 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): BidiGenerateContentRealtimeInputConfig, merged_realtime_input_config, ) + finalized_follow_up = self._finalize_gemini_live_setup( + model, cast(dict[str, Any], follow_up_setup) + ) + # Skip if the follow-up setup is identical to the one already sent. + # The final session.update from Pipecat's _create_response (after history + # items) matches the pre-history session.update we intentionally sent + # before content; sending a duplicate at that point would risk a 1007. + if finalized_follow_up == original_setup: + verbose_logger.debug( + "Gemini Realtime: Skipping duplicate follow-up session.update (no changes)" + ) + return [] verbose_logger.debug( "Gemini Realtime: Forwarding session.update as follow-up setup" ) - return [ - json.dumps( - { - "setup": self._finalize_gemini_live_setup( - model, cast(Dict[str, Any], follow_up_setup) - ) - } - ) - ] + return [json.dumps({"setup": finalized_follow_up})] def _handle_conversation_item(self, json_message: dict) -> List[str]: """ @@ -563,11 +599,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): item = json_message.get("item", {}) item_type = item.get("type") - # Handle function call output (tool response) if item_type == "function_call_output": return self._handle_function_call_output(item) - - # Handle regular text content return self._handle_user_text_content(item) def _handle_function_call_output(self, item: dict) -> List[str]: @@ -579,10 +612,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): f"Gemini Realtime: Transforming function_call_output for call_id={call_id}" ) - # Parse the output to get the result. Gemini's - # functionResponses[].response field is a Struct, so it must be a - # dict; wrap any non-dict (primitives, lists, invalid JSON) under a - # `result` key. + # Gemini functionResponses[].response must be a dict; wrap non-dicts. try: parsed_output = json.loads(output) if isinstance(output, str) else output except json.JSONDecodeError: @@ -593,11 +623,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): else {"result": parsed_output} ) - # Look up the function name from stored mapping. Keep the entry so a - # client SDK that retries function_call_output (or sends it twice for - # the same tool call) still produces a Gemini toolResponse with the - # required ``name`` field; refresh the LRU position so an active - # call_id stays warm across long sessions. + # Keep the entry (don't delete) so retried tool responses still find the name. function_name = self._tool_call_id_to_name.get(call_id) if function_name: self._tool_call_id_to_name.move_to_end(call_id) @@ -607,7 +633,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): "This may cause Gemini to reject the response." ) - # Build Gemini toolResponse format function_response: dict[str, Any] = {"response": output_dict} if self._include_function_response_id() and call_id: function_response["id"] = call_id @@ -632,7 +657,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): if not text: return [] - # Build clientContent message with turns (proper Gemini Live API format) client_content_message = { "clientContent": { "turns": [{"role": "user", "parts": [{"text": text}]}], @@ -661,21 +685,17 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): messages: List[str] = [] msg_type = json_message.get("type") - ## HANDLE SESSION UPDATE — translate to Gemini setup ## if msg_type == "session.update": return self._handle_session_update( json_message, model, session_configuration_request ) - ## HANDLE response.create — Gemini responds automatically; nothing to forward ## if msg_type == "response.create": - return [] + return [] # Gemini responds automatically; nothing to forward - ## HANDLE conversation.item.create — extract user text or function call output ## if msg_type == "conversation.item.create": return self._handle_conversation_item(json_message) - ## HANDLE INPUT AUDIO BUFFER - use realtimeInput for audio streaming ## if msg_type == "input_audio_buffer.append": realtime_input_dict["audio"] = HttpxBlobType( mimeType=self.get_audio_mime_type(), data=json_message["audio"] @@ -701,15 +721,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) if msg_type == "input_audio_buffer.clear": - # Local OpenAI buffer op — nothing to forward to Gemini Live. - verbose_logger.debug( - "Gemini Realtime: input_audio_buffer.clear is a local buffer op" - ) - return [] + return [] # local buffer op, nothing to forward - # Unknown/unsupported OpenAI event type — drop silently rather than - # forwarding raw JSON as text input to the model. - return [] + return [] # unknown/unsupported event type def transform_session_created_event( self, @@ -742,11 +756,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): if _system_instruction is not None and isinstance(_system_instruction, str): session["instructions"] = _system_instruction if _model is not None and isinstance(_model, str): - # Normalise to bare model name for OpenAI compatibility. - # Vertex AI uses a full resource path: - # projects/{project}/locations/{location}/publishers/google/models/{model} - # Google AI Studio uses: - # models/{model} + # Strip Vertex/AI Studio path prefixes to expose the bare model name. if "/models/" in _model: session["model"] = _model.split("/models/")[-1] elif _model.startswith("models/"): @@ -800,8 +810,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): _max_output_tokens = generation_config.get("maxOutputTokens") response_items: List[OpenAIRealtimeEvents] = [] - - ## - return response.created response_created = OpenAIRealtimeStreamResponseBaseObject( type="response.created", event_id="event_{}".format(uuid.uuid4()), @@ -1015,28 +1023,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return returned_items def _consume_usage_metadata_for_response_done(self, frame: dict) -> Optional[dict]: - """Return the ``usageMetadata`` to attribute to a ``response.done``. + """Pop usageMetadata from the frame (authoritative) or drain the pending buffer. - Gemini Live emits ``usageMetadata`` either alongside the closing - frame (``serverContent.turnComplete`` / ``toolCall``) or as a - standalone frame between turns. The standalone form would otherwise - be discarded by the no-op branch in ``transform_realtime_response`` - and the consumed tokens silently dropped from spend/budget - accounting. ``_pending_usage_metadata`` buffers any such standalone - frames so the next emitted ``response.done`` carries the deferred - token counts. - - Returns the in-frame ``usageMetadata`` if present (and clears the - buffer since the in-frame counts are the authoritative attribution - for this turn), otherwise returns the buffered counts. ``None`` is - returned when neither is available so the caller can fall back to - ``get_empty_usage()``. + Uses pop so a frame with both ``toolCall`` and ``turnComplete`` can't + attribute the same counts to two response.done events. """ - # ``pop`` (rather than ``get``) so a single Gemini frame containing - # multiple closing keys (e.g. both ``toolCall`` and - # ``serverContent.turnComplete``) cannot attribute the same - # ``usageMetadata`` to two ``response.done`` events and double-count - # tokens in spend/budget accounting. in_frame = frame.pop("usageMetadata", None) if isinstance(frame, dict) else None if isinstance(in_frame, dict): self._pending_usage_metadata = None @@ -1051,12 +1042,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): response_id: Optional[str] = None, output_item_id: Optional[str] = None, ) -> List[OpenAIRealtimeFunctionCallArgumentsDone]: - """ - Transform Gemini toolCall message to OpenAI function call events. - - Converts Gemini's functionCalls format to OpenAI's response.function_call_arguments.done events. - Also stores call_id → name mapping for later use in function_call_output responses. - """ function_calls = tool_call_message.get("functionCalls", []) resolved_response_id = response_id or f"resp_{uuid.uuid4()}" resolved_output_item_id = output_item_id or f"item_{uuid.uuid4()}" @@ -1070,9 +1055,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): call_id = fc.get("id", "") or f"call_{uuid.uuid4().hex[:16]}" name = fc.get("name", "") - # Store call_id → name mapping for round-trip. Use an LRU so - # repeated function_call_output lookups (retries) still hit, while - # sessions with many tool calls don't grow the dict unboundedly. if call_id and name: self._tool_call_id_to_name[call_id] = name self._tool_call_id_to_name.move_to_end(call_id) @@ -1121,13 +1103,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) any_delta_chunk = True if not any_delta_chunk: - current_delta_chunks = ( - None # reset current_delta_chunks if no delta chunks - ) + current_delta_chunks = None else: if ( transformed_message["type"] == "response.output_text.delta" - ): # ONLY ACCUMULATE TEXT DELTA CHUNKS - AUDIO WILL CAUSE SERVER MEMORY ISSUES + ): # audio deltas are not accumulated (memory) if current_delta_chunks is None: current_delta_chunks = [] current_delta_chunks.append( @@ -1157,9 +1137,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) any_item_chunk = True if not any_item_chunk: - current_item_chunks = ( - None # reset current_item_chunks if no item chunks - ) + current_item_chunks = None else: if transformed_message["type"] == "response.output_item.done": if current_item_chunks is None: @@ -1428,9 +1406,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) returned_message: List[OpenAIRealtimeEvents] = [] - # Handle transcription events that arrive independently from model - # content. Gemini sends inputTranscription / outputTranscription - # inside serverContent, separately from modelTurn / turnComplete. server_content = json_message.get("serverContent") if isinstance(server_content, dict): input_tx = server_content.get("inputTranscription") @@ -1466,8 +1441,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): delta_type="audio", ) ) - # Emit as the GA event name; _GA_TO_BETA_EVENT_TYPES translates - # this back to response.audio_transcript.delta for beta clients. returned_message.append( cast( OpenAIRealtimeEvents, @@ -1484,11 +1457,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) ) - # If serverContent only contained transcription(s) and no model - # content, mark it as already handled so the main loop skips it - # (map_openai_event would raise on an unknown serverContent - # subkey). Fall through so sibling top-level keys such as - # ``toolCall`` are still processed in the main loop. + # Mark transcription-only serverContent as handled so the main loop + # skips it; sibling keys like toolCall are still processed below. _model_content_keys = { "modelTurn", "turnComplete", @@ -1502,11 +1472,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): server_content_handled = False tool_call_handled = False - # Snapshot the items so handlers below can safely mutate - # ``json_message`` (e.g. ``_consume_usage_metadata_for_response_done`` - # pops ``usageMetadata`` to prevent a single frame from attributing - # the same token counts to two ``response.done`` events). - for key, value in list(json_message.items()): + for key, value in list( + json_message.items() + ): # snapshot: handlers may mutate json_message # Skip sibling metadata keys (e.g. ``usageMetadata``) that can # accompany a primary payload like ``toolCall`` or ``serverContent``. # ``map_openai_event`` raises ValueError on unknown keys, which @@ -1533,23 +1501,14 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) returned_message.append(transformed_message) elif openai_event == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE: - # Handle toolCall from Gemini. If the payload has no function - # calls, emit nothing — an orphaned response.created/done pair - # with no output items would confuse OpenAI-compatible clients. - # Mark the key as intentionally consumed (mirroring - # ``server_content_handled``) so any sibling keys in the same - # frame are still processed by the rest of the loop and the - # post-loop guard doesn't treat the no-op as fatal. if not value.get("functionCalls"): + # Empty toolCall — mark consumed so the post-loop guard doesn't raise. tool_call_handled = True continue if current_conversation_id is None: current_conversation_id = f"conv_{uuid.uuid4()}" - # Extract session-level response metadata once so both - # response.created and response.done can include matching - # modalities/temperature/max_output_tokens fields. session_setup: BidiGenerateContentSetup = {} if session_configuration_request is not None: try: @@ -1571,16 +1530,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) ] - # Emit response.created preamble if this is the first event in the response if current_response_id is None: current_response_id = f"resp_{uuid.uuid4()}" current_output_item_id = f"item_{uuid.uuid4()}" - - # Mirror the audio/text path: include modalities, - # temperature, and max_output_tokens on response.created so - # spec-compliant clients see consistent response metadata - # regardless of whether the response starts with content or - # a tool call. returned_message.append( { "type": "response.created", @@ -1608,7 +1560,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): response_id=current_response_id, output_item_id=current_output_item_id, ) - # Emit output_item.added and conversation.item.created for each function call for idx, tool_call in enumerate(tool_call_events): item_id = tool_call["item_id"] function_call_item: OpenAIRealtimeStreamResponseOutputItem = { @@ -1620,7 +1571,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): "name": tool_call["name"], "arguments": tool_call["arguments"], } - # response.output_item.added returned_message.append( OpenAIRealtimeStreamResponseOutputItemAdded( type="response.output_item.added", @@ -1634,14 +1584,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): }, ) ) - # conversation.item.added — Pipecat 1.3.x registers the - # call_id into _pending_function_calls inside - # _handle_evt_conversation_item_added, which is triggered - # by this event (NOT by response.output_item.added and NOT - # by the old conversation.item.created which Pipecat 1.3.x - # does not handle). Without this event the subsequent - # response.function_call_arguments.done finds an empty - # pending-calls dict and drops the tool invocation silently. + # conversation.item.added is required for Pipecat 1.3.x to + # register the call_id into _pending_function_calls before + # response.function_call_arguments.done fires. returned_message.append( cast( OpenAIRealtimeEvents, @@ -1657,13 +1602,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): }, ) ) - # response.function_call_arguments.delta — Gemini delivers - # the full arguments string in a single toolCall frame - # rather than streaming partial chunks, so emit one delta - # carrying the complete payload before the matching - # ``.done`` event. Spec-compliant OpenAI Realtime SDK - # clients accumulate ``delta.delta`` and rely on at least - # one delta before ``.done``. + # Gemini delivers args in one shot; emit a single delta before .done + # so clients that accumulate deltas get the full payload. returned_message.append( cast( OpenAIRealtimeEvents, @@ -1678,12 +1618,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): }, ) ) - # response.function_call_arguments.done returned_message.append(tool_call) - # response.output_item.done — pass a fresh copy so - # downstream handlers that mutate the item dict (e.g. the - # beta-protocol translator) don't corrupt the references - # used by sibling events sharing the same function_call_item. + # Fresh copy — downstream handlers may mutate the item dict. returned_message.append( OpenAIRealtimeOutputItemDone( type="response.output_item.done", @@ -1694,18 +1630,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) ) - # response.done - close the response so clients can submit tool - # results. Mirror the non-tool-call RESPONSE_DONE path: if Gemini - # delivered ``usageMetadata`` alongside this ``toolCall`` frame, - # propagate the real token counts so spend/budget accounting - # records the tokens consumed by the tool-call turn. Standalone - # ``usageMetadata`` frames emitted in a separate WebSocket frame - # are buffered on the instance so the next ``response.done`` - # picks them up (otherwise an authenticated client could drive - # tool-call turns whose token usage is recorded as zero, - # bypassing budgets). Falls back to an empty usage block when - # neither is available (OpenAI-compatible clients expect - # ``usage`` to always be present on response.done). resolved_tool_call_usage_metadata = ( self._consume_usage_metadata_for_response_done(json_message) ) @@ -1766,11 +1690,23 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): int, tool_call_max_output_tokens ) returned_message.append(tool_call_done_event) - # Reset IDs so the next model turn (after tool results) starts a - # fresh response with its own response.created preamble. current_output_item_id = None current_response_id = None elif openai_event == OpenAIRealtimeEventTypes.RESPONSE_DONE: + _has_pending_function_call = current_item_chunks and any( + chunk.get("item", {}).get("type") == "function_call" + for chunk in current_item_chunks + ) + if current_response_id is None and _has_pending_function_call: + # Trailing bare turnComplete after a toolCall (Vertex emits ~5 + # bookkeeping tokens before the follow-up answer). Suppress the + # empty response.done so collect_until("response.done") clients + # don't stop prematurely; buffer usage for the next real turn. + standalone_usage_metadata = json_message.get("usageMetadata") + if isinstance(standalone_usage_metadata, dict): + self._pending_usage_metadata = standalone_usage_metadata + server_content_handled = True + continue transformed_response_done_event = self.transform_response_done_event( message=BidiGenerateContentServerMessage(**json_message), # type: ignore current_response_id=current_response_id, @@ -1779,10 +1715,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): output_items=None, ) returned_message.append(transformed_response_done_event) - # Reset IDs so a subsequent turn (e.g. a `toolCall` arriving in - # a later WebSocket frame after `turnComplete`) starts a fresh - # response with its own `response.created` preamble instead of - # reusing the just-completed response ID. current_output_item_id = None current_response_id = None elif ( @@ -1791,11 +1723,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA or openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DONE ): - # Pass the locally-updated state (rather than the original - # input snapshot) so that prior iterations of this loop — - # e.g. a tool-call or response.done that just reset - # current_response_id/current_output_item_id to None — are - # honoured by the modality handler. + # Use locally-updated state so prior loop iterations' ID resets are visible. _modality_input: RealtimeResponseTransformInput = { **realtime_response_transform_input, "current_output_item_id": current_output_item_id, @@ -1821,15 +1749,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): else: raise ValueError(f"Unknown openai event: {openai_event}") if len(returned_message) == 0: - # A frame whose only top-level keys are sibling metadata (e.g. - # a standalone ``{"usageMetadata": {...}}`` emitted by Gemini - # Live between turns) is not an error — there is just nothing - # to forward to the OpenAI-shaped client. Returning the - # unchanged state keeps the WebSocket alive; raising would - # terminate the session for a benign no-op frame. - # serverContent already consumed by the transcription handler is - # a benign no-op for downstream — treat it like a metadata-only - # key when deciding whether to raise. unhandled_known_keys = [ key for key in json_message @@ -1837,11 +1756,6 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): and not (key == "serverContent" and server_content_handled) and not (key == "toolCall" and tool_call_handled) ] - # Buffer standalone usage metadata so the next response.done can - # attribute the token counts. Without this, an authenticated - # client driving turns whose usageMetadata is emitted in a - # separate frame would have those tokens recorded as zero spend, - # bypassing budget enforcement. standalone_usage_metadata = json_message.get("usageMetadata") if isinstance(standalone_usage_metadata, dict): self._pending_usage_metadata = standalone_usage_metadata @@ -1889,9 +1803,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): } def requires_session_configuration(self) -> bool: - # Default behavior is backwards-compatible: send setup on connect. - # Opt-in to deferred setup for tool-injection flow via: - # litellm.gemini_live_defer_setup = True + # Deferred setup opt-in: litellm.gemini_live_defer_setup = True return not litellm.gemini_live_defer_setup def session_configuration_request(self, model: str) -> str: diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index 6751004f1b1..cebab20adde 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -12,6 +12,7 @@ from litellm.types.realtime import RealtimeQueryParams from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from ....litellm_core_utils.realtime_streaming import ( + RealtimeEventNormalizer, RealTimeStreaming, client_sent_openai_beta_realtime_header, ) @@ -95,6 +96,14 @@ class OpenAIRealtime(OpenAIChatCompletion): url = url.copy_with(params=query_params) return str(url) + def _make_event_normalizer(self) -> Optional[RealtimeEventNormalizer]: + """Return a per-session GA event normalizer, or None for passthrough. + + Subclasses (e.g. XAIRealtime) override this to supply a provider-specific + normalizer instance. + """ + return None + async def async_realtime( self, model: str, @@ -165,6 +174,7 @@ class OpenAIRealtime(OpenAIChatCompletion): if (query_params or {}).get("intent") == "transcription" else None ), + event_normalizer=self._make_event_normalizer(), ) await realtime_streaming.bidirectional_forward() diff --git a/litellm/llms/vertex_ai/realtime/transformation.py b/litellm/llms/vertex_ai/realtime/transformation.py index 1fe9f15c9f0..9f339d3dc29 100644 --- a/litellm/llms/vertex_ai/realtime/transformation.py +++ b/litellm/llms/vertex_ai/realtime/transformation.py @@ -198,7 +198,7 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): setup_config.setdefault("inputAudioTranscription", {}) setup_config.setdefault("outputAudioTranscription", {}) - return setup_config + return self._finalize_gemini_live_setup(model, setup_config) def transform_realtime_request( self, diff --git a/litellm/llms/xai/realtime/handler.py b/litellm/llms/xai/realtime/handler.py index eab19f4a6c8..9dac5d945fd 100644 --- a/litellm/llms/xai/realtime/handler.py +++ b/litellm/llms/xai/realtime/handler.py @@ -10,6 +10,7 @@ This requires websockets, and is currently only supported on LiteLLM Proxy. from litellm.constants import XAI_API_BASE from ...openai.realtime.handler import OpenAIRealtime +from .transformation import XAIRealtimeNormalizer class XAIRealtime(OpenAIRealtime): @@ -28,6 +29,10 @@ class XAIRealtime(OpenAIRealtime): """xAI uses a different API base URL.""" return XAI_API_BASE + def _make_event_normalizer(self) -> XAIRealtimeNormalizer: + """Return a fresh per-session XAI normalizer instance.""" + return XAIRealtimeNormalizer() + def _get_additional_headers( self, api_key: str, diff --git a/litellm/llms/xai/realtime/transformation.py b/litellm/llms/xai/realtime/transformation.py new file mode 100644 index 00000000000..f7981dc547a --- /dev/null +++ b/litellm/llms/xai/realtime/transformation.py @@ -0,0 +1,295 @@ +""" +xAI Grok Voice realtime event normalizer. + +xAI's Grok Voice realtime API is structurally OpenAI-compatible but ships +several wire-format quirks that cause strict GA clients (e.g. pipecat's +``OpenAIRealtimeLLMService``) to crash before they can process tool calls: + + - ``ping`` keepalive events (unknown to GA clients) + - ``usage: {}`` on ``response.created`` / ``response.done`` + - ``role: "tool"`` on ``conversation.item.added`` function_call items + - Missing ``output_index`` / ``content_index`` on streaming response events + - Missing ``part`` on ``response.content_part.done`` + +``XAIRealtimeNormalizer`` is plugged into ``RealTimeStreaming`` at handler +construction time (see ``handler.py``) so all normalization is isolated here +and ``RealTimeStreaming`` stays provider-agnostic. +""" + +from typing import Any, Optional + + +class XAIRealtimeNormalizer: + """Per-session normalizer that fixes xAI Grok Voice wire-format quirks.""" + + # --------------------------------------------------------------------------- + # Event-type sets used by the index-injection logic + # --------------------------------------------------------------------------- + _EVENTS_NEEDING_OUTPUT_INDEX = frozenset( + [ + "response.output_item.added", + "response.output_item.done", + "response.content_part.added", + "response.content_part.done", + "response.output_text.delta", + "response.output_text.done", + "response.output_audio_transcript.delta", + "response.output_audio_transcript.done", + "response.output_audio.delta", + "response.output_audio.done", + "response.function_call_arguments.delta", + "response.function_call_arguments.done", + ] + ) + _EVENTS_NEEDING_CONTENT_INDEX = frozenset( + [ + "response.content_part.added", + "response.content_part.done", + "response.output_text.delta", + "response.output_text.done", + "response.output_audio_transcript.delta", + "response.output_audio_transcript.done", + "response.output_audio.delta", + "response.output_audio.done", + ] + ) + + def __init__(self) -> None: + # Cache content-part objects keyed by (response_id, item_id, content_index) + # so that ``response.content_part.done`` events missing ``part`` can be + # back-filled from earlier ``content_part.added`` / delta-done events. + self._content_part_by_key: dict[tuple, dict[str, Any]] = {} + + # --------------------------------------------------------------------------- + # Public interface consumed by RealTimeStreaming + # --------------------------------------------------------------------------- + + def should_drop(self, event: object) -> bool: + """Return True for provider-specific keepalives unknown to GA clients.""" + return isinstance(event, dict) and event.get("type") == "ping" + + def normalize(self, event: dict) -> dict: + """Apply all xAI normalization passes in order.""" + event = self._normalize_content_part_events(event) + event_type = event.get("type") or "" + event = self._normalize_conversation_item_added(event, event_type) + event = self._inject_missing_indices(event, event_type) + event = self._normalize_response_usage_event(event, event_type) + return event + + def patch_outgoing_session(self, session: dict) -> dict: + """Patch a client ``session.update`` payload before forwarding to xAI. + + Unlike OpenAI, xAI does not default ``turn_detection.create_response`` + to ``True`` for ``server_vad``. Clients such as Pipecat omit the field, + which leaves VAD detecting speech but never auto-creating a response. + Only fill the default when the client did not set ``create_response``. + """ + session = dict(session) + self._default_server_vad_create_response(session) + return session + + @staticmethod + def _default_server_vad_create_response(session: dict) -> None: + turn_detection = session.get("turn_detection") + if isinstance(turn_detection, dict): + XAIRealtimeNormalizer._ensure_server_vad_create_response(turn_detection) + + audio = session.get("audio") + if isinstance(audio, dict): + audio_input = audio.get("input") + if isinstance(audio_input, dict): + nested_td = audio_input.get("turn_detection") + if isinstance(nested_td, dict): + XAIRealtimeNormalizer._ensure_server_vad_create_response(nested_td) + + @staticmethod + def _ensure_server_vad_create_response(turn_detection: dict) -> None: + if ( + turn_detection.get("type") == "server_vad" + and "create_response" not in turn_detection + ): + turn_detection["create_response"] = True + + # --------------------------------------------------------------------------- + # Pass 1: content-part caching and back-fill + # --------------------------------------------------------------------------- + + @staticmethod + def _content_part_key(event: dict) -> tuple: + return ( + event.get("response_id"), + event.get("item_id"), + event.get("content_index", 0), + ) + + def _remember_content_part(self, event: dict) -> None: + part = event.get("part") + if isinstance(part, dict): + self._content_part_by_key[self._content_part_key(event)] = part + + def _update_content_part_field( + self, event: dict, *, part_type: str, field: str, value: object + ) -> None: + if value is None: + return + key = self._content_part_key(event) + existing = self._content_part_by_key.get(key) + if not isinstance(existing, dict): + updated = {"type": part_type, field: value} + else: + updated = { + **existing, + "type": existing.get("type", part_type), + field: value, + } + self._content_part_by_key[key] = updated + + def _resolve_content_part(self, event: dict) -> dict[str, Any]: + part = event.get("part") + if isinstance(part, dict): + return part + cached = self._content_part_by_key.get(self._content_part_key(event)) + if isinstance(cached, dict): + return cached + return {"type": "audio", "transcript": ""} + + def _normalize_content_part_events(self, event: dict) -> dict: + event_type = event.get("type") + + if event_type == "response.content_part.added": + self._remember_content_part(event) + if not isinstance(event.get("part"), dict): + return {**event, "part": self._resolve_content_part(event)} + return event + + if event_type == "response.output_text.done": + self._update_content_part_field( + event, part_type="text", field="text", value=event.get("text") + ) + return event + + if event_type == "response.output_audio_transcript.done": + self._update_content_part_field( + event, + part_type="audio", + field="transcript", + value=event.get("transcript"), + ) + return event + + if event_type == "response.content_part.done": + self._remember_content_part(event) + if not isinstance(event.get("part"), dict): + return {**event, "part": self._resolve_content_part(event)} + return event + + return event + + # --------------------------------------------------------------------------- + # Pass 2: conversation.item.added role normalisation + # --------------------------------------------------------------------------- + + @staticmethod + def _normalize_conversation_item_added(event: dict, event_type: str) -> dict: + """Map ``role: "tool"`` → ``role: "assistant"`` on function_call items. + + xAI uses ``role: "tool"`` which is not in the GA-allowed set + ("user" | "assistant" | "system"). + """ + if event_type != "conversation.item.added": + return event + item = event.get("item") + if not isinstance(item, dict): + return event + if item.get("role") == "tool": + return {**event, "item": {**item, "role": "assistant"}} + return event + + # --------------------------------------------------------------------------- + # Pass 3: inject missing output_index / content_index + # --------------------------------------------------------------------------- + + def _inject_missing_indices(self, event: dict, event_type: str) -> dict: + """Inject ``output_index`` / ``content_index`` defaults when absent. + + xAI omits both fields on every streaming response event; pydantic GA + clients require them as non-optional ints. Defaulting to 0 is correct + for single-turn single-item responses and harmless for well-formed events. + """ + needs_output = event_type in self._EVENTS_NEEDING_OUTPUT_INDEX + needs_content = event_type in self._EVENTS_NEEDING_CONTENT_INDEX + if not needs_output and not needs_content: + return event + patch: dict[str, Any] = {} + if needs_output and "output_index" not in event: + patch["output_index"] = 0 + if needs_content and "content_index" not in event: + patch["content_index"] = 0 + if not patch: + return event + return {**event, **patch} + + # --------------------------------------------------------------------------- + # Pass 4: response usage normalisation + # --------------------------------------------------------------------------- + + @staticmethod + def _default_ga_usage() -> dict[str, Any]: + default_details: dict[str, Any] = { + "cached_tokens": 0, + "text_tokens": 0, + "audio_tokens": 0, + } + return { + "total_tokens": 0, + "input_tokens": 0, + "output_tokens": 0, + "input_token_details": default_details.copy(), + "output_token_details": default_details.copy(), + } + + @staticmethod + def _normalize_usage( + usage: object, *, empty_as_null: bool + ) -> Optional[dict[str, Any]]: + """Coerce a usage object into the full OpenAI GA shape. + + ``empty_as_null=True`` for ``response.created`` (usage optional). + ``empty_as_null=False`` for ``response.done`` (e2e tests assert non-null). + """ + if not isinstance(usage, dict): + return None + if not usage: + return None if empty_as_null else XAIRealtimeNormalizer._default_ga_usage() + default_details: dict[str, Any] = { + "cached_tokens": 0, + "text_tokens": 0, + "audio_tokens": 0, + } + normalized: dict[str, Any] = { + "total_tokens": usage.get("total_tokens", 0), + "input_tokens": usage.get("input_tokens", 0), + "output_tokens": usage.get("output_tokens", 0), + "input_token_details": default_details.copy(), + "output_token_details": default_details.copy(), + } + for key in ("input_token_details", "output_token_details"): + details = usage.get(key) + if isinstance(details, dict): + normalized[key] = {**default_details, **details} + return normalized + + def _normalize_response_usage_event(self, event: dict, event_type: str) -> dict: + if event_type not in ("response.created", "response.done"): + return event + response = event.get("response") + if not isinstance(response, dict) or "usage" not in response: + return event + normalized_usage = self._normalize_usage( + response.get("usage"), + empty_as_null=event_type == "response.created", + ) + if normalized_usage is response.get("usage"): + return event + return {**event, "response": {**response, "usage": normalized_usage}} diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 1fa6cb1ec71..5cf5af8a22a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -16614,7 +16614,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "gemini_native_audio": true }, "gemini/gemini-live-2.5-flash-preview-native-audio-09-2025": { "cache_read_input_token_cost": 7.5e-08, @@ -16665,7 +16666,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "gemini_native_audio": true }, "gemini-2.5-flash-lite-preview-06-17": { "deprecation_date": "2025-11-18", @@ -42438,7 +42440,8 @@ "audio" ], "supports_audio_input": true, - "supports_audio_output": true + "supports_audio_output": true, + "gemini_native_audio": true }, "gemini-2.5-flash-native-audio-preview-09-2025": { "input_cost_per_audio_token": 1e-06, @@ -42462,7 +42465,8 @@ "audio" ], "supports_audio_input": true, - "supports_audio_output": true + "supports_audio_output": true, + "gemini_native_audio": true }, "gemini-2.5-flash-native-audio-preview-12-2025": { "input_cost_per_audio_token": 1e-06, @@ -42486,7 +42490,8 @@ "audio" ], "supports_audio_input": true, - "supports_audio_output": true + "supports_audio_output": true, + "gemini_native_audio": true }, "gemini-3.1-flash-live-preview": { "input_cost_per_audio_token": 3e-06, @@ -42518,7 +42523,8 @@ "supports_audio_output": true, "supports_function_calling": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "gemini_audio_only_live": true }, "gemini/gemini-2.5-flash-native-audio-latest": { "input_cost_per_audio_token": 1e-06, @@ -42544,7 +42550,8 @@ "supports_audio_input": true, "supports_audio_output": true, "tpm": 250000, - "rpm": 10 + "rpm": 10, + "gemini_native_audio": true }, "gemini/gemini-2.5-flash-native-audio-preview-09-2025": { "input_cost_per_audio_token": 1e-06, @@ -42570,7 +42577,8 @@ "supports_audio_input": true, "supports_audio_output": true, "tpm": 250000, - "rpm": 10 + "rpm": 10, + "gemini_native_audio": true }, "gemini/gemini-2.5-flash-native-audio-preview-12-2025": { "input_cost_per_audio_token": 1e-06, @@ -42596,7 +42604,8 @@ "supports_audio_input": true, "supports_audio_output": true, "tpm": 250000, - "rpm": 10 + "rpm": 10, + "gemini_native_audio": true }, "gemini/gemini-3.1-flash-live-preview": { "input_cost_per_audio_token": 3e-06, @@ -42630,7 +42639,8 @@ "supports_vision": true, "supports_web_search": true, "tpm": 250000, - "rpm": 10 + "rpm": 10, + "gemini_audio_only_live": true }, "gemini-2.5-flash-preview-tts": { "input_cost_per_token": 3e-07, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 5d2ed4244e0..ed1b2fb7db9 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -16693,7 +16693,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "gemini_native_audio": true }, "gemini/gemini-live-2.5-flash-preview-native-audio-09-2025": { "cache_read_input_token_cost": 7.5e-08, @@ -16744,7 +16745,8 @@ "search_context_size_low": 0.035, "search_context_size_medium": 0.035, "search_context_size_high": 0.035 - } + }, + "gemini_native_audio": true }, "gemini-2.5-flash-lite-preview-06-17": { "deprecation_date": "2025-11-18", @@ -42676,7 +42678,8 @@ "audio" ], "supports_audio_input": true, - "supports_audio_output": true + "supports_audio_output": true, + "gemini_native_audio": true }, "gemini-2.5-flash-native-audio-preview-09-2025": { "input_cost_per_audio_token": 1e-06, @@ -42700,7 +42703,8 @@ "audio" ], "supports_audio_input": true, - "supports_audio_output": true + "supports_audio_output": true, + "gemini_native_audio": true }, "gemini-2.5-flash-native-audio-preview-12-2025": { "input_cost_per_audio_token": 1e-06, @@ -42724,7 +42728,8 @@ "audio" ], "supports_audio_input": true, - "supports_audio_output": true + "supports_audio_output": true, + "gemini_native_audio": true }, "gemini-3.1-flash-live-preview": { "input_cost_per_audio_token": 3e-06, @@ -42756,7 +42761,8 @@ "supports_audio_output": true, "supports_function_calling": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "gemini_audio_only_live": true }, "gemini/gemini-2.5-flash-native-audio-latest": { "input_cost_per_audio_token": 1e-06, @@ -42782,7 +42788,8 @@ "supports_audio_input": true, "supports_audio_output": true, "tpm": 250000, - "rpm": 10 + "rpm": 10, + "gemini_native_audio": true }, "gemini/gemini-2.5-flash-native-audio-preview-09-2025": { "input_cost_per_audio_token": 1e-06, @@ -42808,7 +42815,8 @@ "supports_audio_input": true, "supports_audio_output": true, "tpm": 250000, - "rpm": 10 + "rpm": 10, + "gemini_native_audio": true }, "gemini/gemini-2.5-flash-native-audio-preview-12-2025": { "input_cost_per_audio_token": 1e-06, @@ -42834,7 +42842,8 @@ "supports_audio_input": true, "supports_audio_output": true, "tpm": 250000, - "rpm": 10 + "rpm": 10, + "gemini_native_audio": true }, "gemini/gemini-3.1-flash-live-preview": { "input_cost_per_audio_token": 3e-06, @@ -42868,7 +42877,8 @@ "supports_vision": true, "supports_web_search": true, "tpm": 250000, - "rpm": 10 + "rpm": 10, + "gemini_audio_only_live": true }, "gemini-2.5-flash-preview-tts": { "input_cost_per_token": 3e-07, diff --git a/tests/e2e/gateway/litellm-config.yml b/tests/e2e/gateway/litellm-config.yml index e059ac62429..ecc2c044039 100644 --- a/tests/e2e/gateway/litellm-config.yml +++ b/tests/e2e/gateway/litellm-config.yml @@ -25,9 +25,10 @@ general_settings: database_url: os.environ/DATABASE_URL control_plane_url: os.environ/CONTROL_PLANE_URL alerts: ["email"] + proxy_budget_rescheduler_min_time: 15 proxy_budget_rescheduler_max_time: 20 - + # fallbacks: [{"gpt-4": ["anthropic.claude-3-5-sonnet-20240620-v1:0"]}] #Configure fallbacks for context window exeeded errors (In this example, we will fall back to Claude Sonnet if over 8000 tokens, which is gpt-4's limit) # default_fallbacks: ["anthropic.claude-3-haiku-20240307-v1:0"] #Configure fallbacks for any error for every model (the above fallback configurations override this one) # environment_variables: @@ -132,14 +133,52 @@ model_list: model: gemini/gemini-2-embedding api_key: os.environ/GEMINI_API_KEY - # realtime models - model_name: openai-realtime litellm_params: - model: openai/realtime-2 + model: openai/gpt-realtime api_key: os.environ/OPENAI_API_KEY model_info: mode: realtime + - model_name: azure-realtime + litellm_params: + model: azure/gpt-realtime-2 + api_key: os.environ/AZURE_API_KEY + api_base: os.environ/AZURE_API_BASE + api_version: "2025-08-28" + realtime_protocol: GA # Possible values: "GA"/ "v1", "beta" + model_info: + mode: realtime + + - model_name: gemini-realtime + litellm_params: + model: gemini/gemini-3.1-flash-live-preview + api_key: os.environ/GEMINI_API_KEY + model_info: + mode: realtime + + - model_name: vertex-realtime + litellm_params: + model: vertex_ai/gemini-live-2.5-flash-preview-native-audio-09-2025 + vertex_project: os.environ/VERTEXAI_PROJECT + vertex_location: us-central1 + vertex_credentials: os.environ/VERTEXAI_CREDENTIALS + model_info: + mode: realtime + + - model_name: bedrock-realtime + litellm_params: + model: bedrock/amazon.nova-sonic-v1:0 + aws_region_name: us-east-1 + model_info: + mode: realtime + + - model_name: xai-realtime + litellm_params: + model: xai/grok-voice-latest + api_key: os.environ/XAI_API_KEY + model_info: + mode: realtime - model_name: rust-ocr-mistral litellm_params: model: mistral/mistral-ocr-latest diff --git a/tests/e2e/models.py b/tests/e2e/models.py index fbeb3d44fa5..83b7d148957 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -197,6 +197,7 @@ class CustomPricing(BaseModel): overrides, and /model/info echoes the rates the proxy resolved.""" model_config = ConfigDict(extra="ignore") + mode: str | None = None input_cost_per_token: float | None = None output_cost_per_token: float | None = None cache_read_input_token_cost: float | None = None diff --git a/tests/e2e/realtime/REALTIME_COVERAGE_MATRIX.md b/tests/e2e/realtime/REALTIME_COVERAGE_MATRIX.md new file mode 100644 index 00000000000..8624475d0de --- /dev/null +++ b/tests/e2e/realtime/REALTIME_COVERAGE_MATRIX.md @@ -0,0 +1,55 @@ +# Realtime e2e coverage + +Live tests for the proxy realtime websocket endpoint (`/v1/realtime`). One +GA-speaking websocket client drives every provider; the proxy normalizes each +provider's stream into the OpenAI GA event schema, so the same assertions hold +across providers and only the model alias changes. + +## What is asserted + +For each configured provider, `test_text_conversation` checks the session +lifecycle (`session.created`, then `session.update` echoed by `session.updated`), +the canonical response sequence (`response.created`, `response.output_item.added`, +through `response.done`), that the streamed deltas reconstruct a non-empty +transcript, and that `response.done` carries normalized usage. + +`test_tool_call_round_trip` checks the full tool path: the model emits a +normalized `response.function_call_arguments.done` with valid JSON arguments and +a matching `function_call` output item, the test sends a `function_call_output` +back, and the follow-up response incorporates the result (the temperature 72 +appears). + +`test_realtime_pipecat_e2e` is a realism layer that drives the same providers +through pipecat's GA `OpenAIRealtimeLLMService` (base_url pointed at the proxy) +rather than speaking the protocol by hand. Its assertions are coarse (the tool +callback fired, assistant text was produced); the raw-websocket suite is the +source of truth. It skips unless `pipecat-ai` is installed +(`uv pip install "pipecat-ai[openai]"`). + +## Provider status + +| provider | model alias | status | +|----------|-------------|--------| +| openai | `openai-realtime` | covered (in gateway config) | +| gemini | `gemini-realtime` | covered (in gateway config; needs Gemini Live API access) | +| azure | `azure-realtime` | gap: add to gateway config + AZURE creds | +| vertex_ai | `vertex-realtime` | gap: add to gateway config + Vertex creds | +| bedrock | `bedrock-realtime` | gap: add to gateway config + AWS creds | +| xai | `xai-realtime` | gap: add to gateway config + XAI_API_KEY | + +A provider whose alias is not present in the proxy's `/model/info` skips (skip on +environment). To enable one, add a `model_info.mode: realtime` entry under that +alias to `tests/e2e/gateway/litellm-config.yml` and give the proxy the +provider's credentials; the test then runs with no code change. + +## Running + +Start a proxy with the gateway config and the provider keys set in its +environment, then + +``` +uv run pytest tests/e2e/realtime/ -v +``` + +Tests skip when no proxy answers `GET /health/liveliness` at `LITELLM_PROXY_URL` +(default `http://localhost:4000`). diff --git a/tests/e2e/realtime/conftest.py b/tests/e2e/realtime/conftest.py new file mode 100644 index 00000000000..4a5c4837a1a --- /dev/null +++ b/tests/e2e/realtime/conftest.py @@ -0,0 +1,20 @@ +"""Realtime suite's `client` fixture. + +The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker +live in the parent tests/e2e/conftest.py. RealtimeClient holds the shared +Gateway, so the `resources` fixture cleans up keys this suite creates. +""" + +import pytest + +from realtime_client import RealtimeClient, build_client + + +@pytest.fixture(scope="session") +def client() -> RealtimeClient: + return build_client() + + +@pytest.fixture(scope="session") +def configured_models(client: RealtimeClient) -> frozenset[str]: + return client.configured_models() diff --git a/tests/e2e/realtime/fixtures/weather_question_24k.wav b/tests/e2e/realtime/fixtures/weather_question_24k.wav new file mode 100644 index 00000000000..c1c043f3874 Binary files /dev/null and b/tests/e2e/realtime/fixtures/weather_question_24k.wav differ diff --git a/tests/e2e/realtime/pipecat_service.py b/tests/e2e/realtime/pipecat_service.py new file mode 100644 index 00000000000..2b9074f474d --- /dev/null +++ b/tests/e2e/realtime/pipecat_service.py @@ -0,0 +1,73 @@ +"""Shared pipecat realtime service for the proxy realtime e2e suite. + +Both pipecat suites drive the proxy through pipecat's GA realtime service. The +stock ``OpenAIRealtimeLLMService`` sends websocket keepalive pings at its default +interval, and the LiteLLM proxy does not answer them, so the connection is closed +with a 1011 before the run completes. ``LiteLLMRealtimeLLMService`` carries the +three overrides from bot.py needed to talk to the proxy, the keepalive-disabling +``_connect`` being the load-bearing one for every provider. + +Importing this module skips the collecting test when pipecat is not installed: + + uv pip install "pipecat-ai[openai]" +""" + +# pipecat is an optional, dynamically typed dependency loaded behind importorskip, +# so its symbols are Unknown to the type checker; relax those rules for this file. +# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false, reportUnknownArgumentType=false, reportAttributeAccessIssue=false, reportUntypedBaseClass=false, reportUnknownParameterType=false, reportMissingParameterType=false + +import pytest + +pytest.importorskip("pipecat", reason="pipecat-ai not installed") + +from pipecat.services.openai.realtime.llm import ( # noqa: E402 + OpenAIRealtimeLLMService, +) +from websockets.asyncio.client import connect as websocket_connect # noqa: E402 + + +class LiteLLMRealtimeLLMService(OpenAIRealtimeLLMService): + """Minimal LiteLLM-aware realtime service for tests. + + Three overrides carried from bot.py: + 1. _connect - disables websockets keepalive pings (LiteLLM proxy + does not respond to pings, causing 1011 errors). + 2. _create_response - sends session.update with tools BEFORE history + items so that Gemini's deferred-setup logic in the + proxy can include tools in the very first setup + message it forwards to the backend. + 3. _handle_evt_session_created - immediately marks the session ready + without waiting for a session.updated echo (the + LiteLLM Gemini bridge does not send one). + """ + + async def _connect(self) -> None: + if self._websocket: + return + try: + # self.base_url already carries the ?model= the proxy routes on: + # the parent __init__ sets self.base_url = f"{base_url}?model={settings.model}" + # before _connect runs, so passing it through preserves the query param. + self._websocket = await websocket_connect( + uri=self.base_url, + additional_headers={"Authorization": f"Bearer {self.api_key}"}, + ping_interval=None, + close_timeout=10, + max_size=None, + ) + self._receive_task = self.create_task(self._receive_task_handler()) + except Exception as exc: + await self.push_error(error_msg=f"Error connecting: {exc}", exception=exc) + self._websocket = None + + async def _create_response(self) -> None: + if self._llm_needs_conversation_setup and self._context: + await self._send_session_update() + await super()._create_response() + + async def _handle_evt_session_created(self, evt: object) -> None: + await self._send_session_update() + self._api_session_ready = True + if self._run_llm_when_api_session_ready: + self._run_llm_when_api_session_ready = False + await self._create_response() diff --git a/tests/e2e/realtime/realtime_client.py b/tests/e2e/realtime/realtime_client.py new file mode 100644 index 00000000000..07dfd76108b --- /dev/null +++ b/tests/e2e/realtime/realtime_client.py @@ -0,0 +1,303 @@ +"""Client for realtime e2e tests over the proxy's /v1/realtime websocket. + +The proxy normalizes every provider's realtime stream into the OpenAI GA event +schema toward the client, so one GA-speaking websocket client validates every +provider and only the model alias changes. Every other e2e suite is HTTP-only; +this is the one suite that opens a websocket, using websockets.sync so it stays +synchronous like the rest of tests/e2e. Sent and received events are pydantic +models, matching the suite's no-raw-dicts rule. +""" + +from __future__ import annotations + +import time +from collections.abc import Generator +from contextlib import contextmanager +from dataclasses import dataclass +from typing import Any, TypeVar +from urllib.parse import urlencode + +import pytest +from pydantic import BaseModel, ConfigDict +from websockets.sync.client import connect +from websockets.sync.connection import Connection + +from e2e_config import PROXY_BASE_URL +from e2e_gateway import Gateway, build_gateway + +_M = TypeVar("_M", bound=BaseModel) + + +def _ws_base_url() -> str: + for scheme, ws_scheme in (("https://", "wss://"), ("http://", "ws://")): + if PROXY_BASE_URL.startswith(scheme): + return ws_scheme + PROXY_BASE_URL[len(scheme) :] + return PROXY_BASE_URL + + +def realtime_ws_url(model: str) -> str: + return f"{_ws_base_url()}/v1/realtime?{urlencode({'model': model})}" + + +@dataclass(frozen=True, slots=True) +class RealtimeProvider: + id: str + model: str + + +PROVIDERS = ( + RealtimeProvider("openai", "openai-realtime"), + RealtimeProvider("azure", "azure-realtime"), + RealtimeProvider("gemini", "gemini-realtime"), + RealtimeProvider("vertex_ai", "vertex-realtime"), + # RealtimeProvider("bedrock", "bedrock-realtime"), # TODO: Enable this when Bedrock is passing + RealtimeProvider("xai", "xai-realtime"), +) + + +def skip_if_unconfigured( + provider: RealtimeProvider, configured: frozenset[str] +) -> None: + if provider.model not in configured: + pytest.skip(f"{provider.model} not configured on proxy") + + +# ---- sent events ------------------------------------------------------- + + +class JsonSchemaProperty(BaseModel): + type: str + + +class JsonSchema(BaseModel): + type: str = "object" + properties: dict[str, JsonSchemaProperty] + required: list[str] + + +class FunctionTool(BaseModel): + type: str = "function" + name: str + description: str + parameters: JsonSchema + + +class SessionConfig(BaseModel): + modalities: list[str] = ["text"] + instructions: str | None = None + tools: list[FunctionTool] | None = None + tool_choice: str | None = None + + +class SessionUpdate(BaseModel): + type: str = "session.update" + session: SessionConfig + + +class InputTextContent(BaseModel): + type: str = "input_text" + text: str + + +class MessageItem(BaseModel): + type: str = "message" + role: str = "user" + content: list[InputTextContent] + + +class FunctionCallOutputItem(BaseModel): + type: str = "function_call_output" + call_id: str + output: str + + +class ConversationItemCreate(BaseModel): + type: str = "conversation.item.create" + item: MessageItem | FunctionCallOutputItem + + +class ResponseCreate(BaseModel): + type: str = "response.create" + + +def user_message(text: str) -> ConversationItemCreate: + return ConversationItemCreate( + item=MessageItem(content=[InputTextContent(text=text)]) + ) + + +# ---- received events --------------------------------------------------- + + +class ServerEnvelope(BaseModel): + model_config = ConfigDict(extra="allow") + type: str = "" + + +class DeltaEvent(BaseModel): + type: str + delta: str = "" + + +class FunctionCallArgumentsDone(BaseModel): + type: str + call_id: str + arguments: str + + +class ContentPart(BaseModel): + model_config = ConfigDict(extra="allow") + text: str | None = None + transcript: str | None = None + + +class OutputItem(BaseModel): + model_config = ConfigDict(extra="allow") + type: str | None = None + name: str | None = None + call_id: str | None = None + content: list[ContentPart] | None = None + + +class OutputItemDone(BaseModel): + type: str + item: OutputItem + + +class ResponsePayload(BaseModel): + model_config = ConfigDict(extra="allow") + usage: dict[str, Any] | None = None + output: list[OutputItem] | None = None + + +class ResponseDone(BaseModel): + type: str + response: ResponsePayload + + +@dataclass(frozen=True, slots=True) +class ReceivedEvent: + type: str + payload: str + + +def events_of_type( + events: tuple[ReceivedEvent, ...], event_type: str +) -> tuple[ReceivedEvent, ...]: + return tuple(e for e in events if e.type == event_type) + + +def parse_last( + events: tuple[ReceivedEvent, ...], event_type: str, model: type[_M] +) -> _M | None: + matches = events_of_type(events, event_type) + return model.model_validate_json(matches[-1].payload) if matches else None + + +_TEXT_DELTA_TYPES = ( + # beta protocol (OpenAI-Beta: realtime=v1) + "response.text.delta", + "response.audio_transcript.delta", + # GA protocol (default toward the proxy) + "response.output_text.delta", + "response.output_audio_transcript.delta", +) + + +def _text_from_response_done(events: tuple[ReceivedEvent, ...]) -> str: + done = parse_last(events, "response.done", ResponseDone) + if done is None or not done.response.output: + return "" + parts: list[str] = [] + for item in done.response.output: + if item.type != "message": + continue + for content in item.content or []: + if content.text: + parts.append(content.text) + elif content.transcript: + parts.append(content.transcript) + return "".join(parts) + + +def transcript(events: tuple[ReceivedEvent, ...]) -> str: + for delta_type in _TEXT_DELTA_TYPES: + text = "".join( + DeltaEvent.model_validate_json(e.payload).delta + for e in events_of_type(events, delta_type) + ) + if text: + return text + return _text_from_response_done(events) + + +def function_call_item(events: tuple[ReceivedEvent, ...]) -> OutputItem | None: + for event in events_of_type(events, "response.output_item.done"): + item = OutputItemDone.model_validate_json(event.payload).item + if item.type == "function_call": + return item + return None + + +# ---- session + client -------------------------------------------------- + + +def _as_text(message: str | bytes) -> str: + return message.decode("utf-8") if isinstance(message, bytes) else message + + +@dataclass(frozen=True, slots=True) +class RealtimeSession: + connection: Connection + + def send(self, event: BaseModel) -> None: + self.connection.send(event.model_dump_json(by_alias=True, exclude_none=True)) + + def collect_until( + self, stop_type: str, *, timeout: float + ) -> tuple[ReceivedEvent, ...]: + deadline = time.monotonic() + timeout + collected: list[ReceivedEvent] = [] + while time.monotonic() < deadline: + try: + text = _as_text( + self.connection.recv(timeout=deadline - time.monotonic()) + ) + except TimeoutError: + break + event = ReceivedEvent( + type=ServerEnvelope.model_validate_json(text).type, payload=text + ) + collected.append(event) + if event.type == stop_type: + return tuple(collected) + raise TimeoutError( + f"no {stop_type!r} within {timeout}s; got {[e.type for e in collected]}" + ) + + +@dataclass(frozen=True, slots=True) +class RealtimeClient: + gateway: Gateway + + def configured_models(self) -> frozenset[str]: + return frozenset( + entry.model_name + for entry in self.gateway.model_info() + if entry.model_info.mode == "realtime" + ) + + @contextmanager + def connect( + self, *, key: str, model: str, timeout: float = 15.0 + ) -> Generator[RealtimeSession, None, None]: + with connect( + realtime_ws_url(model), + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=timeout, + ) as connection: + yield RealtimeSession(connection=connection) + + +def build_client() -> RealtimeClient: + return RealtimeClient(gateway=build_gateway()) diff --git a/tests/e2e/realtime/test_realtime_e2e.py b/tests/e2e/realtime/test_realtime_e2e.py new file mode 100644 index 00000000000..01356900141 --- /dev/null +++ b/tests/e2e/realtime/test_realtime_e2e.py @@ -0,0 +1,149 @@ +"""Live e2e for the proxy realtime websocket (/v1/realtime). + +Each test opens a websocket through the proxy, speaks the OpenAI GA realtime +event schema, and asserts the proxy normalizes the provider's stream into that +schema: the session lifecycle, the canonical response event sequence with a +reconstructed transcript and usage, and a full tool-call round-trip (call -> +tool result -> a follow-up response that uses the result). + +One GA-speaking client validates every provider; only the model alias changes. A +provider whose realtime alias is not configured on the proxy skips (skip on +environment); once it is configured, a protocol failure is a hard failure. See +REALTIME_COVERAGE_MATRIX.md. +""" + +import pytest +from pydantic import BaseModel + +from realtime_client import ( + PROVIDERS, + ConversationItemCreate, + FunctionCallArgumentsDone, + FunctionCallOutputItem, + FunctionTool, + JsonSchema, + JsonSchemaProperty, + RealtimeClient, + RealtimeProvider, + ResponseCreate, + ResponseDone, + SessionConfig, + SessionUpdate, + function_call_item, + parse_last, + skip_if_unconfigured, + transcript, + user_message, +) + +pytestmark = pytest.mark.e2e + +PROVIDER_PARAMS = [pytest.param(p, id=p.id) for p in PROVIDERS] + +WEATHER_TOOL = FunctionTool( + name="get_weather", + description="Get the current temperature in Fahrenheit for a given city.", + parameters=JsonSchema( + properties={"city": JsonSchemaProperty(type="string")}, required=["city"] + ), +) + + +class WeatherArgs(BaseModel): + city: str + + +class WeatherResult(BaseModel): + city: str + temperature_f: int + + +@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +def test_text_conversation( + client: RealtimeClient, + scoped_key: str, + configured_models: frozenset[str], + provider: RealtimeProvider, +) -> None: + skip_if_unconfigured(provider, configured_models) + + with client.connect(key=scoped_key, model=provider.model) as session: + created = session.collect_until("session.created", timeout=20) + assert created[-1].type == "session.created" + + session.send( + SessionUpdate( + session=SessionConfig( + instructions="You are a terse assistant. Reply in one short sentence." + ) + ) + ) + session.collect_until("session.updated", timeout=20) + + session.send(user_message("Say the single word hello.")) + session.send(ResponseCreate()) + events = session.collect_until("response.done", timeout=60) + + types = {e.type for e in events} + assert "response.created" in types + assert "response.output_item.added" in types + + assert transcript(events).strip() != "" + + done = parse_last(events, "response.done", ResponseDone) + assert done is not None + assert done.response.usage is not None, "response.done missing normalized usage" + + +@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +def test_tool_call_round_trip( + client: RealtimeClient, + scoped_key: str, + configured_models: frozenset[str], + provider: RealtimeProvider, +) -> None: + skip_if_unconfigured(provider, configured_models) + + with client.connect(key=scoped_key, model=provider.model) as session: + session.collect_until("session.created", timeout=20) + session.send( + SessionUpdate( + session=SessionConfig( + tools=[WEATHER_TOOL], + tool_choice="auto", + instructions=( + "Use the get_weather tool when asked about weather. " + "After receiving the result, state the temperature." + ), + ) + ) + ) + session.collect_until("session.updated", timeout=20) + + session.send(user_message("What's the weather in Paris right now?")) + session.send(ResponseCreate()) + first = session.collect_until("response.done", timeout=60) + + args_event = parse_last( + first, "response.function_call_arguments.done", FunctionCallArgumentsDone + ) + assert args_event is not None, "model did not emit a function call" + args = WeatherArgs.model_validate_json(args_event.arguments) + + item = function_call_item(first) + assert item is not None, "no completed function_call output item" + assert item.name == "get_weather" + assert item.call_id == args_event.call_id + + tool_result = WeatherResult(city=args.city, temperature_f=72) + session.send( + ConversationItemCreate( + item=FunctionCallOutputItem( + call_id=args_event.call_id, output=tool_result.model_dump_json() + ) + ) + ) + session.send(ResponseCreate()) + second = session.collect_until("response.done", timeout=60) + + assert "72" in transcript(second), "follow-up did not use the tool result" diff --git a/tests/e2e/realtime/test_realtime_pipecat_audio_e2e.py b/tests/e2e/realtime/test_realtime_pipecat_audio_e2e.py new file mode 100644 index 00000000000..d9d7c744f66 --- /dev/null +++ b/tests/e2e/realtime/test_realtime_pipecat_audio_e2e.py @@ -0,0 +1,349 @@ +"""Pipecat audio + server-VAD smoke tests for the proxy realtime websocket. + +Exercises the bot.py LiteLLMRealtimeLLMService pattern (simplified): + - custom _connect (no keepalive pings) + - _create_response override (tools session.update sent before history) + - _handle_evt_session_created (immediate session ready without waiting for + session.updated echo) + +Three test scenarios per provider: + test_pipecat_server_vad – session configured with server-VAD settings; + bot receives a text prompt and produces a reply. + test_pipecat_audio_output – same pipeline, asserts at least one + TTSAudioRawFrame with non-empty audio bytes. + test_pipecat_server_vad_audio_input – streams a real PCM16 audio fixture through + the pipeline without LLMRunFrame; server VAD + detects end-of-speech and auto-creates a response. +""" + +# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false +# pyright: reportUnknownArgumentType=false, reportAttributeAccessIssue=false +# pyright: reportUntypedBaseClass=false, reportUnknownParameterType=false +# pyright: reportMissingParameterType=false + +import asyncio +import wave +from pathlib import Path + +import pytest + +from realtime_client import ( + PROVIDERS, + RealtimeProvider, + _ws_base_url, + skip_if_unconfigured, +) + +pytestmark = pytest.mark.e2e + +pytest.importorskip("pipecat", reason="pipecat-ai not installed") + +from pipecat.adapters.schemas.function_schema import FunctionSchema # noqa: E402 +from pipecat.adapters.schemas.tools_schema import ToolsSchema # noqa: E402 +from pipecat.frames.frames import ( # noqa: E402 + EndFrame, + Frame, + InputAudioRawFrame, + LLMRunFrame, + TranscriptionFrame, + TTSAudioRawFrame, + TTSTextFrame, +) +from pipecat.pipeline.pipeline import Pipeline # noqa: E402 +from pipecat.pipeline.runner import PipelineRunner # noqa: E402 +from pipecat.pipeline.task import PipelineTask # noqa: E402 +from pipecat.processors.aggregators.llm_context import LLMContext # noqa: E402 +from pipecat.processors.aggregators.llm_response_universal import ( # noqa: E402 + LLMContextAggregatorPair, +) +from pipecat.processors.frame_processor import ( + FrameDirection, + FrameProcessor, +) # noqa: E402 +from pipecat.services.llm_service import FunctionCallParams # noqa: E402 +from pipecat.services.openai.realtime import events as rt_events # noqa: E402 +from pipecat.services.openai.realtime.llm import OpenAIRealtimeLLMService # noqa: E402 + +from pipecat_service import LiteLLMRealtimeLLMService # noqa: E402 + +PROVIDER_PARAMS = [pytest.param(p, id=p.id) for p in PROVIDERS] + +# PCM16 24 kHz mono WAV of "What is the weather in Paris?" (generated via macOS +# `say` and resampled with audioop). Used by the server-VAD audio-input test. +_FIXTURES_DIR = Path(__file__).parent / "fixtures" +WEATHER_WAV = _FIXTURES_DIR / "weather_question_24k.wav" + +# How long to stream silence after the speech ends so server VAD has time to +# detect the end-of-turn and fire a response. +_VAD_TAIL_SILENCE_MS = 1500 + +WEATHER_TOOL = ToolsSchema( + standard_tools=[ + FunctionSchema( + name="get_weather", + description="Get current temperature in Fahrenheit for a city.", + properties={"city": {"type": "string", "description": "City name."}}, + required=["city"], + ) + ] +) + +# Server-VAD session properties matching bot.py defaults. +SERVER_VAD_SETTINGS = rt_events.SessionProperties( + output_modalities=["audio"], + audio=rt_events.AudioConfiguration( + input=rt_events.AudioInput( + noise_reduction=rt_events.InputAudioNoiseReduction(type="near_field"), + turn_detection=rt_events.TurnDetection( + type="server_vad", + threshold=0.8, + prefix_padding_ms=300, + silence_duration_ms=700, + ), + ) + ), +) + + +# --------------------------------------------------------------------------- +# Helper frame-capture processor +# --------------------------------------------------------------------------- + + +class _CaptureFrames(FrameProcessor): + def __init__(self) -> None: + super().__init__() + self.texts: list[str] = [] + self.audio_bytes: int = 0 + + async def process_frame(self, frame: Frame, direction: FrameDirection) -> None: + await super().process_frame(frame, direction) + if isinstance(frame, TTSTextFrame): + self.texts.append(frame.text) + elif isinstance(frame, TTSAudioRawFrame): + self.audio_bytes += len(frame.audio) + await self.push_frame(frame, direction) + + +# --------------------------------------------------------------------------- +# Shared pipeline runner +# --------------------------------------------------------------------------- + + +async def _run_pipeline( + key: str, + model: str, + *, + prompt: str = "What is the weather in Paris?", + timeout: float = 45.0, +) -> tuple[bool, bool, int]: + """Run a minimal pipecat pipeline and return (tool_called, got_text, audio_bytes).""" + tool_called = asyncio.Event() + + async def get_weather(params: FunctionCallParams) -> None: + tool_called.set() + city = (params.arguments or {}).get("city", "Paris") + await params.result_callback({"city": city, "temperature_f": 72}) + + llm = LiteLLMRealtimeLLMService( + api_key=key, + base_url=f"{_ws_base_url()}/v1/realtime", + settings=OpenAIRealtimeLLMService.Settings( + model=model, + system_instruction=( + "You are a helpful assistant. " + "When asked about the weather, always call the get_weather tool. " + "Never guess temperatures." + ), + session_properties=SERVER_VAD_SETTINGS, + ), + ) + llm.register_function("get_weather", get_weather) + + context = LLMContext(tools=WEATHER_TOOL) + aggregator = LLMContextAggregatorPair(context) + capture = _CaptureFrames() + task = PipelineTask( + Pipeline([aggregator.user(), llm, capture, aggregator.assistant()]) + ) + + await task.queue_frames( + [ + TranscriptionFrame(prompt, user_id="e2e", timestamp=""), + LLMRunFrame(), + ] + ) + try: + await asyncio.wait_for(PipelineRunner().run(task), timeout=timeout) + except asyncio.TimeoutError: + await task.queue_frame(EndFrame()) + + return tool_called.is_set(), bool(capture.texts), capture.audio_bytes + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +def test_pipecat_server_vad( + scoped_key: str, + configured_models: frozenset[str], + provider: RealtimeProvider, +) -> None: + """Session is configured with server-VAD; bot must respond to a text prompt.""" + skip_if_unconfigured(provider, configured_models) + + tool_called, got_text, _ = asyncio.run(_run_pipeline(scoped_key, provider.model)) + + assert tool_called, "get_weather tool was not invoked" + assert got_text, "no assistant text frames produced" + + +@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +def test_pipecat_audio_output( + scoped_key: str, + configured_models: frozenset[str], + provider: RealtimeProvider, +) -> None: + """Bot must produce at least one non-empty TTS audio frame.""" + skip_if_unconfigured(provider, configured_models) + + _, got_text, audio_bytes = asyncio.run( + _run_pipeline( + scoped_key, + provider.model, + prompt="Say hello in one short sentence.", + timeout=30.0, + ) + ) + + assert got_text, "no assistant text frames produced" + assert audio_bytes > 0, "no TTS audio bytes received" + + +# --------------------------------------------------------------------------- +# Audio-input pipeline: streams a WAV fixture, server VAD fires the response +# --------------------------------------------------------------------------- + + +def _load_wav_chunks(path: Path, chunk_ms: int = 20) -> tuple[list[bytes], int]: + """Read a PCM16 mono WAV and split into ``chunk_ms``-sized byte chunks.""" + with wave.open(str(path), "rb") as wf: + assert wf.getnchannels() == 1, "fixture must be mono" + assert wf.getsampwidth() == 2, "fixture must be 16-bit PCM" + sample_rate = wf.getframerate() + frames_per_chunk = sample_rate * chunk_ms // 1000 + chunks = [] + while True: + data = wf.readframes(frames_per_chunk) + if not data: + break + chunks.append(data) + return chunks, sample_rate + + +async def _run_audio_input_pipeline( + key: str, + model: str, + *, + timeout: float = 60.0, +) -> tuple[bool, int]: + """Stream a WAV fixture as InputAudioRawFrame; return (got_text, audio_bytes). + + No LLMRunFrame is sent — server VAD is expected to detect end-of-speech + and auto-trigger a response. + + Audio is streamed at real-time pace (20 ms per chunk) after the session is + ready. Pre-queuing all frames at once floods the VAD buffer before the + backend is even connected and prevents speech_stopped from firing. + """ + chunks, sample_rate = _load_wav_chunks(WEATHER_WAV) + chunk_duration_s = 0.020 # 20 ms per chunk + + # Append silence after speech so VAD has enough quiet to fire. + silence_frames = sample_rate * _VAD_TAIL_SILENCE_MS // 1000 + silence_chunk = b"\x00" * silence_frames * 2 # 16-bit zero samples + chunks.append(silence_chunk) + + llm = LiteLLMRealtimeLLMService( + api_key=key, + base_url=f"{_ws_base_url()}/v1/realtime", + settings=OpenAIRealtimeLLMService.Settings( + model=model, + system_instruction=( + "You are a helpful assistant. " + "When asked about the weather, always call the get_weather tool. " + "Never guess temperatures." + ), + session_properties=SERVER_VAD_SETTINGS, + ), + ) + + tool_called = asyncio.Event() + + async def get_weather(params: FunctionCallParams) -> None: + tool_called.set() + city = (params.arguments or {}).get("city", "Paris") + await params.result_callback({"city": city, "temperature_f": 72}) + + llm.register_function("get_weather", get_weather) + + context = LLMContext(tools=WEATHER_TOOL) + aggregator = LLMContextAggregatorPair(context) + capture = _CaptureFrames() + task = PipelineTask( + Pipeline([aggregator.user(), llm, capture, aggregator.assistant()]) + ) + + async def _stream_audio() -> None: + # Wait for the LLM session to be ready before streaming so audio + # doesn't arrive before the backend WebSocket is connected. + for _ in range(100): + if getattr(llm, "_api_session_ready", False): + break + await asyncio.sleep(0.1) + + for chunk in chunks: + await task.queue_frame( + InputAudioRawFrame(audio=chunk, sample_rate=sample_rate, num_channels=1) + ) + await asyncio.sleep(chunk_duration_s) + + async def _run() -> None: + await asyncio.gather( + PipelineRunner().run(task), + _stream_audio(), + ) + + try: + await asyncio.wait_for(_run(), timeout=timeout) + except asyncio.TimeoutError: + await task.queue_frame(EndFrame()) + + return bool(capture.texts), capture.audio_bytes + + +@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +def test_pipecat_server_vad_audio_input( + scoped_key: str, + configured_models: frozenset[str], + provider: RealtimeProvider, +) -> None: + """Stream a real PCM16 WAV fixture; server VAD must detect speech end and respond. + + This exercises the full audio path: InputAudioRawFrame → input_audio_buffer.append + → server-VAD turn detection → response.create (auto) → assistant reply. + No LLMRunFrame is sent — the response must be triggered entirely by VAD. + """ + if not WEATHER_WAV.exists(): + pytest.skip(f"audio fixture not found: {WEATHER_WAV}") + skip_if_unconfigured(provider, configured_models) + + got_text, audio_bytes = asyncio.run( + _run_audio_input_pipeline(scoped_key, provider.model) + ) + + assert got_text, "server VAD did not trigger a response (no assistant text)" + assert audio_bytes > 0, "no TTS audio bytes received" diff --git a/tests/e2e/realtime/test_realtime_pipecat_e2e.py b/tests/e2e/realtime/test_realtime_pipecat_e2e.py new file mode 100644 index 00000000000..1068c54fdec --- /dev/null +++ b/tests/e2e/realtime/test_realtime_pipecat_e2e.py @@ -0,0 +1,134 @@ +"""Live pipecat smoke for the proxy realtime websocket. + +A realism layer on top of test_realtime_e2e: instead of speaking the GA protocol +by hand, drive the proxy through the shared LiteLLMRealtimeLLMService (pipecat's +GA service with proxy-specific overrides, keepalive pings disabled) with its +base_url pointed at the proxy and the model swapped per provider. It confirms the +audio and function-call wiring survives the round-trip. Assertions are coarse; +the raw-websocket suite is the source of truth. + +The harness is synchronous, so each test stays a normal sync function and drives +the async pipecat pipeline with asyncio.run. Skips unless pipecat is installed: + + uv pip install "pipecat-ai[openai]" + +Known caveat: pipecat tool calling over the realtime service has been flaky +upstream (pipecat-ai/pipecat#2544). A failure here with the matching raw-ws tool +test passing points at pipecat, not litellm. +""" + +# pipecat is an optional, dynamically typed dependency loaded behind importorskip, +# so its symbols are Unknown to the type checker; relax those rules for this file. +# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false, reportUnknownArgumentType=false, reportAttributeAccessIssue=false, reportUntypedBaseClass=false, reportUnknownParameterType=false, reportMissingParameterType=false + +import asyncio + +import pytest + +from realtime_client import ( + PROVIDERS, + RealtimeProvider, + _ws_base_url, + skip_if_unconfigured, +) + +pytestmark = pytest.mark.e2e + +pytest.importorskip("pipecat", reason="pipecat-ai not installed") + +from pipecat.adapters.schemas.function_schema import FunctionSchema # noqa: E402 +from pipecat.adapters.schemas.tools_schema import ToolsSchema # noqa: E402 +from pipecat.frames.frames import ( # noqa: E402 + EndFrame, + Frame, + LLMRunFrame, + TranscriptionFrame, + TTSTextFrame, +) +from pipecat.pipeline.pipeline import Pipeline # noqa: E402 +from pipecat.pipeline.runner import PipelineRunner # noqa: E402 +from pipecat.pipeline.task import PipelineTask # noqa: E402 +from pipecat.processors.aggregators.llm_context import LLMContext # noqa: E402 +from pipecat.processors.aggregators.llm_response_universal import ( # noqa: E402 + LLMContextAggregatorPair, +) +from pipecat.processors.frame_processor import ( # noqa: E402 + FrameDirection, + FrameProcessor, +) +from pipecat.services.llm_service import FunctionCallParams # noqa: E402 + +from pipecat_service import LiteLLMRealtimeLLMService # noqa: E402 + +PROVIDER_PARAMS = [pytest.param(p, id=p.id) for p in PROVIDERS] + +WEATHER_TOOL = ToolsSchema( + standard_tools=[ + FunctionSchema( + name="get_weather", + description="Get the current temperature in Fahrenheit for a city.", + properties={"city": {"type": "string"}}, + required=["city"], + ) + ] +) + + +class _CaptureText(FrameProcessor): + def __init__(self) -> None: + super().__init__() + self.texts: list[str] = [] + + async def process_frame(self, frame: Frame, direction: FrameDirection) -> None: + await super().process_frame(frame, direction) + if isinstance(frame, (TTSTextFrame, TranscriptionFrame)): + self.texts.append(frame.text) + await self.push_frame(frame, direction) + + +async def _run_pipeline(key: str, model: str) -> tuple[bool, bool]: + tool_called = asyncio.Event() + + async def get_weather(params: FunctionCallParams) -> None: + tool_called.set() + await params.result_callback({"city": "Paris", "temperature_f": 72}) + + llm = LiteLLMRealtimeLLMService( + api_key=key, base_url=f"{_ws_base_url()}/v1/realtime", model=model + ) + llm.register_function("get_weather", get_weather) + + context = LLMContext(tools=WEATHER_TOOL) + aggregator = LLMContextAggregatorPair(context) + capture = _CaptureText() + task = PipelineTask( + Pipeline([aggregator.user(), llm, capture, aggregator.assistant()]) + ) + + await task.queue_frames( + [ + TranscriptionFrame( + "What's the weather in Paris?", user_id="e2e", timestamp="" + ), + LLMRunFrame(), + ] + ) + try: + await asyncio.wait_for(PipelineRunner().run(task), timeout=45) + except asyncio.TimeoutError: + await task.queue_frame(EndFrame()) + return tool_called.is_set(), bool(capture.texts) + + +@pytest.mark.parametrize("provider", PROVIDER_PARAMS) +def test_pipecat_tool_smoke( + scoped_key: str, + configured_models: frozenset[str], + provider: RealtimeProvider, +) -> None: + skip_if_unconfigured(provider, configured_models) + + tool_called, produced_text = asyncio.run(_run_pipeline(scoped_key, provider.model)) + + assert tool_called, "pipecat did not invoke the get_weather callback" + assert produced_text, "pipecat produced no assistant text frames" diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 2c164f2169c..d23fc1a6778 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -17,6 +17,7 @@ from litellm.litellm_core_utils.realtime_streaming import ( RealTimeStreaming, client_sent_openai_beta_realtime_header, ) +from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import ( OpenAIRealtimeStreamResponseBaseObject, @@ -201,6 +202,347 @@ async def test_backend_to_client_skips_non_utf8_binary_frames(): assert isinstance(client_ws.send_text.call_args_list[0].args[0], str) +def _xai_streaming(client_ws=None, backend_ws=None, logging_obj=None): + """Helper: RealTimeStreaming wired with XAIRealtimeNormalizer.""" + return RealTimeStreaming( + client_ws or MagicMock(), + backend_ws or MagicMock(), + logging_obj or MagicMock(), + event_normalizer=XAIRealtimeNormalizer(), + ) + + +# --------------------------------------------------------------------------- +# XAIRealtimeNormalizer unit tests +# --------------------------------------------------------------------------- + + +def test_xai_normalizer_drops_ping(): + n = XAIRealtimeNormalizer() + assert n.should_drop({"type": "ping", "event_id": "x"}) + assert not n.should_drop({"type": "session.created", "session": {}}) + + +def test_xai_normalizer_converts_empty_response_created_usage_to_null(): + n = XAIRealtimeNormalizer() + event = { + "type": "response.created", + "response": { + "id": "r1", + "object": "realtime.response", + "output": [], + "status": "in_progress", + "status_details": None, + "usage": {}, + }, + } + assert n.normalize(event)["response"]["usage"] is None + + +def test_xai_normalizer_converts_empty_response_done_usage_to_defaults(): + n = XAIRealtimeNormalizer() + event = { + "type": "response.done", + "response": { + "id": "r1", + "object": "realtime.response", + "output": [], + "status": "completed", + "status_details": None, + "usage": {}, + }, + } + usage = n.normalize(event)["response"]["usage"] + assert usage is not None + assert usage["total_tokens"] == 0 + assert usage["input_token_details"]["text_tokens"] == 0 + + +def test_xai_normalizer_fills_partial_response_usage(): + n = XAIRealtimeNormalizer() + event = { + "type": "response.done", + "response": { + "id": "r1", + "object": "realtime.response", + "output": [], + "status": "completed", + "status_details": None, + "usage": {"total_tokens": 12, "input_tokens": 5, "output_tokens": 7}, + }, + } + usage = n.normalize(event)["response"]["usage"] + assert usage["total_tokens"] == 12 + assert usage["input_token_details"]["text_tokens"] == 0 + assert usage["output_token_details"]["audio_tokens"] == 0 + + +def test_xai_normalizer_injects_missing_content_part_done_part(): + n = XAIRealtimeNormalizer() + n._update_content_part_field( + {"response_id": "r1", "item_id": "i1", "content_index": 0, "transcript": "hi"}, + part_type="audio", + field="transcript", + value="hi", + ) + event = { + "type": "response.content_part.done", + "response_id": "r1", + "item_id": "i1", + "content_index": 0, + "output_index": 0, + } + assert n.normalize(event)["part"] == {"type": "audio", "transcript": "hi"} + + +def test_xai_normalizer_conversation_item_tool_role_becomes_assistant(): + n = XAIRealtimeNormalizer() + event = { + "type": "conversation.item.added", + "event_id": "e1", + "previous_item_id": None, + "item": { + "id": "i1", + "object": "realtime.item", + "type": "function_call", + "status": "in_progress", + "role": "tool", + "call_id": "c1", + "name": "get_weather", + "arguments": "", + }, + } + normalized = n.normalize(event) + assert normalized["item"]["role"] == "assistant" + assert normalized["item"]["name"] == "get_weather" + + +def test_xai_normalizer_injects_output_index_on_function_call_delta(): + n = XAIRealtimeNormalizer() + event = { + "type": "response.function_call_arguments.delta", + "event_id": "e1", + "item_id": "i1", + "response_id": "r1", + "delta": '{"city":"Paris"}', + "call_id": "c1", + "previous_item_id": None, + } + normalized = n.normalize(event) + assert normalized["output_index"] == 0 + assert "content_index" not in normalized + + +def test_xai_normalizer_injects_both_indices_on_content_part_done(): + n = XAIRealtimeNormalizer() + event = { + "type": "response.content_part.done", + "event_id": "e2", + "item_id": "i1", + "response_id": "r1", + "previous_item_id": None, + } + normalized = n.normalize(event) + assert normalized["output_index"] == 0 + assert normalized["content_index"] == 0 + + +def test_xai_normalizer_preserves_existing_indices(): + n = XAIRealtimeNormalizer() + event = { + "type": "response.function_call_arguments.done", + "event_id": "e3", + "item_id": "i1", + "response_id": "r1", + "output_index": 2, + "call_id": "c1", + "arguments": '{"city":"Paris"}', + } + assert n.normalize(event)["output_index"] == 2 + + +def test_xai_patch_outgoing_session_defaults_create_response_flat(): + n = XAIRealtimeNormalizer() + session = { + "turn_detection": { + "type": "server_vad", + "threshold": 0.8, + "silence_duration_ms": 700, + } + } + patched = n.patch_outgoing_session(session) + assert patched["turn_detection"]["create_response"] is True + + +def test_xai_patch_outgoing_session_defaults_create_response_nested(): + n = XAIRealtimeNormalizer() + session = { + "audio": { + "input": { + "turn_detection": { + "type": "server_vad", + "threshold": 0.8, + } + } + } + } + patched = n.patch_outgoing_session(session) + assert patched["audio"]["input"]["turn_detection"]["create_response"] is True + + +def test_xai_patch_outgoing_session_respects_explicit_create_response_false(): + n = XAIRealtimeNormalizer() + session = { + "turn_detection": { + "type": "server_vad", + "create_response": False, + } + } + patched = n.patch_outgoing_session(session) + assert patched["turn_detection"]["create_response"] is False + + +def test_xai_patch_outgoing_session_ignores_non_server_vad(): + n = XAIRealtimeNormalizer() + session = {"turn_detection": {"type": "semantic_vad"}} + patched = n.patch_outgoing_session(session) + assert "create_response" not in patched["turn_detection"] + + +# --------------------------------------------------------------------------- +# Integration: RealTimeStreaming with XAIRealtimeNormalizer +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_backend_to_client_drops_ping_events(): + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + backend_ws = MagicMock() + backend_ws.recv = AsyncMock( + side_effect=[ + json.dumps( + {"type": "ping", "event_id": "evt_ping", "timestamp": 1782214899793} + ).encode(), + json.dumps({"type": "session.created", "session": {}}).encode(), + ConnectionClosed(None, None), + ] + ) + logging_obj = MagicMock() + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + streaming = _xai_streaming(client_ws, backend_ws, logging_obj) + + await streaming.backend_to_client_send_messages() + + assert client_ws.send_text.call_count == 1 + sent = json.loads(client_ws.send_text.call_args_list[0].args[0]) + assert sent["type"] == "session.created" + + +@pytest.mark.asyncio +async def test_backend_to_client_normalizes_empty_response_usage(): + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + backend_ws = MagicMock() + backend_ws.recv = AsyncMock( + side_effect=[ + json.dumps( + { + "type": "response.created", + "response": { + "id": "r1", + "object": "realtime.response", + "output": [], + "status": "in_progress", + "status_details": "unimplemented", + "usage": {}, + }, + } + ).encode(), + ConnectionClosed(None, None), + ] + ) + logging_obj = MagicMock() + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + streaming = _xai_streaming(client_ws, backend_ws, logging_obj) + + await streaming.backend_to_client_send_messages() + + sent = json.loads(client_ws.send_text.call_args_list[0].args[0]) + assert sent["response"]["usage"] is None + + +@pytest.mark.asyncio +async def test_backend_to_client_beta_receives_normalized_events(): + client_ws = MagicMock() + client_ws.scope = {"headers": [(b"openai-beta", b"realtime=v1")]} + client_ws.send_text = AsyncMock() + backend_ws = MagicMock() + backend_ws.recv = AsyncMock( + side_effect=[ + json.dumps( + { + "type": "response.function_call_arguments.delta", + "event_id": "e1", + "item_id": "i1", + "response_id": "r1", + "delta": '{"city":"Paris"}', + "call_id": "c1", + "previous_item_id": None, + } + ).encode(), + ConnectionClosed(None, None), + ] + ) + logging_obj = MagicMock() + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + streaming = _xai_streaming(client_ws, backend_ws, logging_obj) + + await streaming.backend_to_client_send_messages() + + sent = json.loads(client_ws.send_text.call_args_list[0].args[0]) + assert sent["type"] == "response.function_call_arguments.delta" + assert sent["output_index"] == 0 + + +@pytest.mark.asyncio +async def test_backend_to_client_stores_normalized_events_for_logging(): + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + backend_ws = MagicMock() + backend_ws.recv = AsyncMock( + side_effect=[ + json.dumps( + { + "type": "response.done", + "response": { + "id": "r1", + "object": "realtime.response", + "output": [], + "status": "completed", + "status_details": None, + "usage": {}, + }, + } + ).encode(), + ConnectionClosed(None, None), + ] + ) + logging_obj = MagicMock() + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + streaming = _xai_streaming(client_ws, backend_ws, logging_obj) + + await streaming.backend_to_client_send_messages() + + sent = json.loads(client_ws.send_text.call_args_list[0].args[0]) + assert sent["response"]["usage"]["total_tokens"] == 0 + assert streaming.messages[0]["response"]["usage"]["total_tokens"] == 0 + + @pytest.mark.asyncio async def test_client_ack_messages_keeps_beta_session_shape_for_beta_clients(): client_ws = MagicMock() @@ -849,6 +1191,47 @@ async def test_client_ack_caches_setup_to_prevent_duplicate_session_update_setup assert "setup" in sent_payload +@pytest.mark.asyncio +async def test_failed_content_send_does_not_block_later_setup(): + """A content frame whose backend send fails must not flip + ``_content_sent_after_setup``; otherwise a later setup is silently dropped + even though the backend never received any content.""" + websocket = MagicMock() + backend_ws = MagicMock() + logging_obj = MagicMock() + + provider_config = MagicMock() + provider_config.transform_realtime_request = MagicMock( + side_effect=lambda m, *a, **k: [m] + ) + provider_config.is_setup_message = MagicMock(side_effect=lambda obj: "setup" in obj) + provider_config.is_content_message = MagicMock( + side_effect=lambda obj: obj.get("type") == "conversation.item.create" + ) + + backend_ws.send = AsyncMock(side_effect=[ConnectionClosed(None, None), None]) + + streaming = RealTimeStreaming( + websocket=websocket, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + + content = json.dumps({"type": "conversation.item.create", "item": {}}) + with pytest.raises(ConnectionClosed): + await streaming._send_to_backend(content) + + assert streaming._content_sent_after_setup is False + + setup = json.dumps({"setup": {"model": "models/gemini-2.5-flash"}}) + sent = await streaming._send_to_backend(setup) + + assert sent is True + assert backend_ws.send.await_args_list[-1].args[0] == setup + + def test_collect_session_tools_from_session_update(): """ Test that tools from session.update events are collected. diff --git a/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py b/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py index 85da855f8ef..0f9edb0faf8 100644 --- a/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py +++ b/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py @@ -516,6 +516,75 @@ def test_gemini_session_update_defaults_to_audio_modality(): assert setup_payload["generationConfig"]["responseModalities"] == ["AUDIO"] +@pytest.mark.parametrize( + "model", + [ + "gemini-2.5-flash-native-audio", + "gemini-3.1-flash-live-preview", + "gemini/gemini-3.1-flash-live-preview", + ], +) +def test_gemini_audio_only_live_models_coerce_text_modality_to_audio(model, patch_gemini_audio_cost_map_entries): + """Regression: TEXT-only responseModalities causes 1007 on audio-only Live models.""" + config = GeminiRealtimeConfig() + session_update = { + "type": "session.update", + "session": { + "modalities": ["text"], + "instructions": "You are a terse assistant.", + }, + } + + messages = config.transform_realtime_request( + json.dumps(session_update), + model, + session_configuration_request=None, + ) + + setup = json.loads(messages[0])["setup"] + assert setup["generationConfig"]["responseModalities"] == ["AUDIO"] + + +def test_gemini_audio_only_live_models_drop_text_from_text_audio_combo(patch_gemini_audio_cost_map_entries): + config = GeminiRealtimeConfig() + session_update = { + "type": "session.update", + "session": { + "modalities": ["text", "audio"], + "instructions": "Be concise.", + }, + } + + messages = config.transform_realtime_request( + json.dumps(session_update), + "gemini-3.1-flash-live-preview", + session_configuration_request=None, + ) + + setup = json.loads(messages[0])["setup"] + assert setup["generationConfig"]["responseModalities"] == ["AUDIO"] + + +def test_gemini_non_live_model_preserves_text_modality(): + config = GeminiRealtimeConfig() + session_update = { + "type": "session.update", + "session": { + "modalities": ["text"], + "instructions": "You are a terse assistant.", + }, + } + + messages = config.transform_realtime_request( + json.dumps(session_update), + "gemini-2.5-flash", + session_configuration_request=None, + ) + + setup = json.loads(messages[0])["setup"] + assert setup["generationConfig"]["responseModalities"] == ["TEXT"] + + def test_gemini_requires_session_configuration_feature_flag(monkeypatch): config = GeminiRealtimeConfig() @@ -1199,7 +1268,7 @@ def test_gemini_subsequent_session_update_forwards_tools_merged_with_original_se assert follow_up["inputAudioTranscription"] == {} -def test_gemini_realtime_pipecat_ga_session_voice_and_tools(): +def test_gemini_realtime_pipecat_ga_session_voice_and_tools(patch_gemini_audio_cost_map_entries): """Pipecat OpenAIRealtimeSessionProperties: output_modalities, nested tools, and audio.output.voice (e.g. Kore) must map into Gemini setup.""" config = GeminiRealtimeConfig() @@ -1732,3 +1801,210 @@ def test_gemini_in_frame_usage_metadata_clears_pending_buffer(): assert usage["output_tokens"] == 2 assert usage["total_tokens"] == 5 assert config._pending_usage_metadata is None + + +def test_gemini_post_tool_bare_turn_complete_followed_by_answer(): + """After a tool call, Gemini Live can emit a bare ``turnComplete`` (with + usage but no model content) before the follow-up answer stream. That bare + ``turnComplete`` may produce an extra ``response.done``; Pipecat is tolerant + of that because ``_process_completed_function_calls`` is idempotent (the + pending call queue is empty by the time the second ``response.done`` arrives). + The important thing is that the post-tool answer is correctly generated.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_post_tool_bare_turn_complete" + + session_configuration_request = json.dumps( + { + "setup": { + "model": "gemini-2.5-flash-native-audio", + "generationConfig": {"responseModalities": ["AUDIO"]}, + } + } + ) + base_input = { + "session_configuration_request": session_configuration_request, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + } + + tool_result = config.transform_realtime_response( + json.dumps( + { + "toolCall": { + "functionCalls": [ + { + "id": "call_post_tool", + "name": "get_weather", + "args": {"city": "Paris"}, + } + ] + } + } + ), + "gemini-2.5-flash-native-audio", + logging_obj, + realtime_response_transform_input=base_input, + ) + assert tool_result["response"][-1]["type"] == "response.done" + + bare_turn_complete = config.transform_realtime_response( + json.dumps( + { + "serverContent": {"turnComplete": True}, + "usageMetadata": { + "promptTokenCount": 30, + "responseTokenCount": 5, + "totalTokenCount": 35, + }, + } + ), + "gemini-2.5-flash-native-audio", + logging_obj, + realtime_response_transform_input={ + **base_input, + "current_output_item_id": tool_result["current_output_item_id"], + "current_response_id": tool_result["current_response_id"], + "current_conversation_id": tool_result["current_conversation_id"], + "current_delta_chunks": tool_result["current_delta_chunks"], + "current_item_chunks": tool_result["current_item_chunks"], + "current_delta_type": tool_result["current_delta_type"], + }, + ) + # The bare turnComplete must not surface as a response.done because clients + # that use collect_until("response.done") would stop collecting prematurely + # before the real follow-up answer arrives. + assert bare_turn_complete["response"] == [] + + post_tool_answer = config.transform_realtime_response( + json.dumps( + { + "serverContent": { + "outputTranscription": {"text": "The temperature is 72."}, + "modelTurn": { + "parts": [ + { + "inlineData": { + "mimeType": "audio/pcm", + "data": "audio-chunk", + } + } + ] + }, + } + } + ), + "gemini-2.5-flash-native-audio", + logging_obj, + realtime_response_transform_input={ + **base_input, + "current_output_item_id": bare_turn_complete["current_output_item_id"], + "current_response_id": bare_turn_complete["current_response_id"], + "current_conversation_id": bare_turn_complete["current_conversation_id"], + "current_delta_chunks": bare_turn_complete["current_delta_chunks"], + "current_item_chunks": bare_turn_complete["current_item_chunks"], + "current_delta_type": bare_turn_complete["current_delta_type"], + }, + ) + assert post_tool_answer["response"][0]["type"] == "response.created" + transcript_delta = next( + event + for event in post_tool_answer["response"] + if event["type"] == "response.output_audio_transcript.delta" + ) + assert "72" in transcript_delta["delta"] + + final_turn = config.transform_realtime_response( + json.dumps({"serverContent": {"turnComplete": True}}), + "gemini-2.5-flash-native-audio", + logging_obj, + realtime_response_transform_input={ + **base_input, + "current_output_item_id": post_tool_answer["current_output_item_id"], + "current_response_id": post_tool_answer["current_response_id"], + "current_conversation_id": post_tool_answer["current_conversation_id"], + "current_delta_chunks": post_tool_answer["current_delta_chunks"], + "current_item_chunks": post_tool_answer["current_item_chunks"], + "current_delta_type": post_tool_answer["current_delta_type"], + }, + ) + response_done = next( + event + for event in final_turn["response"] + if event["type"] == "response.done" + ) + assert response_done["response"]["status"] == "completed" + + +@pytest.fixture(autouse=False) +def patch_gemini_audio_cost_map_entries(monkeypatch): + """Inject gemini_native_audio / gemini_audio_only_live into the cost map. + + litellm.model_cost is fetched from main branch at import time, so in CI + the fields may not exist yet. Patch locally so these tests are + self-contained. + """ + native_audio_models = [ + "gemini-2.5-flash-native-audio", + "gemini-2.5-flash-native-audio-latest", + "gemini/gemini-2.5-flash-native-audio-latest", + ] + flash_live_models = [ + "gemini-3.1-flash-live-preview", + "gemini/gemini-3.1-flash-live-preview", + ] + for m in native_audio_models: + entry = dict(litellm.model_cost.get(m, {})) + entry["gemini_native_audio"] = True + monkeypatch.setitem(litellm.model_cost, m, entry) + for m in flash_live_models: + entry = dict(litellm.model_cost.get(m, {})) + entry["gemini_audio_only_live"] = True + monkeypatch.setitem(litellm.model_cost, m, entry) + + +@pytest.mark.parametrize( + "model,expected", + [ + ("gemini-3.1-flash-live-preview", True), + ("gemini/gemini-3.1-flash-live-preview", True), + ("gemini-2.5-flash-native-audio-latest", True), + ("gemini/gemini-2.5-flash-native-audio-latest", True), + ("gemini-2.0-flash", False), + ("gemini-2.5-flash", False), + ], +) +def test_is_audio_only_live_model_uses_cost_map( + model, expected, patch_gemini_audio_cost_map_entries +): + assert GeminiRealtimeConfig._is_audio_only_live_model(model) == expected + + +@pytest.mark.parametrize( + "model,expected", + [ + ("gemini-2.5-flash-native-audio-latest", True), + ("gemini/gemini-2.5-flash-native-audio-latest", True), + ("gemini-3.1-flash-live-preview", False), + ("gemini/gemini-3.1-flash-live-preview", False), + ("gemini-2.0-flash", False), + ], +) +def test_is_native_audio_model_uses_cost_map( + model, expected, patch_gemini_audio_cost_map_entries +): + assert GeminiRealtimeConfig._is_native_audio_model(model) == expected + + +def test_is_setup_message_and_is_content_message(): + config = GeminiRealtimeConfig() + assert config.is_setup_message({"setup": {}}) is True + assert config.is_setup_message({"realtimeInput": {}}) is False + assert config.is_content_message({"realtimeInput": {}}) is True + assert config.is_content_message({"clientContent": {}}) is True + assert config.is_content_message({"toolResponse": {}}) is True + assert config.is_content_message({"setup": {}}) is False diff --git a/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py b/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py index 1f171496cce..f11b00d204d 100644 --- a/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py @@ -112,7 +112,7 @@ def test_vertex_session_update_defaults_to_audio_modality(): messages = cfg.transform_realtime_request( json.dumps(session_update), - "gemini-live-2.5-flash-native-audio", + "gemini-live-2.5-flash-preview-native-audio-09-2025", session_configuration_request=None, ) assert len(messages) == 1 @@ -120,7 +120,50 @@ def test_vertex_session_update_defaults_to_audio_modality(): assert setup_payload["generationConfig"]["responseModalities"] == ["AUDIO"] -def test_vertex_session_update_normalizes_ga_remapped_fields(): +_NATIVE_AUDIO_MODEL = "gemini-live-2.5-flash-preview-native-audio-09-2025" + + +@pytest.fixture(autouse=False) +def patch_native_audio_cost_map_entry(monkeypatch): + """Inject gemini_native_audio into the cost map for the test model. + + litellm.model_cost is fetched from main branch at import time, so in CI + the field may not exist yet. Patch it locally so these unit tests remain + self-contained and don't depend on the remote cost map state. + """ + entry = dict(litellm.model_cost.get(_NATIVE_AUDIO_MODEL, {})) + entry["gemini_native_audio"] = True + monkeypatch.setitem(litellm.model_cost, _NATIVE_AUDIO_MODEL, entry) + + +def test_vertex_audio_only_live_model_coerces_text_modality_to_audio( + patch_native_audio_cost_map_entry, +): + """Regression: TEXT-only responseModalities causes 1007 on native-audio Live models.""" + cfg = VertexAIRealtimeConfig( + access_token="tok", project="my-proj", location="us-central1" + ) + session_update = { + "type": "session.update", + "session": { + "modalities": ["text"], + "instructions": "You are a terse assistant.", + }, + } + + messages = cfg.transform_realtime_request( + json.dumps(session_update), + _NATIVE_AUDIO_MODEL, + session_configuration_request=None, + ) + + setup = json.loads(messages[0])["setup"] + assert setup["generationConfig"]["responseModalities"] == ["AUDIO"] + + +def test_vertex_session_update_normalizes_ga_remapped_fields( + patch_native_audio_cost_map_entry, +): """GA-format clients send ``output_modalities`` and nested ``audio.input.transcription`` / ``audio.input.turn_detection``. These must be normalised back to the flat beta keys before ``map_openai_params`` @@ -146,13 +189,13 @@ def test_vertex_session_update_normalizes_ga_remapped_fields(): messages = cfg.transform_realtime_request( json.dumps(session_update), - "gemini-live-2.5-flash-native-audio", + _NATIVE_AUDIO_MODEL, session_configuration_request=None, ) assert len(messages) == 1 setup_payload = json.loads(messages[0])["setup"] - assert setup_payload["generationConfig"]["responseModalities"] == ["TEXT"] + assert setup_payload["generationConfig"]["responseModalities"] == ["AUDIO"] assert setup_payload["inputAudioTranscription"] == {} assert ( setup_payload["realtimeInputConfig"]["automaticActivityDetection"][ @@ -309,7 +352,7 @@ def test_vertex_warns_when_dropping_guardrail_turn_detection_update(caplog): with caplog.at_level(logging.WARNING, logger="LiteLLM"): result = cfg.transform_realtime_request( json.dumps(session_update), - "gemini-live-2.5-flash-native-audio", + "gemini-live-2.5-flash-preview-native-audio-09-2025", session_configuration_request=json.dumps({"setup": {"model": "x"}}), ) @@ -338,7 +381,7 @@ def test_vertex_does_not_warn_when_dropping_non_guardrail_session_update(caplog) with caplog.at_level(logging.WARNING, logger="LiteLLM"): cfg.transform_realtime_request( json.dumps(session_update), - "gemini-live-2.5-flash-native-audio", + "gemini-live-2.5-flash-preview-native-audio-09-2025", session_configuration_request=json.dumps({"setup": {"model": "x"}}), ) @@ -348,6 +391,7 @@ def test_vertex_does_not_warn_when_dropping_non_guardrail_session_update(caplog) ) +@pytest.mark.asyncio async def test_async_realtime_does_not_forward_client_query_params_to_vertex_backend( monkeypatch, ): @@ -364,7 +408,7 @@ async def test_async_realtime_does_not_forward_client_query_params_to_vertex_bac access_token="tok", project="my-proj", location="us-central1" ) - captured = {} + captured: dict = {} def fake_connect(url, *args, **kwargs): captured["url"] = url @@ -372,18 +416,22 @@ async def test_async_realtime_does_not_forward_client_query_params_to_vertex_bac monkeypatch.setattr(websockets, "connect", fake_connect) - await BaseLLMHTTPHandler().async_realtime( - model="gemini-live-2.5-flash-native-audio", - websocket=AsyncMock(), - logging_obj=MagicMock(), - provider_config=cfg, - headers={}, - query_params={ - "model": "gemini-live-2.5-flash-native-audio", - "intent": "chat", - }, - ) + try: + await BaseLLMHTTPHandler().async_realtime( + model="gemini-live-2.5-flash-preview-native-audio-09-2025", + websocket=AsyncMock(), + logging_obj=MagicMock(), + provider_config=cfg, + headers={}, + query_params={ + "model": "gemini-live-2.5-flash-preview-native-audio-09-2025", + "intent": "chat", + }, + ) + except (RuntimeError, Exception): + pass + assert "url" in captured, "websockets.connect was never called" assert "?" not in captured["url"] assert "model=" not in captured["url"] assert "intent=" not in captured["url"] @@ -407,7 +455,7 @@ def test_vertex_function_call_output_omits_id(): }, } ), - "gemini-live-2.5-flash-native-audio", + "gemini-live-2.5-flash-preview-native-audio-09-2025", session_configuration_request="existing", ) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index cf7afc1af68..ecde529bbfa 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -836,6 +836,8 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "supports_assistant_prefill": {"type": "boolean"}, "supports_audio_input": {"type": "boolean"}, "supports_audio_output": {"type": "boolean"}, + "gemini_native_audio": {"type": "boolean"}, + "gemini_audio_only_live": {"type": "boolean"}, "supports_embedding_image_input": {"type": "boolean"}, "supports_code_execution": {"type": "boolean"}, "supports_file_search": {"type": "boolean"},