fix(guardrails): buffer and mask raw Gemini SSE streams in Presidio output masking and fail closed on unknown stream shapes

With presidio_filter_scope output or both, a Google :streamGenerateContent
stream reached the Presidio streaming hook as raw SSE bytes that only the
Anthropic path knew how to read, so the Gemini frames were forwarded with
the PII the guardrail was configured to mask, while the same request
without streaming came back masked.

Raw Gemini streams are now assembled, masked, and re-emitted as one
provider-shaped frame the way Anthropic streams already were, and any raw
stream whose shape neither assembler can read is withheld with a 500
instead of being passed through unmasked. Whether masking changed the
response is decided on the whole assembled response, so a masked tool-call
argument is re-emitted too.
This commit is contained in:
mateo-berri 2026-09-26 13:50:41 -07:00
parent 5ac640e49d
commit a2efc1c321
5 changed files with 393 additions and 149 deletions

View file

@ -2,7 +2,8 @@
`/v1/messages` streams reach a guardrail's `async_post_call_streaming_iterator_hook` as raw SSE
frames rather than chunk objects, which `stream_chunk_builder` cannot assemble. These helpers let a
hook scan such a stream, and re-emit it when the guardrail rewrote the response.
hook scan such a stream, and re-emit it when the guardrail rewrote the response. The raw SSE
parsing (`joined_sse_stream`, `parsed_sse_events`) is shared with the Gemini sibling module.
"""
from __future__ import annotations
@ -32,19 +33,21 @@ def is_raw_sse_stream(all_chunks: Sequence[object]) -> bool:
return any(isinstance(chunk, (str, bytes)) for chunk in all_chunks)
def _joined_sse_stream(all_chunks: Sequence[object]) -> str | None:
def joined_sse_stream(all_chunks: Sequence[object]) -> str | None:
"""The raw frames decoded as one text, with every SSE line ending (CRLF, LF, CR) folded to LF."""
raw: Final = b"".join(
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8")
for chunk in all_chunks
if isinstance(chunk, (str, bytes))
)
try:
return codecs.getincrementaldecoder("utf-8")().decode(raw, final=False)
decoded: Final = codecs.getincrementaldecoder("utf-8")().decode(raw, final=False)
except UnicodeDecodeError:
return None
return decoded.replace("\r\n", "\n").replace("\r", "\n")
def _parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]:
def parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]:
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
)
@ -60,7 +63,7 @@ def _anthropic_message_start(sse_stream: str) -> Mapping[str, object] | None:
return next(
(
message
for event_data in _parsed_sse_events(sse_stream)
for event_data in parsed_sse_events(sse_stream)
if event_data.get("type") == "message_start" and isinstance(message := event_data.get("message"), dict)
),
None,
@ -75,10 +78,10 @@ def is_anthropic_sse_stream(all_chunks: Sequence[object]) -> bool:
stream raw too. Reading its frames as Anthropic ones would refuse the response in a wire format
its client cannot parse, so the surface is decided on the event types actually present.
"""
sse_stream: Final = _joined_sse_stream(all_chunks)
sse_stream: Final = joined_sse_stream(all_chunks)
if sse_stream is None:
return False
return any(event.get("type") in _ANTHROPIC_EVENT_TYPES for event in _parsed_sse_events(sse_stream))
return any(event.get("type") in _ANTHROPIC_EVENT_TYPES for event in parsed_sse_events(sse_stream))
def assemble_anthropic_sse_stream(
@ -95,7 +98,7 @@ def assemble_anthropic_sse_stream(
AnthropicPassthroughLoggingHandler,
)
sse_stream: Final = _joined_sse_stream(all_chunks)
sse_stream: Final = joined_sse_stream(all_chunks)
if sse_stream is None:
return None
message_start: Final = _anthropic_message_start(sse_stream)
@ -157,10 +160,10 @@ def is_sse_error_stream(all_chunks: Sequence[object]) -> bool:
# A stream mixing typed chunks with an error frame still carries content to scan, and the
# frames-only join below would drop exactly the part that has to be scanned
return False
sse_stream: Final = _joined_sse_stream(all_chunks)
sse_stream: Final = joined_sse_stream(all_chunks)
if sse_stream is None:
return False
events: Final = _parsed_sse_events(sse_stream)
events: Final = parsed_sse_events(sse_stream)
return len(events) > 0 and all(
event.get("type") == "error" or isinstance(event.get("error"), Mapping) for event in events
)

View file

@ -0,0 +1,65 @@
"""Gemini SSE <-> ModelResponse conversion for guardrail streaming hooks.
The Google ``:streamGenerateContent`` route relays the upstream ``data:`` frames as raw bytes, so a
guardrail's ``async_post_call_streaming_iterator_hook`` cannot read them as chunk objects. These
helpers assemble such a stream into a ModelResponse and re-emit it once the guardrail rewrote it.
"""
from __future__ import annotations
import json
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Final
from litellm.proxy.guardrails.anthropic_sse import joined_sse_stream, parsed_sse_events
from litellm.types.utils import ModelResponse
_GEMINI_RESPONSE_KEYS: Final = frozenset({"candidates", "usageMetadata", "promptFeedback"})
@dataclass(frozen=True, slots=True)
class _StreamParserLogging:
"""The one attribute the Gemini stream parser reads off its logging object."""
optional_params: Mapping[str, object]
def is_gemini_sse_stream(all_chunks: Sequence[object]) -> bool:
sse_stream: Final = joined_sse_stream(all_chunks)
if sse_stream is None:
return False
return any(not _GEMINI_RESPONSE_KEYS.isdisjoint(event) for event in parsed_sse_events(sse_stream))
def assemble_gemini_sse_stream(all_chunks: Sequence[object]) -> ModelResponse | None:
"""Assemble raw Gemini SSE frames into a ModelResponse.
A frame the Gemini parser rejects (a mid-stream ``error`` payload, say) raises the parser's own
error so the caller reports the upstream failure rather than a guardrail one.
"""
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ModelResponseIterator
from litellm.main import stream_chunk_builder
sse_stream: Final = joined_sse_stream(all_chunks)
if sse_stream is None:
return None
parser: Final = ModelResponseIterator(
streaming_response=None,
sync_stream=False,
logging_obj=_StreamParserLogging(optional_params={}), # pyright: ignore[reportArgumentType] # the parser reads only optional_params, and a native Gemini request carries no legacy `functions`
)
chunks: Final = tuple(
parsed for event in parsed_sse_events(sse_stream) if (parsed := parser.chunk_parser(dict(event))) is not None
)
if not chunks:
return None
assembled: Final = stream_chunk_builder(chunks=list(chunks))
return assembled if isinstance(assembled, ModelResponse) else None
def gemini_sse_chunks_from_response(assembled: ModelResponse) -> tuple[bytes, ...]:
from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter
body: Final = GoogleGenAIAdapter().translate_completion_to_generate_content(response=assembled)
return (f"data: {json.dumps(body)}\n\n".encode(),)

View file

@ -12,11 +12,11 @@ import asyncio
import json
import re
import threading
from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaitable, Iterator, Sequence
from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaitable, Callable, Iterator, Sequence
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeAlias, TypedDict, cast
import aiohttp
from typing_extensions import NotRequired, ReadOnly
@ -46,7 +46,12 @@ from litellm.proxy.guardrails.anthropic_sse import (
anthropic_sse_chunks_from_response,
assemble_anthropic_sse_stream,
is_anthropic_sse_stream,
model_response_text,
is_sse_error_stream,
)
from litellm.proxy.guardrails.gemini_sse import (
assemble_gemini_sse_stream,
gemini_sse_chunks_from_response,
is_gemini_sse_stream,
)
from litellm.types.guardrails import (
GuardrailEventHooks,
@ -97,9 +102,6 @@ def _json_escaped_len(text: str) -> int:
return len(json.dumps(text).encode("utf-8")) - 2 # strip the surrounding quotes
_MAX_FIRST_SSE_FRAME_BYTES: Final = 64 * 1024
@dataclass(frozen=True, slots=True)
class _SsePreface:
"""Complete leading SSE frames with no ``data:`` line, relayed verbatim before the stream shape is decided."""
@ -139,8 +141,7 @@ async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGener
``data:`` line) as they complete, and join raw ``bytes`` chunks until they
hold one complete SSE event with a data line, so the stream shape is
decided on a whole frame rather than a transport fragment. Everything
after that first frame is forwarded untouched. The byte cap can only be
reached by a single unterminated frame.
after that first frame is forwarded untouched.
"""
pending = b""
try:
@ -154,7 +155,7 @@ async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGener
if preface:
yield _SsePreface(preface)
pending = classifiable + tail
if classifiable or len(pending) >= _MAX_FIRST_SSE_FRAME_BYTES:
if classifiable:
break
else:
if pending:
@ -169,6 +170,18 @@ async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGener
yield chunk
_RawSseReEmit: TypeAlias = Callable[[ModelResponse], tuple[bytes, ...]]
def _assemble_raw_sse_stream(chunks: Sequence[object]) -> tuple[ModelResponse | None, _RawSseReEmit] | None:
"""The stream assembled by its surface, with that surface's re-emitter; None for an unknown surface."""
if is_anthropic_sse_stream(chunks):
return assemble_anthropic_sse_stream(chunks, restore_identity=True), anthropic_sse_chunks_from_response
if is_gemini_sse_stream(chunks):
return assemble_gemini_sse_stream(chunks), gemini_sse_chunks_from_response
return None
class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
user_api_key_cache = None
ad_hoc_recognizers: list[str] | None = None
@ -1442,18 +1455,13 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
elif isinstance(chunk, _SsePreface):
yield chunk.raw
elif isinstance(chunk, bytes):
first_frame_is_anthropic = (
not passthrough_due_to_unknown_stream_shape
and not all_chunks
and is_anthropic_sse_stream((chunk,))
)
if not first_frame_is_anthropic:
if all_chunks or passthrough_due_to_unknown_stream_shape or is_sse_error_stream((chunk,)):
passthrough_due_to_unknown_stream_shape = (
passthrough_due_to_unknown_stream_shape or not all_chunks
)
yield chunk
continue
for masked_chunk in await self._mask_anthropic_sse_stream(chunk, stream, request_data):
for masked_chunk in await self._mask_raw_sse_stream(chunk, stream, request_data):
yield masked_chunk
return
else:
@ -1465,7 +1473,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
if passthrough_due_to_unknown_stream_shape:
verbose_proxy_logger.warning(
"Presidio apply_to_output: streaming response was not a parsed chat completion stream "
"(raw non-Anthropic SSE passthrough or /v1/responses events). "
"(an error frame from an earlier guardrail, /v1/responses events, or a mixed stream). "
"Output PII masking was skipped for this response."
)
return
@ -1487,23 +1495,42 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
for chunk in all_chunks:
yield chunk
async def _mask_anthropic_sse_stream(
async def _mask_raw_sse_stream(
self, first_chunk: bytes, rest: AsyncIterator[object], request_data: dict
) -> tuple[object, ...]:
"""The whole raw SSE stream masked as one response, or a raised refusal when it cannot be read.
Raw frames are buffered to the end because PII can span frames, so no frame is forwarded
before the joined text was scanned. A stream whose surface is unknown, or that its surface's
assembler cannot rebuild, is withheld rather than forwarded unmasked.
"""
rest_chunks: Final = [chunk async for chunk in rest] # mutable-ok: tuple() cannot consume an async iterator
chunks: Final = (first_chunk, *rest_chunks)
assembled: Final = assemble_anthropic_sse_stream(chunks, restore_identity=True)
if assembled is None:
verbose_proxy_logger.warning(
"Presidio apply_to_output: raw SSE stream could not be assembled into a response. "
"Output PII masking was skipped for this response."
surface: Final = _assemble_raw_sse_stream(chunks)
if surface is None:
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=(
"output PII masking cannot read this streaming response shape, "
"so the response was withheld instead of being forwarded unmasked"
),
status_code=500,
)
return chunks
original_text: Final = model_response_text(assembled)
assembled, re_emit = surface
if assembled is None:
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=(
"output PII masking could not assemble the streaming response, "
"so the response was withheld instead of being forwarded unmasked"
),
status_code=500,
)
before: Final = assembled.model_dump()
await self._process_response_for_pii(response=assembled, request_data=request_data, mode="mask")
if model_response_text(assembled) == original_text:
if assembled.model_dump() == before:
return chunks
return anthropic_sse_chunks_from_response(assembled)
return re_emit(assembled)
@staticmethod
def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: dict[str, str]) -> bytes:

View file

@ -2,6 +2,7 @@ import json
import re
import signal
import threading
import time
import uuid
from collections.abc import Callable, Iterator, Mapping
from concurrent.futures import ThreadPoolExecutor
@ -257,67 +258,93 @@ def gemini_provider(reply: Reply) -> Callable[[Request], Reply]:
return provider
def test_native_gemini_first_frame_reaches_caller_before_upstream_sends_the_second(
gateway: Gateway, tmp_path: Path
) -> None:
gate: Final = threading.Event()
first: Final = gemini_frame("first ")
second: Final = gemini_frame("second ")
provider: Final = gemini_provider(
Reply(content_type="text/event-stream", chunks=(first, second), gate_after_first=gate)
)
def test_native_gemini_name_split_across_frames_is_masked_as_one_frame(gateway: Gateway, tmp_path: Path) -> None:
"""The analyzer only matches the whole name, so a per-frame scan would forward both halves."""
first, last = PERSON.split(" ")
frames: Final = (gemini_frame(f"{first} "), gemini_frame(f"{last} designed it."))
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames, pause_between_chunks=0.2))
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert received.status == 200, received.text
assert PERSON not in received.text, received.text
assert gemini_texts(b"".join(received.frames)) == (f"{MASK} designed it.",)
analyzed: Final = rig.analyzer.drain()
assert len(analyzed) == 1 and json.loads(analyzed[0].body)["text"] == f"{PERSON} designed it."
assert len(rig.anonymizer.drain()) == 1
def test_native_gemini_no_frame_reaches_caller_before_upstream_finishes(gateway: Gateway, tmp_path: Path) -> None:
gate: Final = threading.Event()
frames: Final = (gemini_frame(f"{PERSON} "), gemini_frame("designed it."))
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames, gate_after_first=gate))
with presidio_rig(gateway, tmp_path, provider) as rig:
armed_at: Final = time.monotonic()
threading.Timer(1.0, gate.set).start()
with rig.gateway.client.stream(
"POST", rig.gemini_path(), json=rig.gemini_body(), headers={"Authorization": f"Bearer {rig.gateway.key}"}
) as response:
assert response.status_code == 200, response.read().decode()
chunks: Final = response.iter_raw()
arrived: Final = next(chunks)
assert gemini_texts(arrived) == ("first ",), f"first chunk while upstream is gated: {arrived!r}"
gate.set()
arrived_at: Final = time.monotonic()
rest: Final = b"".join(chunks)
assert gemini_texts(rest) == ("second ",), rest
assert arrived_at - armed_at >= 1.0, f"a frame reached the caller while the upstream was gated: {arrived!r}"
assert gemini_texts(arrived + rest) == (f"{MASK} designed it.",), (arrived + rest).decode()
assert PERSON not in (arrived + rest).decode()
assert len(rig.upstream.drain()) == 1
assert rig.analyzer.drain() == () and rig.anonymizer.drain() == ()
def test_native_gemini_frames_received_before_upstream_abort_reach_caller(gateway: Gateway, tmp_path: Path) -> None:
def test_native_gemini_unrecognized_stream_shape_is_withheld(gateway: Gateway, tmp_path: Path) -> None:
frames: Final = (b'data: {"unexpected": "' + PERSON.encode() + b'"}\r\n\r\n',)
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames))
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert PERSON not in received.text, received.text
assert "cannot read this streaming response shape" in received.text, received.text
assert rig.analyzer.drain() == ()
def test_native_gemini_upstream_abort_mid_stream_returns_an_error_and_no_frame(
gateway: Gateway, tmp_path: Path
) -> None:
"""The stream is read whole before its first byte goes out, so an upstream abort still gets an error status."""
frames: Final = (gemini_frame(f"chunk {index} from {PERSON}. ") for index in range(3))
provider: Final = gemini_provider(
Reply(content_type="text/event-stream", chunks=tuple(frames), abort_after=2, pause_between_chunks=0.2)
)
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert received.status == 200, received.text
*frames_before_abort, trailer = data_payloads(b"".join(received.frames))
assert [gemini_text(frame) for frame in frames_before_abort] == [
f"chunk 0 from {PERSON}. ",
f"chunk 1 from {PERSON}. ",
], received.text
assert "candidates" not in trailer and json.dumps(trailer).count('"code": "500"') == 1, received.text
assert received.status == 500, received.text
assert PERSON not in received.text, received.text
error: Final = json.loads(received.text)["error"]
assert error["code"] == "500" and "candidates" not in received.text, received.text
assert len(rig.upstream.drain()) == 1
def test_native_gemini_first_frame_split_into_transport_fragments_streams_every_byte(
def test_native_gemini_first_frame_split_into_transport_fragments_is_still_masked(
gateway: Gateway, tmp_path: Path
) -> None:
first: Final = gemini_frame(f"fragmented {PERSON}")
first: Final = gemini_frame(f"fragmented {PERSON} ")
second: Final = gemini_frame("whole")
chunks: Final = (first[:7], first[7:19], first[19:], second)
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=chunks))
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert received.status == 200, received.text
assert gemini_texts(b"".join(received.frames)) == (f"fragmented {PERSON}", "whole")
assert PERSON not in received.text, received.text
assert gemini_texts(b"".join(received.frames)) == (f"fragmented {MASK} whole",)
def test_native_gemini_non_json_frame_passes_through_unchanged(gateway: Gateway, tmp_path: Path) -> None:
def test_native_gemini_non_json_frame_in_a_stream_without_pii_is_replayed_unchanged(
gateway: Gateway, tmp_path: Path
) -> None:
frames: Final = (b"data: not json at all\r\n\r\n", gemini_frame("after"))
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames))
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert received.status == 200, received.text
assert received.text.replace("\r\n", "\n") == b"".join(frames).decode().replace("\r\n", "\n")
assert len(rig.analyzer.drain()) == 1
def test_native_gemini_empty_stream_returns_200_with_no_body(gateway: Gateway, tmp_path: Path) -> None:
@ -328,14 +355,14 @@ def test_native_gemini_empty_stream_returns_200_with_no_body(gateway: Gateway, t
assert received.text == ""
def test_native_gemini_streams_while_presidio_analyzer_is_down(gateway: Gateway, tmp_path: Path) -> None:
def test_native_gemini_stream_fails_closed_when_analyzer_is_down(gateway: Gateway, tmp_path: Path) -> None:
frames: Final = (gemini_frame(f"{PERSON} one. "), gemini_frame("two."))
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames))
with presidio_rig(gateway, tmp_path, provider, analyze=broken) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert received.status == 200, received.text
assert gemini_texts(b"".join(received.frames)) == (f"{PERSON} one. ", "two.")
assert rig.analyzer.drain() == ()
assert PERSON not in received.text, received.text
assert "Presidio analyzer" in received.text, received.text
assert rig.anonymizer.drain() == ()
def test_native_gemini_unauthenticated_request_is_rejected_before_upstream(gateway: Gateway, tmp_path: Path) -> None:
@ -526,17 +553,13 @@ def test_mixed_burst_survives_anonymizer_outage_and_recovers(gateway: Gateway, t
def gemini_call(index: int) -> tuple[str, str, int]:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
return (
"gemini",
f"g{index}",
received.status if gemini_texts(b"".join(received.frames)) == (f"{PERSON} ", "designed it.") else -1,
)
masked: Final = gemini_texts(b"".join(received.frames)) == (f"{MASK} designed it.",)
return ("gemini", f"g{index}", -1 if PERSON in received.text else (1 if masked else 0))
def anthropic_call(index: int) -> tuple[str, str, int]:
body: Final = {**rig.messages_body(), "messages": [{"role": "user", "content": f"a{index}"}]}
received: Final = rig.stream("/v1/messages", body)
leaked: Final = PERSON in received.text
return ("anthropic", f"a{index}", -1 if leaked else (1 if MASK in received.text else 0))
return ("anthropic", f"a{index}", -1 if PERSON in received.text else (1 if MASK in received.text else 0))
def phase(offset: int) -> tuple[tuple[str, str, int], ...]:
with ThreadPoolExecutor(max_workers=12) as pool:
@ -552,13 +575,9 @@ def test_mixed_burst_survives_anonymizer_outage_and_recovers(gateway: Gateway, t
outage.clear()
healthy_after: Final = phase(200)
for name, results in (("before", healthy_before), ("during", during), ("after", healthy_after)):
assert all(status == 200 for kind, _, status in results if kind == "gemini"), (name, results)
assert all(status == 1 for kind, _, status in healthy_before + healthy_after if kind == "anthropic"), (
healthy_before,
healthy_after,
)
assert all(status == 0 for kind, _, status in during if kind == "anthropic"), during
for name, results in (("before", healthy_before), ("after", healthy_after)):
assert all(outcome == 1 for _, _, outcome in results), (name, results)
assert all(outcome == 0 for _, _, outcome in during), during
identities: Final = tuple(identity for _, identity, _ in healthy_before + during + healthy_after)
assert len(identities) == len(set(identities)) == 36

View file

@ -1449,11 +1449,10 @@ from litellm.types.utils import ModelResponseStream
@pytest.mark.asyncio
async def test_streaming_with_bytes_chunks_does_not_crash(mock_user_api_key):
async def test_streaming_unrecognized_raw_frame_ahead_of_typed_chunks_is_withheld(mock_user_api_key):
"""
Regression test: async_post_call_streaming_iterator_hook should
gracefully handle raw bytes in the stream instead of crashing with
'bytes' object has no attribute 'id'.
A raw frame the hook cannot place on a known surface is not a chunk it can
mask, so the response is refused instead of being forwarded unscanned.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
@ -1462,7 +1461,7 @@ async def test_streaming_with_bytes_chunks_does_not_crash(mock_user_api_key):
)
async def mock_stream():
yield b'data: {"id":"chatcmpl-1"}\n\n' # raw bytes
yield b'data: {"id":"chatcmpl-1"}\n\n'
yield ModelResponseStream(
id="chatcmpl-1",
choices=[],
@ -1470,18 +1469,22 @@ async def test_streaming_with_bytes_chunks_does_not_crash(mock_user_api_key):
model="gpt-4",
object="chat.completion.chunk",
system_fingerprint=None,
) # proper chunk
)
chunks = []
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_user_api_key,
response=mock_stream(),
request_data={},
):
chunks.append(chunk)
# Should not crash, should produce at least one valid chunk
assert len(chunks) >= 1
async def collect():
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_user_api_key,
response=mock_stream(),
request_data={},
):
chunks.append(chunk)
with pytest.raises(GuardrailRaisedException, match="cannot read this streaming response shape"):
await collect()
assert chunks == []
def test_entity_deny_list_filters_detections():
@ -2114,10 +2117,10 @@ async def test_anthropic_native_response_non_text_blocks_untouched():
@pytest.mark.asyncio
async def test_streaming_bytes_chunks_are_yielded_not_discarded():
async def test_streaming_partial_anthropic_stream_without_message_start_is_withheld():
"""
Regression test: bytes chunks (Anthropic native SSE) should be yielded
through the streaming hook, not silently discarded.
A raw Anthropic stream that cannot be assembled (no message_start) is
refused with a clear error rather than forwarded unmasked or dropped silently.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
@ -2132,15 +2135,19 @@ async def test_streaming_bytes_chunks_are_yielded_not_discarded():
mock_user_api_key = UserAPIKeyAuth(api_key="test-key")
chunks = []
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_user_api_key,
response=mock_stream(),
request_data={},
):
chunks.append(chunk)
assert any(isinstance(c, bytes) for c in chunks), "bytes chunks must not be discarded"
assert byte_chunk in chunks
async def collect():
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_user_api_key,
response=mock_stream(),
request_data={},
):
chunks.append(chunk)
with pytest.raises(GuardrailRaisedException, match="could not assemble the streaming response"):
await collect()
assert chunks == []
@pytest.mark.asyncio
@ -2524,13 +2531,135 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_without_pii_are_for
assert collected == byte_chunks
def _gemini_sse(text: str) -> bytes:
def _gemini_sse(text: str, terminator: bytes = b"\n\n") -> bytes:
payload = {"candidates": [{"content": {"parts": [{"text": text}], "role": "model"}, "index": 0}]}
return f"data: {json.dumps(payload)}\n\n".encode()
return b"data: " + json.dumps(payload).encode() + terminator
def _gemini_texts(chunks: list[object]) -> list[str]:
frames = b"".join(chunk for chunk in chunks if isinstance(chunk, bytes)).decode()
return [
json.loads(line[6:])["candidates"][0]["content"]["parts"][0]["text"]
for line in frames.splitlines()
if line.startswith("data: ")
]
async def _collect_masked_output(guardrail: _OPTIONAL_PresidioPIIMasking, stream, collected: list[object]) -> None:
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
response=stream,
request_data={},
):
collected.append(chunk)
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incrementally_until_upstream_aborts():
async def test_apply_to_output_streaming_gemini_name_split_across_frames_is_masked_as_one_response():
"""
The name only exists once the frames are joined, so a per-frame scan would
forward both halves unmasked. The analyzer and anonymizer are the in-process
fake, which finds a person only in two adjacent capitalized words.
"""
frames = [_gemini_sse("The architect was John"), _gemini_sse(" Smith, per the record.")]
async def mock_stream():
for frame in frames:
yield frame
collected: list[object] = []
async with TestServer(_fake_presidio_app()) as server:
guardrail = _OPTIONAL_PresidioPIIMasking(
apply_to_output=True,
presidio_analyzer_api_base=str(server.make_url("/")),
presidio_anonymizer_api_base=str(server.make_url("/")),
pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK},
)
await _collect_masked_output(guardrail, mock_stream(), collected)
await guardrail._close_http_session()
assert "John Smith" not in b"".join(collected).decode()
assert _gemini_texts(collected) == ["The architect was <PERSON>, per the record."]
@pytest.mark.asyncio
async def test_apply_to_output_streaming_anthropic_tool_call_arguments_only_masking_is_re_emitted():
"""
Masking can rewrite a tool call's arguments while the assistant text stays the
same, and replaying the original frames in that case would leak the arguments.
"""
byte_chunks = [
_anthropic_sse(
"message_start",
{"type": "message_start", "message": {"id": "msg_1", "model": "claude", "content": [], "usage": {}}},
),
_anthropic_sse(
"content_block_start",
{
"type": "content_block_start",
"index": 0,
"content_block": {"type": "tool_use", "id": "toolu_1", "name": "lookup", "input": {}},
},
),
_anthropic_sse(
"content_block_delta",
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "input_json_delta", "partial_json": '{"person": "John Smith"}'},
},
),
_anthropic_sse("content_block_stop", {"type": "content_block_stop", "index": 0}),
_anthropic_sse("message_delta", {"type": "message_delta", "delta": {"stop_reason": "tool_use"}, "usage": {}}),
_anthropic_sse("message_stop", {"type": "message_stop"}),
]
async def mock_stream():
for chunk in byte_chunks:
yield chunk
collected: list[object] = []
async with TestServer(_fake_presidio_app()) as server:
guardrail = _OPTIONAL_PresidioPIIMasking(
apply_to_output=True,
presidio_analyzer_api_base=str(server.make_url("/")),
presidio_anonymizer_api_base=str(server.make_url("/")),
pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK},
)
await _collect_masked_output(guardrail, mock_stream(), collected)
await guardrail._close_http_session()
joined = b"".join(collected).decode()
assert "John Smith" not in joined, joined
assert "<PERSON>" in joined, joined
assert joined.count("event: message_start") == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("terminator", [b"\n\n", b"\r\n\r\n"])
async def test_apply_to_output_streaming_gemini_stream_without_pii_is_replayed_byte_for_byte(terminator):
frames = [_gemini_sse("nothing personal ", terminator), _gemini_sse("in here.", terminator)]
async def mock_stream():
for frame in frames:
yield frame
collected: list[object] = []
async with TestServer(_fake_presidio_app()) as server:
guardrail = _OPTIONAL_PresidioPIIMasking(
apply_to_output=True,
presidio_analyzer_api_base=str(server.make_url("/")),
presidio_anonymizer_api_base=str(server.make_url("/")),
pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK},
)
await _collect_masked_output(guardrail, mock_stream(), collected)
await guardrail._close_http_session()
assert collected == frames
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_upstream_abort_mid_stream_forwards_no_frame():
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
@ -2544,18 +2673,31 @@ async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incremen
yield frame
raise ConnectionError("upstream closed mid-stream")
async def collect() -> None:
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
response=mock_stream(),
request_data={},
):
collected.append(chunk)
with pytest.raises(ConnectionError):
await collect()
await _collect_masked_output(guardrail, mock_stream(), collected)
assert collected == frames
assert collected == []
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_sse_bytes_fail_closed_when_presidio_is_unreachable():
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
presidio_analyzer_api_base="http://127.0.0.1:9",
presidio_anonymizer_api_base="http://127.0.0.1:9",
)
frames = [_gemini_sse("Hello John Smith"), _gemini_sse(" from the record.")]
collected: list[object] = []
async def mock_stream():
for frame in frames:
yield frame
with pytest.raises(Exception, match="Presidio PII analysis failed"):
await _collect_masked_output(guardrail, mock_stream(), collected)
assert collected == []
@pytest.mark.asyncio
@ -2846,7 +2988,7 @@ async def test_apply_to_output_streaming_comment_only_stream_is_forwarded_unchan
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_first_frame_split_across_transport_chunks_streams_incrementally():
async def test_apply_to_output_streaming_gemini_first_frame_split_across_transport_chunks_is_still_masked():
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
@ -2860,50 +3002,38 @@ async def test_apply_to_output_streaming_gemini_first_frame_split_across_transpo
yield first[:20]
yield first[20:]
yield second
raise ConnectionError("upstream closed mid-stream")
async def collect() -> None:
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
response=mock_stream(),
request_data={},
):
collected.append(chunk)
await _collect_masked_output(guardrail, mock_stream(), collected)
with pytest.raises(ConnectionError):
await collect()
assert collected == [first, second]
assert "John Smith" not in b"".join(collected).decode()
assert _gemini_texts(collected) == ["<PERSON>"]
@pytest.mark.asyncio
async def test_apply_to_output_streaming_unterminated_first_frame_is_released_once_it_exceeds_the_cap():
@pytest.mark.parametrize(
"frame",
[
b"data: not json at all\n\n",
b'data: {"unexpected": true}\n\n',
b"data: " + b"x" * (70 * 1024) + b"\n",
],
ids=["non-json", "unknown-surface", "unterminated"],
)
async def test_apply_to_output_streaming_unrecognized_raw_sse_stream_is_withheld(frame: bytes):
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
mock_redacted_text={"text": "<PERSON>"},
)
piece = b"data: " + b"x" * 1023 + b"\n"
pieces_to_cap = -(-(64 * 1024) // len(piece))
released_at: list[int] = []
collected: list[object] = []
async def mock_stream():
for index in range(pieces_to_cap * 4):
if collected:
released_at.append(index)
yield piece
yield frame
collected: list[object] = []
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
response=mock_stream(),
request_data={},
):
collected.append(chunk)
with pytest.raises(GuardrailRaisedException, match="cannot read this streaming response shape"):
await _collect_masked_output(guardrail, mock_stream(), collected)
assert released_at, "nothing reached the caller before the upstream finished"
assert released_at[0] == pieces_to_cap, released_at[:3]
assert b"".join(collected) == piece * (pieces_to_cap * 4)
assert collected == []
@pytest.mark.asyncio