diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 32cddc0ef58..1d3ef2e10c2 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1311,6 +1311,118 @@ def _add_guardrails_from_key_or_team_metadata( data[metadata_variable_name]["guardrails"] = list(combined_guardrails) +def _add_guardrails_from_policies_in_metadata( + key_metadata: Optional[dict], + team_metadata: Optional[dict], + data: dict, + metadata_variable_name: str, +) -> None: + """ + Helper to resolve guardrails from policies attached to key/team metadata. + + This function: + 1. Gets policy names from key and team metadata + 2. Resolves guardrails from those policies (including inheritance) + 3. Adds resolved guardrails to request metadata + + Args: + key_metadata: The key metadata dictionary to check for policies + team_metadata: The team metadata dictionary to check for policies + data: The request data to update + metadata_variable_name: The name of the metadata field in data + """ + from litellm._logging import verbose_proxy_logger + from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_resolver import PolicyResolver + from litellm.proxy.utils import _premium_user_check + from litellm.types.proxy.policy_engine import PolicyMatchContext + + # Collect policy names from key and team metadata + policy_names: set = set() + + # Add key-level policies first + if key_metadata and "policies" in key_metadata: + if ( + isinstance(key_metadata["policies"], list) + and len(key_metadata["policies"]) > 0 + ): + _premium_user_check() + policy_names.update(key_metadata["policies"]) + + # Add team-level policies + if team_metadata and "policies" in team_metadata: + if ( + isinstance(team_metadata["policies"], list) + and len(team_metadata["policies"]) > 0 + ): + _premium_user_check() + policy_names.update(team_metadata["policies"]) + + if not policy_names: + return + + verbose_proxy_logger.debug( + f"Policy engine: resolving guardrails from key/team policies: {policy_names}" + ) + + # Check if policy registry is initialized + registry = get_policy_registry() + if not registry.is_initialized(): + verbose_proxy_logger.debug( + "Policy engine not initialized, skipping policy resolution from metadata" + ) + return + + # Build context for policy resolution (model from request data) + context = PolicyMatchContext(model=data.get("model")) + + # Get all policies from registry + all_policies = registry.get_all_policies() + + # Resolve guardrails from the specified policies + resolved_guardrails: set = set() + for policy_name in policy_names: + if registry.has_policy(policy_name): + resolved_policy = PolicyResolver.resolve_policy_guardrails( + policy_name=policy_name, + policies=all_policies, + context=context, + ) + resolved_guardrails.update(resolved_policy.guardrails) + verbose_proxy_logger.debug( + f"Policy engine: resolved guardrails from policy '{policy_name}': {resolved_policy.guardrails}" + ) + else: + verbose_proxy_logger.warning( + f"Policy engine: policy '{policy_name}' not found in registry" + ) + + if not resolved_guardrails: + return + + # Add resolved guardrails to request metadata + if metadata_variable_name not in data: + data[metadata_variable_name] = {} + + existing_guardrails = data[metadata_variable_name].get("guardrails", []) + if not isinstance(existing_guardrails, list): + existing_guardrails = [] + + # Combine existing guardrails with policy-resolved guardrails (no duplicates) + combined = set(existing_guardrails) + combined.update(resolved_guardrails) + data[metadata_variable_name]["guardrails"] = list(combined) + + # Store applied policies in metadata for tracking + if "applied_policies" not in data[metadata_variable_name]: + data[metadata_variable_name]["applied_policies"] = [] + data[metadata_variable_name]["applied_policies"].extend(list(policy_names)) + + verbose_proxy_logger.debug( + f"Policy engine: added guardrails from key/team policies to request metadata: {list(resolved_guardrails)}" + ) + + def move_guardrails_to_metadata( data: dict, _metadata_variable_name: str, @@ -1321,6 +1433,7 @@ def move_guardrails_to_metadata( - If guardrails set on API Key metadata then sets guardrails on request metadata - If guardrails not set on API key, then checks request metadata + - Adds guardrails from policies attached to key/team metadata - Adds guardrails from policy engine based on team/key/model context """ # Check key-level guardrails @@ -1331,6 +1444,16 @@ def move_guardrails_to_metadata( metadata_variable_name=_metadata_variable_name, ) + ######################################################################################### + # Add guardrails from policies attached to key/team metadata + ######################################################################################### + _add_guardrails_from_policies_in_metadata( + key_metadata=user_api_key_dict.metadata, + team_metadata=user_api_key_dict.team_metadata, + data=data, + metadata_variable_name=_metadata_variable_name, + ) + ######################################################################################### # Add guardrails from policy engine based on team/key/model context ######################################################################################### diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 994f6ee862c..64dfc6a5d8f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -381,6 +381,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( router as pass_through_router, ) +from litellm.proxy.policy_engine.policy_endpoints import router as policy_crud_router from litellm.proxy.prompts.prompt_endpoints import router as prompts_router from litellm.proxy.public_endpoints import router as public_endpoints_router from litellm.proxy.rag_endpoints.endpoints import router as rag_router @@ -3801,6 +3802,9 @@ class ProxyConfig: if self._should_load_db_object(object_type="guardrails"): await self._init_guardrails_in_db(prisma_client=prisma_client) + if self._should_load_db_object(object_type="policies"): + await self._init_policies_in_db(prisma_client=prisma_client) + if self._should_load_db_object(object_type="vector_stores"): await self._init_vector_stores_in_db(prisma_client=prisma_client) @@ -4024,6 +4028,36 @@ class ProxyConfig: ) ) + async def _init_policies_in_db(self, prisma_client: PrismaClient): + """ + Initialize policies and policy attachments from database into the in-memory registries. + """ + from litellm.proxy.policy_engine.attachment_registry import ( + get_attachment_registry, + ) + from litellm.proxy.policy_engine.policy_registry import get_policy_registry + + try: + # Get the global singleton instances + policy_registry = get_policy_registry() + attachment_registry = get_attachment_registry() + + # Sync policies from DB to in-memory registry + await policy_registry.sync_policies_from_db(prisma_client=prisma_client) + + # Sync attachments from DB to in-memory registry + await attachment_registry.sync_attachments_from_db(prisma_client=prisma_client) + + verbose_proxy_logger.debug( + "Successfully synced policies and attachments from DB" + ) + except Exception as e: + verbose_proxy_logger.exception( + "litellm.proxy.proxy_server.py::ProxyConfig:_init_policies_in_db - {}".format( + str(e) + ) + ) + async def _init_vector_stores_in_db(self, prisma_client: PrismaClient): from litellm.vector_stores.vector_store_registry import VectorStoreRegistry @@ -10702,6 +10736,7 @@ app.include_router(caching_router) app.include_router(analytics_router) app.include_router(guardrails_router) app.include_router(policy_router) +app.include_router(policy_crud_router) app.include_router(search_tool_management_router) app.include_router(prompts_router) app.include_router(callback_management_endpoints_router)