mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
feat: add auth_filter guardrail type and implementation
- Add AUTH_FILTER to SupportedGuardrailIntegrations enum - Add post_auth_check to GuardrailEventHooks enum - Implement AuthFilterGuardrail extending CustomCodeGuardrail - Supports custom authentication enrichment/validation post-DB lookup - Reuses custom_code sandbox with HTTP primitives - Includes hot-reload support
This commit is contained in:
parent
a9058bb584
commit
36e591ae21
4 changed files with 746 additions and 18 deletions
328
litellm/proxy/guardrails/guardrail_hooks/auth_filter/README.md
Normal file
328
litellm/proxy/guardrails/guardrail_hooks/auth_filter/README.md
Normal file
|
|
@ -0,0 +1,328 @@
|
|||
# Auth Filter Guardrail
|
||||
|
||||
The auth_filter guardrail enables custom authentication enrichment and validation logic that runs **after** LiteLLM's standard authentication completes.
|
||||
|
||||
## When to Use
|
||||
|
||||
Use auth_filter when you need to:
|
||||
- Validate authenticated users against external systems
|
||||
- Enrich the auth object with additional metadata
|
||||
- Apply organization-specific access rules
|
||||
- Call external APIs for compliance checks
|
||||
- Block requests based on custom logic
|
||||
|
||||
## Execution Flow
|
||||
|
||||
```
|
||||
1. Standard LiteLLM Auth (DB lookup) → UserAPIKeyAuth
|
||||
2. Auth Filter Guardrails Execute → Enrich/Validate/Block
|
||||
3. Continue to LLM Request
|
||||
```
|
||||
|
||||
The auth_filter receives the **validated** UserAPIKeyAuth object from standard authentication, along with the original request and API key.
|
||||
|
||||
## Configuration
|
||||
|
||||
### Basic Example
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
- guardrail_name: "org-validator"
|
||||
litellm_params:
|
||||
guardrail: "auth_filter"
|
||||
mode: "post_auth_check"
|
||||
custom_code: |
|
||||
def auth_filter(request, api_key, user_api_key_auth):
|
||||
# Access organization ID from authenticated user
|
||||
org_id = user_api_key_auth.organization_id
|
||||
|
||||
if org_id == "restricted-org":
|
||||
return block("Organization access restricted")
|
||||
|
||||
return allow()
|
||||
```
|
||||
|
||||
### Async Example with External API
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
- guardrail_name: "compliance-check"
|
||||
litellm_params:
|
||||
guardrail: "auth_filter"
|
||||
mode: "post_auth_check"
|
||||
custom_code: |
|
||||
async def auth_filter(request, api_key, user_api_key_auth):
|
||||
# Call external compliance API
|
||||
org_id = user_api_key_auth.organization_id
|
||||
|
||||
response = await http_post(
|
||||
"https://compliance.internal/validate",
|
||||
body={"org_id": org_id, "endpoint": request.url.path}
|
||||
)
|
||||
|
||||
if not response["success"]:
|
||||
return block("Compliance validation failed")
|
||||
|
||||
# Enrich with compliance session
|
||||
user_api_key_auth.metadata["compliance_session"] = response["body"]["session_id"]
|
||||
return modify(user_api_key_auth=user_api_key_auth)
|
||||
```
|
||||
|
||||
### Reading Request Headers
|
||||
|
||||
```yaml
|
||||
custom_code: |
|
||||
def auth_filter(request, api_key, user_api_key_auth):
|
||||
# Extract custom headers
|
||||
department = request.headers.get("X-Department")
|
||||
environment = request.headers.get("X-Environment", "prod")
|
||||
|
||||
# Apply department-specific rules
|
||||
if department == "healthcare" and environment == "prod":
|
||||
# Additional validation for healthcare in production
|
||||
if not user_api_key_auth.metadata.get("hipaa_certified"):
|
||||
return block("HIPAA certification required for healthcare")
|
||||
|
||||
return allow()
|
||||
```
|
||||
|
||||
## Function Signature
|
||||
|
||||
```python
|
||||
def auth_filter(request, api_key, user_api_key_auth):
|
||||
"""
|
||||
Args:
|
||||
request (Request): FastAPI Request object with headers, query params, body
|
||||
api_key (str): The API key used for authentication
|
||||
user_api_key_auth (UserAPIKeyAuth): Validated auth object from standard auth
|
||||
|
||||
Returns:
|
||||
- allow() - Continue without modification
|
||||
- block(reason) - Reject with 403 error
|
||||
- modify(user_api_key_auth=obj) - Return enriched auth object
|
||||
"""
|
||||
```
|
||||
|
||||
### Async Support
|
||||
|
||||
Use `async def` when making external HTTP requests:
|
||||
|
||||
```python
|
||||
async def auth_filter(request, api_key, user_api_key_auth):
|
||||
response = await http_get("https://api.example.com/validate")
|
||||
# ...
|
||||
```
|
||||
|
||||
## Available Primitives
|
||||
|
||||
Auth filters run in the same sandbox as custom_code guardrails, with access to:
|
||||
|
||||
### HTTP Requests
|
||||
- `http_get(url, headers=None, timeout=30)`
|
||||
- `http_post(url, body=None, headers=None, timeout=30)`
|
||||
- `http_request(method, url, body=None, headers=None, timeout=30)`
|
||||
|
||||
### Result Actions
|
||||
- `allow()` - Continue without changes
|
||||
- `block(reason)` - Reject with error message
|
||||
- `modify(user_api_key_auth=obj)` - Return modified auth object
|
||||
|
||||
### Regex
|
||||
- `regex_match(text, pattern)`
|
||||
- `regex_replace(text, pattern, replacement)`
|
||||
- `regex_find_all(text, pattern)`
|
||||
|
||||
### JSON
|
||||
- `json_parse(text)`
|
||||
- `json_stringify(obj)`
|
||||
- `json_schema_valid(data, schema)`
|
||||
|
||||
### Text Utilities
|
||||
- `contains(text, substring)`
|
||||
- `lower(text)`, `upper(text)`, `trim(text)`
|
||||
- `word_count(text)`, `char_count(text)`
|
||||
|
||||
### Safe Builtins
|
||||
- `len()`, `str()`, `int()`, `float()`, `bool()`
|
||||
- `list()`, `dict()`, `True`, `False`, `None`
|
||||
|
||||
## UserAPIKeyAuth Fields
|
||||
|
||||
The `user_api_key_auth` object contains:
|
||||
|
||||
### Core Fields
|
||||
- `api_key`: The API key
|
||||
- `token`: Token identifier
|
||||
- `key_name`, `key_alias`: Key identifiers
|
||||
|
||||
### User/Team/Org
|
||||
- `user_id`, `user_email`, `user_role`
|
||||
- `team_id`, `team_alias`
|
||||
- `org_id`, `organization_id`
|
||||
|
||||
### Access Control
|
||||
- `models`: Allowed models
|
||||
- `team_models`: Team-level models
|
||||
- `permissions`: User permissions
|
||||
- `allowed_routes`: Allowed API routes
|
||||
|
||||
### Budgets & Limits
|
||||
- `max_budget`, `soft_budget`
|
||||
- `tpm_limit`, `rpm_limit` (tokens/requests per minute)
|
||||
- `team_max_budget`, `user_max_budget`
|
||||
|
||||
### Metadata
|
||||
- `metadata`: Dict for custom data
|
||||
|
||||
## Hot-Reload Support
|
||||
|
||||
Auth filters support hot-reload without proxy restart:
|
||||
|
||||
1. **Update via API:**
|
||||
```bash
|
||||
curl -X PUT http://localhost:4000/guardrails/{guardrail_id} \
|
||||
-H "Authorization: Bearer $ADMIN_KEY" \
|
||||
-d '{"custom_code": "..."}'
|
||||
```
|
||||
|
||||
2. **Update via Database:**
|
||||
Update the `custom_code` field in `LiteLLM_GuardrailsTable`
|
||||
|
||||
Changes take effect immediately on the next request.
|
||||
|
||||
## Error Handling
|
||||
|
||||
### Compilation Errors
|
||||
|
||||
If custom code has syntax errors, they're caught at initialization:
|
||||
|
||||
```python
|
||||
# Bad syntax - will fail at startup
|
||||
def auth_filter(request, api_key, user_api_key_auth) # Missing colon
|
||||
return allow()
|
||||
```
|
||||
|
||||
Error message: `"Syntax error in custom code: ..."`
|
||||
|
||||
### Runtime Errors
|
||||
|
||||
Uncaught exceptions are logged and converted to 500 errors:
|
||||
|
||||
```python
|
||||
# Runtime error example
|
||||
def auth_filter(request, api_key, user_api_key_auth):
|
||||
x = 1 / 0 # ZeroDivisionError
|
||||
```
|
||||
|
||||
**Best Practice:** Use try/except in your code:
|
||||
|
||||
```python
|
||||
def auth_filter(request, api_key, user_api_key_auth):
|
||||
try:
|
||||
# Your logic
|
||||
result = some_operation()
|
||||
except Exception as e:
|
||||
return block(f"Validation failed: {e}")
|
||||
return allow()
|
||||
```
|
||||
|
||||
## Use Cases
|
||||
|
||||
### 1. Organization-Based Access Control
|
||||
|
||||
```python
|
||||
def auth_filter(request, api_key, user_api_key_auth):
|
||||
org_id = user_api_key_auth.organization_id
|
||||
|
||||
# Different rules per org
|
||||
if org_id == "free-tier":
|
||||
allowed_models = ["gpt-3.5-turbo"]
|
||||
if request.url.path not in ["/chat/completions"]:
|
||||
return block("Free tier: only chat completions allowed")
|
||||
elif org_id == "enterprise":
|
||||
# Enterprise has full access
|
||||
pass
|
||||
else:
|
||||
return block("Unknown organization")
|
||||
|
||||
return allow()
|
||||
```
|
||||
|
||||
### 2. External Compliance Validation
|
||||
|
||||
```python
|
||||
async def auth_filter(request, api_key, user_api_key_auth):
|
||||
response = await http_post(
|
||||
"https://compliance.internal/validate",
|
||||
body={
|
||||
"user_id": user_api_key_auth.user_id,
|
||||
"org_id": user_api_key_auth.organization_id,
|
||||
"endpoint": request.url.path
|
||||
}
|
||||
)
|
||||
|
||||
if response["status_code"] != 200:
|
||||
return block("Compliance check failed")
|
||||
|
||||
compliance_data = response["body"]
|
||||
if not compliance_data.get("approved"):
|
||||
return block(f"Access denied: {compliance_data.get('reason')}")
|
||||
|
||||
# Enrich with compliance session
|
||||
user_api_key_auth.metadata["compliance_session_id"] = compliance_data["session_id"]
|
||||
user_api_key_auth.metadata["compliance_level"] = compliance_data["level"]
|
||||
|
||||
return modify(user_api_key_auth=user_api_key_auth)
|
||||
```
|
||||
|
||||
### 3. Header-Based Department Routing
|
||||
|
||||
```python
|
||||
def auth_filter(request, api_key, user_api_key_auth):
|
||||
department = request.headers.get("X-Department")
|
||||
|
||||
if not department:
|
||||
return block("X-Department header required")
|
||||
|
||||
# Store department in metadata for downstream use
|
||||
user_api_key_auth.metadata["department"] = department
|
||||
|
||||
# Department-specific model restrictions
|
||||
if department == "research":
|
||||
# Research can use expensive models
|
||||
pass
|
||||
elif department == "customer-service":
|
||||
# Customer service limited to fast models
|
||||
user_api_key_auth.models = ["gpt-3.5-turbo", "gpt-4o-mini"]
|
||||
else:
|
||||
return block(f"Unknown department: {department}")
|
||||
|
||||
return modify(user_api_key_auth=user_api_key_auth)
|
||||
```
|
||||
|
||||
## Security Notes
|
||||
|
||||
- Auth filters run in a sandboxed environment
|
||||
- No file I/O access
|
||||
- No arbitrary imports
|
||||
- HTTP requests have timeout limits (30s default, 60s max)
|
||||
- Exceptions are caught and logged, not exposed to clients
|
||||
|
||||
## Debugging
|
||||
|
||||
Enable debug logging to see auth filter execution:
|
||||
|
||||
```python
|
||||
# In custom code
|
||||
def auth_filter(request, api_key, user_api_key_auth):
|
||||
# Use print() for debugging (goes to proxy logs)
|
||||
print(f"Auth filter: checking user {user_api_key_auth.user_id}")
|
||||
print(f"Organization: {user_api_key_auth.organization_id}")
|
||||
|
||||
return allow()
|
||||
```
|
||||
|
||||
View logs with:
|
||||
```bash
|
||||
tail -f litellm_proxy.log | grep "Auth filter"
|
||||
```
|
||||
|
|
@ -0,0 +1,91 @@
|
|||
"""
|
||||
Auth Filter Guardrail Hook
|
||||
|
||||
This module provides a guardrail for custom authentication enrichment and validation.
|
||||
The auth_filter runs AFTER standard authentication, receiving the validated
|
||||
UserAPIKeyAuth object and allowing custom code to enrich or validate it.
|
||||
|
||||
Configuration:
|
||||
guardrails:
|
||||
- guardrail_name: "my-auth-filter"
|
||||
litellm_params:
|
||||
guardrail: "auth_filter"
|
||||
mode: "post_auth_check"
|
||||
custom_code: |
|
||||
async def auth_filter(request, api_key, user_api_key_auth):
|
||||
# Custom validation logic
|
||||
org_id = user_api_key_auth.organization_id
|
||||
if org_id == "restricted":
|
||||
return block("Access denied")
|
||||
return allow()
|
||||
|
||||
The custom code has access to HTTP primitives for external API calls:
|
||||
- http_get(url, headers=None, timeout=30)
|
||||
- http_post(url, body=None, headers=None, timeout=30)
|
||||
- http_request(method, url, body=None, headers=None, timeout=30)
|
||||
|
||||
And other primitives from the custom_code sandbox:
|
||||
- regex_match, regex_replace, json_parse, json_stringify, etc.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict
|
||||
|
||||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
|
||||
from .auth_filter_guardrail import AuthFilterGuardrail
|
||||
|
||||
|
||||
def initialize_guardrail(litellm_params: Any, guardrail: Dict[str, Any]) -> AuthFilterGuardrail:
|
||||
"""
|
||||
Initialize and register the auth filter guardrail.
|
||||
|
||||
This function is called by the guardrail registry during startup or when
|
||||
a new auth_filter guardrail is created.
|
||||
|
||||
Args:
|
||||
litellm_params: Configuration parameters (includes custom_code, mode, etc.)
|
||||
guardrail: The guardrail configuration dict
|
||||
|
||||
Returns:
|
||||
AuthFilterGuardrail instance
|
||||
|
||||
Raises:
|
||||
CustomCodeCompilationError: If the custom code fails to compile
|
||||
"""
|
||||
guardrail_name = guardrail.get("guardrail_name", "auth_filter")
|
||||
custom_code = getattr(litellm_params, "custom_code", "")
|
||||
|
||||
if not custom_code:
|
||||
raise ValueError("auth_filter guardrail requires 'custom_code' parameter")
|
||||
|
||||
instance = AuthFilterGuardrail(
|
||||
guardrail_name=guardrail_name,
|
||||
custom_code=custom_code,
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=getattr(litellm_params, "default_on", True),
|
||||
)
|
||||
|
||||
# Register with LiteLLM callback manager for hot-reload support
|
||||
import litellm
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(instance)
|
||||
|
||||
return instance
|
||||
|
||||
|
||||
# Register this guardrail with the system
|
||||
guardrail_initializer_registry = {
|
||||
SupportedGuardrailIntegrations.AUTH_FILTER.value: initialize_guardrail,
|
||||
}
|
||||
|
||||
guardrail_class_registry = {
|
||||
SupportedGuardrailIntegrations.AUTH_FILTER.value: AuthFilterGuardrail,
|
||||
}
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AuthFilterGuardrail",
|
||||
"initialize_guardrail",
|
||||
"guardrail_initializer_registry",
|
||||
"guardrail_class_registry",
|
||||
]
|
||||
|
|
@ -0,0 +1,313 @@
|
|||
"""
|
||||
Auth filter guardrail for LiteLLM.
|
||||
|
||||
This module provides a guardrail that executes user-defined Python-like code
|
||||
to implement custom authentication enrichment/validation logic. The code runs
|
||||
in a sandboxed environment with access to LiteLLM-provided primitives for
|
||||
HTTP requests and other operations.
|
||||
|
||||
The auth_filter runs AFTER standard authentication completes, receiving the
|
||||
validated UserAPIKeyAuth object. It can enrich, validate, or block based on
|
||||
custom logic.
|
||||
|
||||
Example custom code (sync):
|
||||
|
||||
def auth_filter(request, api_key, user_api_key_auth):
|
||||
'''Check organization-specific rules'''
|
||||
org_id = user_api_key_auth.organization_id
|
||||
if org_id == "restricted-org":
|
||||
return block("Organization access restricted")
|
||||
return allow()
|
||||
|
||||
Example custom code (async with HTTP):
|
||||
|
||||
async def auth_filter(request, api_key, user_api_key_auth):
|
||||
'''Call external validation API'''
|
||||
org_id = user_api_key_auth.organization_id
|
||||
|
||||
response = await http_post(
|
||||
"https://api.example.com/validate-org",
|
||||
body={"org_id": org_id}
|
||||
)
|
||||
|
||||
if not response["success"] or not response["body"].get("valid"):
|
||||
return block("Organization validation failed")
|
||||
|
||||
# Enrich with external data
|
||||
user_api_key_auth.metadata["validation_session"] = response["body"]["session_id"]
|
||||
return modify(user_api_key_auth=user_api_key_auth)
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional, Type, Union
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.custom_code.custom_code_guardrail import (
|
||||
CustomCodeCompilationError, CustomCodeExecutionError, CustomCodeGuardrail)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.custom_code.primitives import \
|
||||
get_custom_code_primitives
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import \
|
||||
GuardrailConfigModel
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class AuthFilterGuardrailError(Exception):
|
||||
"""Raised when auth filter guardrail execution fails."""
|
||||
|
||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None) -> None:
|
||||
super().__init__(message)
|
||||
self.details = details or {}
|
||||
|
||||
|
||||
class AuthFilterGuardrailConfigModel(GuardrailConfigModel):
|
||||
"""Configuration parameters for the auth filter guardrail."""
|
||||
|
||||
custom_code: str
|
||||
"""The Python-like code containing the auth_filter function."""
|
||||
|
||||
|
||||
class AuthFilterGuardrail(CustomCodeGuardrail):
|
||||
"""
|
||||
Guardrail that executes user-defined auth filter code.
|
||||
|
||||
The auth_filter runs AFTER standard authentication and receives:
|
||||
- request: Original FastAPI Request object
|
||||
- api_key: The API key used for authentication
|
||||
- user_api_key_auth: The validated UserAPIKeyAuth object
|
||||
|
||||
Users write an `auth_filter(request, api_key, user_api_key_auth)` function
|
||||
that returns one of:
|
||||
- allow() - continue without modification
|
||||
- block(reason) - reject with a message
|
||||
- modify(user_api_key_auth={...}) - return enriched auth object
|
||||
|
||||
Example:
|
||||
async def auth_filter(request, api_key, user_api_key_auth):
|
||||
org_id = user_api_key_auth.organization_id
|
||||
if org_id == "restricted":
|
||||
return block("Access denied")
|
||||
return allow()
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
custom_code: str,
|
||||
guardrail_name: Optional[str] = "auth_filter",
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""
|
||||
Initialize the auth filter guardrail.
|
||||
|
||||
Args:
|
||||
custom_code: The source code containing auth_filter function
|
||||
guardrail_name: Name of this guardrail instance
|
||||
**kwargs: Additional arguments passed to CustomGuardrail
|
||||
"""
|
||||
self.custom_code = custom_code
|
||||
self._compiled_function: Optional[Any] = None
|
||||
self._compile_lock = threading.Lock()
|
||||
self._compile_error: Optional[str] = None
|
||||
|
||||
# Auth filter only supports post_auth_check event hook
|
||||
supported_event_hooks = [GuardrailEventHooks.post_auth_check]
|
||||
|
||||
# Initialize parent - do NOT call super().__init__() which would
|
||||
# call _compile_custom_code() with wrong function name
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
||||
CustomGuardrail.__init__(
|
||||
self,
|
||||
guardrail_name=guardrail_name,
|
||||
supported_event_hooks=supported_event_hooks,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Compile the code on initialization
|
||||
self._compile_custom_code()
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> Optional[Type[GuardrailConfigModel]]:
|
||||
"""Returns the config model for the UI."""
|
||||
return AuthFilterGuardrailConfigModel
|
||||
|
||||
def _compile_custom_code(self) -> None:
|
||||
"""
|
||||
Compile the custom code and extract the auth_filter function.
|
||||
|
||||
The code runs in a sandboxed environment with only the allowed primitives.
|
||||
"""
|
||||
with self._compile_lock:
|
||||
if self._compiled_function is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
# Create a restricted execution environment
|
||||
# Only include our safe primitives
|
||||
exec_globals = get_custom_code_primitives().copy()
|
||||
|
||||
# Execute the user code in the restricted environment
|
||||
exec(compile(self.custom_code, "<auth_filter>", "exec"), exec_globals)
|
||||
|
||||
# Extract the auth_filter function
|
||||
if "auth_filter" not in exec_globals:
|
||||
raise CustomCodeCompilationError(
|
||||
"Custom code must define an 'auth_filter' function. "
|
||||
"Expected signature: auth_filter(request, api_key, user_api_key_auth)"
|
||||
)
|
||||
|
||||
auth_fn = exec_globals["auth_filter"]
|
||||
if not callable(auth_fn):
|
||||
raise CustomCodeCompilationError(
|
||||
"'auth_filter' must be a callable function"
|
||||
)
|
||||
|
||||
self._compiled_function = auth_fn
|
||||
verbose_proxy_logger.debug(
|
||||
f"Auth filter guardrail '{self.guardrail_name}' compiled successfully"
|
||||
)
|
||||
|
||||
except SyntaxError as e:
|
||||
self._compile_error = f"Syntax error in custom code: {e}"
|
||||
raise CustomCodeCompilationError(self._compile_error) from e
|
||||
except CustomCodeCompilationError:
|
||||
raise
|
||||
except Exception as e:
|
||||
self._compile_error = f"Failed to compile custom code: {e}"
|
||||
raise CustomCodeCompilationError(self._compile_error) from e
|
||||
|
||||
async def execute_auth_filter(
|
||||
self,
|
||||
request: Request,
|
||||
api_key: str,
|
||||
user_api_key_auth: UserAPIKeyAuth,
|
||||
) -> Union[UserAPIKeyAuth, None]:
|
||||
"""
|
||||
Execute auth filter with the custom_auth-compatible signature.
|
||||
|
||||
This method calls the user-defined auth_filter function and processes
|
||||
its result.
|
||||
|
||||
Args:
|
||||
request: FastAPI Request object
|
||||
api_key: The API key used for authentication
|
||||
user_api_key_auth: The validated UserAPIKeyAuth object from standard auth
|
||||
|
||||
Returns:
|
||||
- UserAPIKeyAuth: Modified/enriched auth object (replaces original)
|
||||
- None: Allow without modification (continue with original)
|
||||
|
||||
Raises:
|
||||
HTTPException: If the filter blocks the request
|
||||
CustomCodeExecutionError: If execution fails
|
||||
"""
|
||||
if self._compiled_function is None:
|
||||
raise AuthFilterGuardrailError("Auth filter not compiled")
|
||||
|
||||
if self._compile_error:
|
||||
raise AuthFilterGuardrailError(
|
||||
f"Auth filter has compilation error: {self._compile_error}"
|
||||
)
|
||||
|
||||
try:
|
||||
# Call the user's auth_filter function
|
||||
result = self._compiled_function(request, api_key, user_api_key_auth)
|
||||
|
||||
# Handle async functions
|
||||
if asyncio.iscoroutine(result):
|
||||
result = await result
|
||||
|
||||
# Process the result
|
||||
return self._process_auth_result(result, user_api_key_auth)
|
||||
|
||||
except HTTPException:
|
||||
# Re-raise HTTP exceptions from block() calls
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
f"Auth filter '{self.guardrail_name}' execution failed: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
raise CustomCodeExecutionError(
|
||||
f"Auth filter execution failed: {str(e)}",
|
||||
details={"guardrail_name": self.guardrail_name, "error": str(e)},
|
||||
) from e
|
||||
|
||||
def _process_auth_result(
|
||||
self, result: Any, original_auth: UserAPIKeyAuth
|
||||
) -> Union[UserAPIKeyAuth, None]:
|
||||
"""
|
||||
Convert auth_filter result to auth response.
|
||||
|
||||
Expected result formats from custom code:
|
||||
- allow() -> {"action": "allow"} -> return None (no changes)
|
||||
- block(reason) -> {"action": "block", "reason": str} -> raise HTTPException
|
||||
- modify(user_api_key_auth=obj) -> {"action": "modify", "user_api_key_auth": dict} -> return modified auth
|
||||
|
||||
Args:
|
||||
result: Result from user's auth_filter function
|
||||
original_auth: The original UserAPIKeyAuth object
|
||||
|
||||
Returns:
|
||||
- UserAPIKeyAuth: If modified
|
||||
- None: If allow (no changes)
|
||||
|
||||
Raises:
|
||||
HTTPException: If blocked
|
||||
"""
|
||||
# If result is not a dict, assume allow
|
||||
if not isinstance(result, dict):
|
||||
return None
|
||||
|
||||
action = result.get("action", "allow")
|
||||
|
||||
if action == "allow":
|
||||
return None # No changes, continue with original
|
||||
|
||||
elif action == "block":
|
||||
reason = result.get("reason", "Blocked by auth filter")
|
||||
raise HTTPException(status_code=403, detail={"error": reason})
|
||||
|
||||
elif action == "modify":
|
||||
# Return modified UserAPIKeyAuth
|
||||
if "user_api_key_auth" in result:
|
||||
modified_data = result["user_api_key_auth"]
|
||||
# If it's already a UserAPIKeyAuth, return it
|
||||
if isinstance(modified_data, UserAPIKeyAuth):
|
||||
return modified_data
|
||||
# Otherwise try to construct from dict
|
||||
elif isinstance(modified_data, dict):
|
||||
# Merge with original to preserve required fields
|
||||
merged = original_auth.model_dump()
|
||||
merged.update(modified_data)
|
||||
return UserAPIKeyAuth(**merged)
|
||||
|
||||
# Default: allow without changes
|
||||
return None
|
||||
|
||||
def update_custom_code(self, new_code: str) -> None:
|
||||
"""
|
||||
Update the custom code and recompile.
|
||||
|
||||
This method is called by the hot-reload mechanism when the guardrail
|
||||
configuration changes in the database.
|
||||
|
||||
Args:
|
||||
new_code: The new custom code to compile
|
||||
"""
|
||||
with self._compile_lock:
|
||||
self.custom_code = new_code
|
||||
self._compiled_function = None
|
||||
self._compile_error = None
|
||||
self._compile_custom_code()
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Auth filter guardrail '{self.guardrail_name}' hot-reloaded successfully"
|
||||
)
|
||||
|
|
@ -5,24 +5,18 @@ from typing import Any, Dict, List, Literal, Optional, Union
|
|||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
from typing_extensions import Required, TypedDict
|
||||
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.enkryptai import (
|
||||
EnkryptAIGuardrailConfigs,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import (
|
||||
GraySwanGuardrailConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.ibm import (
|
||||
IBMGuardrailsBaseConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
|
||||
ContentFilterCategoryConfig,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.qualifire import (
|
||||
QualifireGuardrailConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import (
|
||||
ToolPermissionGuardrailConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.enkryptai import \
|
||||
EnkryptAIGuardrailConfigs
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import \
|
||||
GraySwanGuardrailConfigModel
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.ibm import \
|
||||
IBMGuardrailsBaseConfigModel
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import \
|
||||
ContentFilterCategoryConfig
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.qualifire import \
|
||||
QualifireGuardrailConfigModel
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import \
|
||||
ToolPermissionGuardrailConfigModel
|
||||
|
||||
"""
|
||||
Pydantic object defining how to set guardrails on litellm proxy
|
||||
|
|
@ -40,6 +34,7 @@ guardrails:
|
|||
|
||||
class SupportedGuardrailIntegrations(Enum):
|
||||
APORIA = "aporia"
|
||||
AUTH_FILTER = "auth_filter"
|
||||
BEDROCK = "bedrock"
|
||||
GUARDRAILS_AI = "guardrails_ai"
|
||||
LAKERA = "lakera"
|
||||
|
|
@ -754,6 +749,7 @@ class guardrailConfig(TypedDict):
|
|||
class GuardrailEventHooks(str, Enum):
|
||||
pre_call = "pre_call"
|
||||
post_call = "post_call"
|
||||
post_auth_check = "post_auth_check"
|
||||
during_call = "during_call"
|
||||
logging_only = "logging_only"
|
||||
pre_mcp_call = "pre_mcp_call"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue