mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Guardrails - LiteLLM Content Filter - add support for running content filters on images (#18044)
* feat(litellm_content_filter.py): add support for content filtering categories make it easy for proxy admin to prevent messages about violence, self harm or illegal weapons going through litellm * feat: initial commit adding bias detection allows admin to block inappropriate content about sexual orientation, etc. * refactor: simplify content_filter.py use a more exhaustive set of keywords, instead of guessing at potential phrases user can use * feat(content_filter.py): add new denied topics for in-built content filter guardrails allow user to automatically block content relating to certain categories from being sent to the LLML * refactor(content-filter): document new params to litellm content filter * feat(ui/): litellm content filter - select content categories on ui * docs: update documentation * docs(litellm_content_filter.md): document new content filters * feat: initial commit adding support for inappropriate images via litellm content filter * feat(content_filter.py): support blocking images containing blocked content prevent images which contain disallowed content from being sent to the llm api * docs(litellm_content_filter.md): document new image capabilities of litellm_content_filter * fix: fix expected error code
This commit is contained in:
parent
26cd2c4473
commit
365762596b
9 changed files with 341 additions and 413 deletions
|
|
@ -392,6 +392,85 @@ for chunk in response:
|
|||
# Emails automatically masked in real-time
|
||||
```
|
||||
|
||||
## Image Content Filtering
|
||||
|
||||
Content filter can analyze images by generating descriptions and applying filters to the text descriptions.
|
||||
|
||||
:::warning
|
||||
|
||||
This can introduce significant latency to the request - depending on the speed of the vision-capable model.
|
||||
|
||||
This is because, each request containing images will be sent to the vision-capable model to generate a description.
|
||||
|
||||
:::
|
||||
|
||||
### Configuration
|
||||
|
||||
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
model_list:
|
||||
- model_name: gpt-4-vision
|
||||
litellm_params:
|
||||
model: openai/gpt-4-vision-preview
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "image-filter"
|
||||
litellm_params:
|
||||
guardrail: litellm_content_filter
|
||||
mode: "pre_call"
|
||||
image_model: "gpt-4-vision" # value is `model_name` of the vision-capable model
|
||||
|
||||
# Apply same filters to image descriptions
|
||||
categories:
|
||||
- category: "harmful_violence"
|
||||
enabled: true
|
||||
action: "BLOCK"
|
||||
severity_threshold: "medium"
|
||||
|
||||
patterns:
|
||||
- pattern_type: "prebuilt"
|
||||
pattern_name: "email"
|
||||
action: "MASK"
|
||||
```
|
||||
|
||||
### How It Works
|
||||
|
||||
1. Image is sent to the vision model to generate a text description
|
||||
2. Content filters are applied to the description
|
||||
3. If harmful content is detected, request is blocked with context about the image
|
||||
|
||||
**Example:**
|
||||
|
||||
```python
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(
|
||||
api_key="sk-1234",
|
||||
base_url="http://localhost:4000"
|
||||
)
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="gpt-4-vision",
|
||||
messages=[{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "What's in this image?"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}
|
||||
]
|
||||
}],
|
||||
extra_body={"guardrails": ["image-filter"]}
|
||||
)
|
||||
```
|
||||
|
||||
If the image description contains filtered content, you'll get:
|
||||
|
||||
```json
|
||||
{
|
||||
"error": "Content blocked: harmful_violence category keyword 'weapon' detected (severity: high) (Image description): The image shows..."
|
||||
}
|
||||
```
|
||||
|
||||
## Customizing Redaction Tags
|
||||
|
||||
When using the `MASK` action, sensitive content is replaced with redaction tags. You can customize how these tags appear.
|
||||
|
|
@ -703,211 +782,6 @@ guardrails:
|
|||
### 5. Compliance
|
||||
Ensure regulatory compliance by filtering sensitive data types:
|
||||
|
||||
```yaml
|
||||
patterns:
|
||||
- pattern_type: "prebuilt"
|
||||
pattern_name: "visa"
|
||||
action: "BLOCK"
|
||||
- pattern_type: "prebuilt"
|
||||
pattern_name: "us_ssn"
|
||||
action: "BLOCK"
|
||||
```
|
||||
|
||||
## Best Practices for Bias Detection
|
||||
|
||||
### Choosing the Right Severity Threshold
|
||||
|
||||
Bias detection requires careful tuning to avoid blocking legitimate content:
|
||||
|
||||
**High Threshold (Recommended for most use cases)**
|
||||
- Blocks only explicit discriminatory language
|
||||
- Lower false positives
|
||||
- Allows nuanced discussions about identity, diversity, and social issues
|
||||
- Good for: Public-facing applications, education, research
|
||||
|
||||
**Medium Threshold**
|
||||
- Blocks stereotypes and generalizations
|
||||
- Balanced approach
|
||||
- May catch some edge cases in legitimate discourse
|
||||
- Good for: Consumer applications, internal tools, moderated environments
|
||||
|
||||
**Low Threshold**
|
||||
- Strictest filtering
|
||||
- Blocks even borderline language
|
||||
- Higher false positives but maximum safety
|
||||
- Good for: Youth-focused applications, highly controlled environments
|
||||
|
||||
### Testing Your Bias Filters
|
||||
|
||||
Always test with realistic use cases:
|
||||
|
||||
```yaml
|
||||
# Test legitimate discussions (should NOT be blocked)
|
||||
- "Our company has a gender diversity initiative"
|
||||
- "Research shows racial disparities in healthcare"
|
||||
- "We support LGBTQ+ rights and equality"
|
||||
- "Religious freedom is a fundamental right"
|
||||
|
||||
# Test discriminatory content (SHOULD be blocked)
|
||||
- "Women are too emotional to lead"
|
||||
- "All [group] are [negative stereotype]"
|
||||
- "Being gay is unnatural"
|
||||
- "[Religious group] are all extremists"
|
||||
```
|
||||
|
||||
### Monitoring and Iteration
|
||||
|
||||
1. **Log blocked requests** to review false positives
|
||||
2. **Add exceptions** for legitimate terms in your domain
|
||||
3. **Adjust severity thresholds** based on your audience
|
||||
4. **Use custom category files** for domain-specific bias patterns
|
||||
|
||||
### Cultural and Linguistic Considerations
|
||||
|
||||
The prebuilt categories focus on English and common patterns. For other languages or cultural contexts:
|
||||
|
||||
1. Create custom category files with region-specific terms
|
||||
2. Consult with native speakers and cultural experts
|
||||
3. Include local slurs and stereotypes
|
||||
4. Adjust severity based on regional norms
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### False Positives with Bias Detection
|
||||
|
||||
**Issue:** Legitimate discussions about diversity, identity, or social issues are being blocked
|
||||
|
||||
**Solutions:**
|
||||
|
||||
1. **Raise severity threshold:**
|
||||
```yaml
|
||||
categories:
|
||||
- category: "bias_racial"
|
||||
severity_threshold: "high" # Only explicit discrimination
|
||||
```
|
||||
|
||||
2. **Add domain-specific exceptions:**
|
||||
```yaml
|
||||
# my_custom_bias_gender.yaml
|
||||
exceptions:
|
||||
- "gender pay gap"
|
||||
- "gender diversity"
|
||||
- "women in tech"
|
||||
- "gender equality"
|
||||
- "dei initiative"
|
||||
- "inclusion program"
|
||||
```
|
||||
|
||||
3. **Review what was blocked:**
|
||||
Check error details to understand what triggered the block:
|
||||
```json
|
||||
{
|
||||
"error": "Content blocked: bias_gender category keyword 'women' detected (severity: medium)",
|
||||
"category": "bias_gender",
|
||||
"keyword": "women",
|
||||
"severity": "medium"
|
||||
}
|
||||
```
|
||||
|
||||
If this is a false positive, add "women in leadership" or other legitimate phrases to exceptions in your custom category file.
|
||||
|
||||
### False Positives with Categories
|
||||
|
||||
**Issue:** Legitimate content is being blocked by category filters
|
||||
|
||||
**Solution 1:** Adjust severity threshold to only block high-severity items:
|
||||
```yaml
|
||||
categories:
|
||||
- category: "harmful_violence"
|
||||
enabled: true
|
||||
action: "BLOCK"
|
||||
severity_threshold: "high" # Only block explicit harmful content
|
||||
```
|
||||
|
||||
**Solution 2:** Add exceptions to your custom category file:
|
||||
```yaml
|
||||
# my_custom_violence.yaml
|
||||
exceptions:
|
||||
- "crime statistics"
|
||||
- "documentary"
|
||||
- "news report"
|
||||
- "historical context"
|
||||
```
|
||||
|
||||
**Solution 3:** Use a custom category file with your own curated keyword list:
|
||||
```yaml
|
||||
categories:
|
||||
- category: "harmful_violence"
|
||||
enabled: true
|
||||
action: "BLOCK"
|
||||
category_file: "/path/to/my_violence_keywords.yaml"
|
||||
```
|
||||
|
||||
### Category Not Loading
|
||||
|
||||
**Issue:** Category is not being applied
|
||||
|
||||
**Checklist:**
|
||||
1. Verify category is enabled: `enabled: true`
|
||||
2. Check category name matches file: `harmful_self_harm.yaml`
|
||||
3. Check file exists in `litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/categories/`
|
||||
4. Review logs for loading errors: `litellm --config config.yaml --detailed_debug`
|
||||
|
||||
### Keyword Not Matching
|
||||
|
||||
**Issue:** Expected keyword is not being detected
|
||||
|
||||
**Solutions:**
|
||||
|
||||
1. **For single words:** Ensure the keyword appears as a whole word. The system uses word boundary matching, so "men" won't match "recommend".
|
||||
|
||||
2. **For multi-word phrases:** Use the exact phrase as it should appear. Multi-word keywords are matched as substrings (case-insensitive), so "harm myself" will match "I want to harm myself" or "harming myself".
|
||||
|
||||
3. **Check exceptions:** If your keyword is in the exceptions list, it won't be detected. Review the category file's exceptions section.
|
||||
|
||||
4. **Verify severity threshold:** Lower severity keywords won't match if your threshold is set too high. For example, if a keyword has `severity: "low"` but your `severity_threshold: "high"`, it won't match.
|
||||
|
||||
### Too Many False Negatives
|
||||
|
||||
**Issue:** Harmful content is not being caught
|
||||
|
||||
**Solution 1:** Lower severity threshold:
|
||||
```yaml
|
||||
severity_threshold: "low" # Catch more but may increase false positives
|
||||
```
|
||||
|
||||
**Solution 2:** Add custom keywords for your specific use case:
|
||||
```yaml
|
||||
categories:
|
||||
- category: "harmful_violence"
|
||||
enabled: true
|
||||
category_file: "/path/to/enhanced_violence.yaml"
|
||||
|
||||
# In enhanced_violence.yaml, add domain-specific keywords
|
||||
keywords:
|
||||
- keyword: "your specific harmful phrase"
|
||||
severity: "high"
|
||||
- keyword: "another harmful term"
|
||||
severity: "medium"
|
||||
```
|
||||
|
||||
### Pattern Not Matching
|
||||
|
||||
**Issue:** Regex pattern isn't detecting expected content
|
||||
|
||||
**Solution:** Test your regex pattern:
|
||||
```python
|
||||
import re
|
||||
pattern = r'\b[A-Z]{3}-\d{4}\b'
|
||||
test_text = "Employee ID: ABC-1234"
|
||||
print(re.search(pattern, test_text)) # Should match
|
||||
```
|
||||
|
||||
### Multiple Pattern Matches
|
||||
|
||||
**Issue:** Text contains multiple sensitive patterns
|
||||
|
||||
**Solution:** Guardrail checks in this order: categories (keywords), regex patterns, then blocked words. Order by priority:
|
||||
```yaml
|
||||
# Categories checked first (high priority)
|
||||
# Category keywords are matched first
|
||||
|
|
@ -917,13 +791,12 @@ categories:
|
|||
|
||||
# Then regex patterns
|
||||
patterns:
|
||||
- pattern_type: "prebuilt"
|
||||
pattern_name: "visa"
|
||||
action: "BLOCK"
|
||||
- pattern_type: "prebuilt"
|
||||
pattern_name: "us_ssn"
|
||||
action: "BLOCK"
|
||||
|
||||
# Then simple blocked keywords
|
||||
blocked_words:
|
||||
- keyword: "confidential"
|
||||
action: "MASK"
|
||||
```
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -162,6 +162,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
url = image_url.get("url")
|
||||
if url:
|
||||
images_to_check.append(url)
|
||||
elif isinstance(image_url, str):
|
||||
images_to_check.append(image_url)
|
||||
|
||||
# Extract tool calls (typically in assistant messages)
|
||||
tool_calls = message.get("tool_calls", None)
|
||||
|
|
|
|||
|
|
@ -3,6 +3,10 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/gpt-3.5-turbo
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
- model_name: gpt-4o
|
||||
litellm_params:
|
||||
model: openai/gpt-4o
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
- model_name: claude-sonnet-4-5-20250929
|
||||
litellm_params:
|
||||
model: anthropic/claude-sonnet-4-5-20250929
|
||||
|
|
@ -28,11 +32,7 @@ guardrails:
|
|||
mode: "pre_call"
|
||||
default_on: true
|
||||
# Model configuration
|
||||
guardrail_model:
|
||||
- model: "gpt-4o" # From model_list
|
||||
supported_multimodal_content: # Supported content types = images, documents, audio, video, text
|
||||
- images
|
||||
- documents
|
||||
image_model: "claude-sonnet-4-5-20250929"
|
||||
|
||||
categories:
|
||||
- category: "harmful_self_harm"
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import litellm
|
||||
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
|
||||
|
|
@ -7,10 +7,15 @@ from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_fil
|
|||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm import Router
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
||||
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
|
||||
def initialize_guardrail(
|
||||
litellm_params: "LitellmParams",
|
||||
guardrail: "Guardrail",
|
||||
llm_router: Optional["Router"] = None,
|
||||
):
|
||||
"""
|
||||
Initialize the Content Filter Guardrail.
|
||||
|
||||
|
|
@ -22,6 +27,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
Initialized ContentFilterGuardrail instance
|
||||
"""
|
||||
guardrail_name = guardrail.get("guardrail_name")
|
||||
|
||||
if not guardrail_name:
|
||||
raise ValueError("Content Filter: guardrail_name is required")
|
||||
|
||||
|
|
@ -34,6 +40,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
default_on=litellm_params.default_on or False,
|
||||
categories=getattr(litellm_params, "categories", None),
|
||||
severity_threshold=getattr(litellm_params, "severity_threshold", "medium"),
|
||||
llm_router=llm_router,
|
||||
image_model=getattr(litellm_params, "image_model", None),
|
||||
)
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(content_filter_guardrail)
|
||||
|
|
|
|||
|
|
@ -5,12 +5,12 @@ This guardrail provides regex pattern matching and keyword filtering
|
|||
to detect and block/mask sensitive content.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import re
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
|
|
@ -18,11 +18,13 @@ from typing import (
|
|||
Pattern,
|
||||
Tuple,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
import yaml
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm import Router
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
||||
|
|
@ -94,6 +96,8 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
keyword_redaction_tag: Optional[str] = None,
|
||||
categories: Optional[List[ContentFilterCategoryConfig]] = None,
|
||||
severity_threshold: str = "medium",
|
||||
llm_router: Optional[Router] = None,
|
||||
image_model: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
|
|
@ -129,7 +133,8 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
)
|
||||
self.keyword_redaction_tag = keyword_redaction_tag or self.KEYWORD_REDACTION_STR
|
||||
self.severity_threshold = severity_threshold
|
||||
|
||||
self.llm_router = llm_router
|
||||
self.image_model = image_model
|
||||
# Store loaded categories
|
||||
self.loaded_categories: Dict[str, CategoryConfig] = {}
|
||||
self.category_keywords: Dict[str, Tuple[str, str, ContentFilterAction]] = (
|
||||
|
|
@ -211,8 +216,9 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
|
||||
enabled = cat_config.get("enabled", True)
|
||||
action = cat_config.get("action")
|
||||
severity_threshold = cat_config.get(
|
||||
"severity_threshold", self.severity_threshold
|
||||
severity_threshold = (
|
||||
cat_config.get("severity_threshold", self.severity_threshold)
|
||||
or self.severity_threshold
|
||||
)
|
||||
custom_file = cat_config.get("category_file")
|
||||
|
||||
|
|
@ -502,6 +508,121 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
return (keyword, action, description)
|
||||
return None
|
||||
|
||||
def _filter_single_text(self, text: str) -> str:
|
||||
"""
|
||||
Apply all content filtering checks to a single text.
|
||||
|
||||
This method performs:
|
||||
1. Category keyword checks
|
||||
2. Regex pattern checks
|
||||
3. Blocked word checks
|
||||
|
||||
Args:
|
||||
text: Text to filter
|
||||
|
||||
Returns:
|
||||
Filtered text (with masking applied if action is MASK)
|
||||
|
||||
Raises:
|
||||
HTTPException: If sensitive content is detected and action is BLOCK
|
||||
"""
|
||||
# Collect all exceptions from loaded categories
|
||||
all_exceptions = []
|
||||
for category in self.loaded_categories.values():
|
||||
all_exceptions.extend(category.exceptions)
|
||||
|
||||
# Check category keywords
|
||||
category_keyword_match = self._check_category_keywords(text, all_exceptions)
|
||||
if category_keyword_match:
|
||||
keyword, category, severity, action = category_keyword_match
|
||||
if action == ContentFilterAction.BLOCK:
|
||||
error_msg = (
|
||||
f"Content blocked: {category} category keyword '{keyword}' detected "
|
||||
f"(severity: {severity})"
|
||||
)
|
||||
verbose_proxy_logger.warning(error_msg)
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": error_msg,
|
||||
"category": category,
|
||||
"keyword": keyword,
|
||||
"severity": severity,
|
||||
},
|
||||
)
|
||||
elif action == ContentFilterAction.MASK:
|
||||
# Replace keyword with redaction tag
|
||||
text = re.sub(
|
||||
re.escape(keyword),
|
||||
self.keyword_redaction_tag,
|
||||
text,
|
||||
flags=re.IGNORECASE,
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
f"Masked category keyword '{keyword}' from {category} (severity: {severity})"
|
||||
)
|
||||
|
||||
# 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:
|
||||
continue
|
||||
|
||||
if action == ContentFilterAction.BLOCK:
|
||||
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},
|
||||
)
|
||||
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)
|
||||
verbose_proxy_logger.info(
|
||||
f"Masked all {pattern_name} matches in content"
|
||||
)
|
||||
|
||||
# Check blocked words - iterate through ALL blocked words
|
||||
# to ensure all matching keywords are processed, not just the first one
|
||||
text_lower = text.lower()
|
||||
for keyword, (action, description) in self.blocked_words.items():
|
||||
if keyword not in text_lower:
|
||||
continue
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Blocked word '{keyword}' found with action {action}"
|
||||
)
|
||||
|
||||
if action == ContentFilterAction.BLOCK:
|
||||
error_msg = f"Content blocked: keyword '{keyword}' detected"
|
||||
if description:
|
||||
error_msg += f" ({description})"
|
||||
verbose_proxy_logger.warning(error_msg)
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": error_msg,
|
||||
"keyword": keyword,
|
||||
"description": description,
|
||||
},
|
||||
)
|
||||
elif action == ContentFilterAction.MASK:
|
||||
# Replace keyword with redaction tag (case-insensitive)
|
||||
text = re.sub(
|
||||
re.escape(keyword),
|
||||
self.keyword_redaction_tag,
|
||||
text,
|
||||
flags=re.IGNORECASE,
|
||||
)
|
||||
# Update text_lower after masking to avoid re-matching
|
||||
text_lower = text.lower()
|
||||
verbose_proxy_logger.info(f"Masked keyword '{keyword}' in content")
|
||||
|
||||
return text
|
||||
|
||||
def _mask_content(self, text: str, pattern_name: str) -> str:
|
||||
"""
|
||||
Mask sensitive content in text.
|
||||
|
|
@ -544,110 +665,72 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
HTTPException: If sensitive content is detected and action is BLOCK
|
||||
"""
|
||||
texts = inputs.get("texts", [])
|
||||
images = inputs.get("images", [])
|
||||
if images and self.image_model and self.llm_router:
|
||||
tasks = []
|
||||
for image in images:
|
||||
task = self.llm_router.acompletion(
|
||||
model=self.image_model,
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": "Describe the image in detail.",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": image}},
|
||||
],
|
||||
},
|
||||
],
|
||||
stream=False,
|
||||
)
|
||||
tasks.append(task)
|
||||
|
||||
responses = await asyncio.gather(*tasks)
|
||||
descriptions = []
|
||||
for response in responses:
|
||||
if response.choices[0].message.content:
|
||||
image_description = response.choices[0].message.content
|
||||
verbose_proxy_logger.debug(
|
||||
f"Image description: {image_description}"
|
||||
)
|
||||
descriptions.append(image_description)
|
||||
else:
|
||||
verbose_proxy_logger.warning("No image description found")
|
||||
|
||||
# Apply content filtering to image descriptions
|
||||
verbose_proxy_logger.debug(
|
||||
f"ContentFilterGuardrail: Applying guardrail to {len(descriptions)} image description(s)"
|
||||
)
|
||||
for description in descriptions:
|
||||
# This will raise HTTPException if BLOCK action is triggered
|
||||
try:
|
||||
self._filter_single_text(description)
|
||||
except HTTPException as e:
|
||||
# e.detail can be a string or dict
|
||||
if isinstance(e.detail, dict) and "error" in e.detail:
|
||||
detail_dict = cast(Dict[str, Any], e.detail)
|
||||
detail_dict["error"] = (
|
||||
detail_dict["error"]
|
||||
+ " (Image description): "
|
||||
+ description
|
||||
)
|
||||
elif isinstance(e.detail, str):
|
||||
e.detail = e.detail + " (Image description): " + description
|
||||
else:
|
||||
e.detail = (
|
||||
"Content blocked: Image description detected" + description
|
||||
)
|
||||
raise e
|
||||
verbose_proxy_logger.debug(
|
||||
f"ContentFilterGuardrail: Applying guardrail to {len(texts)} text(s)"
|
||||
)
|
||||
|
||||
processed_texts = []
|
||||
|
||||
for text in texts:
|
||||
# Collect all exceptions from loaded categories
|
||||
all_exceptions = []
|
||||
for category in self.loaded_categories.values():
|
||||
all_exceptions.extend(category.exceptions)
|
||||
|
||||
# Check category keywords
|
||||
category_keyword_match = self._check_category_keywords(text, all_exceptions)
|
||||
if category_keyword_match:
|
||||
keyword, category, severity, action = category_keyword_match
|
||||
if action == ContentFilterAction.BLOCK:
|
||||
error_msg = (
|
||||
f"Content blocked: {category} category keyword '{keyword}' detected "
|
||||
f"(severity: {severity})"
|
||||
)
|
||||
verbose_proxy_logger.warning(error_msg)
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": error_msg,
|
||||
"category": category,
|
||||
"keyword": keyword,
|
||||
"severity": severity,
|
||||
},
|
||||
)
|
||||
elif action == ContentFilterAction.MASK:
|
||||
# Replace keyword with redaction tag
|
||||
text = re.sub(
|
||||
re.escape(keyword),
|
||||
self.keyword_redaction_tag,
|
||||
text,
|
||||
flags=re.IGNORECASE,
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
f"Masked category keyword '{keyword}' from {category} (severity: {severity})"
|
||||
)
|
||||
|
||||
# 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:
|
||||
continue
|
||||
|
||||
if action == ContentFilterAction.BLOCK:
|
||||
error_msg = f"Content blocked: {pattern_name} pattern detected"
|
||||
verbose_proxy_logger.warning(error_msg)
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
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)
|
||||
verbose_proxy_logger.info(
|
||||
f"Masked all {pattern_name} matches in content"
|
||||
)
|
||||
|
||||
# Check blocked words - iterate through ALL blocked words
|
||||
# to ensure all matching keywords are processed, not just the first one
|
||||
text_lower = text.lower()
|
||||
for keyword, (action, description) in self.blocked_words.items():
|
||||
if keyword not in text_lower:
|
||||
continue
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Blocked word '{keyword}' found with action {action}"
|
||||
)
|
||||
|
||||
if action == ContentFilterAction.BLOCK:
|
||||
error_msg = f"Content blocked: keyword '{keyword}' detected"
|
||||
if description:
|
||||
error_msg += f" ({description})"
|
||||
verbose_proxy_logger.warning(error_msg)
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": error_msg,
|
||||
"keyword": keyword,
|
||||
"description": description,
|
||||
},
|
||||
)
|
||||
elif action == ContentFilterAction.MASK:
|
||||
# Replace keyword with redaction tag (case-insensitive)
|
||||
text = re.sub(
|
||||
re.escape(keyword),
|
||||
self.keyword_redaction_tag,
|
||||
text,
|
||||
flags=re.IGNORECASE,
|
||||
)
|
||||
# Update text_lower after masking to avoid re-matching
|
||||
text_lower = text.lower()
|
||||
verbose_proxy_logger.info(f"Masked keyword '{keyword}' in content")
|
||||
|
||||
processed_texts.append(text)
|
||||
filtered_text = self._filter_single_text(text)
|
||||
processed_texts.append(filtered_text)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"ContentFilterGuardrail: Guardrail applied successfully"
|
||||
|
|
@ -655,68 +738,6 @@ class ContentFilterGuardrail(CustomGuardrail):
|
|||
inputs["texts"] = processed_texts
|
||||
return inputs
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
request_data: dict,
|
||||
) -> AsyncGenerator[ModelResponseStream, None]:
|
||||
"""
|
||||
Streaming hook to check each chunk as it's yielded.
|
||||
|
||||
This implementation checks each chunk individually and yields it immediately,
|
||||
allowing for low-latency streaming with content filtering.
|
||||
|
||||
Args:
|
||||
user_api_key_dict: User API key authentication
|
||||
response: Async generator of response chunks
|
||||
request_data: Original request data
|
||||
|
||||
Yields:
|
||||
Checked and potentially masked chunks
|
||||
|
||||
Raises:
|
||||
HTTPException: If chunk content should be blocked
|
||||
"""
|
||||
verbose_proxy_logger.debug(
|
||||
"ContentFilterGuardrail: Running streaming check (per-chunk mode)"
|
||||
)
|
||||
|
||||
# Process each chunk individually
|
||||
async for chunk in response:
|
||||
if isinstance(chunk, ModelResponseStream):
|
||||
for choice in chunk.choices:
|
||||
if hasattr(choice, "delta") and choice.delta.content:
|
||||
if isinstance(choice.delta.content, str):
|
||||
# Check the chunk content using apply_guardrail
|
||||
try:
|
||||
guardrailed_inputs = await self.apply_guardrail(
|
||||
inputs={"texts": [choice.delta.content]},
|
||||
input_type="response",
|
||||
request_data=request_data,
|
||||
)
|
||||
processed_texts = guardrailed_inputs.get("texts", [])
|
||||
processed_content = (
|
||||
processed_texts[0]
|
||||
if processed_texts
|
||||
else choice.delta.content
|
||||
)
|
||||
if processed_content != choice.delta.content:
|
||||
choice.delta.content = processed_content
|
||||
verbose_proxy_logger.debug(
|
||||
"ContentFilterGuardrail: Modified streaming chunk"
|
||||
)
|
||||
except HTTPException as e:
|
||||
# If content should be blocked, raise immediately
|
||||
verbose_proxy_logger.warning(
|
||||
f"ContentFilterGuardrail: Blocked streaming chunk: {e.detail}"
|
||||
)
|
||||
raise
|
||||
|
||||
yield chunk
|
||||
|
||||
verbose_proxy_logger.debug("ContentFilterGuardrail: Streaming check completed")
|
||||
|
||||
@staticmethod
|
||||
def get_config_model():
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from datetime import datetime, timezone
|
|||
from typing import Any, Dict, List, Optional, Type, cast
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -399,6 +400,7 @@ class InMemoryGuardrailHandler:
|
|||
self,
|
||||
guardrail: Guardrail,
|
||||
config_file_path: Optional[str] = None,
|
||||
llm_router: Optional["Router"] = None,
|
||||
) -> Optional[Guardrail]:
|
||||
"""
|
||||
Initialize a guardrail from a dictionary and add it to the litellm callback manager
|
||||
|
|
@ -447,7 +449,16 @@ class InMemoryGuardrailHandler:
|
|||
initializer = guardrail_initializer_registry.get(guardrail_type)
|
||||
|
||||
if initializer:
|
||||
custom_guardrail_callback = initializer(litellm_params, guardrail)
|
||||
# Try to call with llm_router first, fall back to without if it fails
|
||||
import inspect
|
||||
|
||||
sig = inspect.signature(initializer)
|
||||
if "llm_router" in sig.parameters:
|
||||
custom_guardrail_callback = initializer(
|
||||
litellm_params, guardrail, llm_router # type: ignore
|
||||
)
|
||||
else:
|
||||
custom_guardrail_callback = initializer(litellm_params, guardrail)
|
||||
elif isinstance(guardrail_type, str) and "." in guardrail_type:
|
||||
custom_guardrail_callback = self.initialize_custom_guardrail(
|
||||
guardrail=cast(dict, guardrail),
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
from typing import Dict, List, Optional, cast
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_utils.callback_utils import initialize_callbacks_on_proxy
|
||||
|
||||
|
|
@ -18,6 +19,7 @@ Map guardrail_name: <pre_call>, <post_call>, during_call
|
|||
def init_guardrails_v2(
|
||||
all_guardrails: List[Dict],
|
||||
config_file_path: Optional[str] = None,
|
||||
llm_router: Optional[Router] = None,
|
||||
):
|
||||
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
|
||||
|
||||
|
|
@ -27,6 +29,7 @@ def init_guardrails_v2(
|
|||
initialized_guardrail = IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(
|
||||
guardrail=cast(Guardrail, guardrail),
|
||||
config_file_path=config_file_path,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
if initialized_guardrail:
|
||||
guardrail_list.append(initialized_guardrail)
|
||||
|
|
|
|||
|
|
@ -297,9 +297,7 @@ from litellm.proxy.management_endpoints.customer_endpoints import (
|
|||
from litellm.proxy.management_endpoints.internal_user_endpoints import (
|
||||
router as internal_user_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import (
|
||||
user_update,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import user_update
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
delete_verification_tokens,
|
||||
duration_in_seconds,
|
||||
|
|
@ -353,9 +351,7 @@ from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router
|
|||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
router as openai_files_router,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
set_files_config,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import set_files_config
|
||||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
passthrough_endpoint_router,
|
||||
)
|
||||
|
|
@ -450,9 +446,7 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
|
|||
LiteLLM_UpperboundKeyGenerateParams,
|
||||
)
|
||||
from litellm.types.realtime import RealtimeQueryParams
|
||||
from litellm.types.router import (
|
||||
DeploymentTypedDict,
|
||||
)
|
||||
from litellm.types.router import DeploymentTypedDict
|
||||
from litellm.types.router import ModelInfo as RouterModelInfo
|
||||
from litellm.types.router import (
|
||||
RouterGeneralSettings,
|
||||
|
|
@ -565,9 +559,7 @@ else:
|
|||
ui_link = f"{server_root_path}/ui"
|
||||
fallback_login_link = f"{server_root_path}/fallback/login"
|
||||
model_hub_link = f"{server_root_path}/ui/model_hub_table"
|
||||
ui_message = (
|
||||
f"👉 [```LiteLLM Admin Panel on /ui```]({ui_link}). Create, Edit Keys with SSO. Having issues? Try [```Fallback Login```]({fallback_login_link})"
|
||||
)
|
||||
ui_message = f"👉 [```LiteLLM Admin Panel on /ui```]({ui_link}). Create, Edit Keys with SSO. Having issues? Try [```Fallback Login```]({fallback_login_link})"
|
||||
ui_message += "\n\n💸 [```LiteLLM Model Cost Map```](https://models.litellm.ai/)."
|
||||
|
||||
ui_message += f"\n\n🔎 [```LiteLLM Model Hub```]({model_hub_link}). See available models on the proxy. [**Docs**](https://docs.litellm.ai/docs/proxy/ai_hub)"
|
||||
|
|
@ -649,10 +641,10 @@ async def _initialize_shared_aiohttp_session():
|
|||
connector_kwargs["limit"] = AIOHTTP_CONNECTOR_LIMIT
|
||||
if AIOHTTP_CONNECTOR_LIMIT_PER_HOST > 0:
|
||||
connector_kwargs["limit_per_host"] = AIOHTTP_CONNECTOR_LIMIT_PER_HOST
|
||||
|
||||
|
||||
connector = TCPConnector(**connector_kwargs)
|
||||
session = ClientSession(connector=connector)
|
||||
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"SESSION REUSE: Created shared aiohttp session for connection pooling (ID: {id(session)}, "
|
||||
f"limit={AIOHTTP_CONNECTOR_LIMIT}, limit_per_host={AIOHTTP_CONNECTOR_LIMIT_PER_HOST})"
|
||||
|
|
@ -1151,6 +1143,7 @@ if docs_url != "/" and root_redirect_url is not None:
|
|||
async def root_redirect():
|
||||
return RedirectResponse(url=root_redirect_url) # type: ignore[arg-type]
|
||||
|
||||
|
||||
from typing import Dict
|
||||
|
||||
user_api_base = None
|
||||
|
|
@ -1734,7 +1727,7 @@ async def _run_background_health_check():
|
|||
else:
|
||||
# Use a system identifier for background health checks
|
||||
checked_by = "background_health_check"
|
||||
|
||||
|
||||
start_time = time_module.time()
|
||||
asyncio.create_task(
|
||||
_save_background_health_checks_to_db(
|
||||
|
|
@ -2425,7 +2418,9 @@ class ProxyConfig:
|
|||
# Initialize global polling via cache settings
|
||||
global polling_via_cache_enabled, polling_cache_ttl
|
||||
background_mode = value.get("background_mode", {})
|
||||
polling_via_cache_enabled = background_mode.get("polling_via_cache", False)
|
||||
polling_via_cache_enabled = background_mode.get(
|
||||
"polling_via_cache", False
|
||||
)
|
||||
polling_cache_ttl = background_mode.get("ttl", 3600)
|
||||
verbose_proxy_logger.debug(
|
||||
f"{blue_color_code} Initialized polling via cache: enabled={polling_via_cache_enabled}, ttl={polling_cache_ttl}{reset_color_code}"
|
||||
|
|
@ -2720,7 +2715,9 @@ class ProxyConfig:
|
|||
guardrails_v2 = config.get("guardrails", None)
|
||||
if guardrails_v2:
|
||||
init_guardrails_v2(
|
||||
all_guardrails=guardrails_v2, config_file_path=config_file_path
|
||||
all_guardrails=guardrails_v2,
|
||||
config_file_path=config_file_path,
|
||||
llm_router=router,
|
||||
)
|
||||
|
||||
## Prompt settings
|
||||
|
|
@ -4522,7 +4519,7 @@ class ProxyStartupEvent:
|
|||
### MONITOR SPEND LOGS QUEUE (queue-size-based job) ###
|
||||
if general_settings.get("disable_spend_logs", False) is False:
|
||||
from litellm.proxy.utils import _monitor_spend_logs_queue
|
||||
|
||||
|
||||
# Start background task to monitor spend logs queue size
|
||||
asyncio.create_task(
|
||||
_monitor_spend_logs_queue(
|
||||
|
|
@ -5383,7 +5380,9 @@ async def embeddings( # noqa: PLR0915
|
|||
# check if provider accept list of tokens as input - e.g. for langchain integration
|
||||
if llm_router is not None and data.get("model") in router_model_names:
|
||||
# Use router's O(1) lookup instead of O(N) iteration through llm_model_list
|
||||
deployment = llm_router.get_deployment_by_model_group_name(model_group_name=data["model"])
|
||||
deployment = llm_router.get_deployment_by_model_group_name(
|
||||
model_group_name=data["model"]
|
||||
)
|
||||
if deployment is not None:
|
||||
litellm_params = deployment.get("litellm_params", {}) or {}
|
||||
litellm_model = litellm_params.get("model", "")
|
||||
|
|
@ -5663,10 +5662,12 @@ async def audio_speech(
|
|||
if "gemini" in request_model_lower and (
|
||||
"tts" in request_model_lower or "preview-tts" in request_model_lower
|
||||
):
|
||||
media_type = "audio/wav" # Gemini TTS returns WAV format after conversion
|
||||
media_type = (
|
||||
"audio/wav" # Gemini TTS returns WAV format after conversion
|
||||
)
|
||||
|
||||
return StreamingResponse(
|
||||
_audio_speech_chunk_generator(response), # type: ignore[arg-type]
|
||||
_audio_speech_chunk_generator(response), # type: ignore[arg-type]
|
||||
media_type=media_type,
|
||||
headers=custom_headers, # type: ignore
|
||||
)
|
||||
|
|
@ -8638,6 +8639,7 @@ async def login_v2(request: Request): # noqa: PLR0915
|
|||
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
)
|
||||
|
||||
|
||||
@app.get("/onboarding/get_token", include_in_schema=False)
|
||||
async def onboarding(invite_link: str, request: Request):
|
||||
"""
|
||||
|
|
@ -9730,11 +9732,11 @@ async def get_config(): # noqa: PLR0915
|
|||
_litellm_settings = config_data.get("litellm_settings", {})
|
||||
_general_settings = config_data.get("general_settings", {})
|
||||
environment_variables = config_data.get("environment_variables", {})
|
||||
|
||||
|
||||
_success_callbacks = _litellm_settings.get("success_callback", [])
|
||||
_failure_callbacks = _litellm_settings.get("failure_callback", [])
|
||||
_success_and_failure_callbacks = _litellm_settings.get("callbacks", [])
|
||||
|
||||
|
||||
_data_to_return = []
|
||||
"""
|
||||
[
|
||||
|
|
@ -9750,15 +9752,23 @@ async def get_config(): # noqa: PLR0915
|
|||
]
|
||||
|
||||
"""
|
||||
|
||||
|
||||
for _callback in _success_callbacks:
|
||||
_data_to_return.append(process_callback(_callback, "success", environment_variables))
|
||||
|
||||
_data_to_return.append(
|
||||
process_callback(_callback, "success", environment_variables)
|
||||
)
|
||||
|
||||
for _callback in _failure_callbacks:
|
||||
_data_to_return.append(process_callback(_callback, "failure", environment_variables))
|
||||
|
||||
_data_to_return.append(
|
||||
process_callback(_callback, "failure", environment_variables)
|
||||
)
|
||||
|
||||
for _callback in _success_and_failure_callbacks:
|
||||
_data_to_return.append(process_callback(_callback, "success_and_failure", environment_variables))
|
||||
_data_to_return.append(
|
||||
process_callback(
|
||||
_callback, "success_and_failure", environment_variables
|
||||
)
|
||||
)
|
||||
|
||||
# Check if slack alerting is on
|
||||
_alerting = _general_settings.get("alerting", [])
|
||||
|
|
|
|||
|
|
@ -196,7 +196,7 @@ class TestContentFilterGuardrail:
|
|||
input_type="request",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "us_ssn" in str(exc_info.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -501,7 +501,7 @@ class TestContentFilterGuardrail:
|
|||
):
|
||||
pass
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "us_ssn" in str(exc_info.value.detail)
|
||||
|
||||
def test_init_with_plain_dicts(self):
|
||||
|
|
@ -669,7 +669,7 @@ class TestContentFilterGuardrail:
|
|||
input_type="request",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "danger_word" in str(exc_info.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue