Merge pull request #21203 from jquinter/fix-too-many-statements-pre-call-utils

refactor: extract helper functions to reduce statement count in add_guardrails_from_policy_engine
This commit is contained in:
jquinter 2026-02-14 14:55:16 -03:00 committed by GitHub
commit a68d61ad76
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -40,10 +40,12 @@ service_logger_obj = ServiceLogging() # used for tracking latency on OTEL
if TYPE_CHECKING:
from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig
from litellm.types.proxy.policy_engine import PolicyMatchContext
ProxyConfig = _ProxyConfig
else:
ProxyConfig = Any
PolicyMatchContext = Any
def parse_cache_control(cache_control):
@ -1515,75 +1517,26 @@ def move_guardrails_to_metadata(
] = request_body_guardrail_config
def add_guardrails_from_policy_engine(
def _match_and_track_policies(
data: dict,
metadata_variable_name: str,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
context: "PolicyMatchContext",
request_body_policies: Any,
) -> tuple[list[str], dict[str, str]]:
"""
Add guardrails from the policy engine based on request context.
Match policies via attachments and request body, track them in metadata.
This function:
1. Extracts "policies" from request body (if present) for dynamic policy application
2. Gets matching policies based on team_alias, key_alias, and model (via attachments)
3. Combines dynamic policies with attachment-based policies
4. Resolves guardrails from all policies (including inheritance)
5. Adds guardrails to request metadata
6. Tracks applied policies in metadata for response headers
7. Removes "policies" from request body so it's not forwarded to LLM provider
Args:
data: The request data to update
metadata_variable_name: The name of the metadata field in data
user_api_key_dict: The user's API key authentication info
Returns:
Tuple of (applied_policy_names, policy_reasons)
"""
from litellm._logging import verbose_proxy_logger
from litellm.proxy.common_utils.callback_utils import (
add_policy_sources_to_metadata,
add_policy_to_applied_policies_header,
)
from litellm.proxy.common_utils.http_parsing_utils import (
get_tags_from_request_body,
)
from litellm.proxy.policy_engine.attachment_registry import (
get_attachment_registry,
)
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
from litellm.types.proxy.policy_engine import PolicyMatchContext
# Extract dynamic policies from request body (if present)
# These will be combined with attachment-based policies
request_body_policies = data.pop("policies", None)
registry = get_policy_registry()
verbose_proxy_logger.debug(
f"Policy engine: registry initialized={registry.is_initialized()}, "
f"policy_count={len(registry.get_all_policies())}"
)
if not registry.is_initialized():
verbose_proxy_logger.debug(
"Policy engine not initialized, skipping policy matching"
)
return
# Extract tags using the shared helper (handles metadata / litellm_metadata,
# top-level tags, deduplication, and type filtering).
all_tags = get_tags_from_request_body(data) or None
context = PolicyMatchContext(
team_alias=user_api_key_dict.team_alias,
key_alias=user_api_key_dict.key_alias,
model=data.get("model"),
tags=all_tags,
)
verbose_proxy_logger.debug(
f"Policy engine: matching policies for context team_alias={context.team_alias}, "
f"key_alias={context.key_alias}, model={context.model}, tags={context.tags}"
)
# Get matching policies via attachments (with match reasons for attribution)
attachment_registry = get_attachment_registry()
@ -1591,7 +1544,6 @@ def add_guardrails_from_policy_engine(
context
)
matching_policy_names = [m["policy_name"] for m in matches_with_reasons]
# Build reasons map: {"hipaa-policy": "tag:healthcare", ...}
policy_reasons = {m["policy_name"]: m["matched_via"] for m in matches_with_reasons}
verbose_proxy_logger.debug(
@ -1607,7 +1559,7 @@ def add_guardrails_from_policy_engine(
)
if not all_policy_names:
return
return [], {}
# Filter to only policies whose conditions match the context
applied_policy_names = PolicyMatcher.get_policies_with_matching_conditions(
@ -1635,6 +1587,18 @@ def add_guardrails_from_policy_engine(
request_data=data, policy_sources=applied_reasons
)
return applied_policy_names, policy_reasons
def _apply_resolved_guardrails_to_metadata(
data: dict,
metadata_variable_name: str,
context: "PolicyMatchContext",
) -> None:
"""Apply resolved guardrails and pipelines to request metadata."""
from litellm._logging import verbose_proxy_logger
from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
# Resolve guardrails from matching policies
resolved_guardrails = PolicyResolver.resolve_guardrails_for_context(context=context)
@ -1683,6 +1647,72 @@ def add_guardrails_from_policy_engine(
)
def add_guardrails_from_policy_engine(
data: dict,
metadata_variable_name: str,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""
Add guardrails from the policy engine based on request context.
This function:
1. Extracts "policies" from request body (if present) for dynamic policy application
2. Gets matching policies based on team_alias, key_alias, and model (via attachments)
3. Combines dynamic policies with attachment-based policies
4. Resolves guardrails from all policies (including inheritance)
5. Adds guardrails to request metadata
6. Tracks applied policies in metadata for response headers
7. Removes "policies" from request body so it's not forwarded to LLM provider
Args:
data: The request data to update
metadata_variable_name: The name of the metadata field in data
user_api_key_dict: The user's API key authentication info
"""
from litellm._logging import verbose_proxy_logger
from litellm.proxy.common_utils.http_parsing_utils import (
get_tags_from_request_body,
)
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
from litellm.types.proxy.policy_engine import PolicyMatchContext
# Extract dynamic policies from request body (if present)
request_body_policies = data.pop("policies", None)
registry = get_policy_registry()
verbose_proxy_logger.debug(
f"Policy engine: registry initialized={registry.is_initialized()}, "
f"policy_count={len(registry.get_all_policies())}"
)
if not registry.is_initialized():
verbose_proxy_logger.debug(
"Policy engine not initialized, skipping policy matching"
)
return
# Extract tags and build context
all_tags = get_tags_from_request_body(data) or None
context = PolicyMatchContext(
team_alias=user_api_key_dict.team_alias,
key_alias=user_api_key_dict.key_alias,
model=data.get("model"),
tags=all_tags,
)
verbose_proxy_logger.debug(
f"Policy engine: matching policies for context team_alias={context.team_alias}, "
f"key_alias={context.key_alias}, model={context.model}, tags={context.tags}"
)
# Match and track policies based on attachments and request body
_match_and_track_policies(data, context, request_body_policies)
# Always resolve and apply guardrails, even if no policies matched above.
# PolicyResolver does its own independent matching and inheritance resolution,
# so guardrails can still be applied via inherited parent policies.
_apply_resolved_guardrails_to_metadata(data, metadata_variable_name, context)
def add_provider_specific_headers_to_request(
data: dict,
headers: dict,