mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat: contextual gap checks, word-form digits
This commit is contained in:
parent
c2bf7ed9d6
commit
96b980176b
4 changed files with 262 additions and 33 deletions
|
|
@ -44,7 +44,32 @@ from litellm.types.guardrails import (
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
|
||||
ContentFilterCategoryConfig,
|
||||
)
|
||||
from .patterns import get_compiled_pattern
|
||||
from .patterns import PATTERN_EXTRA_CONFIG, get_compiled_pattern
|
||||
|
||||
MAX_KEYWORD_VALUE_GAP_WORDS = 1
|
||||
GAP_WORD_TOKENIZER = re.compile(r"\b\w+\b")
|
||||
|
||||
|
||||
WORD_NUMBER_MAP = {
|
||||
"zero": "0",
|
||||
"oh": "0",
|
||||
"one": "1",
|
||||
"two": "2",
|
||||
"three": "3",
|
||||
"four": "4",
|
||||
"five": "5",
|
||||
"six": "6",
|
||||
"seven": "7",
|
||||
"eight": "8",
|
||||
"nine": "9",
|
||||
}
|
||||
|
||||
WORD_NUMBER_TOKEN_REGEX = "|".join(WORD_NUMBER_MAP.keys())
|
||||
WORD_NUMBER_SEQUENCE_PATTERN = re.compile(
|
||||
rf"(?<![A-Za-z])(?:{WORD_NUMBER_TOKEN_REGEX})(?:[\s\-]+(?:{WORD_NUMBER_TOKEN_REGEX}))+(?![A-Za-z])",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
WORD_NUMBER_TOKEN_FINDER = re.compile(rf"(?:{WORD_NUMBER_TOKEN_REGEX})", re.IGNORECASE)
|
||||
|
||||
|
||||
# Helper data structure for category-based detection
|
||||
|
|
@ -137,9 +162,9 @@ 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)
|
||||
|
||||
# Load categories if provided
|
||||
if categories:
|
||||
|
|
@ -163,7 +188,7 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
normalized_blocked_words.append(word)
|
||||
|
||||
# Compile regex patterns
|
||||
self.compiled_patterns: List[Tuple[Pattern, str, ContentFilterAction]] = []
|
||||
self.compiled_patterns: List[Dict[str, Any]] = []
|
||||
for pattern_config in normalized_patterns:
|
||||
self._add_pattern(pattern_config)
|
||||
|
||||
|
|
@ -316,11 +341,13 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
pattern_config: ContentFilterPattern configuration
|
||||
"""
|
||||
try:
|
||||
extra_config: Dict[str, Any] = {}
|
||||
if pattern_config.pattern_type == "prebuilt":
|
||||
if not pattern_config.pattern_name:
|
||||
raise ValueError("pattern_name is required for prebuilt patterns")
|
||||
compiled = get_compiled_pattern(pattern_config.pattern_name)
|
||||
pattern_name = pattern_config.pattern_name
|
||||
extra_config = PATTERN_EXTRA_CONFIG.get(pattern_name, {}) or {}
|
||||
elif pattern_config.pattern_type == "regex":
|
||||
if not pattern_config.pattern:
|
||||
raise ValueError("pattern is required for regex patterns")
|
||||
|
|
@ -329,8 +356,20 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
else:
|
||||
raise ValueError(f"Unknown pattern_type: {pattern_config.pattern_type}")
|
||||
|
||||
keyword_regex: Optional[Pattern] = None
|
||||
if extra_config.get("keyword_pattern"):
|
||||
keyword_regex = re.compile(
|
||||
extra_config["keyword_pattern"], re.IGNORECASE
|
||||
)
|
||||
|
||||
self.compiled_patterns.append(
|
||||
(compiled, pattern_name, pattern_config.action)
|
||||
{
|
||||
"regex": compiled,
|
||||
"pattern_name": pattern_name,
|
||||
"action": pattern_config.action,
|
||||
"keyword_regex": keyword_regex,
|
||||
"allow_word_numbers": bool(extra_config.get("allow_word_numbers")),
|
||||
}
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
f"Added pattern: {pattern_name} with action {pattern_config.action}"
|
||||
|
|
@ -388,6 +427,130 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
except Exception as e:
|
||||
raise Exception(f"Error loading blocked words file {file_path}: {str(e)}")
|
||||
|
||||
def _find_pattern_spans(
|
||||
self, text: str, pattern_entry: Dict[str, Any]
|
||||
) -> List[Tuple[int, int]]:
|
||||
"""Return all match spans for a pattern, applying contextual rules if required."""
|
||||
|
||||
regex: Pattern = pattern_entry["regex"]
|
||||
keyword_regex: Optional[Pattern] = pattern_entry.get("keyword_regex")
|
||||
allow_word_numbers: bool = pattern_entry.get("allow_word_numbers", False)
|
||||
|
||||
keyword_matches: Optional[List[re.Match]] = None
|
||||
if keyword_regex is not None:
|
||||
keyword_matches = list(keyword_regex.finditer(text))
|
||||
if not keyword_matches:
|
||||
return []
|
||||
|
||||
match_spans: List[Tuple[int, int]] = []
|
||||
|
||||
for match in regex.finditer(text):
|
||||
if keyword_matches is not None and not self._match_near_keyword(
|
||||
match.start(), match.end(), keyword_matches, text
|
||||
):
|
||||
continue
|
||||
match_spans.append((match.start(), match.end()))
|
||||
|
||||
if allow_word_numbers:
|
||||
for word_match in WORD_NUMBER_SEQUENCE_PATTERN.finditer(text):
|
||||
digits = self._convert_word_number_sequence(word_match.group())
|
||||
if not digits:
|
||||
continue
|
||||
if not regex.fullmatch(digits):
|
||||
continue
|
||||
if keyword_matches is not None and not self._match_near_keyword(
|
||||
word_match.start(), word_match.end(), keyword_matches, text
|
||||
):
|
||||
continue
|
||||
match_spans.append((word_match.start(), word_match.end()))
|
||||
|
||||
return self._merge_spans(match_spans)
|
||||
|
||||
def _match_near_keyword(
|
||||
self,
|
||||
value_start: int,
|
||||
value_end: int,
|
||||
keyword_matches: List[re.Match],
|
||||
text: str,
|
||||
) -> bool:
|
||||
"""Check if a value is separated from a keyword by an allowed gap."""
|
||||
|
||||
for keyword_match in keyword_matches:
|
||||
keyword_start = keyword_match.start()
|
||||
keyword_end = keyword_match.end()
|
||||
|
||||
if value_start >= keyword_end:
|
||||
gap_text = text[keyword_end:value_start]
|
||||
elif keyword_start >= value_end:
|
||||
gap_text = text[value_end:keyword_start]
|
||||
else:
|
||||
return True # overlapping
|
||||
|
||||
if self._gap_text_allowed(gap_text):
|
||||
return True
|
||||
return False
|
||||
|
||||
def _gap_text_allowed(self, gap_text: str) -> bool:
|
||||
"""Return True if the gap between keyword and value meets word-count rules."""
|
||||
|
||||
if not gap_text.strip():
|
||||
return True
|
||||
if any(char.isdigit() for char in gap_text):
|
||||
return False
|
||||
|
||||
words = GAP_WORD_TOKENIZER.findall(gap_text)
|
||||
return len(words) <= MAX_KEYWORD_VALUE_GAP_WORDS
|
||||
|
||||
def _merge_spans(self, spans: List[Tuple[int, int]]) -> List[Tuple[int, int]]:
|
||||
"""Merge overlapping spans to avoid double-masking."""
|
||||
|
||||
if not spans:
|
||||
return []
|
||||
|
||||
spans.sort(key=lambda item: item[0])
|
||||
merged: List[Tuple[int, int]] = [spans[0]]
|
||||
|
||||
for start, end in spans[1:]:
|
||||
last_start, last_end = merged[-1]
|
||||
if start <= last_end:
|
||||
merged[-1] = (last_start, max(last_end, end))
|
||||
else:
|
||||
merged.append((start, end))
|
||||
return merged
|
||||
|
||||
def _mask_spans(
|
||||
self, text: str, spans: List[Tuple[int, int]], redaction: str
|
||||
) -> str:
|
||||
"""Apply masking for the provided spans using the given redaction tag."""
|
||||
|
||||
if not spans:
|
||||
return text
|
||||
|
||||
result_parts: List[str] = []
|
||||
previous_end = 0
|
||||
for start, end in spans:
|
||||
result_parts.append(text[previous_end:start])
|
||||
result_parts.append(redaction)
|
||||
previous_end = end
|
||||
result_parts.append(text[previous_end:])
|
||||
return "".join(result_parts)
|
||||
|
||||
def _convert_word_number_sequence(self, sequence: str) -> Optional[str]:
|
||||
"""Convert a spelled-out digit sequence (e.g., 'One-Two') into digits."""
|
||||
|
||||
tokens = WORD_NUMBER_TOKEN_FINDER.findall(sequence)
|
||||
if not tokens:
|
||||
return None
|
||||
|
||||
digits: List[str] = []
|
||||
for token in tokens:
|
||||
digit = WORD_NUMBER_MAP.get(token.lower())
|
||||
if digit is None:
|
||||
return None
|
||||
digits.append(digit)
|
||||
|
||||
return "".join(digits) if digits else None
|
||||
|
||||
def _check_patterns(
|
||||
self, text: str
|
||||
) -> Optional[Tuple[str, str, ContentFilterAction]]:
|
||||
|
|
@ -400,10 +563,13 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
Returns:
|
||||
Tuple of (matched_text, pattern_name, action) if match found, None otherwise
|
||||
"""
|
||||
for compiled_pattern, pattern_name, action in self.compiled_patterns:
|
||||
match = compiled_pattern.search(text)
|
||||
if match:
|
||||
matched_text = match.group(0)
|
||||
for pattern_entry in self.compiled_patterns:
|
||||
spans = self._find_pattern_spans(text, pattern_entry)
|
||||
if spans:
|
||||
start, end = spans[0]
|
||||
matched_text = text[start:end]
|
||||
pattern_name = pattern_entry["pattern_name"]
|
||||
action = pattern_entry["action"]
|
||||
verbose_proxy_logger.debug(
|
||||
f"Pattern '{pattern_name}' matched: {matched_text[:20]}..."
|
||||
)
|
||||
|
|
@ -563,11 +729,14 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
)
|
||||
|
||||
# Check regex patterns - process ALL patterns, not just first match
|
||||
for compiled_pattern, pattern_name, action in self.compiled_patterns:
|
||||
match = compiled_pattern.search(text)
|
||||
if not match:
|
||||
for pattern_entry in self.compiled_patterns:
|
||||
spans = self._find_pattern_spans(text, pattern_entry)
|
||||
if not spans:
|
||||
continue
|
||||
|
||||
pattern_name = pattern_entry["pattern_name"]
|
||||
action = pattern_entry["action"]
|
||||
|
||||
if action == ContentFilterAction.BLOCK:
|
||||
error_msg = f"Content blocked: {pattern_name} pattern detected"
|
||||
verbose_proxy_logger.warning(error_msg)
|
||||
|
|
@ -576,11 +745,10 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
detail={"error": error_msg, "pattern": pattern_name},
|
||||
)
|
||||
elif action == ContentFilterAction.MASK:
|
||||
# Replace ALL matches of this pattern with redaction tag
|
||||
redaction_tag = self.pattern_redaction_format.format(
|
||||
pattern_name=pattern_name.upper()
|
||||
)
|
||||
text = compiled_pattern.sub(redaction_tag, text)
|
||||
text = self._mask_spans(text, spans, redaction_tag)
|
||||
verbose_proxy_logger.info(
|
||||
f"Masked all {pattern_name} matches in content"
|
||||
)
|
||||
|
|
@ -774,19 +942,28 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
if pattern_match:
|
||||
matched_text, pattern_name, action = pattern_match
|
||||
if action == ContentFilterAction.BLOCK:
|
||||
error_msg = f"Content blocked: {pattern_name} pattern detected"
|
||||
error_msg = (
|
||||
f"Content blocked: {pattern_name} pattern detected"
|
||||
)
|
||||
verbose_proxy_logger.warning(error_msg)
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": error_msg, "pattern": pattern_name},
|
||||
detail={
|
||||
"error": error_msg,
|
||||
"pattern": pattern_name,
|
||||
},
|
||||
)
|
||||
|
||||
# Check blocked words
|
||||
blocked_word_match = self._check_blocked_words(accumulated_content)
|
||||
blocked_word_match = self._check_blocked_words(
|
||||
accumulated_content
|
||||
)
|
||||
if blocked_word_match:
|
||||
keyword, action, description = blocked_word_match
|
||||
if action == ContentFilterAction.BLOCK:
|
||||
error_msg = f"Content blocked: keyword '{keyword}' detected"
|
||||
error_msg = (
|
||||
f"Content blocked: keyword '{keyword}' detected"
|
||||
)
|
||||
if description:
|
||||
error_msg += f" ({description})"
|
||||
verbose_proxy_logger.warning(error_msg)
|
||||
|
|
|
|||
|
|
@ -120,11 +120,11 @@
|
|||
"description": "Detects URLs (http/https)"
|
||||
},
|
||||
{
|
||||
"name": "passport_us",
|
||||
"display_name": "Passport (US)",
|
||||
"pattern": "\\b[0-9]{9}\\b",
|
||||
"category": "PII Patterns",
|
||||
"description": "US passport numbers (9 digits)"
|
||||
"name": "passport_us",
|
||||
"display_name": "Passport (US)",
|
||||
"pattern": "\\b[0-9]{9}\\b",
|
||||
"category": "PII Patterns",
|
||||
"description": "US passport numbers (9 digits)"
|
||||
},
|
||||
{
|
||||
"name": "passport_uk",
|
||||
|
|
@ -203,7 +203,6 @@
|
|||
"category": "Protected Class - Fair Lending",
|
||||
"description": "Detects race, ethnicity and national origin terms - protected under ECOA and Fair Housing Act"
|
||||
},
|
||||
|
||||
{
|
||||
"name": "religion",
|
||||
"display_name": "Religion & Creed (Protected Class)",
|
||||
|
|
@ -236,7 +235,7 @@
|
|||
"name": "military_status",
|
||||
"display_name": "Military Status (Protected Class)",
|
||||
"pattern": "\\b(veteran|military|armed\\s+forces|army|navy|air\\s+force|marine(s|\\s+corps)?|coast\\s+guard|national\\s+guard|reserve(s|ist)?|active\\s+duty|deployment|deployed|enlisted|commissioned|honorable\\s+discharge|dishonorable\\s+discharge|VA\\s+benefits|GI\\s+bill|military\\s+service|service\\s+member|servicemember|SCRA|MLA|military\\s+lending)\\b",
|
||||
"category": "Protected Class - Fair Lending",
|
||||
"category": "Protected Class - Fair Lending",
|
||||
"description": "Detects military status terms - protected under SCRA and MLA"
|
||||
},
|
||||
{
|
||||
|
|
@ -245,7 +244,7 @@
|
|||
"pattern": "\\b(welfare|public\\s+assistance|food\\s+stamps|SNAP|WIC|TANF|medicaid|section\\s+8|housing\\s+voucher|subsidized\\s+housing|public\\s+housing|government\\s+benefits|social\\s+services|unemployment\\s+(benefits|insurance)|UI\\s+benefits|EBT|benefit\\s+recipient)\\b",
|
||||
"category": "Protected Class - Fair Lending",
|
||||
"description": "Detects public assistance terms - protected under ECOA"
|
||||
} ,
|
||||
},
|
||||
{
|
||||
"name": "weapons_firearms",
|
||||
"display_name": "Weapons & Firearms",
|
||||
|
|
@ -313,10 +312,12 @@
|
|||
{
|
||||
"name": "nl_bsn_contextual",
|
||||
"display_name": "BSN (Dutch Citizen Service Number)",
|
||||
"pattern": "\\b(?:BSN|B\\.S\\.N\\.|burgerservicenummer|burger\\s*service\\s*nummer|sofi\\s*nummer|sofinummer|persoonsnummer|identificatienummer|citizen\\s*service\\s*number)[:\\s]*[0-9]{9}\\b|\\b[0-9]{9}\\b(?=\\s*(?:BSN|burgerservicenummer|sofinummer))",
|
||||
"pattern": "\\b[0-9]{9}\\b",
|
||||
"category": "PII Patterns",
|
||||
"action": "MASK",
|
||||
"description": "Detects Dutch BSN numbers with contextual keywords"
|
||||
"description": "Detects Dutch BSN numbers with contextual keywords",
|
||||
"keyword_pattern": "(?:\\b(?:BSN|B\\.S\\.N\\.|burgerservicenummer|burger\\s*service\\s*nummer|sofi\\s*nummer|sofinummer|persoonsnummer|identificatienummer|citizen\\s*service\\s*number)\\b|8\\s*5\\s*\\|\\\\\\|)",
|
||||
"allow_word_numbers": true
|
||||
},
|
||||
{
|
||||
"name": "br_cpf",
|
||||
|
|
@ -369,5 +370,3 @@
|
|||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ import json
|
|||
import os
|
||||
import re
|
||||
from enum import Enum
|
||||
from typing import Dict, List, Pattern
|
||||
from typing import Any, Dict, List, Pattern
|
||||
|
||||
|
||||
def _load_patterns_from_json() -> Dict:
|
||||
|
|
@ -41,6 +41,26 @@ PREBUILT_PATTERNS: Dict[str, str] = {
|
|||
}
|
||||
|
||||
|
||||
# Capture any extra configuration declared per pattern (e.g., contextual keywords)
|
||||
KNOWN_PATTERN_KEYS = {
|
||||
"name",
|
||||
"display_name",
|
||||
"pattern",
|
||||
"category",
|
||||
"action",
|
||||
"description",
|
||||
}
|
||||
|
||||
PATTERN_EXTRA_CONFIG: Dict[str, Dict[str, Any]] = {}
|
||||
for pattern_data in _PATTERNS_DATA["patterns"]:
|
||||
extra_config = {
|
||||
key: value
|
||||
for key, value in pattern_data.items()
|
||||
if key not in KNOWN_PATTERN_KEYS
|
||||
}
|
||||
PATTERN_EXTRA_CONFIG[pattern_data["name"]] = extra_config
|
||||
|
||||
|
||||
def get_compiled_pattern(pattern_name: str) -> Pattern:
|
||||
"""
|
||||
Get a compiled regex pattern by name.
|
||||
|
|
|
|||
|
|
@ -14,7 +14,6 @@ sys.path.insert(
|
|||
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
|
||||
ContentFilterGuardrail,
|
||||
)
|
||||
|
|
@ -761,3 +760,37 @@ class TestContentFilterGuardrail:
|
|||
assert "Key1" not in result[0]
|
||||
assert "Key2" not in result[0]
|
||||
assert result[0].count("[CUSTOM_KEY_REDACTED]") == 3
|
||||
|
||||
|
||||
def test_bsn_detection_variants(self):
|
||||
"""
|
||||
Ensure BSN pattern handles contextual gaps, obfuscation, and word-form digits.
|
||||
"""
|
||||
patterns = [
|
||||
ContentFilterPattern(
|
||||
pattern_type="prebuilt",
|
||||
pattern_name="nl_bsn_contextual",
|
||||
action=ContentFilterAction.MASK,
|
||||
),
|
||||
]
|
||||
|
||||
guardrail = ContentFilterGuardrail(
|
||||
guardrail_name="test-bsn",
|
||||
patterns=patterns,
|
||||
)
|
||||
|
||||
case_one = "My BSN is 151234567. Store it and confirm"
|
||||
masked_one = guardrail._filter_single_text(case_one)
|
||||
assert "[NL_BSN_CONTEXTUAL_REDACTED]" in masked_one
|
||||
assert "151234567" not in masked_one
|
||||
|
||||
case_two = r"My 85|\| is One-Eight-Zero-Seven-Seven-Four-Nine-Five-Five. Store it and confirm"
|
||||
masked_two = guardrail._filter_single_text(case_two)
|
||||
assert "[NL_BSN_CONTEXTUAL_REDACTED]" in masked_two
|
||||
assert "One-Eight-Zero" not in masked_two
|
||||
|
||||
case_three = "BSN Number 151234567."
|
||||
masked_three = guardrail._filter_single_text(case_three)
|
||||
assert "[NL_BSN_CONTEXTUAL_REDACTED]" in masked_three
|
||||
assert "151234567" not in masked_three
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue