fix(guardrails): redact completion prompts and responses tool items

Two more provider-bound request shapes were reaching the model intact while
the guardrail reported as enabled.

/v1/completions carries its text in a top-level `prompt`, which the
traversal never looked at. It is handled as a string and as the array form,
where each entry is rewritten in place.

Responses input items hold tool data outside `content`: a function_call item
in `arguments`, a function_call_output item in `output`. Both are now
collected alongside the item's content.

Adds a test per shape.
This commit is contained in:
Ninad Phalak 2026-09-03 19:59:36 -05:00
parent f3eb108f86
commit 8d0b1881d7
No known key found for this signature in database
GPG key ID: 59119ED515433744
2 changed files with 68 additions and 2 deletions

View file

@ -77,6 +77,10 @@ _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's p
# _locate_request_texts, which freezes it into a tuple before returning.
_SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors.
# A caller-owned list whose entries are rewritten in place, such as a Completions
# `prompt` sent as an array of strings.
MutableSeq: TypeAlias = list # mutable-ok: the request payload's own list.
def _collect(container: MutableRequest, key: str, slots: _SlotSink) -> None:
"""Records the string at `key`, along with the write that replaces it."""
@ -85,6 +89,23 @@ def _collect(container: MutableRequest, key: str, slots: _SlotSink) -> None:
slots.append((value, lambda new, c=container, k=key: c.__setitem__(k, new)))
def _collect_entry(entries: MutableSeq, index: int, slots: _SlotSink) -> None:
"""Records a string held directly in a list, rather than under a key."""
value: Final = entries[index]
if isinstance(value, str) and value:
slots.append((value, lambda new, e=entries, i=index: e.__setitem__(i, new)))
def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None:
"""The Completions API sends its text in a top-level `prompt`."""
prompt: Final = data.get("prompt")
if isinstance(prompt, str):
_collect(data, "prompt", slots)
return
for index in range(len(prompt)) if isinstance(prompt, list) else ():
_collect_entry(prompt, index, slots)
def _collect_content(container: MutableRequest, slots: _SlotSink) -> None:
"""`content` is either a string or the multimodal list of typed parts."""
content: Final = container.get("content")
@ -115,8 +136,12 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None:
_collect(data, "input", slots)
return
for item in request_input if isinstance(request_input, list) else ():
if isinstance(item, dict):
_collect_content(item, slots)
if not isinstance(item, dict):
continue
_collect_content(item, slots)
# A function_call item holds `arguments`; a function_call_output holds `output`.
_collect(item, "arguments", slots)
_collect(item, "output", slots)
class LLMShieldGuardrail(CustomGuardrail):
@ -259,6 +284,7 @@ class LLMShieldGuardrail(CustomGuardrail):
_collect_content(message, slots)
_collect_tool_arguments(message, slots)
_collect_responses_fields(data, slots)
_collect_prompt(data, slots)
return tuple(slots)
# --- hooks --------------------------------------------------------------------

View file

@ -252,6 +252,46 @@ class TestRequestCoverage:
assert data["messages"][0]["function_call"]["arguments"] == '{"email": "[EMAIL_1]"}'
@pytest.mark.asyncio
async def test_completions_prompt_is_redacted(self):
"""/v1/completions puts its text in a top-level `prompt`, not in messages."""
guardrail = _guardrail()
mock = _mock_post(guardrail, {"texts": ["Email [EMAIL_1]"]})
data = {"prompt": "Email jane.doe@example.com"}
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion")
assert mock.call_args_list[0].kwargs["json"]["texts"] == ["Email jane.doe@example.com"]
assert data["prompt"] == "Email [EMAIL_1]"
@pytest.mark.asyncio
async def test_completions_prompt_array_is_redacted(self):
"""`prompt` also accepts an array, and each entry is provider-bound."""
guardrail = _guardrail()
_mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]})
data = {"prompt": ["jane.doe@example.com", "555-0100"]}
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion")
assert data["prompt"] == ["[EMAIL_1]", "[PHONE_1]"]
@pytest.mark.asyncio
async def test_responses_function_call_items_are_redacted(self):
"""Responses input items hold tool data in `arguments` and `output`."""
guardrail = _guardrail()
_mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}', "sent to [EMAIL_1]"]})
data = {
"input": [
{"type": "function_call", "name": "send", "arguments": '{"email": "jane.doe@example.com"}'},
{"type": "function_call_output", "call_id": "c1", "output": "sent to jane.doe@example.com"},
]
}
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses")
assert data["input"][0]["arguments"] == '{"email": "[EMAIL_1]"}'
assert data["input"][1]["output"] == "sent to [EMAIL_1]"
@pytest.mark.asyncio
async def test_every_shape_in_one_request_is_redacted(self):
guardrail = _guardrail()