diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 4b1dc1a66d1..bd6406c6241 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -347,6 +347,7 @@ class RealTimeStreaming: if self._content_sent_after_setup: verbose_logger.debug("Dropping follow-up setup after content was already sent to backend") continue + msg = self._maybe_inject_guardrail_auto_response_disable(msg) await self.backend_ws.send(msg) # type: ignore[union-attr, attr-defined] self._cache_session_configuration_request(msg) sent = True @@ -623,6 +624,36 @@ class RealTimeStreaming: if sent: self._guardrail_turn_detection_update_sent = True + def _maybe_inject_guardrail_auto_response_disable(self, setup_message: str) -> str: + """Fold the transcription-guardrail auto-response disable into the setup. + + Gemini/Vertex Live reject a second ``setup`` (1007), so the guardrail's + ``automaticActivityDetection.disabled=true`` cannot be delivered as a + follow-up session.update; it must live in the one-and-only setup, or a + ``realtime_input_transcription`` guardrail is bypassed (the model + auto-responds before the proxy can gate the turn). Applies only to the + bidi ``setup`` shape; OpenAI sessions accept follow-up updates and so are + left untouched (handled by ``_maybe_send_guardrail_turn_detection_update``). + """ + if self._guardrail_turn_detection_update_sent: + return setup_message + if not self._has_audio_transcription_guardrails(): + return setup_message + try: + obj = json.loads(setup_message) + except (json.JSONDecodeError, TypeError): + return setup_message + setup = obj.get("setup") if isinstance(obj, dict) else None + if not isinstance(setup, dict): + return setup_message + automatic = setup.setdefault("realtimeInputConfig", {}).setdefault("automaticActivityDetection", {}) + automatic["disabled"] = True + self._guardrail_turn_detection_update_sent = True + verbose_logger.debug( + "Realtime: folded automaticActivityDetection.disabled=true into setup for transcription-guardrail gating" + ) + return json.dumps(obj) + def _has_realtime_guardrails_for_event_hooks( self, event_hooks: List[Any], diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index d8f4856b4ad..289dce7b366 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -5688,6 +5688,58 @@ class BaseLLMHTTPHandler: new_query = parsed.query + ("&" if parsed.query else "") + urlencode(extras) return urlunparse(parsed._replace(query=new_query)) + @staticmethod + async def _open_realtime_backend_ws( + websockets_module: Any, + url: str, + headers: dict, + ssl_context: Any, + *, + open_timeout: float = 8.0, + max_attempts: int = 3, + ) -> Any: + """Open the backend realtime websocket, retrying a hung open handshake. + + The upstream Live handshake (e.g. Gemini Live) intermittently hangs on + open; waiting longer never recovers a hung attempt, but a fresh attempt + almost always connects in ~1s. So bound each attempt with ``open_timeout`` + and retry, instead of surfacing one slow handshake to the caller as a + fatal 1011. A bounded attempt that timed out already spaced out the + retry, so no extra backoff is needed. Deterministic rejections (auth / + handshake status) are not retried. + """ + # Handshake-status rejections are deterministic (auth / 4xx): retrying + # cannot help and the caller must see the upstream status, not a generic + # 1011. websockets <15 raises InvalidStatusCode, >=15 raises InvalidStatus. + deterministic_errors = tuple( + exc + for exc in ( + getattr(websockets_module.exceptions, "InvalidStatus", None), + getattr(websockets_module.exceptions, "InvalidStatusCode", None), + ) + if exc is not None + ) + last_exc: Optional[BaseException] = None + for _ in range(max_attempts): + try: + return await websockets_module.connect( + url, + additional_headers=headers, + max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, + ssl=ssl_context, + open_timeout=open_timeout, + ) + except deterministic_errors: + raise + except ( + TimeoutError, + OSError, + websockets_module.exceptions.WebSocketException, + ) as e: + last_exc = e + assert last_exc is not None # loop only exits via return or a captured exc + raise last_exc + async def async_realtime( self, model: str, @@ -5720,20 +5772,8 @@ class BaseLLMHTTPHandler: ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) ssl_context.check_hostname = False ssl_context.verify_mode = ssl.CERT_NONE - async with websockets.connect( # type: ignore - url, - additional_headers=headers, - max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, - ssl=ssl_context, - ) as backend_ws: - # Auto-send session setup if the provider requires it - # (e.g. Gemini/Vertex AI Live needs a `setup` message before any realtime_input) - _session_config: Optional[str] = None - if provider_config.requires_session_configuration(): - _session_config = provider_config.session_configuration_request(model) - if _session_config: - await backend_ws.send(_session_config) - + backend_ws = await self._open_realtime_backend_ws(websockets, url, headers, ssl_context) + async with backend_ws: _request_data: Dict[str, Any] = {} if litellm_metadata: _request_data["litellm_metadata"] = litellm_metadata @@ -5749,8 +5789,22 @@ class BaseLLMHTTPHandler: model if (query_params or {}).get("intent") == "transcription" else None ), ) - if _session_config: - realtime_streaming.session_configuration_request = _session_config + + # Auto-send session setup if the provider requires it (e.g. + # Gemini/Vertex AI Live needs a `setup` before any realtime_input). + # Build the streaming handler first so a transcription guardrail's + # auto-response disable can be folded into this one setup: Gemini + # rejects a second setup, so a follow-up disable would be dropped + # and the guardrail bypassed. + _session_config: Optional[str] = None + if provider_config.requires_session_configuration(): + _session_config = provider_config.session_configuration_request(model) + if _session_config: + _session_config = realtime_streaming._maybe_inject_guardrail_auto_response_disable( + _session_config + ) + await backend_ws.send(_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. diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index 0ff4190a6d3..bc2145fd832 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -406,17 +406,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): 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. + the full setup with all configuration is sent. Gemini Live accepts setup + as the first-and-only client message, so every later session.update is + dropped rather than forwarded as a second setup (which Gemini rejects + with a 1007, tearing the session down). To carry tools/instructions, send + them on the first session.update before any conversation content. """ session_payload = json_message.get("session") or {} # Normalize GA-remapped fields (``output_modalities``, @@ -437,70 +431,27 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): verbose_logger.debug("Gemini Realtime: Sending initial setup with tools to backend") return [json.dumps({"setup": self._finalize_gemini_live_setup(model, 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", {}), + # Gemini Live accepts exactly one ``setup`` message: the first and only + # client message. A second ``setup`` closes the socket with + # ``1007 Request contains an invalid argument``, so a session.update + # after the initial setup must not be forwarded as a follow-up setup. + # Every GA client (pipecat included) sends several session.updates while + # configuring the session; forwarding a second one tears the session down + # before the first turn, which surfaces to callers as silence after the + # first response, reconnect/retry latency churn, and 1011 errors. Drop + # it. The Vertex subclass already drops subsequent setups for this exact + # reason; the constraint is identical on AI Studio. + client_turn_detection = self._extract_turn_detection(session_payload) + if isinstance(client_turn_detection, dict) and client_turn_detection.get("create_response") is False: + verbose_logger.warning( + "Gemini Realtime: Dropping subsequent session.update " + "(turn_detection.create_response=False) — Gemini Live rejects a " + "second setup message, so audio-transcription guardrails cannot " + "suppress the model's auto-response mid-session." ) - 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, - ) - finalized_follow_up = self._finalize_gemini_live_setup(model, cast(dict[str, Any], follow_up_setup)) - # Skip if the follow-up setup is identical to the one already sent. - # The final session.update from Pipecat's _create_response (after history - # items) matches the pre-history session.update we intentionally sent - # before content; sending a duplicate at that point would risk a 1007. - if finalized_follow_up == original_setup: - verbose_logger.debug("Gemini Realtime: Skipping duplicate follow-up session.update (no changes)") - return [] - verbose_logger.debug("Gemini Realtime: Forwarding session.update as follow-up setup") - return [json.dumps({"setup": finalized_follow_up})] + else: + verbose_logger.debug("Gemini Realtime: Ignoring session.update (setup already sent)") + return [] def _handle_conversation_item(self, json_message: dict) -> List[str]: """ 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 d23fc1a6778..2ad9b919a1f 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -8,9 +8,7 @@ from websockets.exceptions import ConnectionClosed import litellm -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.realtime_streaming import ( @@ -43,9 +41,7 @@ def test_realtime_streaming_store_message(): streaming = RealTimeStreaming(websocket, backend_ws, logging_obj) # Test 1: Session created event (string input) - session_created_msg = json.dumps( - {"type": "session.created", "session": {"id": "test-session"}} - ) + session_created_msg = json.dumps({"type": "session.created", "session": {"id": "test-session"}}) streaming.store_message(session_created_msg) assert len(streaming.messages) == 1 assert "session" in streaming.messages[0] @@ -70,9 +66,7 @@ def test_realtime_streaming_store_message(): streaming.store_message(invalid_msg) # Test 4: Message type not in logged events - streaming.logged_real_time_event_types = [ - "session.created" - ] # Only log session.created + streaming.logged_real_time_event_types = ["session.created"] # Only log session.created other_msg = json.dumps( { "type": "response.done", @@ -85,9 +79,7 @@ def test_realtime_streaming_store_message(): def test_remap_beta_session_to_ga_normalizes_modalities_and_audio(): - out = RealTimeStreaming._remap_beta_session_to_ga( - {"modalities": ["audio", "text"], "voice": "alloy"} - ) + out = RealTimeStreaming._remap_beta_session_to_ga({"modalities": ["audio", "text"], "voice": "alloy"}) assert out["type"] == "realtime" assert out["output_modalities"] == ["audio"] assert out["audio"]["output"]["voice"] == "alloy" @@ -126,13 +118,11 @@ def test_make_disable_auto_response_message_produces_ga_shape(): assert msg["type"] == "session.update" session = msg["session"] - assert ( - session.get("type") == "realtime" - ), "GA session.update must include session.type='realtime'" + assert session.get("type") == "realtime", "GA session.update must include session.type='realtime'" # turn_detection must NOT be at the flat beta location - assert ( - "turn_detection" not in session - ), "turn_detection must not be at the top-level session (beta shape); use audio.input" + assert "turn_detection" not in session, ( + "turn_detection must not be at the top-level session (beta shape); use audio.input" + ) # turn_detection must be nested under audio.input td = session["audio"]["input"]["turn_detection"] assert td["type"] == "server_vad" @@ -151,9 +141,7 @@ def test_make_disable_auto_response_message_produces_beta_shape_for_beta_clients assert msg["type"] == "session.update" session = msg["session"] - assert session == { - "turn_detection": {"type": "server_vad", "create_response": False} - } + assert session == {"turn_detection": {"type": "server_vad", "create_response": False}} @pytest.mark.asyncio @@ -421,9 +409,7 @@ async def test_backend_to_client_drops_ping_events(): backend_ws = MagicMock() backend_ws.recv = AsyncMock( side_effect=[ - json.dumps( - {"type": "ping", "event_id": "evt_ping", "timestamp": 1782214899793} - ).encode(), + json.dumps({"type": "ping", "event_id": "evt_ping", "timestamp": 1782214899793}).encode(), json.dumps({"type": "session.created", "session": {}}).encode(), ConnectionClosed(None, None), ] @@ -604,9 +590,7 @@ async def test_client_ack_messages_keeps_beta_session_shape_for_beta_backend(): backend_ws.send = AsyncMock() logging_obj = MagicMock() logging_obj.pre_call = MagicMock() - streaming = RealTimeStreaming( - client_ws, backend_ws, logging_obj, backend_uses_beta_protocol=True - ) + streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj, backend_uses_beta_protocol=True) await streaming.client_ack_messages() @@ -629,10 +613,7 @@ def test_translate_event_to_beta_renames_delta_types(): def test_translate_event_to_beta_drops_conversation_item_done(): - assert ( - RealTimeStreaming._translate_event_to_beta({"type": "conversation.item.done"}) - is None - ) + assert RealTimeStreaming._translate_event_to_beta({"type": "conversation.item.done"}) is None @pytest.mark.asyncio @@ -878,9 +859,7 @@ async def test_transcription_session_captures_usage_and_skips_response_create(): "type": "session.created", "session": { "type": "transcription", - "audio": { - "input": {"transcription": {"model": "gpt-realtime-whisper"}} - }, + "audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}}, }, } ).encode() @@ -894,9 +873,7 @@ async def test_transcription_session_captures_usage_and_skips_response_create(): ).encode() backend_ws = MagicMock() - backend_ws.recv = AsyncMock( - side_effect=[session_created, completed, ConnectionClosed(None, None)] - ) + backend_ws.recv = AsyncMock(side_effect=[session_created, completed, ConnectionClosed(None, None)]) backend_ws.send = AsyncMock() logging_obj = MagicMock() @@ -910,9 +887,7 @@ async def test_transcription_session_captures_usage_and_skips_response_create(): assert streaming._is_transcription_session is True captured = [ - m - for m in streaming.messages - if m.get("type") == "conversation.item.input_audio_transcription.completed" + m for m in streaming.messages if m.get("type") == "conversation.item.input_audio_transcription.completed" ] assert len(captured) == 1, "completed usage event must be captured for cost" assert captured[0]["usage"]["seconds"] == 12.0 @@ -921,12 +896,10 @@ async def test_transcription_session_captures_usage_and_skips_response_create(): client_ws.send_text.assert_any_call(completed.decode()) # No response.create — transcription sessions have no assistant turn. - sent_to_backend = [ - json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args - ] - assert all( - e.get("type") != "response.create" for e in sent_to_backend - ), f"transcription session must not trigger response.create, got: {sent_to_backend}" + sent_to_backend = [json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args] + assert all(e.get("type") != "response.create" for e in sent_to_backend), ( + f"transcription session must not trigger response.create, got: {sent_to_backend}" + ) @pytest.mark.asyncio @@ -958,9 +931,7 @@ async def test_non_transcription_completed_event_still_triggers_response_create( await streaming.backend_to_client_send_messages() assert streaming._is_transcription_session is False - sent_to_backend = [ - json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args - ] + sent_to_backend = [json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args] assert any(e.get("type") == "response.create" for e in sent_to_backend) @@ -1095,15 +1066,11 @@ def test_detect_transcription_session_from_backend_transcription_session_events( """Backend transcription_session.created/updated events flag the session.""" streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) assert streaming._is_transcription_session is False - streaming._detect_transcription_session_from_backend( - {"type": "transcription_session.created"} - ) + streaming._detect_transcription_session_from_backend({"type": "transcription_session.created"}) assert streaming._is_transcription_session is True streaming2 = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) - streaming2._detect_transcription_session_from_backend( - {"type": "transcription_session.updated"} - ) + streaming2._detect_transcription_session_from_backend({"type": "transcription_session.updated"}) assert streaming2._is_transcription_session is True @@ -1134,9 +1101,7 @@ def test_capture_transcription_usage_deduplicates_when_already_stored(): streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) # Add the event type to the default logged list so _should_store_message returns True. - streaming.logged_real_time_event_types = [ - "conversation.item.input_audio_transcription.completed" - ] + streaming.logged_real_time_event_types = ["conversation.item.input_audio_transcription.completed"] event = { "type": "conversation.item.input_audio_transcription.completed", "usage": {"type": "duration", "seconds": 5.0}, @@ -1201,9 +1166,7 @@ async def test_failed_content_send_does_not_block_later_setup(): logging_obj = MagicMock() provider_config = MagicMock() - provider_config.transform_realtime_request = MagicMock( - side_effect=lambda m, *a, **k: [m] - ) + provider_config.transform_realtime_request = MagicMock(side_effect=lambda m, *a, **k: [m]) provider_config.is_setup_message = MagicMock(side_effect=lambda obj: "setup" in obj) provider_config.is_content_message = MagicMock( side_effect=lambda obj: obj.get("type") == "conversation.item.create" @@ -1347,9 +1310,7 @@ async def test_log_messages_includes_tools_in_model_call_details(): logging_obj.success_handler = MagicMock() streaming = RealTimeStreaming(websocket, backend_ws, logging_obj) - streaming.session_tools = [ - {"type": "function", "name": "get_weather", "description": "Get weather"} - ] + streaming.session_tools = [{"type": "function", "name": "get_weather", "description": "Get weather"}] streaming.tool_calls = [ { "id": "call_1", @@ -1378,14 +1339,10 @@ async def test_realtime_guardrail_blocks_prompt_injection(): # Simple guardrail that blocks anything with "system update" class PromptInjectionGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): for text in inputs.get("texts", []): if "system update" in text.lower(): - raise ValueError( - "⚠️ Prompt injection detected. Request blocked by guardrail." - ) + raise ValueError("⚠️ Prompt injection detected. Request blocked by guardrail.") return inputs guardrail = PromptInjectionGuardrail( @@ -1428,39 +1385,25 @@ async def test_realtime_guardrail_blocks_prompt_injection(): # violation message. There should be exactly ONE response.create (the # guardrail-triggered one), preceded by a response.cancel and a # conversation.item.create carrying the violation text. - sent_to_backend = [ - json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args - ] - response_cancels = [ - e for e in sent_to_backend if e.get("type") == "response.cancel" - ] - assert ( - len(response_cancels) == 1 - ), f"Guardrail should send response.cancel, got: {response_cancels}" - guardrail_items = [ - e for e in sent_to_backend if e.get("type") == "conversation.item.create" - ] - assert ( - len(guardrail_items) == 1 - ), f"Guardrail should inject a conversation.item.create with violation message, got: {guardrail_items}" - response_creates = [ - e for e in sent_to_backend if e.get("type") == "response.create" - ] - assert ( - len(response_creates) == 1 - ), f"Guardrail should send exactly one response.create to voice the violation, got: {response_creates}" + sent_to_backend = [json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args] + response_cancels = [e for e in sent_to_backend if e.get("type") == "response.cancel"] + assert len(response_cancels) == 1, f"Guardrail should send response.cancel, got: {response_cancels}" + guardrail_items = [e for e in sent_to_backend if e.get("type") == "conversation.item.create"] + assert len(guardrail_items) == 1, ( + f"Guardrail should inject a conversation.item.create with violation message, got: {guardrail_items}" + ) + response_creates = [e for e in sent_to_backend if e.get("type") == "response.create"] + assert len(response_creates) == 1, ( + f"Guardrail should send exactly one response.create to voice the violation, got: {response_creates}" + ) # ASSERT 2: error event was sent directly to the client WebSocket - sent_to_client = [ - json.loads(c.args[0]) for c in client_ws.send_text.call_args_list if c.args - ] + sent_to_client = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list if c.args] error_events = [e for e in sent_to_client if e.get("type") == "error"] - assert ( - len(error_events) == 1 - ), f"Expected one error event sent to client, got: {sent_to_client}" - assert ( - error_events[0]["error"]["type"] == "guardrail_violation" - ), f"Expected guardrail_violation error type, got: {error_events[0]}" + assert len(error_events) == 1, f"Expected one error event sent to client, got: {sent_to_client}" + assert error_events[0]["error"]["type"] == "guardrail_violation", ( + f"Expected guardrail_violation error type, got: {error_events[0]}" + ) litellm.callbacks = [] # cleanup @@ -1476,9 +1419,7 @@ async def test_realtime_guardrail_allows_clean_transcript(): from litellm.types.guardrails import GuardrailEventHooks class PromptInjectionGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): for text in inputs.get("texts", []): if "system update" in text.lower(): raise ValueError("⚠️ Prompt injection detected.") @@ -1518,15 +1459,9 @@ async def test_realtime_guardrail_allows_clean_transcript(): await streaming.backend_to_client_send_messages() # ASSERT: response.create WAS sent to backend (clean transcript) - sent_to_backend = [ - json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args - ] - response_creates = [ - e for e in sent_to_backend if e.get("type") == "response.create" - ] - assert ( - len(response_creates) == 1 - ), f"Clean transcript should trigger response.create, got: {sent_to_backend}" + sent_to_backend = [json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args] + response_creates = [e for e in sent_to_backend if e.get("type") == "response.create"] + assert len(response_creates) == 1, f"Clean transcript should trigger response.create, got: {sent_to_backend}" litellm.callbacks = [] # cleanup @@ -1545,9 +1480,7 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error(): from litellm.types.guardrails import GuardrailEventHooks class BlockingGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): texts = inputs.get("texts", []) for text in texts: if "@" in text: @@ -1581,9 +1514,7 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error(): "type": "conversation.item.create", "item": { "role": "user", - "content": [ - {"type": "input_text", "text": "My email is test@example.com"} - ], + "content": [{"type": "input_text", "text": "My email is test@example.com"}], }, } ) @@ -1613,8 +1544,7 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error(): forwarded_items = [ json.loads(m) for m in sent_to_backend - if isinstance(m, str) - and json.loads(m).get("type") == "conversation.item.create" + if isinstance(m, str) and json.loads(m).get("type") == "conversation.item.create" ] # Filter out guardrail-injected items (contain "Say exactly the following message") original_items = [ @@ -1626,9 +1556,7 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error(): if isinstance(c, dict) ) ] - assert ( - len(original_items) == 0 - ), f"Blocked item should not be forwarded to backend, got: {original_items}" + assert len(original_items) == 0, f"Blocked item should not be forwarded to backend, got: {original_items}" litellm.callbacks = [] # cleanup @@ -1647,9 +1575,7 @@ async def test_realtime_function_call_output_guardrail_blocks_and_returns_error( from litellm.types.guardrails import GuardrailEventHooks class BlockingGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): texts = inputs.get("texts", []) for text in texts: if "@" in text: @@ -1715,9 +1641,9 @@ async def test_realtime_function_call_output_guardrail_blocks_and_returns_error( # 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: {forwarded_tool_outputs}" + assert len(forwarded_tool_outputs) == 1, ( + f"Sanitized function_call_output should be forwarded, got: {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"] @@ -1736,9 +1662,7 @@ async def test_realtime_function_call_output_guardrail_allows_clean_output(): from litellm.types.guardrails import GuardrailEventHooks class BlockingGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): return inputs guardrail = BlockingGuardrail( @@ -1788,9 +1712,7 @@ async def test_realtime_function_call_output_guardrail_allows_clean_output(): 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}" + assert len(forwarded) == 1, f"Clean function_call_output should be forwarded, got: {forwarded}" litellm.callbacks = [] # cleanup @@ -1806,9 +1728,7 @@ async def test_realtime_text_input_guardrail_uses_pre_call_mode(): from litellm.types.guardrails import GuardrailEventHooks class DummyGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): return inputs guardrail = DummyGuardrail( @@ -1823,13 +1743,13 @@ async def test_realtime_text_input_guardrail_uses_pre_call_mode(): logging_obj = MagicMock() streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) - assert ( - streaming._has_realtime_guardrails() is True - ), "pre_call guardrail should be recognized as a realtime guardrail" + assert streaming._has_realtime_guardrails() is True, ( + "pre_call guardrail should be recognized as a realtime guardrail" + ) # pre_call-only guardrails gate typed user messages / tool output, not audio VAD. - assert ( - streaming._has_audio_transcription_guardrails() is False - ), "pre_call-only guardrail must not disable server_vad auto-response" + assert streaming._has_audio_transcription_guardrails() is False, ( + "pre_call-only guardrail must not disable server_vad auto-response" + ) litellm.callbacks = [] # cleanup @@ -1847,9 +1767,7 @@ async def test_realtime_session_created_injects_session_update_for_audio_guardra from litellm.types.guardrails import GuardrailEventHooks class AudioGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): return inputs guardrail = AudioGuardrail( @@ -1862,14 +1780,10 @@ async def test_realtime_session_created_injects_session_update_for_audio_guardra client_ws = MagicMock() client_ws.send_text = AsyncMock() - session_created_event = json.dumps( - {"type": "session.created", "session": {"id": "sess_abc"}} - ).encode() + session_created_event = json.dumps({"type": "session.created", "session": {"id": "sess_abc"}}).encode() backend_ws = MagicMock() - backend_ws.recv = AsyncMock( - side_effect=[session_created_event, ConnectionClosed(None, None)] - ) + backend_ws.recv = AsyncMock(side_effect=[session_created_event, ConnectionClosed(None, None)]) backend_ws.send = AsyncMock() logging_obj = MagicMock() @@ -1880,32 +1794,20 @@ async def test_realtime_session_created_injects_session_update_for_audio_guardra await streaming.backend_to_client_send_messages() # session.created must be forwarded to the client - sent_to_client = [ - json.loads(c.args[0]) for c in client_ws.send_text.call_args_list if c.args - ] - session_created_events = [ - e for e in sent_to_client if e.get("type") == "session.created" - ] - assert ( - len(session_created_events) == 1 - ), f"session.created should be forwarded to client, got: {sent_to_client}" + sent_to_client = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list if c.args] + session_created_events = [e for e in sent_to_client if e.get("type") == "session.created"] + assert len(session_created_events) == 1, f"session.created should be forwarded to client, got: {sent_to_client}" # session.update must be sent to the backend AFTER session.created was forwarded - sent_to_backend = [ - json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args - ] + sent_to_backend = [json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args] session_updates = [e for e in sent_to_backend if e.get("type") == "session.update"] - assert ( - len(session_updates) == 1 - ), f"Expected one session.update injected to backend, got: {sent_to_backend}" + assert len(session_updates) == 1, f"Expected one session.update injected to backend, got: {sent_to_backend}" # GA shape: turn_detection must be nested under audio.input, not at top-level session injected_session = session_updates[0]["session"] - assert ( - injected_session["type"] == "realtime" - ), "GA session.update must include session.type='realtime'" - assert ( - injected_session["audio"]["input"]["turn_detection"]["create_response"] is False - ), "GA session.update must nest turn_detection under audio.input" + assert injected_session["type"] == "realtime", "GA session.update must include session.type='realtime'" + assert injected_session["audio"]["input"]["turn_detection"]["create_response"] is False, ( + "GA session.update must nest turn_detection under audio.input" + ) litellm.callbacks = [] # cleanup @@ -1921,9 +1823,7 @@ async def test_realtime_session_created_does_not_inject_session_update_for_pre_c from litellm.types.guardrails import GuardrailEventHooks class PreCallGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): return inputs guardrail = PreCallGuardrail( @@ -1936,14 +1836,10 @@ async def test_realtime_session_created_does_not_inject_session_update_for_pre_c client_ws = MagicMock() client_ws.send_text = AsyncMock() - session_created_event = json.dumps( - {"type": "session.created", "session": {"id": "sess_xyz"}} - ).encode() + session_created_event = json.dumps({"type": "session.created", "session": {"id": "sess_xyz"}}).encode() backend_ws = MagicMock() - backend_ws.recv = AsyncMock( - side_effect=[session_created_event, ConnectionClosed(None, None)] - ) + backend_ws.recv = AsyncMock(side_effect=[session_created_event, ConnectionClosed(None, None)]) backend_ws.send = AsyncMock() logging_obj = MagicMock() @@ -1953,13 +1849,9 @@ async def test_realtime_session_created_does_not_inject_session_update_for_pre_c streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) await streaming.backend_to_client_send_messages() - sent_to_backend = [ - json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args - ] + sent_to_backend = [json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args] session_updates = [e for e in sent_to_backend if e.get("type") == "session.update"] - assert ( - len(session_updates) == 0 - ), f"pre_call-only guardrail must not inject session.update, got: {sent_to_backend}" + assert len(session_updates) == 0, f"pre_call-only guardrail must not inject session.update, got: {sent_to_backend}" litellm.callbacks = [] # cleanup @@ -1972,9 +1864,7 @@ async def test_pre_call_and_post_call_guardrails_do_not_disable_server_vad(): from litellm.types.guardrails import GuardrailEventHooks class ModelArmorStyleGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): return inputs litellm.callbacks = [ @@ -2021,9 +1911,7 @@ async def test_end_session_after_n_fails_closes_connection(): """ class BadWordGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): for text in inputs.get("texts", []): if "blocked" in text.lower(): raise ValueError("Content blocked by guardrail.") @@ -2057,9 +1945,7 @@ async def test_end_session_after_n_fails_closes_connection(): streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) await streaming.backend_to_client_send_messages() - assert ( - backend_ws.close.called - ), "Expected backend_ws.close() to be called after 2 violations" + assert backend_ws.close.called, "Expected backend_ws.close() to be called after 2 violations" assert streaming._violation_count == 2 litellm.callbacks = [] # cleanup @@ -2073,9 +1959,7 @@ async def test_on_violation_end_session_closes_on_first_fail(): """ class TopicGuardrail(CustomGuardrail): - async def apply_guardrail( - self, inputs, request_data, input_type, logging_obj=None - ): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): for text in inputs.get("texts", []): if "stock" in text.lower(): raise ValueError("Topic not allowed: financial advice.") @@ -2108,9 +1992,7 @@ async def test_on_violation_end_session_closes_on_first_fail(): streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) await streaming.backend_to_client_send_messages() - assert ( - backend_ws.close.called - ), "Expected session to close immediately with on_violation=end_session" + assert backend_ws.close.called, "Expected session to close immediately with on_violation=end_session" assert streaming._violation_count == 1 litellm.callbacks = [] # cleanup @@ -2122,9 +2004,7 @@ async def test_provider_path_suppresses_duplicate_session_created_after_syntheti client_ws.send_text = AsyncMock() backend_ws = MagicMock() - backend_ws.recv = AsyncMock( - side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)] - ) + backend_ws.recv = AsyncMock(side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)]) backend_ws.send = AsyncMock() provider_config = MagicMock() @@ -2165,9 +2045,9 @@ async def test_provider_path_suppresses_duplicate_session_created_after_syntheti 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}" + 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 @@ -2176,9 +2056,7 @@ async def test_duplicate_session_created_still_triggers_guardrail_turn_detection client_ws.send_text = AsyncMock() backend_ws = MagicMock() - backend_ws.recv = AsyncMock( - side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)] - ) + backend_ws.recv = AsyncMock(side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)]) backend_ws.send = AsyncMock() provider_config = MagicMock() @@ -2227,9 +2105,7 @@ async def test_duplicate_session_created_still_triggers_guardrail_turn_detection 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 - ) + assert injected_session["audio"]["input"]["turn_detection"]["create_response"] is False @pytest.mark.asyncio @@ -2245,9 +2121,7 @@ async def test_guardrail_update_respects_idempotency_flag(): logging_obj.success_handler = MagicMock() provider_config = MagicMock() - provider_config.transform_realtime_request = MagicMock( - side_effect=lambda msg, model, session_config: [msg] - ) + provider_config.transform_realtime_request = MagicMock(side_effect=lambda msg, model, session_config: [msg]) streaming = RealTimeStreaming( websocket=client_ws, @@ -2324,9 +2198,9 @@ async def test_guardrail_turn_detection_injected_into_first_session_update_defer 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") + 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 @@ -2385,9 +2259,9 @@ async def test_guardrail_turn_detection_injection_tolerates_non_dict_value( 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") + 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 @@ -2398,13 +2272,7 @@ async def test_guardrail_turn_detection_injection_tolerates_non_dict_value( "client_session", [ {"turn_detection": {"type": "server_vad", "create_response": True}}, - { - "audio": { - "input": { - "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( @@ -2457,9 +2325,9 @@ async def test_subsequent_session_update_cannot_reenable_vad_when_guardrails_act 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") + 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 @@ -2493,9 +2361,7 @@ async def test_follow_up_setup_updates_cached_session_configuration_request(): } } ) - provider_config.transform_realtime_request = MagicMock( - return_value=[follow_up_setup] - ) + provider_config.transform_realtime_request = MagicMock(return_value=[follow_up_setup]) streaming = RealTimeStreaming( websocket=client_ws, @@ -2527,9 +2393,7 @@ async def test_deferred_setup_buffers_audio_until_backend_setup_complete(monkeyp client_ws = MagicMock() audio_msg = json.dumps({"type": "input_audio_buffer.append", "audio": "AA=="}) - client_ws.receive_text = AsyncMock( - side_effect=[audio_msg, ConnectionClosed(None, None)] - ) + client_ws.receive_text = AsyncMock(side_effect=[audio_msg, ConnectionClosed(None, None)]) backend_ws = MagicMock() backend_ws.send = AsyncMock() logging_obj = MagicMock() @@ -2562,12 +2426,8 @@ async def test_deferred_setup_sends_session_update_before_buffered_audio(monkeyp client_ws = MagicMock() audio_msg = json.dumps({"type": "input_audio_buffer.append", "audio": "AA=="}) - session_update = json.dumps( - {"type": "session.update", "session": {"modalities": ["audio"]}} - ) - client_ws.receive_text = AsyncMock( - side_effect=[audio_msg, session_update, ConnectionClosed(None, None)] - ) + session_update = json.dumps({"type": "session.update", "session": {"modalities": ["audio"]}}) + client_ws.receive_text = AsyncMock(side_effect=[audio_msg, session_update, ConnectionClosed(None, None)]) backend_ws = MagicMock() backend_ws.send = AsyncMock() logging_obj = MagicMock() @@ -2596,12 +2456,8 @@ async def test_deferred_setup_flush_buffers_audio_received_during_flush(): client_ws = MagicMock() client_ws.send_text = AsyncMock() - new_audio_msg = json.dumps( - {"type": "input_audio_buffer.append", "audio": "new-audio"} - ) - client_ws.receive_text = AsyncMock( - side_effect=[new_audio_msg, ConnectionClosed(None, None)] - ) + new_audio_msg = json.dumps({"type": "input_audio_buffer.append", "audio": "new-audio"}) + client_ws.receive_text = AsyncMock(side_effect=[new_audio_msg, ConnectionClosed(None, None)]) backend_ws = MagicMock() logging_obj = MagicMock() @@ -2631,9 +2487,7 @@ async def test_deferred_setup_flush_buffers_audio_received_during_flush(): provider_config=provider_config, model="gemini-live-2.5-flash-native-audio", ) - old_audio_msg = json.dumps( - {"type": "input_audio_buffer.append", "audio": "old-audio"} - ) + old_audio_msg = json.dumps({"type": "input_audio_buffer.append", "audio": "old-audio"}) streaming._pending_messages_until_setup = [old_audio_msg] streaming._pending_messages_byte_total = len(old_audio_msg.encode("utf-8")) @@ -2649,9 +2503,7 @@ async def test_deferred_setup_flush_buffers_audio_received_during_flush(): return True streaming._send_to_backend = send_to_backend # type: ignore[method-assign] - setup_task = asyncio.create_task( - streaming._handle_provider_config_message(json.dumps({"setupComplete": {}})) - ) + setup_task = asyncio.create_task(streaming._handle_provider_config_message(json.dumps({"setupComplete": {}}))) await asyncio.wait_for(first_flush_started.wait(), timeout=1) await streaming.client_ack_messages() @@ -2677,9 +2529,7 @@ async def test_deferred_setup_flush_retains_unsent_messages_after_send_failure() json.dumps({"type": "input_audio_buffer.commit"}), ] streaming._pending_messages_until_setup = list(buffered_messages) - streaming._pending_messages_byte_total = sum( - len(message.encode("utf-8")) for message in buffered_messages - ) + streaming._pending_messages_byte_total = sum(len(message.encode("utf-8")) for message in buffered_messages) streaming._send_to_backend = AsyncMock( # type: ignore[method-assign] side_effect=Exception("transient") ) @@ -2687,9 +2537,7 @@ async def test_deferred_setup_flush_retains_unsent_messages_after_send_failure() await streaming._flush_pending_messages_until_setup() assert streaming._pending_messages_until_setup == buffered_messages - assert streaming._pending_messages_byte_total == sum( - len(message.encode("utf-8")) for message in buffered_messages - ) + assert streaming._pending_messages_byte_total == sum(len(message.encode("utf-8")) for message in buffered_messages) streaming._send_to_backend = AsyncMock(return_value=True) # type: ignore[method-assign] @@ -2729,9 +2577,7 @@ async def test_deferred_setup_flushes_audio_on_backend_session_created(monkeypat provider_config=config, model="gemini-live-2.5-flash-native-audio", ) - streaming._pending_messages_until_setup.append( - json.dumps({"type": "input_audio_buffer.append", "audio": "AA=="}) - ) + streaming._pending_messages_until_setup.append(json.dumps({"type": "input_audio_buffer.append", "audio": "AA=="})) await streaming.backend_to_client_send_messages() @@ -2753,9 +2599,7 @@ async def test_deferred_setup_caps_non_audio_buffered_messages(monkeypatch): client_ws = MagicMock() client_ws.receive_text = AsyncMock( - side_effect=[audio_msg] - + [flood_msg] * (cap + 50) - + [ConnectionClosed(None, None)] + side_effect=[audio_msg] + [flood_msg] * (cap + 50) + [ConnectionClosed(None, None)] ) backend_ws = MagicMock() backend_ws.send = AsyncMock() @@ -2774,9 +2618,7 @@ async def test_deferred_setup_caps_non_audio_buffered_messages(monkeypatch): backend_ws.send.assert_not_called() assert len(streaming._pending_messages_until_setup) == cap - assert ( - streaming._pending_messages_byte_total <= RealTimeStreaming._MAX_BUFFERED_BYTES - ) + assert streaming._pending_messages_byte_total <= RealTimeStreaming._MAX_BUFFERED_BYTES @pytest.mark.asyncio @@ -2786,14 +2628,10 @@ async def test_deferred_setup_caps_non_audio_buffered_bytes(monkeypatch): from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig audio_msg = json.dumps({"type": "input_audio_buffer.append", "audio": "AA=="}) - big_non_audio = json.dumps( - {"type": "foo", "data": "x" * (RealTimeStreaming._MAX_BUFFERED_BYTES + 1)} - ) + big_non_audio = json.dumps({"type": "foo", "data": "x" * (RealTimeStreaming._MAX_BUFFERED_BYTES + 1)}) client_ws = MagicMock() - client_ws.receive_text = AsyncMock( - side_effect=[audio_msg, big_non_audio, ConnectionClosed(None, None)] - ) + client_ws.receive_text = AsyncMock(side_effect=[audio_msg, big_non_audio, ConnectionClosed(None, None)]) backend_ws = MagicMock() backend_ws.send = AsyncMock() logging_obj = MagicMock() @@ -2809,9 +2647,7 @@ async def test_deferred_setup_caps_non_audio_buffered_bytes(monkeypatch): await streaming.client_ack_messages() assert streaming._pending_messages_until_setup == [audio_msg] - assert ( - streaming._pending_messages_byte_total <= RealTimeStreaming._MAX_BUFFERED_BYTES - ) + assert streaming._pending_messages_byte_total <= RealTimeStreaming._MAX_BUFFERED_BYTES def _beta_client_ws(): @@ -2889,13 +2725,9 @@ def test_translate_event_to_beta_remaps_response_done_output_content_types(): @pytest.mark.asyncio async def test_beta_client_receives_translated_audio_delta(): client_ws = _beta_client_ws() - frame = json.dumps( - {"type": "response.output_audio.delta", "delta": "QUJD", "event_id": "e1"} - ) + frame = json.dumps({"type": "response.output_audio.delta", "delta": "QUJD", "event_id": "e1"}) backend_ws = MagicMock() - backend_ws.recv = AsyncMock( - side_effect=[frame.encode(), ConnectionClosed(None, None)] - ) + backend_ws.recv = AsyncMock(side_effect=[frame.encode(), ConnectionClosed(None, None)]) logging_obj = MagicMock() logging_obj.async_success_handler = AsyncMock() logging_obj.success_handler = MagicMock() @@ -2912,13 +2744,9 @@ async def test_beta_client_receives_translated_audio_delta(): @pytest.mark.asyncio async def test_ga_client_receives_raw_passthrough(): client_ws = _ga_client_ws() - frame = json.dumps( - {"type": "response.output_audio.delta", "delta": "QUJD", "event_id": "e1"} - ) + frame = json.dumps({"type": "response.output_audio.delta", "delta": "QUJD", "event_id": "e1"}) backend_ws = MagicMock() - backend_ws.recv = AsyncMock( - side_effect=[frame.encode(), ConnectionClosed(None, None)] - ) + backend_ws.recv = AsyncMock(side_effect=[frame.encode(), ConnectionClosed(None, None)]) logging_obj = MagicMock() logging_obj.async_success_handler = AsyncMock() logging_obj.success_handler = MagicMock() @@ -2938,9 +2766,7 @@ async def test_beta_client_non_translated_event_forwarded_raw(): client_ws = _beta_client_ws() frame = json.dumps({"type": "error", "error": {"message": "boom"}}) backend_ws = MagicMock() - backend_ws.recv = AsyncMock( - side_effect=[frame.encode(), ConnectionClosed(None, None)] - ) + backend_ws.recv = AsyncMock(side_effect=[frame.encode(), ConnectionClosed(None, None)]) logging_obj = MagicMock() logging_obj.async_success_handler = AsyncMock() logging_obj.success_handler = MagicMock() @@ -2957,9 +2783,7 @@ async def test_beta_client_drops_conversation_item_done(): client_ws = _beta_client_ws() frame = json.dumps({"type": "conversation.item.done", "item": {"id": "i1"}}) backend_ws = MagicMock() - backend_ws.recv = AsyncMock( - side_effect=[frame.encode(), ConnectionClosed(None, None)] - ) + backend_ws.recv = AsyncMock(side_effect=[frame.encode(), ConnectionClosed(None, None)]) logging_obj = MagicMock() logging_obj.async_success_handler = AsyncMock() logging_obj.success_handler = MagicMock() @@ -2974,9 +2798,7 @@ def test_store_message_skips_pydantic_for_unlogged_audio_delta(): """Audio deltas are not in DefaultLoggedRealTimeEventTypes; store_message must skip the Pydantic build entirely (no append, no validation).""" streaming = _streaming_with(_ga_client_ws()) - with patch( - "litellm.litellm_core_utils.realtime_streaming.OpenAIRealtimeStreamResponseBaseObject" - ) as base_obj: + with patch("litellm.litellm_core_utils.realtime_streaming.OpenAIRealtimeStreamResponseBaseObject") as base_obj: streaming.store_message({"type": "response.output_audio.delta", "delta": "x"}) base_obj.assert_not_called() assert streaming.messages == [] @@ -2985,13 +2807,9 @@ def test_store_message_skips_pydantic_for_unlogged_audio_delta(): @pytest.mark.asyncio async def test_audio_delta_frame_parsed_at_most_once(): client_ws = _beta_client_ws() - frame = json.dumps( - {"type": "response.output_audio.delta", "delta": "QUJD", "event_id": "e1"} - ) + frame = json.dumps({"type": "response.output_audio.delta", "delta": "QUJD", "event_id": "e1"}) backend_ws = MagicMock() - backend_ws.recv = AsyncMock( - side_effect=[frame.encode(), ConnectionClosed(None, None)] - ) + backend_ws.recv = AsyncMock(side_effect=[frame.encode(), ConnectionClosed(None, None)]) logging_obj = MagicMock() logging_obj.async_success_handler = AsyncMock() logging_obj.success_handler = MagicMock() @@ -3019,9 +2837,7 @@ def test_collapse_buffered_audio_messages_applies_clear_semantics(): new = json.dumps({"type": "input_audio_buffer.append", "audio": "new"}) commit = json.dumps({"type": "input_audio_buffer.commit"}) - collapsed = RealTimeStreaming._collapse_buffered_audio_messages( - [old, cleared, new, commit] - ) + collapsed = RealTimeStreaming._collapse_buffered_audio_messages([old, cleared, new, commit]) assert collapsed == [new, commit] @@ -3064,3 +2880,68 @@ async def test_deferred_setup_clear_drops_appends_when_buffered(): streaming._buffer_pending_message_until_setup(new_audio) assert streaming._pending_messages_until_setup == [new_audio] + + +def _transcription_guardrail(): + """A minimal real CustomGuardrail registered for the realtime transcript hook.""" + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.types.guardrails import GuardrailEventHooks + + class _TranscriptionGuardrail(CustomGuardrail): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + return inputs + + return _TranscriptionGuardrail( + guardrail_name="test_transcription_guard", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + ) + + +def test_setup_folds_in_auto_response_disable_when_transcription_guardrail_active(): + """Gemini rejects a second setup, so a transcription guardrail's auto-response + disable must be folded into the one-and-only setup; otherwise the model + auto-responds and the guardrail is bypassed.""" + import litellm + + litellm.callbacks = [_transcription_guardrail()] + try: + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + setup = json.dumps( + { + "setup": { + "model": "models/gemini-3.1-flash-live-preview", + "generationConfig": {"responseModalities": ["AUDIO"]}, + "inputAudioTranscription": {}, + } + } + ) + out = json.loads(streaming._maybe_inject_guardrail_auto_response_disable(setup)) + aad = out["setup"]["realtimeInputConfig"]["automaticActivityDetection"] + assert aad["disabled"] is True + finally: + litellm.callbacks = [] + + +def test_setup_unchanged_without_transcription_guardrail(): + import litellm + + litellm.callbacks = [] + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + setup = json.dumps({"setup": {"model": "x", "generationConfig": {"responseModalities": ["AUDIO"]}}}) + out = streaming._maybe_inject_guardrail_auto_response_disable(setup) + assert json.loads(out) == json.loads(setup) + + +def test_non_bidi_setup_left_untouched_for_followup_capable_providers(): + """OpenAI realtime accepts a follow-up session.update, so a non-bidi message + (no top-level 'setup' key) must be left untouched even with a guardrail on.""" + import litellm + + litellm.callbacks = [_transcription_guardrail()] + try: + streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock()) + msg = json.dumps({"type": "session.update", "session": {"instructions": "hi"}}) + assert streaming._maybe_inject_guardrail_auto_response_disable(msg) == msg + finally: + litellm.callbacks = [] diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 64ae30daa70..d57d115fa4e 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -6,9 +6,7 @@ from unittest.mock import AsyncMock, Mock, patch import httpx import pytest -sys.path.insert( - 0, os.path.abspath("../../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../../..")) # Adds the parent directory to the system path import litellm from litellm.integrations.code_interpreter_interception.handler import ( CodeInterpreterInterceptionLogger, @@ -39,9 +37,7 @@ def test_prepare_fake_stream_request(): } fake_stream = True - result_stream, result_data = handler._prepare_fake_stream_request( - stream=stream, data=data, fake_stream=fake_stream - ) + result_stream, result_data = handler._prepare_fake_stream_request(stream=stream, data=data, fake_stream=fake_stream) # Verify that stream is set to False assert result_stream is False @@ -60,9 +56,7 @@ def test_prepare_fake_stream_request(): } fake_stream = False - result_stream, result_data = handler._prepare_fake_stream_request( - stream=stream, data=data, fake_stream=fake_stream - ) + result_stream, result_data = handler._prepare_fake_stream_request(stream=stream, data=data, fake_stream=fake_stream) # Verify that stream remains True assert result_stream is True @@ -77,9 +71,7 @@ def test_prepare_fake_stream_request(): data = {"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}]} fake_stream = True - result_stream, result_data = handler._prepare_fake_stream_request( - stream=stream, data=data, fake_stream=fake_stream - ) + result_stream, result_data = handler._prepare_fake_stream_request(stream=stream, data=data, fake_stream=fake_stream) # Verify that stream is set to False assert result_stream is False @@ -165,9 +157,7 @@ def test_response_api_handler_runs_agentic_hooks_in_sync_path(monkeypatch): assert response is final_response hook_mock.assert_awaited_once() assert hook_mock.call_args.kwargs["api_surface"] == "responses" - assert hook_mock.call_args.kwargs["messages"] == [ - {"role": "user", "content": "hi"} - ] + assert hook_mock.call_args.kwargs["messages"] == [{"role": "user", "content": "hi"}] def test_response_api_handler_runs_responses_pre_call_hook_before_transform(): @@ -226,9 +216,7 @@ def test_response_api_handler_runs_responses_pre_call_hook_before_transform(): tools = transform_kwargs["response_api_optional_request_params"]["tools"] assert not any(tool.get("type") == "code_interpreter" for tool in tools) assert any( - tool.get("type") == "function" - and tool.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME - for tool in tools + tool.get("type") == "function" and tool.get("name") == LITELLM_CODE_EXECUTION_TOOL_NAME for tool in tools ) hook_litellm_params = transform_kwargs["litellm_params"] assert hook_litellm_params.get(_ACTIVE_KEY) is True @@ -363,9 +351,7 @@ def test_fingerprint_agentic_tools_is_deterministic(): tools_a = {"tool_calls": [{"id": "1", "input": {"q": "abc"}, "name": "web_search"}]} tools_b = {"tool_calls": [{"name": "web_search", "input": {"q": "abc"}, "id": "1"}]} - assert handler._fingerprint_agentic_tools( - tools_a - ) == handler._fingerprint_agentic_tools(tools_b) + assert handler._fingerprint_agentic_tools(tools_a) == handler._fingerprint_agentic_tools(tools_b) @pytest.mark.asyncio @@ -587,9 +573,7 @@ async def test_async_anthropic_messages_handler_forwards_router_model_info(): mock_logging_obj.update_from_kwargs.assert_called_once() call_kwargs = mock_logging_obj.update_from_kwargs.call_args litellm_params_arg = ( - call_kwargs.kwargs.get( - "litellm_params", call_kwargs[1].get("litellm_params", {}) - ) + call_kwargs.kwargs.get("litellm_params", call_kwargs[1].get("litellm_params", {})) if call_kwargs.kwargs else call_kwargs[1].get("litellm_params", {}) ) @@ -672,10 +656,7 @@ def test_google_genai_streaming_hidden_params_model_info_and_router_fallback(): response_headers=httpx.Headers({"x-ratelimit-remaining": "10"}), ) assert from_model_info["model_id"] == "info-id" - assert ( - from_model_info["api_base"] - == "https://generativelanguage.googleapis.com/v1beta" - ) + assert from_model_info["api_base"] == "https://generativelanguage.googleapis.com/v1beta" assert isinstance(from_model_info["additional_headers"], dict) from_router = _google_genai_streaming_hidden_params( @@ -729,9 +710,7 @@ def test_async_delete_responses_omits_body_for_azure(): assert "json" not in captured assert "data" not in captured - assert captured["url"].endswith( - "/openai/responses/resp_xyz?api-version=2025-03-01-preview" - ) + assert captured["url"].endswith("/openai/responses/resp_xyz?api-version=2025-03-01-preview") def test_sync_delete_responses_omits_body_for_azure(): @@ -749,9 +728,7 @@ def test_sync_delete_responses_omits_body_for_azure(): assert "json" not in captured assert "data" not in captured - assert captured["url"].endswith( - "/openai/responses/resp_xyz?api-version=2025-03-01-preview" - ) + assert captured["url"].endswith("/openai/responses/resp_xyz?api-version=2025-03-01-preview") def _content_type(headers: dict) -> str: @@ -889,9 +866,7 @@ async def test_anthropic_post_retry_reserializes_mutated_body(): prebuilt = _json.dumps(request_body) err_resp = Mock() - http_error = httpx.HTTPStatusError( - "bad", request=Mock(), response=Mock(status_code=400) - ) + http_error = httpx.HTTPStatusError("bad", request=Mock(), response=Mock(status_code=400)) err_resp.raise_for_status = Mock(side_effect=http_error) ok_resp = Mock() ok_resp.raise_for_status = Mock(return_value=None) @@ -903,9 +878,7 @@ async def test_anthropic_post_retry_reserializes_mutated_body(): provider_config = Mock() provider_config.max_retry_on_anthropic_messages_http_error = 2 - provider_config.should_retry_anthropic_messages_on_http_error = Mock( - return_value=True - ) + provider_config.should_retry_anthropic_messages_on_http_error = Mock(return_value=True) provider_config.transform_anthropic_messages_request_on_http_error = _mutate # Re-sign returns no signed body (native anthropic path) -> must re-dump. provider_config.sign_request = Mock(return_value=({}, None)) @@ -968,9 +941,7 @@ def _make_responses_handler_call(signed_body): provider_config = MagicMock() provider_config.validate_environment.return_value = {} - provider_config.get_complete_url.return_value = ( - "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" - ) + provider_config.get_complete_url.return_value = "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" provider_config.transform_responses_api_request.return_value = {"input": "hi"} provider_config.should_fake_stream.return_value = False provider_config.sign_request.return_value = ({"X-Signed": "1"}, signed_body) @@ -1026,9 +997,7 @@ def test_responses_handler_signs_after_fake_stream_prep_strips_stream(): provider_config = MagicMock() provider_config.validate_environment.return_value = {} - provider_config.get_complete_url.return_value = ( - "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" - ) + provider_config.get_complete_url.return_value = "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" provider_config.transform_responses_api_request.return_value = { "input": "hi", "stream": True, @@ -1091,9 +1060,7 @@ def _make_compact_handler_call(signed_body, is_async): compact_url = "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses/compact" provider_config = MagicMock() provider_config.validate_environment.return_value = {} - provider_config.get_complete_url.return_value = ( - "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" - ) + provider_config.get_complete_url.return_value = "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" provider_config.transform_compact_response_api_request.return_value = ( compact_url, {"model": "openai.gpt-5.5", "input": "hi"}, @@ -1127,9 +1094,7 @@ def _make_compact_handler_call(signed_body, is_async): def test_compact_handler_sends_json_when_not_signed(): """No-op provider on compact (signed_body is None) -> posts json=data, no data= bytes.""" - provider_config, kwargs = _make_compact_handler_call( - signed_body=None, is_async=False - ) + provider_config, kwargs = _make_compact_handler_call(signed_body=None, is_async=False) provider_config.sign_request.assert_called_once() assert kwargs.get("json") == {"model": "openai.gpt-5.5", "input": "hi"} assert "data" not in kwargs @@ -1148,9 +1113,7 @@ def test_compact_handler_sends_signed_bytes_when_signed(): assert "json" not in kwargs assert kwargs["headers"] == {"X-Signed": "1"} # signing must use the compact endpoint as api_base, not the create URL - assert provider_config.sign_request.call_args.kwargs["api_base"].endswith( - "/openai/v1/responses/compact" - ) + assert provider_config.sign_request.call_args.kwargs["api_base"].endswith("/openai/v1/responses/compact") def test_async_compact_handler_sends_signed_bytes_when_signed(): @@ -1165,8 +1128,93 @@ def test_async_compact_handler_sends_signed_bytes_when_signed(): def test_async_compact_handler_sends_json_when_not_signed(): """Async no-op provider on compact -> posts json=data, no data= bytes.""" - _provider_config, kwargs = _make_compact_handler_call( - signed_body=None, is_async=True - ) + _provider_config, kwargs = _make_compact_handler_call(signed_body=None, is_async=True) assert kwargs.get("json") == {"model": "openai.gpt-5.5", "input": "hi"} assert "data" not in kwargs + + +class _FakeWSExceptions: + class WebSocketException(Exception): + pass + + class InvalidStatusCode(WebSocketException): + def __init__(self) -> None: + super().__init__("HTTP 403") + + # websockets>=15 raises InvalidStatus (not InvalidStatusCode) for a rejected + # client handshake; both must be treated as deterministic. + class InvalidStatus(WebSocketException): + def __init__(self) -> None: + super().__init__("HTTP 401") + + +class _FakeWebsocketsModule: + """Stand-in for the ``websockets`` module so the realtime backend-open retry + can be exercised without a real network handshake (dependency injection, + no monkeypatching).""" + + def __init__(self, outcomes): + # outcomes: list where each item is either an Exception to raise or a + # sentinel object to return as the "connected" websocket. + self._outcomes = list(outcomes) + self.exceptions = _FakeWSExceptions + self.attempts = 0 + self.open_timeouts: list = [] + + async def connect(self, *args, **kwargs): + self.attempts += 1 + self.open_timeouts.append(kwargs.get("open_timeout")) + outcome = self._outcomes.pop(0) + if isinstance(outcome, Exception): + raise outcome + return outcome + + +@pytest.mark.asyncio +async def test_realtime_backend_open_retries_then_succeeds(): + """A hung/slow open handshake is retried; a later fresh attempt connects. + + Regression for intermittent ``1011 timed out during opening handshake``: + the proxy used to surface a single slow upstream handshake to the caller as + a fatal 1011 with no retry. + """ + sentinel = object() + fake = _FakeWebsocketsModule([TimeoutError("timed out during opening handshake"), sentinel]) + + result = await BaseLLMHTTPHandler._open_realtime_backend_ws( + fake, "wss://backend.example/live", {"Authorization": "Bearer x"}, None + ) + + assert result is sentinel + assert fake.attempts == 2 + # Each attempt must be bounded by a finite open_timeout (not the default/None). + assert all(t is not None and t > 0 for t in fake.open_timeouts) + + +@pytest.mark.asyncio +async def test_realtime_backend_open_raises_after_max_attempts(): + """When every attempt times out, the final error propagates (so the caller + still closes the client socket) rather than looping forever.""" + fake = _FakeWebsocketsModule([TimeoutError("hang")] * 2) + + with pytest.raises(TimeoutError): + await BaseLLMHTTPHandler._open_realtime_backend_ws(fake, "wss://backend.example/live", {}, None, max_attempts=2) + + assert fake.attempts == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "rejection", + [_FakeWSExceptions.InvalidStatusCode, _FakeWSExceptions.InvalidStatus], +) +async def test_realtime_backend_open_does_not_retry_auth_failure(rejection): + """A deterministic handshake-status rejection (auth/4xx) must not be retried; + retrying cannot help and the upstream status must surface, not a 1011. Both + the websockets<15 (InvalidStatusCode) and >=15 (InvalidStatus) shapes apply.""" + fake = _FakeWebsocketsModule([rejection()]) + + with pytest.raises(_FakeWSExceptions.WebSocketException): + await BaseLLMHTTPHandler._open_realtime_backend_ws(fake, "wss://backend.example/live", {}, None) + + assert fake.attempts == 1 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 0f9edb0faf8..02004d5c8a8 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 @@ -6,9 +6,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest -sys.path.insert( - 0, os.path.abspath("../../../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../../../..")) # Adds the parent directory to the system path import litellm from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig @@ -86,10 +84,7 @@ def test_session_created_does_not_overwrite_session_configuration_request(): ) # Must keep original setup payload (with "setup"), not overwrite with session.created event. - assert ( - transformed["session_configuration_request"] - == session_configuration_request_str - ) + 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] @@ -146,19 +141,13 @@ def test_gemini_realtime_transformation_content_delta(): print(transformed_message) ## assert all instances of 'event_id' are unique - event_ids = [ - event["event_id"] for event in transformed_message if "event_id" in event - ] + event_ids = [event["event_id"] for event in transformed_message if "event_id" in event] assert len(event_ids) == len(set(event_ids)) ## assert all instances of 'response_id' are the same - response_ids = [ - event["response_id"] for event in transformed_message if "response_id" in event - ] + response_ids = [event["response_id"] for event in transformed_message if "response_id" in event] assert len(set(response_ids)) == 1 ## assert all instances of 'output_item_id' are the same - output_item_ids = [ - event["item_id"] for event in transformed_message if "item_id" in event - ] + output_item_ids = [event["item_id"] for event in transformed_message if "item_id" in event] assert len(set(output_item_ids)) == 1 @@ -172,9 +161,7 @@ def test_gemini_model_turn_event_mapping(): openai_event = config.map_model_turn_event(model_turn_event) assert openai_event == OpenAIRealtimeEventTypes.RESPONSE_TEXT_DELTA - model_turn_event = { - "parts": [{"inlineData": {"mimeType": "audio/pcm", "data": "..."}}] - } + model_turn_event = {"parts": [{"inlineData": {"mimeType": "audio/pcm", "data": "..."}}]} openai_event = config.map_model_turn_event(model_turn_event) assert openai_event == OpenAIRealtimeEventTypes.RESPONSE_AUDIO_DELTA @@ -205,13 +192,7 @@ def test_gemini_realtime_transformation_audio_delta(): session_configuration_request_str = json.dumps(session_configuration_request) audio_delta_event = { - "serverContent": { - "modelTurn": { - "parts": [ - {"inlineData": {"mimeType": "audio/pcm", "data": "my-audio-data"}} - ] - } - } + "serverContent": {"modelTurn": {"parts": [{"inlineData": {"mimeType": "audio/pcm", "data": "my-audio-data"}}]}} } result = config.transform_realtime_response( @@ -235,10 +216,7 @@ def test_gemini_realtime_transformation_audio_delta(): contains_audio_delta = False for response in responses: - if ( - response["type"] - == OpenAIRealtimeEventTypes.RESPONSE_OUTPUT_AUDIO_DELTA.value - ): + if response["type"] == OpenAIRealtimeEventTypes.RESPONSE_OUTPUT_AUDIO_DELTA.value: contains_audio_delta = True break assert contains_audio_delta, "Expected audio delta event" @@ -257,11 +235,7 @@ def test_gemini_output_audio_transcript_delta_uses_active_response_ids(): event = { "serverContent": { "outputTranscription": {"text": "Hello from Gemini."}, - "modelTurn": { - "parts": [ - {"inlineData": {"mimeType": "audio/pcm", "data": "my-audio-data"}} - ] - }, + "modelTurn": {"parts": [{"inlineData": {"mimeType": "audio/pcm", "data": "my-audio-data"}}]}, } } @@ -281,19 +255,11 @@ def test_gemini_output_audio_transcript_delta_uses_active_response_ids(): ) responses = result["response"] - response_created = next( - response for response in responses if response["type"] == "response.created" - ) + response_created = next(response for response in responses if response["type"] == "response.created") transcript_delta = next( - response - for response in responses - if response["type"] == "response.output_audio_transcript.delta" - ) - audio_delta = next( - response - for response in responses - if response["type"] == "response.output_audio.delta" + response for response in responses if response["type"] == "response.output_audio_transcript.delta" ) + audio_delta = next(response for response in responses if response["type"] == "response.output_audio.delta") assert transcript_delta["response_id"] == response_created["response"]["id"] assert transcript_delta["response_id"] == audio_delta["response_id"] @@ -339,10 +305,7 @@ def test_gemini_realtime_transformation_generation_complete(): contains_audio_done_event = False for response in responses: - if ( - response["type"] - == OpenAIRealtimeEventTypes.RESPONSE_OUTPUT_AUDIO_DONE.value - ): + if response["type"] == OpenAIRealtimeEventTypes.RESPONSE_OUTPUT_AUDIO_DONE.value: contains_audio_done_event = True break assert contains_audio_done_event, "Expected audio done event" @@ -412,9 +375,7 @@ def test_gemini_realtime_tool_call_transformation(): function_call_event = event break - assert ( - function_call_event is not None - ), "Expected function_call_arguments.done event" + 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" @@ -516,6 +477,109 @@ def test_gemini_session_update_defaults_to_audio_modality(): assert setup_payload["generationConfig"]["responseModalities"] == ["AUDIO"] +def test_gemini_subsequent_session_update_is_dropped_not_resent_as_setup(): + """Regression: Gemini Live accepts exactly one ``setup`` message; a second + one closes the socket with ``1007 Request contains an invalid argument``. + + Once the initial setup has been sent (``session_configuration_request`` is + set), a follow-up session.update must be dropped rather than forwarded as + another setup. Forwarding it tore the session down before the first turn, + which surfaced to callers as silence after the first response and 1011s. + """ + config = GeminiRealtimeConfig() + initial_setup = json.dumps( + { + "setup": { + "model": "models/gemini-2.5-flash", + "generationConfig": {"responseModalities": ["AUDIO"]}, + "inputAudioTranscription": {}, + } + } + ) + follow_up = { + "type": "session.update", + "session": {"instructions": "Updated instructions", "temperature": 0.4}, + } + + messages = config.transform_realtime_request( + json.dumps(follow_up), + "gemini-2.5-flash", + session_configuration_request=initial_setup, + ) + + assert messages == [], ( + "a session.update after the initial setup must be dropped, never " + "forwarded as a second setup (Gemini Live rejects it with 1007)" + ) + + +def test_gemini_subsequent_session_update_with_new_tools_is_dropped(): + """Regression: even a follow-up session.update that *differs* from the initial + setup (e.g. registers tools after connect) must be dropped. + + The previous dedup only skipped follow-ups identical to the initial setup; a + changed one was merged and re-sent as a second setup, still hitting the 1007. + Tools must instead ride on the first session.update. + """ + config = GeminiRealtimeConfig() + initial_setup = json.dumps( + { + "setup": { + "model": "models/gemini-2.5-flash", + "generationConfig": {"responseModalities": ["AUDIO"]}, + } + } + ) + follow_up = { + "type": "session.update", + "session": { + "tools": [ + { + "type": "function", + "function": { + "name": "terminate_call", + "description": "End the call.", + "parameters": {"type": "object", "properties": {}}, + }, + } + ] + }, + } + + messages = config.transform_realtime_request( + json.dumps(follow_up), + "gemini-2.5-flash", + session_configuration_request=initial_setup, + ) + + assert messages == [] + + +def test_gemini_subsequent_guardrail_session_update_dropped_with_warning(caplog): + """A dropped follow-up carrying ``turn_detection.create_response=False`` (the + transcription-guardrail signal) is still dropped, but warns so operators know + the guardrail cannot gate the model's auto-response mid-session on Gemini. + """ + import logging + + config = GeminiRealtimeConfig() + initial_setup = json.dumps({"setup": {"model": "models/gemini-2.5-flash"}}) + follow_up = { + "type": "session.update", + "session": {"turn_detection": {"type": "server_vad", "create_response": False}}, + } + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + messages = config.transform_realtime_request( + json.dumps(follow_up), + "gemini-2.5-flash", + session_configuration_request=initial_setup, + ) + + assert messages == [] + assert any("Dropping subsequent session.update" in record.message for record in caplog.records) + + @pytest.mark.parametrize( "model", [ @@ -687,9 +751,7 @@ def test_gemini_realtime_function_call_output_transformation(): "gemini-2.5-flash", session_configuration_request="existing", ) - retry_response = json.loads(retry_messages[0])["toolResponse"]["functionResponses"][ - 0 - ] + retry_response = json.loads(retry_messages[0])["toolResponse"]["functionResponses"][0] assert retry_response["name"] == "get_weather" @@ -703,9 +765,7 @@ def test_gemini_realtime_user_text_transformation(): "item": { "type": "message", "role": "user", - "content": [ - {"type": "input_text", "text": "What's the weather in London?"} - ], + "content": [{"type": "input_text", "text": "What's the weather in London?"}], }, } @@ -784,11 +844,7 @@ def test_gemini_realtime_multi_tool_calls_have_unique_item_ids(): }, ) - responses = [ - ev - for ev in result["response"] - if ev.get("type") == "response.function_call_arguments.done" - ] + 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" @@ -954,13 +1010,7 @@ def test_gemini_tool_call_resets_ids_for_post_tool_model_turn(): 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."}]} - } - } - ), + json.dumps({"serverContent": {"modelTurn": {"parts": [{"text": "The weather is sunny."}]}}}), "gemini-2.5-flash", logging_obj, realtime_response_transform_input={ @@ -977,9 +1027,7 @@ def test_gemini_tool_call_resets_ids_for_post_tool_model_turn(): 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"] - ) + assert post_tool_result["current_response_id"] == post_tool_events[0]["response"]["id"] def test_gemini_empty_tool_call_does_not_crash_websocket(): @@ -1092,9 +1140,7 @@ def test_gemini_tool_call_response_done_includes_usage_from_sibling_metadata(): }, ) - response_done = next( - ev for ev in result["response"] if ev.get("type") == "response.done" - ) + 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 @@ -1138,9 +1184,7 @@ def test_gemini_tool_call_response_done_falls_back_to_empty_usage(): }, ) - response_done = next( - ev for ev in result["response"] if ev.get("type") == "response.done" - ) + 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 @@ -1214,60 +1258,6 @@ def test_gemini_function_call_output_includes_name(): 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_realtime_pipecat_ga_session_voice_and_tools(patch_gemini_audio_cost_map_entries): """Pipecat OpenAIRealtimeSessionProperties: output_modalities, nested tools, and audio.output.voice (e.g. Kore) must map into Gemini setup.""" @@ -1314,9 +1304,7 @@ def test_gemini_realtime_pipecat_ga_session_voice_and_tools(patch_gemini_audio_c # Native-audio Live rejects speechConfig on setup (see _finalize_gemini_live_setup). assert "speechConfig" not in setup.get("generationConfig", {}) assert setup["tools"][0]["function_declarations"][0]["name"] == "terminate_call" - assert ( - setup["realtimeInputConfig"]["automaticActivityDetection"]["disabled"] is False - ) + assert setup["realtimeInputConfig"]["automaticActivityDetection"]["disabled"] is False def test_gemini_realtime_pipecat_semantic_vad_omits_realtime_input_config(): @@ -1409,116 +1397,6 @@ def test_gemini_input_audio_buffer_commit_maps_to_activity_end_when_manual_vad() assert json.loads(messages[0]) == {"realtimeInput": {"activityEnd": True}} -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, @@ -1669,9 +1547,7 @@ def test_gemini_standalone_usage_metadata_is_attributed_to_next_tool_call_respon }, ) - response_done = next( - ev for ev in tool_call_result["response"] if ev.get("type") == "response.done" - ) + 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 @@ -1733,11 +1609,7 @@ def test_gemini_standalone_usage_metadata_is_attributed_to_next_response_done(): }, ) - response_done = next( - ev - for ev in turn_complete_result["response"] - if ev.get("type") == "response.done" - ) + 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 @@ -1793,9 +1665,7 @@ def test_gemini_in_frame_usage_metadata_clears_pending_buffer(): }, ) - response_done = next( - ev for ev in result["response"] if ev.get("type") == "response.done" - ) + 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 @@ -1912,9 +1782,7 @@ def test_gemini_post_tool_bare_turn_complete_followed_by_answer(): ) assert post_tool_answer["response"][0]["type"] == "response.created" transcript_delta = next( - event - for event in post_tool_answer["response"] - if event["type"] == "response.output_audio_transcript.delta" + event for event in post_tool_answer["response"] if event["type"] == "response.output_audio_transcript.delta" ) assert "72" in transcript_delta["delta"] @@ -1932,11 +1800,7 @@ def test_gemini_post_tool_bare_turn_complete_followed_by_answer(): "current_delta_type": post_tool_answer["current_delta_type"], }, ) - response_done = next( - event - for event in final_turn["response"] - if event["type"] == "response.done" - ) + response_done = next(event for event in final_turn["response"] if event["type"] == "response.done") assert response_done["response"]["status"] == "completed" @@ -1978,9 +1842,7 @@ def patch_gemini_audio_cost_map_entries(monkeypatch): ("gemini-2.5-flash", False), ], ) -def test_is_audio_only_live_model_uses_cost_map( - model, expected, patch_gemini_audio_cost_map_entries -): +def test_is_audio_only_live_model_uses_cost_map(model, expected, patch_gemini_audio_cost_map_entries): assert GeminiRealtimeConfig._is_audio_only_live_model(model) == expected @@ -1994,9 +1856,7 @@ def test_is_audio_only_live_model_uses_cost_map( ("gemini-2.0-flash", False), ], ) -def test_is_native_audio_model_uses_cost_map( - model, expected, patch_gemini_audio_cost_map_entries -): +def test_is_native_audio_model_uses_cost_map(model, expected, patch_gemini_audio_cost_map_entries): assert GeminiRealtimeConfig._is_native_audio_model(model) == expected