diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index 72f49a7ea84..eea3b83d744 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -984,6 +984,78 @@ async def test_provider_config_completed_transcription_without_guardrails_does_n 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(): """A client session.update with type=transcription flags the session.""" streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock())