fix(guardrails/peyeeye): redact text_completion prompt + embeddings input

`async_pre_call_hook` advertises support for `text_completion` and
`embeddings` but only walked `data["messages"]`, so PII placed in the
`/v1/completions` `prompt` or embeddings `input` flowed to the model
unredacted. Walk those fields on the way out and rehydrate
`TextCompletionResponse.choices[].text` on the way back. Embeddings
have no text response, so nothing to do post-call there.
This commit is contained in:
Tim 2026-05-10 14:25:52 -05:00
parent 5860e66730
commit 8db871f6e8
2 changed files with 155 additions and 5 deletions

View file

@ -139,13 +139,15 @@ class PEyeEyeGuardrail(CustomGuardrail):
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return data
messages: List[Dict[str, Any]] = data.get("messages") or []
text_parts = list(_iter_message_text(messages))
# Walk every text-bearing input we know about — chat ``messages``,
# tool_call arguments, ``text_completion`` ``prompt``, and embeddings
# ``input``. Anything we don't extract here would bypass redaction.
text_parts = list(_iter_data_text(data))
if not text_parts:
return data
redacted_texts, session_id = await self._redact_batch(
[t for _, _, t in text_parts]
[t for _, t in text_parts]
)
if len(redacted_texts) != len(text_parts):
raise PEyeEyeGuardrailAPIError(
@ -153,8 +155,8 @@ class PEyeEyeGuardrail(CustomGuardrail):
f"{len(text_parts)} inputs; refusing to forward partially-"
"redacted data"
)
for (msg_idx, part_path, _), redacted in zip(text_parts, redacted_texts):
_set_message_text(messages[msg_idx], part_path, redacted)
for (locator, _), redacted in zip(text_parts, redacted_texts):
_set_data_text(data, locator, redacted)
if session_id:
cache_key = self._cache_key(data, user_api_key_dict)
@ -230,6 +232,15 @@ class PEyeEyeGuardrail(CustomGuardrail):
fn.arguments = new_args
except Exception:
pass
elif isinstance(response, litellm.TextCompletionResponse):
# /v1/completions returns choices[].text rather than choices[].message.
for choice in response.choices:
text = getattr(choice, "text", None)
if isinstance(text, str) and text:
try:
choice.text = await self._rehydrate(text, session_id)
except Exception:
pass
# Clean up: drop the stateful session server-side. Stateless
# ``skey_…`` blobs hold no server-side state, so skip the DELETE.
@ -363,6 +374,64 @@ class PEyeEyeGuardrail(CustomGuardrail):
# ---------------------------------------------------------------- text helpers
def _iter_data_text(data: Dict[str, Any]):
"""Yield ``(locator, text)`` for every text-bearing input in ``data``.
Covers chat ``messages`` (incl. tool_call arguments), ``text_completion``
``prompt``, and ``embeddings`` ``input``. ``locator`` is opaque to callers
and is consumed by ``_set_data_text`` to write the redacted value back.
"""
messages = data.get("messages")
if isinstance(messages, list):
for msg_idx, sub_path, text in _iter_message_text(messages):
yield ("messages", msg_idx, sub_path), text
prompt = data.get("prompt")
if isinstance(prompt, str):
if prompt:
yield ("prompt",), prompt
elif isinstance(prompt, list):
for j, p in enumerate(prompt):
if isinstance(p, str) and p:
yield ("prompt", j), p
inp = data.get("input")
if isinstance(inp, str):
if inp:
yield ("input",), inp
elif isinstance(inp, list):
for j, v in enumerate(inp):
if isinstance(v, str) and v:
yield ("input", j), v
def _set_data_text(data: Dict[str, Any], locator, value: str) -> None:
head = locator[0]
if head == "messages":
_, msg_idx, sub_path = locator
messages = data.get("messages") or []
if isinstance(msg_idx, int) and msg_idx < len(messages):
_set_message_text(messages[msg_idx], sub_path, value)
return
if head == "prompt":
if len(locator) == 1:
data["prompt"] = value
return
_, j = locator
prompt = data.get("prompt")
if isinstance(prompt, list) and isinstance(j, int) and j < len(prompt):
prompt[j] = value
return
if head == "input":
if len(locator) == 1:
data["input"] = value
return
_, j = locator
inp = data.get("input")
if isinstance(inp, list) and isinstance(j, int) and j < len(inp):
inp[j] = value
def _iter_message_text(messages: List[Dict[str, Any]]):
"""Yield (message_index, part_path, text) for every text-bearing chunk.

View file

@ -427,6 +427,87 @@ async def test_post_call_rehydrates_tool_call_arguments():
assert out.choices[0].message.tool_calls[0].function.arguments == '{"to":"alice@acme.com"}'
@pytest.mark.asyncio
async def test_pre_call_redacts_text_completion_prompt():
"""text_completion `prompt` (str or list[str]) must be redacted."""
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]", "ping [EMAIL_2]"],
"session_id": "ses_p",
}
)
)
cache = DualCache()
data = {
"prompt": ["hi alice@acme.com", "ping bob@acme.com"],
"litellm_call_id": "p-1",
}
out = await g.async_pre_call_hook(
UserAPIKeyAuth(api_key="x"), cache, data, "text_completion"
)
assert out["prompt"] == ["hi [EMAIL_1]", "ping [EMAIL_2]"]
global_cache.delete_cache("peyeeye_session:x:p-1")
@pytest.mark.asyncio
async def test_pre_call_redacts_embeddings_input():
"""embeddings `input` (str or list[str]) must be redacted."""
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": ["[EMAIL_1]"], "session_id": "ses_e"})
)
cache = DualCache()
data = {"input": "alice@acme.com", "litellm_call_id": "e-1"}
out = await g.async_pre_call_hook(
UserAPIKeyAuth(api_key="x"), cache, data, "embeddings"
)
assert out["input"] == "[EMAIL_1]"
global_cache.delete_cache("peyeeye_session:x:e-1")
@pytest.mark.asyncio
async def test_post_call_rehydrates_text_completion_response():
"""TextCompletionResponse.choices[].text must be rehydrated."""
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": "Reply to alice@acme.com", "replaced": 1})
)
g.async_handler.delete = AsyncMock()
user = UserAPIKeyAuth(api_key="x")
global_cache.set_cache("peyeeye_session:x:tc-2", "ses_tc", ttl=60)
response = litellm.TextCompletionResponse()
response.choices = [
litellm.utils.TextChoices(
finish_reason="stop", index=0, text="Reply to [EMAIL_1]"
)
]
out = await g.async_post_call_success_hook(
{"litellm_call_id": "tc-2"}, user, response
)
assert out.choices[0].text == "Reply 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."""