This commit is contained in:
devin-ai-integration[bot] 2026-09-30 10:27:42 -04:00 • committed by GitHub
commit a7fa010ec3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 950 additions and 155 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
@ -12,6 +13,8 @@ import json
from collections.abc import Mapping, Sequence
from typing import Final
from pydantic import TypeAdapter, ValidationError
from litellm.types.utils import Choices, ModelResponse
_ANTHROPIC_EVENT_TYPES: Final = frozenset(
@ -26,25 +29,29 @@ _ANTHROPIC_EVENT_TYPES: Final = frozenset(
"error",
}
)
_CONTENT_FREE_PAYLOADS: Final = frozenset({"", "[DONE]"})
_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object])
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,
)
@ -52,15 +59,49 @@ def _parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]:
return tuple(
event_data
for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses
if (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing
if isinstance(event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event), Mapping) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing
)
def _data_payload(event: str) -> str | None:
lines: Final = tuple(line.strip() for line in event.splitlines())
return next((line[len("data:") :].strip() for line in lines if line.startswith("data:")), None)
def _is_unreadable_payload(payload: str) -> bool:
if payload in _CONTENT_FREE_PAYLOADS:
return False
try:
_JSON_OBJECT.validate_json(payload)
except ValidationError:
return True
return False
def has_unreadable_sse_frames(all_chunks: Sequence[object]) -> bool:
"""Whether a ``data:`` payload is neither blank, ``[DONE]``, nor a JSON object.
``parsed_sse_events`` drops such a frame silently, which suits the assemblers and not a
masking hook: a frame it cannot read is one it cannot scan, so the hook withholds the stream
instead of replaying it.
"""
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
)
sse_stream: Final = joined_sse_stream(all_chunks)
if sse_stream is None:
return True
events: Final = AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same splitter the assembler uses
payloads: Final = tuple(_data_payload(event) for event in events)
return any(_is_unreadable_payload(payload) for payload in payloads if payload is not None)
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 +116,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 +136,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 +198,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,235 @@
"""Gemini SSE frame masking 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 fold such a stream into the one response body a non-streaming call would have returned,
rewrite its text and function-call arguments in place, and re-emit it as one frame, followed by any
upstream error frame kept as its own terminal frame. Every other field (thought signatures,
function-call ids, model version, response id, safety ratings, usage) is carried through untouched,
because a client echoes the model turn back on its next request and Gemini 3 rejects a function
call whose thought signature is missing.
"""
from __future__ import annotations
import json
from collections.abc import Awaitable, Callable, Iterable, Iterator, Mapping, Sequence
from dataclasses import dataclass
from itertools import accumulate, chain, groupby
from types import MappingProxyType
from typing import Final, TypeAlias
from pydantic import TypeAdapter, ValidationError
from litellm.proxy.guardrails.anthropic_sse import joined_sse_stream, parsed_sse_events
_GEMINI_RESPONSE_KEYS: Final = frozenset({"candidates", "usageMetadata", "promptFeedback"})
_TEXT_RUN_PART_KEYS: Final = frozenset({"text", "thought", "thoughtSignature"})
_EMPTY_UNSIGNED_TEXT_PART: Final = MappingProxyType({"text": ""})
_UNREADABLE_ARGUMENTS: Final = "left the function call arguments unreadable as JSON after masking"
TextMasker: TypeAlias = Callable[[str], Awaitable[str]] # mutable-ok: Callable params
_JsonObject: TypeAlias = Mapping[str, object]
_JSON_OBJECT: Final = TypeAdapter(_JsonObject)
_JSON_ARRAY: Final = TypeAdapter(tuple[object, ...])
@dataclass(frozen=True, slots=True)
class GeminiStreamUnchanged:
"""Masking rewrote nothing, so the original frames are replayed as they arrived."""
@dataclass(frozen=True, slots=True)
class GeminiStreamMasked:
frames: tuple[bytes, ...]
@dataclass(frozen=True, slots=True)
class GeminiStreamUnreadable:
reason: str
GeminiStreamMaskResult: TypeAlias = GeminiStreamUnchanged | GeminiStreamMasked | GeminiStreamUnreadable
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))
async def mask_gemini_sse_stream(all_chunks: Sequence[object], mask_text: TextMasker) -> GeminiStreamMaskResult:
"""The stream's frames folded into one response, with every text and function-call argument masked.
Text arrives split across frames, so the fragments of one text part are joined before they are
scanned; PII that only exists once the fragments meet is otherwise forwarded in halves. Top-level
keys and candidate keys take the last frame's value, and candidates are merged by index. An
upstream error frame stays its own terminal frame after the masked response, as it arrived.
"""
sse_stream: Final = joined_sse_stream(all_chunks)
if sse_stream is None:
return GeminiStreamUnreadable("could not decode the streaming response as UTF-8")
events: Final = parsed_sse_events(sse_stream)
candidates: Final = _merged_candidates(tuple(_candidates_of(events)))
masked_candidates: Final = tuple([await _masked_candidate(candidate, mask_text) for candidate in candidates])
unreadable: Final = next((item for item in masked_candidates if isinstance(item, GeminiStreamUnreadable)), None)
if unreadable is not None:
return unreadable
if masked_candidates == candidates:
return GeminiStreamUnchanged()
response: Final = MappingProxyType({**_without_error(_later_wins(events)), "candidates": masked_candidates})
return GeminiStreamMasked((_frame(response), *map(_frame, _error_frames(events))))
def _frame(event: _JsonObject) -> bytes:
return f"data: {_json_text(event)}\n\n".encode()
def _without_error(event: _JsonObject) -> _JsonObject:
return MappingProxyType({key: value for key, value in event.items() if key != "error"})
def _error_frames(events: Sequence[_JsonObject]) -> Iterator[_JsonObject]:
return (MappingProxyType({"error": event["error"]}) for event in events if "error" in event)
def _json_text(value: object) -> str:
return json.dumps(value, ensure_ascii=False, default=dict)
def _json_object(value: object) -> _JsonObject | None:
try:
return _JSON_OBJECT.validate_python(value)
except ValidationError:
return None
def _json_objects(values: Iterable[object]) -> Iterator[_JsonObject]:
return (parsed for parsed in map(_json_object, values) if parsed is not None)
def _json_array(value: object) -> tuple[object, ...]:
try:
return _JSON_ARRAY.validate_python(value)
except ValidationError:
return ()
def _later_wins(mappings: Sequence[_JsonObject]) -> _JsonObject:
return MappingProxyType(dict(chain.from_iterable(mapping.items() for mapping in mappings)))
def _candidates_of(events: Sequence[_JsonObject]) -> Iterator[_JsonObject]:
for event in events:
yield from _json_objects(_json_array(event.get("candidates")))
def _merged_candidates(candidates: Sequence[_JsonObject]) -> tuple[_JsonObject, ...]:
indexes: Final = tuple(dict.fromkeys(_candidate_index(candidate) for candidate in candidates))
return tuple(
_merged_candidate(tuple(candidate for candidate in candidates if _candidate_index(candidate) == index))
for index in indexes
)
def _candidate_index(candidate: _JsonObject) -> int:
index: Final = candidate.get("index")
return index if isinstance(index, int) else 0
def _merged_candidate(fragments: Sequence[_JsonObject]) -> _JsonObject:
merged: Final = _later_wins(fragments)
contents: Final = tuple(_json_objects(fragment.get("content") for fragment in fragments))
if not contents:
return merged
content: Final = MappingProxyType({**_later_wins(contents), "parts": _merged_parts(tuple(_parts_of(contents)))})
return MappingProxyType({**merged, "content": content})
def _parts_of(contents: Sequence[_JsonObject]) -> Iterator[object]:
for content in contents:
yield from _json_array(content.get("parts"))
def _merged_parts(parts: Sequence[object]) -> tuple[object, ...]:
"""Adjacent fragments of one text part joined into that part; every other part kept as it came.
The empty unsigned text part Gemini streams on its final frame is dropped, since the
non-streaming body carries none and Gemini rejects it when a client echoes it back.
"""
run_ids: Final = tuple(run_id for run_id, _ in accumulate(parts, _next_run, initial=(0, None)))[1:]
merged: Final = tuple(
_merged_run(tuple(part for part, _ in run)) for _, run in groupby(zip(parts, run_ids), key=lambda pair: pair[1])
)
return tuple(part for part in merged if part != _EMPTY_UNSIGNED_TEXT_PART)
def _next_run(state: tuple[int, object], part: object) -> tuple[int, object]:
run_id, previous = state
return (run_id if _continues_text_run(previous, part) else run_id + 1, part)
def _continues_text_run(previous: object, part: object) -> bool:
"""Whether ``part`` is the next fragment of the text part ``previous`` started.
A fragment carrying a thought signature closes its run, since the signature belongs to the
part it arrived on and one merged part can only carry one.
"""
previous_fragment: Final = _text_fragment(previous)
fragment: Final = _text_fragment(part)
if previous_fragment is None or fragment is None:
return False
return (
bool(previous_fragment.get("thought")) == bool(fragment.get("thought"))
and "thoughtSignature" not in previous_fragment
)
def _text_fragment(part: object) -> _JsonObject | None:
fragment: Final = _json_object(part)
if fragment is None or not _TEXT_RUN_PART_KEYS.issuperset(fragment) or not isinstance(fragment.get("text"), str):
return None
return fragment
def _merged_run(run: Sequence[object]) -> object:
"""A run longer than one part holds text fragments only, since nothing else continues a run."""
if len(run) == 1:
return run[0]
fragments: Final = tuple(_json_objects(run))
text: Final = "".join(str(fragment["text"]) for fragment in fragments)
return MappingProxyType({**fragments[0], **fragments[-1], "text": text})
async def _masked_candidate(candidate: _JsonObject, mask_text: TextMasker) -> _JsonObject | GeminiStreamUnreadable:
content: Final = _json_object(candidate.get("content"))
if content is None:
return candidate
parts: Final = _json_array(content.get("parts"))
masked_parts: Final = tuple([await _masked_part(part, mask_text) for part in parts])
unreadable: Final = next((item for item in masked_parts if isinstance(item, GeminiStreamUnreadable)), None)
if unreadable is not None:
return unreadable
return MappingProxyType({**candidate, "content": MappingProxyType({**content, "parts": masked_parts})})
async def _masked_part(part: object, mask_text: TextMasker) -> object | GeminiStreamUnreadable:
fragment: Final = _json_object(part)
if fragment is None:
return part
text: Final = fragment.get("text")
if isinstance(text, str):
return MappingProxyType({**fragment, "text": await mask_text(text)}) if text else part
function_call: Final = _json_object(fragment.get("functionCall"))
if function_call is None:
return part
arguments: Final = _json_object(function_call.get("args"))
if arguments is None:
return part
masked_arguments: Final = await mask_text(_json_text(arguments))
try:
parsed_arguments: Final = _JSON_OBJECT.validate_json(masked_arguments)
except ValidationError:
return GeminiStreamUnreadable(_UNREADABLE_ARGUMENTS)
return MappingProxyType({**fragment, "functionCall": MappingProxyType({**function_call, "args": parsed_arguments})})

View file

@ -19,7 +19,7 @@ from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast
import aiohttp
from typing_extensions import NotRequired, ReadOnly
from typing_extensions import NotRequired, ReadOnly, assert_never
import litellm
from litellm import get_secret
@ -45,8 +45,16 @@ from litellm.proxy.common_utils.sse_keepalive import split_complete_sse_frames
from litellm.proxy.guardrails.anthropic_sse import (
anthropic_sse_chunks_from_response,
assemble_anthropic_sse_stream,
has_unreadable_sse_frames,
is_anthropic_sse_stream,
model_response_text,
is_sse_error_stream,
)
from litellm.proxy.guardrails.gemini_sse import (
GeminiStreamMasked,
GeminiStreamUnchanged,
GeminiStreamUnreadable,
is_gemini_sse_stream,
mask_gemini_sse_stream,
)
from litellm.types.guardrails import (
GuardrailEventHooks,
@ -97,9 +105,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 +144,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 +158,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:
@ -1444,18 +1448,10 @@ 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:
passthrough_due_to_unknown_stream_shape = (
passthrough_due_to_unknown_stream_shape or not all_chunks
)
if all_chunks or passthrough_due_to_unknown_stream_shape:
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:
@ -1467,7 +1463,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
@ -1489,24 +1485,71 @@ 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, ...]:
rest_chunks: Final = [chunk async for chunk in rest] # mutable-ok: tuple() cannot consume an async iterator
"""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, and the surface is decided on the whole stream rather
than on its first frame. A stream that is nothing but the refusal an earlier guardrail in
the chain emitted is forwarded as it arrived. A stream carrying a frame the parser cannot
read, whose surface is unknown, or that its surface's assembler cannot rebuild, is withheld
rather than forwarded unmasked.
"""
rest_chunks: Final = tuple([chunk async for chunk in rest])
chunks: Final = (first_chunk, *rest_chunks)
if has_unreadable_sse_frames(chunks):
raise self._withheld_stream_error("could not read every streamed frame")
if is_sse_error_stream(chunks):
return chunks
if is_anthropic_sse_stream(chunks):
return await self._mask_anthropic_sse_stream(chunks, request_data)
if is_gemini_sse_stream(chunks):
return await self._mask_gemini_sse_stream(chunks, request_data)
raise self._withheld_stream_error("cannot read this streaming response shape")
async def _mask_anthropic_sse_stream(self, chunks: tuple[object, ...], request_data: dict) -> tuple[object, ...]:
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."
)
return chunks
original_text: Final = model_response_text(assembled)
raise self._withheld_stream_error("could not assemble the streaming response")
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)
async def _mask_gemini_sse_stream(self, chunks: tuple[object, ...], request_data: dict) -> tuple[object, ...]:
"""Gemini frames masked in place, so thought signatures and function-call ids reach the client.
The frames are rewritten rather than round-tripped through a chat-completions response
because that translation drops the fields a Gemini client has to echo on its next turn.
"""
presidio_config: Final = self.get_presidio_settings_from_request_data(request_data or {})
async def mask_text(text: str) -> str:
return await self.check_pii(
text=text, output_parse_pii=False, presidio_config=presidio_config, request_data=request_data
)
result: Final = await mask_gemini_sse_stream(chunks, mask_text)
match result:
case GeminiStreamUnchanged():
return chunks
case GeminiStreamMasked(frames=frames):
return frames
case GeminiStreamUnreadable(reason=reason):
raise self._withheld_stream_error(reason)
case _:
assert_never(result)
def _withheld_stream_error(self, reason: str) -> GuardrailRaisedException:
return GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=f"output PII masking {reason}, so the response was withheld instead of being forwarded unmasked",
status_code=500,
)
@staticmethod
def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: dict[str, str]) -> bytes:
try:

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:
frames: Final = (b"data: not json at all\r\n\r\n", gemini_frame("after"))
def test_native_gemini_stream_with_an_unreadable_frame_is_withheld(gateway: Gateway, tmp_path: Path) -> None:
"""A frame the parser cannot read is one the guardrail cannot scan, so nothing around it is replayed either."""
frames: Final = (gemini_frame(f"the architect was {PERSON}"), b"data: not json at all\r\n\r\n", gemini_frame("."))
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 received.status == 500, received.text
assert PERSON not in received.text, received.text
assert "could not read every streamed frame" in received.text, received.text
assert rig.analyzer.drain() == ()
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
@ -2528,13 +2535,397 @@ 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: ")
]
def _gemini_frame(*parts: dict, finish_reason: str | None = None, **top_level: object) -> bytes:
candidate = {"content": {"parts": list(parts), "role": "model"}, "index": 0}
payload = {
"candidates": [candidate if finish_reason is None else {**candidate, "finishReason": finish_reason}],
**top_level,
}
return b"data: " + json.dumps(payload).encode() + b"\n\n"
def _gemini_frames(chunks: list[object]) -> list[dict]:
frames = b"".join(chunk for chunk in chunks if isinstance(chunk, bytes)).decode()
return [json.loads(line[6:]) for line in frames.splitlines() if line.startswith("data: ")]
def _gemini_fake_masking_guardrail(server: TestServer) -> _OPTIONAL_PresidioPIIMasking:
return _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},
)
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):
function_call = {
"functionCall": {"id": "call_1", "name": "lookup", "args": {"city": "Paris"}},
"thoughtSignature": "sig",
}
frames = [
_gemini_sse("nothing personal ", terminator),
_gemini_sse("in here.", terminator),
b"data: "
+ json.dumps({"candidates": [{"content": {"parts": [function_call], "role": "model"}}]}).encode()
+ 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_function_call_keeps_its_signature_and_id_while_its_arguments_are_masked():
"""
A Gemini client echoes the streamed model turn on its next request, and Gemini 3 rejects
a function call whose thought signature is missing, so masking has to rewrite the frames
in place rather than rebuild them from a chat-completions response that drops those fields.
"""
frames = [
_gemini_frame(
{"text": "emailing John"}, usageMetadata={"totalTokenCount": 1}, modelVersion="m", responseId="r1"
),
_gemini_frame({"text": " Smith now."}, usageMetadata={"totalTokenCount": 2}, modelVersion="m", responseId="r1"),
_gemini_frame(
{
"functionCall": {"id": "call_1", "name": "send_email", "args": {"to": "John Smith", "urgent": True}},
"thoughtSignature": "sig-fc",
},
usageMetadata={"totalTokenCount": 3},
modelVersion="m",
responseId="r1",
),
_gemini_frame(
{"text": ""}, finish_reason="STOP", usageMetadata={"totalTokenCount": 9}, modelVersion="m", responseId="r1"
),
]
async def mock_stream():
for frame in frames:
yield frame
collected: list[object] = []
async with TestServer(_fake_presidio_app()) as server:
guardrail = _gemini_fake_masking_guardrail(server)
await _collect_masked_output(guardrail, mock_stream(), collected)
await guardrail._close_http_session()
assert "John Smith" not in b"".join(collected).decode()
(body,) = _gemini_frames(collected)
assert body["candidates"] == [
{
"content": {
"parts": [
{"text": "emailing <PERSON> now."},
{
"functionCall": {
"id": "call_1",
"name": "send_email",
"args": {"to": "<PERSON>", "urgent": True},
},
"thoughtSignature": "sig-fc",
},
],
"role": "model",
},
"index": 0,
"finishReason": "STOP",
}
]
assert body["usageMetadata"] == {"totalTokenCount": 9}
assert body["modelVersion"] == "m"
assert body["responseId"] == "r1"
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_thought_part_and_trailing_signature_survive_the_merge():
"""
A thought part is never merged into the visible text that follows it, and the signature
Gemini sends on a trailing empty text part lands on the text it signs.
"""
frames = [
_gemini_frame({"text": "thinking about John Smith", "thought": True}),
_gemini_frame({"text": "hello John"}),
_gemini_frame({"text": " Smith,"}),
_gemini_frame({"text": "", "thoughtSignature": "sig-text"}, finish_reason="STOP"),
]
async def mock_stream():
for frame in frames:
yield frame
collected: list[object] = []
async with TestServer(_fake_presidio_app()) as server:
guardrail = _gemini_fake_masking_guardrail(server)
await _collect_masked_output(guardrail, mock_stream(), collected)
await guardrail._close_http_session()
(body,) = _gemini_frames(collected)
assert body["candidates"][0]["content"]["parts"] == [
{"text": "thinking about <PERSON>", "thought": True},
{"text": "hello <PERSON>,", "thoughtSignature": "sig-text"},
]
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_masked_function_call_arguments_that_no_longer_parse_are_withheld():
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
mock_redacted_text={"text": "<PERSON>"},
)
frames = [
_gemini_frame({"functionCall": {"name": "lookup", "args": {"person": "John Smith"}}}, finish_reason="STOP")
]
collected: list[object] = []
async def mock_stream():
for frame in frames:
yield frame
with pytest.raises(GuardrailRaisedException) as raised:
await _collect_masked_output(guardrail, mock_stream(), collected)
assert raised.value.status_code == 500
assert collected == []
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_error_frame_mid_stream_keeps_its_own_terminal_frame():
"""
Gemini can end a stream with an error frame after content frames, and a client
reads the failure off that terminal frame, so the masked re-emit keeps it separate.
"""
frames = [
_gemini_frame({"text": "The architect was John Smith."}),
b'data: {"error": {"code": 503, "message": "overloaded", "status": "UNAVAILABLE"}}\n\n',
]
collected: list[object] = []
async def mock_stream():
for frame in frames:
yield frame
async with TestServer(_fake_presidio_app()) as server:
guardrail = _gemini_fake_masking_guardrail(server)
await _collect_masked_output(guardrail, mock_stream(), collected)
await guardrail._close_http_session()
masked, error = _gemini_frames(collected)
assert masked["candidates"][0]["content"]["parts"] == [{"text": "The architect was <PERSON>."}]
assert "error" not in masked
assert error == {"error": {"code": 503, "message": "overloaded", "status": "UNAVAILABLE"}}
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_content_behind_a_leading_error_frame_is_still_masked():
"""
The stream surface is decided on the whole buffered stream, not on its first frame, so an
error frame ahead of content frames does not switch masking off for the content behind it.
"""
frames = [
b'data: {"error": {"code": 429, "message": "quota", "status": "RESOURCE_EXHAUSTED"}}\n\n',
_gemini_frame({"text": "The architect was John Smith."}),
]
collected: list[object] = []
async def mock_stream():
for frame in frames:
yield frame
async with TestServer(_fake_presidio_app()) as server:
guardrail = _gemini_fake_masking_guardrail(server)
await _collect_masked_output(guardrail, mock_stream(), collected)
await guardrail._close_http_session()
masked, error = _gemini_frames(collected)
assert masked["candidates"][0]["content"]["parts"] == [{"text": "The architect was <PERSON>."}]
assert error == {"error": {"code": 429, "message": "quota", "status": "RESOURCE_EXHAUSTED"}}
assert not any(b"John Smith" in chunk for chunk in collected if isinstance(chunk, bytes))
@pytest.mark.asyncio
async def test_apply_to_output_streaming_error_only_raw_stream_is_forwarded_as_it_arrived():
"""
A post_call chain hands this hook the refusal frames an earlier guardrail emitted. They carry
nothing to scan, and withholding or rewriting them would hide the refusal the client is owed.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
mock_redacted_text={"text": "<PERSON>"},
)
frames = [
b'data: {"error": {"message": "Violated guardrail policy", "type": "guardrail_error"}}\n\n',
b"data: [DONE]\n\n",
]
collected: list[object] = []
async def mock_stream():
for frame in frames:
yield frame
await _collect_masked_output(guardrail, mock_stream(), collected)
assert collected == frames
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_every_candidate_is_merged_by_index_and_masked():
"""
Gemini interleaves every candidate's fragments across frames, so each candidate
is joined on its own index and scanned, not only the first one.
"""
def frame(*texts: tuple[int, str]) -> bytes:
candidates = [{"content": {"parts": [{"text": text}], "role": "model"}, "index": index} for index, text in texts]
return b"data: " + json.dumps({"candidates": candidates}).encode() + b"\n\n"
frames = [
frame((0, "The architect was John"), (1, "The engineer was Ada")),
frame((0, " Smith."), (1, " Lovelace.")),
]
collected: list[object] = []
async def mock_stream():
for chunk in frames:
yield chunk
async with TestServer(_fake_presidio_app()) as server:
guardrail = _gemini_fake_masking_guardrail(server)
await _collect_masked_output(guardrail, mock_stream(), collected)
await guardrail._close_http_session()
(masked,) = _gemini_frames(collected)
assert [candidate["content"]["parts"] for candidate in masked["candidates"]] == [
[{"text": "The architect was <PERSON>."}],
[{"text": "The engineer was <PERSON>."}],
]
@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,
@ -2548,18 +2939,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
@ -2787,13 +3191,13 @@ async def test_apply_to_output_streaming_leading_keepalive_is_forwarded_before_u
@pytest.mark.asyncio
async def test_apply_to_output_streaming_leading_comments_over_the_frame_cap_are_still_masked():
async def test_apply_to_output_streaming_leading_comments_of_any_size_are_still_masked():
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
mock_redacted_text={"text": "<PERSON>"},
)
keepalives = [b": keepalive\n\n" * 512] * 12 # ~72 KiB of complete comment frames, over the 64 KiB cap
keepalives = [b": keepalive\n\n" * 512] * 12 # ~72 KiB of complete comment frames
byte_chunks = [
*keepalives[:-1],
keepalives[-1]
@ -2850,7 +3254,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,
@ -2864,50 +3268,103 @@ 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", "reason"),
[
(b"data: not json at all\n\n", "could not read every streamed frame"),
(b"data: 42\n\n", "could not read every streamed frame"),
(b'data: {"unexpected": true}\n\n', "cannot read this streaming response shape"),
(b"data: " + b"x" * (70 * 1024) + b"\n", "could not read every streamed frame"),
],
ids=["non-json", "json-scalar", "unknown-surface", "unterminated"],
)
async def test_apply_to_output_streaming_unrecognized_raw_sse_stream_is_withheld(frame: bytes, reason: str):
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
with pytest.raises(GuardrailRaisedException, match=reason):
await _collect_masked_output(guardrail, mock_stream(), collected)
assert collected == []
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_stream_with_an_unreadable_frame_is_withheld():
"""
The parser skips a frame it cannot read, and replaying the originals around it
would forward that frame unscanned, so the whole stream is withheld instead.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
mock_redacted_text={"text": "<PERSON>"},
)
frames = [
_gemini_frame({"text": "The architect was"}),
b"data: not json at all\n\n",
_gemini_frame({"text": " Ada."}, finish_reason="STOP"),
]
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)
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)
async def mock_stream():
for frame in frames:
yield frame
with pytest.raises(GuardrailRaisedException, match="could not read every streamed frame"):
await _collect_masked_output(guardrail, mock_stream(), collected)
assert collected == []
@pytest.mark.asyncio
async def test_apply_to_output_streaming_anthropic_stream_with_an_unreadable_frame_is_withheld():
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
mock_redacted_text={"text": "Hello"},
)
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": "text", "text": ""}},
),
b"event: content_block_delta\ndata: not json at all\n\n",
_anthropic_sse(
"content_block_delta",
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "Hello"}},
),
_anthropic_sse("content_block_stop", {"type": "content_block_stop", "index": 0}),
_anthropic_sse("message_stop", {"type": "message_stop"}),
]
collected: list[object] = []
async def mock_stream():
for chunk in byte_chunks:
yield chunk
with pytest.raises(GuardrailRaisedException, match="could not read every streamed frame"):
await _collect_masked_output(guardrail, mock_stream(), collected)
assert collected == []
@pytest.mark.asyncio