diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index d15d23f8eea..d6181bd8527 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py index 6c83833c668..e9dfa96e371 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py @@ -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, diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index b9c99eaabfb..37326bcd8cd 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -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): diff --git a/tests/test_litellm/realtime/test_output_guardrail.py b/tests/test_litellm/realtime/test_output_guardrail.py new file mode 100644 index 00000000000..5e5878f3597 --- /dev/null +++ b/tests/test_litellm/realtime/test_output_guardrail.py @@ -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}")