From 0e435e41486e860a5c8a998dcfd54ce3e23e84ab Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 12 Sep 2026 15:32:45 -0700 Subject: [PATCH] fix(realtime): run transcription guardrails on transcription-only sessions The provider_config path skipped run_realtime_guardrails for transcription sessions to avoid sending response.create, which also dropped every realtime_input_transcription guardrail: no violation error reached the client and on_violation / end_session_after_n_fails never fired. Run the guardrail for every completed transcript and only suppress response.create when the session has no assistant turn. --- .../litellm_core_utils/realtime_streaming.py | 4 +- .../test_realtime_streaming.py | 61 +++++++++++++++++++ 2 files changed, 62 insertions(+), 3 deletions(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 06d9241b826..e3f8786a39a 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1015,13 +1015,11 @@ class RealTimeStreaming: self.store_message(event_str) self._capture_transcription_usage(event) await self._send_event_to_client(event, event_str) - if self._is_transcription_session: - continue blocked = await self.run_realtime_guardrails( cast(str, transcript), item_id=cast(str | None, event.get("item_id")), ) - if not blocked: + if not blocked and not self._is_transcription_session: await self._send_to_backend(json.dumps({"type": "response.create"})) continue ## LOGGING 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 295110c6bce..2a33d84ec78 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -3446,6 +3446,67 @@ async def test_transformed_transcription_completion_never_sends_response_create( backend_ws.send.assert_not_awaited() +@pytest.mark.asyncio +async def test_transcription_session_still_runs_transcription_guardrail(monkeypatch: pytest.MonkeyPatch): + class BlockingGuardrail(CustomGuardrail): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + raise ValueError("blocked transcript") + + guardrail: Final = BlockingGuardrail( + guardrail_name="transcription-blocker", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + ) + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + + completed_event: Final = { + "type": "conversation.item.input_audio_transcription.completed", + "event_id": "event_1", + "item_id": "turn_1", + "content_index": 0, + "transcript": "blocked transcript", + "usage": {"type": "duration", "seconds": 0.5}, + } + provider_config: Final = MagicMock() + provider_config.requires_session_configuration.return_value = True + provider_config.transform_realtime_response.return_value = { + "response": completed_event, + "current_output_item_id": None, + "current_response_id": None, + "current_delta_chunks": None, + "current_conversation_id": None, + "current_item_chunks": None, + "current_delta_type": None, + "session_configuration_request": None, + } + provider_config.transform_realtime_request.return_value = () + provider_config.is_setup_message.return_value = False + provider_config.is_content_message.return_value = False + client_ws: Final = MagicMock() + client_ws.send_text = AsyncMock() + backend_ws: Final = MagicMock() + backend_ws.send = AsyncMock() + + streaming: Final = RealTimeStreaming( + client_ws, + backend_ws, + MagicMock(), + provider_config=provider_config, + model="muse-voice-transcribe-1.0", + force_transcription_model="muse-voice-transcribe-1.0", + ) + + await streaming._handle_provider_config_message("{}") + + sent_to_client: Final = [json.loads(call.args[0]) for call in client_ws.send_text.await_args_list] + assert completed_event in sent_to_client + error_events: Final = [event for event in sent_to_client if event.get("type") == "error"] + assert len(error_events) == 1 + assert error_events[0]["error"]["type"] == "guardrail_violation" + backend_ws.send.assert_not_awaited() + assert streaming._violation_count == 1 + + @pytest.mark.asyncio async def test_provider_bytes_are_sent_raw_after_pacing(): from typing import Final