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 ae7599c312c..7d988fe3f4b 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 @@ -1859,20 +1859,36 @@ class ContentFilterGuardrail(CustomGuardrail): """ 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. + Returns True when ``event_hook`` selects a response-side hook + (``post_call`` or ``during_call``). 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. + + Supports all three accepted ``event_hook`` shapes: + - single ``GuardrailEventHooks`` value + - ``list`` of ``GuardrailEventHooks`` (opt-in if any member matches) + - ``Mode`` (tag-routed): opt-in if any tag value or the default + resolves to a response hook, since the configured routing could + send matching requests to a response-side mode """ - response_hooks = { - GuardrailEventHooks.post_call, - GuardrailEventHooks.during_call, + response_hook_values = { + GuardrailEventHooks.post_call.value, + GuardrailEventHooks.during_call.value, } hook = self.event_hook if isinstance(hook, list): - return any(h in response_hooks for h in hook) - return hook in response_hooks + return any(getattr(h, "value", h) in response_hook_values for h in hook) + if isinstance(hook, Mode): + candidates: List[Union[str, List[str]]] = list(hook.tags.values()) + if hook.default is not None: + candidates.append(hook.default) + for value in candidates: + items = value if isinstance(value, list) else [value] + if any(v in response_hook_values for v in items): + return True + return False + return getattr(hook, "value", hook) in response_hook_values async def async_post_call_streaming_iterator_hook( self, 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 0c477e56f66..d9e05313c12 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 @@ -22,6 +22,7 @@ from litellm.types.guardrails import ( ContentFilterAction, ContentFilterPattern, GuardrailEventHooks, + Mode, ) from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( ContentFilterCategoryConfig, @@ -2194,3 +2195,54 @@ class TestStreamingHookRespectsEventHook: assert "test@example.com" in content assert "[EMAIL_REDACTED]" not in content + + @pytest.mark.asyncio + async def test_mode_tag_routing_with_post_call_masks(self): + """Mode tag routing: any tag resolving to post_call opts into response scanning.""" + guardrail = ContentFilterGuardrail( + guardrail_name="test-mode-tag-post-call", + patterns=self._patterns(), + event_hook=Mode( + tags={"sensitive": "post_call", "internal": "pre_call"}, + default="pre_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_mode_tag_routing_all_pre_call_does_not_mask(self): + """Mode tag routing: all tags pre_call and default pre_call must not scan response.""" + guardrail = ContentFilterGuardrail( + guardrail_name="test-mode-tag-all-pre-call", + patterns=self._patterns(), + event_hook=Mode( + tags={"a": "pre_call", "b": "pre_call"}, + default="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_mode_default_post_call_masks(self): + """Mode with default=post_call opts into response scanning.""" + guardrail = ContentFilterGuardrail( + guardrail_name="test-mode-default-post-call", + patterns=self._patterns(), + event_hook=Mode( + tags={"a": "pre_call"}, + default="post_call", + ), + ) + + content = await self._collect_stream(guardrail) + + assert "test@example.com" not in content + assert "[EMAIL_REDACTED]" in content