mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
5cd1126870
commit
1c6f714187
3 changed files with 1152 additions and 57 deletions
|
|
@ -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
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue