diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py index 8eb49602647..d79ad248067 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py @@ -38,6 +38,10 @@ def initialize_guardrail( patterns=litellm_params.patterns, blocked_words=litellm_params.blocked_words, blocked_words_file=litellm_params.blocked_words_file, + pattern_redaction_format=getattr( + litellm_params, "pattern_redaction_format", None + ), + keyword_redaction_tag=getattr(litellm_params, "keyword_redaction_tag", None), event_hook=litellm_params.mode, # type: ignore default_on=litellm_params.default_on or False, categories=getattr(litellm_params, "categories", None), 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 bb079ea6580..0b35bbde18a 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 @@ -4,7 +4,7 @@ Tests for the Content Filter Guardrail import os import sys -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch import pytest @@ -17,11 +17,15 @@ from fastapi import HTTPException from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import ( ContentFilterGuardrail, ) +from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( + initialize_guardrail, +) from litellm.types.guardrails import ( BlockedWord, ContentFilterAction, ContentFilterPattern, GuardrailEventHooks, + LitellmParams, ) from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( ContentFilterCategoryConfig, @@ -76,6 +80,43 @@ class TestContentFilterGuardrail: assert "secret_project" in guardrail.blocked_words assert guardrail.blocked_words["secret_project"][0] == ContentFilterAction.BLOCK + @pytest.mark.asyncio + async def test_initializer_preserves_custom_redaction_formats(self): + litellm_params = LitellmParams( + guardrail="litellm_content_filter", + mode="pre_call", + patterns=[ + ContentFilterPattern( + pattern_type="regex", + pattern=r"1[3-9]\d{9}", + name="phone_number", + action=ContentFilterAction.MASK, + ) + ], + blocked_words=[ + BlockedWord( + keyword="secret_project", + action=ContentFilterAction.MASK, + ) + ], + pattern_redaction_format="<<{pattern_name}>>", + keyword_redaction_tag="<>", + ) + + with patch("litellm.logging_callback_manager.add_litellm_callback"): + guardrail = initialize_guardrail( + litellm_params=litellm_params, + guardrail={"guardrail_name": "test-content-filter"}, + ) + + result = await guardrail.apply_guardrail( + inputs={"texts": ["Call 13912345678 about secret_project"]}, + request_data={}, + input_type="request", + ) + + assert result["texts"] == ["Call <> about <>"] + def test_check_patterns_ssn(self): """ Test SSN pattern detection