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>
This commit is contained in:
joshua-berri 2026-09-28 17:34:38 -07:00 • committed by GitHub
parent 3572d359a1
commit 98c710c411
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 222 additions and 21 deletions

View file

@ -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:

View file

@ -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,
)

View file

@ -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 "<PERSON>" 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 <PERSON> and my colleague is <PERSON>."
assert later[3]["content"] == "Now compare against <PERSON> 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": {"<PERSON_1>": "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 "<CREDIT_CARD" in masked
assert len(anonymized) == 1
assert len(analyzed) == 1
assert analyzed[0]["text"] == CHUNK_MARKER_ONE
@pytest.mark.asyncio
@pytest.mark.parametrize("has_tokens", [False, True])
async def test_restoration_never_contacts_presidio(has_tokens: bool) -> 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": {"<CREDIT_CARD_1>": 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": ["<CREDIT_CARD_1>", ""]}, request_data=data, input_type="response"
)
assert result["texts"] == [CHUNK_MARKER_ONE if has_tokens else "<CREDIT_CARD_1>", ""]
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="<PERSON_1>"))])
data: Final = {"metadata": {"pii_tokens": {"<PERSON_1>": "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"

View file

@ -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,

View file

@ -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