fix(presidio): stream SSE output incrementally instead of buffering the whole response (#31503)

The Presidio streaming post-call hooks (_stream_apply_output_masking for
apply_to_output and _stream_pii_unmasking for output_parse_pii) collected every
upstream chunk, reassembled the full completion with stream_chunk_builder at
end-of-stream, ran Presidio over it, then emitted one reconstructed SSE chunk.
Time-to-first-token collapsed to the total generation time and token-by-token
streaming was lost whenever Presidio output handling was enabled. With the
default presidio_filter_scope both, an apply_to_output masking instance is always
created, so even the unmask configuration buffered the stream.

Both paths now transform and forward chunks as they arrive. The unmask path
replaces placeholder tokens per chunk, holding back only the trailing run that
could still grow into a token so a placeholder split across SSE chunks
(<PER + SON_1>) is still rewritten atomically. The mask path emits a prefix only
when masking it in isolation matches the corresponding prefix of masking the
whole buffer, with a lookahead margin still buffered past the cut, so an entity
straddling the cut is detected and held until complete; past
_PRESIDIO_STREAM_MAX_BUFFER the run is bounded without splitting an entity.
Tool-call and legacy function-call argument fragments are accumulated per choice
and transformed once the choice closes, content is buffered independently per
choice index for correct n>1 streaming, raw Anthropic SSE bytes and /v1/responses
events pass through with any held content flushed first so events never reorder,
and a masking error redacts only the affected chunk (fail closed, keeping
finish_reason) while the stream continues.

Resolves LIT-3222
This commit is contained in:
Yassin Kortam 2026-06-30 22:59:18 +03:00 • committed by GitHub
parent d4c33b2b59
commit 94936a3922
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 1086 additions and 155 deletions

View file

@ -17,6 +17,8 @@ from typing import (
TYPE_CHECKING,
Any,
AsyncGenerator,
Awaitable,
Callable,
Dict,
List,
Literal,
@ -54,7 +56,14 @@ from litellm.types.proxy.guardrails.guardrail_hooks.presidio import (
PresidioAnalyzeRequest,
PresidioAnalyzeResponseItem,
)
from litellm.types.utils import GuardrailStatus, StreamingChoices
from litellm.types.utils import (
ChatCompletionDeltaToolCall,
Delta,
Function,
FunctionCall,
GuardrailStatus,
StreamingChoices,
)
from litellm.utils import (
EmbeddingResponse,
ImageResponse,
@ -62,6 +71,17 @@ 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
@ -93,6 +113,10 @@ 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
@ -1048,81 +1072,353 @@ 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
``<PERSON_1>``."""
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)."""
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
presidio_config = self.get_presidio_settings_from_request_data(request_data or {})
all_chunks: List[ModelResponseStream] = []
passthrough_due_to_unknown_stream_shape = False
try:
async for chunk in response:
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
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 "
"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,
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,
mode="mask",
)
mock_response_stream = convert_model_response_to_streaming(assembled_model_response)
yield mock_response_stream
async def emit_content(text: str, terminal: bool) -> tuple[str, str]:
return await self._mask_emit_decision(text, terminal, transform)
except Exception as e:
verbose_proxy_logger.error(f"Error masking streaming PII output: {str(e)}")
for chunk in all_chunks:
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
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
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:
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."
)
except Exception as e:
if self._is_guardrail_intervention(e):
raise
verbose_proxy_logger.error(f"Error masking streaming PII output: {str(e)}")
@staticmethod
def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: Dict[str, str]) -> bytes:
try:
@ -1183,74 +1479,56 @@ 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", {})
remaining_chunks: List[ModelResponseStream] = []
saw_non_chat_chunk = False
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
try:
async for chunk in response:
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]
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
)
continue
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:
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:
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
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
tail = await self._build_tail_chunk(last_chunk, content_buffers, tool_acc, func_acc, transform)
if tail is not None:
yield tail
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,
@ -1282,17 +1560,6 @@ 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)

View file

@ -19,7 +19,7 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.presidio import (
_OPTIONAL_PresidioPIIMasking,
)
from litellm.exceptions import GuardrailRaisedException
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType
from litellm.types.utils import Choices, Message, ModelResponse
@ -2207,11 +2207,12 @@ async def test_apply_to_output_streaming_unknown_events_passthrough():
@pytest.mark.asyncio
async def test_apply_to_output_streaming_mixed_chunks_flushes_and_warns():
async def test_apply_to_output_streaming_mixed_chunks_preserve_order():
"""
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.
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.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
@ -2238,26 +2239,14 @@ async def test_apply_to_output_streaming_mixed_chunks_flushes_and_warns():
mock_user_api_key = UserAPIKeyAuth(api_key="test-key")
received = []
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)
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)
# 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)
assert received == [model_chunk, response_completed]
# ---------------------------------------------------------------------------
@ -2849,3 +2838,678 @@ 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": {"<PERSON_1>": "John Smith"}}}
pieces = ["Hello ", "<PERSON_1>", " 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("<PERSON_1>" 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 (``<PER`` + ``SON_1>``) 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": {"<PERSON_1>": "John Smith"}}}
pieces = ["Hi ", "<PER", "SON_1>", "!"]
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 "<PER" not in reassembled and "SON_1>" 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": {"<PERSON_1>": "John", "<PERSON_2>": "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("<PERSON", "<PER")
yield two_choice_chunk("_1> 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 ``<EMAIL...>``).
"""
guardrail = _OPTIONAL_PresidioPIIMasking(mock_testing=True, output_parse_pii=True)
request_data = {
"metadata": {"pii_tokens": {"<EMAIL_ADDRESS_1>": "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": "<EMAIL_ADD')
yield tool_chunk(args='RESS_1>"}')
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(), request_data)
tool_calls = [
tc
for chunk in collected
for choice in chunk.choices
for tc in (getattr(choice.delta, "tool_calls", None) or [])
]
assert len(tool_calls) == 1
assert tool_calls[0].id == "call_1"
assert tool_calls[0].function.name == "send_email"
assert tool_calls[0].function.arguments == '{"to": "real@example.com"}'
@pytest.mark.asyncio
async def test_mask_streaming_is_incremental_per_sentence():
"""
apply_to_output streaming must mask and emit a completed sentence once enough
following context confirms it, instead of buffering the whole completion into
one chunk. A small safety margin keeps the test content short.
"""
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):
return text.replace("secret@example.com", "<EMAIL>")
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 <EMAIL>. 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", "<EMAIL>")
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": "<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", "<EMAIL>")
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 <EMAIL>"
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", "<US_SSN>")
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 "<US_SSN>" 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", "<EMAIL>")
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 "<EMAIL>" 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", "<EMAIL>")
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": "<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", "<EMAIL>")
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 "<EMAIL>" 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", "<PERSON>")
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 "<PERSON>" 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", "<EMAIL>")
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 "<EMAIL>" 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": {"<PERSON_1>": "Jane"}}}
async def stream():
yield _content_chunk("Hi <PERSON") # "<PERSON" held as a partial token
yield b"data: {\"type\": \"ping\"}\n\n"
collected = await _drive(guardrail, stream(), request_data)
bytes_index = next(
i for i, chunk in enumerate(collected) if isinstance(chunk, bytes)
)
held_index = next(
i
for i, chunk in enumerate(collected)
if not isinstance(chunk, bytes)
and any("Jane" in (getattr(c.delta, "content", None) or "") for c in chunk.choices)
)
assert held_index < bytes_index, "raw bytes were sent before held chat content"