mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(guardrails): mask Responses API input in Presidio pre-call hook
The Presidio PII pre-call hook read only data["messages"] and returned early when it was absent, so /v1/responses requests (which carry the prompt in data["input"]) were never masked. Route text extraction and in-place rewrite through the shared _content_utils helpers (iter_message_text / walk_user_text) so both messages and input are masked consistently. Refs #30728 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
ac7c2dc0d7
commit
8dcd6a45e7
2 changed files with 105 additions and 53 deletions
|
|
@ -43,6 +43,10 @@ from litellm.integrations.custom_guardrail import (
|
|||
log_guardrail_information,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails._content_utils import (
|
||||
iter_message_text,
|
||||
walk_user_text,
|
||||
)
|
||||
from litellm.types.guardrails import (
|
||||
GuardrailEventHooks,
|
||||
LitellmParams,
|
||||
|
|
@ -746,66 +750,35 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
content_safety = data.get("content_safety", None)
|
||||
verbose_proxy_logger.debug("content_safety: %s", content_safety)
|
||||
presidio_config = self.get_presidio_settings_from_request_data(data)
|
||||
messages = data.get("messages", None)
|
||||
if messages is None:
|
||||
|
||||
# Collect every text fragment from BOTH `messages` and the
|
||||
# Responses-API `input` field. A hook that only reads
|
||||
# `data["messages"]` silently skips `/v1/responses` input; the
|
||||
# shared `_content_utils` helpers normalise both request shapes.
|
||||
fragments = list(dict.fromkeys(iter_message_text(data)))
|
||||
if not fragments:
|
||||
return data
|
||||
tasks = []
|
||||
task_mappings: List[Tuple[int, Optional[int]]] = (
|
||||
[]
|
||||
) # Track (message_index, content_index) for each task
|
||||
|
||||
for msg_idx, m in enumerate(messages):
|
||||
content = m.get("content", None)
|
||||
if content is None:
|
||||
continue
|
||||
if isinstance(content, str):
|
||||
tasks.append(
|
||||
self.check_pii(
|
||||
text=content,
|
||||
output_parse_pii=self.output_parse_pii,
|
||||
presidio_config=presidio_config,
|
||||
request_data=data,
|
||||
)
|
||||
# Mask each unique fragment via the analyzer in parallel.
|
||||
masked = await asyncio.gather(
|
||||
*[
|
||||
self.check_pii(
|
||||
text=fragment,
|
||||
output_parse_pii=self.output_parse_pii,
|
||||
presidio_config=presidio_config,
|
||||
request_data=data,
|
||||
)
|
||||
task_mappings.append(
|
||||
(msg_idx, None)
|
||||
) # None indicates string content
|
||||
elif isinstance(content, list):
|
||||
for content_idx, c in enumerate(content):
|
||||
text_str = c.get("text", None)
|
||||
if text_str is None:
|
||||
continue
|
||||
tasks.append(
|
||||
self.check_pii(
|
||||
text=text_str,
|
||||
output_parse_pii=self.output_parse_pii,
|
||||
presidio_config=presidio_config,
|
||||
request_data=data,
|
||||
)
|
||||
)
|
||||
task_mappings.append((msg_idx, int(content_idx)))
|
||||
for fragment in fragments
|
||||
]
|
||||
)
|
||||
mask_map = dict(zip(fragments, masked))
|
||||
|
||||
responses = await asyncio.gather(*tasks)
|
||||
|
||||
# Map responses back to the correct message and content item
|
||||
for task_idx, r in enumerate(responses):
|
||||
mapping = task_mappings[task_idx]
|
||||
msg_idx = cast(int, mapping[0])
|
||||
content_idx_optional = cast(Optional[int], mapping[1])
|
||||
content = messages[msg_idx].get("content", None)
|
||||
if content is None:
|
||||
continue
|
||||
if isinstance(content, str) and content_idx_optional is None:
|
||||
messages[msg_idx][
|
||||
"content"
|
||||
] = r # replace content with redacted string
|
||||
elif isinstance(content, list) and content_idx_optional is not None:
|
||||
messages[msg_idx]["content"][content_idx_optional]["text"] = r
|
||||
# Rewrite the request body in place across `messages` and `input`.
|
||||
walk_user_text(data, lambda text: mask_map.get(text, text))
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Presidio PII Masking: Redacted pii message: {data['messages']}"
|
||||
"Presidio PII Masking: redacted request body (messages + input)"
|
||||
)
|
||||
data["messages"] = messages
|
||||
return data
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
|
|||
|
|
@ -476,6 +476,85 @@ async def test_no_messages_field(presidio_guardrail, mock_user_api_key, mock_cac
|
|||
print("✓ No messages field test passed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_input_string_is_masked(
|
||||
presidio_guardrail, mock_user_api_key, mock_cache
|
||||
):
|
||||
"""The Responses API carries the prompt in data['input'] (string), which
|
||||
must be masked like chat `messages` (issue #30728)."""
|
||||
test_data = {
|
||||
"input": "My email is test@example.com and card 4111-1111-1111-1111",
|
||||
"model": "gpt-4o",
|
||||
}
|
||||
|
||||
async def mock_check_pii(text, output_parse_pii, presidio_config, request_data):
|
||||
return text.replace("test@example.com", "[EMAIL]").replace(
|
||||
"4111-1111-1111-1111", "[CREDIT_CARD]"
|
||||
)
|
||||
|
||||
presidio_guardrail.check_pii = mock_check_pii
|
||||
result = await presidio_guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=mock_user_api_key,
|
||||
cache=mock_cache,
|
||||
data=test_data,
|
||||
call_type="aresponses",
|
||||
)
|
||||
assert result["input"] == "My email is [EMAIL] and card [CREDIT_CARD]"
|
||||
assert "test@example.com" not in result["input"]
|
||||
assert "4111-1111-1111-1111" not in result["input"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_input_role_messages_are_masked(
|
||||
presidio_guardrail, mock_user_api_key, mock_cache
|
||||
):
|
||||
"""Responses API `input` given as a list of role messages must be masked."""
|
||||
test_data = {
|
||||
"input": [{"role": "user", "content": "Contact me at test@example.com"}],
|
||||
"model": "gpt-4o",
|
||||
}
|
||||
|
||||
async def mock_check_pii(text, output_parse_pii, presidio_config, request_data):
|
||||
return text.replace("test@example.com", "[EMAIL]")
|
||||
|
||||
presidio_guardrail.check_pii = mock_check_pii
|
||||
result = await presidio_guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=mock_user_api_key,
|
||||
cache=mock_cache,
|
||||
data=test_data,
|
||||
call_type="aresponses",
|
||||
)
|
||||
assert result["input"][0]["content"] == "Contact me at [EMAIL]"
|
||||
assert "test@example.com" not in result["input"][0]["content"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_both_messages_and_input_are_masked(
|
||||
presidio_guardrail, mock_user_api_key, mock_cache
|
||||
):
|
||||
"""When both `messages` and `input` are present, both must be masked."""
|
||||
test_data = {
|
||||
"messages": [{"role": "user", "content": "msg card 4111-1111-1111-1111"}],
|
||||
"input": "input email test@example.com",
|
||||
"model": "gpt-4o",
|
||||
}
|
||||
|
||||
async def mock_check_pii(text, output_parse_pii, presidio_config, request_data):
|
||||
return text.replace("4111-1111-1111-1111", "[CREDIT_CARD]").replace(
|
||||
"test@example.com", "[EMAIL]"
|
||||
)
|
||||
|
||||
presidio_guardrail.check_pii = mock_check_pii
|
||||
result = await presidio_guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=mock_user_api_key,
|
||||
cache=mock_cache,
|
||||
data=test_data,
|
||||
call_type="aresponses",
|
||||
)
|
||||
assert result["messages"][0]["content"] == "msg card [CREDIT_CARD]"
|
||||
assert result["input"] == "input email [EMAIL]"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_logging_hook_multimodal_message_format(presidio_guardrail):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue