fix(policy_engine): union list-gated stages across all callbacks sharing a guardrail name

This commit is contained in:
mateo-berri 2026-08-31 16:34:29 -07:00
parent 55e71a53d2
commit 808b5e90c5
3 changed files with 98 additions and 16 deletions

View file

@ -3012,7 +3012,8 @@ def _apply_resolved_guardrails_to_metadata(
# Combine existing guardrails with policy-resolved guardrails (no duplicates).
# Drop a name only when pipelines cover every stage it supports: a pipeline
# manages just its own mode, so the guardrail must stay in the flat list for
# its stages outside that mode (mode-scoped execution skips prevent doubles).
# its stages outside that mode. The executor's pre_call skips keyed on
# _pipeline_managed_guardrails prevent double runs at that stage.
combined = set(existing_guardrails)
combined.update(resolved_guardrails)
combined -= fully_pipeline_covered_guardrails

View file

@ -11,6 +11,7 @@ Handles:
from collections.abc import Sequence
from typing import Final
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.policy_engine import (
@ -29,21 +30,7 @@ def _plain_stage(hook: object) -> str | None:
return None
def _list_gated_stages_for_guardrail(guardrail_name: str) -> frozenset[str] | None:
"""
The lifecycle stages where this guardrail only runs if its name is in the
request's flat guardrails list, or None when they cannot be determined
statically (enterprise tag-based Mode hooks, whose stages depend on
request tags).
An unregistered name or an event_hook of None yields an empty set: the
flat list never gates such a guardrail, so dropping the name from the
list cannot suppress anything.
"""
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
callback: Final = PipelineExecutor.find_guardrail_callback(guardrail_name)
event_hook: Final = callback.event_hook if callback is not None else None
def _gated_stages_for_event_hook(event_hook: object) -> frozenset[str] | None:
if event_hook is None:
return frozenset()
hooks: Final = tuple(event_hook) if isinstance(event_hook, list) else (event_hook,)
@ -53,6 +40,35 @@ def _list_gated_stages_for_guardrail(guardrail_name: str) -> frozenset[str] | No
return frozenset(stage for stage in stages if stage is not None and stage != GuardrailEventHooks.logging_only.value)
def _list_gated_stages_for_guardrail(guardrail_name: str) -> frozenset[str] | None:
"""
The lifecycle stages where this guardrail only runs if its name is in the
request's flat guardrails list, or None when they cannot be determined
statically (enterprise tag-based Mode hooks, whose stages depend on
request tags).
Stages are unioned across every registered callback carrying the name:
one guardrail_name can map to several callbacks with different hooks
(e.g. Presidio registers a post_call output-masking sibling alongside the
configured one, and duplicate-name deployments are supported for load
balancing), and stripping the name gates all of them.
An unregistered name or an event_hook of None yields an empty set: the
flat list never gates such a guardrail, so dropping the name from the
list cannot suppress anything.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
stage_sets: Final = tuple(
_gated_stages_for_event_hook(callback.event_hook)
for callback in litellm.callbacks
if isinstance(callback, CustomGuardrail) and callback.guardrail_name == guardrail_name
)
if any(stages is None for stages in stage_sets):
return None
return frozenset(stage for stages in stage_sets if stages is not None for stage in stages)
class PolicyResolver:
"""
Resolves the final list of guardrails from policies.

View file

@ -4397,6 +4397,71 @@ async def test_pipeline_keeps_guardrail_with_stage_no_pipeline_can_cover(monkeyp
assert data["metadata"]["guardrails"] == ["word_guard"]
@pytest.mark.asyncio
async def test_pipeline_keeps_guardrail_when_sibling_callback_adds_stages(monkeypatch):
"""
One guardrail_name can map to several registered callbacks (Presidio adds
a post_call output-masking sibling; duplicate-name deployments are a
supported load-balancing setup). Coverage must union stages across all of
them, so a pre_call pipeline does not strip a name whose sibling still
needs the flat list at post_call.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
input_callback = CustomGuardrail(guardrail_name="word_guard", event_hook="pre_call")
output_sibling = CustomGuardrail(guardrail_name="word_guard", event_hook=GuardrailEventHooks.post_call)
policy_registry, attachment_registry = _policy_engine_pipeline_registries(
{"input-pipeline": _word_guard_pipeline_policy("pre_call")},
monkeypatch,
callbacks=[input_callback, output_sibling],
)
data = {"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], "metadata": {}}
try:
await add_guardrails_from_policy_engine(
data=data,
metadata_variable_name="metadata",
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
)
finally:
_reset_policy_engine_registries(policy_registry, attachment_registry)
assert data["metadata"]["guardrails"] == ["word_guard"]
@pytest.mark.asyncio
async def test_pipeline_keeps_guardrail_with_tag_based_mode(monkeypatch):
"""
A tag-based Mode event_hook resolves its stages per request, so coverage
cannot be determined statically and the name must stay in the flat list.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import Mode
guardrail = CustomGuardrail(
guardrail_name="word_guard",
event_hook=Mode(tags={"team": "security"}, default="pre_call"),
)
policy_registry, attachment_registry = _policy_engine_pipeline_registries(
{"input-pipeline": _word_guard_pipeline_policy("pre_call")},
monkeypatch,
callbacks=[guardrail],
)
data = {"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], "metadata": {}}
try:
await add_guardrails_from_policy_engine(
data=data,
metadata_variable_name="metadata",
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
)
finally:
_reset_policy_engine_registries(policy_registry, attachment_registry)
assert data["metadata"]["guardrails"] == ["word_guard"]
@pytest.mark.asyncio
async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_pops_from_data():
"""