fix: policy_crud_router

This commit is contained in:
Ishaan Jaffer 2026-01-23 09:26:43 -08:00
parent 1a3cb2a1d3
commit cbd73ddded
2 changed files with 158 additions and 0 deletions

View file

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

View file

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