mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(guardrails): gather tool-call PII checks + recurse nested Anthropic tool inputs
Address two review findings:
1. Pre-call hook: collect tool-call check_pii coroutines and dispatch
via asyncio.gather instead of sequential awaits, matching the
parallel pattern used by the content-masking path above.
2. Anthropic response handler: replace flat top-level-only string loop
with recursive traversal so PII inside nested dicts/lists within
tool_use input (e.g. {"patient": {"ssn": "..."}}) is also masked.
Adds a _mask_nested_strings helper and a new test for nested inputs.
This commit is contained in:
parent
c208335a3d
commit
a23b13ae87
2 changed files with 109 additions and 30 deletions
|
|
@ -802,10 +802,10 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
elif isinstance(content, list) and content_idx_optional is not None:
|
||||
messages[msg_idx]["content"][content_idx_optional]["text"] = r
|
||||
|
||||
# Also mask PII in tool_call / function_call arguments. The content loop
|
||||
# above only inspects message content, so PII embedded in tool-call
|
||||
# arguments would otherwise reach the provider unmasked -- this mirrors the
|
||||
# response-side handling in _process_response_for_pii.
|
||||
# Also mask PII in tool_call / function_call arguments. Gather
|
||||
# coroutines to match the parallel dispatch used for content above.
|
||||
tool_tasks: List[Any] = []
|
||||
tool_targets: List[Any] = []
|
||||
for m in messages:
|
||||
if not isinstance(m, dict):
|
||||
continue
|
||||
|
|
@ -818,22 +818,32 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
if isinstance(function, dict) and isinstance(
|
||||
function.get("arguments"), str
|
||||
):
|
||||
function["arguments"] = await self.check_pii(
|
||||
text=function["arguments"],
|
||||
output_parse_pii=self.output_parse_pii,
|
||||
presidio_config=presidio_config,
|
||||
request_data=data,
|
||||
tool_tasks.append(
|
||||
self.check_pii(
|
||||
text=function["arguments"],
|
||||
output_parse_pii=self.output_parse_pii,
|
||||
presidio_config=presidio_config,
|
||||
request_data=data,
|
||||
)
|
||||
)
|
||||
tool_targets.append((function, "arguments"))
|
||||
function_call = m.get("function_call")
|
||||
if isinstance(function_call, dict) and isinstance(
|
||||
function_call.get("arguments"), str
|
||||
):
|
||||
function_call["arguments"] = await self.check_pii(
|
||||
text=function_call["arguments"],
|
||||
output_parse_pii=self.output_parse_pii,
|
||||
presidio_config=presidio_config,
|
||||
request_data=data,
|
||||
tool_tasks.append(
|
||||
self.check_pii(
|
||||
text=function_call["arguments"],
|
||||
output_parse_pii=self.output_parse_pii,
|
||||
presidio_config=presidio_config,
|
||||
request_data=data,
|
||||
)
|
||||
)
|
||||
tool_targets.append((function_call, "arguments"))
|
||||
if tool_tasks:
|
||||
tool_results = await asyncio.gather(*tool_tasks)
|
||||
for (container, key), result in zip(tool_targets, tool_results):
|
||||
container[key] = result
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Presidio PII Masking: Redacted pii message: {data['messages']}"
|
||||
|
|
@ -1015,6 +1025,38 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
break
|
||||
return text
|
||||
|
||||
async def _mask_nested_strings(
|
||||
self,
|
||||
obj: Any,
|
||||
mode: Literal["mask", "unmask"],
|
||||
pii_tokens: Dict[str, str],
|
||||
presidio_config: Optional[dict],
|
||||
request_data: dict,
|
||||
) -> Any:
|
||||
"""Recursively mask/unmask every string leaf in a JSON-like structure."""
|
||||
if isinstance(obj, str):
|
||||
if mode == "unmask":
|
||||
return self._unmask_pii_text(obj, pii_tokens)
|
||||
return await self.check_pii(
|
||||
text=obj,
|
||||
output_parse_pii=False,
|
||||
presidio_config=presidio_config,
|
||||
request_data=request_data,
|
||||
)
|
||||
if isinstance(obj, dict):
|
||||
for key in list(obj.keys()):
|
||||
obj[key] = await self._mask_nested_strings(
|
||||
obj[key], mode, pii_tokens, presidio_config, request_data
|
||||
)
|
||||
return obj
|
||||
if isinstance(obj, list):
|
||||
for i, item in enumerate(obj):
|
||||
obj[i] = await self._mask_nested_strings(
|
||||
item, mode, pii_tokens, presidio_config, request_data
|
||||
)
|
||||
return obj
|
||||
return obj
|
||||
|
||||
@staticmethod
|
||||
def _is_anthropic_message_response(response: Any) -> bool:
|
||||
"""Check if the response is an Anthropic native message dict."""
|
||||
|
|
@ -1066,23 +1108,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
request_data=request_data,
|
||||
)
|
||||
elif block_type == "tool_use":
|
||||
# Mirror the OpenAI-format path (_process_response_for_pii), which masks
|
||||
# tool_call arguments: a tool_use block's `input` dict can carry PII the
|
||||
# model produced, and must be masked too -- not just sibling text blocks.
|
||||
tool_input = block.get("input")
|
||||
if isinstance(tool_input, dict):
|
||||
for key, value in list(tool_input.items()):
|
||||
if not isinstance(value, str):
|
||||
continue
|
||||
if mode == "unmask":
|
||||
tool_input[key] = self._unmask_pii_text(value, pii_tokens)
|
||||
elif mode == "mask":
|
||||
tool_input[key] = await self.check_pii(
|
||||
text=value,
|
||||
output_parse_pii=False,
|
||||
presidio_config=presidio_config,
|
||||
request_data=request_data,
|
||||
)
|
||||
if isinstance(tool_input, (dict, list)):
|
||||
await self._mask_nested_strings(
|
||||
tool_input, mode, pii_tokens, presidio_config, request_data
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
|
|
|||
|
|
@ -147,3 +147,52 @@ async def test_anthropic_response_masks_pii_in_tool_use_input(presidio_guardrail
|
|||
), "tool_use input should be masked"
|
||||
assert tool_block["input"]["ssn"] == "<US_SSN>"
|
||||
assert tool_block["input"]["note"] == "customer", "non-PII values unchanged"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_response_masks_nested_tool_use_input(presidio_guardrail):
|
||||
"""
|
||||
PII inside nested dicts/lists within tool_use input must also be masked.
|
||||
"""
|
||||
response = {
|
||||
"id": "msg_02",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "tu_2",
|
||||
"name": "save_patient",
|
||||
"input": {
|
||||
"patient": {
|
||||
"ssn": "078-05-1120",
|
||||
"name": "Alice",
|
||||
},
|
||||
"addresses": [
|
||||
{"street": "123 Main", "phone": "078-05-1120"}
|
||||
],
|
||||
"code": 42,
|
||||
},
|
||||
},
|
||||
],
|
||||
"model": "claude-sonnet-4-20250514",
|
||||
"stop_reason": "end_turn",
|
||||
}
|
||||
|
||||
async def mock_check_pii(text, output_parse_pii, presidio_config, request_data):
|
||||
return text.replace("078-05-1120", "<US_SSN>")
|
||||
|
||||
presidio_guardrail.check_pii = mock_check_pii
|
||||
|
||||
result = await presidio_guardrail._process_anthropic_response_for_pii(
|
||||
response=response,
|
||||
request_data={},
|
||||
mode="mask",
|
||||
)
|
||||
|
||||
tool_input = result["content"][0]["input"]
|
||||
assert tool_input["patient"]["ssn"] == "<US_SSN>", "nested dict value masked"
|
||||
assert tool_input["patient"]["name"] == "Alice", "non-PII nested value unchanged"
|
||||
assert tool_input["addresses"][0]["phone"] == "<US_SSN>", "list-nested value masked"
|
||||
assert tool_input["addresses"][0]["street"] == "123 Main", "non-PII list value unchanged"
|
||||
assert tool_input["code"] == 42, "non-string values unchanged"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue