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:
Krrish Dholakia 2026-02-18 15:03:02 -08:00
parent a9058bb584
commit 36e591ae21
4 changed files with 746 additions and 18 deletions

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

View file

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

View file

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

View file

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