diff --git a/litellm/proxy/guardrails/guardrail_hooks/peyeeye/peyeeye.py b/litellm/proxy/guardrails/guardrail_hooks/peyeeye/peyeeye.py index 9a9776a1610..4631bafaa49 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/peyeeye/peyeeye.py +++ b/litellm/proxy/guardrails/guardrail_hooks/peyeeye/peyeeye.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_peyeeye.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_peyeeye.py index 9359681c158..52270172c99 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_peyeeye.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_peyeeye.py @@ -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."""