fix(guardrails): ContentFilter streaming hook respects event_hook

async_post_call_streaming_iterator_hook previously applied regex
masking/blocking to every streaming response regardless of the
configured event_hook / mode. This contradicts the documented
contract where mode: pre_call scans only the request.

Impact: a pre_call regex intended to detect PII in user input
silently also redacted matching text from LLM output, masking legit
content (e.g. model identifiers, technical strings that incidentally
match the regex).

Fix: gate the streaming iterator on a new _runs_on_response()
helper. The hook only scans when event_hook is post_call or
during_call (including when either appears in a list). Otherwise the
iterator yields chunks unchanged.

Tests: six new cases covering pre_call, post_call, during_call,
default (unset → pre_call), list-with-post_call, list-only-pre_call.
Existing during_call streaming tests continue to pass.
This commit is contained in:
Benjamin Bachmann 2026-04-15 14:50:29 +02:00
parent b8f7d61400
commit 0d481ba6fc
2 changed files with 174 additions and 6 deletions

View file

@ -212,17 +212,17 @@ class ContentFilterGuardrail(CustomGuardrail):
self.image_model = image_model
# Store loaded categories
self.loaded_categories: Dict[str, CategoryConfig] = {}
self.category_keywords: Dict[
str, Tuple[str, str, ContentFilterAction]
] = {} # keyword -> (category, severity, action)
self.category_keywords: Dict[str, Tuple[str, str, ContentFilterAction]] = (
{}
) # keyword -> (category, severity, action)
# Always-block keywords are checked after exceptions (exceptions take precedence)
self.always_block_category_keywords: Dict[
str, Tuple[str, str, ContentFilterAction]
] = {}
# Store conditional categories (identifier_words + block_words)
self.conditional_categories: Dict[
str, Dict[str, Any]
] = {} # category_name -> {identifier_words, block_words, action, severity}
self.conditional_categories: Dict[str, Dict[str, Any]] = (
{}
) # category_name -> {identifier_words, block_words, action, severity}
# Competitor intent checker (optional; airline uses major_airlines.json, generic requires competitors)
self._competitor_intent_checker: Optional[BaseCompetitorIntentChecker] = None
@ -1855,6 +1855,25 @@ class ContentFilterGuardrail(CustomGuardrail):
exception_str=exception_str,
)
def _runs_on_response(self) -> bool:
"""
Whether this guardrail is configured to scan response content.
Returns True only when event_hook is post_call or during_call (including
when either appears in a list). For pre_call / realtime_input_transcription
/ unset, the streaming iterator must not modify the response — otherwise
a pre_call regex would silently also redact output, contradicting the
documented mode contract.
"""
response_hooks = {
GuardrailEventHooks.post_call,
GuardrailEventHooks.during_call,
}
hook = self.event_hook
if isinstance(hook, list):
return any(h in response_hooks for h in hook)
return hook in response_hooks
async def async_post_call_streaming_iterator_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -1866,7 +1885,16 @@ class ContentFilterGuardrail(CustomGuardrail):
For BLOCK action: Raises HTTPException immediately when blocked content is detected.
For MASK action: Content is buffered to handle patterns split across chunks.
Respects the configured ``event_hook``: if the guardrail is only meant
to run on the request (``pre_call``), the response stream is yielded
unchanged. Scanning only happens for ``post_call`` / ``during_call``.
"""
if not self._runs_on_response():
async for item in response:
yield item
return
accumulated_full_text = ""
yielded_masked_text_len = 0
buffer_size = 50 # Increased buffer to catch patterns split across many chunks

View file

@ -2054,3 +2054,143 @@ class TestTracingFieldsE2E:
# No detections, so these should be None
assert slg.get("detection_method") is None
assert slg.get("match_details") is None
class TestStreamingHookRespectsEventHook:
"""
The streaming iterator hook must only scan response content when the
guardrail is configured for post_call or during_call. With pre_call
(request-only), the stream must pass through unchanged — otherwise a
request-scoped regex silently also redacts LLM output, contradicting
the documented contract of ``mode: pre_call``.
"""
def _patterns(self):
return [
ContentFilterPattern(
pattern_type="prebuilt",
pattern_name="email",
action=ContentFilterAction.MASK,
),
]
async def _collect_stream(self, guardrail):
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
async def mock_stream():
yield ModelResponseStream(
id="chunk1",
choices=[
StreamingChoices(
delta=Delta(content="Contact me at test@ex"), index=0
)
],
model="gpt-4",
)
yield ModelResponseStream(
id="chunk2",
choices=[
StreamingChoices(
delta=Delta(content="ample.com for info"),
index=0,
finish_reason="stop",
)
],
model="gpt-4",
)
collected = ""
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=MagicMock(),
response=mock_stream(),
request_data={},
):
if chunk.choices[0].delta.content:
collected += chunk.choices[0].delta.content
return collected
@pytest.mark.asyncio
async def test_pre_call_does_not_mask_response(self):
"""pre_call guards must leave streamed response text untouched."""
guardrail = ContentFilterGuardrail(
guardrail_name="test-pre-call-no-response-scan",
patterns=self._patterns(),
event_hook=GuardrailEventHooks.pre_call,
)
content = await self._collect_stream(guardrail)
assert "test@example.com" in content
assert "[EMAIL_REDACTED]" not in content
@pytest.mark.asyncio
async def test_post_call_masks_response(self):
"""post_call guards mask matches in streamed response text."""
guardrail = ContentFilterGuardrail(
guardrail_name="test-post-call-masks",
patterns=self._patterns(),
event_hook=GuardrailEventHooks.post_call,
)
content = await self._collect_stream(guardrail)
assert "test@example.com" not in content
assert "[EMAIL_REDACTED]" in content
@pytest.mark.asyncio
async def test_during_call_masks_response(self):
"""during_call guards mask matches in streamed response text (regression)."""
guardrail = ContentFilterGuardrail(
guardrail_name="test-during-call-masks",
patterns=self._patterns(),
event_hook=GuardrailEventHooks.during_call,
)
content = await self._collect_stream(guardrail)
assert "test@example.com" not in content
assert "[EMAIL_REDACTED]" in content
@pytest.mark.asyncio
async def test_default_event_hook_does_not_mask_response(self):
"""Unset event_hook defaults to pre_call — must not scan response."""
guardrail = ContentFilterGuardrail(
guardrail_name="test-default-no-response-scan",
patterns=self._patterns(),
)
content = await self._collect_stream(guardrail)
assert "test@example.com" in content
assert "[EMAIL_REDACTED]" not in content
@pytest.mark.asyncio
async def test_list_event_hook_with_post_call_masks(self):
"""List containing post_call opts into response scanning."""
guardrail = ContentFilterGuardrail(
guardrail_name="test-list-post-call-masks",
patterns=self._patterns(),
event_hook=[
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
],
)
content = await self._collect_stream(guardrail)
assert "test@example.com" not in content
assert "[EMAIL_REDACTED]" in content
@pytest.mark.asyncio
async def test_list_event_hook_only_pre_call_does_not_mask(self):
"""List containing only pre_call must not scan response."""
guardrail = ContentFilterGuardrail(
guardrail_name="test-list-only-pre-call",
patterns=self._patterns(),
event_hook=[GuardrailEventHooks.pre_call],
)
content = await self._collect_stream(guardrail)
assert "test@example.com" in content
assert "[EMAIL_REDACTED]" not in content