fix(presidio): review-round hardening for chunked analyze

- measure the chunk budget on the JSON-serialized text (non-ASCII escapes
  expand beyond raw UTF-8, so a raw-byte budget could still exceed the
  analyzer body limit)
- share the chunk fan-out semaphore per event loop and instance instead of
  per call, so many oversized blocks cannot multiply concurrent analyzer
  calls
- apply configured score thresholds and deny list per chunk BEFORE overlap
  resolution, so a below-threshold span cannot displace a detection the
  thresholds keep

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Yucheng Zhu 2026-08-27 01:19:56 -07:00
parent 2468cc9e6c
commit 795fa3554a
3 changed files with 270 additions and 168 deletions

View file

@ -68,6 +68,18 @@ class _PresidioAnonymizeResponse(TypedDict):
items: ReadOnly[NotRequired[list[_PresidioAnonymizeItem]]] items: ReadOnly[NotRequired[list[_PresidioAnonymizeItem]]]
_LoopSemaphores = dict[asyncio.AbstractEventLoop, asyncio.Semaphore]
def _json_escaped_len(text: str) -> int:
"""
Byte length of ``text`` as it appears serialized inside the JSON request
body sent to Presidio (``json.dumps`` escapes non-ASCII characters, so a
3-byte UTF-8 character can occupy 6+ bytes on the wire).
"""
return len(json.dumps(text).encode("utf-8")) - 2 # strip the surrounding quotes
class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
user_api_key_cache = None user_api_key_cache = None
ad_hoc_recognizers: list[str] | None = None ad_hoc_recognizers: list[str] | None = None
@ -141,6 +153,10 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
# Loop-bound session cache for background threads # Loop-bound session cache for background threads
self._loop_sessions: dict[asyncio.AbstractEventLoop, aiohttp.ClientSession] = {} self._loop_sessions: dict[asyncio.AbstractEventLoop, aiohttp.ClientSession] = {}
# Per-loop semaphores bounding chunked-analyze fan-out across ALL
# concurrent oversized blocks/requests on this instance, not per call
self._loop_chunk_semaphores: _LoopSemaphores = {} # mutable-ok: per-loop semaphore cache
if mock_testing is True: # for testing purposes only if mock_testing is True: # for testing purposes only
return return
@ -302,7 +318,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
text text
and len(text) > 1 and len(text) > 1
and self.mock_redacted_text is None and self.mock_redacted_text is None
and len(text.encode("utf-8")) > self.presidio_analyze_chunk_size_bytes and _json_escaped_len(text) > self.presidio_analyze_chunk_size_bytes
): ):
return await self._analyze_text_chunked( return await self._analyze_text_chunked(
text=text, text=text,
@ -434,9 +450,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
""" """
Analyze an oversized text by splitting it into overlapping chunks. Analyze an oversized text by splitting it into overlapping chunks.
Each chunk is at most ``presidio_analyze_chunk_size_bytes`` UTF-8 bytes, Each chunk serializes to at most ``presidio_analyze_chunk_size_bytes``
so every /analyze call stays below the analyzer deployment's request bytes inside the JSON request body, so every /analyze call stays below
body limit; per-chunk results are remapped onto the original text and the analyzer deployment's request body limit; per-chunk results are remapped onto the original text and
merged. Raises exactly like a single ``analyze_text`` call if any chunk merged. Raises exactly like a single ``analyze_text`` call if any chunk
fails. fails.
@ -454,9 +470,14 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
self.presidio_analyze_chunk_size_bytes, self.presidio_analyze_chunk_size_bytes,
len(text_chunks), len(text_chunks),
) )
# Bound the fan-out so a single oversized request cannot saturate the # Bound the fan-out so oversized requests cannot saturate the analyzer.
# analyzer; excess chunks wait here instead of piling onto the pool. # The semaphore is shared per event loop across every chunked call on
analyze_semaphore: Final = asyncio.Semaphore(PRESIDIO_ANALYZE_CHUNK_CONCURRENCY) # this instance, so many oversized blocks in one request (or many
# concurrent requests) still hold at most this many analyzer calls in
# flight. On the proxy's main thread the shared-session lock in
# _get_session_iterator additionally serializes the HTTP calls; the
# bound matters for loop-bound sessions (background threads).
analyze_semaphore: Final = self._get_chunk_semaphore()
async def _analyze_chunk_bounded( async def _analyze_chunk_bounded(
chunk_text: str, chunk_text: str,
@ -479,9 +500,27 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
# analyze_text only returns a non-list shape when mock_redacted_text # analyze_text only returns a non-list shape when mock_redacted_text
# is set, and the chunked path is never entered in that case. # is set, and the chunked path is never entered in that case.
typed_result = cast("list[PresidioAnalyzeResponseItem]", result) # cast-ok: gather() erases element type typed_result = cast("list[PresidioAnalyzeResponseItem]", result) # cast-ok: gather() erases element type
chunk_results.append(typed_result) # Apply the configured score thresholds and deny list BEFORE the
# overlap merge: a below-threshold detection must not win overlap
# resolution against one the thresholds would keep. The same filter
# runs again downstream in check_pii, where it is a no-op for the
# already-filtered items.
filtered_result = self.filter_analyze_results_by_score(analyze_results=typed_result)
chunk_results.append(
cast("list[PresidioAnalyzeResponseItem]", filtered_result) # cast-ok: list input yields list
)
return self._merge_chunked_analyze_results(text_chunks=text_chunks, chunk_results=chunk_results) return self._merge_chunked_analyze_results(text_chunks=text_chunks, chunk_results=chunk_results)
def _get_chunk_semaphore(self) -> asyncio.Semaphore:
"""Per-event-loop semaphore shared by all chunked analyze calls on this instance."""
loop: Final = asyncio.get_running_loop()
existing: Final = self._loop_chunk_semaphores.get(loop)
if existing is not None:
return existing
created: Final = asyncio.Semaphore(PRESIDIO_ANALYZE_CHUNK_CONCURRENCY)
self._loop_chunk_semaphores[loop] = created
return created
@staticmethod @staticmethod
def _coerce_analyze_chunk_size(value: int | None) -> int: def _coerce_analyze_chunk_size(value: int | None) -> int:
""" """
@ -489,9 +528,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
Non-positive values would either bypass chunking entirely or degenerate Non-positive values would either bypass chunking entirely or degenerate
it into per-character splits (silently disabling detection), so they are it into per-character splits (silently disabling detection), so they are
replaced by the default; values below 4 bytes (the widest UTF-8 replaced by the default; values below 4 bytes are floored to 4 and the
character) are floored to 4 so a single character always fits in a splitter always emits at least one character per chunk, so the chunked
chunk and the chunked path can never re-enter itself. path can never re-enter itself.
""" """
if not value or value <= 0: if not value or value <= 0:
return DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES return DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES
@ -504,7 +543,10 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
overlap_chars: int, overlap_chars: int,
) -> Sequence[tuple[int, str]]: ) -> Sequence[tuple[int, str]]:
""" """
Split ``text`` into chunks of at most ``chunk_size_bytes`` UTF-8 bytes. Split ``text`` into chunks whose JSON-serialized form is at most
``chunk_size_bytes`` bytes (the analyzer body limit applies to the
JSON request body, where non-ASCII characters are escaped and larger
than their raw UTF-8 encoding).
Consecutive chunks overlap by up to ``overlap_chars`` characters so a Consecutive chunks overlap by up to ``overlap_chars`` characters so a
PII entity up to that length lying across a chunk boundary is still PII entity up to that length lying across a chunk boundary is still
@ -518,15 +560,23 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
text_len: Final = len(text) text_len: Final = len(text)
start = 0 # rebind-ok: chunk cursor advances across the loop start = 0 # rebind-ok: chunk cursor advances across the loop
while start < text_len: while start < text_len:
# Byte-truncate a char-count-bounded slice, then drop the at most # Serialized length of a character is at least 1 byte, so a slice
# one trailing character the truncation split, so every chunk ends # of chunk_size_bytes characters is a sufficient search window.
# on a character boundary and holds at most chunk_size_bytes.
candidate = text[start : start + chunk_size_bytes] candidate = text[start : start + chunk_size_bytes]
chunk = candidate.encode("utf-8")[:chunk_size_bytes].decode("utf-8", errors="ignore") if _json_escaped_len(candidate) <= chunk_size_bytes:
if not chunk: chunk = candidate
# chunk_size_bytes is below one character's UTF-8 width; emit a else:
# single character rather than an empty chunk. # Largest prefix whose serialized form fits the budget.
chunk = candidate[:1] low, high = 1, len(candidate)
while low < high:
mid = (low + high + 1) // 2
if _json_escaped_len(candidate[:mid]) <= chunk_size_bytes:
low = mid
else:
high = mid - 1
# low >= 1 keeps the loop advancing even when a single
# character serializes over a (floored, tiny) budget.
chunk = candidate[:low]
end = start + len(chunk) end = start + len(chunk)
chunks.append((start, chunk)) chunks.append((start, chunk))
if end >= text_len: if end >= text_len:

View file

@ -100,11 +100,15 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail):
presidio_anonymizer_api_base=litellm_params.presidio_anonymizer_api_base, presidio_anonymizer_api_base=litellm_params.presidio_anonymizer_api_base,
presidio_language=litellm_params.presidio_language, presidio_language=litellm_params.presidio_language,
presidio_entities_deny_list=litellm_params.presidio_entities_deny_list, presidio_entities_deny_list=litellm_params.presidio_entities_deny_list,
presidio_analyze_chunk_size_bytes=litellm_params.presidio_analyze_chunk_size_bytes,
apply_to_output=False, apply_to_output=False,
) )
params.update(overrides) params.update(overrides)
callback: Final = _OPTIONAL_PresidioPIIMasking(**params) # Passed outside the heterogeneous params dict so the argument keeps
# its precise int | None type.
callback: Final = _OPTIONAL_PresidioPIIMasking(
presidio_analyze_chunk_size_bytes=litellm_params.presidio_analyze_chunk_size_bytes,
**params,
)
litellm.logging_callback_manager.add_litellm_callback(callback) litellm.logging_callback_manager.add_litellm_callback(callback)
return callback return callback

View file

@ -22,9 +22,7 @@ from litellm.types.utils import Choices, Message, ModelResponse
from litellm.exceptions import BlockedPiiEntityError from litellm.exceptions import BlockedPiiEntityError
def _make_mock_session_iterator( def _make_mock_session_iterator(json_response, status=200, content_type="application/json", text_response=""):
json_response, status=200, content_type="application/json", text_response=""
):
"""Create a mock _get_session_iterator that yields a session returning json_response.""" """Create a mock _get_session_iterator that yields a session returning json_response."""
@asynccontextmanager @asynccontextmanager
@ -100,9 +98,7 @@ def mock_cache():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_multimodal_message_format_completion_call_type( async def test_multimodal_message_format_completion_call_type(presidio_guardrail, mock_user_api_key, mock_cache):
presidio_guardrail, mock_user_api_key, mock_cache
):
""" """
Test Presidio PII masking with multimodal message format (content as list) Test Presidio PII masking with multimodal message format (content as list)
for completion call type. for completion call type.
@ -247,9 +243,7 @@ async def test_multimodal_message_format_anthropic_messages_call_type(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_multimodal_message_multiple_content_items( async def test_multimodal_message_multiple_content_items(presidio_guardrail, mock_user_api_key, mock_cache):
presidio_guardrail, mock_user_api_key, mock_cache
):
""" """
Test Presidio PII masking with multiple content items in the content list. Test Presidio PII masking with multiple content items in the content list.
""" """
@ -303,9 +297,7 @@ async def test_multimodal_message_multiple_content_items(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_mixed_string_and_list_content( async def test_mixed_string_and_list_content(presidio_guardrail, mock_user_api_key, mock_cache):
presidio_guardrail, mock_user_api_key, mock_cache
):
""" """
Test Presidio PII masking with mixed string and list content formats. Test Presidio PII masking with mixed string and list content formats.
""" """
@ -370,9 +362,7 @@ async def test_mixed_string_and_list_content(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_content_list_without_text_field( async def test_content_list_without_text_field(presidio_guardrail, mock_user_api_key, mock_cache):
presidio_guardrail, mock_user_api_key, mock_cache
):
""" """
Test Presidio PII masking gracefully handles content items without text field Test Presidio PII masking gracefully handles content items without text field
(e.g., image content items). (e.g., image content items).
@ -629,9 +619,7 @@ async def test_logging_hook_masks_the_response_too(presidio_guardrail):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_logging_only_does_not_mask_pre_call_request( async def test_logging_only_does_not_mask_pre_call_request(mock_user_api_key, mock_cache):
mock_user_api_key, mock_cache
):
""" """
A guardrail configured with `logging_only` must only mask PII for logs/traces, A guardrail configured with `logging_only` must only mask PII for logs/traces,
never for the request sent to the model. `async_pre_call_hook` should leave the never for the request sent to the model. `async_pre_call_hook` should leave the
@ -718,9 +706,7 @@ async def test_presidio_sets_guardrail_information_in_request_data():
assert "metadata" in request_data assert "metadata" in request_data
assert "standard_logging_guardrail_information" in request_data["metadata"] assert "standard_logging_guardrail_information" in request_data["metadata"]
guardrail_info_list = request_data["metadata"][ guardrail_info_list = request_data["metadata"]["standard_logging_guardrail_information"]
"standard_logging_guardrail_information"
]
assert isinstance(guardrail_info_list, list) assert isinstance(guardrail_info_list, list)
assert len(guardrail_info_list) > 0 assert len(guardrail_info_list) > 0
@ -847,20 +833,14 @@ async def test_presidio_filter_scope_initializer(monkeypatch):
import litellm.proxy.guardrails.guardrail_hooks.presidio as presidio_mod import litellm.proxy.guardrails.guardrail_hooks.presidio as presidio_mod
import litellm.proxy.guardrails.guardrail_initializers as gi import litellm.proxy.guardrails.guardrail_initializers as gi
monkeypatch.setattr( monkeypatch.setattr(presidio_mod, "_OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False)
presidio_mod, "_OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False monkeypatch.setattr(gi, "_OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False)
)
monkeypatch.setattr(
gi, "_OPTIONAL_PresidioPIIMasking", DummyGuardrail, raising=False
)
# input-only # input-only
created.clear() created.clear()
from litellm.proxy.guardrails.guardrail_initializers import initialize_presidio from litellm.proxy.guardrails.guardrail_initializers import initialize_presidio
params_input = LitellmParams( params_input = LitellmParams(guardrail="presidio", mode="pre_call", presidio_filter_scope="input")
guardrail="presidio", mode="pre_call", presidio_filter_scope="input"
)
guardrail_dict = {"guardrail_name": "g1"} guardrail_dict = {"guardrail_name": "g1"}
cb = initialize_presidio(params_input, guardrail_dict) cb = initialize_presidio(params_input, guardrail_dict)
assert cb is created[0] assert cb is created[0]
@ -868,18 +848,14 @@ async def test_presidio_filter_scope_initializer(monkeypatch):
# output-only # output-only
created.clear() created.clear()
params_output = LitellmParams( params_output = LitellmParams(guardrail="presidio", mode="pre_call", presidio_filter_scope="output")
guardrail="presidio", mode="pre_call", presidio_filter_scope="output"
)
cb = initialize_presidio(params_output, guardrail_dict) cb = initialize_presidio(params_output, guardrail_dict)
assert len(created) == 1 assert len(created) == 1
assert created[0].apply_to_output is True assert created[0].apply_to_output is True
# both -> expect two callbacks (input + output) # both -> expect two callbacks (input + output)
created.clear() created.clear()
params_both = LitellmParams( params_both = LitellmParams(guardrail="presidio", mode="pre_call", presidio_filter_scope="both")
guardrail="presidio", mode="pre_call", presidio_filter_scope="both"
)
cb = initialize_presidio(params_both, guardrail_dict) cb = initialize_presidio(params_both, guardrail_dict)
assert len(created) == 2 assert len(created) == 2
assert any(not c.apply_to_output for c in created) assert any(not c.apply_to_output for c in created)
@ -887,9 +863,7 @@ async def test_presidio_filter_scope_initializer(monkeypatch):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_empty_content_handling( async def test_empty_content_handling(presidio_guardrail, mock_user_api_key, mock_cache):
presidio_guardrail, mock_user_api_key, mock_cache
):
""" """
Test that Presidio handles empty content gracefully. Test that Presidio handles empty content gracefully.
@ -945,9 +919,7 @@ async def test_empty_content_handling(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_whitespace_only_content( async def test_whitespace_only_content(presidio_guardrail, mock_user_api_key, mock_cache):
presidio_guardrail, mock_user_api_key, mock_cache
):
""" """
Test that Presidio handles whitespace-only content gracefully. Test that Presidio handles whitespace-only content gracefully.
@ -1142,9 +1114,7 @@ async def test_analyze_text_list_with_non_dict_items():
"invalid_string_item", "invalid_string_item",
{"entity_type": "EMAIL", "start": 10, "end": 25, "score": 0.85}, {"entity_type": "EMAIL", "start": 10, "end": 25, "score": 0.85},
] ]
with patch.object( with patch.object(presidio, "_get_session_iterator", _make_mock_session_iterator(json_response)):
presidio, "_get_session_iterator", _make_mock_session_iterator(json_response)
):
result = await presidio.analyze_text( result = await presidio.analyze_text(
text="some text", text="some text",
presidio_config=None, presidio_config=None,
@ -1156,9 +1126,7 @@ async def test_analyze_text_list_with_non_dict_items():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_tool_calling_complete_scenario( async def test_tool_calling_complete_scenario(presidio_guardrail, mock_user_api_key, mock_cache):
presidio_guardrail, mock_user_api_key, mock_cache
):
""" """
Test complete tool calling scenario with PII in user message. Test complete tool calling scenario with PII in user message.
@ -1224,9 +1192,7 @@ def test_filter_drops_low_score_detection():
mock_testing=True, mock_testing=True,
presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8}, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8},
) )
analyze_results = [ analyze_results = [{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.7, "start": 0, "end": 4}]
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.7, "start": 0, "end": 4}
]
filtered = guardrail.filter_analyze_results_by_score(analyze_results) filtered = guardrail.filter_analyze_results_by_score(analyze_results)
assert filtered == [] assert filtered == []
@ -1240,9 +1206,7 @@ def test_filter_preserves_high_score_detection():
mock_testing=True, mock_testing=True,
presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8}, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.8},
) )
analyze_results = [ analyze_results = [{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.9, "start": 0, "end": 4}]
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.9, "start": 0, "end": 4}
]
filtered = guardrail.filter_analyze_results_by_score(analyze_results) filtered = guardrail.filter_analyze_results_by_score(analyze_results)
assert len(filtered) == 1 assert len(filtered) == 1
@ -1379,15 +1343,11 @@ def test_blocking_respects_threshold_filter():
presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.9}, presidio_score_thresholds={PiiEntityType.CREDIT_CARD: 0.9},
) )
low_score_results = [ low_score_results = [{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.7, "start": 0, "end": 4}]
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.7, "start": 0, "end": 4}
]
filtered = guardrail.filter_analyze_results_by_score(low_score_results) filtered = guardrail.filter_analyze_results_by_score(low_score_results)
guardrail.raise_exception_if_blocked_entities_detected(filtered) guardrail.raise_exception_if_blocked_entities_detected(filtered)
high_score_results = [ high_score_results = [{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.95, "start": 0, "end": 4}]
{"entity_type": PiiEntityType.CREDIT_CARD, "score": 0.95, "start": 0, "end": 4}
]
filtered_high = guardrail.filter_analyze_results_by_score(high_score_results) filtered_high = guardrail.filter_analyze_results_by_score(high_score_results)
with pytest.raises(BlockedPiiEntityError): with pytest.raises(BlockedPiiEntityError):
guardrail.raise_exception_if_blocked_entities_detected(filtered_high) guardrail.raise_exception_if_blocked_entities_detected(filtered_high)
@ -1448,9 +1408,7 @@ async def test_get_session_iterator_thread_safety(presidio_guardrail):
# Run the background thread test # Run the background thread test
bg_future = asyncio.Future() bg_future = asyncio.Future()
t = threading.Thread( t = threading.Thread(target=thread_target, args=(asyncio.get_running_loop(), bg_future))
target=thread_target, args=(asyncio.get_running_loop(), bg_future)
)
t.start() t.start()
t.join() t.join()
@ -1659,9 +1617,7 @@ async def test_anonymize_text_non_json_content_type():
) )
with patch.object(guardrail, "_get_session_iterator", mock_iterator): with patch.object(guardrail, "_get_session_iterator", mock_iterator):
with pytest.raises( with pytest.raises(Exception, match="Presidio anonymizer returned non-JSON Content-Type"):
Exception, match="Presidio anonymizer returned non-JSON Content-Type"
):
await guardrail.anonymize_text( await guardrail.anonymize_text(
text="Hello world", text="Hello world",
analyze_results=[{"start": 0, "end": 5, "entity_type": "PERSON"}], analyze_results=[{"start": 0, "end": 5, "entity_type": "PERSON"}],
@ -1719,9 +1675,7 @@ async def test_pii_tokens_stored_in_metadata_not_top_level(presidio_guardrail):
mock_cache = DualCache() mock_cache = DualCache()
test_data = { test_data = {
"messages": [ "messages": [{"role": "user", "content": "My name is John and my phone is 555-123-4567"}],
{"role": "user", "content": "My name is John and my phone is 555-123-4567"}
],
"model": "claude-haiku-4-5-20251001", "model": "claude-haiku-4-5-20251001",
"metadata": {}, "metadata": {},
} }
@ -1870,9 +1824,7 @@ async def test_metadata_none_does_not_crash():
) )
# No pii_tokens to unmask, so content stays as-is # No pii_tokens to unmask, so content stays as-is
assert ( assert response.choices[0].message.content == f"Hello {token_key}, how can I help you?"
response.choices[0].message.content == f"Hello {token_key}, how can I help you?"
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@ -2049,9 +2001,7 @@ async def test_anthropic_native_response_unmasking():
response=anthropic_response, response=anthropic_response,
) )
assert result["content"][0]["text"] == ( assert result["content"][0]["text"] == ("Hello John Smith, your number is 555-123-4567.")
"Hello John Smith, your number is 555-123-4567."
)
@pytest.mark.asyncio @pytest.mark.asyncio
@ -2170,9 +2120,7 @@ async def test_streaming_bytes_chunks_are_yielded_not_discarded():
): ):
chunks.append(chunk) chunks.append(chunk)
assert any( assert any(isinstance(c, bytes) for c in chunks), "bytes chunks must not be discarded"
isinstance(c, bytes) for c in chunks
), "bytes chunks must not be discarded"
assert byte_chunk in chunks assert byte_chunk in chunks
@ -2282,9 +2230,7 @@ async def test_apply_to_output_streaming_mixed_chunks_flushes_and_warns():
mock_user_api_key = UserAPIKeyAuth(api_key="test-key") mock_user_api_key = UserAPIKeyAuth(api_key="test-key")
received = [] received = []
with patch( with patch("litellm.proxy.guardrails.guardrail_hooks.presidio.verbose_proxy_logger") as mock_logger:
"litellm.proxy.guardrails.guardrail_hooks.presidio.verbose_proxy_logger"
) as mock_logger:
async for chunk in guardrail.async_post_call_streaming_iterator_hook( async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_user_api_key, user_api_key_dict=mock_user_api_key,
response=mock_stream(), response=mock_stream(),
@ -2396,9 +2342,7 @@ async def test_apply_to_output_streaming_bytes_only_logs_warning():
mock_user_api_key = UserAPIKeyAuth(api_key="test-key") mock_user_api_key = UserAPIKeyAuth(api_key="test-key")
collected = [] collected = []
with patch( with patch("litellm.proxy.guardrails.guardrail_hooks.presidio.verbose_proxy_logger") as mock_logger:
"litellm.proxy.guardrails.guardrail_hooks.presidio.verbose_proxy_logger"
) as mock_logger:
async for chunk in guardrail.async_post_call_streaming_iterator_hook( async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_user_api_key, user_api_key_dict=mock_user_api_key,
response=mock_stream(), response=mock_stream(),
@ -2521,10 +2465,7 @@ async def test_output_parse_pii_streaming_responses_completed_event_unmasked(
collected.append(chunk) collected.append(chunk)
assert collected == [completed_event] assert collected == [completed_event]
assert ( assert collected[0].response.output[0].content[0].text == "Reach me at john@example.com today."
collected[0].response.output[0].content[0].text
== "Reach me at john@example.com today."
)
@pytest.mark.asyncio @pytest.mark.asyncio
@ -2587,9 +2528,7 @@ async def test_anonymize_text_uses_correct_positions_no_parse_pii():
original text using those positions, which produces garbled output original text using those positions, which produces garbled output
with remnants of original PII data. with remnants of original PII data.
""" """
original_text = ( original_text = "My name is John Smith, my email is john@example.com, phone 555-867-5309"
"My name is John Smith, my email is john@example.com, phone 555-867-5309"
)
# Positions as returned by the analyzer (reference original text) # Positions as returned by the analyzer (reference original text)
analyze_results = [ analyze_results = [
{"end": 51, "entity_type": "EMAIL_ADDRESS", "score": 1.0, "start": 35}, {"end": 51, "entity_type": "EMAIL_ADDRESS", "score": 1.0, "start": 35},
@ -2644,9 +2583,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>" expected = "My name is <PERSON>, my email is <EMAIL_ADDRESS>, phone <PHONE_NUMBER>"
assert ( assert result == expected, (
result == expected f"anonymize_text produced garbled output with PII remnants.\nExpected: {expected!r}\nGot: {result!r}"
), f"anonymize_text produced garbled output with PII remnants.\nExpected: {expected!r}\nGot: {result!r}" )
assert masked_entity_count == { assert masked_entity_count == {
"PERSON": 1, "PERSON": 1,
"EMAIL_ADDRESS": 1, "EMAIL_ADDRESS": 1,
@ -2665,9 +2604,7 @@ async def test_anonymize_text_uses_correct_positions_with_parse_pii():
tokens and the pii_tokens mapping, not positions from anonymizer items tokens and the pii_tokens mapping, not positions from anonymizer items
(which reference the anonymized output text). (which reference the anonymized output text).
""" """
original_text = ( original_text = "My name is John Smith, my email is john@example.com, phone 555-867-5309"
"My name is John Smith, my email is john@example.com, phone 555-867-5309"
)
analyze_results = [ analyze_results = [
{"end": 51, "entity_type": "EMAIL_ADDRESS", "score": 1.0, "start": 35}, {"end": 51, "entity_type": "EMAIL_ADDRESS", "score": 1.0, "start": 35},
{"end": 21, "entity_type": "PERSON", "score": 0.85, "start": 11}, {"end": 21, "entity_type": "PERSON", "score": 0.85, "start": 11},
@ -2783,17 +2720,13 @@ def test_unmask_sse_bytes_chunk_ignores_non_text_delta():
def test_unmask_sse_bytes_chunk_handles_malformed_json(): def test_unmask_sse_bytes_chunk_handles_malformed_json():
chunk = b"data: {not valid json}\n\n" chunk = b"data: {not valid json}\n\n"
result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk( result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, {"<PERSON_1>": "Bobby"})
chunk, {"<PERSON_1>": "Bobby"}
)
assert result == chunk assert result == chunk
def test_unmask_sse_bytes_chunk_handles_unicode_decode_error(): def test_unmask_sse_bytes_chunk_handles_unicode_decode_error():
chunk = b"\xff\xfe invalid utf-8" chunk = b"\xff\xfe invalid utf-8"
result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk( result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(chunk, {"<PERSON_1>": "Bobby"})
chunk, {"<PERSON_1>": "Bobby"}
)
assert result == chunk assert result == chunk
@ -2827,9 +2760,7 @@ def test_unmask_sse_bytes_chunk_handles_crlf_line_endings():
} }
crlf_chunk = ("data: " + json.dumps(event) + "\r\ndata: [DONE]\r\n").encode("utf-8") crlf_chunk = ("data: " + json.dumps(event) + "\r\ndata: [DONE]\r\n").encode("utf-8")
result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk( result = _OPTIONAL_PresidioPIIMasking._unmask_sse_bytes_chunk(crlf_chunk, pii_tokens)
crlf_chunk, pii_tokens
)
decoded = result.decode("utf-8") decoded = result.decode("utf-8")
parsed = json.loads(decoded.split("data: ", 1)[1].split("\n")[0].strip()) parsed = json.loads(decoded.split("data: ", 1)[1].split("\n")[0].strip())
@ -2948,10 +2879,7 @@ def _make_marker_session_iterator(
if url.endswith("analyze"): if url.endswith("analyze"):
recorded_analyze_payloads.append(payload) recorded_analyze_payloads.append(payload)
text = payload["text"] text = payload["text"]
if ( if analyzer_body_limit_bytes is not None and len(text.encode("utf-8")) > analyzer_body_limit_bytes:
analyzer_body_limit_bytes is not None
and len(text.encode("utf-8")) > analyzer_body_limit_bytes
):
return MockResponse( return MockResponse(
413, 413,
{ {
@ -2973,9 +2901,7 @@ def _make_marker_session_iterator(
if recorded_anonymize_payloads is not None: if recorded_anonymize_payloads is not None:
recorded_anonymize_payloads.append(payload) recorded_anonymize_payloads.append(payload)
text = payload["text"] text = payload["text"]
items = sorted( items = sorted(payload["analyzer_results"], key=lambda r: r["start"], reverse=True)
payload["analyzer_results"], key=lambda r: r["start"], reverse=True
)
for r in items: for r in items:
text = text[: r["start"]] + "<" + r["entity_type"] + ">" + text[r["end"] :] text = text[: r["start"]] + "<" + r["entity_type"] + ">" + text[r["end"] :]
return MockResponse( return MockResponse(
@ -3015,9 +2941,7 @@ def _oversized_marker_text():
def test_split_text_for_analysis_offsets_and_byte_budget(): def test_split_text_for_analysis_offsets_and_byte_budget():
text = " ".join(f"word{i}" for i in range(200)) text = " ".join(f"word{i}" for i in range(200))
chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis( chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=100, overlap_chars=20)
text=text, chunk_size_bytes=100, overlap_chars=20
)
assert len(chunks) > 1 assert len(chunks) > 1
for offset, chunk in chunks: for offset, chunk in chunks:
assert len(chunk.encode("utf-8")) <= 100 assert len(chunk.encode("utf-8")) <= 100
@ -3032,9 +2956,7 @@ def test_split_text_for_analysis_offsets_and_byte_budget():
def test_split_text_for_analysis_multibyte_characters(): def test_split_text_for_analysis_multibyte_characters():
text = "émoji🙂 çafé " * 120 text = "émoji🙂 çafé " * 120
chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis( chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=64, overlap_chars=8)
text=text, chunk_size_bytes=64, overlap_chars=8
)
assert len(chunks) > 1 assert len(chunks) > 1
for offset, chunk in chunks: for offset, chunk in chunks:
assert len(chunk.encode("utf-8")) <= 64 assert len(chunk.encode("utf-8")) <= 64
@ -3044,9 +2966,7 @@ def test_split_text_for_analysis_multibyte_characters():
def test_split_text_for_analysis_under_budget_returns_single_chunk(): def test_split_text_for_analysis_under_budget_returns_single_chunk():
text = "short text" text = "short text"
chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis( chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=100, overlap_chars=20)
text=text, chunk_size_bytes=100, overlap_chars=20
)
assert chunks == [(0, text)] assert chunks == [(0, text)]
@ -3055,12 +2975,8 @@ async def test_analyze_text_single_call_when_under_limit():
guardrail = _chunking_guardrail(chunk_size_bytes=10_000) guardrail = _chunking_guardrail(chunk_size_bytes=10_000)
payloads = [] payloads = []
text = f"my card is {CHUNK_MARKER_ONE} thanks" text = f"my card is {CHUNK_MARKER_ONE} thanks"
with patch.object( with patch.object(guardrail, "_get_session_iterator", _make_marker_session_iterator(payloads)):
guardrail, "_get_session_iterator", _make_marker_session_iterator(payloads) results = await guardrail.analyze_text(text=text, presidio_config=None, request_data={})
):
results = await guardrail.analyze_text(
text=text, presidio_config=None, request_data={}
)
assert len(payloads) == 1 assert len(payloads) == 1
assert payloads[0]["text"] == text assert payloads[0]["text"] == text
assert len(results) == 1 assert len(results) == 1
@ -3088,9 +3004,7 @@ async def test_analyze_text_chunks_oversized_text_and_remaps_offsets():
"_get_session_iterator", "_get_session_iterator",
_make_marker_session_iterator(payloads, analyzer_body_limit_bytes=100), _make_marker_session_iterator(payloads, analyzer_body_limit_bytes=100),
): ):
results = await guardrail.analyze_text( results = await guardrail.analyze_text(text=text, presidio_config=None, request_data={})
text=text, presidio_config=None, request_data={}
)
assert len(payloads) > 1 assert len(payloads) > 1
for payload in payloads: for payload in payloads:
assert len(payload["text"].encode("utf-8")) <= 100 assert len(payload["text"].encode("utf-8")) <= 100
@ -3118,9 +3032,7 @@ async def test_check_pii_masks_oversized_text_with_chunking():
recorded_anonymize_payloads=anonymize_payloads, recorded_anonymize_payloads=anonymize_payloads,
), ),
): ):
masked = await guardrail.check_pii( masked = await guardrail.check_pii(text=text, output_parse_pii=False, presidio_config=None, request_data={})
text=text, output_parse_pii=False, presidio_config=None, request_data={}
)
assert CHUNK_MARKER_ONE not in masked assert CHUNK_MARKER_ONE not in masked
assert CHUNK_MARKER_TWO not in masked assert CHUNK_MARKER_TWO not in masked
assert masked.count("<CREDIT_CARD>") == 2 assert masked.count("<CREDIT_CARD>") == 2
@ -3177,9 +3089,7 @@ async def test_analyze_text_chunked_failure_stays_fail_closed():
_make_marker_session_iterator(payloads, analyzer_body_limit_bytes=10), _make_marker_session_iterator(payloads, analyzer_body_limit_bytes=10),
): ):
with pytest.raises(GuardrailRaisedException, match="HTTP 413"): with pytest.raises(GuardrailRaisedException, match="HTTP 413"):
await guardrail.analyze_text( await guardrail.analyze_text(text=text, presidio_config=None, request_data={})
text=text, presidio_config=None, request_data={}
)
def test_presidio_analyze_chunk_size_default_and_validation(): def test_presidio_analyze_chunk_size_default_and_validation():
@ -3188,14 +3098,10 @@ def test_presidio_analyze_chunk_size_default_and_validation():
guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True) guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True)
assert guardrail.presidio_analyze_chunk_size_bytes == DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES assert guardrail.presidio_analyze_chunk_size_bytes == DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES
nonpositive = _OPTIONAL_PresidioPIIMasking( nonpositive = _OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=-5)
mock_testing=True, presidio_analyze_chunk_size_bytes=-5
)
assert nonpositive.presidio_analyze_chunk_size_bytes == DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES assert nonpositive.presidio_analyze_chunk_size_bytes == DEFAULT_PRESIDIO_ANALYZE_CHUNK_SIZE_BYTES
custom = _OPTIONAL_PresidioPIIMasking( custom = _OPTIONAL_PresidioPIIMasking(mock_testing=True, presidio_analyze_chunk_size_bytes=1234)
mock_testing=True, presidio_analyze_chunk_size_bytes=1234
)
assert custom.presidio_analyze_chunk_size_bytes == 1234 assert custom.presidio_analyze_chunk_size_bytes == 1234
@ -3276,9 +3182,7 @@ async def test_tiny_chunk_size_with_multibyte_text_terminates():
guardrail = _chunking_guardrail(chunk_size_bytes=1) guardrail = _chunking_guardrail(chunk_size_bytes=1)
assert guardrail.presidio_analyze_chunk_size_bytes == 4 assert guardrail.presidio_analyze_chunk_size_bytes == 4
payloads = [] payloads = []
with patch.object( with patch.object(guardrail, "_get_session_iterator", _make_marker_session_iterator(payloads)):
guardrail, "_get_session_iterator", _make_marker_session_iterator(payloads)
):
results = await guardrail.analyze_text( results = await guardrail.analyze_text(
text="\U0001f642\U0001f642\U0001f642ab", presidio_config=None, request_data={} text="\U0001f642\U0001f642\U0001f642ab", presidio_config=None, request_data={}
) )
@ -3332,3 +3236,147 @@ async def test_chunked_analyze_concurrency_is_bounded():
await guardrail.analyze_text(text="a" * 400, presidio_config=None, request_data={}) await guardrail.analyze_text(text="a" * 400, presidio_config=None, request_data={})
assert state["peak"] >= 2 assert state["peak"] >= 2
assert state["peak"] <= PRESIDIO_ANALYZE_CHUNK_CONCURRENCY assert state["peak"] <= PRESIDIO_ANALYZE_CHUNK_CONCURRENCY
def test_split_text_accounts_for_json_body_expansion():
"""Non-ASCII text expands under JSON escaping; the budget must apply to the
serialized form or a chunk can still exceed the analyzer body limit."""
import json as json_module
text = "これは個人情報テストです。" * 200 # 3-byte UTF-8 chars, 6-byte escapes
budget = 1000
chunks = _OPTIONAL_PresidioPIIMasking._split_text_for_analysis(text=text, chunk_size_bytes=budget, overlap_chars=8)
assert len(chunks) > 1
for offset, chunk in chunks:
assert len(json_module.dumps(chunk).encode("utf-8")) - 2 <= budget
assert text[offset : offset + len(chunk)] == chunk
# full coverage: last chunk reaches the end of the text
last_offset, last_chunk = chunks[-1]
assert last_offset + len(last_chunk) == len(text)
@pytest.mark.asyncio
async def test_chunked_analyze_applies_score_threshold_before_merge():
"""A below-threshold long span must not win overlap resolution against an
above-threshold detection of the same type (it would then be dropped by the
downstream threshold filter, leaving the entity unmasked)."""
guardrail = _chunking_guardrail(
chunk_size_bytes=100,
presidio_score_thresholds={"CREDIT_CARD": 0.6},
)
marker_text = "x" * 40 + CHUNK_MARKER_ONE + "x" * 80 # single chunked text
@asynccontextmanager
async def mock_iterator():
class MockResponse:
status = 200
content_type = "application/json"
headers = {"Content-Type": "application/json"}
def __init__(self, body):
self._body = body
async def text(self):
import json as json_module
return json_module.dumps(self._body)
async def json(self):
return self._body
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
class MockSession:
def post(self, url, json=None, headers=None):
text = json["text"]
idx = text.find(CHUNK_MARKER_ONE)
if idx == -1:
return MockResponse([])
return MockResponse(
[
# long, below-threshold span engulfing the marker
{
"entity_type": "CREDIT_CARD",
"start": max(idx - 5, 0),
"end": idx + len(CHUNK_MARKER_ONE) + 5,
"score": 0.3,
},
# the true, above-threshold detection
{
"entity_type": "CREDIT_CARD",
"start": idx,
"end": idx + len(CHUNK_MARKER_ONE),
"score": 0.9,
},
]
)
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
yield MockSession()
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
results = await guardrail.analyze_text(text=marker_text, presidio_config=None, request_data={})
kept = [r for r in results if r.get("entity_type") == "CREDIT_CARD"]
assert any(r.get("score") == 0.9 for r in kept), kept
assert all(r.get("score") != 0.3 for r in kept), kept
@pytest.mark.asyncio
async def test_chunk_fanout_bound_is_shared_across_concurrent_calls():
"""The chunk semaphore is per event loop and instance, so several oversized
blocks analyzed concurrently share ONE bound instead of getting 8 each."""
from litellm.constants import PRESIDIO_ANALYZE_CHUNK_CONCURRENCY
guardrail = _chunking_guardrail(chunk_size_bytes=10)
state = {"active": 0, "peak": 0}
@asynccontextmanager
async def mock_iterator():
class MockResponse:
status = 200
content_type = "application/json"
headers = {"Content-Type": "application/json"}
async def text(self):
return "[]"
async def json(self):
state["active"] += 1
state["peak"] = max(state["peak"], state["active"])
await asyncio.sleep(0.005)
state["active"] -= 1
return []
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
class MockSession:
def post(self, url, json=None, headers=None):
return MockResponse()
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
yield MockSession()
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
await asyncio.gather(
*(guardrail.analyze_text(text="a" * 400, presidio_config=None, request_data={}) for _ in range(4))
)
assert state["peak"] >= 2
assert state["peak"] <= PRESIDIO_ANALYZE_CHUNK_CONCURRENCY