diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index d2fbb26bb02..a9e538fcfcd 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -806,6 +806,20 @@ class RealTimeStreaming: return self._has_realtime_guardrails_for_event_hooks([GuardrailEventHooks.realtime_input_transcription]) + def _proxy_drives_turns(self) -> bool: + """True when the proxy, not the backend's server-VAD, starts the assistant's response. + + With a ``realtime_input_transcription`` guardrail the proxy sets + ``turn_detection.create_response: false`` on the backend and sends ``response.create`` + itself after the transcript passed the guardrail. Without such a guardrail the backend + auto-responds as soon as the caller stops speaking, so a ``response.create`` from the + proxy would start a second response for the same turn; the backend rejects it with + ``conversation_already_has_active_response`` (#31726). + + Transcription-only sessions have no assistant response at all. + """ + return not self._is_transcription_session and self._has_audio_transcription_guardrails() + async def run_realtime_guardrails( self, transcript: str, @@ -1017,7 +1031,9 @@ class RealTimeStreaming: cast(str, transcript), item_id=cast(str | None, event.get("item_id")), ) - if not blocked and not self._is_transcription_session: + # Send response.create only if the proxy disabled the backend's auto-response + # (transcript guardrail configured). Otherwise the backend created it already. + if not blocked and self._proxy_drives_turns(): await self._send_to_backend(json.dumps({"type": "response.create"})) continue ## LOGGING @@ -1068,7 +1084,9 @@ class RealTimeStreaming: transcript, item_id=event_obj.get("item_id"), ) - if not blocked: + # Send response.create only if the proxy disabled the backend's auto-response + # (transcript guardrail configured). Otherwise the backend created it already. + if not blocked and self._proxy_drives_turns(): await self._send_to_backend(json.dumps({"type": "response.create"})) return True return False diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index 7e6d4d24905..eea3b83d744 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -902,10 +902,12 @@ async def test_transcription_session_captures_usage_and_skips_response_create(): @pytest.mark.asyncio -async def test_non_transcription_completed_event_still_triggers_response_create(): +async def test_completed_transcription_without_guardrails_does_not_inject_response_create(): """ - Regression guard: a normal (non-transcription) session with no guardrails must - keep triggering response.create on a completed transcription event. + Regression guard for #31726. Without a ``realtime_input_transcription`` guardrail the + backend's server-VAD auto-response stays on, so the backend already created this turn's + response. The proxy must not send its own ``response.create`` (the backend would reject it + with ``conversation_already_has_active_response``), but must still forward the transcript. """ client_ws = MagicMock() client_ws.send_text = AsyncMock() @@ -931,7 +933,127 @@ async def test_non_transcription_completed_event_still_triggers_response_create( 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] - assert any(e.get("type") == "response.create" for e in sent_to_backend) + assert all(e.get("type") != "response.create" for e in sent_to_backend), sent_to_backend + forwarded = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list if c.args] + assert any(e.get("type") == "conversation.item.input_audio_transcription.completed" for e in forwarded) + + +@pytest.mark.asyncio +async def test_provider_config_completed_transcription_without_guardrails_does_not_inject_response_create(): + """Same contract on the provider_config path (OpenAI / Gemini / Vertex transformed backends).""" + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + + completed_event = { + "type": "conversation.item.input_audio_transcription.completed", + "transcript": "hi", + "item_id": "item_1", + } + provider_config = MagicMock() + provider_config.transform_realtime_response = MagicMock( + return_value={ + "response": [completed_event], + "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.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + streaming = RealTimeStreaming( + client_ws, + backend_ws, + logging_obj, + provider_config=provider_config, + model="gpt-realtime", + ) + await streaming._handle_provider_config_message(json.dumps(completed_event)) + + assert backend_ws.send.await_count == 0 + forwarded = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list if c.args] + assert any(e.get("type") == "conversation.item.input_audio_transcription.completed" for e in forwarded) + + +@pytest.mark.asyncio +async def test_provider_config_completed_transcription_with_guardrail_injects_response_create( + monkeypatch: pytest.MonkeyPatch, +): + """Other half of the pair on the provider_config path: with a ``realtime_input_transcription`` + guardrail the proxy disabled the backend's auto-response, so after a clean transcript it must + send exactly one ``response.create`` (mirrors ``test_realtime_guardrail_allows_clean_transcript`` + for the raw path).""" + import litellm + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.types.guardrails import GuardrailEventHooks + + class AudioGuardrail(CustomGuardrail): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + return inputs + + guardrail = AudioGuardrail( + guardrail_name="audio-guardrail", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + ) + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + + client_ws = MagicMock() + client_ws.send_text = AsyncMock() + + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + + completed_event = { + "type": "conversation.item.input_audio_transcription.completed", + "transcript": "What are the opening hours tomorrow?", + "item_id": "item_1", + } + provider_config = MagicMock() + provider_config.transform_realtime_response = MagicMock( + return_value={ + "response": [completed_event], + "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, + } + ) + # Pass-through transform so the assertion sees the exact frame the proxy chose to send. + provider_config.transform_realtime_request = MagicMock(side_effect=lambda message, *_: (message,)) + provider_config.is_setup_message.return_value = False + provider_config.is_content_message.return_value = False + + logging_obj = MagicMock() + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + streaming = RealTimeStreaming( + client_ws, + backend_ws, + logging_obj, + provider_config=provider_config, + model="gpt-realtime", + ) + await streaming._handle_provider_config_message(json.dumps(completed_event)) + + 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"Guardrail-gated turn must trigger response.create, got: {sent_to_backend}" + forwarded = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list if c.args] + assert any(e.get("type") == "conversation.item.input_audio_transcription.completed" for e in forwarded) def test_client_session_update_marks_transcription_session():