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:
Krish Dholakia 2025-12-18 16:46:14 +05:30 • committed by GitHub
parent 26cd2c4473
commit 365762596b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 341 additions and 413 deletions

View file

@ -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"
```

View file

@ -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)

View file

@ -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"

View file

@ -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)

View file

@ -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 (

View file

@ -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),

View file

@ -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)

View file

@ -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", [])

View file

@ -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