From 917df893cb36124107e224fd8313a172a01e5a66 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Thu, 22 Jan 2026 15:07:57 -0800 Subject: [PATCH] test_add_guardrails_from_policy_engine --- litellm/proxy/common_utils/callback_utils.py | 26 ++++++ litellm/proxy/litellm_pre_call_utils.py | 84 ++++++++++++++++++- .../proxy/test_litellm_pre_call_utils.py | 58 +++++++++++++ 3 files changed, 167 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index cb434da55b3..c0cff84bacd 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -380,6 +380,11 @@ def get_logging_caching_headers(request_data: Dict) -> Optional[Dict]: _metadata["applied_guardrails"] ) + if "applied_policies" in _metadata: + headers["x-litellm-applied-policies"] = ",".join( + _metadata["applied_policies"] + ) + if "semantic-similarity" in _metadata: headers["x-litellm-semantic-similarity"] = str(_metadata["semantic-similarity"]) @@ -406,6 +411,27 @@ def add_guardrail_to_applied_guardrails_header( request_data["metadata"] = _metadata +def add_policy_to_applied_policies_header( + request_data: Dict, policy_name: Optional[str] +): + """ + Add a policy name to the applied_policies list in request metadata. + + This is used to track which policies were applied to a request, + similar to how applied_guardrails tracks guardrails. + """ + if policy_name is None: + return + _metadata = request_data.get("metadata", None) or {} + if "applied_policies" in _metadata: + if policy_name not in _metadata["applied_policies"]: + _metadata["applied_policies"].append(policy_name) + else: + _metadata["applied_policies"] = [policy_name] + # Ensure metadata is set back to request_data (important when metadata didn't exist) + request_data["metadata"] = _metadata + + def add_guardrail_response_to_standard_logging_object( litellm_logging_obj: Optional["LiteLLMLogging"], guardrail_response: StandardLoggingGuardrailInformation, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 1fbd8ee72c2..a868d91cd13 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1082,13 +1082,20 @@ async def add_litellm_data_to_request( # noqa: PLR0915 if disabled_callbacks and isinstance(disabled_callbacks, list): data["litellm_disabled_callbacks"] = disabled_callbacks - # Guardrails + # Guardrails from key/team metadata move_guardrails_to_metadata( data=data, _metadata_variable_name=_metadata_variable_name, user_api_key_dict=user_api_key_dict, ) + # Guardrails from policy engine + add_guardrails_from_policy_engine( + data=data, + metadata_variable_name=_metadata_variable_name, + user_api_key_dict=user_api_key_dict, + ) + # Team Model Aliases _update_model_if_team_alias_exists( data=data, @@ -1351,6 +1358,81 @@ def move_guardrails_to_metadata( ] = request_body_guardrail_config +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. Gets matching policies based on team_alias, key_alias, and model + 2. Resolves guardrails from matching policies (including inheritance) + 3. Adds guardrails to request metadata + 4. Tracks applied policies in metadata for response headers + + 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.proxy.common_utils.callback_utils import ( + add_policy_to_applied_policies_header, + ) + 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 + + registry = get_policy_registry() + if not registry.is_initialized(): + return + + # Build context from request + context = PolicyMatchContext( + team_alias=user_api_key_dict.team_alias, + key_alias=user_api_key_dict.key_alias, + model=data.get("model"), + ) + + # Get matching policies + policies = registry.get_all_policies() + matching_policy_names = PolicyMatcher.get_matching_policies( + policies=policies, context=context + ) + + if not matching_policy_names: + return + + # Track applied policies in metadata + for policy_name in matching_policy_names: + add_policy_to_applied_policies_header( + request_data=data, policy_name=policy_name + ) + + # Resolve guardrails from matching policies + resolved_guardrails = PolicyResolver.resolve_guardrails_for_context( + context=context, policies=policies + ) + + 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) + + def add_provider_specific_headers_to_request( data: dict, headers: dict, diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 133fc07d340..c0a4fd2e5a2 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -16,6 +16,7 @@ from litellm.proxy.litellm_pre_call_utils import ( _get_dynamic_logging_metadata, _get_enforced_params, _update_model_if_key_alias_exists, + add_guardrails_from_policy_engine, add_litellm_data_to_request, check_if_token_is_service_account, ) @@ -1477,3 +1478,60 @@ async def test_embedding_header_forwarding_without_model_group_config(): finally: # Restore original model_group_settings litellm.model_group_settings = original_model_group_settings + + +def test_add_guardrails_from_policy_engine(): + """ + Test that add_guardrails_from_policy_engine adds guardrails from matching policies + and tracks applied policies in metadata. + """ + from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.types.proxy.policy_engine import Policy, PolicyGuardrails, PolicyScope + + # Setup test data + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + team_alias="healthcare-team", + key_alias="my-key", + ) + + # Setup mock policies in the registry (directly set parsed Policy objects) + registry = get_policy_registry() + registry._policies = { + "global-baseline": Policy( + guardrails=PolicyGuardrails(add=["pii_blocker"]), + scope=PolicyScope(teams=["*"]), + ), + "healthcare": Policy( + guardrails=PolicyGuardrails(add=["hipaa_audit"]), + scope=PolicyScope(teams=["healthcare-team"]), + ), + } + registry._initialized = True + + # Call the function + add_guardrails_from_policy_engine( + data=data, + metadata_variable_name="metadata", + user_api_key_dict=user_api_key_dict, + ) + + # Verify guardrails were added + assert "guardrails" in data["metadata"] + assert "pii_blocker" in data["metadata"]["guardrails"] + assert "hipaa_audit" in data["metadata"]["guardrails"] + + # Verify applied policies were tracked + assert "applied_policies" in data["metadata"] + assert "global-baseline" in data["metadata"]["applied_policies"] + assert "healthcare" in data["metadata"]["applied_policies"] + + # Clean up registry + registry._policies = {} + registry._initialized = False