mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(policy_engine): keep cross-mode guardrail stages active when a pipeline manages one mode
This commit is contained in:
parent
5a0ed05765
commit
e18c19f7ab
3 changed files with 226 additions and 3 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue