From 8d0b1881d7b028e1eb6b09c1bb76f1527923768f Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 19:59:36 -0500 Subject: [PATCH] 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. --- .../guardrail_hooks/llm_shield/llm_shield.py | 30 +++++++++++++- .../guardrail_hooks/test_llm_shield.py | 40 +++++++++++++++++++ 2 files changed, 68 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index 172de67020b..8989fabc321 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -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 -------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py index 45d45301858..ffd6ede28b6 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -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()