fix(policy_engine): keep cross-mode guardrail stages active when a pipeline manages one mode

This commit is contained in:
mateo-berri 2026-08-31 16:03:36 -07:00
parent 5a0ed05765
commit e18c19f7ab
3 changed files with 226 additions and 3 deletions

View file

@ -2991,6 +2991,7 @@ def _apply_resolved_guardrails_to_metadata(
# Track pipeline-managed guardrails to exclude from independent execution
pipeline_managed_guardrails: set = set()
fully_pipeline_covered_guardrails: Final = PolicyResolver.get_guardrails_fully_covered_by_pipelines(pipelines)
if pipelines:
pipeline_managed_guardrails = PolicyResolver.get_pipeline_managed_guardrails(pipelines)
data[metadata_variable_name]["_guardrail_pipelines"] = pipelines
@ -3008,11 +3009,13 @@ def _apply_resolved_guardrails_to_metadata(
if not isinstance(existing_guardrails, list):
existing_guardrails = []
# Combine existing guardrails with policy-resolved guardrails (no duplicates)
# Exclude pipeline-managed guardrails from the flat list
# 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).
combined = set(existing_guardrails)
combined.update(resolved_guardrails)
combined -= pipeline_managed_guardrails
combined -= fully_pipeline_covered_guardrails
data[metadata_variable_name]["guardrails"] = list(combined)
verbose_proxy_logger.debug("Policy engine: added guardrails to request metadata: %s", list(combined))

View file

@ -8,9 +8,11 @@ Handles:
- Combining guardrails from multiple matching policies
"""
from collections.abc import Sequence
from typing import Final
from litellm._logging import verbose_proxy_logger
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.policy_engine import (
GuardrailPipeline,
Policy,
@ -19,6 +21,38 @@ from litellm.types.proxy.policy_engine import (
)
def _plain_stage(hook: object) -> str | None:
if isinstance(hook, GuardrailEventHooks):
return hook.value
if isinstance(hook, str):
return hook
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
if event_hook is None:
return frozenset()
hooks: Final = tuple(event_hook) if isinstance(event_hook, list) else (event_hook,)
stages: Final = tuple(_plain_stage(hook) for hook in hooks)
if any(stage is None for stage in stages):
return None
return frozenset(stage for stage in stages if stage is not None and stage != GuardrailEventHooks.logging_only.value)
class PolicyResolver:
"""
Resolves the final list of guardrails from policies.
@ -253,6 +287,34 @@ class PolicyResolver:
managed.add(step.guardrail)
return managed
@staticmethod
def get_guardrails_fully_covered_by_pipelines(
pipelines: Sequence[tuple[str, GuardrailPipeline]],
) -> frozenset[str]:
"""
Guardrail names whose every list-gated stage is covered by the mode of
a pipeline naming them.
Only these may be dropped from the request's flat guardrails list. A
pipeline manages just its own mode, so a guardrail that also runs in a
stage no pipeline covers must stay in the list or that stage's
independent activation is silently suppressed.
"""
managed_names: Final = frozenset(
step.guardrail for _policy_name, pipeline in pipelines for step in pipeline.steps
)
return frozenset(
name
for name in managed_names
if (gated_stages := _list_gated_stages_for_guardrail(name)) is not None
and gated_stages
<= frozenset(
pipeline.mode
for _policy_name, pipeline in pipelines
if any(step.guardrail == name for step in pipeline.steps)
)
)
@staticmethod
def get_all_resolved_policies(
policies: dict[str, Policy] | None = None,

View file

@ -4226,6 +4226,164 @@ async def test_add_guardrails_from_policy_engine():
attachment_registry._initialized = False
def _policy_engine_pipeline_registries(policies, monkeypatch, callbacks):
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
from litellm.types.proxy.policy_engine import PolicyAttachment
monkeypatch.setattr(litellm, "callbacks", callbacks)
policy_registry = get_policy_registry()
policy_registry._policies = policies
policy_registry._initialized = True
attachment_registry = get_attachment_registry()
attachment_registry._attachments = [PolicyAttachment(policy=name, scope="*") for name in policies]
attachment_registry._initialized = True
return policy_registry, attachment_registry
def _reset_policy_engine_registries(policy_registry, attachment_registry):
policy_registry._policies = {}
policy_registry._initialized = False
attachment_registry._attachments = []
attachment_registry._initialized = False
def _word_guard_pipeline_policy(mode: str):
from litellm.types.proxy.policy_engine import (
GuardrailPipeline,
PipelineStep,
Policy,
PolicyGuardrails,
)
return Policy(
guardrails=PolicyGuardrails(add=["word_guard"]),
pipeline=GuardrailPipeline(mode=mode, steps=[PipelineStep(guardrail="word_guard")]),
)
@pytest.mark.asyncio
async def test_pipeline_keeps_cross_mode_guardrail_in_flat_list(monkeypatch):
"""
Regression for a mode-blind strip: a guardrail supporting pre_call and
post_call, activated independently by one policy while another policy's
pre_call pipeline names it, must stay in the flat guardrails list so its
post_call stage keeps running (LIT-6536).
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.proxy.policy_engine import Policy, PolicyGuardrails
guardrail = CustomGuardrail(guardrail_name="word_guard", event_hook=["pre_call", "post_call"])
policy_registry, attachment_registry = _policy_engine_pipeline_registries(
{
"activate-guard": Policy(guardrails=PolicyGuardrails(add=["word_guard"])),
"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"]
assert data["metadata"]["_pipeline_managed_guardrails"] == {"word_guard"}
@pytest.mark.asyncio
async def test_pipelines_covering_every_stage_strip_guardrail_from_flat_list(monkeypatch):
"""
When pipelines cover every stage the guardrail supports, the name is
dropped from the flat list so nothing runs it independently.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
guardrail = CustomGuardrail(guardrail_name="word_guard", event_hook=["pre_call", "post_call"])
policy_registry, attachment_registry = _policy_engine_pipeline_registries(
{
"input-pipeline": _word_guard_pipeline_policy("pre_call"),
"output-pipeline": _word_guard_pipeline_policy("post_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"] == []
@pytest.mark.asyncio
async def test_pipeline_strips_guardrail_with_no_registered_callback(monkeypatch):
"""
A pipeline-managed name with no registered callback is dropped from the
flat list: the name is inert there, so the pre-LIT-6536 behavior stands.
"""
policy_registry, attachment_registry = _policy_engine_pipeline_registries(
{"input-pipeline": _word_guard_pipeline_policy("pre_call")},
monkeypatch,
callbacks=[],
)
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"] == []
@pytest.mark.asyncio
async def test_pipeline_keeps_guardrail_with_stage_no_pipeline_can_cover(monkeypatch):
"""
A guardrail with a during_call stage can never be fully covered, since
pipelines only have pre_call and post_call modes, so it stays listed.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
guardrail = CustomGuardrail(guardrail_name="word_guard", event_hook=["pre_call", "during_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():
"""