From 0d481ba6fc76e236864591bb1d3c027019520163 Mon Sep 17 00:00:00 2001 From: Benjamin Bachmann Date: Wed, 15 Apr 2026 14:50:29 +0200 Subject: [PATCH] fix(guardrails): ContentFilter streaming hook respects event_hook MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- .../litellm_content_filter/content_filter.py | 40 ++++- .../content_filter/test_content_filter.py | 140 ++++++++++++++++++ 2 files changed, 174 insertions(+), 6 deletions(-) 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 e4da1c1ae77..ae7599c312c 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 @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py index a1d3eb152bb..0c477e56f66 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py @@ -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