diff --git a/litellm/proxy/guardrails/guardrail_hooks/auth_filter/README.md b/litellm/proxy/guardrails/guardrail_hooks/auth_filter/README.md new file mode 100644 index 00000000000..e1ec3a57cb1 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/auth_filter/README.md @@ -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" +``` diff --git a/litellm/proxy/guardrails/guardrail_hooks/auth_filter/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/auth_filter/__init__.py new file mode 100644 index 00000000000..279c2eb3c5f --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/auth_filter/__init__.py @@ -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", +] diff --git a/litellm/proxy/guardrails/guardrail_hooks/auth_filter/auth_filter_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/auth_filter/auth_filter_guardrail.py new file mode 100644 index 00000000000..70f5a89de63 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/auth_filter/auth_filter_guardrail.py @@ -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, "", "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" + ) diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index bbbd60dab7c..ac0bfe64743 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -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"