fix(guardrails/peyeeye): namespace cache key + redact tool_call args

Two issues raised by Veria AI on the open PR:

- caller-controlled rehydration cache key: `litellm_call_id` comes from
  the inbound `x-litellm-call-id` header, so two callers can collide on
  the same key and rehydrate each other's PII. Namespace the cache entry
  with the authenticated key (server-side) in addition to the call id.

- PII bypass through tool call arguments: `_iter_message_text` only
  walked `message.content`, so PII inside `tool_calls[].function.arguments`
  or legacy `function_call.arguments` flowed to the LLM unredacted. Walk
  those fields too, and mirror the coverage on the response side so
  placeholders the model echoes into tool args get rehydrated.
This commit is contained in:
Tim 2026-05-10 14:07:35 -05:00
parent c35b0e3714
commit 5860e66730
2 changed files with 230 additions and 17 deletions

View file

@ -157,7 +157,7 @@ class PEyeEyeGuardrail(CustomGuardrail):
_set_message_text(messages[msg_idx], part_path, redacted)
if session_id:
cache_key = self._cache_key(data)
cache_key = self._cache_key(data, user_api_key_dict)
try:
global_cache.set_cache(
cache_key, session_id, ttl=SESSION_CACHE_TTL_SECONDS
@ -181,7 +181,7 @@ class PEyeEyeGuardrail(CustomGuardrail):
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return response
cache_key = self._cache_key(data)
cache_key = self._cache_key(data, user_api_key_dict)
try:
session_id = global_cache.get_cache(cache_key)
except Exception:
@ -210,6 +210,26 @@ class PEyeEyeGuardrail(CustomGuardrail):
else:
new_parts.append(part)
message.content = new_parts
# Mirror the pre-call coverage: if the model echoes placeholders
# into tool_call arguments, rehydrate those too.
tool_calls = getattr(message, "tool_calls", None)
if isinstance(tool_calls, list):
for tc in tool_calls:
fn = getattr(tc, "function", None) or (
tc.get("function") if isinstance(tc, dict) else None
)
if fn is None:
continue
args = getattr(fn, "arguments", None) if not isinstance(fn, dict) else fn.get("arguments")
if isinstance(args, str) and args:
new_args = await self._rehydrate(args, session_id)
if isinstance(fn, dict):
fn["arguments"] = new_args
else:
try:
fn.arguments = new_args
except Exception:
pass
# Clean up: drop the stateful session server-side. Stateless
# ``skey_…`` blobs hold no server-side state, so skip the DELETE.
@ -230,8 +250,18 @@ class PEyeEyeGuardrail(CustomGuardrail):
# --------------------------------------------------------------- internals
@staticmethod
def _cache_key(data: dict) -> str:
return f"peyeeye_session:{data.get('litellm_call_id') or id(data)}"
def _cache_key(data: dict, user_api_key_dict: UserAPIKeyAuth) -> str:
# ``litellm_call_id`` is sourced from the inbound ``x-litellm-call-id``
# header, so it is caller-controlled. Namespace by the authenticated
# key (server-controlled) so two callers can't collide on the same key
# and rehydrate each other's PII.
call_id = data.get("litellm_call_id") or id(data)
auth_ns = (
getattr(user_api_key_dict, "api_key", None)
or getattr(user_api_key_dict, "token", None)
or "anon"
)
return f"peyeeye_session:{auth_ns}:{call_id}"
async def _redact_batch(self, texts: List[str]) -> tuple[List[str], Optional[str]]:
"""Redact a batch of texts in a single peyeeye session.
@ -336,8 +366,13 @@ class PEyeEyeGuardrail(CustomGuardrail):
def _iter_message_text(messages: List[Dict[str, Any]]):
"""Yield (message_index, part_path, text) for every text-bearing chunk.
``part_path`` is either ``"content"`` for a plain string message or an
int index into the multimodal content list.
``part_path`` identifies where the text lives so ``_set_message_text``
can write the redacted value back:
* ``"content"`` — plain string ``content``
* ``("content", j)`` — the ``j``-th item of a multimodal content list
* ``("tool_call", k)`` — ``tool_calls[k].function.arguments``
* ``"function_call"`` — legacy ``function_call.arguments``
"""
for i, msg in enumerate(messages):
if not isinstance(msg, dict):
@ -351,15 +386,49 @@ def _iter_message_text(messages: List[Dict[str, Any]]):
if isinstance(part, dict) and part.get("type") == "text":
text = part.get("text", "")
if text:
yield i, j, text
yield i, ("content", j), text
# Tool calls carry model-visible text in ``function.arguments``; if we
# leave them alone a caller can put PII there and bypass redaction.
tool_calls = msg.get("tool_calls")
if isinstance(tool_calls, list):
for k, tc in enumerate(tool_calls):
if not isinstance(tc, dict):
continue
fn = tc.get("function")
if isinstance(fn, dict):
args = fn.get("arguments")
if isinstance(args, str) and args:
yield i, ("tool_call", k), args
fc = msg.get("function_call")
if isinstance(fc, dict):
args = fc.get("arguments")
if isinstance(args, str) and args:
yield i, "function_call", args
def _set_message_text(message: Dict[str, Any], part_path, value: str) -> None:
if part_path == "content":
message["content"] = value
return
parts = message.get("content")
if isinstance(parts, list) and isinstance(part_path, int) and part_path < len(parts):
part = parts[part_path]
if isinstance(part, dict):
part["text"] = value
if part_path == "function_call":
fc = message.get("function_call")
if isinstance(fc, dict):
fc["arguments"] = value
return
if isinstance(part_path, tuple) and len(part_path) == 2:
kind, idx = part_path
if kind == "content":
parts = message.get("content")
if isinstance(parts, list) and isinstance(idx, int) and idx < len(parts):
part = parts[idx]
if isinstance(part, dict):
part["text"] = value
return
if kind == "tool_call":
tool_calls = message.get("tool_calls")
if isinstance(tool_calls, list) and isinstance(idx, int) and idx < len(tool_calls):
tc = tool_calls[idx]
if isinstance(tc, dict):
fn = tc.get("function")
if isinstance(fn, dict):
fn["arguments"] = value

View file

@ -101,8 +101,8 @@ async def test_pre_call_redacts_messages_and_caches_session():
global_cache,
)
assert out["messages"][0]["content"] == "hi [EMAIL_1]"
assert global_cache.get_cache("peyeeye_session:call-1") == "ses_abc"
global_cache.delete_cache("peyeeye_session:call-1")
assert global_cache.get_cache("peyeeye_session:x:call-1") == "ses_abc"
global_cache.delete_cache("peyeeye_session:x:call-1")
@pytest.mark.asyncio
@ -131,8 +131,8 @@ async def test_pre_call_stateless_returns_skey():
)
sent_body = g.async_handler.post.call_args.kwargs["json"]
assert sent_body["session"] == "stateless"
assert global_cache.get_cache("peyeeye_session:call-2") == "skey_xyz"
global_cache.delete_cache("peyeeye_session:call-2")
assert global_cache.get_cache("peyeeye_session:x:call-2") == "skey_xyz"
global_cache.delete_cache("peyeeye_session:x:call-2")
@pytest.mark.asyncio
@ -190,7 +190,7 @@ async def test_pre_and_post_call_roundtrip_uses_shared_cache():
)
assert out.choices[0].message.content == "Reply to alice@acme.com"
g.async_handler.delete.assert_awaited()
assert global_cache.get_cache("peyeeye_session:rt-1") is None
assert global_cache.get_cache("peyeeye_session:x:rt-1") is None
@pytest.mark.asyncio
@ -283,6 +283,150 @@ async def test_pre_call_raises_on_length_mismatch():
)
@pytest.mark.asyncio
async def test_cache_key_isolated_per_authenticated_key():
"""Two callers sharing a litellm_call_id must not share a session entry."""
from litellm.proxy.guardrails.guardrail_hooks.peyeeye.peyeeye import (
global_cache,
)
g = PEyeEyeGuardrail(peyeeye_api_key="pk", guardrail_name="t")
g.async_handler = MagicMock()
g.async_handler.post = AsyncMock(
side_effect=[
_ok({"text": ["[EMAIL_1]"], "session_id": "ses_attacker"}),
_ok({"text": ["[EMAIL_1]"], "session_id": "ses_victim"}),
]
)
cache = DualCache()
shared_call_id = "shared-call-id"
await g.async_pre_call_hook(
UserAPIKeyAuth(api_key="attacker"),
cache,
{"messages": [{"role": "user", "content": "alice@acme.com"}],
"litellm_call_id": shared_call_id},
"completion",
)
await g.async_pre_call_hook(
UserAPIKeyAuth(api_key="victim"),
cache,
{"messages": [{"role": "user", "content": "bob@acme.com"}],
"litellm_call_id": shared_call_id},
"completion",
)
assert global_cache.get_cache(f"peyeeye_session:attacker:{shared_call_id}") == "ses_attacker"
assert global_cache.get_cache(f"peyeeye_session:victim:{shared_call_id}") == "ses_victim"
global_cache.delete_cache(f"peyeeye_session:attacker:{shared_call_id}")
global_cache.delete_cache(f"peyeeye_session:victim:{shared_call_id}")
@pytest.mark.asyncio
async def test_pre_call_redacts_tool_call_arguments():
"""tool_calls[].function.arguments must not bypass redaction."""
from litellm.proxy.guardrails.guardrail_hooks.peyeeye.peyeeye import (
global_cache,
)
g = PEyeEyeGuardrail(peyeeye_api_key="pk", guardrail_name="t")
g.async_handler = MagicMock()
g.async_handler.post = AsyncMock(
return_value=_ok(
{
"text": [
"hi [EMAIL_1]",
'{"to":"[EMAIL_1]"}',
'{"to":"[EMAIL_2]"}',
],
"session_id": "ses_tc",
}
)
)
cache = DualCache()
data = {
"messages": [
{"role": "user", "content": "hi alice@acme.com"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "send_email",
"arguments": '{"to":"alice@acme.com"}',
},
}
],
},
{
"role": "assistant",
"content": None,
"function_call": {
"name": "send_email",
"arguments": '{"to":"bob@acme.com"}',
},
},
],
"litellm_call_id": "tc-1",
}
out = await g.async_pre_call_hook(
UserAPIKeyAuth(api_key="x"), cache, data, "completion"
)
sent = g.async_handler.post.call_args.kwargs["json"]["text"]
assert "alice@acme.com" in sent[0]
assert "alice@acme.com" in sent[1]
assert "bob@acme.com" in sent[2]
assert out["messages"][1]["tool_calls"][0]["function"]["arguments"] == '{"to":"[EMAIL_1]"}'
assert out["messages"][2]["function_call"]["arguments"] == '{"to":"[EMAIL_2]"}'
global_cache.delete_cache("peyeeye_session:x:tc-1")
@pytest.mark.asyncio
async def test_post_call_rehydrates_tool_call_arguments():
"""If the model echoes placeholders into tool_call args, swap them back."""
from litellm.proxy.guardrails.guardrail_hooks.peyeeye.peyeeye import (
global_cache,
)
g = PEyeEyeGuardrail(peyeeye_api_key="pk", guardrail_name="t")
g.async_handler = MagicMock()
g.async_handler.post = AsyncMock(
return_value=_ok(
{"text": '{"to":"alice@acme.com"}', "replaced": 1}
)
)
g.async_handler.delete = AsyncMock()
user = UserAPIKeyAuth(api_key="x")
global_cache.set_cache("peyeeye_session:x:tc-out", "ses_tc", ttl=60)
response = litellm.ModelResponse()
msg = litellm.utils.Message(content=None, role="assistant")
msg.tool_calls = [
litellm.utils.ChatCompletionMessageToolCall(
id="call_1",
type="function",
function=litellm.utils.Function(
name="send_email", arguments='{"to":"[EMAIL_1]"}'
),
)
]
response.choices = [
litellm.utils.Choices(finish_reason="stop", index=0, message=msg)
]
out = await g.async_post_call_success_hook(
{"litellm_call_id": "tc-out"}, user, response
)
assert out.choices[0].message.tool_calls[0].function.arguments == '{"to":"alice@acme.com"}'
@pytest.mark.asyncio
async def test_pre_call_raises_on_unexpected_response_shape():
"""If /v1/redact returns neither str nor list for `text`, refuse to forward."""