diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index a44e6cff1f7..85799c729ad 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -899,9 +899,11 @@ class RealTimeStreaming: ): session = msg_obj.setdefault("session", {}) if isinstance(session, dict): - session.setdefault("turn_detection", {})[ - "create_response" - ] = False + existing_td = session.get("turn_detection") + if not isinstance(existing_td, dict): + existing_td = {} + existing_td["create_response"] = False + session["turn_detection"] = existing_td message = json.dumps(msg_obj) self._guardrail_turn_detection_update_sent = True verbose_logger.debug( 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 c62ef6e23b0..68b31ae5bfa 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -1368,3 +1368,61 @@ async def test_guardrail_turn_detection_injected_into_first_session_update_defer assert injected_turn_detection is not None assert injected_turn_detection["create_response"] is False assert streaming._guardrail_turn_detection_update_sent is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("existing_turn_detection", [None, "auto", 42, ["server_vad"]]) +async def test_guardrail_turn_detection_injection_tolerates_non_dict_value( + existing_turn_detection, +): + """Client-supplied non-dict turn_detection must not crash client_ack_messages.""" + client_ws = AsyncMock() + client_ws.receive_text = AsyncMock( + side_effect=[ + json.dumps({ + "type": "session.update", + "session": { + "modalities": ["text", "audio"], + "turn_detection": existing_turn_detection, + }, + }), + ConnectionClosed(None, None), + ] + ) + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + + logging_obj = MagicMock() + logging_obj.litellm_trace_id = "trace_1" + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + provider_config = MagicMock() + transformed_messages = [] + def mock_transform(msg, model, session_config): + transformed_messages.append((msg, session_config)) + return [msg] + provider_config.transform_realtime_request = MagicMock(side_effect=mock_transform) + + streaming = RealTimeStreaming( + websocket=client_ws, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign] + + await streaming.client_ack_messages() + + assert len(transformed_messages) == 1 + transformed_msg, _ = transformed_messages[0] + msg_obj = json.loads(transformed_msg) + session_obj = msg_obj["session"] + injected_turn_detection = ( + session_obj.get("turn_detection") + or session_obj.get("audio", {}).get("input", {}).get("turn_detection") + ) + assert isinstance(injected_turn_detection, dict) + assert injected_turn_detection["create_response"] is False + assert streaming._guardrail_turn_detection_update_sent is True