From 98c710c41176b68326af2b37174ba40ebac0d8c4 Mon Sep 17 00:00:00 2001 From: joshua-berri Date: Mon, 28 Sep 2026 17:34:38 -0700 Subject: [PATCH] fix(guardrails): preserve Presidio output selection and restoration (#43401) * fix(guardrails): preserve Presidio output callback intent and tag selection * test(guardrails): verify Presidio callback stages after registry updates * fix(guardrails): preserve standalone Presidio token behavior --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> --- .../guardrails/guardrail_hooks/presidio.py | 11 +- .../guardrails/guardrail_initializers.py | 20 ++- .../guardrail_hooks/test_presidio.py | 129 +++++++++++++++++- .../guardrails/test_guardrail_registry.py | 14 +- .../proxy/guardrails/test_init_guardrails.py | 69 +++++++++- 5 files changed, 222 insertions(+), 21 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index f5e24c501f1..94750f08a9e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -200,6 +200,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): presidio_score_thresholds: dict[PiiEntityType | str, float] | None = None, presidio_entities_deny_list: list[PiiEntityType | str] | None = None, presidio_analyze_chunk_size_bytes: int | None = None, + _callback_role: Literal["scan", "restore"] | None = None, **kwargs, ): if logging_only is True: @@ -214,11 +215,12 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self.mock_redacted_text = mock_redacted_text self.output_parse_pii = output_parse_pii or False self.apply_to_output = apply_to_output + self._callback_role = _callback_role # When output_parse_pii or apply_to_output is enabled, the guardrail must # also run on post_call to unmask/mask the response. Expand the event_hook # so should_run_guardrail returns True for both pre_call and post_call. - if (self.output_parse_pii or self.apply_to_output) and not logging_only: + if _callback_role is None and (self.output_parse_pii or self.apply_to_output) and not logging_only: current_hook: Final = self.event_hook if isinstance(current_hook, str) and current_hook != "post_call": self.event_hook = cast(list[GuardrailEventHooks], [current_hook, "post_call"]) @@ -1710,13 +1712,14 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): """ texts: Final = inputs.get("texts", []) - # When input_type is "response" and pii_tokens are available, - # unmask the text instead of masking it. metadata: Final = (request_data.get("metadata") or {}) if request_data else {} pii_tokens: Final = metadata.get("pii_tokens", {}) new_texts: Final = [] - if input_type == "response" and pii_tokens: + if input_type == "response" and ( + self._callback_role == "restore" + or (self._callback_role is None and not self.apply_to_output and pii_tokens) + ): for text in texts: new_texts.append(self._unmask_pii_text(text, pii_tokens)) else: diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index c422902d30d..31688b2e903 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -115,6 +115,20 @@ def _is_mcp_only_mode(mode: str | list[str] | Mode) -> bool: return bool(hooks) and all(hook in _MCP_EVENT_HOOKS for hook in hooks) +def _presidio_output_mode(mode: str | list[str] | Mode, *, include_mcp: bool) -> str | list[str] | Mode: + def output_hooks(hooks: str | list[str]) -> list[str]: + if not hooks or (not include_mcp and _is_mcp_only_mode(hooks)): + return [] + return [GuardrailEventHooks.post_call.value] + + if isinstance(mode, Mode): + return Mode( + tags={tag: output_hooks(hooks) for tag, hooks in mode.tags.items()}, + default=output_hooks(mode.default) if mode.default is not None else None, + ) + return output_hooks(mode) + + def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> tuple[CustomGuardrail, ...]: from litellm.proxy.guardrails.guardrail_hooks.presidio import ( _OPTIONAL_PresidioPIIMasking, @@ -140,6 +154,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> presidio_language=litellm_params.presidio_language, presidio_entities_deny_list=litellm_params.presidio_entities_deny_list, apply_to_output=False, + _callback_role="scan", ) params.update(overrides) # Passed outside the heterogeneous params dict so the argument keeps @@ -155,7 +170,8 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> unmask_output_callback: Final = ( _make_presidio_callback( output_parse_pii=True, - event_hook=GuardrailEventHooks.post_call.value, + event_hook=_presidio_output_mode(litellm_params.mode, include_mcp=True), + _callback_role="restore", ) if run_input and litellm_params.output_parse_pii else None @@ -163,7 +179,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> mask_output_callback: Final = ( _make_presidio_callback( apply_to_output=True, - event_hook=GuardrailEventHooks.post_call.value, + event_hook=_presidio_output_mode(litellm_params.mode, include_mcp=explicit_filter_scope is not None), output_parse_pii=False, mask_response_content=True, ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index 0a4ffbaef26..a5625e45d75 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -8,7 +8,7 @@ import copy import json import re from contextlib import asynccontextmanager -from typing import Final +from typing import Final, Literal from unittest.mock import MagicMock, patch from aiohttp import web @@ -2275,15 +2275,17 @@ async def test_apply_to_output_streaming_mixed_chunks_flushes_and_warns(): @pytest.mark.asyncio -async def test_apply_guardrail_unmask_on_response(): +@pytest.mark.parametrize("output_parse_pii", [False, True]) +async def test_apply_guardrail_unmask_on_response(output_parse_pii: bool) -> None: """ When input_type is 'response' and pii_tokens exist, apply_guardrail should unmask text instead of masking it. """ guardrail = _OPTIONAL_PresidioPIIMasking( guardrail_name="test_presidio", - output_parse_pii=True, + output_parse_pii=output_parse_pii, mock_testing=True, + mock_redacted_text={"text": "unexpected scan", "items": []}, ) request_data = { @@ -2312,12 +2314,14 @@ async def test_apply_guardrail_unmask_on_response(): @pytest.mark.asyncio -async def test_apply_guardrail_masks_on_request(): +@pytest.mark.parametrize("input_type", ["request", "response"]) +async def test_standalone_scans_without_restoration_tokens(input_type: Literal["request", "response"]) -> None: """ - When input_type is 'request', apply_guardrail should mask as before. + Standalone callbacks retain scanning without tokens, including MCP results. """ guardrail = _OPTIONAL_PresidioPIIMasking( guardrail_name="test_presidio", + event_hook="post_mcp_call", output_parse_pii=True, mock_testing=True, ) @@ -2330,7 +2334,7 @@ async def test_apply_guardrail_masks_on_request(): result = await guardrail.apply_guardrail( inputs={"texts": ["Hello John Smith"]}, request_data={"model": "gpt-4o", "metadata": {}}, - input_type="request", + input_type=input_type, ) assert "" in result["texts"][0] @@ -4171,3 +4175,116 @@ async def test_pii_masking_replays_a_byte_identical_prefix_across_turns(mock_use assert json.dumps(later[: len(earlier)], sort_keys=True) == json.dumps(earlier, sort_keys=True) assert earlier[1]["content"] == "My name is and my colleague is ." assert later[3]["content"] == "Now compare against too." + + +@pytest.mark.asyncio +@pytest.mark.parametrize("surface", ["mcp_arguments", "mcp_result", "llm_output"]) +@pytest.mark.parametrize("action", [PiiAction.MASK, PiiAction.BLOCK]) +@pytest.mark.parametrize("has_tokens", [False, True]) +async def test_initialized_presidio_scans_selected_surface(surface: str, action: PiiAction, has_tokens: bool) -> None: + from mcp.types import CallToolResult, TextContent + + from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import MCPGuardrailTranslationHandler + from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails + from litellm.proxy.guardrails.guardrail_initializers import initialize_presidio + + params: Final = LitellmParams( + guardrail="presidio", + mode="post_mcp_call" if surface == "mcp_result" else "pre_mcp_call", + default_on=True, + output_parse_pii=True, + presidio_filter_scope="output" if surface == "llm_output" else "input", + presidio_analyzer_api_base="http://test-analyzer/", + presidio_anonymizer_api_base="http://test-anonymizer/", + pii_entities_config={"CREDIT_CARD": action}, + ) + callback: Final = initialize_presidio(params, {"guardrail_name": "selected_surface"})[0] + data: Final = { + "metadata": {"pii_tokens": {"": "Somebody"} if has_tokens else {}}, + "mcp_tool_name": "echo", + "mcp_arguments": {"text": CHUNK_MARKER_ONE}, + "guardrail_to_apply": callback, + } + result: Final = CallToolResult(content=[TextContent(type="text", text=CHUNK_MARKER_ONE)]) + answer: Final = ModelResponse(choices=[Choices(message=Message(role="assistant", content=CHUNK_MARKER_ONE))]) + analyzed: Final = [] + anonymized: Final = [] + + async def dispatch() -> None: + if surface == "mcp_arguments": + await MCPGuardrailTranslationHandler().process_input_messages(data, callback) + elif surface == "mcp_result": + await MCPGuardrailTranslationHandler().process_output_response(result, callback, request_data=data) + else: + await UnifiedLLMGuardrails().async_post_call_success_hook( + data, UserAPIKeyAuth(request_route="/v1/chat/completions"), answer + ) + + with patch.object( + callback, + "_get_session_iterator", + _make_marker_session_iterator(analyzed, recorded_anonymize_payloads=anonymized), + ): + if action == PiiAction.BLOCK: + with pytest.raises(BlockedPiiEntityError): + await dispatch() + assert anonymized == [] + assert data["mcp_arguments"]["text"] == CHUNK_MARKER_ONE + assert result.content[0].text == CHUNK_MARKER_ONE + assert answer.choices[0].message.content == CHUNK_MARKER_ONE + else: + await dispatch() + masked: Final = ( + data["mcp_arguments"]["text"] + if surface == "mcp_arguments" + else result.content[0].text + if surface == "mcp_result" + else answer.choices[0].message.content + ) + assert CHUNK_MARKER_ONE not in masked + assert " None: + from litellm.proxy.guardrails.guardrail_initializers import initialize_presidio + + params: Final = LitellmParams( + guardrail="presidio", + mode="pre_mcp_call", + output_parse_pii=True, + presidio_analyzer_api_base="http://test-analyzer/", + presidio_anonymizer_api_base="http://test-anonymizer/", + ) + callback: Final = initialize_presidio(params, {"guardrail_name": "restore_only"})[1] + analyzed: Final = [] + data: Final = {"metadata": {"pii_tokens": {"": CHUNK_MARKER_ONE} if has_tokens else {}}} + with patch.object(callback, "_get_session_iterator", _make_marker_session_iterator(analyzed)): + result: Final = await callback.apply_guardrail( + inputs={"texts": ["", ""]}, request_data=data, input_type="response" + ) + assert result["texts"] == [CHUNK_MARKER_ONE if has_tokens else "", ""] + assert analyzed == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("event_hook", ["pre_call", ["pre_call"], ["pre_call", "post_call"]]) +async def test_standalone_restoration_preserves_post_call_selection(event_hook: str | list[str]) -> None: + from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails + + callback: Final = _OPTIONAL_PresidioPIIMasking( + event_hook=event_hook, + default_on=True, + output_parse_pii=True, + mock_testing=True, + ) + response: Final = ModelResponse(choices=[Choices(message=Message(role="assistant", content=""))]) + data: Final = {"metadata": {"pii_tokens": {"": "Jane"}}, "guardrail_to_apply": callback} + await UnifiedLLMGuardrails().async_post_call_success_hook( + data, UserAPIKeyAuth(request_route="/v1/chat/completions"), response + ) + assert response.choices[0].message.content == "Jane" diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 836668de0c8..022fe85c779 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -615,7 +615,8 @@ def test_presidio_siblings_are_tracked_and_deleted_together(): siblings = handler.guardrail_id_to_sibling_callbacks[PRESIDIO_SIBLINGS_GID] assert primary is registered[0] assert siblings == tuple(registered[1:]) - assert [sibling.event_hook for sibling in siblings] == [GuardrailEventHooks.post_call] * 2 + assert not primary.should_run_guardrail({}, GuardrailEventHooks.post_call) + assert all(sibling.should_run_guardrail({}, GuardrailEventHooks.post_call) for sibling in siblings) for cb_list in lists[1:]: cb_list.extend(registered) @@ -643,11 +644,12 @@ def test_update_in_memory_guardrail_rebuilds_presidio_siblings_and_keeps_their_s roles_before = [ (callback.apply_to_output, callback.output_parse_pii, callback.event_hook) for callback in tracked ] - assert roles_before == [ - (False, True, [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call]), - (False, True, GuardrailEventHooks.post_call), - (True, False, GuardrailEventHooks.post_call), - ] + assert [ + callback for callback in tracked if callback.should_run_guardrail({}, GuardrailEventHooks.pre_call) + ] == tracked[:1] + assert [ + callback for callback in tracked if callback.should_run_guardrail({}, GuardrailEventHooks.post_call) + ] == tracked[1:] updated = Guardrail( guardrail_id=PRESIDIO_SIBLINGS_GID, diff --git a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py index 39f9f9458b7..79d91db902c 100644 --- a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py @@ -1,4 +1,5 @@ import json +from typing import Literal from unittest.mock import MagicMock, patch import pytest @@ -7,7 +8,7 @@ import pytest from litellm.proxy.guardrails.guardrail_hooks.custom_code.custom_code_guardrail import CustomCodeCompilationError from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 -from litellm.types.guardrails import SupportedGuardrailIntegrations +from litellm.types.guardrails import Mode, SupportedGuardrailIntegrations def test_initialize_presidio_guardrail(): @@ -211,13 +212,15 @@ def test_initialize_presidio_forwards_analyze_chunk_size_bytes(): (["pre_mcp_call", "post_mcp_call"], None, False), ({"tags": {"team:mcp": "pre_mcp_call"}, "default": ["pre_mcp_call", "post_mcp_call"]}, None, False), ({"tags": {"team:mcp": ["pre_mcp_call"]}, "default": "pre_call"}, None, True), - ({"tags": {}}, None, True), + ({"tags": {}}, None, False), ("pre_mcp_call", "both", True), ("pre_mcp_call", "output", True), ("pre_call", None, True), ], ) -async def test_initialize_presidio_mcp_only_mode_skips_post_call_output_scan(mode, filter_scope, expect_output_scanned): +async def test_initialize_presidio_mcp_only_mode_skips_post_call_output_scan( + mode, filter_scope, expect_output_scanned, monkeypatch +): """Regression: an MCP-only Presidio guardrail used to also scan the LLM response on post_call, so a blocked MCP tool call that the model repeated in its answer turned the whole request into an HTTP 400 instead of a 200.""" @@ -225,6 +228,7 @@ async def test_initialize_presidio_mcp_only_mode_skips_post_call_output_scan(mod from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import Choices, Message, ModelResponse + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) llm_answer = "Call me at 415-555-2671" litellm_params = { "guardrail": SupportedGuardrailIntegrations.PRESIDIO.value, @@ -431,3 +435,62 @@ def test_init_guardrails_v2_skips_guardrail_with_malformed_advisory_template(): } assert "broken_lakera_template" not in guardrail_names assert "healthy_presidio" in guardrail_names + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "mode,tags,restore,scope,tokens,expected,expected_calls", + [ + ({"tags": {"team:mcp": "pre_mcp_call"}, "default": "pre_call"}, ["team:mcp"], False, None, {}, "raw", 0), + ({"tags": {"team:mcp": "pre_mcp_call"}, "default": "pre_call"}, ["other"], False, None, {}, "masked", 1), + ({"tags": {"team:mcp": "pre_mcp_call"}, "default": "pre_call"}, [], False, None, {}, "masked", 1), + ({"tags": {"team:mcp": "pre_mcp_call"}}, [], False, None, {}, "raw", 0), + ("pre_mcp_call", [], True, None, {}, "raw", 1), + ("pre_mcp_call", [], True, None, {"restored": "twice", "raw": "restored"}, "restored", 1), + ("pre_mcp_call", [], False, "output", {"raw": "restored"}, "masked", 1), + ({"tags": {"team:mcp": "pre_mcp_call"}}, ["team:mcp"], False, "output", {}, "masked", 1), + ({"tags": {"team:mcp": "pre_mcp_call"}}, [], False, "output", {}, "raw", 0), + ], +) +async def test_presidio_initialized_output_dispatch( + mode: str | list[str] | Mode, + tags: list[str], + restore: bool, + scope: Literal["input", "output", "both"] | None, + tokens: dict[str, str], + expected: str, + expected_calls: int, + monkeypatch: pytest.MonkeyPatch, +) -> None: + from typing import Final + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails + from litellm.proxy.guardrails.guardrail_initializers import initialize_presidio + from litellm.types.guardrails import GuardrailEventHooks, LitellmParams + from litellm.types.utils import Choices, Message, ModelResponse + + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + params: Final = LitellmParams( + guardrail="presidio", + mode=mode, + default_on=True, + output_parse_pii=restore, + presidio_filter_scope=scope, + presidio_analyzer_api_base="https://example.invalid/analyze", + presidio_anonymizer_api_base="https://example.invalid/anonymize", + mock_redacted_text={"text": "masked", "items": []}, + ) + callbacks: Final = initialize_presidio(params, {"guardrail_name": "output_dispatch"}) + data: Final = {"metadata": {"tags": tags, "pii_tokens": tokens}} + response: Final = ModelResponse(choices=[Choices(message=Message(role="assistant", content="raw"), index=0)]) + selected: Final = tuple( + callback for callback in callbacks if callback.should_run_guardrail(data, GuardrailEventHooks.post_call) + ) + for callback in selected: + data["guardrail_to_apply"] = callback + await UnifiedLLMGuardrails().async_post_call_success_hook( + data, UserAPIKeyAuth(request_route="/v1/chat/completions"), response + ) + assert response.choices[0].message.content == expected + assert len(selected) == expected_calls