test updates

This commit is contained in:
Ishaan Jaffer 2026-01-22 17:18:51 -08:00
parent b77fe50ed1
commit 1a48932b19
5 changed files with 914 additions and 120 deletions

View file

@ -0,0 +1,289 @@
"""
Unit tests for AttachmentRegistry - tests policy attachment management.
Tests:
- Loading attachments from config
- Getting attached policies for a context
- Global scope attachments
- Team/key/model specific attachments
"""
import pytest
from litellm.proxy.policy_engine.attachment_registry import (
AttachmentRegistry,
get_attachment_registry,
)
from litellm.types.proxy.policy_engine import (
PolicyAttachment,
PolicyMatchContext,
)
class TestAttachmentRegistryLoading:
"""Test loading attachments from configuration."""
def test_load_attachments_simple(self):
"""Test loading simple attachments."""
registry = AttachmentRegistry()
config = [
{"policy": "global-baseline", "scope": "*"},
{"policy": "healthcare-policy", "teams": ["healthcare-team"]},
]
registry.load_attachments(config)
assert registry.is_initialized()
assert len(registry.get_all_attachments()) == 2
def test_load_attachments_with_multiple_scopes(self):
"""Test loading attachments with multiple scope types."""
registry = AttachmentRegistry()
config = [
{"policy": "global-baseline", "scope": "*"},
{"policy": "team-policy", "teams": ["team-a", "team-b"]},
{"policy": "key-policy", "keys": ["dev-key-*"]},
{"policy": "model-policy", "models": ["gpt-4", "gpt-4-turbo"]},
]
registry.load_attachments(config)
assert len(registry.get_all_attachments()) == 4
def test_load_attachments_empty_list(self):
"""Test loading empty attachments list."""
registry = AttachmentRegistry()
registry.load_attachments([])
assert registry.is_initialized()
assert len(registry.get_all_attachments()) == 0
def test_clear_attachments(self):
"""Test clearing attachments."""
registry = AttachmentRegistry()
registry.load_attachments([{"policy": "test", "scope": "*"}])
assert registry.is_initialized()
registry.clear()
assert not registry.is_initialized()
assert len(registry.get_all_attachments()) == 0
class TestGetAttachedPolicies:
"""Test getting attached policies for a context."""
def test_global_scope_matches_all(self):
"""Test global scope (*) matches all contexts."""
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "global-baseline", "scope": "*"},
])
context = PolicyMatchContext(
team_alias="any-team", key_alias="any-key", model="any-model"
)
attached = registry.get_attached_policies(context)
assert "global-baseline" in attached
def test_team_specific_attachment(self):
"""Test team-specific attachment matches only that team."""
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "healthcare-policy", "teams": ["healthcare-team"]},
])
# Match
context = PolicyMatchContext(
team_alias="healthcare-team", key_alias="key", model="gpt-4"
)
attached = registry.get_attached_policies(context)
assert "healthcare-policy" in attached
# No match
context_other = PolicyMatchContext(
team_alias="finance-team", key_alias="key", model="gpt-4"
)
attached_other = registry.get_attached_policies(context_other)
assert "healthcare-policy" not in attached_other
def test_key_pattern_attachment(self):
"""Test key pattern attachment matches wildcard."""
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "dev-policy", "keys": ["dev-key-*"]},
])
# Match
context = PolicyMatchContext(
team_alias="team", key_alias="dev-key-123", model="gpt-4"
)
attached = registry.get_attached_policies(context)
assert "dev-policy" in attached
# No match
context_prod = PolicyMatchContext(
team_alias="team", key_alias="prod-key-123", model="gpt-4"
)
attached_prod = registry.get_attached_policies(context_prod)
assert "dev-policy" not in attached_prod
def test_model_specific_attachment(self):
"""Test model-specific attachment."""
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "gpt4-policy", "models": ["gpt-4", "gpt-4-turbo"]},
])
# Match
context = PolicyMatchContext(
team_alias="team", key_alias="key", model="gpt-4"
)
attached = registry.get_attached_policies(context)
assert "gpt4-policy" in attached
# No match
context_other = PolicyMatchContext(
team_alias="team", key_alias="key", model="gpt-3.5"
)
attached_other = registry.get_attached_policies(context_other)
assert "gpt4-policy" not in attached_other
def test_multiple_attachments_match(self):
"""Test multiple attachments can match same context."""
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "global-baseline", "scope": "*"},
{"policy": "healthcare-policy", "teams": ["healthcare-team"]},
{"policy": "gpt4-policy", "models": ["gpt-4"]},
])
context = PolicyMatchContext(
team_alias="healthcare-team", key_alias="key", model="gpt-4"
)
attached = registry.get_attached_policies(context)
assert "global-baseline" in attached
assert "healthcare-policy" in attached
assert "gpt4-policy" in attached
assert len(attached) == 3
def test_no_duplicate_policies(self):
"""Test same policy attached multiple ways doesn't duplicate."""
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "multi-policy", "scope": "*"},
{"policy": "multi-policy", "teams": ["healthcare-team"]},
])
context = PolicyMatchContext(
team_alias="healthcare-team", key_alias="key", model="gpt-4"
)
attached = registry.get_attached_policies(context)
# Should only appear once
assert attached.count("multi-policy") == 1
class TestIsPolicyAttached:
"""Test is_policy_attached method."""
def test_policy_is_attached(self):
"""Test checking if a specific policy is attached."""
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "global-baseline", "scope": "*"},
])
context = PolicyMatchContext(team_alias="team", key_alias="key", model="gpt-4")
assert registry.is_policy_attached("global-baseline", context) is True
assert registry.is_policy_attached("other-policy", context) is False
class TestGetAttachmentsForPolicy:
"""Test getting attachments for a specific policy."""
def test_get_attachments_for_policy(self):
"""Test getting all attachments for a policy."""
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "multi-policy", "scope": "*"},
{"policy": "multi-policy", "teams": ["team-a"]},
{"policy": "other-policy", "teams": ["team-b"]},
])
attachments = registry.get_attachments_for_policy("multi-policy")
assert len(attachments) == 2
attachments_other = registry.get_attachments_for_policy("other-policy")
assert len(attachments_other) == 1
class TestAddAndRemoveAttachments:
"""Test adding and removing individual attachments."""
def test_add_attachment(self):
"""Test adding a single attachment."""
registry = AttachmentRegistry()
registry.load_attachments([])
attachment = PolicyAttachment(policy="new-policy", scope="*")
registry.add_attachment(attachment)
assert len(registry.get_all_attachments()) == 1
def test_remove_attachments_for_policy(self):
"""Test removing all attachments for a policy."""
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "policy-a", "scope": "*"},
{"policy": "policy-a", "teams": ["team-a"]},
{"policy": "policy-b", "teams": ["team-b"]},
])
removed = registry.remove_attachments_for_policy("policy-a")
assert removed == 2
assert len(registry.get_all_attachments()) == 1
assert len(registry.get_attachments_for_policy("policy-a")) == 0
class TestPolicyAttachmentModel:
"""Test PolicyAttachment model methods."""
def test_is_global(self):
"""Test is_global method."""
global_attachment = PolicyAttachment(policy="test", scope="*")
assert global_attachment.is_global() is True
team_attachment = PolicyAttachment(policy="test", teams=["team-a"])
assert team_attachment.is_global() is False
def test_to_policy_scope_global(self):
"""Test converting global attachment to PolicyScope."""
attachment = PolicyAttachment(policy="test", scope="*")
scope = attachment.to_policy_scope()
assert scope.get_teams() == ["*"]
assert scope.get_keys() == ["*"]
assert scope.get_models() == ["*"]
def test_to_policy_scope_specific(self):
"""Test converting specific attachment to PolicyScope."""
attachment = PolicyAttachment(
policy="test",
teams=["team-a", "team-b"],
keys=["key-*"],
models=["gpt-4"],
)
scope = attachment.to_policy_scope()
assert scope.teams == ["team-a", "team-b"]
assert scope.keys == ["key-*"]
assert scope.models == ["gpt-4"]
class TestGlobalSingleton:
"""Test global singleton behavior."""
def test_get_attachment_registry_singleton(self):
"""Test get_attachment_registry returns same instance."""
registry1 = get_attachment_registry()
registry2 = get_attachment_registry()
assert registry1 is registry2

View file

@ -0,0 +1,289 @@
"""
Unit tests for ConditionEvaluator - tests AWS IAM-style condition evaluation.
Tests:
- Condition operators (equals, in, prefix, not_equals, not_in)
- PolicyCondition evaluation against request context
- Metadata condition evaluation
"""
import pytest
from litellm.proxy.policy_engine.condition_evaluator import ConditionEvaluator
from litellm.types.proxy.policy_engine import (
ConditionOperator,
PolicyCondition,
PolicyMatchContext,
)
class TestConditionOperatorEvaluation:
"""Test individual condition operator evaluation."""
def test_equals_operator_match(self):
"""Test equals operator matches exact value."""
operator = ConditionOperator(equals="gpt-4")
assert ConditionEvaluator.evaluate_operator(operator, "gpt-4") is True
def test_equals_operator_no_match(self):
"""Test equals operator does not match different value."""
operator = ConditionOperator(equals="gpt-4")
assert ConditionEvaluator.evaluate_operator(operator, "gpt-3.5") is False
def test_in_operator_match(self):
"""Test in operator matches value in list."""
operator = ConditionOperator(in_=["gpt-4", "gpt-4-turbo", "gpt-4o"])
assert ConditionEvaluator.evaluate_operator(operator, "gpt-4") is True
assert ConditionEvaluator.evaluate_operator(operator, "gpt-4-turbo") is True
def test_in_operator_no_match(self):
"""Test in operator does not match value not in list."""
operator = ConditionOperator(in_=["gpt-4", "gpt-4-turbo"])
assert ConditionEvaluator.evaluate_operator(operator, "gpt-3.5") is False
def test_prefix_operator_match(self):
"""Test prefix operator matches value starting with prefix."""
operator = ConditionOperator(prefix="bedrock/")
assert ConditionEvaluator.evaluate_operator(operator, "bedrock/claude-3") is True
assert ConditionEvaluator.evaluate_operator(operator, "bedrock/llama") is True
def test_prefix_operator_no_match(self):
"""Test prefix operator does not match value not starting with prefix."""
operator = ConditionOperator(prefix="bedrock/")
assert ConditionEvaluator.evaluate_operator(operator, "openai/gpt-4") is False
def test_not_equals_operator_match(self):
"""Test not_equals operator matches when value is different."""
operator = ConditionOperator(not_equals="gpt-3.5")
assert ConditionEvaluator.evaluate_operator(operator, "gpt-4") is True
def test_not_equals_operator_no_match(self):
"""Test not_equals operator does not match when value is same."""
operator = ConditionOperator(not_equals="gpt-4")
assert ConditionEvaluator.evaluate_operator(operator, "gpt-4") is False
def test_not_in_operator_match(self):
"""Test not_in operator matches when value not in list."""
operator = ConditionOperator(not_in=["gpt-3.5", "gpt-3.5-turbo"])
assert ConditionEvaluator.evaluate_operator(operator, "gpt-4") is True
def test_not_in_operator_no_match(self):
"""Test not_in operator does not match when value in list."""
operator = ConditionOperator(not_in=["gpt-4", "gpt-4-turbo"])
assert ConditionEvaluator.evaluate_operator(operator, "gpt-4") is False
def test_none_value_with_positive_operators(self):
"""Test None value does not match positive operators."""
assert ConditionEvaluator.evaluate_operator(
ConditionOperator(equals="gpt-4"), None
) is False
assert ConditionEvaluator.evaluate_operator(
ConditionOperator(in_=["gpt-4"]), None
) is False
assert ConditionEvaluator.evaluate_operator(
ConditionOperator(prefix="gpt"), None
) is False
def test_none_value_with_negative_operators(self):
"""Test None value matches negative operators (None is not equal to anything)."""
assert ConditionEvaluator.evaluate_operator(
ConditionOperator(not_equals="gpt-4"), None
) is True
assert ConditionEvaluator.evaluate_operator(
ConditionOperator(not_in=["gpt-4"]), None
) is True
def test_empty_operator_matches_any(self):
"""Test empty operator (no conditions) matches any value."""
operator = ConditionOperator()
assert ConditionEvaluator.evaluate_operator(operator, "anything") is True
assert ConditionEvaluator.evaluate_operator(operator, None) is True
class TestPolicyConditionEvaluation:
"""Test PolicyCondition evaluation against request context."""
def test_model_condition_match(self):
"""Test model condition matches."""
condition = PolicyCondition(
model=ConditionOperator(in_=["gpt-4", "gpt-4-turbo"])
)
context = PolicyMatchContext(team_alias="team", key_alias="key", model="gpt-4")
assert ConditionEvaluator.evaluate(condition, context) is True
def test_model_condition_no_match(self):
"""Test model condition does not match."""
condition = PolicyCondition(
model=ConditionOperator(in_=["gpt-4", "gpt-4-turbo"])
)
context = PolicyMatchContext(team_alias="team", key_alias="key", model="gpt-3.5")
assert ConditionEvaluator.evaluate(condition, context) is False
def test_team_condition_match(self):
"""Test team condition matches."""
condition = PolicyCondition(
team=ConditionOperator(prefix="healthcare-")
)
context = PolicyMatchContext(
team_alias="healthcare-research", key_alias="key", model="gpt-4"
)
assert ConditionEvaluator.evaluate(condition, context) is True
def test_team_condition_no_match(self):
"""Test team condition does not match."""
condition = PolicyCondition(
team=ConditionOperator(prefix="healthcare-")
)
context = PolicyMatchContext(
team_alias="finance-team", key_alias="key", model="gpt-4"
)
assert ConditionEvaluator.evaluate(condition, context) is False
def test_key_condition_match(self):
"""Test key condition matches."""
condition = PolicyCondition(
key=ConditionOperator(equals="production-key")
)
context = PolicyMatchContext(
team_alias="team", key_alias="production-key", model="gpt-4"
)
assert ConditionEvaluator.evaluate(condition, context) is True
def test_multiple_conditions_all_match(self):
"""Test multiple conditions all must match (AND logic)."""
condition = PolicyCondition(
model=ConditionOperator(in_=["gpt-4", "gpt-4-turbo"]),
team=ConditionOperator(prefix="healthcare-"),
)
context = PolicyMatchContext(
team_alias="healthcare-research", key_alias="key", model="gpt-4"
)
assert ConditionEvaluator.evaluate(condition, context) is True
def test_multiple_conditions_one_fails(self):
"""Test multiple conditions - if one fails, all fails."""
condition = PolicyCondition(
model=ConditionOperator(in_=["gpt-4", "gpt-4-turbo"]),
team=ConditionOperator(prefix="healthcare-"),
)
# Model matches but team doesn't
context = PolicyMatchContext(
team_alias="finance-team", key_alias="key", model="gpt-4"
)
assert ConditionEvaluator.evaluate(condition, context) is False
def test_none_condition_always_matches(self):
"""Test None condition always matches."""
context = PolicyMatchContext(team_alias="any", key_alias="any", model="any")
assert ConditionEvaluator.evaluate(None, context) is True
class TestMetadataConditionEvaluation:
"""Test metadata condition evaluation."""
def test_metadata_condition_match(self):
"""Test metadata condition matches."""
condition = PolicyCondition(
metadata={
"environment": ConditionOperator(equals="production"),
}
)
context = PolicyMatchContext(team_alias="team", key_alias="key", model="gpt-4")
metadata = {"environment": "production"}
assert ConditionEvaluator.evaluate(condition, context, metadata) is True
def test_metadata_condition_no_match(self):
"""Test metadata condition does not match."""
condition = PolicyCondition(
metadata={
"environment": ConditionOperator(equals="production"),
}
)
context = PolicyMatchContext(team_alias="team", key_alias="key", model="gpt-4")
metadata = {"environment": "staging"}
assert ConditionEvaluator.evaluate(condition, context, metadata) is False
def test_metadata_condition_missing_field(self):
"""Test metadata condition with missing field does not match."""
condition = PolicyCondition(
metadata={
"environment": ConditionOperator(equals="production"),
}
)
context = PolicyMatchContext(team_alias="team", key_alias="key", model="gpt-4")
metadata = {"other_field": "value"}
assert ConditionEvaluator.evaluate(condition, context, metadata) is False
def test_metadata_condition_with_model_condition(self):
"""Test combining metadata and model conditions."""
condition = PolicyCondition(
model=ConditionOperator(in_=["gpt-4"]),
metadata={
"environment": ConditionOperator(equals="production"),
},
)
context = PolicyMatchContext(team_alias="team", key_alias="key", model="gpt-4")
metadata = {"environment": "production"}
assert ConditionEvaluator.evaluate(condition, context, metadata) is True
# Model matches but metadata doesn't
metadata_staging = {"environment": "staging"}
assert ConditionEvaluator.evaluate(condition, context, metadata_staging) is False
class TestEvaluateAllConditions:
"""Test evaluate_all_conditions helper."""
def test_all_conditions_match(self):
"""Test all conditions match returns True."""
conditions = [
PolicyCondition(model=ConditionOperator(equals="gpt-4")),
PolicyCondition(team=ConditionOperator(prefix="healthcare-")),
]
context = PolicyMatchContext(
team_alias="healthcare-team", key_alias="key", model="gpt-4"
)
assert ConditionEvaluator.evaluate_all_conditions(conditions, context) is True
def test_one_condition_fails(self):
"""Test one condition fails returns False."""
conditions = [
PolicyCondition(model=ConditionOperator(equals="gpt-4")),
PolicyCondition(team=ConditionOperator(prefix="healthcare-")),
]
context = PolicyMatchContext(
team_alias="finance-team", key_alias="key", model="gpt-4"
)
assert ConditionEvaluator.evaluate_all_conditions(conditions, context) is False
def test_empty_conditions_returns_true(self):
"""Test empty conditions list returns True."""
context = PolicyMatchContext(team_alias="any", key_alias="any", model="any")
assert ConditionEvaluator.evaluate_all_conditions([], context) is True
class TestEvaluateAnyCondition:
"""Test evaluate_any_condition helper."""
def test_any_condition_matches(self):
"""Test any condition matches returns True."""
conditions = [
PolicyCondition(model=ConditionOperator(equals="gpt-4")),
PolicyCondition(model=ConditionOperator(equals="gpt-3.5")),
]
context = PolicyMatchContext(team_alias="team", key_alias="key", model="gpt-4")
assert ConditionEvaluator.evaluate_any_condition(conditions, context) is True
def test_no_condition_matches(self):
"""Test no condition matches returns False."""
conditions = [
PolicyCondition(model=ConditionOperator(equals="gpt-4")),
PolicyCondition(model=ConditionOperator(equals="gpt-3.5")),
]
context = PolicyMatchContext(team_alias="team", key_alias="key", model="claude-3")
assert ConditionEvaluator.evaluate_any_condition(conditions, context) is False
def test_empty_conditions_returns_true(self):
"""Test empty conditions list returns True."""
context = PolicyMatchContext(team_alias="any", key_alias="any", model="any")
assert ConditionEvaluator.evaluate_any_condition([], context) is True

View file

@ -1,84 +1,96 @@
"""
Unit tests for PolicyMatcher - tests wildcard pattern matching for policies.
Unit tests for PolicyMatcher - tests wildcard pattern matching via attachments.
Tests:
- Wildcard matching (*, prefix-*)
- Scope matching (teams, keys, models)
- Scope matching via attachments (teams, keys, models)
"""
import pytest
from litellm.proxy.policy_engine.attachment_registry import AttachmentRegistry
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
from litellm.types.proxy.policy_engine import (
Policy,
PolicyGuardrails,
PolicyMatchContext,
PolicyScope,
)
class TestPolicyMatcherGetMatchingPolicies:
"""Test getting matching policies from a set of policies."""
class TestPolicyMatcherPatternMatching:
"""Test pattern matching utilities."""
def test_get_matching_policies_by_team(self):
"""Test matching policies by team alias."""
policies = {
"healthcare": Policy(
guardrails=PolicyGuardrails(add=["hipaa_audit"]),
scope=PolicyScope(teams=["healthcare-team"]),
),
}
def test_matches_pattern_exact(self):
"""Test exact pattern matching."""
assert PolicyMatcher.matches_pattern("healthcare-team", ["healthcare-team"]) is True
assert PolicyMatcher.matches_pattern("finance-team", ["healthcare-team"]) is False
# Match
context = PolicyMatchContext(team_alias="healthcare-team", key_alias="k", model="gpt-4")
assert "healthcare" in PolicyMatcher.get_matching_policies(policies=policies, context=context)
def test_matches_pattern_wildcard(self):
"""Test wildcard pattern matching."""
assert PolicyMatcher.matches_pattern("any-team", ["*"]) is True
assert PolicyMatcher.matches_pattern("dev-key-123", ["dev-key-*"]) is True
assert PolicyMatcher.matches_pattern("prod-key-123", ["dev-key-*"]) is False
# No match
context = PolicyMatchContext(team_alias="finance-team", key_alias="k", model="gpt-4")
assert len(PolicyMatcher.get_matching_policies(policies=policies, context=context)) == 0
def test_matches_pattern_none_value(self):
"""Test None value only matches '*'."""
assert PolicyMatcher.matches_pattern(None, ["*"]) is True
assert PolicyMatcher.matches_pattern(None, ["specific"]) is False
def test_get_matching_policies_by_model_wildcard(self):
"""Test matching policies by model with wildcard pattern."""
policies = {
"bedrock-only": Policy(
guardrails=PolicyGuardrails(add=["pii_blocker"]),
scope=PolicyScope(models=["bedrock/*"]),
),
}
# Match - bedrock model
context = PolicyMatchContext(team_alias="t", key_alias="k", model="bedrock/claude-3")
assert "bedrock-only" in PolicyMatcher.get_matching_policies(policies=policies, context=context)
class TestPolicyMatcherScopeMatching:
"""Test scope matching against context."""
# No match - different provider
context = PolicyMatchContext(team_alias="t", key_alias="k", model="openai/gpt-4")
assert len(PolicyMatcher.get_matching_policies(policies=policies, context=context)) == 0
def test_scope_matches_all_fields(self):
"""Test scope matches when all fields match."""
scope = PolicyScope(teams=["healthcare-team"], keys=["*"], models=["gpt-4"])
context = PolicyMatchContext(team_alias="healthcare-team", key_alias="any-key", model="gpt-4")
assert PolicyMatcher.scope_matches(scope, context) is True
def test_get_matching_policies_by_key_pattern(self):
"""Test matching policies by key alias pattern."""
policies = {
"dev-keys": Policy(
guardrails=PolicyGuardrails(add=["toxicity_filter"]),
scope=PolicyScope(keys=["dev-key-*"]),
),
}
def test_scope_does_not_match_team(self):
"""Test scope doesn't match when team doesn't match."""
scope = PolicyScope(teams=["healthcare-team"], keys=["*"], models=["*"])
context = PolicyMatchContext(team_alias="finance-team", key_alias="any-key", model="gpt-4")
assert PolicyMatcher.scope_matches(scope, context) is False
# Match
context = PolicyMatchContext(team_alias="t", key_alias="dev-key-123", model="gpt-4")
assert "dev-keys" in PolicyMatcher.get_matching_policies(policies=policies, context=context)
# No match
context = PolicyMatchContext(team_alias="t", key_alias="prod-key-123", model="gpt-4")
assert len(PolicyMatcher.get_matching_policies(policies=policies, context=context)) == 0
def test_get_matching_policies_global_wildcard(self):
"""Test global policy with '*' matches everything."""
policies = {
"global": Policy(
guardrails=PolicyGuardrails(add=["pii_blocker"]),
scope=PolicyScope(teams=["*"], keys=["*"], models=["*"]),
),
}
def test_scope_matches_with_wildcard_patterns(self):
"""Test scope matches with wildcard patterns."""
scope = PolicyScope(teams=["*"], keys=["dev-key-*"], models=["bedrock/*"])
context = PolicyMatchContext(team_alias="any-team", key_alias="dev-key-123", model="bedrock/claude-3")
assert PolicyMatcher.scope_matches(scope, context) is True
def test_scope_global_wildcard(self):
"""Test global scope with all wildcards."""
scope = PolicyScope(teams=["*"], keys=["*"], models=["*"])
context = PolicyMatchContext(team_alias="any-team", key_alias="any-key", model="any-model")
assert "global" in PolicyMatcher.get_matching_policies(policies=policies, context=context)
assert PolicyMatcher.scope_matches(scope, context) is True
class TestPolicyMatcherWithAttachments:
"""Test getting matching policies via attachments."""
def test_get_matching_policies_via_attachments(self):
"""Test matching policies through attachment registry."""
# Create and configure attachment registry
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "healthcare-policy", "teams": ["healthcare-team"]},
{"policy": "global-policy", "scope": "*"},
])
# Test matching via the registry directly
context = PolicyMatchContext(team_alias="healthcare-team", key_alias="k", model="gpt-4")
attached = registry.get_attached_policies(context)
assert "healthcare-policy" in attached
assert "global-policy" in attached
def test_get_matching_policies_no_match(self):
"""Test no policies match when attachments don't match context."""
registry = AttachmentRegistry()
registry.load_attachments([
{"policy": "healthcare-policy", "teams": ["healthcare-team"]},
])
context = PolicyMatchContext(team_alias="finance-team", key_alias="k", model="gpt-4")
attached = registry.get_attached_policies(context)
assert "healthcare-policy" not in attached

View file

@ -1,96 +1,309 @@
"""
Unit tests for PolicyResolver - tests guardrail resolution for request contexts.
Unit tests for PolicyResolver - tests guardrail resolution.
Tests:
- Inheritance chain resolution
- Inheritance with add/remove
- Conditional statements
"""
import pytest
from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
from litellm.types.proxy.policy_engine import (
ConditionOperator,
Policy,
PolicyCondition,
PolicyGuardrails,
PolicyMatchContext,
PolicyScope,
PolicyStatement,
)
class TestPolicyMatcherGetMatchingPolicies:
"""Test resolve_guardrails_for_context - the main entry point."""
class TestPolicyResolverInheritance:
"""Test resolve_policy_guardrails - inheritance and add/remove."""
def test_resolve_guardrails_simple_match(self):
"""Test resolving guardrails for a simple matching policy."""
def test_resolve_simple_policy(self):
"""Test resolving guardrails for a simple policy."""
policies = {
"global": Policy(
guardrails=PolicyGuardrails(add=["pii_blocker", "toxicity_filter"]),
scope=PolicyScope(teams=["*"]),
),
}
context = PolicyMatchContext(team_alias="any-team", key_alias="k", model="gpt-4")
guardrails = PolicyResolver.resolve_guardrails_for_context(
context=context, policies=policies
resolved = PolicyResolver.resolve_policy_guardrails(
policy_name="global", policies=policies
)
assert set(guardrails) == {"pii_blocker", "toxicity_filter"}
assert set(resolved.guardrails) == {"pii_blocker", "toxicity_filter"}
assert resolved.inheritance_chain == ["global"]
def test_resolve_guardrails_with_inheritance(self):
def test_resolve_with_inheritance(self):
"""Test child policy inherits and adds guardrails from parent."""
policies = {
"base": Policy(
guardrails=PolicyGuardrails(add=["pii_blocker"]),
scope=PolicyScope(teams=["*"]),
),
"healthcare": Policy(
inherit="base",
guardrails=PolicyGuardrails(add=["hipaa_audit"]),
scope=PolicyScope(teams=["healthcare-team"]),
),
}
context = PolicyMatchContext(team_alias="healthcare-team", key_alias="k", model="gpt-4")
guardrails = PolicyResolver.resolve_guardrails_for_context(
context=context, policies=policies
resolved = PolicyResolver.resolve_policy_guardrails(
policy_name="healthcare", policies=policies
)
# Both base and healthcare match, healthcare inherits from base
assert set(guardrails) == {"pii_blocker", "hipaa_audit"}
# Healthcare inherits pii_blocker from base and adds hipaa_audit
assert set(resolved.guardrails) == {"pii_blocker", "hipaa_audit"}
assert resolved.inheritance_chain == ["base", "healthcare"]
def test_resolve_guardrails_with_remove(self):
"""Test child policy can remove guardrails from parent in its inheritance chain."""
def test_resolve_with_remove(self):
"""Test child policy can remove guardrails from parent."""
policies = {
"base": Policy(
guardrails=PolicyGuardrails(add=["pii_blocker", "phi_blocker"]),
scope=PolicyScope(teams=["internal-only"]), # Does NOT match dev-team
),
"dev": Policy(
inherit="base",
guardrails=PolicyGuardrails(add=["toxicity_filter"], remove=["phi_blocker"]),
scope=PolicyScope(teams=["dev-team"]), # Only this matches
),
}
# Only dev policy matches (base scope doesn't match)
context = PolicyMatchContext(team_alias="dev-team", key_alias="k", model="gpt-4")
guardrails = PolicyResolver.resolve_guardrails_for_context(
context=context, policies=policies
resolved = PolicyResolver.resolve_policy_guardrails(
policy_name="dev", policies=policies
)
# dev inherits pii_blocker from base, adds toxicity_filter, removes phi_blocker
assert "pii_blocker" in guardrails
assert "toxicity_filter" in guardrails
assert "phi_blocker" not in guardrails
assert "pii_blocker" in resolved.guardrails
assert "toxicity_filter" in resolved.guardrails
assert "phi_blocker" not in resolved.guardrails
def test_resolve_guardrails_no_match(self):
"""Test returns empty list when no policies match."""
def test_resolve_deep_inheritance_chain(self):
"""Test multi-level inheritance chain."""
policies = {
"healthcare": Policy(
guardrails=PolicyGuardrails(add=["hipaa_audit"]),
scope=PolicyScope(teams=["healthcare-team"]),
"root": Policy(
guardrails=PolicyGuardrails(add=["root_guardrail"]),
),
"middle": Policy(
inherit="root",
guardrails=PolicyGuardrails(add=["middle_guardrail"]),
),
"leaf": Policy(
inherit="middle",
guardrails=PolicyGuardrails(add=["leaf_guardrail"]),
),
}
context = PolicyMatchContext(team_alias="finance-team", key_alias="k", model="gpt-4")
guardrails = PolicyResolver.resolve_guardrails_for_context(
context=context, policies=policies
resolved = PolicyResolver.resolve_policy_guardrails(
policy_name="leaf", policies=policies
)
assert guardrails == []
assert set(resolved.guardrails) == {"root_guardrail", "middle_guardrail", "leaf_guardrail"}
assert resolved.inheritance_chain == ["root", "middle", "leaf"]
class TestPolicyResolverWithStatements:
"""Test resolve_policy_guardrails with conditional statements."""
def test_statement_condition_matches(self):
"""Test statement guardrails are added when condition matches."""
policies = {
"conditional-policy": Policy(
guardrails=PolicyGuardrails(add=["base_guardrail"]),
statements=[
PolicyStatement(
sid="GPT4Safety",
guardrails=["toxicity_filter"],
condition=PolicyCondition(
model=ConditionOperator(in_=["gpt-4", "gpt-4-turbo"])
),
),
],
),
}
# GPT-4 should get both base and statement guardrails
context = PolicyMatchContext(team_alias="team", key_alias="k", model="gpt-4")
resolved = PolicyResolver.resolve_policy_guardrails(
policy_name="conditional-policy",
policies=policies,
context=context,
)
assert "base_guardrail" in resolved.guardrails
assert "toxicity_filter" in resolved.guardrails
def test_statement_condition_does_not_match(self):
"""Test statement guardrails are NOT added when condition doesn't match."""
policies = {
"conditional-policy": Policy(
guardrails=PolicyGuardrails(add=["base_guardrail"]),
statements=[
PolicyStatement(
sid="GPT4Safety",
guardrails=["toxicity_filter"],
condition=PolicyCondition(
model=ConditionOperator(in_=["gpt-4", "gpt-4-turbo"])
),
),
],
),
}
# GPT-3.5 should only get base guardrails, not statement guardrails
context = PolicyMatchContext(team_alias="team", key_alias="k", model="gpt-3.5")
resolved = PolicyResolver.resolve_policy_guardrails(
policy_name="conditional-policy",
policies=policies,
context=context,
)
assert "base_guardrail" in resolved.guardrails
assert "toxicity_filter" not in resolved.guardrails
def test_multiple_statements_some_match(self):
"""Test multiple statements where only some match."""
policies = {
"multi-statement": Policy(
guardrails=PolicyGuardrails(add=["base"]),
statements=[
PolicyStatement(
sid="GPT4Only",
guardrails=["gpt4_guardrail"],
condition=PolicyCondition(
model=ConditionOperator(equals="gpt-4")
),
),
PolicyStatement(
sid="HealthcareOnly",
guardrails=["hipaa_audit"],
condition=PolicyCondition(
team=ConditionOperator(prefix="healthcare-")
),
),
],
),
}
# Healthcare team with GPT-4 should get all guardrails
context = PolicyMatchContext(
team_alias="healthcare-team", key_alias="k", model="gpt-4"
)
resolved = PolicyResolver.resolve_policy_guardrails(
policy_name="multi-statement",
policies=policies,
context=context,
)
assert "base" in resolved.guardrails
assert "gpt4_guardrail" in resolved.guardrails
assert "hipaa_audit" in resolved.guardrails
# Finance team with GPT-4 should only get base + gpt4_guardrail
context_finance = PolicyMatchContext(
team_alias="finance-team", key_alias="k", model="gpt-4"
)
resolved_finance = PolicyResolver.resolve_policy_guardrails(
policy_name="multi-statement",
policies=policies,
context=context_finance,
)
assert "base" in resolved_finance.guardrails
assert "gpt4_guardrail" in resolved_finance.guardrails
assert "hipaa_audit" not in resolved_finance.guardrails
def test_statement_with_no_condition_always_applies(self):
"""Test statement with no condition always applies."""
policies = {
"always-policy": Policy(
guardrails=PolicyGuardrails(add=["base"]),
statements=[
PolicyStatement(
sid="AlwaysApply",
guardrails=["always_guardrail"],
condition=None, # No condition = always applies
),
],
),
}
context = PolicyMatchContext(team_alias="any", key_alias="any", model="any")
resolved = PolicyResolver.resolve_policy_guardrails(
policy_name="always-policy",
policies=policies,
context=context,
)
assert "base" in resolved.guardrails
assert "always_guardrail" in resolved.guardrails
def test_inheritance_with_statements(self):
"""Test inheritance works with statements."""
policies = {
"base": Policy(
guardrails=PolicyGuardrails(add=["pii_blocker"]),
),
"child": Policy(
inherit="base",
guardrails=PolicyGuardrails(add=["child_guardrail"]),
statements=[
PolicyStatement(
sid="ConditionalStatement",
guardrails=["conditional_guardrail"],
condition=PolicyCondition(
model=ConditionOperator(equals="gpt-4")
),
),
],
),
}
context = PolicyMatchContext(team_alias="any-team", key_alias="k", model="gpt-4")
resolved = PolicyResolver.resolve_policy_guardrails(
policy_name="child",
policies=policies,
context=context,
)
# Should have: inherited pii_blocker, child's child_guardrail, and conditional_guardrail
assert "pii_blocker" in resolved.guardrails
assert "child_guardrail" in resolved.guardrails
assert "conditional_guardrail" in resolved.guardrails
def test_inheritance_with_remove_and_statements(self):
"""Test inheritance with remove still works alongside statements."""
policies = {
"base": Policy(
guardrails=PolicyGuardrails(add=["pii_blocker", "phi_blocker"]),
),
"child": Policy(
inherit="base",
guardrails=PolicyGuardrails(
add=["child_guardrail"],
remove=["phi_blocker"], # Remove phi_blocker from parent
),
statements=[
PolicyStatement(
sid="Conditional",
guardrails=["conditional_guardrail"],
condition=PolicyCondition(
model=ConditionOperator(equals="gpt-4")
),
),
],
),
}
context = PolicyMatchContext(team_alias="any-team", key_alias="k", model="gpt-4")
resolved = PolicyResolver.resolve_policy_guardrails(
policy_name="child",
policies=policies,
context=context,
)
assert "pii_blocker" in resolved.guardrails # Inherited
assert "phi_blocker" not in resolved.guardrails # Removed
assert "child_guardrail" in resolved.guardrails # Added by child
assert "conditional_guardrail" in resolved.guardrails # From statement

View file

@ -4,7 +4,6 @@ Unit tests for PolicyValidator - tests policy configuration validation.
Tests validation of:
- Inheritance chains (parent exists, no circular deps)
- Guardrail names exist in registry
- Model names exist in router
"""
from unittest.mock import MagicMock, patch
@ -15,7 +14,6 @@ from litellm.proxy.policy_engine.policy_validator import PolicyValidator
from litellm.types.proxy.policy_engine import (
Policy,
PolicyGuardrails,
PolicyScope,
PolicyValidationErrorType,
)
@ -30,7 +28,6 @@ class TestPolicyValidator:
"child": Policy(
inherit="nonexistent-parent",
guardrails=PolicyGuardrails(add=["hipaa_audit"]),
scope=PolicyScope(teams=["healthcare-team"]),
),
}
@ -49,7 +46,6 @@ class TestPolicyValidator:
policies = {
"test-policy": Policy(
guardrails=PolicyGuardrails(add=["nonexistent_guardrail"]),
scope=PolicyScope(teams=["*"]),
),
}
@ -67,28 +63,23 @@ class TestPolicyValidator:
)
@pytest.mark.asyncio
async def test_validate_invalid_model(self):
"""Test that referencing non-existent model warns."""
async def test_validate_valid_policy(self):
"""Test that a valid policy passes validation."""
policies = {
"test-policy": Policy(
"base": Policy(
guardrails=PolicyGuardrails(add=["pii_blocker"]),
scope=PolicyScope(models=["nonexistent-model"]),
),
"child": Policy(
inherit="base",
guardrails=PolicyGuardrails(add=["toxicity_filter"]),
),
}
# Mock the router with known model names
mock_router = MagicMock()
mock_router.model_names = {"gpt-4", "gpt-3.5-turbo"}
# Mock pattern_router to return empty list (no pattern matches)
mock_router.pattern_router.get_deployments_by_pattern.return_value = []
validator = PolicyValidator(prisma_client=None, llm_router=mock_router)
with patch.object(validator, "get_available_guardrails", return_value={"pii_blocker"}):
validator = PolicyValidator(prisma_client=None)
with patch.object(
validator, "get_available_guardrails", return_value={"pii_blocker", "toxicity_filter"}
):
result = await validator.validate_policies(policies=policies, validate_db=False)
# Model validation is a warning, not an error
assert any(
w.error_type == PolicyValidationErrorType.INVALID_MODEL
and w.value == "nonexistent-model"
for w in result.warnings
)
assert result.valid is True
assert len(result.errors) == 0