mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: policy_crud_router
This commit is contained in:
parent
1a3cb2a1d3
commit
cbd73ddded
2 changed files with 158 additions and 0 deletions
|
|
@ -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
|
||||
#########################################################################################
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue