From ef5d05f137d719924a8b1116ca6c0ea17ea6e2c3 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Sat, 27 Jun 2026 20:22:20 -0700 Subject: [PATCH] fix(realtime): stop second Gemini Live setup, retry hung handshake, close guardrail bypass (#31519) * fix(realtime): stop sending a second Gemini Live setup on follow-up session.update Gemini Live (BidiGenerateContent) accepts setup as the first-and-only client message; a second setup closes the socket with 1007 Request contains an invalid argument. The AI Studio Gemini path forwarded every client session.update after the first as a follow-up setup, and GA clients (pipecat) send several while configuring the session, so the second one tore the session down before the first turn. Callers saw silence after the first response, exponential per-turn latency from reconnect/retry churn, and intermittent 1011 errors. Drop subsequent session.updates instead of resending setup, matching what the Vertex subclass already does. Tools and instructions must ride on the first session.update before any conversation content. Adds regression tests covering the plain follow-up, a follow-up that adds tools (the case the previous identical-only dedup still forwarded), and the guardrail create_response=False warning path. * fix(realtime): retry the backend open handshake instead of failing with 1011 The upstream Live API open handshake (e.g. Gemini Live) intermittently hangs; waiting longer never recovers a hung attempt, but a fresh attempt almost always connects in ~1s. The proxy opened the backend websocket once with the default open_timeout and no retry, so a single slow handshake surfaced to the caller as a fatal 1011 internal error and dropped the call. Bound each open attempt with a short open_timeout and retry; a bounded attempt that already timed out spaces out the next try, so no backoff is needed. Deterministic handshake-status rejections (auth/4xx) are not retried, and the retry only ever wraps the open, never a live session. Adds tests for retry-then-succeed, raise-after-max-attempts, and no-retry-on-auth-failure. * fix(realtime): close guardrail bypass + surface handshake status; drop obsolete tests Three review fixes on the Gemini Live realtime path. Transcription-guardrail bypass: Gemini Live rejects a second setup (1007), so once the initial setup is sent the guardrail's automaticActivityDetection.disabled=true can no longer be delivered as a follow-up session.update. With that follow-up now dropped, the model's auto-response stayed enabled and a realtime_input_transcription guardrail was bypassed (the model answered before the proxy could gate the turn). Fold the disable into the one-and-only setup instead: the handler injects it into the auto-sent setup (gemini_live_defer_setup false) and _send_to_backend injects it into the deferred first setup. OpenAI sessions accept follow-up updates and are left untouched. Backend handshake status: the open-retry treated only InvalidStatusCode as deterministic; websockets>=15 raises InvalidStatus for a rejected client handshake, so a 401/403 fell into the broad WebSocketException branch and was retried before the caller closed the client with 1011 instead of the upstream status. Treat both as non-retryable. Obsolete tests: the four tests asserting a follow-up session.update is merged and re-sent as a second setup asserted behavior that crashes Gemini Live with 1007 (verified directly against the API). Removed; the drop is covered by new regression tests. * style(realtime): reformat changed files to ruff line-length 120 Post-merge with litellm_internal_staging, which unified ruff format width to 120 (#31518). The realtime change set was formatted at 88, so the changed lines tripped the whole-file ruff format check. Reformat with ruff 0.15.3 at the repo's 120 width; no logic changes. --- .../litellm_core_utils/realtime_streaming.py | 31 ++ litellm/llms/custom_httpx/llm_http_handler.py | 86 ++- .../llms/gemini/realtime/transformation.py | 99 +--- .../test_realtime_streaming.py | 507 +++++++----------- .../custom_httpx/test_llm_http_handler.py | 164 ++++-- .../test_gemini_realtime_transformation.py | 404 +++++--------- 6 files changed, 558 insertions(+), 733 deletions(-) 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