mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
c35b0e3714
commit
5860e66730
2 changed files with 230 additions and 17 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue