diff --git a/litellm/proxy/policy_engine/policy_resolver.py b/litellm/proxy/policy_engine/policy_resolver.py new file mode 100644 index 00000000000..b8686680c12 --- /dev/null +++ b/litellm/proxy/policy_engine/policy_resolver.py @@ -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