mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
3572d359a1
commit
98c710c411
5 changed files with 222 additions and 21 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue