fix(realtime): guardrails with pre_call/post_call mode now work on realtime WebSocket (#22161)

* fix(realtime): guardrails with pre_call/post_call mode now work on realtime WebSocket; return error directly to consumer

* fix(realtime guardrails): address code review feedback

- Restore session.update injection for audio/VAD path, but only when
  realtime_input_transcription guardrails are configured (not pre_call).
  Forward session.created to the client first so no error arrives before
  the client sees the session.
- Change _swallow_next_response_create bool to int counter so consecutive
  blocked items are handled correctly.
- Extract _build_litellm_metadata() helper to eliminate duplicated
  metadata-building logic across OpenAI/Azure/XAI provider branches.
- Plumb litellm_metadata and user_api_key_dict to Azure and XAI handlers
  so guardrails work for those providers too.
- Add tests for session.update injection, no-inject for pre_call-only,
  and consecutive-block counter.

* simplify: remove response.create swallowing after guardrail block

When an item is blocked, the error event is already sent to the client.
The subsequent response.create from the client is fine to forward through —
the LLM may respond to previous context which is acceptable behavior.
Removing the swallow counter eliminates unnecessary state tracking.
This commit is contained in:
Ishaan Jaff 2026-02-25 23:43:13 -08:00 • committed by Sameer Kankute
parent 1ed4023700
commit 587ad3ccb1
5 changed files with 321 additions and 89 deletions

View file

@ -43,6 +43,7 @@ class RealTimeStreaming:
provider_config: Optional[BaseRealtimeConfig] = None,
model: str = "",
user_api_key_dict: Optional[Any] = None,
request_data: Optional[Dict] = None,
):
self.websocket = websocket
self.backend_ws = backend_ws
@ -68,6 +69,7 @@ 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
self.request_data: Dict = request_data or {}
def _should_store_message(
self,
@ -231,14 +233,40 @@ class RealTimeStreaming:
await self.backend_ws.send(message)
def _has_realtime_guardrails(self) -> bool:
"""Return True if any callback is registered for realtime_input_transcription."""
"""Return True if any callback is registered for realtime guardrail event types."""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
_realtime_event_types = [
GuardrailEventHooks.realtime_input_transcription,
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
]
return any(
isinstance(cb, CustomGuardrail)
and any(
cb.should_run_guardrail(
data=self.request_data,
event_type=et,
)
for et in _realtime_event_types
)
for cb in litellm.callbacks
)
def _has_audio_transcription_guardrails(self) -> bool:
"""Return True if any callback needs to run on audio transcriptions (VAD path).
When this returns True, we inject a session.update to disable the LLM's
auto-response so the guardrail can gate it first.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
return any(
isinstance(cb, CustomGuardrail)
and cb.should_run_guardrail(
data={},
data=self.request_data,
event_type=GuardrailEventHooks.realtime_input_transcription,
)
for cb in litellm.callbacks
@ -258,17 +286,25 @@ class RealTimeStreaming:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
_realtime_event_types = [
GuardrailEventHooks.realtime_input_transcription,
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
]
_check_data = {**self.request_data, "transcript": transcript}
_already_run: set = set()
for callback in litellm.callbacks:
if not isinstance(callback, CustomGuardrail):
continue
if (
callback.should_run_guardrail(
data={"transcript": transcript},
event_type=GuardrailEventHooks.realtime_input_transcription,
)
is not True
if id(callback) in _already_run:
continue
if not any(
callback.should_run_guardrail(data=_check_data, event_type=et)
for et in _realtime_event_types
):
continue
_already_run.add(id(callback))
try:
await callback.apply_guardrail(
inputs={"texts": [transcript], "images": []},
@ -293,20 +329,15 @@ class RealTimeStreaming:
safe_msg = str(detail)
else:
safe_msg = str(e) or "I'm sorry, that request was blocked by the content filter."
# Cancel any in-flight response before speaking the warning.
# This handles the race where create_response fired before we could intercept.
await self._send_to_backend(json.dumps({"type": "response.cancel"}))
# Ask the model to speak the warning — TTS audio plays naturally in the client
await self._send_to_backend(
# Return the error directly to the WebSocket consumer.
await self.websocket.send_text(
json.dumps(
{
"type": "response.create",
"response": {
"modalities": ["text", "audio"],
"instructions": (
f"Say exactly and only: \"{safe_msg}\". "
"Do not add anything else."
),
"type": "error",
"error": {
"type": "guardrail_violation",
"message": safe_msg,
"code": "content_policy_violation",
},
}
)
@ -348,25 +379,25 @@ class RealTimeStreaming:
if isinstance(transformed_response, list)
else [transformed_response]
)
for event in events:
## GUARDRAIL: inject create_response=false on session.created
if isinstance(event, dict) and event.get("type") == "session.created":
if self._has_realtime_guardrails():
await self._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {
"turn_detection": {
"type": "server_vad",
"create_response": False,
}
},
}
)
)
for event in events:
event_str = json.dumps(event)
## For audio/VAD guardrail path: forward session.created first, then inject.
if (
isinstance(event, dict)
and event.get("type") == "session.created"
and self._has_audio_transcription_guardrails()
):
self.store_message(event_str)
await self.websocket.send_text(event_str)
await self._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {"turn_detection": {"create_response": False}},
}
)
)
continue
## GUARDRAIL: run on transcription events in provider_config path too
if (
isinstance(event, dict)
@ -397,27 +428,26 @@ class RealTimeStreaming:
try:
event_obj = json.loads(raw_response)
if event_obj.get("type") == "session.created":
# If any realtime guardrails are registered, proactively
# set create_response=false so the LLM never auto-responds
# before our guardrail has a chance to run.
if self._has_realtime_guardrails():
await self._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {
"turn_detection": {
"type": "server_vad",
"create_response": False,
}
},
}
)
)
verbose_logger.debug(
"[realtime guardrail] injected create_response=false into session"
# For audio/VAD guardrail path: once the session is ready, tell the backend
# not to auto-respond after VAD detects end-of-speech. We send the
# session.created to the client FIRST so the client is always in sync, then
# inject the session.update so a potential error from the backend doesn't
# arrive before the client sees session.created.
if (
event_obj.get("type") == "session.created"
and self._has_audio_transcription_guardrails()
):
self.store_message(raw_response)
await self.websocket.send_text(raw_response)
await self._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {"turn_detection": {"create_response": False}},
}
)
)
return True
if (
event_obj.get("type")

View file

@ -6,13 +6,13 @@ This requires websockets, and is currently only supported on LiteLLM Proxy.
from typing import Any, Optional, cast
from litellm._logging import verbose_proxy_logger
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ....litellm_core_utils.realtime_streaming import RealTimeStreaming
from ....llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
from ..azure import AzureChatCompletion
from litellm._logging import verbose_proxy_logger
# BACKEND_WS_URL = "ws://localhost:8080/v1/realtime?model=gpt-4o-realtime-preview-2024-10-01"
@ -77,6 +77,8 @@ class AzureOpenAIRealtime(AzureChatCompletion):
client: Optional[Any] = None,
timeout: Optional[float] = None,
realtime_protocol: Optional[str] = None,
user_api_key_dict: Optional[Any] = None,
litellm_metadata: Optional[dict] = None,
):
import websockets
from websockets.asyncio.client import ClientConnection
@ -101,7 +103,11 @@ class AzureOpenAIRealtime(AzureChatCompletion):
ssl=ssl_context,
) as backend_ws:
realtime_streaming = RealTimeStreaming(
websocket, cast(ClientConnection, backend_ws), logging_obj
websocket,
cast(ClientConnection, backend_ws),
logging_obj,
user_api_key_dict=user_api_key_dict,
request_data={"litellm_metadata": litellm_metadata or {}},
)
await realtime_streaming.bidirectional_forward()

View file

@ -99,6 +99,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
timeout: Optional[float] = None,
query_params: Optional[RealtimeQueryParams] = None,
user_api_key_dict: Optional[Any] = None,
litellm_metadata: Optional[dict] = None,
**kwargs: Any,
):
import websockets
@ -142,6 +143,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
cast(ClientConnection, backend_ws),
logging_obj,
user_api_key_dict=user_api_key_dict,
request_data={"litellm_metadata": litellm_metadata or {}},
)
await realtime_streaming.bidirectional_forward()

View file

@ -32,6 +32,15 @@ vertex_llm_base = VertexBase()
base_llm_http_handler = BaseLLMHTTPHandler()
def _build_litellm_metadata(kwargs: dict) -> dict:
"""Build the litellm_metadata dict for guardrail checking (internal only, not forwarded to provider)."""
metadata: dict = {**(kwargs.get("litellm_metadata") or {})}
guardrails = (kwargs.get("metadata") or {}).get("guardrails") or kwargs.get("guardrails") or []
if guardrails:
metadata["guardrails"] = guardrails
return metadata
@wrapper_client
async def _arealtime(
model: str,
@ -134,6 +143,8 @@ async def _arealtime(
timeout=timeout,
logging_obj=litellm_logging_obj,
realtime_protocol=realtime_protocol,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
)
elif _custom_llm_provider == "openai":
api_base = (
@ -160,6 +171,7 @@ async def _arealtime(
timeout=timeout,
query_params=query_params,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
)
elif _custom_llm_provider == "bedrock":
# Extract AWS parameters from kwargs
@ -217,6 +229,8 @@ async def _arealtime(
client=None,
timeout=timeout,
query_params=query_params,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
)
elif _custom_llm_provider == "vertex_ai":
vertex_credentials = (

View file

@ -416,32 +416,32 @@ async def test_realtime_guardrail_blocks_prompt_injection():
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
await streaming.backend_to_client_send_messages()
# ASSERT 1: no bare response.create was sent to backend (injection blocked).
# The only response.create allowed is the warning one (has "instructions" field).
# ASSERT 1: no response.create was sent to backend (injection blocked).
sent_to_backend = [
json.loads(c.args[0])
for c in backend_ws.send.call_args_list
if c.args
]
bare_response_creates = [
response_creates = [
e for e in sent_to_backend
if e.get("type") == "response.create"
and "instructions" not in e.get("response", {})
]
assert len(bare_response_creates) == 0, (
f"Guardrail should prevent bare response.create for injected content, "
f"but got: {bare_response_creates}"
assert len(response_creates) == 0, (
f"Guardrail should prevent response.create for injected content, "
f"but got: {response_creates}"
)
# ASSERT 2: warning response.create was sent to backend (to speak the block message)
warning_creates = [
e for e in sent_to_backend
if e.get("type") == "response.create"
and "instructions" in e.get("response", {})
# ASSERT 2: error event was sent directly to the client WebSocket
sent_to_client = [
json.loads(c.args[0]) for c in client_ws.send_text.call_args_list
if c.args
]
assert len(warning_creates) > 0, (
f"Backend should receive a response.create with warning instructions, "
f"but got: {sent_to_backend}"
error_events = [e for e in sent_to_client if e.get("type") == "error"]
assert len(error_events) == 1, (
f"Expected one error event sent to client, got: {sent_to_client}"
)
assert error_events[0]["error"]["type"] == "guardrail_violation", (
f"Expected guardrail_violation error type, got: {error_events[0]}"
)
litellm.callbacks = [] # cleanup
@ -514,11 +514,91 @@ async def test_realtime_guardrail_allows_clean_transcript():
@pytest.mark.asyncio
async def test_realtime_session_created_injects_create_response_false():
async def test_realtime_text_input_guardrail_blocks_and_returns_error():
"""
Test that when session.created arrives from the backend and realtime guardrails
are registered, the proxy injects a session.update with create_response=False
so the LLM never auto-responds before the guardrail runs.
Test that when conversation.item.create arrives with text that triggers a guardrail,
the proxy blocks it (doesn't forward to backend) and returns an error event directly
to the client WebSocket.
"""
from fastapi import HTTPException
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
class BlockingGuardrail(CustomGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
texts = inputs.get("texts", [])
for text in texts:
if "@" in text:
raise HTTPException(
status_code=403,
detail={"error": "email address detected"},
)
return inputs
guardrail = BlockingGuardrail(
guardrail_name="email-blocker",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
litellm.callbacks = [guardrail]
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
backend_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None))
logging_obj = MagicMock()
logging_obj.pre_call = MagicMock()
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
item_create_msg = json.dumps({
"type": "conversation.item.create",
"item": {
"role": "user",
"content": [{"type": "input_text", "text": "My email is test@example.com"}],
},
})
# Simulate the client sending a conversation.item.create with an email
client_ws.receive_text = AsyncMock(
side_effect=[
item_create_msg,
Exception("connection closed"), # stop the loop
]
)
await streaming.client_ack_messages()
# ASSERT: error event was sent to client
assert client_ws.send_text.called, "Expected error to be sent to client websocket"
sent_texts = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list]
error_events = [e for e in sent_texts if e.get("type") == "error"]
assert len(error_events) == 1, f"Expected one error event, got: {sent_texts}"
assert error_events[0]["error"]["type"] == "guardrail_violation"
# ASSERT: blocked item was NOT forwarded to the backend
sent_to_backend = [c.args[0] for c in backend_ws.send.call_args_list if c.args]
forwarded_items = [
json.loads(m) for m in sent_to_backend
if isinstance(m, str) and json.loads(m).get("type") == "conversation.item.create"
]
assert len(forwarded_items) == 0, (
f"Blocked item should not be forwarded to backend, got: {forwarded_items}"
)
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_text_input_guardrail_uses_pre_call_mode():
"""
Test that _has_realtime_guardrails returns True for a guardrail configured with
pre_call mode (not just realtime_input_transcription).
"""
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -529,7 +609,46 @@ async def test_realtime_session_created_injects_create_response_false():
return inputs
guardrail = DummyGuardrail(
guardrail_name="dummy",
guardrail_name="pre-call-guardrail",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
litellm.callbacks = [guardrail]
client_ws = MagicMock()
backend_ws = MagicMock()
logging_obj = MagicMock()
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
assert streaming._has_realtime_guardrails() is True, (
"pre_call guardrail should be recognized as a realtime guardrail"
)
# pre_call guardrail should NOT trigger the audio/VAD session.update injection
assert streaming._has_audio_transcription_guardrails() is False, (
"pre_call guardrail should not trigger audio transcription guardrail path"
)
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_session_created_injects_session_update_for_audio_guardrail():
"""
Test that when an audio transcription guardrail is configured, a session.created
event from the backend triggers a session.update injection (create_response: false)
AFTER forwarding session.created to the client. This prevents the LLM from
auto-responding before the guardrail can run on the transcript.
"""
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
class AudioGuardrail(CustomGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
return inputs
guardrail = AudioGuardrail(
guardrail_name="audio-guardrail",
event_hook=GuardrailEventHooks.realtime_input_transcription,
default_on=True,
)
@ -538,34 +657,95 @@ async def test_realtime_session_created_injects_create_response_false():
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
session_created_event = json.dumps({"type": "session.created"}).encode()
session_created_event = json.dumps(
{"type": "session.created", "session": {"id": "sess_abc"}}
).encode()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[
session_created_event,
ConnectionClosed(None, None),
]
side_effect=[session_created_event, ConnectionClosed(None, None)]
)
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
await streaming.backend_to_client_send_messages()
# ASSERT: proxy injected session.update with create_response=False to backend
# session.created must be forwarded to the client
sent_to_client = [
json.loads(c.args[0]) for c in client_ws.send_text.call_args_list if c.args
]
session_created_events = [e for e in sent_to_client if e.get("type") == "session.created"]
assert len(session_created_events) == 1, (
f"session.created should be forwarded to client, got: {sent_to_client}"
)
# session.update must be sent to the backend AFTER session.created was forwarded
sent_to_backend = [
json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args
]
session_updates = [e for e in sent_to_backend if e.get("type") == "session.update"]
assert len(session_updates) == 1, (
f"Expected proxy to inject session.update, got: {sent_to_backend}"
f"Expected one session.update injected to backend, got: {sent_to_backend}"
)
td = session_updates[0]["session"]["turn_detection"]
assert td["create_response"] is False, (
f"Expected create_response=False, got: {td}"
assert session_updates[0]["session"]["turn_detection"]["create_response"] is False
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_session_created_no_injection_for_pre_call_only():
"""
Test that when only a pre_call guardrail is configured (no audio transcription),
session.created does NOT trigger the session.update injection.
"""
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
class PreCallGuardrail(CustomGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
return inputs
guardrail = PreCallGuardrail(
guardrail_name="pre-call-only",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
litellm.callbacks = [guardrail]
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
session_created_event = json.dumps(
{"type": "session.created", "session": {"id": "sess_xyz"}}
).encode()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[session_created_event, ConnectionClosed(None, None)]
)
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
await streaming.backend_to_client_send_messages()
# No session.update should be injected
sent_to_backend = [
json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args
]
session_updates = [e for e in sent_to_backend if e.get("type") == "session.update"]
assert len(session_updates) == 0, (
f"pre_call guardrail should NOT inject session.update, got: {sent_to_backend}"
)
litellm.callbacks = [] # cleanup