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