From a0b26d2c3c498f6c11ec4548aa986183eec18f20 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Tue, 30 Jun 2026 14:37:11 -0700 Subject: [PATCH] =?UTF-8?q?Revert=20"fix(presidio):=20stream=20SSE=20outpu?= =?UTF-8?q?t=20incrementally=20instead=20of=20buffering=20t=E2=80=A6"=20(#?= =?UTF-8?q?31764)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This reverts commit 94936a3922aeb8aa9a4d5928e63ea6f15be6f098. --- .../guardrails/guardrail_hooks/presidio.py | 523 ++++--------- .../guardrail_hooks/test_presidio.py | 712 +----------------- 2 files changed, 152 insertions(+), 1083 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index e60b6233038..95876a55eab 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -17,8 +17,6 @@ from typing import ( TYPE_CHECKING, Any, AsyncGenerator, - Awaitable, - Callable, Dict, List, Literal, @@ -56,14 +54,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.presidio import ( PresidioAnalyzeRequest, PresidioAnalyzeResponseItem, ) -from litellm.types.utils import ( - ChatCompletionDeltaToolCall, - Delta, - Function, - FunctionCall, - GuardrailStatus, - StreamingChoices, -) +from litellm.types.utils import GuardrailStatus, StreamingChoices from litellm.utils import ( EmbeddingResponse, ImageResponse, @@ -71,17 +62,6 @@ from litellm.utils import ( ModelResponseStream, ) -# Trailing context (chars) the streaming output-masking path keeps buffered past -# a sentence boundary before emitting, so a PII entity that straddles the -# boundary is seen in full by Presidio and is never split across two analyze -# calls. It bounds the largest single entity the incremental path can mask -# without leaking; an entity longer than this could still be split. -_PRESIDIO_STREAM_MARGIN = 96 -# Hard cap on buffered un-emitted output. Past this with no sentence boundary, -# stable prefixes are flushed; if stability cannot be proven, the ambiguous -# prefix is dropped while retaining the trailing margin. -_PRESIDIO_STREAM_MAX_BUFFER = 2000 - class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): user_api_key_cache = None @@ -113,10 +93,6 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): self.mock_redacted_text = mock_redacted_text self.output_parse_pii = output_parse_pii or False self.apply_to_output = apply_to_output - # Streaming output-masking safety window; instance attributes so tests can - # exercise incremental flushing with short content (see _mask_emit_decision). - self._stream_mask_margin = _PRESIDIO_STREAM_MARGIN - self._stream_mask_max_buffer = _PRESIDIO_STREAM_MAX_BUFFER # When output_parse_pii or apply_to_output is enabled, the guardrail must # also run on post_call to unmask/mask the response. Expand the event_hook @@ -1072,352 +1048,80 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): ) return response - @staticmethod - def _unmask_hold_len(text: str, token_keys: Any) -> int: - """Length of the trailing run of ``text`` that could still grow into a - PII placeholder token, so the unmask path holds it until the next chunk - completes (or aborts) the token instead of emitting a half-written - ````.""" - keys = tuple(token_keys) - if not text or not keys: - return 0 - longest = max(len(key) for key in keys) - for start in range(max(0, len(text) - (longest - 1)), len(text)): - suffix = text[start:] - if any(key.startswith(suffix) for key in keys if len(suffix) < len(key)): - return len(text) - start - return 0 - - @staticmethod - def _mask_boundaries(text: str) -> tuple[int, ...]: - """Candidate flush points: a newline, or a sentence terminator already - followed by whitespace. A terminator at the very end of the buffer is - excluded because the next chunk may continue the token (``jane.`` + - ``doe@example.com``); it becomes a boundary once the whitespace arrives. - A boundary is only a *candidate* here; ``_mask_emit_decision`` still - confirms via a stability check that no entity straddles it.""" - return tuple( - i + 1 - for i in range(len(text)) - if text[i] == "\n" or (text[i] in ".!?" and i + 1 < len(text) and text[i + 1].isspace()) - ) - - async def _mask_emit_decision( - self, - buffer: str, - terminal: bool, - transform: "Callable[[str], Awaitable[str]]", - ) -> "tuple[str, str]": - """Decide how much of ``buffer`` is safe to mask and emit now, returning - ``(masked_emit, hold_raw)``. - - A sentence boundary is not trusted blindly (it can fall inside a name - with an initial or an address spanning a newline). Instead a prefix is - emitted only when masking it in isolation matches the corresponding - prefix of masking the whole buffer, with at least ``_PRESIDIO_STREAM_MARGIN`` - characters of lookahead still buffered past the cut. That guarantees any - entity overlapping the cut is present in full when the buffer is analyzed, - so a straddling entity makes the prefixes differ and the cut is held. - Past ``_PRESIDIO_STREAM_MAX_BUFFER`` with no sentence boundary the buffer - first tries stable forced cuts and then drops the ambiguous prefix while - retaining the trailing margin, so a failed stability check cannot grow - the held buffer without bound.""" - if terminal: - return (await transform(buffer) if buffer else ""), "" - margin = self._stream_mask_margin - forced_cut = ( - len(buffer) - margin if len(buffer) > self._stream_mask_max_buffer and len(buffer) > margin else None - ) - cuts = [index for index in self._mask_boundaries(buffer) if len(buffer) - index >= margin] - if forced_cut is not None: - verbose_proxy_logger.warning( - "Presidio apply_to_output: buffered %d streamed chars with no " - "sentence boundary; bounding held stream state.", - len(buffer), - ) - cuts.append(forced_cut) - cuts.extend(index for index in range(forced_cut, len(buffer)) if buffer[index].isspace()) - if cuts: - masked_full = await transform(buffer) - for index in sorted(set(cuts), reverse=True): - masked_prefix = await transform(buffer[:index]) - if masked_full.startswith(masked_prefix): - return masked_prefix, buffer[index:] - if forced_cut is not None: - return "", buffer[forced_cut:] - return "", buffer - - @staticmethod - def _accumulate_tool_calls( - tool_acc: dict[int, dict[int, dict[str, Optional[str]]]], - choice_index: int, - tool_calls: list[Any], - ) -> None: - choice_acc = tool_acc.setdefault(choice_index, {}) # mutable-ok: streaming tool-call accumulator - for tool_call in tool_calls: - entry = choice_acc.setdefault( # mutable-ok: streaming tool-call accumulator - getattr(tool_call, "index", 0) or 0, - {"id": None, "type": None, "name": None, "args": ""}, - ) - if getattr(tool_call, "id", None): - entry["id"] = tool_call.id - if getattr(tool_call, "type", None): - entry["type"] = tool_call.type - function = getattr(tool_call, "function", None) - if function is not None: - if getattr(function, "name", None): - entry["name"] = function.name - arguments = getattr(function, "arguments", None) - if isinstance(arguments, str): - entry["args"] = (entry["args"] or "") + arguments - - @staticmethod - def _accumulate_function_call( - func_acc: dict[int, dict[str, Optional[str]]], - choice_index: int, - function_call: Any, - ) -> None: - entry = func_acc.setdefault( # mutable-ok: streaming function-call accumulator - choice_index, {"name": None, "args": ""} - ) - if getattr(function_call, "name", None): - entry["name"] = function_call.name - arguments = getattr(function_call, "arguments", None) - if isinstance(arguments, str): - entry["args"] = (entry["args"] or "") + arguments - - @staticmethod - async def _build_tool_calls( - choice_acc: dict[int, dict[str, Optional[str]]], - transform: "Callable[[str], Awaitable[str]]", - ) -> list[ChatCompletionDeltaToolCall]: - return [ - ChatCompletionDeltaToolCall( - index=tool_index, - id=entry["id"], - type=entry["type"], - function=Function( - name=entry["name"], - arguments=(await transform(entry["args"]) if entry["args"] else ""), - ), - ) - for tool_index, entry in sorted(choice_acc.items()) - ] - - @staticmethod - async def _build_function_call( - entry: Optional[dict[str, Optional[str]]], - transform: "Callable[[str], Awaitable[str]]", - ) -> Optional[FunctionCall]: - if entry is None: - return None - return FunctionCall( - name=entry["name"], - arguments=await transform(entry["args"]) if entry["args"] else "", - ) - - async def _rewrite_chat_chunk( - self, - chunk: ModelResponseStream, - content_buffers: dict[int, str], - tool_acc: dict[int, dict[int, dict[str, Optional[str]]]], - func_acc: dict[int, dict[str, Optional[str]]], - transform: "Callable[[str], Awaitable[str]]", - emit_content: "Callable[[str, bool], Awaitable[tuple[str, str]]]", - ) -> None: - """Transform one streaming chat chunk in place: text content is masked / - unmasked and emitted as soon as ``emit_content`` deems a prefix safe (it - returns the already-transformed text to emit plus the raw remainder to - hold), while tool-call and function-call argument fragments are - accumulated and emitted, fully transformed, on the chunk that closes the - choice.""" - for choice in chunk.choices: - index = getattr(choice, "index", 0) - delta = getattr(choice, "delta", None) - if delta is None: - continue - terminal = bool(getattr(choice, "finish_reason", None)) - - tool_calls = getattr(delta, "tool_calls", None) - if tool_calls: - self._accumulate_tool_calls(tool_acc, index, tool_calls) - delta.tool_calls = None - function_call = getattr(delta, "function_call", None) - if function_call is not None: - self._accumulate_function_call(func_acc, index, function_call) - delta.function_call = None - - raw_content = getattr(delta, "content", None) - content = raw_content if isinstance(raw_content, str) else None - if content is not None or terminal: - emitted, hold = await emit_content(content_buffers.pop(index, "") + (content or ""), terminal) - if hold: - content_buffers[index] = hold - if emitted: - delta.content = emitted - else: - delta.content = None if content is None else "" - - if terminal: - built_tool_calls = await self._build_tool_calls(tool_acc.get(index, {}), transform) - built_function_call = await self._build_function_call(func_acc.get(index), transform) - if built_tool_calls: - delta.tool_calls = built_tool_calls - if built_function_call is not None: - delta.function_call = built_function_call - tool_acc.pop(index, None) - func_acc.pop(index, None) - - @staticmethod - async def _build_tail_chunk( - template: Optional[ModelResponseStream], - content_buffers: dict[int, str], - tool_acc: dict[int, dict[int, dict[str, Optional[str]]]], - func_acc: dict[int, dict[str, Optional[str]]], - transform: "Callable[[str], Awaitable[str]]", - ) -> Optional[ModelResponseStream]: - """Flush any content / tool-call state still held when a stream ends - without a finish-reason chunk to attach it to.""" - if template is None: - return None - cls = _OPTIONAL_PresidioPIIMasking - choices: list[StreamingChoices] = [] - for index in sorted(set(content_buffers) | set(tool_acc) | set(func_acc)): - held = content_buffers.get(index, "") - masked_content = await transform(held) if held else None - built_tool_calls = await cls._build_tool_calls(tool_acc.get(index, {}), transform) - built_function_call = await cls._build_function_call(func_acc.get(index), transform) - if masked_content is None and not built_tool_calls and built_function_call is None: - continue - choices.append( - StreamingChoices( - index=index, - delta=Delta( - content=masked_content, - tool_calls=built_tool_calls or None, - function_call=built_function_call, - ), - ) - ) - if not choices: - return None - return ModelResponseStream( - id=getattr(template, "id", None), - created=getattr(template, "created", None), - model=getattr(template, "model", None), - object="chat.completion.chunk", - choices=choices, - ) - - @staticmethod - def _redacted_chunk(chunk: ModelResponseStream) -> ModelResponseStream: - """Fail closed when masking a chunk raises: rebuild it with empty content - but its original ``finish_reason`` and choice indices preserved, so - possibly-unmasked PII never reaches the client yet a terminal chunk still - carries the completion signal instead of being dropped.""" - return ModelResponseStream( - id=chunk.id, - created=chunk.created, - model=chunk.model, - object="chat.completion.chunk", - choices=[ - StreamingChoices( - index=choice.index, - delta=Delta(content=None), - finish_reason=choice.finish_reason, - ) - for choice in chunk.choices - ], - ) - async def _stream_apply_output_masking( self, response: Any, request_data: dict, ) -> AsyncGenerator[Union[ModelResponseStream, bytes], None]: """Apply Presidio masking to streaming output (apply_to_output=True path).""" - presidio_config = self.get_presidio_settings_from_request_data(request_data or {}) + from litellm.llms.base_llm.base_model_iterator import ( + convert_model_response_to_streaming, + ) + from litellm.main import stream_chunk_builder + from litellm.types.utils import ModelResponse - async def transform(text: str) -> str: - return await self.check_pii( - text=text, - output_parse_pii=False, - presidio_config=presidio_config, - request_data=request_data, - ) - - async def emit_content(text: str, terminal: bool) -> tuple[str, str]: - return await self._mask_emit_decision(text, terminal, transform) - - async def flush_held() -> Optional[ModelResponseStream]: - """Build the held-content tail, failing closed (drop held content) - on a masking error instead of letting it abort the whole stream.""" - try: - return await self._build_tail_chunk(last_chunk, content_buffers, tool_acc, func_acc, transform) - except Exception as e: - if self._is_guardrail_intervention(e): - raise - verbose_proxy_logger.error(f"Error masking streaming PII tail: {str(e)}") - return None - - content_buffers: dict[int, str] = {} - tool_acc: dict[int, dict[int, dict[str, Optional[str]]]] = {} - func_acc: dict[int, dict[str, Optional[str]]] = {} - last_chunk: Optional[ModelResponseStream] = None - masked_any_content = False - saw_unmaskable_shape = False + all_chunks: List[ModelResponseStream] = [] + passthrough_due_to_unknown_stream_shape = False try: async for chunk in response: - if not isinstance(chunk, ModelResponseStream): - # Flush buffered masked content before forwarding a non-chat - # shape (raw bytes / a /v1/responses event) so the client - # never sees a later event ahead of earlier masked text. - tail = await flush_held() - if tail is not None: - yield tail - content_buffers.clear() - tool_acc.clear() - func_acc.clear() - saw_unmaskable_shape = True + if isinstance(chunk, ModelResponseStream): + if passthrough_due_to_unknown_stream_shape: + yield chunk + else: + all_chunks.append(chunk) + elif isinstance(chunk, bytes): + yield chunk # type: ignore[misc] + continue + else: + if all_chunks: + # Flush buffered chunks and switch to transparent passthrough for this stream shape. + # NOTE: these buffered chunks are emitted unmasked because this + # stream mixed chunk types and cannot be safely reconstructed. + verbose_proxy_logger.warning( + "Presidio apply_to_output: mixed stream detected (ModelResponseStream + unknown event). " + "Flushing %d buffered chunks without PII masking and switching to transparent passthrough.", + len(all_chunks), + ) + for buffered_chunk in all_chunks: + yield buffered_chunk + all_chunks = [] + passthrough_due_to_unknown_stream_shape = True yield chunk - continue - masked_any_content = True - last_chunk = chunk - try: - await self._rewrite_chat_chunk( - chunk, - content_buffers, - tool_acc, - func_acc, - transform, - emit_content, - ) - except Exception as e: - if self._is_guardrail_intervention(e): - raise - # Fail closed: a transient masking error redacts this chunk's - # content (so possibly-unmasked PII never reaches the client) - # but keeps its finish_reason and keeps the stream flowing, - # rather than truncating the response or dropping a terminal - # chunk's completion signal. - verbose_proxy_logger.error(f"Error masking streaming PII chunk: {str(e)}") - yield self._redacted_chunk(chunk) - continue - yield chunk - - tail = await flush_held() - if tail is not None: - yield tail - if not masked_any_content and saw_unmaskable_shape: + if passthrough_due_to_unknown_stream_shape: + verbose_proxy_logger.warning( + "Presidio apply_to_output: streaming response contained unknown event objects " + "(e.g. /v1/responses events). Output PII masking was skipped for this response." + ) + return + if not all_chunks: verbose_proxy_logger.warning( "Presidio apply_to_output: streaming response contained no " - "maskable chat content (e.g. raw SSE bytes or /v1/responses " - "events). Output PII masking was skipped for this response." + "ModelResponseStream chunks (e.g. raw SSE bytes or an empty " + "upstream stream). Output PII masking was skipped for this " + "response." ) + return + + assembled_model_response = stream_chunk_builder(chunks=all_chunks, messages=request_data.get("messages")) + + if not isinstance(assembled_model_response, ModelResponse): + for chunk in all_chunks: + yield chunk + return + + await self._process_response_for_pii( + response=assembled_model_response, + request_data=request_data, + mode="mask", + ) + + mock_response_stream = convert_model_response_to_streaming(assembled_model_response) + yield mock_response_stream + except Exception as e: - if self._is_guardrail_intervention(e): - raise verbose_proxy_logger.error(f"Error masking streaming PII output: {str(e)}") + for chunk in all_chunks: + yield chunk @staticmethod def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: Dict[str, str]) -> bytes: @@ -1479,56 +1183,74 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): request_data: dict, ) -> AsyncGenerator[Union[ModelResponseStream, bytes], None]: """Apply PII unmasking to streaming output (output_parse_pii=True path).""" + from litellm.llms.base_llm.base_model_iterator import ( + convert_model_response_to_streaming, + ) + 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", {}) - async def transform(text: str) -> str: - return self._unmask_pii_text(text, pii_tokens) - - async def emit_content(text: str, terminal: bool) -> tuple[str, str]: - if terminal: - return (await transform(text) if text else ""), "" - hold = self._unmask_hold_len(text, pii_tokens.keys()) - emit_raw, held = text[: len(text) - hold], text[len(text) - hold :] - return (await transform(emit_raw) if emit_raw else ""), held - - content_buffers: dict[int, str] = {} - tool_acc: dict[int, dict[int, dict[str, Optional[str]]]] = {} - func_acc: dict[int, dict[str, Optional[str]]] = {} - last_chunk: Optional[ModelResponseStream] = None + remaining_chunks: List[ModelResponseStream] = [] + saw_non_chat_chunk = False try: async for chunk in response: - if isinstance(chunk, bytes): - tail = await self._build_tail_chunk(last_chunk, content_buffers, tool_acc, func_acc, transform) - if tail is not None: - yield tail - content_buffers.clear() - tool_acc.clear() - func_acc.clear() - yield ( # type: ignore[misc] - self._unmask_sse_bytes_chunk(chunk, pii_tokens) if pii_tokens else chunk - ) + if isinstance(chunk, ModelResponseStream): + if saw_non_chat_chunk: + yield chunk + else: + remaining_chunks.append(chunk) + elif isinstance(chunk, bytes): + if pii_tokens: + yield self._unmask_sse_bytes_chunk(chunk, pii_tokens) # type: ignore[misc] + else: + yield chunk # type: ignore[misc] continue - if not isinstance(chunk, ModelResponseStream): - tail = await self._build_tail_chunk(last_chunk, content_buffers, tool_acc, func_acc, transform) - if tail is not None: - yield tail - content_buffers.clear() - tool_acc.clear() - func_acc.clear() - if getattr(chunk, "type", None) == "response.completed" and pii_tokens: + else: + # /v1/responses events: unmask response.completed text in-place. + # A mixed stream can't be reassembled, so flush buffered chat + # chunks in order before passthrough instead of dropping them. + if remaining_chunks and not saw_non_chat_chunk: + for buffered_chunk in remaining_chunks: + yield buffered_chunk + remaining_chunks = [] + chunk_type = getattr(chunk, "type", None) + if chunk_type == "response.completed" and pii_tokens: self._unmask_responses_api_completed_chunk(chunk, pii_tokens) + saw_non_chat_chunk = True yield chunk - continue - last_chunk = chunk - await self._rewrite_chat_chunk(chunk, content_buffers, tool_acc, func_acc, transform, emit_content) - yield chunk - tail = await self._build_tail_chunk(last_chunk, content_buffers, tool_acc, func_acc, transform) - if tail is not None: - yield tail + if saw_non_chat_chunk: + return + + if not remaining_chunks: + return + + assembled_model_response = stream_chunk_builder( + chunks=remaining_chunks, messages=request_data.get("messages") + ) + + if not isinstance(assembled_model_response, ModelResponse): + for chunk in remaining_chunks: + yield chunk + return + + self._preserve_usage_from_last_chunk(assembled_model_response, remaining_chunks) + + await self._process_response_for_pii( + response=assembled_model_response, + request_data=request_data, + mode="unmask", + ) + + mock_response_stream = convert_model_response_to_streaming(assembled_model_response) + yield mock_response_stream + except Exception as e: verbose_proxy_logger.error(f"Error in PII streaming processing: {str(e)}") + for chunk in remaining_chunks: + yield chunk async def async_post_call_streaming_iterator_hook( # type: ignore[override] self, @@ -1560,6 +1282,17 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): async for chunk in self._stream_pii_unmasking(response, request_data): yield chunk + @staticmethod + def _preserve_usage_from_last_chunk( + assembled_model_response: Any, + chunks: List[Any], + ) -> None: + """Copy usage metadata from the last chunk when stream_chunk_builder misses it.""" + if not getattr(assembled_model_response, "usage", None) and chunks: + last_chunk_usage = getattr(chunks[-1], "usage", None) + if last_chunk_usage: + setattr(assembled_model_response, "usage", last_chunk_usage) + def get_presidio_settings_from_request_data(self, data: dict) -> Optional[PresidioPerRequestConfig]: if "metadata" in data: _metadata = data.get("metadata", None) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index ce4d6d4a8b0..253d989f203 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -19,7 +19,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.presidio import ( _OPTIONAL_PresidioPIIMasking, ) -from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException +from litellm.exceptions import GuardrailRaisedException from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType from litellm.types.utils import Choices, Message, ModelResponse @@ -2207,12 +2207,11 @@ async def test_apply_to_output_streaming_unknown_events_passthrough(): @pytest.mark.asyncio -async def test_apply_to_output_streaming_mixed_chunks_preserve_order(): +async def test_apply_to_output_streaming_mixed_chunks_flushes_and_warns(): """ - Regression test for mixed stream shape: a ModelResponseStream chat chunk - followed by an unknown responses-style event must be forwarded in order. - Incremental masking forwards chat chunks as they arrive, so a responses - event after them does not buffer or drop anything. + Regression test for mixed stream shape: + a buffered ModelResponseStream chunk followed by unknown responses-style + events should be preserved, and masking skip should be visible via warnings. """ guardrail = _OPTIONAL_PresidioPIIMasking( mock_testing=True, @@ -2239,14 +2238,26 @@ async def test_apply_to_output_streaming_mixed_chunks_preserve_order(): mock_user_api_key = UserAPIKeyAuth(api_key="test-key") received = [] - async for chunk in guardrail.async_post_call_streaming_iterator_hook( - user_api_key_dict=mock_user_api_key, - response=mock_stream(), - request_data={}, - ): - received.append(chunk) + with patch( + "litellm.proxy.guardrails.guardrail_hooks.presidio.verbose_proxy_logger" + ) as mock_logger: + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=mock_user_api_key, + response=mock_stream(), + request_data={}, + ): + received.append(chunk) - assert received == [model_chunk, response_completed] + # Preserve original ordering across mixed stream types. + assert received == [model_chunk, response_completed] + + # Two warnings are expected: + # 1) mixed stream detected + unmasked flush + # 2) passthrough mode skipped output masking + assert mock_logger.warning.call_count == 2 + warning_messages = [call.args[0] for call in mock_logger.warning.call_args_list] + assert any("mixed stream detected" in msg for msg in warning_messages) + assert any("unknown event objects" in msg for msg in warning_messages) # --------------------------------------------------------------------------- @@ -2838,678 +2849,3 @@ async def test_stream_pii_unmasking_passthrough_when_no_tokens(mock_user_api_key chunks.append(chunk) assert chunks == [raw_chunk] - - -# --------------------------------------------------------------------------- -# LIT-3222: incremental SSE streaming for Presidio output masking / unmasking -# --------------------------------------------------------------------------- - -from litellm.types.utils import ( - ChatCompletionDeltaToolCall, - Delta, - Function, - StreamingChoices, -) - - -def _content_chunk(text, index=0, finish_reason=None): - return ModelResponseStream( - id="chatcmpl-lit3222", - created=1, - model="gpt-4o-mini", - object="chat.completion.chunk", - choices=[ - StreamingChoices( - index=index, delta=Delta(content=text), finish_reason=finish_reason - ) - ], - ) - - -def _non_empty_content(chunks): - out = [] - for chunk in chunks: - for choice in chunk.choices: - piece = getattr(choice.delta, "content", None) - if piece: - out.append(piece) - return out - - -async def _drive(guardrail, stream, request_data): - collected = [] - async for chunk in guardrail.async_post_call_streaming_iterator_hook( - user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), - response=stream, - request_data=request_data, - ): - collected.append(chunk) - return collected - - -@pytest.mark.asyncio -async def test_unmask_streaming_is_incremental_not_buffered(): - """ - output_parse_pii streaming must forward each content chunk as it arrives - (unmasked), not collapse the whole completion into a single end-of-stream - chunk. The buffering implementation yielded exactly one content chunk. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) - request_data = {"metadata": {"pii_tokens": {"": "John Smith"}}} - - pieces = ["Hello ", "", " is here."] - - async def stream(): - for i, piece in enumerate(pieces): - yield _content_chunk(piece) - yield _content_chunk("", finish_reason="stop") - - collected = await _drive(guardrail, stream(), request_data) - content = _non_empty_content(collected) - - assert len(content) >= 3, f"expected progressive chunks, got {content}" - assert "".join(content) == "Hello John Smith is here." - assert all("" not in piece for piece in content) - - -@pytest.mark.asyncio -async def test_unmask_streaming_token_split_across_chunks(): - """ - A placeholder token split across SSE chunks (````) must - still be unmasked atomically via the cross-chunk carry buffer. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) - request_data = {"metadata": {"pii_tokens": {"": "John Smith"}}} - - pieces = ["Hi ", "", "!"] - - async def stream(): - for piece in pieces: - yield _content_chunk(piece) - yield _content_chunk("", finish_reason="stop") - - collected = await _drive(guardrail, stream(), request_data) - reassembled = "".join(_non_empty_content(collected)) - - assert reassembled == "Hi John Smith!" - assert "" not in reassembled - - -@pytest.mark.asyncio -async def test_unmask_streaming_independent_per_choice_buffers(): - """ - With n>1 each choice keeps its own carry buffer, so a token split across - chunks on choice 1 does not corrupt choice 0 and vice versa. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) - request_data = { - "metadata": { - "pii_tokens": {"": "John", "": "Jane"} - } - } - - def two_choice_chunk(c0, c1): - return ModelResponseStream( - id="chatcmpl-lit3222", - created=1, - model="gpt-4o-mini", - object="chat.completion.chunk", - choices=[ - StreamingChoices(index=0, delta=Delta(content=c0)), - StreamingChoices(index=1, delta=Delta(content=c1)), - ], - ) - - async def stream(): - yield two_choice_chunk(" ok", "SON_2>!") - yield two_choice_chunk("", "") - - collected = [] - async for chunk in guardrail.async_post_call_streaming_iterator_hook( - user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), - response=stream(), - request_data=request_data, - ): - collected.append(chunk) - - per_choice = {0: "", 1: ""} - for chunk in collected: - for choice in chunk.choices: - if choice.delta.content: - per_choice[choice.index] += choice.delta.content - - assert per_choice[0] == "John ok" - assert per_choice[1] == "Jane!" - - -@pytest.mark.asyncio -async def test_unmask_streaming_tool_call_arguments_unmasked_at_finish(): - """ - Tool-call argument fragments carrying a placeholder token must be - reassembled and unmasked (the tool would otherwise receive ````). - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) - request_data = { - "metadata": {"pii_tokens": {"": "real@example.com"}} - } - - def tool_chunk(*, id=None, name=None, args, finish_reason=None): - return ModelResponseStream( - id="chatcmpl-lit3222", - created=1, - model="gpt-4o-mini", - object="chat.completion.chunk", - choices=[ - StreamingChoices( - index=0, - delta=Delta( - tool_calls=[ - ChatCompletionDeltaToolCall( - index=0, - id=id, - type="function" if id else None, - function=Function(name=name, arguments=args), - ) - ] - ), - finish_reason=finish_reason, - ) - ], - ) - - async def stream(): - yield tool_chunk(id="call_1", name="send_email", args="") - yield tool_chunk(args='{"to": "") - - guardrail.check_pii = mock_check_pii - - pieces = ["My email is ", "secret@example.com. ", "Call me later."] - - async def stream(): - for piece in pieces: - yield _content_chunk(piece) - yield _content_chunk("", finish_reason="stop") - - collected = await _drive(guardrail, stream(), {"metadata": {}}) - content = _non_empty_content(collected) - - assert len(content) >= 2, f"expected per-sentence chunks, got {content}" - reassembled = "".join(content) - assert reassembled == "My email is . Call me later." - assert "secret@example.com" not in reassembled - - -@pytest.mark.asyncio -async def test_mask_streaming_tool_call_arguments_masked_at_finish(): - """ - Model-generated PII inside streamed tool-call arguments must be masked - before reaching the client, not passed through unmasked. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - - async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): - return text.replace("secret@example.com", "") - - guardrail.check_pii = mock_check_pii - - def tool_chunk(*, id=None, name=None, args, finish_reason=None): - return ModelResponseStream( - id="chatcmpl-lit3222", - created=1, - model="gpt-4o-mini", - object="chat.completion.chunk", - choices=[ - StreamingChoices( - index=0, - delta=Delta( - tool_calls=[ - ChatCompletionDeltaToolCall( - index=0, - id=id, - type="function" if id else None, - function=Function(name=name, arguments=args), - ) - ] - ), - finish_reason=finish_reason, - ) - ], - ) - - async def stream(): - yield tool_chunk(id="call_1", name="save", args='{"email": "sec') - yield tool_chunk(args='ret@example.com"}') - yield ModelResponseStream( - id="chatcmpl-lit3222", - created=1, - model="gpt-4o-mini", - object="chat.completion.chunk", - choices=[ - StreamingChoices(index=0, delta=Delta(), finish_reason="tool_calls") - ], - ) - - collected = await _drive(guardrail, stream(), {"metadata": {}}) - - all_args = [ - tc.function.arguments - for chunk in collected - for choice in chunk.choices - for tc in (getattr(choice.delta, "tool_calls", None) or []) - ] - assert all_args == ['{"email": ""}'] - assert all("secret@example.com" not in args for args in all_args) - - -@pytest.mark.asyncio -async def test_mask_streaming_flushes_buffered_content_before_passthrough_event(): - """ - Regression test (Greptile 4/5 finding): in the apply_to_output path, masked - content held in the buffer (no sentence boundary yet) must be flushed BEFORE - a non-chat passthrough event (e.g. a /v1/responses completion) is forwarded, - so the client never observes stream completion ahead of the final text. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - - async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): - return text.replace("secret@example.com", "") - - guardrail.check_pii = mock_check_pii - - class FakeResponsesEvent: - def __init__(self, event_type: str): - self.type = event_type - - completed = FakeResponsesEvent("response.completed") - - async def stream(): - # No sentence terminator -> held in the mask buffer, not yet emitted. - yield _content_chunk("My email is secret@example.com") - yield completed - - collected = await _drive(guardrail, stream(), {"metadata": {}}) - - event_index = collected.index(completed) - masked_indexes = [ - i - for i, chunk in enumerate(collected) - if not isinstance(chunk, FakeResponsesEvent) - and any(getattr(c.delta, "content", None) for c in chunk.choices) - ] - assert masked_indexes, "buffered masked content was never emitted" - assert max(masked_indexes) < event_index, "masked content must precede the event" - - masked_text = "".join( - c.delta.content - for chunk in collected - if not isinstance(chunk, FakeResponsesEvent) - for c in chunk.choices - if getattr(c.delta, "content", None) - ) - assert masked_text == "My email is " - assert "secret@example.com" not in masked_text - - -@pytest.mark.asyncio -async def test_mask_streaming_does_not_split_entity_on_long_unpunctuated_run(): - """ - A long punctuation-free run must never force-flush mid-entity. The old - fixed-window fallback emitted at the last whitespace once the buffer grew - past a cap, so a space-bearing entity (SSN, phone) straddling that cut was - analyzed in two halves and leaked unmasked. Content past the last sentence - boundary is now held until a boundary or end-of-stream so each analyze call - sees the whole entity. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - - async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): - return text.replace("123 45 6789", "") - - guardrail.check_pii = mock_check_pii - - filler = "data " * 90 # >400 chars, spaces only, no .!?\n boundary - async def stream(): - yield _content_chunk(filler + "my ssn is 123 45 6789") - yield _content_chunk("", finish_reason="stop") - - collected = await _drive(guardrail, stream(), {"metadata": {}}) - reassembled = "".join(_non_empty_content(collected)) - - assert "" in reassembled - assert "123 45 6789" not in reassembled - assert "456789" not in reassembled - - -@pytest.mark.asyncio -async def test_mask_streaming_preserves_stream_on_check_pii_error(): - """ - A transient Presidio failure mid-stream must not truncate the response. The - failing run is dropped (fail closed, never leaking the PII it could not mask) - while content that already flushed safely, later chunks, and the finish chunk - still reach the client. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - guardrail._stream_mask_margin = 4 - - async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): - if "boom@example.com" in text: - raise RuntimeError("presidio down") - return text.replace("ok@example.com", "") - - guardrail.check_pii = mock_check_pii - - pieces = [ - "First ok@example.com. ", - "filler text here. ", - "Second boom@example.com. ", - "Third part here.", - ] - - async def stream(): - for piece in pieces: - yield _content_chunk(piece) - yield _content_chunk("", finish_reason="stop") - - collected = await _drive(guardrail, stream(), {"metadata": {}}) - reassembled = "".join(_non_empty_content(collected)) - - assert "" in reassembled - assert "boom@example.com" not in reassembled - assert "Third part" in reassembled, "stream truncated after a masking error" - assert any( - getattr(choice, "finish_reason", None) - for chunk in collected - for choice in getattr(chunk, "choices", []) - ), "finish chunk dropped after a masking error" - - -@pytest.mark.asyncio -async def test_mask_streaming_error_preserves_tool_call_accumulators(): - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - - async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): - if "bad@example.com" in text: - raise RuntimeError("presidio down") - return text.replace("secret@example.com", "") - - guardrail.check_pii = mock_check_pii - - def tool_chunk(*, id=None, name=None, args=None, finish_reason=None): - return ModelResponseStream( - id="chatcmpl-lit3222", - created=1, - model="gpt-4o-mini", - object="chat.completion.chunk", - choices=[ - StreamingChoices( - index=1, - delta=( - Delta( - tool_calls=[ - ChatCompletionDeltaToolCall( - index=0, - id=id, - type="function" if id else None, - function=Function(name=name, arguments=args), - ) - ] - ) - if args is not None - else Delta() - ), - finish_reason=finish_reason, - ) - ], - ) - - async def stream(): - yield tool_chunk( - id="call_1", - name="save", - args='{"email": "secret@example.com"}', - ) - yield _content_chunk("bad@example.com", index=0, finish_reason="stop") - yield tool_chunk(finish_reason="tool_calls") - - collected = await _drive(guardrail, stream(), {"metadata": {}}) - tool_args = [ - tc.function.arguments - for chunk in collected - for choice in chunk.choices - for tc in (getattr(choice.delta, "tool_calls", None) or []) - ] - - assert tool_args == ['{"email": ""}'] - - -@pytest.mark.asyncio -async def test_mask_emit_decision_caps_buffer_when_stability_fails(): - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - guardrail._stream_mask_margin = 3 - guardrail._stream_mask_max_buffer = 6 - - async def unstable_transform(text): - return text[::-1] - - emitted, held = await guardrail._mask_emit_decision( - "abcdefghij", False, unstable_transform - ) - - assert emitted == "" - assert held == "hij" - - -@pytest.mark.parametrize( - "exception", - [ - BlockedPiiEntityError(entity_type="EMAIL_ADDRESS", guardrail_name="presidio"), - GuardrailRaisedException(guardrail_name="presidio", message="invalid response"), - ], -) -@pytest.mark.asyncio -async def test_mask_streaming_propagates_guardrail_interventions(exception): - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - - async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): - raise exception - - guardrail.check_pii = mock_check_pii - - async def stream(): - yield _content_chunk("blocked", finish_reason="stop") - - with pytest.raises(type(exception)): - await _drive(guardrail, stream(), {"metadata": {}}) - - -@pytest.mark.asyncio -async def test_mask_streaming_holds_terminator_at_chunk_end_until_whitespace(): - """ - A sentence terminator at the very end of a chunk is not a safe boundary: the - next chunk may continue the token. "Contact jane." followed by - "doe@example.com" must mask the whole email rather than flushing "jane." and - analyzing the two halves separately, which would leak the address. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - - async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): - return text.replace("jane.doe@example.com", "") - - guardrail.check_pii = mock_check_pii - - async def stream(): - yield _content_chunk("Contact jane.") - yield _content_chunk("doe@example.com") - yield _content_chunk("", finish_reason="stop") - - collected = await _drive(guardrail, stream(), {"metadata": {}}) - reassembled = "".join(_non_empty_content(collected)) - - assert "" in reassembled - assert "jane.doe@example.com" not in reassembled - assert "jane." not in reassembled - - -@pytest.mark.asyncio -async def test_mask_streaming_does_not_split_entity_across_sentence_boundary(): - """ - An entity that straddles a sentence boundary (a name with a middle initial, - an address across a newline) must not be flushed in halves. The stability - check holds the prefix until masking it alone matches masking the whole - buffer, so the straddling entity is analyzed and masked as one unit instead - of leaking the part before the boundary. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - guardrail._stream_mask_margin = 8 - - async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): - return text.replace("John Q. Public", "") - - guardrail.check_pii = mock_check_pii - - async def stream(): - yield _content_chunk("Please greet John Q. Public warmly when they arrive.") - yield _content_chunk("", finish_reason="stop") - - collected = await _drive(guardrail, stream(), {"metadata": {}}) - reassembled = "".join(_non_empty_content(collected)) - - assert "" in reassembled - assert "John Q. Public" not in reassembled - assert "John Q." not in reassembled - - -@pytest.mark.asyncio -async def test_mask_streaming_caps_runaway_buffer_without_splitting_entity(): - """ - A punctuation-free run past the buffer cap must be flushed to bound memory, - but the forced flush still cuts a margin back from the end so a PII value - split exactly at the cap (``secret@`` in one chunk, ``example.com`` in the - next) stays buffered and is masked whole instead of leaking its raw halves. - The early flush plus the terminal flush yields more than one content chunk. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - guardrail._stream_mask_margin = 8 - guardrail._stream_mask_max_buffer = 20 - - async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): - return text.replace("secret@example.com", "") - - guardrail.check_pii = mock_check_pii - - async def stream(): - yield _content_chunk("please email me at secret@") # 26 > cap, email cut - yield _content_chunk("example.com now") - yield _content_chunk("", finish_reason="stop") - - collected = await _drive(guardrail, stream(), {"metadata": {}}) - content = _non_empty_content(collected) - reassembled = "".join(content) - - assert "" in reassembled - assert "secret@example.com" not in reassembled - assert "secret@" not in reassembled - assert len(content) >= 2, f"cap did not flush before end of stream, got {content}" - - -@pytest.mark.asyncio -async def test_mask_streaming_preserves_finish_reason_when_terminal_chunk_fails(): - """ - When the masking call fails on the terminal chunk itself, that chunk must be - redacted in place (content dropped, fail closed) but keep its finish_reason, - so the client still receives the completion signal instead of a stream that - ends without one. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, apply_to_output=True) - guardrail._stream_mask_margin = 4 - - async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): - if "boom@example.com" in text: - raise RuntimeError("presidio down") - return text - - guardrail.check_pii = mock_check_pii - - async def stream(): - yield _content_chunk("Hello world. ") - yield _content_chunk("more text here. ") - yield _content_chunk("boom@example.com", finish_reason="stop") - - collected = await _drive(guardrail, stream(), {"metadata": {}}) - reassembled = "".join(_non_empty_content(collected)) - - assert "boom@example.com" not in reassembled - assert "Hello world" in reassembled - finish_reasons = [ - choice.finish_reason - for chunk in collected - for choice in getattr(chunk, "choices", []) - if getattr(choice, "finish_reason", None) - ] - assert "stop" in finish_reasons, "finish_reason dropped when terminal chunk failed" - - -@pytest.mark.asyncio -async def test_unmask_streaming_flushes_held_content_before_bytes(): - """ - When a held placeholder prefix is buffered and the next upstream item is a - raw SSE byte chunk, the held chat text must be flushed before the bytes so - the client never sees the byte chunk ahead of earlier content. - """ - guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True) - request_data = {"metadata": {"pii_tokens": {"": "Jane"}}} - - async def stream(): - yield _content_chunk("Hi