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.
This commit is contained in:
ryan-crabbe-berri 2026-09-12 15:32:45 -07:00
parent 4647cd1215
commit 0e435e4148
2 changed files with 62 additions and 3 deletions

View file

@ -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

View file

@ -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