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:
AUTHENSOR 2026-06-18 21:14:47 -05:00
parent c208335a3d
commit a23b13ae87
2 changed files with 109 additions and 30 deletions

View file

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

View file

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