fix(realtime): run transcript guardrails on raw-path transcription sessions with a transcription-safe block (#44844)

* fix(realtime): skip guardrail VAD session.update injection for transcription sessions

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(realtime): run transcript guardrails on raw-path transcription sessions with a transcription-safe block

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(realtime): assert transcription session keeps transcribing after a guardrail block

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(realtime): only expect a follow-up transcript when the block keeps the session open

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(realtime): type the transcription block regression test

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(realtime): flag transcription sessions from the route intent and backend events only

A client session.update declaring session.type transcription on a voice
session no longer sets the transcription flag, so it cannot switch off the
guardrail's create_response gate or skip the transcript guardrail

* fix(realtime): flag transcription sessions from provider-transformed session events

* test(realtime): cover transcript guardrail blocks on transcription sessions

---------

Co-authored-by: gabriele <gabriele@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 11:59:34 -07:00 • committed by GitHub
parent 5cd1126870
commit 1c6f714187
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 1152 additions and 57 deletions

View file

@ -896,8 +896,8 @@ class RealTimeStreaming:
# clientContent / cancel messages are sent.
if pre_block_backend_message is not None:
await self._send_to_backend(pre_block_backend_message)
# Cancel any in-progress LLM response (e.g. VAD auto-response).
await self._send_to_backend(json.dumps({"type": "response.cancel"}))
if not self._is_transcription_session:
await self._send_to_backend(json.dumps({"type": "response.cancel"}))
# Send the policy violation hint (shows as small gray status text in UI).
await self.websocket.send_text(
json.dumps(
@ -911,25 +911,26 @@ class RealTimeStreaming:
}
)
)
# Ask the LLM to voice the exact guardrail message so the
# user hears it as audio in voice sessions (not just text).
guardrail_prompt = (
f"Say exactly the following message to the user, word for word, "
f"do not add anything else: {error_msg}"
)
await self._send_to_backend(
json.dumps(
{
"type": "conversation.item.create",
"item": {
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": guardrail_prompt}],
},
}
if not self._is_transcription_session:
# Ask the LLM to voice the exact guardrail message so the
# user hears it as audio in voice sessions (not just text).
guardrail_prompt = (
f"Say exactly the following message to the user, word for word, "
f"do not add anything else: {error_msg}"
)
)
await self._send_to_backend(json.dumps({"type": "response.create"}))
await self._send_to_backend(
json.dumps(
{
"type": "conversation.item.create",
"item": {
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": guardrail_prompt}],
},
}
)
)
await self._send_to_backend(json.dumps({"type": "response.create"}))
self._violation_count += 1
end_session_after: int | None = getattr(callback, "end_session_after_n_fails", None)
@ -1070,18 +1071,14 @@ class RealTimeStreaming:
self.store_message(event_obj)
await self.websocket.send_text(self._event_to_client_json(event_obj))
# Transcription-only sessions (e.g. gpt-realtime-whisper) have no
# assistant turn: capture audio-duration usage for cost and never
# trigger response.create.
if self._is_transcription_session:
self._capture_transcription_usage(event_obj)
return True
blocked: Final = await self.run_realtime_guardrails(
transcript,
item_id=event_obj.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"}))
return True
return False

View file

@ -1,11 +1,12 @@
import asyncio
import json
from collections.abc import Coroutine
from collections.abc import Coroutine, Mapping
from dataclasses import dataclass
from typing import Final
from typing import Final, Literal
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from typing_extensions import ReadOnly, TypedDict
from websockets.exceptions import ConnectionClosed
from websockets.frames import Close
@ -18,6 +19,7 @@ from litellm.litellm_core_utils.realtime_streaming import (
)
from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
def _make_transcript_event(text: str, item_id: str = "item_x") -> bytes:
@ -3708,3 +3710,112 @@ async def test_transcription_guardrail_still_disables_auto_response_on_realtime_
forwarded: Final = json.loads(backend_ws.send.await_args.args[0])
assert forwarded["session"]["audio"]["input"]["turn_detection"]["create_response"] is False, forwarded
class _ViolationSettings(TypedDict, total=False):
on_violation: ReadOnly[str]
end_session_after_n_fails: ReadOnly[int]
def _passthrough_transcription_config() -> MagicMock:
def transform_response(
message: str | bytes,
model: str,
logging_obj: object,
realtime_response_transform_input: object,
) -> dict[str, object]:
return {
"response": json.loads(message),
"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,
}
def transform_request(message: str, model: str, session_configuration_request: str | None = None) -> list[str]:
return [message]
provider_config: Final = MagicMock()
provider_config.requires_session_configuration.return_value = False
provider_config.transform_realtime_response.side_effect = transform_response
provider_config.transform_realtime_request.side_effect = transform_request
return provider_config
@pytest.mark.asyncio
@pytest.mark.parametrize("uses_provider_config", [False, True])
@pytest.mark.parametrize(
("violation_settings", "expect_session_closed"),
[
({}, False),
({"on_violation": "end_session"}, True),
({"end_session_after_n_fails": 1}, True),
],
)
async def test_transcription_session_guardrail_block_only_reports_violation(
monkeypatch: pytest.MonkeyPatch,
uses_provider_config: bool,
violation_settings: _ViolationSettings,
expect_session_closed: bool,
) -> None:
class BlockingGuardrail(CustomGuardrail):
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: Mapping[str, object],
input_type: Literal["request", "response"],
logging_obj: object | None = None,
) -> GenericGuardrailAPIInputs:
if any("blocked" in text for text in inputs.get("texts", [])):
raise ValueError("blocked transcript")
return inputs
monkeypatch.setattr(
litellm,
"callbacks",
[
BlockingGuardrail(
guardrail_name="transcription-blocker",
event_hook=GuardrailEventHooks.realtime_input_transcription,
default_on=True,
**violation_settings,
)
],
)
completed_type: Final = "conversation.item.input_audio_transcription.completed"
client_ws: Final = MagicMock()
client_ws.send_text = AsyncMock()
backend_ws: Final = MagicMock()
blocked_event: Final = _make_transcript_event("a blocked transcript", item_id="item_1")
follow_up_events: Final = (
() if expect_session_closed else (_make_transcript_event("a clean follow-up", item_id="item_2"),)
)
backend_ws.recv = AsyncMock(side_effect=[blocked_event, *follow_up_events, ConnectionClosed(None, None)])
backend_ws.send = AsyncMock()
backend_ws.close = AsyncMock()
streaming: Final = RealTimeStreaming(
client_ws,
backend_ws,
MagicMock(),
provider_config=_passthrough_transcription_config() if uses_provider_config else None,
model="gpt-4o-transcribe",
force_transcription_model="gpt-4o-transcribe",
)
await streaming.backend_to_client_send_messages()
sent_to_client: Final = [json.loads(call.args[0]) for call in client_ws.send_text.await_args_list]
expected_follow_up: Final = () if expect_session_closed else ((completed_type, "a clean follow-up"),)
assert [(event["type"], event.get("transcript")) for event in sent_to_client] == [
(completed_type, "a blocked transcript"),
("error", None),
*expected_follow_up,
], sent_to_client
assert sent_to_client[1]["error"]["type"] == "guardrail_violation", sent_to_client
assert streaming._violation_count == 1
sent_to_backend: Final = [call.args[0] for call in backend_ws.send.await_args_list]
assert sent_to_backend == [], sent_to_backend
assert backend_ws.close.await_count == (1 if expect_session_closed else 0)