diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index ae55b7ab906..af784af11df 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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)) diff --git a/litellm/proxy/policy_engine/policy_resolver.py b/litellm/proxy/policy_engine/policy_resolver.py index 70503f85b03..9df98be1d00 100644 --- a/litellm/proxy/policy_engine/policy_resolver.py +++ b/litellm/proxy/policy_engine/policy_resolver.py @@ -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, diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index ee0e2014951..2338de6d59c 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -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(): """