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.
This commit is contained in:
Ninad Phalak 2026-09-22 22:02:49 -05:00
parent 097cafb4f2
commit 28697377c3
No known key found for this signature in database
GPG key ID: 59119ED515433744
2 changed files with 75 additions and 23 deletions

View file

@ -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):

View file

@ -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]"}'