mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
feat(realtime): add output text guardrails for /realtime WebSocket proxy
Adds a new `realtime_output_text` guardrail hook that buffers `response.text.delta` events and checks the full text on `response.text.done` before forwarding to the client. If a guardrail blocks the response, the buffered deltas are dropped and replaced with an error delta + done pair. If clean, the buffer is flushed and normal forwarding resumes. Works on both the provider-config (Gemini/Bedrock) and raw OpenAI forwarding paths in RealTimeStreaming. The built-in `litellm_content_filter` guardrail is wired up to this new hook via `mode: realtime_output_text` in proxy config.
This commit is contained in:
parent
f78104d34c
commit
046d032084
4 changed files with 446 additions and 1 deletions
|
|
@ -1,7 +1,7 @@
|
|||
import asyncio
|
||||
import concurrent.futures
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -68,6 +68,8 @@ class RealTimeStreaming:
|
|||
self.current_delta_type: Optional[ALL_DELTA_TYPES] = None
|
||||
self.session_configuration_request: Optional[str] = None
|
||||
self.user_api_key_dict = user_api_key_dict
|
||||
# Buffer for response.text.delta events pending output-guardrail check
|
||||
self._pending_output_text_events: List[str] = []
|
||||
|
||||
def _should_store_message(
|
||||
self,
|
||||
|
|
@ -299,6 +301,73 @@ class RealTimeStreaming:
|
|||
return True
|
||||
return False
|
||||
|
||||
def _has_realtime_output_guardrails(self) -> bool:
|
||||
"""Return True if any callback is registered for realtime_output_text."""
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
return any(
|
||||
isinstance(cb, CustomGuardrail)
|
||||
and cb.should_run_guardrail(
|
||||
data={},
|
||||
event_type=GuardrailEventHooks.realtime_output_text,
|
||||
)
|
||||
for cb in litellm.callbacks
|
||||
)
|
||||
|
||||
async def run_realtime_output_guardrails(
|
||||
self, text: str
|
||||
) -> Tuple[bool, str]:
|
||||
"""
|
||||
Run registered guardrails on completed response text.
|
||||
|
||||
Returns (True, error_msg) if blocked, (False, "") if clean.
|
||||
"""
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if not isinstance(callback, CustomGuardrail):
|
||||
continue
|
||||
if (
|
||||
callback.should_run_guardrail(
|
||||
data={"text": text},
|
||||
event_type=GuardrailEventHooks.realtime_output_text,
|
||||
)
|
||||
is not True
|
||||
):
|
||||
continue
|
||||
try:
|
||||
await callback.apply_guardrail(
|
||||
inputs={"texts": [text], "images": []},
|
||||
request_data={"user_api_key_dict": self.user_api_key_dict},
|
||||
input_type="response",
|
||||
)
|
||||
except Exception as e:
|
||||
is_guardrail_block = hasattr(e, "status_code") or isinstance(
|
||||
e, ValueError
|
||||
)
|
||||
if not is_guardrail_block:
|
||||
verbose_logger.exception(
|
||||
"[realtime output guardrail] unexpected error: %s", e
|
||||
)
|
||||
raise
|
||||
detail = getattr(e, "detail", None)
|
||||
if isinstance(detail, dict):
|
||||
safe_msg = detail.get("error") or str(e)
|
||||
elif detail is not None:
|
||||
safe_msg = str(detail)
|
||||
else:
|
||||
safe_msg = (
|
||||
str(e) or "Response blocked by content filter."
|
||||
)
|
||||
verbose_logger.warning(
|
||||
"[realtime output guardrail] BLOCKED output text: %r",
|
||||
text[:80],
|
||||
)
|
||||
return True, safe_msg
|
||||
return False, ""
|
||||
|
||||
async def _handle_provider_config_message(self, raw_response) -> None:
|
||||
"""Process a backend message when a provider_config is set (transformed path)."""
|
||||
returned_object = self.provider_config.transform_realtime_response( # type: ignore[union-attr]
|
||||
|
|
@ -366,6 +435,45 @@ class RealTimeStreaming:
|
|||
json.dumps({"type": "response.create"})
|
||||
)
|
||||
continue
|
||||
## OUTPUT GUARDRAIL: buffer text deltas; check on text done
|
||||
if (
|
||||
isinstance(event, dict)
|
||||
and event.get("type") == "response.text.delta"
|
||||
and self._has_realtime_output_guardrails()
|
||||
):
|
||||
self._pending_output_text_events.append(event_str)
|
||||
continue
|
||||
if isinstance(event, dict) and event.get("type") == "response.text.done":
|
||||
full_text = event.get("text", "")
|
||||
blocked_out, error_msg = await self.run_realtime_output_guardrails(
|
||||
full_text
|
||||
)
|
||||
if blocked_out:
|
||||
self._pending_output_text_events.clear()
|
||||
error_delta_str = json.dumps(
|
||||
{
|
||||
"type": "response.text.delta",
|
||||
"delta": error_msg,
|
||||
"content_index": event.get("content_index", 0),
|
||||
"item_id": event.get("item_id", ""),
|
||||
"output_index": event.get("output_index", 0),
|
||||
"response_id": event.get("response_id", ""),
|
||||
}
|
||||
)
|
||||
error_done = dict(event)
|
||||
error_done["text"] = error_msg
|
||||
error_done_str = json.dumps(error_done)
|
||||
self.store_message(error_done_str)
|
||||
await self.websocket.send_text(error_delta_str)
|
||||
await self.websocket.send_text(error_done_str)
|
||||
else:
|
||||
for pending in self._pending_output_text_events:
|
||||
self.store_message(pending)
|
||||
await self.websocket.send_text(pending)
|
||||
self._pending_output_text_events.clear()
|
||||
self.store_message(event_str)
|
||||
await self.websocket.send_text(event_str)
|
||||
continue
|
||||
## LOGGING
|
||||
self.store_message(event_str)
|
||||
await self.websocket.send_text(event_str)
|
||||
|
|
@ -420,6 +528,44 @@ class RealTimeStreaming:
|
|||
json.dumps({"type": "response.create"})
|
||||
)
|
||||
return True
|
||||
|
||||
if event_obj.get("type") == "response.text.delta":
|
||||
if self._has_realtime_output_guardrails():
|
||||
self._pending_output_text_events.append(raw_response)
|
||||
return True
|
||||
|
||||
if event_obj.get("type") == "response.text.done":
|
||||
full_text = event_obj.get("text", "")
|
||||
blocked_out, error_msg = await self.run_realtime_output_guardrails(
|
||||
full_text
|
||||
)
|
||||
if blocked_out:
|
||||
self._pending_output_text_events.clear()
|
||||
error_delta_str = json.dumps(
|
||||
{
|
||||
"type": "response.text.delta",
|
||||
"delta": error_msg,
|
||||
"content_index": event_obj.get("content_index", 0),
|
||||
"item_id": event_obj.get("item_id", ""),
|
||||
"output_index": event_obj.get("output_index", 0),
|
||||
"response_id": event_obj.get("response_id", ""),
|
||||
}
|
||||
)
|
||||
error_done = dict(event_obj)
|
||||
error_done["text"] = error_msg
|
||||
error_done_str = json.dumps(error_done)
|
||||
self.store_message(error_done_str)
|
||||
await self.websocket.send_text(error_delta_str)
|
||||
await self.websocket.send_text(error_done_str)
|
||||
else:
|
||||
for pending in self._pending_output_text_events:
|
||||
self.store_message(pending)
|
||||
await self.websocket.send_text(pending)
|
||||
self._pending_output_text_events.clear()
|
||||
self.store_message(raw_response)
|
||||
await self.websocket.send_text(raw_response)
|
||||
return True
|
||||
|
||||
except (json.JSONDecodeError, AttributeError):
|
||||
pass
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -194,6 +194,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
GuardrailEventHooks.post_call,
|
||||
GuardrailEventHooks.during_call,
|
||||
GuardrailEventHooks.realtime_input_transcription,
|
||||
GuardrailEventHooks.realtime_output_text,
|
||||
],
|
||||
event_hook=event_hook or GuardrailEventHooks.pre_call,
|
||||
default_on=default_on,
|
||||
|
|
|
|||
|
|
@ -774,6 +774,7 @@ class GuardrailEventHooks(str, Enum):
|
|||
pre_mcp_call = "pre_mcp_call"
|
||||
during_mcp_call = "during_mcp_call"
|
||||
realtime_input_transcription = "realtime_input_transcription"
|
||||
realtime_output_text = "realtime_output_text"
|
||||
|
||||
|
||||
class DynamicGuardrailParams(TypedDict):
|
||||
|
|
|
|||
297
tests/test_litellm/realtime/test_output_guardrail.py
Normal file
297
tests/test_litellm/realtime/test_output_guardrail.py
Normal file
|
|
@ -0,0 +1,297 @@
|
|||
"""
|
||||
Unit tests for realtime output text guardrail.
|
||||
|
||||
Tests that response.text.delta events are buffered and checked on
|
||||
response.text.done, blocking bad output and forwarding clean output.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import unittest
|
||||
from typing import Any, List, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
class _FakeWebSocket:
|
||||
"""Mimics a FastAPI WebSocket (client ↔ proxy)."""
|
||||
|
||||
def __init__(self):
|
||||
self.sent: List[str] = []
|
||||
|
||||
async def send_text(self, text: str):
|
||||
self.sent.append(text)
|
||||
|
||||
|
||||
class _FakeBackendWS:
|
||||
"""Mimics a websockets.ClientConnection (proxy ↔ backend)."""
|
||||
|
||||
def __init__(self):
|
||||
self.sent: List[str] = []
|
||||
|
||||
async def send(self, data: str):
|
||||
self.sent.append(data)
|
||||
|
||||
async def recv(self, decode=True):
|
||||
raise StopAsyncIteration
|
||||
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
|
||||
class _BlockingGuardrail(CustomGuardrail):
|
||||
"""Blocks any text containing 'bomb'."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
guardrail_name="test-content-filter",
|
||||
supported_event_hooks=[GuardrailEventHooks.realtime_output_text],
|
||||
event_hook=GuardrailEventHooks.realtime_output_text,
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
|
||||
texts = inputs.get("texts", [])
|
||||
for text in texts:
|
||||
if "bomb" in text.lower():
|
||||
from fastapi import HTTPException
|
||||
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "Response blocked by content filter."},
|
||||
)
|
||||
return inputs
|
||||
|
||||
|
||||
class _PassthroughGuardrail(CustomGuardrail):
|
||||
"""Always passes."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
guardrail_name="test-passthrough",
|
||||
supported_event_hooks=[GuardrailEventHooks.realtime_output_text],
|
||||
event_hook=GuardrailEventHooks.realtime_output_text,
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
|
||||
return inputs
|
||||
|
||||
|
||||
def _make_streaming(guardrail=None):
|
||||
"""Create a RealTimeStreaming instance with mocked dependencies."""
|
||||
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
|
||||
|
||||
client_ws = _FakeWebSocket()
|
||||
backend_ws = _FakeBackendWS()
|
||||
logging_obj = MagicMock()
|
||||
|
||||
streaming = RealTimeStreaming(
|
||||
websocket=client_ws,
|
||||
backend_ws=backend_ws,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
# Patch store_message to be a no-op (avoids JSON parsing complexity)
|
||||
streaming.store_message = MagicMock()
|
||||
|
||||
if guardrail is not None:
|
||||
import litellm
|
||||
|
||||
litellm.callbacks = [guardrail]
|
||||
else:
|
||||
import litellm
|
||||
|
||||
litellm.callbacks = []
|
||||
|
||||
return streaming, client_ws, backend_ws
|
||||
|
||||
|
||||
# ── Tests ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestOutputGuardrailBlocked:
|
||||
"""response.text.done with blocked word → error replaces original text."""
|
||||
|
||||
def setup_method(self):
|
||||
import litellm
|
||||
|
||||
litellm.callbacks = []
|
||||
|
||||
def teardown_method(self):
|
||||
import litellm
|
||||
|
||||
litellm.callbacks = []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blocked_text_replaces_deltas(self):
|
||||
streaming, client_ws, _ = _make_streaming(guardrail=_BlockingGuardrail())
|
||||
|
||||
delta_event = {
|
||||
"type": "response.text.delta",
|
||||
"delta": "I will tell you how to make a bomb.",
|
||||
"content_index": 0,
|
||||
"item_id": "item_abc",
|
||||
"output_index": 0,
|
||||
"response_id": "resp_xyz",
|
||||
}
|
||||
done_event = {
|
||||
"type": "response.text.done",
|
||||
"text": "I will tell you how to make a bomb.",
|
||||
"content_index": 0,
|
||||
"item_id": "item_abc",
|
||||
"output_index": 0,
|
||||
"response_id": "resp_xyz",
|
||||
}
|
||||
|
||||
events = [delta_event, done_event]
|
||||
for event in events:
|
||||
event_str = json.dumps(event)
|
||||
# Simulate _handle_provider_config_message inner loop
|
||||
etype = event.get("type", "")
|
||||
if etype == "response.text.delta" and streaming._has_realtime_output_guardrails():
|
||||
streaming._pending_output_text_events.append(event_str)
|
||||
elif etype == "response.text.done":
|
||||
full_text = event.get("text", "")
|
||||
blocked, error_msg = await streaming.run_realtime_output_guardrails(full_text)
|
||||
if blocked:
|
||||
streaming._pending_output_text_events.clear()
|
||||
error_delta_str = json.dumps(
|
||||
{
|
||||
"type": "response.text.delta",
|
||||
"delta": error_msg,
|
||||
"content_index": event.get("content_index", 0),
|
||||
"item_id": event.get("item_id", ""),
|
||||
"output_index": event.get("output_index", 0),
|
||||
"response_id": event.get("response_id", ""),
|
||||
}
|
||||
)
|
||||
error_done = dict(event)
|
||||
error_done["text"] = error_msg
|
||||
await client_ws.send_text(error_delta_str)
|
||||
await client_ws.send_text(json.dumps(error_done))
|
||||
|
||||
# The original delta should NOT reach the client
|
||||
sent_texts = [json.loads(e) for e in client_ws.sent]
|
||||
for sent in sent_texts:
|
||||
assert "bomb" not in sent.get("delta", "").lower(), (
|
||||
f"Blocked word leaked in delta: {sent}"
|
||||
)
|
||||
assert "bomb" not in sent.get("text", "").lower(), (
|
||||
f"Blocked word leaked in done event: {sent}"
|
||||
)
|
||||
|
||||
# The client should have received exactly 2 events: error delta + error done
|
||||
assert len(client_ws.sent) == 2, f"Expected 2 events, got {len(client_ws.sent)}"
|
||||
error_delta = json.loads(client_ws.sent[0])
|
||||
error_done = json.loads(client_ws.sent[1])
|
||||
assert error_delta["type"] == "response.text.delta"
|
||||
assert error_done["type"] == "response.text.done"
|
||||
# Error message should mention blocking
|
||||
assert len(error_delta.get("delta", "")) > 0
|
||||
print(f"\n✅ PASS: blocked text replaced with: {error_delta['delta']!r}")
|
||||
|
||||
|
||||
class TestOutputGuardrailClean:
|
||||
"""response.text.done with clean text → buffered deltas flushed normally."""
|
||||
|
||||
def setup_method(self):
|
||||
import litellm
|
||||
|
||||
litellm.callbacks = []
|
||||
|
||||
def teardown_method(self):
|
||||
import litellm
|
||||
|
||||
litellm.callbacks = []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clean_text_flushed(self):
|
||||
streaming, client_ws, _ = _make_streaming(guardrail=_PassthroughGuardrail())
|
||||
|
||||
delta1 = {
|
||||
"type": "response.text.delta",
|
||||
"delta": "Hello, how can I help?",
|
||||
"content_index": 0,
|
||||
"item_id": "item_1",
|
||||
"output_index": 0,
|
||||
"response_id": "resp_1",
|
||||
}
|
||||
done_event = {
|
||||
"type": "response.text.done",
|
||||
"text": "Hello, how can I help?",
|
||||
"content_index": 0,
|
||||
"item_id": "item_1",
|
||||
"output_index": 0,
|
||||
"response_id": "resp_1",
|
||||
}
|
||||
|
||||
# Buffer the delta
|
||||
streaming._pending_output_text_events.append(json.dumps(delta1))
|
||||
|
||||
# Process done event
|
||||
blocked, error_msg = await streaming.run_realtime_output_guardrails(done_event["text"])
|
||||
assert not blocked, f"Expected not blocked, got error: {error_msg}"
|
||||
|
||||
# Flush buffer
|
||||
for pending in streaming._pending_output_text_events:
|
||||
streaming.store_message(pending)
|
||||
await client_ws.send_text(pending)
|
||||
streaming._pending_output_text_events.clear()
|
||||
await client_ws.send_text(json.dumps(done_event))
|
||||
|
||||
# Both delta and done should reach the client
|
||||
assert len(client_ws.sent) == 2
|
||||
d = json.loads(client_ws.sent[0])
|
||||
done = json.loads(client_ws.sent[1])
|
||||
assert d["type"] == "response.text.delta"
|
||||
assert done["type"] == "response.text.done"
|
||||
assert "Hello" in d["delta"]
|
||||
print(f"\n✅ PASS: clean text flushed normally: {d['delta']!r}")
|
||||
|
||||
|
||||
class TestOutputGuardrailNoGuardrail:
|
||||
"""Without output guardrail, deltas pass straight through (no buffering)."""
|
||||
|
||||
def setup_method(self):
|
||||
import litellm
|
||||
|
||||
litellm.callbacks = []
|
||||
|
||||
def teardown_method(self):
|
||||
import litellm
|
||||
|
||||
litellm.callbacks = []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_buffering_without_guardrail(self):
|
||||
streaming, client_ws, _ = _make_streaming(guardrail=None)
|
||||
|
||||
# With no guardrail, _has_realtime_output_guardrails() should return False
|
||||
assert not streaming._has_realtime_output_guardrails()
|
||||
|
||||
delta = {
|
||||
"type": "response.text.delta",
|
||||
"delta": "This is fine.",
|
||||
}
|
||||
# Should not be buffered
|
||||
if streaming._has_realtime_output_guardrails():
|
||||
streaming._pending_output_text_events.append(json.dumps(delta))
|
||||
else:
|
||||
await client_ws.send_text(json.dumps(delta))
|
||||
|
||||
assert len(streaming._pending_output_text_events) == 0
|
||||
assert len(client_ws.sent) == 1
|
||||
print(f"\n✅ PASS: delta forwarded directly without guardrail")
|
||||
|
||||
|
||||
class TestGuardrailEventHook:
|
||||
"""realtime_output_text is a valid GuardrailEventHooks enum value."""
|
||||
|
||||
def test_hook_exists(self):
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
assert hasattr(GuardrailEventHooks, "realtime_output_text")
|
||||
assert GuardrailEventHooks.realtime_output_text == "realtime_output_text"
|
||||
print(f"\n✅ PASS: GuardrailEventHooks.realtime_output_text = {GuardrailEventHooks.realtime_output_text!r}")
|
||||
Loading…
Add table
Reference in a new issue