mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(policy_engine): union list-gated stages across all callbacks sharing a guardrail name
This commit is contained in:
parent
55e71a53d2
commit
808b5e90c5
3 changed files with 98 additions and 16 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue