diff --git a/litellm/__init__.py b/litellm/__init__.py index 7c92623358d..56d516536e8 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -225,6 +225,11 @@ use_chat_completions_url_for_anthropic_messages: bool = bool( route_all_chat_openai_to_responses: bool = ( os.getenv("LITELLM_ROUTE_ALL_CHAT_OPENAI_TO_RESPONSES", "false").lower() == "true" ) # When True, routes all OpenAI /chat/completions requests through the Responses API bridge +# When True, Gemini/Vertex Live setup is deferred until client `session.update`. +# Default False preserves historical behavior (auto-send setup on connect). +gemini_live_defer_setup: bool = ( + os.getenv("LITELLM_GEMINI_LIVE_DEFER_SETUP", "false").lower() == "true" +) use_legacy_interactions_schema: bool = ( os.getenv("LITELLM_USE_LEGACY_INTERACTIONS_SCHEMA", "false").lower() == "true" ) # When True, sends Api-Revision: 2026-05-07 to Google so responses use the legacy `outputs` diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index f9f1bf5c56b..33bb6d7d2ea 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -86,6 +86,12 @@ class RealTimeStreaming: # When a text message is blocked, hold the guardrail reason so the next # response.create can be rewritten to include the failure context. self._pending_guardrail_message: Optional[str] = None + # Track whether session.created has already been sent to the client + # (e.g. synthetic event in deferred setup mode). + self._session_created_sent_to_client: bool = False + # Track whether we have already sent the guardrail turn-detection update + # that disables provider auto-response for transcription guardrails. + self._guardrail_turn_detection_update_sent: bool = False _SESSION_EVENT_TYPES = frozenset(["session.created", "session.updated"]) _AUDIO_FORMAT_MAP: Dict[str, Dict[str, Any]] = { @@ -248,22 +254,52 @@ class RealTimeStreaming: ## SYNC LOGGING executor.submit(self.logging_obj.success_handler(self.messages)) - async def _send_to_backend(self, message: str) -> None: + async def _send_to_backend(self, message: str) -> bool: """Send a message to the backend WebSocket. If a provider_config is set the message is first passed through transform_realtime_request so that provider-specific translation (e.g. dropping session.update for Vertex AI) is applied even for guardrail-injected messages. + + Returns True if at least one message was actually delivered to the + backend, False if the provider transformation produced no output and + the message was effectively dropped. """ if self.provider_config: transformed = self.provider_config.transform_realtime_request( message, self.model, self.session_configuration_request ) + 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] - else: - await self.backend_ws.send(message) # type: ignore[union-attr, attr-defined] + self._cache_session_configuration_request(msg) + sent = True + return sent + await self.backend_ws.send(message) # type: ignore[union-attr, attr-defined] + return True + + def _cache_session_configuration_request(self, transformed_message: str) -> None: + """Store setup payload once sent to backend. + + Updates the cached setup on every successful setup send so follow-up + ``session.update`` messages (which produce a merged setup with new + ``generationConfig`` / ``systemInstruction`` / etc.) are reflected in + the cache used by downstream readers (``transform_session_created_event``, + ``return_new_content_delta_events`` modality lookup, ...). + """ + try: + message_obj = json.loads(transformed_message) + if "setup" in message_obj: + self.session_configuration_request = transformed_message + except (json.JSONDecodeError, TypeError): + return def _make_disable_auto_response_message(self) -> str: """Return a session.update that disables VAD auto-response.""" @@ -280,6 +316,20 @@ class RealTimeStreaming: } return json.dumps({"type": "session.update", "session": session}) + async def _maybe_send_guardrail_turn_detection_update(self) -> None: + """Disable provider auto-response once when transcription guardrails are enabled.""" + if self._guardrail_turn_detection_update_sent: + return + if not self._has_audio_transcription_guardrails(): + return + sent = await self._send_to_backend(self._make_disable_auto_response_message()) + # Only mark as sent when the provider transformation actually delivered + # the update to the backend. Otherwise (e.g. Gemini drops session.update + # after the initial setup), leave the flag unset so future opportunities + # — such as a duplicate session.created — can retry. + if sent: + self._guardrail_turn_detection_update_sent = True + def _has_realtime_guardrails(self) -> bool: """Return True if any callback is registered for realtime guardrail event types.""" from litellm.integrations.custom_guardrail import CustomGuardrail @@ -318,12 +368,20 @@ class RealTimeStreaming: self, transcript: str, item_id: Optional[str] = None, + pre_block_backend_message: Optional[str] = None, ) -> bool: """ Run registered guardrails on a completed speech transcription. Returns True if blocked (synthetic warning already sent to client). Returns False if clean (caller should send response.create to the backend). + + ``pre_block_backend_message`` (if provided) is sent to the backend + BEFORE any of the guardrail's own backend messages when a block is + triggered. This is needed for protocol contracts that require a + specific message to be sent first — e.g. Gemini Live requires a + matching ``toolResponse`` immediately after a ``toolCall`` before any + other client messages can be accepted. """ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.types.guardrails import GuardrailEventHooks @@ -383,6 +441,13 @@ class RealTimeStreaming: getattr(callback, "realtime_violation_message", None) or safe_msg ) + # Deliver any caller-supplied backend message FIRST so that + # protocol contracts requiring a specific ordering (e.g. + # Gemini Live's mandatory ``toolResponse`` after a + # ``toolCall``) are honored before the guardrail's own + # clientContent / cancel messages are sent. + if pre_block_backend_message is not None: + await self._send_to_backend(pre_block_backend_message) # Cancel any in-progress LLM response (e.g. VAD auto-response). await self._send_to_backend(json.dumps({"type": "response.cancel"})) # Send the policy violation hint (shows as small gray status text in UI). @@ -478,16 +543,34 @@ class RealTimeStreaming: else [transformed_response] ) for event in events: + is_session_created_event = ( + isinstance(event, dict) and event.get("type") == "session.created" + ) + if is_session_created_event: + if self._session_created_sent_to_client: + # A synthetic session.created (with placeholder defaults) was + # already forwarded to the client when we connected. The + # provider's real session.created (e.g. emitted from Gemini + # `setupComplete`) carries the authoritative modalities/model + # from the client's session.update. Re-emit it as + # `session.updated` so the client learns the corrected + # configuration without seeing two `session.created` events. + event = {**event, "type": "session.updated"} + else: + self._session_created_sent_to_client = True event_str = json.dumps(event) - ## For audio/VAD guardrail path: forward session.created first, then inject. - if ( - isinstance(event, dict) - and event.get("type") == "session.created" - and self._has_audio_transcription_guardrails() - ): + ## For audio/VAD guardrail path: forward the (possibly retyped) + ## session.created first, then invoke the one-time guardrail + ## turn-detection update. ``_maybe_send_guardrail_turn_detection_update`` + ## is idempotent (gated by ``_guardrail_turn_detection_update_sent``), + ## so duplicate session.created events — including those emitted + ## after a synthetic session.created from ``llm_http_handler`` in + ## deferred-setup mode — still get a single chance to inject the + ## update if a prior attempt was dropped by the provider transform. + if is_session_created_event and self._has_audio_transcription_guardrails(): self.store_message(event_str) await self.websocket.send_text(event_str) - await self._send_to_backend(self._make_disable_auto_response_message()) + await self._maybe_send_guardrail_turn_detection_update() continue ## GUARDRAIL: run on transcription events in provider_config path too if ( @@ -790,12 +873,13 @@ class RealTimeStreaming: item["content"] = new_content return item - async def client_ack_messages(self): + async def client_ack_messages(self): # noqa: PLR0915 try: while True: message = await self.websocket.receive_text() ## GUARDRAIL: intercept conversation.item.create for text-based injection. + guardrail_turn_detection_injected = False try: msg_obj = json.loads(message) msg_type = msg_obj.get("type") @@ -803,7 +887,68 @@ class RealTimeStreaming: if msg_type == "conversation.item.create": # Check user text messages for prompt injection item = msg_obj.get("item", {}) - if item.get("role") == "user": + # Check function_call_output first so a client cannot + # bypass the tool-result guardrail by also setting + # role="user" on a function_call_output item. + if item.get("type") == "function_call_output": + # Tool results are client-controlled and fed to the + # model; check them with the same guardrail used for + # user text so an attacker cannot smuggle blocked + # content into a function_call_output. + output = item.get("output", "") + output_text = ( + output + if isinstance(output, str) + else json.dumps(output) + ) + if output_text: + # Build the sanitized function_call_output up + # front so we can hand it to the guardrail + # runner as the pre-block message. Providers + # that pair every toolCall with a toolResponse + # (e.g. Gemini/Vertex Live) require the + # toolResponse to arrive BEFORE any other + # client message — otherwise the guardrail's + # own clientContent would violate the + # pending-tool-call protocol contract and the + # backend could close the connection before + # the sanitized response ever lands. Dropping + # the blocked item outright would similarly + # leave such providers waiting indefinitely. + # The sanitized payload carries no blocked + # content — only a generic policy marker. + sanitized_msg = json.dumps( + { + **msg_obj, + "item": { + **item, + "output": json.dumps( + { + "error": "Tool output blocked by content policy", + } + ), + }, + } + ) + blocked = await self.run_realtime_guardrails( + output_text, + pre_block_backend_message=sanitized_msg, + ) + if blocked: + # ``_pending_guardrail_message`` is + # intentionally NOT set here. That flag + # exists to swallow the reflexive + # ``response.create`` an OpenAI client + # sends immediately after a user text + # message. In a tool-calling flow the + # client may not send a ``response.create`` + # at all (e.g. Gemini SDKs auto-respond), + # so leaving the flag set would + # incorrectly drop an unrelated + # ``response.create`` from a later + # interaction turn. + continue + elif item.get("role") == "user": content_list = item.get("content", []) texts = [ c.get("text", "") @@ -831,6 +976,89 @@ class RealTimeStreaming: self._pending_guardrail_message = None continue + ## GUARDRAIL: Inject turn_detection into first session.update + # if needed. Done BEFORE the GA remap so the injected + # ``create_response`` rides along with any client-provided + # turn_detection fields (e.g. silence_duration_ms) into the + # nested ``audio.input.turn_detection`` path produced by the + # remap. Doing this after the remap would create a separate + # minimal root-level ``turn_detection`` and silently drop + # the client's nested settings. + if ( + msg_type == "session.update" + and self.session_configuration_request is None + and not self._guardrail_turn_detection_update_sent + and self._has_audio_transcription_guardrails() + ): + session = msg_obj.setdefault("session", {}) + if isinstance(session, dict): + existing_td = session.get("turn_detection") + if not isinstance(existing_td, dict): + existing_td = {} + existing_td["create_response"] = False + session["turn_detection"] = existing_td + message = json.dumps(msg_obj) + guardrail_turn_detection_injected = True + verbose_logger.debug( + "Injected turn_detection into first session.update for audio transcription guardrails" + ) + + ## GUARDRAIL: Force ``create_response`` to False in any + # client-provided ``turn_detection`` so a later + # ``session.update`` cannot re-enable VAD auto-response + # and bypass the transcription guardrail after the + # initial disable. Covers both the flat beta key and the + # nested GA ``audio.input.turn_detection`` shape, since + # the GA remap below also accepts either form. Skipped + # when the injection block above already ran for this + # message, to avoid redundant double-serialization. + if ( + msg_type == "session.update" + and not guardrail_turn_detection_injected + and self._has_audio_transcription_guardrails() + ): + session = msg_obj.get("session") + if isinstance(session, dict): + td_overridden = False + flat_td = session.get("turn_detection") + flat_td_present = flat_td is not None + if flat_td_present: + if not isinstance(flat_td, dict): + flat_td = {} + if flat_td.get("create_response") is not False: + flat_td["create_response"] = False + session["turn_detection"] = flat_td + td_overridden = True + nested_td_present = False + 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 nested_td is not None: + nested_td_present = True + if not isinstance(nested_td, dict): + nested_td = {} + if ( + nested_td.get("create_response") + is not False + ): + nested_td["create_response"] = False + audio_input["turn_detection"] = nested_td + td_overridden = True + # Symmetric with the first-update injection block: + # if the client omitted turn_detection entirely on + # a subsequent session.update, still inject the + # ``create_response: False`` override so the + # transcription guardrail cannot be re-enabled by + # any downstream merge that drops the original + # disable. + if not flat_td_present and not nested_td_present: + session["turn_detection"] = {"create_response": False} + td_overridden = True + if td_overridden: + message = json.dumps(msg_obj) + # GA compatibility: remap beta-style session fields only when # the upstream is in GA mode. Beta upstreams expect the flat # session shape unchanged. @@ -848,17 +1076,20 @@ class RealTimeStreaming: pass ## LOGGING + # Log after any in-place modifications (GA remap, guardrail + # turn_detection injection) so audit logs reflect what we + # actually forward to the backend. self.store_input(message=message) - ## FORWARD TO BACKEND - if self.provider_config: - message = self.provider_config.transform_realtime_request( - message, self.model - ) - for msg in message: - await self.backend_ws.send(msg) # type: ignore[union-attr] - else: - await self.backend_ws.send(message) # type: ignore[union-attr] + ## FORWARD TO BACKEND + # Only mark the guardrail turn_detection update as sent after the + # backend actually accepted the message. Setting the flag earlier + # would permanently disable the injection if ``_send_to_backend`` + # raised — neither this loop nor + # ``_maybe_send_guardrail_turn_detection_update`` would retry. + sent = await self._send_to_backend(message) + if guardrail_turn_detection_injected and sent: + self._guardrail_turn_detection_update_sent = True except Exception as e: verbose_logger.debug(f"Error in client ack messages: {e}") diff --git a/litellm/llms/base_llm/realtime/transformation.py b/litellm/llms/base_llm/realtime/transformation.py index d5531a532b9..0f239b4ad45 100644 --- a/litellm/llms/base_llm/realtime/transformation.py +++ b/litellm/llms/base_llm/realtime/transformation.py @@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, List, Optional, Union import httpx +from litellm.types.llms.openai import OpenAIRealtimeStreamSessionEvents from litellm.types.realtime import ( RealtimeResponseTransformInput, RealtimeResponseTypedDict, @@ -69,6 +70,20 @@ class BaseRealtimeConfig(ABC): ) -> Optional[str]: # message sent to setup the realtime session return None + def transform_session_created_event( + self, + model: str, + logging_session_id: str, + session_configuration_request: Optional[str] = None, + ) -> Optional[Union[dict, OpenAIRealtimeStreamSessionEvents]]: + """ + Optional hook for providers that defer session setup until client `session.update`. + + Return an OpenAI-compatible `session.created` payload when the proxy should + emit a synthetic event immediately after backend websocket connection. + """ + return None + @abstractmethod def transform_realtime_response( self, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index c9ab3c648ac..36f6b3903e1 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -5316,6 +5316,28 @@ class BaseLLMHTTPHandler: ) if _session_config: realtime_streaming.session_configuration_request = _session_config + + # For providers that defer setup until client session.update, optionally + # send synthetic session.created to unblock clients waiting on connect. + if not provider_config.requires_session_configuration(): + synthetic_session = provider_config.transform_session_created_event( + model=model, + logging_session_id=logging_obj.litellm_trace_id, + session_configuration_request=None, + ) + if synthetic_session is not None: + synthetic_session_str = json.dumps(synthetic_session) + # Record before sending so the synthetic session.created is + # captured in the session log alongside provider-driven + # events; without this it would be silently absent from + # success_handler / async_success_handler payloads. + realtime_streaming.store_message(synthetic_session_str) + await websocket.send_text(synthetic_session_str) + realtime_streaming._session_created_sent_to_client = True + verbose_logger.debug( + "Sent synthetic session.created to client to unblock connection" + ) + await realtime_streaming.bidirectional_forward() except websockets.exceptions.InvalidStatusCode as e: # type: ignore diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index 4378db06358..cf1fc75ef10 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -3,8 +3,10 @@ This file contains the transformation logic for the Gemini realtime API. """ import json +from collections import OrderedDict from typing import Any, Dict, List, Optional, Union, cast +import litellm from litellm import verbose_logger from litellm._uuid import uuid from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -29,6 +31,7 @@ from litellm.types.llms.openai import ( OpenAIRealtimeDoneEvent, OpenAIRealtimeEvents, OpenAIRealtimeEventTypes, + OpenAIRealtimeFunctionCallArgumentsDone, OpenAIRealtimeOutputItemDone, OpenAIRealtimeResponseAudioDone, OpenAIRealtimeResponseContentPartAdded, @@ -36,10 +39,12 @@ from litellm.types.llms.openai import ( OpenAIRealtimeResponseDoneObject, OpenAIRealtimeResponseTextDone, OpenAIRealtimeStreamResponseBaseObject, + OpenAIRealtimeStreamResponseOutputItem, OpenAIRealtimeStreamResponseOutputItemAdded, OpenAIRealtimeStreamSession, OpenAIRealtimeStreamSessionEvents, OpenAIRealtimeTurnDetection, + ResponsesAPIStreamEvents, ) from litellm.types.llms.vertex_ai import ( GeminiResponseModalities, @@ -56,15 +61,43 @@ from litellm.utils import get_empty_usage from ..common_utils import encode_unserializable_types, get_api_key_from_env -MAP_GEMINI_FIELD_TO_OPENAI_EVENT: Dict[str, OpenAIRealtimeEventTypes] = { +MAP_GEMINI_FIELD_TO_OPENAI_EVENT: Dict[ + str, Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents] +] = { "setupComplete": OpenAIRealtimeEventTypes.SESSION_CREATED, "serverContent.generationComplete": OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE, "serverContent.turnComplete": OpenAIRealtimeEventTypes.RESPONSE_DONE, "serverContent.interrupted": OpenAIRealtimeEventTypes.RESPONSE_DONE, + "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. +_KNOWN_GEMINI_TOP_LEVEL_KEYS: set = { + map_key.split(".", 1)[0] for map_key in MAP_GEMINI_FIELD_TO_OPENAI_EVENT } 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 + + 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. + self._pending_usage_metadata: Optional[dict] = None + def validate_environment( self, headers: dict, model: str, api_key: Optional[str] = None ) -> dict: @@ -190,10 +223,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) vertex_gemini_config = VertexGeminiConfig() - optional_params["generationConfig"]["tools"] = ( - vertex_gemini_config._map_function( - value=value, optional_params=optional_params - ) + # Tools should be at the top level of setup, not inside generationConfig + optional_params["tools"] = vertex_gemini_config._map_function( + value=value, optional_params=optional_params ) elif key == "input_audio_transcription" and value is not None: optional_params["inputAudioTranscription"] = {} @@ -214,6 +246,272 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): optional_params.pop("generationConfig") return optional_params + @staticmethod + def _extract_turn_detection(session: dict) -> Optional[dict]: + """Extract turn_detection from a session.update payload. + + Handles both the flat beta shape (``session.turn_detection``) and the + GA shape (``session.audio.input.turn_detection``). + """ + if not isinstance(session, dict): + return None + td = session.get("turn_detection") + if isinstance(td, dict): + return td + audio = session.get("audio") + if isinstance(audio, dict): + input_cfg = audio.get("input") + if isinstance(input_cfg, dict): + td = input_cfg.get("turn_detection") + if isinstance(td, dict): + return td + return None + + @staticmethod + def _normalize_session_payload_for_mapping(session: dict) -> dict: + """Normalize GA-remapped session fields back to their beta keys. + + ``map_openai_params`` only recognises the flat OpenAI-beta key names + (``modalities``, ``input_audio_transcription``, ``turn_detection``). + For GA clients the upstream shim renames these into the nested GA + schema (``output_modalities``, ``audio.input.transcription``, + ``audio.input.turn_detection``), which would otherwise be silently + dropped here. Surface them back at the top level so the existing + mapping logic picks them up without duplicating provider-specific + knowledge of the GA schema in ``map_openai_params``. + """ + if not isinstance(session, dict): + return session + + normalized = dict(session) + + if "modalities" not in normalized and "output_modalities" in normalized: + normalized["modalities"] = normalized["output_modalities"] + + audio = normalized.get("audio") + if isinstance(audio, dict): + input_cfg = audio.get("input") + if isinstance(input_cfg, dict): + if ( + "input_audio_transcription" not in normalized + and "transcription" in input_cfg + ): + normalized["input_audio_transcription"] = input_cfg["transcription"] + + extracted_turn_detection = GeminiRealtimeConfig._extract_turn_detection( + normalized + ) + if extracted_turn_detection is not None and not isinstance( + normalized.get("turn_detection"), dict + ): + normalized["turn_detection"] = extracted_turn_detection + + return normalized + + def _handle_session_update( + self, + json_message: dict, + model: str, + session_configuration_request: Optional[str], + ) -> List[str]: + """ + Handle session.update by sending setup to Gemini. + + On the FIRST session.update (when session_configuration_request is None), + the full setup with all configuration is sent. + + Subsequent session.update messages are forwarded as a follow-up setup + with the new fields merged into the original setup. Gemini Live treats + a follow-up BidiGenerateContentSetup as a full session replacement + rather than a partial merge, so we carry forward the previous setup + (tools, generationConfig, inputAudioTranscription, systemInstruction, + ...) and overlay the new fields on top. This preserves the old + behavior where clients could refine the session via session.update + (e.g. add tools after the auto-setup on connect), and also keeps the + guardrail-driven turn_detection update working. + """ + session_payload = json_message.get("session") or {} + # Normalize GA-remapped fields (``output_modalities``, + # nested ``audio.input.transcription``, + # ``audio.input.turn_detection``) back to their flat beta keys so + # ``map_openai_params`` picks them up. Without this, GA clients' + # explicit modality / transcription / turn-detection settings + # would be silently dropped because ``map_openai_params`` only + # recognises the flat OpenAI-beta key names. + session_payload = self._normalize_session_payload_for_mapping(session_payload) + new_overrides = self.map_openai_params( + optional_params={}, non_default_params=session_payload + ) + + if session_configuration_request is None: + generation_config = new_overrides.setdefault("generationConfig", {}) + generation_config.setdefault("responseModalities", ["AUDIO"]) + new_overrides.setdefault("inputAudioTranscription", {}) + new_overrides["model"] = f"models/{model}" + verbose_logger.debug( + "Gemini Realtime: Sending initial setup with tools to backend" + ) + return [json.dumps({"setup": new_overrides})] + + if not new_overrides: + verbose_logger.debug( + "Gemini Realtime: Ignoring session.update (no mappable fields)" + ) + return [] + + try: + original_setup = cast( + BidiGenerateContentSetup, + json.loads(session_configuration_request).get("setup", {}), + ) + except (json.JSONDecodeError, AttributeError): + original_setup = {} + + # Deep-merge ``generationConfig`` and ``realtimeInputConfig`` so a + # partial session.update (e.g. only ``temperature`` or only + # ``modalities``) does not silently drop unrelated sub-keys + # (``responseModalities``, ``maxOutputTokens``, ...) from the original + # setup. + follow_up_setup: BidiGenerateContentSetup = { + **original_setup, + **new_overrides, + "model": f"models/{model}", + } + original_generation_config = original_setup.get("generationConfig") + new_generation_config = new_overrides.get("generationConfig") + if isinstance(original_generation_config, dict) and isinstance( + new_generation_config, dict + ): + follow_up_setup["generationConfig"] = { + **original_generation_config, + **new_generation_config, + } + original_realtime_input_config = original_setup.get("realtimeInputConfig") + new_realtime_input_config = new_overrides.get("realtimeInputConfig") + if isinstance(original_realtime_input_config, dict) and isinstance( + new_realtime_input_config, dict + ): + merged_realtime_input_config = { + **original_realtime_input_config, + **new_realtime_input_config, + } + # Deep-merge ``automaticActivityDetection`` so a partial VAD + # update (e.g. the guardrail-injected ``disabled: True`` from + # ``create_response: False``) does not silently drop unrelated + # knobs like ``silenceDurationMs`` / ``prefixPaddingMs`` from + # the original setup. + original_automatic_activity_detection = original_realtime_input_config.get( + "automaticActivityDetection" + ) + new_automatic_activity_detection = new_realtime_input_config.get( + "automaticActivityDetection" + ) + if isinstance(original_automatic_activity_detection, dict) and isinstance( + new_automatic_activity_detection, dict + ): + merged_realtime_input_config["automaticActivityDetection"] = { + **original_automatic_activity_detection, + **new_automatic_activity_detection, + } + follow_up_setup["realtimeInputConfig"] = cast( + BidiGenerateContentRealtimeInputConfig, + merged_realtime_input_config, + ) + verbose_logger.debug( + "Gemini Realtime: Forwarding session.update as follow-up setup" + ) + return [json.dumps({"setup": follow_up_setup})] + + def _handle_conversation_item(self, json_message: dict) -> List[str]: + """ + Handle conversation.item.create for user text or function call output. + + Converts OpenAI format to Gemini's clientContent (for user text) or + toolResponse (for function outputs). + """ + 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]: + """Transform function_call_output to Gemini toolResponse format.""" + call_id = item.get("call_id", "") + output = item.get("output", "{}") + + verbose_logger.debug( + 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. + try: + parsed_output = json.loads(output) if isinstance(output, str) else output + except json.JSONDecodeError: + parsed_output = output + output_dict = ( + parsed_output + if isinstance(parsed_output, dict) + 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. + 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) + else: + verbose_logger.warning( + f"Gemini Realtime: Function name not found for call_id={call_id}. " + "This may cause Gemini to reject the response." + ) + + # Build Gemini toolResponse format + function_response = { + "id": call_id, + "response": output_dict, + } + if function_name: + function_response["name"] = function_name + + tool_response_message = { + "toolResponse": {"functionResponses": [function_response]} + } + + return [json.dumps(tool_response_message)] + + def _handle_user_text_content(self, item: dict) -> List[str]: + """Transform user text content to Gemini clientContent format.""" + content_list = item.get("content", []) + text_parts = [ + c.get("text", "") + for c in content_list + if isinstance(c, dict) and c.get("type") == "input_text" + ] + text = " ".join(filter(None, text_parts)) + if not text: + return [] + + # Build clientContent message with turns (proper Gemini Live API format) + client_content_message = { + "clientContent": { + "turns": [{"role": "user", "parts": [{"text": text}]}], + "turnComplete": True, + } + } + + return [json.dumps(client_content_message)] + def transform_realtime_request( self, message: str, @@ -233,55 +531,42 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): messages: List[str] = [] msg_type = json_message.get("type") - ## HANDLE SESSION UPDATE — translate to Gemini setup; no realtime_input needed ## + ## HANDLE SESSION UPDATE — translate to Gemini setup ## if msg_type == "session.update": - client_session_configuration_request = self.map_openai_params( - optional_params={}, non_default_params=json_message["session"] + return self._handle_session_update( + json_message, model, session_configuration_request ) - client_session_configuration_request["model"] = f"models/{model}" - messages.append(json.dumps({"setup": client_session_configuration_request})) - return messages ## HANDLE response.create — Gemini responds automatically; nothing to forward ## if msg_type == "response.create": return [] - ## HANDLE INPUT AUDIO BUFFER ## + ## 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"] ) - ## HANDLE conversation.item.create — extract actual user text ## - elif msg_type == "conversation.item.create": - item = json_message.get("item", {}) - content_list = item.get("content", []) - text_parts = [ - c.get("text", "") - for c in content_list - if isinstance(c, dict) and c.get("type") == "input_text" - ] - text = " ".join(filter(None, text_parts)) - if not text: - return [] - realtime_input_dict["text"] = text - else: - # Unknown/unsupported OpenAI event type — drop silently rather than - # forwarding raw JSON as text input to the model. - return [] - if len(realtime_input_dict) != 1: - raise ValueError( - f"Only one argument can be set, got {len(realtime_input_dict)}:" - f" {list(realtime_input_dict.keys())}" + realtime_input_dict = cast( + BidiGenerateContentRealtimeInput, + encode_unserializable_types( + cast(Dict[str, object], realtime_input_dict) + ), ) - realtime_input_dict = cast( - BidiGenerateContentRealtimeInput, - encode_unserializable_types(cast(Dict[str, object], realtime_input_dict)), - ) - - messages.append(json.dumps({"realtime_input": realtime_input_dict})) - return messages + gemini_msg = json.dumps({"realtimeInput": realtime_input_dict}) + verbose_logger.debug( + "Gemini Realtime: Sending audio realtimeInput to backend" + ) + messages.append(gemini_msg) + return messages + # Unknown/unsupported OpenAI event type — drop silently rather than + # forwarding raw JSON as text input to the model. + return [] def transform_session_created_event( self, @@ -300,7 +585,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): generation_config = ( session_configuration_request_dict.get("generationConfig", {}) or {} ) - gemini_modalities = generation_config.get("responseModalities", ["TEXT"]) + gemini_modalities = generation_config.get("responseModalities", ["AUDIO"]) _modalities = [ modality.lower() for modality in cast(List[str], gemini_modalities) ] @@ -352,18 +637,18 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): delta_type: ALL_DELTA_TYPES, session_configuration_request: Optional[str] = None, ) -> List[OpenAIRealtimeEvents]: - if session_configuration_request is None: - raise ValueError( - "session_configuration_request is required for Gemini API calls" - ) - - session_configuration_request_dict: BidiGenerateContentSetup = json.loads( - session_configuration_request - ).get("setup", {}) + session_configuration_request_dict: BidiGenerateContentSetup = {} + if session_configuration_request is not None: + try: + session_configuration_request_dict = json.loads( + session_configuration_request + ).get("setup", {}) + except json.JSONDecodeError: + session_configuration_request_dict = {} generation_config = session_configuration_request_dict.get( "generationConfig", {} ) - gemini_modalities = generation_config.get("responseModalities", ["TEXT"]) + gemini_modalities = generation_config.get("responseModalities", ["AUDIO"]) _modalities = [ modality.lower() for modality in cast(List[str], gemini_modalities) ] @@ -576,6 +861,86 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): returned_items.append(response_output_item_done) return returned_items + def _consume_usage_metadata_for_response_done(self, frame: dict) -> Optional[dict]: + """Return the ``usageMetadata`` to attribute to a ``response.done``. + + 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()``. + """ + # ``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 + return in_frame + buffered = self._pending_usage_metadata + self._pending_usage_metadata = None + return buffered + + def transform_tool_call_events( + self, + tool_call_message: dict, + 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()}" + + verbose_logger.debug( + f"Gemini Realtime: Transforming {len(function_calls)} tool call(s) to OpenAI format" + ) + + events: List[OpenAIRealtimeFunctionCallArgumentsDone] = [] + for idx, fc in enumerate(function_calls): + call_id = fc.get("id", "") + 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) + while len(self._tool_call_id_to_name) > self._TOOL_CALL_ID_TO_NAME_MAX: + self._tool_call_id_to_name.popitem(last=False) + + events.append( + OpenAIRealtimeFunctionCallArgumentsDone( + type="response.function_call_arguments.done", + event_id=f"event_{uuid.uuid4()}", + response_id=resolved_response_id, + item_id=f"{resolved_output_item_id}_tool_{idx}", + output_index=idx, + call_id=call_id, + name=name, + arguments=json.dumps(fc.get("args", {})), + ) + ) + + return events + @staticmethod def get_nested_value(obj: dict, path: str) -> Any: keys = path.split(".") @@ -681,14 +1046,20 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): "generationConfig", {} ) temperature = generation_config.get("temperature") - max_output_tokens = generation_config.get("max_output_tokens") - gemini_modalities = generation_config.get("responseModalities", ["TEXT"]) + max_output_tokens = generation_config.get("maxOutputTokens") + gemini_modalities = generation_config.get("responseModalities", ["AUDIO"]) _modalities = [ modality.lower() for modality in cast(List[str], gemini_modalities) ] - if "usageMetadata" in message: + resolved_usage_metadata = self._consume_usage_metadata_for_response_done( + cast(dict, message) + ) + if resolved_usage_metadata is not None: _chat_completion_usage = VertexGeminiConfig._calculate_usage( - completion_response=message, + completion_response=cast( + BidiGenerateContentServerMessage, + {**cast(dict, message), "usageMetadata": resolved_usage_metadata}, + ), ) else: _chat_completion_usage = get_empty_usage() @@ -716,7 +1087,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): if temperature is not None: response_done_event["response"]["temperature"] = temperature if max_output_tokens is not None: - response_done_event["response"]["max_output_tokens"] = max_output_tokens + response_done_event["response"]["max_output_tokens"] = cast( + int, max_output_tokens + ) return response_done_event @@ -808,13 +1181,18 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def map_openai_event( self, key: str, - value: dict, + value: Any, current_delta_type: Optional[ALL_DELTA_TYPES], - json_message: dict, - ) -> OpenAIRealtimeEventTypes: - model_turn_event = value.get("modelTurn") - generation_complete_event = value.get("generationComplete") - openai_event: Optional[OpenAIRealtimeEventTypes] = None + ) -> Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents]: + if isinstance(value, dict): + model_turn_event = value.get("modelTurn") + generation_complete_event = value.get("generationComplete") + else: + model_turn_event = None + generation_complete_event = None + openai_event: Optional[ + Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents] + ] = None if model_turn_event: # check if model turn event openai_event = self.map_model_turn_event(model_turn_event) elif generation_complete_event: @@ -822,15 +1200,27 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): delta_type=current_delta_type ) else: - # Check if this key or any nested key matches our mapping - for map_key, openai_event in MAP_GEMINI_FIELD_TO_OPENAI_EVENT.items(): - if map_key == key or ( - "." in map_key - and GeminiRealtimeConfig.get_nested_value(json_message, map_key) - is not None - ): - openai_event = openai_event + # Check if this key or any nested key matches our mapping. Use a + # distinct loop variable so we don't shadow ``openai_event`` and + # leak the last dict value when no entry matches. Scope dotted-key + # lookups to the current ``key``/``value`` pair — checking the + # whole ``json_message`` would let a sibling key (e.g. + # ``serverContent.turnComplete``) misclassify the event currently + # being processed (e.g. ``toolCall``). + for map_key, candidate_event in MAP_GEMINI_FIELD_TO_OPENAI_EVENT.items(): + if map_key == key: + openai_event = candidate_event break + if "." in map_key: + prefix, _, nested_path = map_key.partition(".") + if ( + prefix == key + and isinstance(value, dict) + and GeminiRealtimeConfig.get_nested_value(value, nested_path) + is not None + ): + openai_event = candidate_event + break if openai_event is None: raise ValueError(f"Unknown openai event: {key}, value: {value}") return openai_event @@ -854,6 +1244,15 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): message_str = str(message) raise ValueError(f"Invalid JSON message: {message_str}") + verbose_logger.debug( + "Realtime Response Transform: Gemini frame keys=%s", + ( + sorted(json_message.keys()) + if isinstance(json_message, dict) + else type(json_message).__name__ + ), + ) + logging_session_id = logging_obj.litellm_trace_id current_output_item_id = realtime_response_transform_input[ @@ -913,32 +1312,44 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): ) # If serverContent only contained transcription(s) and no model - # content, return early — the main loop would fail on unknown keys. + # 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. _model_content_keys = { "modelTurn", "turnComplete", "interrupted", "generationComplete", } - if not any(k in server_content for k in _model_content_keys): - return { - "response": returned_message, - "current_output_item_id": current_output_item_id, - "current_response_id": current_response_id, - "current_delta_chunks": current_delta_chunks, - "current_conversation_id": current_conversation_id, - "current_item_chunks": current_item_chunks, - "current_delta_type": current_delta_type, - "session_configuration_request": session_configuration_request, - } + server_content_handled = not any( + k in server_content for k in _model_content_keys + ) + else: + server_content_handled = False - for key, value in json_message.items(): + 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()): + # 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 + # would otherwise terminate the WebSocket session. + if key not in _KNOWN_GEMINI_TOP_LEVEL_KEYS: + continue + # serverContent was a transcription-only payload already emitted + # above; skip it here so map_openai_event doesn't raise on the + # missing model-content subkeys. + if key == "serverContent" and server_content_handled: + continue # Check if this key or any nested key matches our mapping openai_event = self.map_openai_event( key=key, value=value, current_delta_type=current_delta_type, - json_message=json_message, ) if openai_event == OpenAIRealtimeEventTypes.SESSION_CREATED: @@ -947,8 +1358,226 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): logging_session_id, realtime_response_transform_input["session_configuration_request"], ) - session_configuration_request = json.dumps(transformed_message) 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"): + 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: + session_setup = json.loads(session_configuration_request).get( + "setup", {} + ) + except (json.JSONDecodeError, TypeError): + session_setup = {} + tool_call_generation_config = ( + session_setup.get("generationConfig", {}) or {} + ) + tool_call_modalities = [ + modality.lower() + for modality in cast( + List[str], + tool_call_generation_config.get( + "responseModalities", ["AUDIO"] + ), + ) + ] + + # 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", + "event_id": f"event_{uuid.uuid4()}", + "response": { + "object": "realtime.response", + "id": current_response_id, + "status": "in_progress", + "output": [], + "conversation_id": current_conversation_id, + "modalities": tool_call_modalities, + "temperature": tool_call_generation_config.get( + "temperature" + ), + "max_output_tokens": tool_call_generation_config.get( + "maxOutputTokens" + ), + }, + } + ) + + tool_call_events = self.transform_tool_call_events( + value, + 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 = { + "id": item_id, + "object": "realtime.item", + "type": "function_call", + "status": "completed", + "call_id": tool_call["call_id"], + "name": tool_call["name"], + "arguments": tool_call["arguments"], + } + # response.output_item.added + returned_message.append( + OpenAIRealtimeStreamResponseOutputItemAdded( + type="response.output_item.added", + event_id=f"event_{uuid.uuid4()}", + response_id=current_response_id, + output_index=idx, + item={ + **function_call_item, + "status": "in_progress", + "arguments": "", + }, + ) + ) + # 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``. + returned_message.append( + cast( + OpenAIRealtimeEvents, + { + "type": "response.function_call_arguments.delta", + "event_id": f"event_{uuid.uuid4()}", + "response_id": current_response_id, + "item_id": item_id, + "output_index": idx, + "call_id": tool_call["call_id"], + "delta": tool_call["arguments"], + }, + ) + ) + # 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. + returned_message.append( + OpenAIRealtimeOutputItemDone( + type="response.output_item.done", + event_id=f"event_{uuid.uuid4()}", + response_id=current_response_id, + output_index=idx, + item={**function_call_item}, + ) + ) + # conversation.item.created + returned_message.append( + OpenAIRealtimeConversationItemCreated( + type="conversation.item.created", + event_id=f"event_{uuid.uuid4()}", + item={**function_call_item}, + ) + ) + + # 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) + ) + if resolved_tool_call_usage_metadata is not None: + _tool_call_chat_completion_usage = ( + VertexGeminiConfig._calculate_usage( + completion_response=cast( + BidiGenerateContentServerMessage, + { + **json_message, + "usageMetadata": resolved_tool_call_usage_metadata, + }, + ), + ) + ) + else: + _tool_call_chat_completion_usage = get_empty_usage() + tool_call_responses_api_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( + _tool_call_chat_completion_usage, + ) + tool_call_done_event = OpenAIRealtimeDoneEvent( + type="response.done", + event_id=f"event_{uuid.uuid4()}", + response=OpenAIRealtimeResponseDoneObject( + id=current_response_id, + object="realtime.response", + status="completed", + output=[ + { + "id": te["item_id"], + "object": "realtime.item", + "type": "function_call", + "status": "completed", + "call_id": te["call_id"], + "name": te["name"], + "arguments": te["arguments"], + } + for te in tool_call_events + ], + conversation_id=current_conversation_id, + modalities=tool_call_modalities, + usage=tool_call_responses_api_usage.model_dump(), + ), + ) + tool_call_temperature = tool_call_generation_config.get("temperature") + if tool_call_temperature is not None: + tool_call_done_event["response"][ + "temperature" + ] = tool_call_temperature + tool_call_max_output_tokens = tool_call_generation_config.get( + "maxOutputTokens" + ) + if tool_call_max_output_tokens is not None: + tool_call_done_event["response"]["max_output_tokens"] = cast( + 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: transformed_response_done_event = self.transform_response_done_event( message=BidiGenerateContentServerMessage(**json_message), # type: ignore @@ -958,16 +1587,37 @@ 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 ( openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA or openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE 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. + _modality_input: RealtimeResponseTransformInput = { + **realtime_response_transform_input, + "current_output_item_id": current_output_item_id, + "current_response_id": current_response_id, + "current_conversation_id": current_conversation_id, + "current_delta_chunks": current_delta_chunks, + "current_item_chunks": current_item_chunks, + "current_delta_type": current_delta_type, + "session_configuration_request": session_configuration_request, + } _returned_message = self.handle_openai_modality_event( openai_event, json_message, - realtime_response_transform_input, + _modality_input, delta_type="text" if "text" in openai_event.value else "audio", ) returned_message.extend(_returned_message["returned_message"]) @@ -979,6 +1629,41 @@ 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 + if key in _KNOWN_GEMINI_TOP_LEVEL_KEYS + 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 + if not unhandled_known_keys: + return { + "response": returned_message, + "current_output_item_id": current_output_item_id, + "current_response_id": current_response_id, + "current_delta_chunks": current_delta_chunks, + "current_conversation_id": current_conversation_id, + "current_item_chunks": current_item_chunks, + "current_delta_type": current_delta_type, + "session_configuration_request": session_configuration_request, + } if isinstance(message, bytes): message_str = message.decode("utf-8", errors="replace") else: @@ -993,6 +1678,13 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): transformed_message=returned_message, current_item_chunks=current_item_chunks, ) + + for msg in returned_message: + event_type = msg.get("type") if isinstance(msg, dict) else "unknown" + verbose_logger.debug( + "Realtime Response Transform: OpenAI event=%s", event_type + ) + return { "response": returned_message, "current_output_item_id": current_output_item_id, @@ -1005,7 +1697,10 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): } def requires_session_configuration(self) -> bool: - return True + # 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 + return not litellm.gemini_live_defer_setup def session_configuration_request(self, model: str) -> str: """ diff --git a/litellm/llms/vertex_ai/realtime/transformation.py b/litellm/llms/vertex_ai/realtime/transformation.py index 2b4746b174e..ea4dbccc8c8 100644 --- a/litellm/llms/vertex_ai/realtime/transformation.py +++ b/litellm/llms/vertex_ai/realtime/transformation.py @@ -14,6 +14,7 @@ Auth: OAuth2 Bearer token (not an API key). import json from typing import List, Optional +from litellm import verbose_logger from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig @@ -26,6 +27,7 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): """ def __init__(self, access_token: str, project: str, location: str) -> None: + super().__init__() self._access_token = access_token self._project = project self._location = location @@ -138,6 +140,62 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): # Request translation # ------------------------------------------------------------------ + def _vertex_model_path(self, model: str) -> str: + """Return the fully-qualified Vertex AI model resource path.""" + return ( + f"projects/{self._project}" + f"/locations/{self._location}" + f"/publishers/google/models/{model}" + ) + + def _build_vertex_ai_setup_config(self, model: str, session_params: dict) -> dict: + """Build Vertex AI setup configuration with proper model path and defaults.""" + # Normalize GA-remapped fields (``output_modalities``, nested + # ``audio.input.transcription``, ``audio.input.turn_detection``) back to + # their flat beta keys so ``map_openai_params`` picks them up. Without + # this, GA clients' explicit modality / transcription / turn-detection + # settings would be silently dropped because ``map_openai_params`` only + # recognises the flat OpenAI-beta key names. + session_params = self._normalize_session_payload_for_mapping(session_params) + setup_config = self.map_openai_params( + optional_params={}, non_default_params=session_params + ) + + # Use full Vertex AI model path + setup_config["model"] = self._vertex_model_path(model) + + # Add Vertex AI specific defaults if not provided + generation_config = setup_config.setdefault("generationConfig", {}) + generation_config.setdefault("responseModalities", ["AUDIO"]) + + # Ensure Vertex defaults for realtimeInputConfig apply even when + # the client provided a partial ``turn_detection`` (e.g. only + # ``silence_duration_ms``). ``map_automatic_turn_detection`` sets + # ``disabled=True`` whenever ``create_response`` is absent or + # ``False``. Force ``disabled=False`` only when the client did + # not explicitly request ``create_response: False`` — that path + # is how transcription guardrails suppress automatic responses, + # and overriding it here would silently bypass the guardrail. + # Vertex Live has no "VAD on, no auto-response" mode, so callers + # that need that behaviour must accept that VAD is off. + client_turn_detection = session_params.get("turn_detection") + client_disabled_auto_response = ( + isinstance(client_turn_detection, dict) + and client_turn_detection.get("create_response") is False + ) + realtime_input_config = setup_config.setdefault("realtimeInputConfig", {}) + automatic_detection = realtime_input_config.setdefault( + "automaticActivityDetection", {} + ) + if not client_disabled_auto_response: + automatic_detection["disabled"] = False + automatic_detection.setdefault("silenceDurationMs", 800) + + setup_config.setdefault("inputAudioTranscription", {}) + setup_config.setdefault("outputAudioTranscription", {}) + + return setup_config + def transform_realtime_request( self, message: str, @@ -147,16 +205,50 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): """ Translate OpenAI realtime client messages to Vertex AI format. - ``session.update`` is intentionally ignored (returns []) because - Vertex AI only accepts a single ``setup`` message at the start of - the connection — sending a second one causes a 1007 close error. - The initial setup (sent automatically before bidirectional_forward) - already includes AUDIO modality and server VAD, so there is nothing - more to configure. + On the first ``session.update`` (when no setup has been sent yet) the + full ``BidiGenerateContentSetup`` is built with Vertex AI's model path + and forwarded. Any later ``session.update`` is dropped: Vertex AI + documents ``setup`` as the first-and-only client message, and a second + ``setup`` closes the connection with a 1007 policy error. """ json_message = json.loads(message) - if json_message.get("type") == "session.update": - # Do not forward as a second setup — Vertex AI rejects it. + msg_type = json_message.get("type") + + if msg_type == "session.update": + if session_configuration_request is None: + setup_config = self._build_vertex_ai_setup_config( + model, json_message.get("session") or {} + ) + gemini_setup_msg = json.dumps({"setup": setup_config}) + + verbose_logger.debug( + "Vertex AI Realtime: Sending initial setup with tools to backend" + ) + return [gemini_setup_msg] + + # A follow-up session.update can't be forwarded as a second setup + # (Vertex Live closes the WebSocket with 1007). If this drop is + # silencing the audio-transcription guardrail's create_response + # disable, surface a warning so operators know the model will + # auto-respond before the guardrail can gate it on Vertex AI. + client_turn_detection = GeminiRealtimeConfig._extract_turn_detection( + json_message.get("session") or {} + ) + if ( + isinstance(client_turn_detection, dict) + and client_turn_detection.get("create_response") is False + ): + verbose_logger.warning( + "Vertex AI Realtime: Dropping subsequent session.update " + "(turn_detection.create_response=False) — Vertex Live " + "rejects a second setup message. Audio-transcription " + "guardrails cannot suppress the model's auto-response on " + "Vertex AI in non-deferred mode." + ) + else: + verbose_logger.debug( + "Vertex AI Realtime: Ignoring session.update (setup already sent)" + ) return [] return super().transform_realtime_request( diff --git a/litellm/types/llms/gemini.py b/litellm/types/llms/gemini.py index 9e3fea1bbbb..8763544facc 100644 --- a/litellm/types/llms/gemini.py +++ b/litellm/types/llms/gemini.py @@ -133,7 +133,7 @@ class BidiGenerateContentSetup(TypedDict, total=False): tools: List[Tools] """The tools to be used for the realtime session.""" - realtimeInputConfig: dict + realtimeInputConfig: BidiGenerateContentRealtimeInputConfig """The realtime config to be used for the realtime session.""" sessionResumption: dict diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index abe58199dfd..14114b22f39 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -79,7 +79,14 @@ from pydantic import ( field_serializer, field_validator, ) -from typing_extensions import Annotated, Dict, Required, TypedDict, override +from typing_extensions import ( + Annotated, + Dict, + NotRequired, + Required, + TypedDict, + override, +) from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject from litellm.types.responses.main import ( @@ -1935,6 +1942,7 @@ class OpenAIRealtimeStreamResponseOutputItemAdded(TypedDict): response_id: str output_index: int item: OpenAIRealtimeStreamResponseOutputItem + event_id: NotRequired[str] class OpenAIRealtimeStreamResponseBaseObject(TypedDict): @@ -2061,6 +2069,17 @@ class OpenAIRealtimeContentPartDone(TypedDict): type: Literal["response.content_part.done"] +class OpenAIRealtimeFunctionCallArgumentsDone(TypedDict): + type: Literal["response.function_call_arguments.done"] + event_id: str + response_id: str + item_id: str + output_index: int + call_id: str + name: str + arguments: str + + class OpenAIRealtimeOutputItemDone(TypedDict): event_id: str item: OpenAIRealtimeStreamResponseOutputItem @@ -2126,6 +2145,7 @@ OpenAIRealtimeEvents = Union[ OpenAIRealtimeResponseAudioDone, OpenAIRealtimeContentPartDone, OpenAIRealtimeOutputItemDone, + OpenAIRealtimeFunctionCallArgumentsDone, OpenAIRealtimeDoneEvent, ] 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 acf8e4d2a14..7913efe8294 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -476,6 +476,50 @@ async def test_transcription_captured_in_backend_to_client(): assert logging_obj.model_call_details["messages"] == streaming.input_messages +@pytest.mark.asyncio +async def test_client_ack_caches_setup_to_prevent_duplicate_session_update_setup(): + websocket = MagicMock() + backend_ws = MagicMock() + logging_obj = MagicMock() + logging_obj.pre_call = MagicMock() + + # Two session.update messages arrive before setupComplete round-trip. + websocket.receive_text = AsyncMock( + side_effect=[ + json.dumps({"type": "session.update", "session": {"tools": []}}), + json.dumps({"type": "session.update", "session": {"tools": []}}), + Exception("client done"), + ] + ) + + provider_config = MagicMock() + + def _transform(message: str, model: str, session_configuration_request=None): + if session_configuration_request is None: + return [json.dumps({"setup": {"model": "models/gemini-2.5-flash"}})] + return [] + + provider_config.transform_realtime_request = MagicMock(side_effect=_transform) + + backend_ws.send = AsyncMock() + + streaming = RealTimeStreaming( + websocket=websocket, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + + await streaming.client_ack_messages() + + # Setup should be forwarded exactly once even with repeated session.update. + assert backend_ws.send.await_count == 1 + assert streaming.session_configuration_request is not None + sent_payload = json.loads(backend_ws.send.await_args_list[0].args[0]) + assert "setup" in sent_payload + + def test_collect_session_tools_from_session_update(): """ Test that tools from session.update events are collected. @@ -879,6 +923,169 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error(): litellm.callbacks = [] # cleanup +@pytest.mark.asyncio +async def test_realtime_function_call_output_guardrail_blocks_and_returns_error(): + """ + Test that a client-supplied function_call_output whose content triggers a + guardrail is blocked: it is not forwarded to the backend, and an error + event is sent to the client. + """ + from fastapi import HTTPException + + import litellm + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.types.guardrails import GuardrailEventHooks + + class BlockingGuardrail(CustomGuardrail): + async def apply_guardrail( + self, inputs, request_data, input_type, logging_obj=None + ): + texts = inputs.get("texts", []) + for text in texts: + if "@" in text: + raise HTTPException( + status_code=403, + detail={"error": "email address detected"}, + ) + return inputs + + guardrail = BlockingGuardrail( + guardrail_name="email-blocker", + event_hook=GuardrailEventHooks.pre_call, + default_on=True, + ) + litellm.callbacks = [guardrail] + + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + backend_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None)) + + logging_obj = MagicMock() + logging_obj.pre_call = MagicMock() + + streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) + + item_create_msg = json.dumps( + { + "type": "conversation.item.create", + "item": { + "type": "function_call_output", + "call_id": "call_123", + "output": "Tool says: my email is test@example.com", + }, + } + ) + + client_ws.receive_text = AsyncMock( + side_effect=[ + item_create_msg, + Exception("connection closed"), + ] + ) + + await streaming.client_ack_messages() + + sent_texts = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list] + error_events = [e for e in sent_texts if e.get("type") == "error"] + assert len(error_events) == 1, f"Expected one error event, got: {sent_texts}" + assert error_events[0]["error"]["type"] == "guardrail_violation" + + sent_to_backend = [c.args[0] for c in backend_ws.send.call_args_list if c.args] + forwarded_tool_outputs = [ + json.loads(m) + for m in sent_to_backend + if isinstance(m, str) + and json.loads(m).get("type") == "conversation.item.create" + and json.loads(m).get("item", {}).get("type") == "function_call_output" + ] + # A sanitized placeholder must reach the backend so providers that pair + # every toolCall with a toolResponse (Gemini/Vertex Live) exit their + # pending-tool-call state instead of stalling. The placeholder must NOT + # contain any of the blocked content. + assert len(forwarded_tool_outputs) == 1, ( + f"Sanitized function_call_output should be forwarded, got: " + f"{forwarded_tool_outputs}" + ) + sanitized_item = forwarded_tool_outputs[0]["item"] + assert sanitized_item["call_id"] == "call_123" + assert "test@example.com" not in sanitized_item["output"] + + litellm.callbacks = [] # cleanup + + +@pytest.mark.asyncio +async def test_realtime_function_call_output_guardrail_allows_clean_output(): + """ + Test that a clean function_call_output passes through and reaches the backend + when guardrails are configured. + """ + import litellm + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.types.guardrails import GuardrailEventHooks + + class BlockingGuardrail(CustomGuardrail): + async def apply_guardrail( + self, inputs, request_data, input_type, logging_obj=None + ): + return inputs + + guardrail = BlockingGuardrail( + guardrail_name="noop", + event_hook=GuardrailEventHooks.pre_call, + default_on=True, + ) + litellm.callbacks = [guardrail] + + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + backend_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None)) + + logging_obj = MagicMock() + logging_obj.pre_call = MagicMock() + + streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) + + item_create_msg = json.dumps( + { + "type": "conversation.item.create", + "item": { + "type": "function_call_output", + "call_id": "call_456", + "output": '{"temperature": 72, "unit": "F"}', + }, + } + ) + + client_ws.receive_text = AsyncMock( + side_effect=[ + item_create_msg, + Exception("connection closed"), + ] + ) + + await streaming.client_ack_messages() + + sent_to_backend = [c.args[0] for c in backend_ws.send.call_args_list if c.args] + forwarded = [ + json.loads(m) + for m in sent_to_backend + if isinstance(m, str) + and json.loads(m).get("type") == "conversation.item.create" + and json.loads(m).get("item", {}).get("type") == "function_call_output" + ] + assert ( + len(forwarded) == 1 + ), f"Clean function_call_output should be forwarded, got: {forwarded}" + + litellm.callbacks = [] # cleanup + + @pytest.mark.asyncio async def test_realtime_text_input_guardrail_uses_pre_call_mode(): """ @@ -1160,3 +1367,406 @@ async def test_on_violation_end_session_closes_on_first_fail(): assert streaming._violation_count == 1 litellm.callbacks = [] # cleanup + + +@pytest.mark.asyncio +async def test_provider_path_suppresses_duplicate_session_created_after_synthetic(): + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + + backend_ws = MagicMock() + backend_ws.recv = AsyncMock( + side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)] + ) + backend_ws.send = AsyncMock() + + provider_config = MagicMock() + provider_config.transform_realtime_response = MagicMock( + return_value={ + "response": [ + { + "type": "session.created", + "event_id": "event_1", + "session": {"id": "sess_1", "modalities": ["audio"]}, + } + ], + "current_output_item_id": None, + "current_response_id": None, + "current_delta_chunks": [], + "current_conversation_id": None, + "current_item_chunks": [], + "current_delta_type": None, + "session_configuration_request": None, + } + ) + + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_1" + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + streaming = RealTimeStreaming( + websocket=client_ws, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + # Simulate synthetic session.created already sent by llm_http_handler. + streaming._session_created_sent_to_client = True + + await streaming.backend_to_client_send_messages() + + sent_payloads = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list] + assert not any( + payload.get("type") == "session.created" for payload in sent_payloads + ), f"Expected duplicate session.created to be suppressed, got: {sent_payloads}" + + +@pytest.mark.asyncio +async def test_duplicate_session_created_still_triggers_guardrail_turn_detection_update(): + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + + backend_ws = MagicMock() + backend_ws.recv = AsyncMock( + side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)] + ) + backend_ws.send = AsyncMock() + + provider_config = MagicMock() + provider_config.transform_realtime_response = MagicMock( + return_value={ + "response": [ + { + "type": "session.created", + "event_id": "event_1", + "session": {"id": "sess_1", "modalities": ["audio"]}, + } + ], + "current_output_item_id": None, + "current_response_id": None, + "current_delta_chunks": [], + "current_conversation_id": None, + "current_item_chunks": [], + "current_delta_type": None, + "session_configuration_request": None, + } + ) + + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_1" + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + streaming = RealTimeStreaming( + websocket=client_ws, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + # Synthetic session.created already sent by llm_http_handler. + streaming._session_created_sent_to_client = True + streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign] + streaming._send_to_backend = AsyncMock() # type: ignore[method-assign] + + await streaming.backend_to_client_send_messages() + + # Duplicate session.created should still cause the one-time guardrail + # turn_detection update to be sent to backend. + assert streaming._send_to_backend.await_count == 1 + sent_update = json.loads(streaming._send_to_backend.await_args_list[0].args[0]) + assert sent_update["type"] == "session.update" + injected_session = sent_update["session"] + assert injected_session["type"] == "realtime" + assert ( + injected_session["audio"]["input"]["turn_detection"]["create_response"] is False + ) + + +@pytest.mark.asyncio +async def test_guardrail_update_respects_idempotency_flag(): + """Verify guardrail turn-detection update uses idempotency flag correctly.""" + client_ws = AsyncMock() + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_1" + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + provider_config = MagicMock() + provider_config.transform_realtime_request = MagicMock( + side_effect=lambda msg, model, session_config: [msg] + ) + + streaming = RealTimeStreaming( + websocket=client_ws, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign] + + # First call should send the update + assert streaming._guardrail_turn_detection_update_sent is False + await streaming._maybe_send_guardrail_turn_detection_update() + assert streaming._guardrail_turn_detection_update_sent is True + assert backend_ws.send.await_count == 1 + + # Second call should be a no-op (idempotent) + await streaming._maybe_send_guardrail_turn_detection_update() + assert backend_ws.send.await_count == 1 # Still 1, not 2 + + +@pytest.mark.asyncio +async def test_guardrail_turn_detection_injected_into_first_session_update_deferred_mode(): + """Verify turn_detection is injected into first session.update in deferred mode.""" + client_ws = AsyncMock() + client_ws.receive_text = AsyncMock( + side_effect=[ + json.dumps( + { + "type": "session.update", + "session": { + "modalities": ["text", "audio"], + "tools": [{"type": "function", "name": "get_weather"}], + }, + } + ), + ConnectionClosed(None, None), + ] + ) + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_1" + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + provider_config = MagicMock() + transformed_messages = [] + + def mock_transform(msg, model, session_config): + transformed_messages.append((msg, session_config)) + return [msg] # Pass through for simplicity + + provider_config.transform_realtime_request = MagicMock(side_effect=mock_transform) + + streaming = RealTimeStreaming( + websocket=client_ws, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign] + + # Simulate first session.update in deferred mode + await streaming.client_ack_messages() + + # Verify turn_detection was injected into the session.update. The + # injection runs before the GA remap, so the create_response flag ends + # up nested under audio.input.turn_detection in the GA-shaped payload. + assert len(transformed_messages) == 1 + transformed_msg, session_config = transformed_messages[0] + msg_obj = json.loads(transformed_msg) + assert msg_obj["type"] == "session.update" + session_obj = msg_obj["session"] + injected_turn_detection = session_obj.get("turn_detection") or session_obj.get( + "audio", {} + ).get("input", {}).get("turn_detection") + assert injected_turn_detection is not None + assert injected_turn_detection["create_response"] is False + assert streaming._guardrail_turn_detection_update_sent is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("existing_turn_detection", [None, "auto", 42, ["server_vad"]]) +async def test_guardrail_turn_detection_injection_tolerates_non_dict_value( + existing_turn_detection, +): + """Client-supplied non-dict turn_detection must not crash client_ack_messages.""" + client_ws = AsyncMock() + client_ws.receive_text = AsyncMock( + side_effect=[ + json.dumps( + { + "type": "session.update", + "session": { + "modalities": ["text", "audio"], + "turn_detection": existing_turn_detection, + }, + } + ), + ConnectionClosed(None, None), + ] + ) + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_1" + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + provider_config = MagicMock() + transformed_messages = [] + + def mock_transform(msg, model, session_config): + transformed_messages.append((msg, session_config)) + return [msg] + + provider_config.transform_realtime_request = MagicMock(side_effect=mock_transform) + + streaming = RealTimeStreaming( + websocket=client_ws, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign] + + await streaming.client_ack_messages() + + assert len(transformed_messages) == 1 + transformed_msg, _ = transformed_messages[0] + msg_obj = json.loads(transformed_msg) + session_obj = msg_obj["session"] + injected_turn_detection = session_obj.get("turn_detection") or session_obj.get( + "audio", {} + ).get("input", {}).get("turn_detection") + assert isinstance(injected_turn_detection, dict) + assert injected_turn_detection["create_response"] is False + assert streaming._guardrail_turn_detection_update_sent is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "client_session", + [ + {"turn_detection": {"type": "server_vad", "create_response": True}}, + { + "audio": { + "input": { + "turn_detection": {"type": "server_vad", "create_response": True} + } + } + }, + ], +) +async def test_subsequent_session_update_cannot_reenable_vad_when_guardrails_active( + client_session, +): + """A subsequent client session.update must not be allowed to flip + ``create_response`` back to True once audio transcription guardrails have + disabled VAD auto-response. Covers both the flat beta shape and the + nested GA ``audio.input.turn_detection`` shape. + """ + client_ws = AsyncMock() + client_ws.receive_text = AsyncMock( + side_effect=[ + json.dumps({"type": "session.update", "session": client_session}), + ConnectionClosed(None, None), + ] + ) + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_1" + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + provider_config = MagicMock() + transformed_messages = [] + + def mock_transform(msg, model, session_config): + transformed_messages.append((msg, session_config)) + return [msg] + + provider_config.transform_realtime_request = MagicMock(side_effect=mock_transform) + + streaming = RealTimeStreaming( + websocket=client_ws, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign] + # Simulate that initial setup + guardrail disable have already happened. + streaming.session_configuration_request = json.dumps({"setup": {"model": "x"}}) + streaming._guardrail_turn_detection_update_sent = True + + await streaming.client_ack_messages() + + assert len(transformed_messages) == 1 + forwarded_msg, _ = transformed_messages[0] + msg_obj = json.loads(forwarded_msg) + session_obj = msg_obj["session"] + forwarded_turn_detection = session_obj.get("turn_detection") or session_obj.get( + "audio", {} + ).get("input", {}).get("turn_detection") + assert isinstance(forwarded_turn_detection, dict) + assert forwarded_turn_detection["create_response"] is False + + +@pytest.mark.asyncio +async def test_follow_up_setup_updates_cached_session_configuration_request(): + """A follow-up setup produced by a subsequent session.update must replace + the cached ``session_configuration_request`` so downstream readers + (e.g. modality lookup in ``response.created``) see the latest config.""" + client_ws = AsyncMock() + client_ws.receive_text = AsyncMock( + side_effect=[ + json.dumps({"type": "session.update", "session": {"tools": []}}), + ConnectionClosed(None, None), + ] + ) + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + + logging_obj = MagicMock() + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + provider_config = MagicMock() + follow_up_setup = json.dumps( + { + "setup": { + "model": "models/gemini-2.5-flash", + "generationConfig": {"responseModalities": ["TEXT"]}, + "tools": [{"function_declarations": []}], + } + } + ) + provider_config.transform_realtime_request = MagicMock( + return_value=[follow_up_setup] + ) + + streaming = RealTimeStreaming( + websocket=client_ws, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + # Simulate that the original auto-setup was already cached. + streaming.session_configuration_request = json.dumps( + { + "setup": { + "model": "models/gemini-2.5-flash", + "generationConfig": {"responseModalities": ["AUDIO"]}, + } + } + ) + + await streaming.client_ack_messages() + + assert streaming.session_configuration_request == follow_up_setup 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 cc0a32d2ce6..cc0adc4277c 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 @@ -20,8 +20,10 @@ def test_gemini_realtime_transformation_session_created(): assert config is not None session_configuration_request = { - "model": "gemini-1.5-flash", - "generationConfig": {"responseModalities": ["TEXT"]}, + "setup": { + "model": "gemini-1.5-flash", + "generationConfig": {"responseModalities": ["TEXT"]}, + } } session_configuration_request_str = json.dumps(session_configuration_request) session_created_message = {"setupComplete": {}} @@ -45,8 +47,54 @@ def test_gemini_realtime_transformation_session_created(): }, ) - print(transformed_message) - assert transformed_message["response"][0]["type"] == "session.created" + session_created = transformed_message["response"][0] + assert session_created["type"] == "session.created" + # Verify the setup-wrapped configuration reaches the modality lookup so + # the synthetic session.created reflects the cached responseModalities. + assert session_created["session"]["modalities"] == ["text"] + + +def test_session_created_does_not_overwrite_session_configuration_request(): + config = GeminiRealtimeConfig() + + session_configuration_request_str = json.dumps( + { + "setup": { + "model": "models/gemini-2.5-flash-native-audio", + "generationConfig": {"responseModalities": ["AUDIO"]}, + } + } + ) + setup_complete_message = {"setupComplete": {}} + + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_123" + + transformed = config.transform_realtime_response( + json.dumps(setup_complete_message), + "gemini-2.5-flash-native-audio", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": session_configuration_request_str, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + # Must keep original setup payload (with "setup"), not overwrite with session.created event. + assert ( + transformed["session_configuration_request"] + == session_configuration_request_str + ) + + # Also verify emitted session.created reflects audio modality from setup payload. + session_created = transformed["response"][0] + assert session_created["type"] == "session.created" + assert "audio" in session_created["session"]["modalities"] def test_gemini_realtime_transformation_content_delta(): @@ -54,8 +102,10 @@ def test_gemini_realtime_transformation_content_delta(): assert config is not None session_configuration_request = { - "model": "gemini-1.5-flash", - "generationConfig": {"responseModalities": ["TEXT"]}, + "setup": { + "model": "gemini-1.5-flash", + "generationConfig": {"responseModalities": ["TEXT"]}, + } } session_configuration_request_str = json.dumps(session_configuration_request) session_created_message = { @@ -147,8 +197,10 @@ def test_gemini_realtime_transformation_audio_delta(): assert config is not None session_configuration_request = { - "model": "gemini-1.5-flash", - "generationConfig": {"responseModalities": ["AUDIO"]}, + "setup": { + "model": "gemini-1.5-flash", + "generationConfig": {"responseModalities": ["AUDIO"]}, + } } session_configuration_request_str = json.dumps(session_configuration_request) @@ -196,8 +248,10 @@ def test_gemini_realtime_transformation_generation_complete(): assert config is not None session_configuration_request = { - "model": "gemini-1.5-flash", - "generationConfig": {"responseModalities": ["AUDIO"]}, + "setup": { + "model": "gemini-1.5-flash", + "generationConfig": {"responseModalities": ["AUDIO"]}, + } } session_configuration_request_str = json.dumps(session_configuration_request) @@ -225,9 +279,9 @@ def test_gemini_realtime_transformation_generation_complete(): contains_audio_done_event = False for response in responses: if response["type"] == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DONE.value: - contains_audio_delta = True + contains_audio_done_event = True break - assert contains_audio_delta, "Expected audio delta event" + assert contains_audio_done_event, "Expected audio done event" def test_gemini_3_1_flash_live_preview_model_cost_map_entry(): @@ -242,3 +296,1211 @@ def test_gemini_3_1_flash_live_preview_model_cost_map_entry(): assert info.get("max_output_tokens") == 65536 assert "video" in info.get("supported_modalities", []) assert info.get("supports_function_calling") is True + + +def test_gemini_realtime_tool_call_transformation(): + """Test transformation of Gemini toolCall to OpenAI function_call_arguments.done format.""" + config = GeminiRealtimeConfig() + + # Gemini toolCall message format + gemini_tool_call = { + "toolCall": { + "functionCalls": [ + { + "id": "call_123", + "name": "get_weather", + "args": {"location": "San Francisco", "unit": "fahrenheit"}, + } + ] + } + } + + gemini_tool_call_str = json.dumps(gemini_tool_call) + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "test-trace-123" + + # Transform the toolCall message + result = config.transform_realtime_response( + gemini_tool_call_str, + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": "item_123", + "current_response_id": "resp_123", + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + print("Tool call transformation result:", json.dumps(result, indent=2)) + + # Verify the transformation + responses = result["response"] + assert len(responses) > 0, "Expected at least one response event" + + # Find the function_call_arguments.done event + function_call_event = None + for event in responses: + if event.get("type") == "response.function_call_arguments.done": + function_call_event = event + break + + assert ( + function_call_event is not None + ), "Expected function_call_arguments.done event" + assert function_call_event["call_id"] == "call_123" + assert function_call_event["name"] == "get_weather" + assert function_call_event["response_id"] == "resp_123" + assert function_call_event["item_id"] == "item_123_tool_0" + assert function_call_event["output_index"] == 0 + + # Verify arguments are properly serialized as JSON string + args = json.loads(function_call_event["arguments"]) + assert args["location"] == "San Francisco" + assert args["unit"] == "fahrenheit" + + +def test_gemini_realtime_session_update_with_tools(): + """Test transformation of OpenAI session.update with tools to Gemini setup format.""" + config = GeminiRealtimeConfig() + + # OpenAI format session update with tools + session_update = { + "type": "session.update", + "session": { + "instructions": "You are a helpful assistant with weather tools.", + "temperature": 0.7, + "max_response_output_tokens": 1024, + "modalities": ["audio"], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the current weather for a location.", + "parameters": { + "type": "object", + "properties": { + "location": { + "type": "string", + "description": "The city name", + }, + "unit": { + "type": "string", + "enum": ["fahrenheit", "celsius"], + }, + }, + "required": ["location"], + }, + }, + } + ], + }, + } + + # Transform to Gemini format (first session.update, so setup should be sent) + messages = config.transform_realtime_request( + json.dumps(session_update), + "gemini-2.5-flash", + session_configuration_request=None, + ) + + assert len(messages) == 1, "Expected one setup message" + + gemini_setup = json.loads(messages[0]) + assert "setup" in gemini_setup + + setup_config = gemini_setup["setup"] + + # Verify tools are at top level, not in generationConfig + assert "tools" in setup_config + assert "tools" not in setup_config.get("generationConfig", {}) + + # Verify tool structure matches Gemini format + tools = setup_config["tools"] + assert len(tools) == 1 + assert "function_declarations" in tools[0] + + function_decl = tools[0]["function_declarations"][0] + assert function_decl["name"] == "get_weather" + assert "Get the current weather" in function_decl["description"] + assert "parameters" in function_decl + + +def test_gemini_session_update_defaults_to_audio_modality(): + config = GeminiRealtimeConfig() + + session_update = { + "type": "session.update", + "session": { + "instructions": "You are a helpful assistant.", + # No modalities on purpose + }, + } + + messages = config.transform_realtime_request( + json.dumps(session_update), + "gemini-2.5-flash", + session_configuration_request=None, + ) + + assert len(messages) == 1 + setup_payload = json.loads(messages[0])["setup"] + assert setup_payload["generationConfig"]["responseModalities"] == ["AUDIO"] + + +def test_gemini_requires_session_configuration_feature_flag(monkeypatch): + config = GeminiRealtimeConfig() + + # Default behavior remains backwards-compatible (auto setup on connect) + monkeypatch.setattr(litellm, "gemini_live_defer_setup", False, raising=False) + assert config.requires_session_configuration() is True + + # Opt-in behavior: defer setup until client sends session.update + monkeypatch.setattr(litellm, "gemini_live_defer_setup", True, raising=False) + assert config.requires_session_configuration() is False + + +def test_gemini_realtime_function_call_output_transformation(): + """Test transformation of OpenAI function_call_output to Gemini toolResponse format. + + Exercises the full production round-trip: a Gemini toolCall arrives first + and populates the call_id -> name mapping, then the OpenAI + function_call_output is transformed and must carry the function name back + to Gemini in functionResponses. + """ + config = GeminiRealtimeConfig() + + # Receive a toolCall from Gemini first to populate the call_id -> name mapping. + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_func_output" + config.transform_realtime_response( + json.dumps( + { + "toolCall": { + "functionCalls": [ + { + "id": "call_123", + "name": "get_weather", + "args": {"location": "San Francisco"}, + } + ] + } + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + assert config._tool_call_id_to_name.get("call_123") == "get_weather" + + # OpenAI format function call output + function_output = { + "type": "conversation.item.create", + "item": { + "type": "function_call_output", + "call_id": "call_123", + "output": json.dumps( + { + "location": "San Francisco", + "temperature": 72, + "unit": "fahrenheit", + "conditions": "sunny", + } + ), + }, + } + + # Transform to Gemini format + messages = config.transform_realtime_request( + json.dumps(function_output), + "gemini-2.5-flash", + session_configuration_request="existing", + ) + + assert len(messages) == 1, "Expected one toolResponse message" + + gemini_response = json.loads(messages[0]) + assert "toolResponse" in gemini_response + + tool_response = gemini_response["toolResponse"] + assert "functionResponses" in tool_response + assert len(tool_response["functionResponses"]) == 1 + + func_response = tool_response["functionResponses"][0] + assert func_response["id"] == "call_123" + assert func_response["name"] == "get_weather" + assert "response" in func_response + assert func_response["response"]["temperature"] == 72 + assert func_response["response"]["conditions"] == "sunny" + + # A retry of the same function_call_output (e.g. a client SDK that + # re-sends the result) must still produce a functionResponses payload + # carrying ``name`` — the call_id → name mapping must not be evicted + # after the first lookup. + retry_messages = config.transform_realtime_request( + json.dumps(function_output), + "gemini-2.5-flash", + session_configuration_request="existing", + ) + retry_response = json.loads(retry_messages[0])["toolResponse"]["functionResponses"][ + 0 + ] + assert retry_response["name"] == "get_weather" + + +def test_gemini_realtime_user_text_transformation(): + """Test transformation of OpenAI user message to Gemini clientContent format.""" + config = GeminiRealtimeConfig() + + # OpenAI format user message + user_message = { + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [ + {"type": "input_text", "text": "What's the weather in London?"} + ], + }, + } + + # Transform to Gemini format + messages = config.transform_realtime_request( + json.dumps(user_message), + "gemini-2.5-flash", + session_configuration_request="existing", + ) + + assert len(messages) == 1, "Expected one clientContent message" + + gemini_message = json.loads(messages[0]) + assert "clientContent" in gemini_message + + client_content = gemini_message["clientContent"] + assert "turns" in client_content + assert len(client_content["turns"]) == 1 + + turn = client_content["turns"][0] + assert turn["role"] == "user" + assert len(turn["parts"]) == 1 + assert turn["parts"][0]["text"] == "What's the weather in London?" + assert client_content["turnComplete"] is True + + +def test_return_new_content_delta_events_without_session_config_does_not_error(): + config = GeminiRealtimeConfig() + + events = config.return_new_content_delta_events( + response_id="resp_1", + output_item_id="item_1", + conversation_id="conv_1", + delta_type="text", + session_configuration_request=None, + ) + + assert len(events) >= 1 + assert events[0]["type"] == "response.created" + + +def test_gemini_realtime_multi_tool_calls_have_unique_item_ids(): + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "test-trace-123" + + gemini_tool_call = { + "toolCall": { + "functionCalls": [ + { + "id": "call_1", + "name": "get_weather", + "args": {"location": "SF"}, + }, + { + "id": "call_2", + "name": "get_weather", + "args": {"location": "NYC"}, + }, + ] + } + } + + result = config.transform_realtime_response( + json.dumps(gemini_tool_call), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": "item_123", + "current_response_id": "resp_123", + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + responses = [ + ev + for ev in result["response"] + if ev.get("type") == "response.function_call_arguments.done" + ] + assert len(responses) == 2 + assert responses[0]["response_id"] == "resp_123" + assert responses[1]["response_id"] == "resp_123" + assert responses[0]["item_id"] == "item_123_tool_0" + assert responses[1]["item_id"] == "item_123_tool_1" + assert responses[0]["item_id"] != responses[1]["item_id"] + assert responses[0]["output_index"] == 0 + assert responses[1]["output_index"] == 1 + + +def test_gemini_session_update_includes_input_audio_transcription_default(): + """Verify _handle_session_update includes inputAudioTranscription default.""" + config = GeminiRealtimeConfig() + session_update = { + "type": "session.update", + "session": { + "modalities": ["text", "audio"], + "tools": [ + { + "type": "function", + "name": "get_weather", + "description": "Get weather", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + }, + } + ], + }, + } + + result = config.transform_realtime_request( + json.dumps(session_update), + "gemini-2.5-flash", + session_configuration_request=None, + ) + + assert len(result) == 1 + setup = json.loads(result[0]) + assert "setup" in setup + assert "inputAudioTranscription" in setup["setup"] + assert setup["setup"]["inputAudioTranscription"] == {} + + +def test_gemini_tool_call_emits_response_created_preamble(): + """Verify response.created is emitted before tool call events when response_id is None.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_123" + + gemini_tool_call = { + "toolCall": { + "functionCalls": [ + { + "id": "call_123", + "name": "get_weather", + "args": {"location": "San Francisco", "unit": "fahrenheit"}, + } + ] + } + } + + # Transform with current_response_id=None to trigger preamble emission + result = config.transform_realtime_response( + json.dumps(gemini_tool_call), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + responses = result["response"] + # Should have: response.created, output_item.added, function_call_arguments.delta, function_call_arguments.done, output_item.done, conversation.item.created, response.done + assert len(responses) >= 7 + assert responses[0]["type"] == "response.created" + assert "response" in responses[0] + assert responses[0]["response"]["status"] == "in_progress" + # response.created on the tool-call path mirrors the audio/text preamble: + # modalities/temperature/max_output_tokens are present so spec-compliant + # clients see consistent response metadata regardless of payload type. + assert "modalities" in responses[0]["response"] + assert "temperature" in responses[0]["response"] + assert "max_output_tokens" in responses[0]["response"] + assert responses[1]["type"] == "response.output_item.added" + assert responses[1]["item"]["type"] == "function_call" + assert responses[1]["item"]["status"] == "in_progress" + assert responses[2]["type"] == "response.function_call_arguments.delta" + assert responses[2]["call_id"] == "call_123" + assert responses[2]["delta"] == responses[3]["arguments"] + assert responses[3]["type"] == "response.function_call_arguments.done" + assert responses[4]["type"] == "response.output_item.done" + assert responses[4]["item"]["type"] == "function_call" + assert responses[4]["item"]["status"] == "completed" + assert responses[5]["type"] == "conversation.item.created" + assert responses[5]["item"]["type"] == "function_call" + assert responses[5]["item"]["status"] == "completed" + assert responses[6]["type"] == "response.done" + assert responses[6]["response"]["status"] == "completed" + assert len(responses[6]["response"]["output"]) == 1 + assert responses[6]["response"]["output"][0]["type"] == "function_call" + assert result["current_output_item_id"] is None + assert result["current_response_id"] is None + + +def test_gemini_tool_call_resets_ids_for_post_tool_model_turn(): + """After tool-call response.done, a subsequent modelTurn must emit response.created.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_123" + + session_configuration_request = json.dumps( + { + "setup": { + "model": "gemini-1.5-flash", + "generationConfig": {"responseModalities": ["TEXT"]}, + } + } + ) + + tool_result = config.transform_realtime_response( + json.dumps( + { + "toolCall": { + "functionCalls": [ + { + "id": "call_123", + "name": "get_weather", + "args": {"location": "San Francisco"}, + } + ] + } + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_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_response_id = tool_result["response"][0]["response"]["id"] + assert tool_result["current_output_item_id"] is None + assert tool_result["current_response_id"] is None + + post_tool_result = config.transform_realtime_response( + json.dumps( + { + "serverContent": { + "modelTurn": {"parts": [{"text": "The weather is sunny."}]} + } + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": session_configuration_request, + "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"], + }, + ) + + post_tool_events = post_tool_result["response"] + assert post_tool_events[0]["type"] == "response.created" + assert post_tool_events[0]["response"]["id"] != tool_response_id + assert ( + post_tool_result["current_response_id"] == post_tool_events[0]["response"]["id"] + ) + + +def test_gemini_empty_tool_call_does_not_crash_websocket(): + """A toolCall payload with no functionCalls must not raise the + 'Unknown message type' guard — that would terminate the WebSocket session + on what is at worst a benign no-op from Gemini.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_empty_tool_call" + + result = config.transform_realtime_response( + json.dumps({"toolCall": {"functionCalls": []}}), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + assert result["response"] == [] + assert result["current_response_id"] is None + assert result["current_output_item_id"] is None + + +def test_gemini_empty_tool_call_with_sibling_usage_metadata_does_not_crash(): + """A toolCall with empty functionCalls alongside a sibling key (e.g. + ``usageMetadata``) must still be handled as a benign no-op: the empty + toolCall is consumed and the metadata sibling is skipped, without + raising ``Unknown message type``.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_empty_tool_call_with_sibling" + + result = config.transform_realtime_response( + json.dumps( + { + "toolCall": {"functionCalls": []}, + "usageMetadata": {"totalTokenCount": 7}, + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": "item_existing", + "current_response_id": "resp_existing", + "current_conversation_id": "conv_existing", + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + assert result["response"] == [] + # In-flight response IDs must survive the benign no-op. + assert result["current_response_id"] == "resp_existing" + assert result["current_output_item_id"] == "item_existing" + + +def test_gemini_tool_call_response_done_includes_usage_from_sibling_metadata(): + """A ``toolCall`` frame with a sibling ``usageMetadata`` must propagate the + real token counts onto the emitted ``response.done`` so spend/budget + accounting records tokens consumed by the tool-call turn — otherwise an + authenticated client can repeatedly drive tool calls with zero spend.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_tool_call_usage" + + result = config.transform_realtime_response( + json.dumps( + { + "toolCall": { + "functionCalls": [ + { + "id": "call_usage", + "name": "get_weather", + "args": {"location": "NYC"}, + } + ] + }, + "usageMetadata": { + "promptTokenCount": 17, + "responseTokenCount": 4, + "totalTokenCount": 21, + }, + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + response_done = next( + ev for ev in result["response"] if ev.get("type") == "response.done" + ) + usage = response_done["response"]["usage"] + assert usage["input_tokens"] == 17 + assert usage["output_tokens"] == 4 + assert usage["total_tokens"] == 21 + + +def test_gemini_tool_call_response_done_falls_back_to_empty_usage(): + """Without sibling ``usageMetadata`` the tool-call ``response.done`` still + carries a valid empty usage block so OpenAI-compatible clients (which + expect ``usage`` on every ``response.done``) don't break.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_tool_call_no_usage" + + result = config.transform_realtime_response( + json.dumps( + { + "toolCall": { + "functionCalls": [ + { + "id": "call_no_usage", + "name": "get_weather", + "args": {"location": "NYC"}, + } + ] + } + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + response_done = next( + ev for ev in result["response"] if ev.get("type") == "response.done" + ) + usage = response_done["response"]["usage"] + assert usage["input_tokens"] == 0 + assert usage["output_tokens"] == 0 + assert usage["total_tokens"] == 0 + + +def test_gemini_function_call_output_includes_name(): + """Verify function_call_output includes name field from stored mapping.""" + config = GeminiRealtimeConfig() + + # First, receive a toolCall from Gemini (this stores the call_id → name mapping) + gemini_tool_call = { + "toolCall": { + "functionCalls": [ + { + "id": "call_123", + "name": "get_weather", + "args": {"location": "San Francisco"}, + } + ] + } + } + + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_123" + + config.transform_realtime_response( + json.dumps(gemini_tool_call), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + # Verify mapping was stored + assert "call_123" in config._tool_call_id_to_name + assert config._tool_call_id_to_name["call_123"] == "get_weather" + + # Now send a function_call_output back (this should include the name) + function_output = { + "type": "conversation.item.create", + "item": { + "type": "function_call_output", + "call_id": "call_123", + "output": json.dumps({"result": "72 degrees"}), + }, + } + + result = config.transform_realtime_request( + json.dumps(function_output), + "gemini-2.5-flash", + session_configuration_request="{}", + ) + + assert len(result) == 1 + tool_response = json.loads(result[0]) + assert "toolResponse" in tool_response + assert "functionResponses" in tool_response["toolResponse"] + assert len(tool_response["toolResponse"]["functionResponses"]) == 1 + + function_response = tool_response["toolResponse"]["functionResponses"][0] + assert function_response["id"] == "call_123" + assert function_response["name"] == "get_weather" # ✅ Name is included + assert "response" in function_response + + +def test_gemini_subsequent_session_update_forwards_tools_merged_with_original_setup(): + """A client session.update sent after the auto-setup must forward tools/ + instructions as a follow-up setup, merged with the original setup so we + don't drop the pre-existing config (model, generationConfig, etc.).""" + config = GeminiRealtimeConfig() + + original_setup = { + "setup": { + "model": "models/gemini-2.5-flash-native-audio", + "generationConfig": {"responseModalities": ["AUDIO"]}, + "inputAudioTranscription": {}, + "systemInstruction": {"role": "user", "parts": [{"text": "original"}]}, + } + } + + session_update = { + "type": "session.update", + "session": { + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather.", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + }, + }, + } + ], + "instructions": "Be concise.", + }, + } + + messages = config.transform_realtime_request( + json.dumps(session_update), + "gemini-2.5-flash-native-audio", + session_configuration_request=json.dumps(original_setup), + ) + + assert len(messages) == 1 + follow_up = json.loads(messages[0])["setup"] + assert "tools" in follow_up + assert follow_up["tools"][0]["function_declarations"][0]["name"] == "get_weather" + # systemInstruction overwritten by client's instructions + assert follow_up["systemInstruction"]["parts"][0]["text"] == "Be concise." + # Original generationConfig / model / inputAudioTranscription preserved + assert follow_up["generationConfig"]["responseModalities"] == ["AUDIO"] + assert follow_up["model"] == "models/gemini-2.5-flash-native-audio" + assert follow_up["inputAudioTranscription"] == {} + + +def test_gemini_subsequent_session_update_with_turn_detection_only_preserves_original_tools(): + """A subsequent session.update carrying only turn_detection (the + guardrail-injected disable) must keep the original tools/generationConfig.""" + config = GeminiRealtimeConfig() + + original_setup = { + "setup": { + "model": "models/gemini-2.5-flash-native-audio", + "generationConfig": {"responseModalities": ["AUDIO"]}, + "inputAudioTranscription": {}, + "tools": [ + { + "function_declarations": [ + {"name": "lookup", "description": "x", "parameters": {}} + ] + } + ], + } + } + + session_update = { + "type": "session.update", + "session": {"turn_detection": {"create_response": False}}, + } + + messages = config.transform_realtime_request( + json.dumps(session_update), + "gemini-2.5-flash-native-audio", + session_configuration_request=json.dumps(original_setup), + ) + + assert len(messages) == 1 + follow_up = json.loads(messages[0])["setup"] + assert follow_up["tools"] == original_setup["setup"]["tools"] + assert ( + follow_up["realtimeInputConfig"]["automaticActivityDetection"]["disabled"] + is True + ) + + +def test_gemini_follow_up_session_update_preserves_response_modalities_on_partial_generation_config(): + """A follow-up session.update that only sets `temperature` (or any other + generationConfig sub-field) must not wipe `responseModalities` from the + original setup.""" + config = GeminiRealtimeConfig() + + original_setup = { + "setup": { + "model": "models/gemini-2.5-flash-native-audio", + "generationConfig": { + "responseModalities": ["AUDIO"], + "maxOutputTokens": 2048, + }, + "inputAudioTranscription": {}, + } + } + + session_update = { + "type": "session.update", + "session": {"temperature": 0.7}, + } + + messages = config.transform_realtime_request( + json.dumps(session_update), + "gemini-2.5-flash-native-audio", + session_configuration_request=json.dumps(original_setup), + ) + + follow_up = json.loads(messages[0])["setup"] + assert follow_up["generationConfig"]["responseModalities"] == ["AUDIO"] + assert follow_up["generationConfig"]["maxOutputTokens"] == 2048 + assert follow_up["generationConfig"]["temperature"] == 0.7 + + +def test_gemini_subsequent_session_update_preserves_automatic_activity_detection_subfields(): + config = GeminiRealtimeConfig() + + original_setup = { + "setup": { + "model": "models/gemini-2.5-flash-native-audio", + "generationConfig": {"responseModalities": ["AUDIO"]}, + "realtimeInputConfig": { + "automaticActivityDetection": { + "disabled": False, + "silenceDurationMs": 500, + "prefixPaddingMs": 100, + } + }, + } + } + + session_update = { + "type": "session.update", + "session": {"turn_detection": {"create_response": False}}, + } + + messages = config.transform_realtime_request( + json.dumps(session_update), + "gemini-2.5-flash-native-audio", + session_configuration_request=json.dumps(original_setup), + ) + + automatic_activity_detection = json.loads(messages[0])["setup"][ + "realtimeInputConfig" + ]["automaticActivityDetection"] + assert automatic_activity_detection["disabled"] is True + assert automatic_activity_detection["silenceDurationMs"] == 500 + assert automatic_activity_detection["prefixPaddingMs"] == 100 + + +def test_gemini_tool_call_id_to_name_evicts_oldest_when_capped(): + """The call_id → name LRU must evict the oldest entry once the cap is + reached so long sessions with many tool calls don't grow unboundedly, + while keeping recently-seen call_ids resolvable for retried + function_call_output messages.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_lru" + + config._TOOL_CALL_ID_TO_NAME_MAX = 4 + + for idx in range(8): + config.transform_realtime_response( + json.dumps( + { + "toolCall": { + "functionCalls": [ + { + "id": f"call_{idx}", + "name": f"fn_{idx}", + "args": {}, + } + ] + } + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + assert len(config._tool_call_id_to_name) == 4 + # Most recent 4 retained; oldest 4 evicted. + assert list(config._tool_call_id_to_name) == [ + "call_4", + "call_5", + "call_6", + "call_7", + ] + + +def test_gemini_standalone_usage_metadata_does_not_crash_websocket(): + """A Gemini frame containing only sibling metadata (e.g. a standalone + ``usageMetadata`` block emitted between turns) must not trip the + ``Unknown message type`` guard — that would terminate the WebSocket + session on a benign no-op frame.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_usage_only" + + result = config.transform_realtime_response( + json.dumps( + { + "usageMetadata": { + "promptTokenCount": 12, + "responseTokenCount": 34, + "totalTokenCount": 46, + } + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": "item_existing", + "current_response_id": "resp_existing", + "current_conversation_id": "conv_existing", + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + assert result["response"] == [] + # State must be returned unchanged so subsequent frames continue the + # in-flight response correctly. + assert result["current_output_item_id"] == "item_existing" + assert result["current_response_id"] == "resp_existing" + assert result["current_conversation_id"] == "conv_existing" + + +def test_gemini_standalone_usage_metadata_is_attributed_to_next_tool_call_response_done(): + """A standalone ``usageMetadata`` frame emitted between turns must not + silently drop the consumed tokens. The next tool-call ``response.done`` + must carry those token counts so an authenticated client cannot drive + tool-call turns whose token usage is recorded as zero, bypassing + spend/budget accounting.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_standalone_usage_then_tool_call" + + standalone_result = config.transform_realtime_response( + json.dumps( + { + "usageMetadata": { + "promptTokenCount": 31, + "responseTokenCount": 9, + "totalTokenCount": 40, + } + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + assert standalone_result["response"] == [] + + tool_call_result = config.transform_realtime_response( + json.dumps( + { + "toolCall": { + "functionCalls": [ + { + "id": "call_buffered", + "name": "get_weather", + "args": {"location": "NYC"}, + } + ] + } + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + response_done = next( + ev for ev in tool_call_result["response"] if ev.get("type") == "response.done" + ) + usage = response_done["response"]["usage"] + assert usage["input_tokens"] == 31 + assert usage["output_tokens"] == 9 + assert usage["total_tokens"] == 40 + # Buffer must be cleared after attribution so a subsequent tool-call + # turn without its own usage does not double-count the previous frame. + assert config._pending_usage_metadata is None + + +def test_gemini_standalone_usage_metadata_is_attributed_to_next_response_done(): + """A standalone ``usageMetadata`` frame must also flow into the normal + (non-tool-call) ``response.done`` path so audio/text turns whose usage + arrives in a separate frame are still billed correctly.""" + config = GeminiRealtimeConfig() + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_standalone_usage_then_turn_complete" + + config.transform_realtime_response( + json.dumps( + { + "usageMetadata": { + "promptTokenCount": 5, + "responseTokenCount": 11, + "totalTokenCount": 16, + } + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + turn_complete_result = config.transform_realtime_response( + json.dumps({"serverContent": {"turnComplete": True}}), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + response_done = next( + ev + for ev in turn_complete_result["response"] + if ev.get("type") == "response.done" + ) + usage = response_done["response"]["usage"] + assert usage["input_tokens"] == 5 + assert usage["output_tokens"] == 11 + assert usage["total_tokens"] == 16 + assert config._pending_usage_metadata is None + + +def test_gemini_in_frame_usage_metadata_clears_pending_buffer(): + """When ``usageMetadata`` arrives in the same frame as the closing + ``toolCall`` / ``turnComplete``, the in-frame counts are authoritative + and any buffered standalone metadata must be discarded so a later + turn's ``response.done`` does not double-count tokens.""" + config = GeminiRealtimeConfig() + config._pending_usage_metadata = { + "promptTokenCount": 99, + "responseTokenCount": 99, + "totalTokenCount": 198, + } + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_in_frame_clears_buffer" + + result = config.transform_realtime_response( + json.dumps( + { + "toolCall": { + "functionCalls": [ + { + "id": "call_in_frame", + "name": "get_weather", + "args": {"location": "NYC"}, + } + ] + }, + "usageMetadata": { + "promptTokenCount": 3, + "responseTokenCount": 2, + "totalTokenCount": 5, + }, + } + ), + "gemini-2.5-flash", + logging_obj, + realtime_response_transform_input={ + "session_configuration_request": None, + "current_output_item_id": None, + "current_response_id": None, + "current_conversation_id": None, + "current_delta_chunks": [], + "current_item_chunks": [], + "current_delta_type": None, + }, + ) + + response_done = next( + ev for ev in result["response"] if ev.get("type") == "response.done" + ) + usage = response_done["response"]["usage"] + assert usage["input_tokens"] == 3 + assert usage["output_tokens"] == 2 + assert usage["total_tokens"] == 5 + assert config._pending_usage_metadata is None 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 1baaf912568..0ad614099de 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 @@ -19,6 +19,7 @@ import websockets.exceptions # registers websockets.exceptions on the websocket sys.path.insert(0, os.path.abspath("../../../../..")) +import litellm from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig # --------------------------------------------------------------------------- @@ -82,6 +83,85 @@ def test_session_configuration_request_model_format(): ) +def test_vertex_requires_session_configuration_feature_flag(monkeypatch): + cfg = VertexAIRealtimeConfig( + access_token="tok", project="my-proj", location="us-central1" + ) + + # Default remains backwards-compatible (auto setup on connect) + monkeypatch.setattr(litellm, "gemini_live_defer_setup", False, raising=False) + assert cfg.requires_session_configuration() is True + + # Opt-in deferred setup for tool-injection flow + monkeypatch.setattr(litellm, "gemini_live_defer_setup", True, raising=False) + assert cfg.requires_session_configuration() is False + + +def test_vertex_session_update_defaults_to_audio_modality(): + cfg = VertexAIRealtimeConfig( + access_token="tok", project="my-proj", location="us-central1" + ) + + session_update = { + "type": "session.update", + "session": { + "instructions": "You are a helpful assistant.", + # No modalities provided on purpose + }, + } + + messages = cfg.transform_realtime_request( + json.dumps(session_update), + "gemini-live-2.5-flash-native-audio", + session_configuration_request=None, + ) + assert len(messages) == 1 + setup_payload = json.loads(messages[0])["setup"] + assert setup_payload["generationConfig"]["responseModalities"] == ["AUDIO"] + + +def test_vertex_session_update_normalizes_ga_remapped_fields(): + """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`` + runs so client preferences aren't silently dropped. + """ + cfg = VertexAIRealtimeConfig( + access_token="tok", project="my-proj", location="us-central1" + ) + + session_update = { + "type": "session.update", + "session": { + "instructions": "Be concise.", + "output_modalities": ["text"], + "audio": { + "input": { + "transcription": {}, + "turn_detection": {"silence_duration_ms": 1500}, + }, + }, + }, + } + + messages = cfg.transform_realtime_request( + json.dumps(session_update), + "gemini-live-2.5-flash-native-audio", + session_configuration_request=None, + ) + assert len(messages) == 1 + setup_payload = json.loads(messages[0])["setup"] + + assert setup_payload["generationConfig"]["responseModalities"] == ["TEXT"] + assert setup_payload["inputAudioTranscription"] == {} + assert ( + setup_payload["realtimeInputConfig"]["automaticActivityDetection"][ + "silenceDurationMs" + ] + == 1500 + ) + + # --------------------------------------------------------------------------- # Round-trip test: text-in / text-out via RealTimeStreaming # --------------------------------------------------------------------------- @@ -208,3 +288,61 @@ async def test_vertex_realtime_text_in_text_out(): # response.done should have been forwarded done_msgs = [m for m in sent_to_client if '"response.done"' in m] assert done_msgs, "Expected response.done to be sent to client" + + +def test_vertex_warns_when_dropping_guardrail_turn_detection_update(caplog): + """A subsequent session.update carrying the guardrail's + ``create_response: False`` cannot be forwarded as a follow-up setup on + Vertex AI (1007). Surface a warning so operators know the auto-response + suppression is being silently dropped.""" + import logging + + cfg = VertexAIRealtimeConfig( + access_token="tok", project="my-proj", location="us-central1" + ) + + session_update = { + "type": "session.update", + "session": {"turn_detection": {"create_response": False}}, + } + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + result = cfg.transform_realtime_request( + json.dumps(session_update), + "gemini-live-2.5-flash-native-audio", + session_configuration_request=json.dumps({"setup": {"model": "x"}}), + ) + + assert result == [] + assert any( + "Vertex AI Realtime" in record.message + and "create_response=False" in record.message + for record in caplog.records + ) + + +def test_vertex_does_not_warn_when_dropping_non_guardrail_session_update(caplog): + """A subsequent session.update without ``create_response: False`` is a + routine drop and should stay at debug level (no warning).""" + import logging + + cfg = VertexAIRealtimeConfig( + access_token="tok", project="my-proj", location="us-central1" + ) + + session_update = { + "type": "session.update", + "session": {"instructions": "Be concise."}, + } + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + cfg.transform_realtime_request( + json.dumps(session_update), + "gemini-live-2.5-flash-native-audio", + session_configuration_request=json.dumps({"setup": {"model": "x"}}), + ) + + assert not any( + "Vertex AI Realtime" in record.message and "session.update" in record.message + for record in caplog.records + )