mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(presidio): unmask PII tokens in Anthropic native SSE streaming bytes (#30028)
* fix(presidio): unmask PII tokens in Anthropic native SSE streaming bytes When output_parse_pii=true on the Anthropic native path (anthropic/claude-*), response chunks arrive as raw bytes in SSE format. _stream_pii_unmasking was yielding those bytes unchanged, so <PERSON_1> tokens were never replaced with the original values before reaching the caller. Add _unmask_sse_bytes_chunk to parse each data: line, find content_block_delta / text_delta events, and apply _unmask_pii_text before re-encoding. Wire it into _stream_pii_unmasking so bytes chunks are unmasked when pii_tokens exist. * fix(presidio): handle CRLF line endings and non-ASCII PII in SSE unmask Strip trailing \r before the [DONE] guard so CRLF-terminated SSE chunks don't bypass it and silently swallow a JSONDecodeError. Add ensure_ascii=False to json.dumps so non-ASCII replacement values like accented names are preserved as UTF-8 on the wire rather than being \uXXXX-escaped. Add regression tests for both cases.
This commit is contained in:
parent
bbbd3ca69d
commit
1bcf4a2691
2 changed files with 197 additions and 6 deletions
|
|
@ -1225,6 +1225,39 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
|
||||
@staticmethod
|
||||
def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: Dict[str, str]) -> bytes:
|
||||
try:
|
||||
text = chunk.decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
return chunk
|
||||
|
||||
result_lines: List[str] = []
|
||||
for line in text.split("\n"):
|
||||
line = line.rstrip("\r")
|
||||
if line.startswith("data: ") and line != "data: [DONE]":
|
||||
raw_json = line[6:]
|
||||
try:
|
||||
event = json.loads(raw_json)
|
||||
delta = event.get("delta") if isinstance(event, dict) else None
|
||||
if (
|
||||
isinstance(delta, dict)
|
||||
and event.get("type") == "content_block_delta"
|
||||
and delta.get("type") == "text_delta"
|
||||
and isinstance(delta.get("text"), str)
|
||||
):
|
||||
unmasked = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(
|
||||
delta["text"], pii_tokens
|
||||
)
|
||||
if unmasked != delta["text"]:
|
||||
event["delta"]["text"] = unmasked
|
||||
line = "data: " + json.dumps(event, ensure_ascii=False)
|
||||
except (json.JSONDecodeError, KeyError, TypeError):
|
||||
pass
|
||||
result_lines.append(line)
|
||||
|
||||
return "\n".join(result_lines).encode("utf-8")
|
||||
|
||||
async def _stream_pii_unmasking(
|
||||
self,
|
||||
response: Any,
|
||||
|
|
@ -1237,13 +1270,19 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
from litellm.main import stream_chunk_builder
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
metadata = (request_data.get("metadata") or {}) if request_data else {}
|
||||
pii_tokens: Dict[str, str] = metadata.get("pii_tokens", {})
|
||||
|
||||
remaining_chunks: List[ModelResponseStream] = []
|
||||
try:
|
||||
async for chunk in response:
|
||||
if isinstance(chunk, ModelResponseStream):
|
||||
remaining_chunks.append(chunk)
|
||||
elif isinstance(chunk, bytes):
|
||||
yield chunk # type: ignore[misc]
|
||||
if pii_tokens:
|
||||
yield self._unmask_sse_bytes_chunk(chunk, pii_tokens) # type: ignore[misc]
|
||||
else:
|
||||
yield chunk # type: ignore[misc]
|
||||
continue
|
||||
|
||||
if not remaining_chunks:
|
||||
|
|
|
|||
|
|
@ -2398,11 +2398,9 @@ async def test_anonymize_text_uses_correct_positions_no_parse_pii():
|
|||
)
|
||||
|
||||
expected = "My name is <PERSON>, my email is <EMAIL_ADDRESS>, phone <PHONE_NUMBER>"
|
||||
assert result == expected, (
|
||||
f"anonymize_text produced garbled output with PII remnants.\n"
|
||||
f"Expected: {expected!r}\n"
|
||||
f"Got: {result!r}"
|
||||
)
|
||||
assert (
|
||||
result == expected
|
||||
), f"anonymize_text produced garbled output with PII remnants.\nExpected: {expected!r}\nGot: {result!r}"
|
||||
assert masked_entity_count == {
|
||||
"PERSON": 1,
|
||||
"EMAIL_ADDRESS": 1,
|
||||
|
|
@ -2495,3 +2493,157 @@ async def test_anonymize_text_uses_correct_positions_with_parse_pii():
|
|||
assert pii_tokens.get("<PERSON_1>") == "John Smith"
|
||||
assert pii_tokens.get("<EMAIL_ADDRESS_2>") == "john@example.com"
|
||||
assert pii_tokens.get("<PHONE_NUMBER_3>") == "555-867-5309"
|
||||
|
||||
|
||||
def test_unmask_sse_bytes_chunk_replaces_text_delta():
|
||||
import json
|
||||
|
||||
pii_tokens = {"<PERSON_1>": "Bobby"}
|
||||
event = {
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "Hello <PERSON_1>, how are you?"},
|
||||
}
|
||||
chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8")
|
||||
|
||||
result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens)
|
||||
|
||||
decoded = result.decode("utf-8")
|
||||
parsed = json.loads(decoded.split("data: ", 1)[1].strip())
|
||||
assert parsed["delta"]["text"] == "Hello Bobby, how are you?"
|
||||
|
||||
|
||||
def test_unmask_sse_bytes_chunk_ignores_non_text_delta():
|
||||
import json
|
||||
|
||||
pii_tokens = {"<PERSON_1>": "Bobby"}
|
||||
|
||||
# message_start event — no delta
|
||||
event = {"type": "message_start", "message": {"id": "msg_01", "role": "assistant"}}
|
||||
chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8")
|
||||
result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens)
|
||||
assert result == chunk
|
||||
|
||||
# input_json_delta — should not be touched
|
||||
event2 = {
|
||||
"type": "content_block_delta",
|
||||
"index": 1,
|
||||
"delta": {"type": "input_json_delta", "partial_json": '{"name": "<PERSON_1>"}'},
|
||||
}
|
||||
chunk2 = ("data: " + json.dumps(event2) + "\n\n").encode("utf-8")
|
||||
result2 = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk2, pii_tokens)
|
||||
assert result2 == chunk2
|
||||
|
||||
|
||||
def test_unmask_sse_bytes_chunk_handles_malformed_json():
|
||||
chunk = b"data: {not valid json}\n\n"
|
||||
result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(
|
||||
chunk, {"<PERSON_1>": "Bobby"}
|
||||
)
|
||||
assert result == chunk
|
||||
|
||||
|
||||
def test_unmask_sse_bytes_chunk_handles_unicode_decode_error():
|
||||
chunk = b"\xff\xfe invalid utf-8"
|
||||
result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(
|
||||
chunk, {"<PERSON_1>": "Bobby"}
|
||||
)
|
||||
assert result == chunk
|
||||
|
||||
|
||||
def test_unmask_sse_bytes_chunk_non_ascii_pii_not_escaped():
|
||||
import json
|
||||
|
||||
pii_tokens = {"<PERSON_1>": "José"}
|
||||
event = {
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "Hello <PERSON_1>!"},
|
||||
}
|
||||
chunk = ("data: " + json.dumps(event) + "\n\n").encode("utf-8")
|
||||
|
||||
result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, pii_tokens)
|
||||
|
||||
decoded = result.decode("utf-8")
|
||||
assert "Jos\\u" not in decoded
|
||||
parsed = json.loads(decoded.split("data: ", 1)[1].strip())
|
||||
assert parsed["delta"]["text"] == "Hello José!"
|
||||
|
||||
|
||||
def test_unmask_sse_bytes_chunk_handles_crlf_line_endings():
|
||||
import json
|
||||
|
||||
pii_tokens = {"<PERSON_1>": "Bobby"}
|
||||
event = {
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "Hi <PERSON_1>!"},
|
||||
}
|
||||
crlf_chunk = ("data: " + json.dumps(event) + "\r\ndata: [DONE]\r\n").encode("utf-8")
|
||||
|
||||
result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(
|
||||
crlf_chunk, pii_tokens
|
||||
)
|
||||
|
||||
decoded = result.decode("utf-8")
|
||||
parsed = json.loads(decoded.split("data: ", 1)[1].split("\n")[0].strip())
|
||||
assert parsed["delta"]["text"] == "Hi Bobby!"
|
||||
assert "data: [DONE]" in decoded
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_pii_unmasking_unmaskes_bytes_chunks(mock_user_api_key):
|
||||
import json
|
||||
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
output_parse_pii=True,
|
||||
)
|
||||
|
||||
pii_tokens = {"<PERSON_1>": "Bobby"}
|
||||
request_data = {"metadata": {"pii_tokens": pii_tokens}}
|
||||
|
||||
def _make_sse_chunk(text: str) -> bytes:
|
||||
event = {
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": text},
|
||||
}
|
||||
return ("data: " + json.dumps(event) + "\n\n").encode("utf-8")
|
||||
|
||||
async def mock_stream():
|
||||
yield _make_sse_chunk("Hello <PERSON_1>!")
|
||||
yield _make_sse_chunk(" How can I help?")
|
||||
|
||||
chunks = []
|
||||
async for chunk in guardrail._stream_pii_unmasking(mock_stream(), request_data):
|
||||
chunks.append(chunk)
|
||||
|
||||
assert len(chunks) == 2
|
||||
first = chunks[0].decode("utf-8")
|
||||
first_event = json.loads(first.split("data: ", 1)[1].strip())
|
||||
assert first_event["delta"]["text"] == "Hello Bobby!"
|
||||
|
||||
second = chunks[1].decode("utf-8")
|
||||
second_event = json.loads(second.split("data: ", 1)[1].strip())
|
||||
assert second_event["delta"]["text"] == " How can I help?"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_pii_unmasking_passthrough_when_no_tokens(mock_user_api_key):
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
output_parse_pii=True,
|
||||
)
|
||||
|
||||
raw_chunk = b"data: {}\n\n"
|
||||
request_data: dict = {"metadata": {}}
|
||||
|
||||
async def mock_stream():
|
||||
yield raw_chunk
|
||||
|
||||
chunks = []
|
||||
async for chunk in guardrail._stream_pii_unmasking(mock_stream(), request_data):
|
||||
chunks.append(chunk)
|
||||
|
||||
assert chunks == [raw_chunk]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue