mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
f3eb108f86
commit
8d0b1881d7
2 changed files with 68 additions and 2 deletions
|
|
@ -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 --------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue