diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 49bcb768600..62ab2d5ae2b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -39,11 +39,15 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.proxy._types import UserAPIKeyAuth from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.utils import GenericGuardrailAPIInputs +from litellm.types.utils import GenericGuardrailAPIInputs, TextChoices if TYPE_CHECKING: from litellm.caching.caching import DualCache from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.types.proxy.guardrails.guardrail_hooks.llm_shield_proxy import ( + LLMShieldProxyGuardrailConfigModel, + ) + from litellm.types.utils import CallTypes, LLMResponseTypes GUARDRAIL_NAME: Final = "llm_shield_proxy" @@ -202,6 +206,15 @@ def _as_array(value: object) -> MutableSeq | None: return value if isinstance(value, list) else None +def _detached(value: object) -> object: + """A deep copy of a reply or chunk, for restoring without touching LiteLLM's own object. + + LiteLLM keeps the object it handed the hooks to fill its response cache and its + logs, so writing restored plaintext into that object would put it there too. + """ + return copy.deepcopy(value) + + def _is_container(value: object) -> bool: """Whether `value` is a JSON object or array, without narrowing it to unknown types.""" return isinstance(value, (dict, list)) @@ -232,11 +245,15 @@ def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None: if prompt_object is not None: # A Responses API PromptObject. `variables` are substituted into the stored # prompt on the provider side, so they are caller text. `id` and `version` - # identify which prompt to use and must arrive unchanged. + # identify which prompt to use and must arrive unchanged. A variable is a string + # or a typed input such as `{"type": "input_text", "text": ...}`. variables: Final = _as_object(prompt_object.get("variables")) if variables is not None: for name in tuple(variables): _collect(variables, name, slots) + typed = _as_object(variables[name]) + if typed is not None: + _collect(typed, "text", slots) return entries: Final = _as_array(prompt) if entries is None: @@ -327,7 +344,9 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged """The Responses API sends text outside `messages`, in `instructions` and `input`. `instructions` is written by the application, not by the caller, so it is - collected into the privileged sink; `input` is the caller's own text. + collected into the privileged sink; `input` is the caller's own text, except for + system and developer items in it, which go to the privileged sink like their Chat + counterparts. """ _collect(data, "instructions", privileged) request_input: Final = data.get("input") @@ -345,10 +364,15 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink, privileged item = _as_object(entry) if item is None: continue - _collect_content(item, slots) - # A function_call item holds `arguments`; a function_call_output holds `output`. + _collect_content(item, privileged if item.get("role") in _PRIVILEGED_ROLES else slots) + # A function_call item holds `arguments`; a function_call_output holds `output`, + # as a string or as a list of input_text parts. A custom_tool_call holds `input` + # and a code_interpreter_call `code` -- the fields the reply side restores. _collect(item, "arguments", slots) _collect(item, "output", slots) + _collect_text_parts(item, "output", slots) + _collect(item, "input", slots) + _collect(item, "code", slots) # A replayed reasoning item carries the model's summary of its own reasoning, # which quotes whatever the conversation contained. _collect_text_parts(item, "summary", slots) @@ -946,6 +970,30 @@ class LLMShieldProxyGuardrail(CustomGuardrail): def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: parent's signature. return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] + @staticmethod + def get_config_model() -> type["LLMShieldProxyGuardrailConfigModel"]: + from litellm.types.proxy.guardrails.guardrail_hooks.llm_shield_proxy import ( + LLMShieldProxyGuardrailConfigModel, + ) + + return LLMShieldProxyGuardrailConfigModel + + async def async_post_call_success_deployment_hook( + self, + request_data: MutableRequest, + response: "LLMResponseTypes", + call_type: "CallTypes | None", + ) -> "LLMResponseTypes | None": + """Leaves the reply alone at the deployment, where LiteLLM caches what this returns. + + The inherited hook restores a model-level guardrail's reply here, before + `litellm/utils.py` writes it to the response cache, so the cache would hold this + caller's plaintext under a key built from the redacted request. The proxy's own + post-call hook runs model-level guardrails too, after the cache write, and + restores the reply there. + """ + return None + # --- transport --------------------------------------------------------------- def _headers(self, session_id: str) -> JsonBody: @@ -1143,9 +1191,19 @@ class LLMShieldProxyGuardrail(CustomGuardrail): user_api_key_dict: UserAPIKeyAuth, response: Any, ) -> Any: - """Restores the original values in a non-streaming response.""" + """Restores the original values in a copy of a non-streaming response. + + The copy is what keeps plaintext out of the response cache. LiteLLM caches the + reply it received from the provider, and on some paths -- a native Anthropic dict, + an in-memory cache -- it stores the object itself rather than a serialised + snapshot. Restoring that object in place would cache this caller's values under a + key built from the redacted request, which another caller's identical-looking + request then hits. Left untouched, the cached reply holds placeholders, and a hit + is restored against the new caller's own vault. + """ if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: return response + response = _detached(response) # rebind-ok: everything below restores the copy. if self._is_anthropic_message_response(response): return await self._restore_anthropic_response(response, data) @@ -1168,6 +1226,10 @@ class LLMShieldProxyGuardrail(CustomGuardrail): for choice in choices: message = getattr(choice, "message", None) if message is None: + # A Completions reply carries its text on the choice itself. + text = _read_field(choice, "text") + if isinstance(text, str) and text: + pending.append((text, functools.partial(_write_field, choice, "text"))) continue content = getattr(message, "content", None) if isinstance(content, str) and content: @@ -1290,14 +1352,18 @@ class LLMShieldProxyGuardrail(CustomGuardrail): for frames in await sse.feed(chunk): yield frames continue + # Chunks are restored as copies, for the reason the non-streaming hook copies: + # LiteLLM keeps the chunks it yielded to assemble the reply it caches and logs, + # so restoring them in place would cache this caller's plaintext. if _responses_event_type(chunk) is not None: - for event in await events.restore(chunk): + for event in await events.restore(_detached(chunk)): yield event continue - last_chunk = chunk - for choice in getattr(chunk, "choices", None) or (): + restored_chunk = _detached(chunk) + last_chunk = restored_chunk + for choice in getattr(restored_chunk, "choices", None) or (): await self._restore_choice(choice, carries, session_id) - yield chunk + yield restored_chunk # A stream that ended early can still leave text held back, in any shape. for frames in await sse.finish(): @@ -1308,7 +1374,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): async for trailing in self._flush_trailing(last_chunk, carries, session_id): yield trailing - async def _restore_choice(self, choice: Any, carries: _CarryWindows, session_id: str) -> None: + async def _restore_choice(self, choice: object, carries: _CarryWindows, session_id: str) -> None: """Restores one choice's delta, advancing that choice's own windows. Content and each tool call are separate token streams, so each gets its own @@ -1317,10 +1383,15 @@ class LLMShieldProxyGuardrail(CustomGuardrail): held back for one stream onto another. """ delta: Final = getattr(choice, "delta", None) - if delta is None: - return index: Final = _choice_index(choice) is_final: Final = bool(getattr(choice, "finish_reason", None)) + if isinstance(choice, TextChoices): + # A Completions stream carries its text on the choice itself, with no delta + # and no tool calls: one window, the content one. + await self._restore_text_window(choice, (index, None), carries, session_id, is_final) + return + if delta is None: + return await self._restore_content_window(delta, (index, None), carries, session_id, is_final) @@ -1333,6 +1404,25 @@ class LLMShieldProxyGuardrail(CustomGuardrail): # after it produces argument JSON the client has already stopped waiting for. await self._flush_finished_choice(delta, index, carries, session_id) + async def _restore_text_window( + self, + choice: object, + key: _CarryKey, + carries: _CarryWindows, + session_id: str, + is_final: bool, + ) -> None: + """Restores a Completions stream choice's `text` through its window.""" + carry: Final = carries.get(key, "") + text: Final = _read_field(choice, "text") + if not isinstance(text, str) or not text: + if not (is_final and carry): + return + emitted, remaining = await self._stream_step(text if isinstance(text, str) else "", carry, is_final, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if emitted or text: + _write_field(choice, "text", emitted) + async def _restore_content_window( self, delta: Any, @@ -1446,7 +1536,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail): chunk = self._chunk_for_choice(last_chunk, choice_index) if chunk is None: continue - if tool_index is None: + if isinstance(chunk.choices[0], TextChoices): + chunk.choices[0].text = text + elif tool_index is None: chunk.choices[0].delta.content = text else: # The copy carried this chunk's own content and tool calls, both already @@ -1470,7 +1562,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): choices: Final[tuple[object, ...]] = tuple(raw_choices) position: Final = next((at for at, choice in enumerate(choices) if _choice_index(choice) == index), 0) kept: Final = raw_choices[position] - if getattr(kept, "delta", None) is None: + if getattr(kept, "delta", None) is None and not isinstance(kept, TextChoices): return None kept.index = index # The terminal signal, if there was one, already went out with the real chunk. diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index b9314590f38..228ac50de54 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -19,7 +19,16 @@ from litellm.types.llms.openai import ( OutputTextDoneEvent, ResponsesAPIStreamEvents, ) -from litellm.types.utils import Choices, Delta, Message, ModelResponse, ModelResponseStream, StreamingChoices +from litellm.types.utils import ( + Choices, + Delta, + Message, + ModelResponse, + ModelResponseStream, + StreamingChoices, + TextChoices, + TextCompletionResponse, +) def _guardrail(**overrides: object) -> LLMShieldProxyGuardrail: @@ -387,6 +396,53 @@ class TestRequestCoverage: assert data["input"][0]["arguments"] == '{"email": "[EMAIL_1]"}' assert data["input"][1]["output"] == "sent to [EMAIL_1]" + @pytest.mark.asyncio + async def test_responses_tool_output_parts_are_redacted(self): + """A function_call_output can carry its result as a list of input_text parts.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["sent to [EMAIL_1]"]}) + + data = { + "input": [ + { + "type": "function_call_output", + "call_id": "c1", + "output": [{"type": "input_text", "text": "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 mock.call_args_list[0].kwargs["json"]["texts"] == ["sent to jane.doe@example.com"] + assert data["input"][0]["output"][0]["text"] == "sent to [EMAIL_1]" + + @pytest.mark.asyncio + async def test_responses_custom_tool_call_input_is_redacted(self): + """A replayed custom_tool_call carries its payload in `input`, not `arguments`.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["email [EMAIL_1]"]}) + + data = { + "input": [ + {"type": "custom_tool_call", "call_id": "c1", "name": "mail", "input": "email 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]["input"] == "email [EMAIL_1]" + assert data["input"][0]["name"] == "mail" + + @pytest.mark.asyncio + async def test_responses_code_interpreter_code_is_redacted(self): + """A replayed code_interpreter_call carries the code the model wrote, which the reply side restores.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["send('[EMAIL_1]')"]}) + + data = {"input": [{"type": "code_interpreter_call", "id": "ci_1", "code": "send('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]["code"] == "send('[EMAIL_1]')" + @pytest.mark.asyncio async def test_anthropic_system_prompt_is_redacted(self): """/v1/messages carries its system prompt at the top level, not in messages.""" @@ -551,6 +607,26 @@ class TestRequestCoverage: assert data["prompt"]["id"] == "pmpt_123" assert data["prompt"]["version"] == "2" + @pytest.mark.asyncio + async def test_responses_typed_prompt_variables_are_redacted(self): + """A variable can be a typed input rather than a string; its `text` is caller text.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = { + "prompt": { + "id": "pmpt_123", + "variables": { + "customer": {"type": "input_text", "text": "jane.doe@example.com"}, + "logo": {"type": "input_image", "image_url": "https://example.com/logo.png"}, + }, + } + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["prompt"]["variables"]["customer"] == {"type": "input_text", "text": "[EMAIL_1]"} + assert data["prompt"]["variables"]["logo"]["image_url"] == "https://example.com/logo.png" + @pytest.mark.asyncio async def test_completions_suffix_is_redacted(self): """LiteLLM forwards the legacy `suffix` to providers that support it.""" @@ -720,9 +796,9 @@ class TestRequestCoverage: call = SimpleNamespace(function=SimpleNamespace(name="notify", arguments='{"to": "[EMAIL_1]"}')) reply = ModelResponse(choices=[Choices(message=Message(content=None, tool_calls=None))]) reply.choices[0].message.tool_calls = [call] - await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply) + restored = await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply) - assert json.loads(call.function.arguments) == {"to": "ops@example.com"} + assert json.loads(restored.choices[0].message.tool_calls[0].function.arguments) == {"to": "ops@example.com"} def test_schema_nesting_past_the_bound_is_refused(self): schema: dict = {"type": "object", "description": "past-the-bound@example.com"} @@ -773,9 +849,11 @@ class TestRestoration: block = SimpleNamespace(text="[EMAIL_1]") response = SimpleNamespace(output=[SimpleNamespace(content=[block])]) - await guardrail.async_post_call_success_hook(data={"messages": []}, user_api_key_dict=None, response=response) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) - assert block.text == "a@b.com" + assert result.output[0].content[0].text == "a@b.com" @pytest.mark.asyncio async def test_anthropic_message_shape_is_restored(self): @@ -1027,6 +1105,24 @@ class TestVaultIsolation: privileged_id = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] assert privileged_id not in json.dumps(data, default=str) + def test_responses_system_and_developer_items_are_privileged(self) -> None: + """Responses `input` carries system and developer turns as items, like Chat messages. + + In the caller's vault, a caller could have the model echo a placeholder out of a + system message they cannot see and receive the plaintext behind it. + """ + data = { + "input": [ + {"role": "system", "content": "S"}, + {"type": "message", "role": "developer", "content": [{"type": "input_text", "text": "D"}]}, + {"role": "user", "content": "U"}, + ] + } + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert [text for text, _ in caller] == ["U"] + assert [text for text, _ in privileged] == ["S", "D"] + class TestFailClosed: @pytest.mark.asyncio @@ -1522,10 +1618,11 @@ class TestResponsesStreamRestoration: response=SimpleNamespace(output=[SimpleNamespace(content=[block]), call]), ) - await _restore_stream(guardrail, [completed]) + (out,) = await _restore_stream(guardrail, [completed]) - assert block["text"] == "Mail a@example.com" - assert call.arguments == '{"to": "a@example.com"}' + restored_block, restored_call = out.response.output[0].content[0], out.response.output[1] + assert restored_block["text"] == "Mail a@example.com" + assert restored_call.arguments == '{"to": "a@example.com"}' @pytest.mark.asyncio async def test_streams_on_different_parts_do_not_share_a_window(self): @@ -1552,9 +1649,9 @@ class TestResponsesStreamRestoration: type="response.reasoning_summary_part.done", item_id="rs_1", output_index=0, summary_index=0, part=part ) - await _restore_stream(guardrail, [event]) + (out,) = await _restore_stream(guardrail, [event]) - assert part.text == "asked about a@example.com" + assert out.part.text == "asked about a@example.com" @pytest.mark.asyncio async def test_mcp_call_arguments_are_restored(self): @@ -1587,3 +1684,144 @@ class TestResponsesStreamRestoration: assert out == [audio] assert shield.urls == [] + + +class TestResponseCacheIsolation: + """The reply LiteLLM caches must keep its placeholders. + + Placeholders are numbered per request, so two callers' redacted requests can be + identical and share a cache key. LiteLLM keeps the provider's reply object -- for a + native Anthropic dict or an in-memory cache, the object itself -- so restoring it in + place would hand one caller's values to the next caller who hits that key. + """ + + @pytest.mark.asyncio + async def test_a_cache_hit_is_restored_against_the_new_callers_vault(self): + cache = litellm.caching.caching.InMemoryCache() + provider_reply = {"type": "message", "content": [{"type": "text", "text": "Repeat [EMAIL_1]"}]} + cache.set_cache("redacted-request", provider_reply) + + alice, _ = _shielded({"[EMAIL_1]": "alice@example.com"}) + to_alice = await alice.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=cache.get_cache("redacted-request") + ) + bob, _ = _shielded({"[EMAIL_1]": "bob@example.com"}) + to_bob = await bob.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=cache.get_cache("redacted-request") + ) + + assert to_alice["content"][0]["text"] == "Repeat alice@example.com" + assert to_bob["content"][0]["text"] == "Repeat bob@example.com" + assert cache.get_cache("redacted-request")["content"][0]["text"] == "Repeat [EMAIL_1]" + + @pytest.mark.asyncio + async def test_a_model_response_is_restored_as_a_copy(self): + """A reply as `acompletion` returns it, hidden params and all.""" + reply = await litellm.acompletion( + model="gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], mock_response="Mail [EMAIL_1]" + ) + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + + restored = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=reply + ) + + assert restored.choices[0].message.content == "Mail a@example.com" + assert reply.choices[0].message.content == "Mail [EMAIL_1]" + + @pytest.mark.asyncio + async def test_responses_reply_is_restored_as_a_copy(self): + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + block = {"type": "output_text", "text": "Mail [EMAIL_1]"} + reply = SimpleNamespace(output=[SimpleNamespace(content=[block])]) + + restored = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=reply + ) + + assert restored.output[0].content[0]["text"] == "Mail a@example.com" + assert block["text"] == "Mail [EMAIL_1]" + + @pytest.mark.asyncio + async def test_stream_chunks_are_restored_as_copies(self): + """LiteLLM assembles the reply it caches from the chunks it yielded.""" + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + chunk = _chunk("Mail [EMAIL_1]", finish_reason="stop") + event = OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=0, + delta="Mail [EMAIL_1]", + ) + + out = await _restore_stream(guardrail, [chunk, event]) + + assert out[0].choices[0].delta.content == "Mail a@example.com" + assert chunk.choices[0].delta.content == "Mail [EMAIL_1]" + assert "Mail [EMAIL_1]" == event.delta + + +class TestCompletionsRestoration: + """`/v1/completions` replies carry their text on the choice, with no message or delta.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @pytest.mark.asyncio + async def test_completions_reply_is_restored(self): + guardrail, _ = _shielded(self.VAULT) + reply = TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAIL_1]", finish_reason="stop")]) + + restored = await guardrail.async_post_call_success_hook( + data={"prompt": "x"}, user_api_key_dict=None, response=reply + ) + + assert restored.choices[0].text == "Mail a@example.com" + + @pytest.mark.asyncio + async def test_completions_stream_is_restored_across_chunks(self): + guardrail, _ = _shielded(self.VAULT) + chunks = [ + TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAI")]), + TextCompletionResponse(choices=[TextChoices(index=0, text="L_1] now", finish_reason="stop")]), + ] + + out = await _restore_stream(guardrail, chunks) + + assert "".join(chunk.choices[0].text for chunk in out) == "Mail a@example.com now" + + @pytest.mark.asyncio + async def test_completions_stream_without_finish_reason_is_flushed(self): + guardrail, _ = _shielded(self.VAULT) + chunks = [TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAIL_1")])] + + out = await _restore_stream(guardrail, chunks) + + assert "".join(chunk.choices[0].text or "" for chunk in out) == "Mail [EMAIL_1" + + +class TestProxyWiring: + def test_dashboard_config_model_is_exposed(self): + """The guardrail garden reads the provider's fields from `get_config_model`.""" + model = LLMShieldProxyGuardrail.get_config_model() + + assert model is not None + assert {"api_key", "api_base"} <= set(model.model_fields) + + @pytest.mark.asyncio + async def test_the_deployment_hook_leaves_the_reply_for_the_cache_untouched(self): + """LiteLLM caches what the deployment hook returns, so restoring there caches plaintext. + + The proxy's post-call hook, which runs after the cache write, restores model-level + guardrails instead. + """ + guardrail, shield = _shielded({"[EMAIL_1]": "a@example.com"}) + reply = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + + result = await guardrail.async_post_call_success_deployment_hook( + request_data={"messages": [], "guardrails": [GUARDRAIL_NAME]}, response=reply, call_type=None + ) + + assert result is None + assert reply.choices[0].message.content == "[EMAIL_1]" + assert shield.urls == []