From 28697377c3bc968e78875997f25e55dcfd007429 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Tue, 22 Sep 2026 22:02:49 -0500 Subject: [PATCH] fix(guardrails): import copy, keep the vault id off the provider, drop recursion Three defects Greptile and veria-ai found on the reopened PR, all real: - `copy.deepcopy` was called in `apply_guardrail` with no `import copy`, a guaranteed NameError on every response carrying tool calls. It landed on 2026-09-13, ten days after the review that rated this branch safe, and no test reached it: every tool-call test covered the request side. Adds the import and a regression test on the response side. - The vault session id was stored in `metadata`, which is forwarded to the provider on /v1/responses. A provider holding the placeholders and the session id can call the shield's rehydrate endpoint and read back the plaintext this guardrail exists to withhold. Moves it to `litellm_metadata`, which is not forwarded, and reads it back from there only. - `_collect_json_leaves` recursed over model-controlled JSON; the repo's recursive_detector gate rejects that. Rewritten with an explicit stack, same depth bound. 52 tests pass. ruff format, ruff-strict and check_type_discipline all clean, with LIT counts identical to the merge base. --- .../llm_shield_proxy/llm_shield_proxy.py | 53 +++++++++++-------- .../guardrail_hooks/test_llm_shield_proxy.py | 45 +++++++++++++++- 2 files changed, 75 insertions(+), 23 deletions(-) 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 e093199578a..9877f746730 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 @@ -5,6 +5,7 @@ # # +-------------------------------------------------------------+ +import copy import os import uuid from collections.abc import AsyncGenerator, Callable, Mapping, Sequence @@ -254,24 +255,28 @@ def _collect_json_leaves(node: object, slots: _SlotSink, depth: int = 0) -> None string, so a value worth restoring can sit at any depth. Bounded by `_MAX_CONTENT_DEPTH` for the same reason the request walk is: the shape is model controlled, and the bound is what stops a crafted one from becoming an unbounded - descent. + descent. Walked with an explicit stack rather than recursively, so a deeply nested + tool input cannot spend stack frames proportional to attacker-chosen depth. """ - if depth > _MAX_CONTENT_DEPTH: - return - if isinstance(node, dict): - for key in tuple(node): - value = node[key] - if isinstance(value, str) and value: - slots.append((value, lambda new, d=node, k=key: d.__setitem__(k, new))) - else: - _collect_json_leaves(value, slots, depth + 1) - return - if isinstance(node, list): - for index, value in enumerate(node): - if isinstance(value, str) and value: - slots.append((value, lambda new, entries=node, i=index: entries.__setitem__(i, new))) - else: - _collect_json_leaves(value, slots, depth + 1) + pending: Final[list] = [(node, depth)] # mutable-ok: local walk stack. + while pending: + current, current_depth = pending.pop() + if current_depth > _MAX_CONTENT_DEPTH: + continue + if isinstance(current, dict): + for key in tuple(current): + value = current[key] + if isinstance(value, str) and value: + slots.append((value, lambda new, d=current, k=key: d.__setitem__(k, new))) + else: + pending.append((value, current_depth + 1)) + continue + if isinstance(current, list): + for index, value in enumerate(current): + if isinstance(value, str) and value: + slots.append((value, lambda new, entries=current, i=index: entries.__setitem__(i, new))) + else: + pending.append((value, current_depth + 1)) def _carry_sort_key(key: tuple) -> tuple: @@ -391,7 +396,11 @@ class LLMShieldProxyGuardrail(CustomGuardrail): caller from reaching another caller's vault. """ session_id: Final = f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" - metadata: Final = data.setdefault("metadata", {}) # mutable-ok: per-request store. + # `litellm_metadata` is proxy-private; `metadata` is forwarded to the provider on + # /v1/responses. The session id is a capability against the vault's rehydrate + # endpoint, so handing it to the provider alongside the placeholders would let the + # provider read back exactly what this guardrail exists to withhold. + metadata: Final = data.setdefault("litellm_metadata", {}) # mutable-ok: per-request store. if isinstance(metadata, dict): metadata[_SESSION_METADATA_KEY] = session_id return session_id @@ -404,7 +413,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail): reply that cannot be restored is a visible placeholder, while trusting a caller-supplied id would hand them someone else's plaintext. """ - metadata: Final = data.get("metadata") + # Read only from `litellm_metadata`, the same proxy-private store `_mint_session_id` + # writes to. A caller can populate `metadata`; they cannot populate this. + metadata: Final = data.get("litellm_metadata") existing: Final = metadata.get(_SESSION_METADATA_KEY) if isinstance(metadata, dict) else None if isinstance(existing, str) and existing.startswith(_VAULT_PREFIX): return existing @@ -602,9 +613,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): slots.append((value, lambda new, i=item, f=field: _write_field(i, f, new))) return tuple(slots) - async def _restore_responses_api_response( - self, response: Any, slots: Sequence[_Slot], data: MutableRequest - ) -> Any: + async def _restore_responses_api_response(self, response: Any, slots: Sequence[_Slot], data: MutableRequest) -> Any: """Puts the original values back into a Responses API reply.""" restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data)) for (_, write), replacement in zip(slots, restored): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 700702f3bb2..bfa8c04c603 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -575,7 +575,25 @@ class TestVaultIsolation: used = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] assert used != "victim-session" - assert data["metadata"]["llm_shield_session_id"] == used + assert data["litellm_metadata"]["llm_shield_session_id"] == used + + @pytest.mark.asyncio + async def test_session_id_is_not_forwarded_to_the_provider(self): + """The vault id is a capability, so it must stay out of provider-visible metadata. + + `metadata` is forwarded upstream on /v1/responses; `litellm_metadata` is not. A + provider holding both the placeholders and the session id could call the shield's + rehydrate endpoint and read back exactly what this guardrail withholds. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"messages": [{"role": "user", "content": "a@b.com"}], "metadata": {}} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + used = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert "llm_shield_session_id" not in data["metadata"] + assert data["litellm_metadata"]["llm_shield_session_id"] == used @pytest.mark.asyncio async def test_restore_ignores_a_foreign_session_id(self): @@ -899,3 +917,28 @@ class TestStreamingRehydration: assert len(chunks) == 3 assert [c.choices[0].delta.content for c in chunks] == ["one ", "two ", "three"] + + +class TestApplyGuardrailToolCalls: + """The unified entry point the UI's Test button and the translation handlers use.""" + + @pytest.mark.asyncio + async def test_response_tool_call_arguments_are_rehydrated(self): + """Regression: this path deep-copied tool calls with `copy` never imported. + + 47 tests passed with a guaranteed NameError here, because every tool-call test + covered the request side and this is the only path that reaches the copy. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["hi", '{"email": "a@b.com"}']}) + + data = {"litellm_metadata": {"llm_shield_session_id": "shield-abc"}} + inputs = { + "texts": ["hi"], + "tool_calls": [{"function": {"name": "send", "arguments": '{"email": "[EMAIL_1]"}'}}], + } + + merged = await guardrail.apply_guardrail(inputs=inputs, request_data=data, input_type="response") + + assert merged["tool_calls"][0]["function"]["arguments"] == '{"email": "a@b.com"}' + assert inputs["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}'