From 817792c6cc3c2cc402b260e38e612a6929286d8c Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 27 Apr 2026 19:05:59 +0530 Subject: [PATCH] fix lint --- .../litellm_core_utils/realtime_streaming.py | 29 +++++++--- litellm/llms/custom_httpx/llm_http_handler.py | 1 + .../test_realtime_streaming.py | 58 +++++++++++++++++++ 3 files changed, 80 insertions(+), 8 deletions(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index c640c1e52ad..3550363601b 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -78,6 +78,9 @@ class RealTimeStreaming: # Track whether session.created has already been sent to the client # (e.g. synthetic event in deferred setup mode). self._session_created_sent_to_client: bool = False + # Track whether we have already sent the guardrail turn-detection update + # that disables provider auto-response for transcription guardrails. + self._guardrail_turn_detection_update_sent: bool = False def _should_store_message( self, @@ -245,6 +248,22 @@ class RealTimeStreaming: except (json.JSONDecodeError, TypeError): return + async def _maybe_send_guardrail_turn_detection_update(self) -> None: + """Disable provider auto-response once when transcription guardrails are enabled.""" + if self._guardrail_turn_detection_update_sent: + return + if not self._has_audio_transcription_guardrails(): + return + await self._send_to_backend( + json.dumps( + { + "type": "session.update", + "session": {"turn_detection": {"create_response": False}}, + } + ) + ) + self._guardrail_turn_detection_update_sent = True + def _has_realtime_guardrails(self) -> bool: """Return True if any callback is registered for realtime guardrail event types.""" from litellm.integrations.custom_guardrail import CustomGuardrail @@ -446,6 +465,7 @@ class RealTimeStreaming: event_str = json.dumps(event) if isinstance(event, dict) and event.get("type") == "session.created": if self._session_created_sent_to_client: + await self._maybe_send_guardrail_turn_detection_update() verbose_logger.debug( "Skipping duplicate session.created from provider stream" ) @@ -459,14 +479,7 @@ class RealTimeStreaming: ): self.store_message(event_str) await self.websocket.send_text(event_str) - await self._send_to_backend( - json.dumps( - { - "type": "session.update", - "session": {"turn_detection": {"create_response": False}}, - } - ) - ) + await self._maybe_send_guardrail_turn_detection_update() continue ## GUARDRAIL: run on transcription events in provider_config path too if ( diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 7149e16fdfe..d0b2d5e3da7 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -5167,6 +5167,7 @@ class BaseLLMHTTPHandler: if synthetic_session is not None: await websocket.send_text(json.dumps(synthetic_session)) realtime_streaming._session_created_sent_to_client = True + await realtime_streaming._maybe_send_guardrail_turn_detection_update() verbose_logger.debug( "Sent synthetic session.created to client to unblock connection" ) 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 78443664d54..a1fcc8385d1 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -1022,3 +1022,61 @@ async def test_provider_path_suppresses_duplicate_session_created_after_syntheti assert not any( payload.get("type") == "session.created" for payload in sent_payloads ), f"Expected duplicate session.created to be suppressed, got: {sent_payloads}" + + +@pytest.mark.asyncio +async def test_duplicate_session_created_still_triggers_guardrail_turn_detection_update(): + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + + backend_ws = MagicMock() + backend_ws.recv = AsyncMock( + side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)] + ) + backend_ws.send = AsyncMock() + + provider_config = MagicMock() + provider_config.transform_realtime_response = MagicMock( + return_value={ + "response": [ + { + "type": "session.created", + "event_id": "event_1", + "session": {"id": "sess_1", "modalities": ["audio"]}, + } + ], + "current_output_item_id": None, + "current_response_id": None, + "current_delta_chunks": [], + "current_conversation_id": None, + "current_item_chunks": [], + "current_delta_type": None, + "session_configuration_request": None, + } + ) + + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_1" + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + streaming = RealTimeStreaming( + websocket=client_ws, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + # Synthetic session.created already sent by llm_http_handler. + streaming._session_created_sent_to_client = True + streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign] + streaming._send_to_backend = AsyncMock() # type: ignore[method-assign] + + await streaming.backend_to_client_send_messages() + + # Duplicate session.created should still cause the one-time guardrail + # turn_detection update to be sent to backend. + assert streaming._send_to_backend.await_count == 1 + sent_update = json.loads(streaming._send_to_backend.await_args_list[0].args[0]) + assert sent_update["type"] == "session.update" + assert sent_update["session"]["turn_detection"]["create_response"] is False