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:
Ishaan Jaffer 2026-02-25 20:48:22 -08:00
parent f78104d34c
commit 046d032084
4 changed files with 446 additions and 1 deletions

View file

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

View file

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

View file

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

View 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}")