diff --git a/docs/my-website/docs/proxy/guardrails/litellm_content_filter.md b/docs/my-website/docs/proxy/guardrails/litellm_content_filter.md index 20bed9a0488..5ba39ba35ed 100644 --- a/docs/my-website/docs/proxy/guardrails/litellm_content_filter.md +++ b/docs/my-website/docs/proxy/guardrails/litellm_content_filter.md @@ -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" ``` + diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index 87e3eece14c..3ab938da520 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -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) diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index e8074c82512..0bdee099720 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -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" diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py index d7f06cf6a53..ec6fc53d3c8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py index 344383af2bb..a74f5088e01 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py @@ -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 ( diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 2d5f07dbf6e..fe53fe3b32b 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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), diff --git a/litellm/proxy/guardrails/init_guardrails.py b/litellm/proxy/guardrails/init_guardrails.py index aeef7040c4b..2ab9ce2ae11 100644 --- a/litellm/proxy/guardrails/init_guardrails.py +++ b/litellm/proxy/guardrails/init_guardrails.py @@ -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: , , 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) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0a1cfda1502..0e591c5e10f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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", []) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py index 0f8b73ee640..3e82c8ed0af 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py @@ -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