mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
feat: init PolicyResolver
This commit is contained in:
parent
715e3b6b01
commit
c96ef32843
1 changed files with 199 additions and 0 deletions
199
litellm/proxy/policy_engine/policy_resolver.py
Normal file
199
litellm/proxy/policy_engine/policy_resolver.py
Normal file
|
|
@ -0,0 +1,199 @@
|
|||
"""
|
||||
Policy Resolver - Resolves final guardrail list from policies.
|
||||
|
||||
Handles:
|
||||
- Inheritance chain resolution
|
||||
- Applying add/remove guardrails
|
||||
- Combining guardrails from multiple matching policies
|
||||
"""
|
||||
|
||||
from typing import Dict, List, Optional, Set
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.types.proxy.policy_engine import (
|
||||
Policy,
|
||||
PolicyMatchContext,
|
||||
ResolvedPolicy,
|
||||
)
|
||||
|
||||
|
||||
class PolicyResolver:
|
||||
"""
|
||||
Resolves the final list of guardrails from policies.
|
||||
|
||||
Handles inheritance chains and add/remove operations.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def resolve_inheritance_chain(
|
||||
policy_name: str,
|
||||
policies: Dict[str, Policy],
|
||||
visited: Optional[Set[str]] = None,
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get the inheritance chain for a policy (from root to policy).
|
||||
|
||||
Args:
|
||||
policy_name: Name of the policy
|
||||
policies: Dictionary of all policies
|
||||
visited: Set of visited policies (for cycle detection)
|
||||
|
||||
Returns:
|
||||
List of policy names from root ancestor to the given policy
|
||||
"""
|
||||
if visited is None:
|
||||
visited = set()
|
||||
|
||||
if policy_name in visited:
|
||||
verbose_proxy_logger.warning(
|
||||
f"Circular inheritance detected for policy '{policy_name}'"
|
||||
)
|
||||
return []
|
||||
|
||||
policy = policies.get(policy_name)
|
||||
if policy is None:
|
||||
return []
|
||||
|
||||
visited.add(policy_name)
|
||||
|
||||
if policy.inherit:
|
||||
parent_chain = PolicyResolver.resolve_inheritance_chain(
|
||||
policy_name=policy.inherit, policies=policies, visited=visited
|
||||
)
|
||||
return parent_chain + [policy_name]
|
||||
|
||||
return [policy_name]
|
||||
|
||||
@staticmethod
|
||||
def resolve_policy_guardrails(
|
||||
policy_name: str,
|
||||
policies: Dict[str, Policy],
|
||||
) -> ResolvedPolicy:
|
||||
"""
|
||||
Resolve the final guardrails for a single policy, including inheritance.
|
||||
|
||||
Args:
|
||||
policy_name: Name of the policy to resolve
|
||||
policies: Dictionary of all policies
|
||||
|
||||
Returns:
|
||||
ResolvedPolicy with final guardrails list
|
||||
"""
|
||||
inheritance_chain = PolicyResolver.resolve_inheritance_chain(
|
||||
policy_name=policy_name, policies=policies
|
||||
)
|
||||
|
||||
# Start with empty set of guardrails
|
||||
guardrails: Set[str] = set()
|
||||
|
||||
# Apply each policy in the chain (from root to leaf)
|
||||
for chain_policy_name in inheritance_chain:
|
||||
policy = policies.get(chain_policy_name)
|
||||
if policy is None:
|
||||
continue
|
||||
|
||||
# Add guardrails
|
||||
for guardrail in policy.guardrails.get_add():
|
||||
guardrails.add(guardrail)
|
||||
|
||||
# Remove guardrails
|
||||
for guardrail in policy.guardrails.get_remove():
|
||||
guardrails.discard(guardrail)
|
||||
|
||||
return ResolvedPolicy(
|
||||
policy_name=policy_name,
|
||||
guardrails=list(guardrails),
|
||||
inheritance_chain=inheritance_chain,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def resolve_guardrails_for_context(
|
||||
context: PolicyMatchContext,
|
||||
policies: Optional[Dict[str, Policy]] = None,
|
||||
) -> List[str]:
|
||||
"""
|
||||
Resolve the final list of guardrails for a request context.
|
||||
|
||||
This:
|
||||
1. Finds all policies that match the context
|
||||
2. Resolves each policy's guardrails (including inheritance)
|
||||
3. Combines all guardrails (union)
|
||||
|
||||
Args:
|
||||
context: The request context
|
||||
policies: Dictionary of all policies (if None, uses global registry)
|
||||
|
||||
Returns:
|
||||
List of guardrail names to apply
|
||||
"""
|
||||
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
|
||||
if policies is None:
|
||||
registry = get_policy_registry()
|
||||
if not registry.is_initialized():
|
||||
return []
|
||||
policies = registry.get_all_policies()
|
||||
|
||||
# Get matching policies
|
||||
matching_policy_names = PolicyMatcher.get_matching_policies(
|
||||
policies=policies, context=context
|
||||
)
|
||||
|
||||
if not matching_policy_names:
|
||||
verbose_proxy_logger.debug(
|
||||
f"No policies match context: team_alias={context.team_alias}, "
|
||||
f"key_alias={context.key_alias}, model={context.model}"
|
||||
)
|
||||
return []
|
||||
|
||||
# Resolve each matching policy and combine guardrails
|
||||
all_guardrails: Set[str] = set()
|
||||
|
||||
for policy_name in matching_policy_names:
|
||||
resolved = PolicyResolver.resolve_policy_guardrails(
|
||||
policy_name=policy_name, policies=policies
|
||||
)
|
||||
all_guardrails.update(resolved.guardrails)
|
||||
verbose_proxy_logger.debug(
|
||||
f"Policy '{policy_name}' contributes guardrails: {resolved.guardrails}"
|
||||
)
|
||||
|
||||
result = list(all_guardrails)
|
||||
verbose_proxy_logger.debug(
|
||||
f"Final guardrails for context: {result}"
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def get_all_resolved_policies(
|
||||
policies: Optional[Dict[str, Policy]] = None,
|
||||
) -> Dict[str, ResolvedPolicy]:
|
||||
"""
|
||||
Resolve all policies and return their final guardrails.
|
||||
|
||||
Useful for debugging and displaying policy configurations.
|
||||
|
||||
Args:
|
||||
policies: Dictionary of all policies (if None, uses global registry)
|
||||
|
||||
Returns:
|
||||
Dictionary mapping policy names to ResolvedPolicy objects
|
||||
"""
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
|
||||
if policies is None:
|
||||
registry = get_policy_registry()
|
||||
if not registry.is_initialized():
|
||||
return {}
|
||||
policies = registry.get_all_policies()
|
||||
|
||||
resolved: Dict[str, ResolvedPolicy] = {}
|
||||
|
||||
for policy_name in policies:
|
||||
resolved[policy_name] = PolicyResolver.resolve_policy_guardrails(
|
||||
policy_name=policy_name, policies=policies
|
||||
)
|
||||
|
||||
return resolved
|
||||
Loading…
Add table
Reference in a new issue