mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
5860e66730
commit
8db871f6e8
2 changed files with 155 additions and 5 deletions
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue