fix(guardrails): rewrite masked Gemini SSE frames in place and withhold streams with unreadable frames

This commit is contained in:
mateo-berri 2026-09-26 15:52:07 -07:00
parent a2efc1c321
commit 95cf660961
5 changed files with 535 additions and 86 deletions

View file

@ -13,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(
@ -27,6 +29,8 @@ _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:
@ -55,10 +59,44 @@ 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(
(

View file

@ -1,28 +1,54 @@
"""Gemini SSE <-> ModelResponse conversion for guardrail streaming hooks.
"""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 assemble such a stream into a ModelResponse and re-emit it once the guardrail rewrote it.
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. 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 Mapping, Sequence
from collections.abc import Awaitable, Callable, Iterable, Iterator, Mapping, Sequence
from dataclasses import dataclass
from typing import Final
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
from litellm.types.utils import ModelResponse
_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 _StreamParserLogging:
"""The one attribute the Gemini stream parser reads off its logging object."""
class GeminiStreamUnchanged:
"""Masking rewrote nothing, so the original frames are replayed as they arrived."""
optional_params: Mapping[str, object]
@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:
@ -32,34 +58,164 @@ def is_gemini_sse_stream(all_chunks: Sequence[object]) -> bool:
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.
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.
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.
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.
"""
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 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({**_later_wins(events), "candidates": masked_candidates})
return GeminiStreamMasked((f"data: {_json_text(response)}\n\n".encode(),))
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
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`
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
)
chunks: Final = tuple(
parsed for event in parsed_sse_events(sse_stream) if (parsed := parser.chunk_parser(dict(event))) is not None
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])
)
if not chunks:
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
assembled: Final = stream_chunk_builder(chunks=list(chunks))
return assembled if isinstance(assembled, ModelResponse) else None
return fragment
def gemini_sse_chunks_from_response(assembled: ModelResponse) -> tuple[bytes, ...]:
from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter
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})
body: Final = GoogleGenAIAdapter().translate_completion_to_generate_content(response=assembled)
return (f"data: {json.dumps(body)}\n\n".encode(),)
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

@ -12,14 +12,14 @@ import asyncio
import json
import re
import threading
from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaitable, Callable, Iterator, Sequence
from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaitable, Iterator, Sequence
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeAlias, TypedDict, cast
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,13 +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,
is_sse_error_stream,
)
from litellm.proxy.guardrails.gemini_sse import (
assemble_gemini_sse_stream,
gemini_sse_chunks_from_response,
GeminiStreamMasked,
GeminiStreamUnchanged,
GeminiStreamUnreadable,
is_gemini_sse_stream,
mask_gemini_sse_stream,
)
from litellm.types.guardrails import (
GuardrailEventHooks,
@ -170,18 +173,6 @@ 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
@ -1501,36 +1492,60 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
"""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.
before the joined text was scanned. 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 = [chunk async for chunk in rest] # mutable-ok: tuple() cannot consume an async iterator
rest_chunks: Final = tuple([chunk async for chunk in rest])
chunks: Final = (first_chunk, *rest_chunks)
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,
)
assembled, re_emit = surface
if has_unreadable_sse_frames(chunks):
raise self._withheld_stream_error("could not read every streamed frame")
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:
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,
)
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 assembled.model_dump() == before:
return chunks
return re_emit(assembled)
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:

View file

@ -335,16 +335,16 @@ def test_native_gemini_first_frame_split_into_transport_fragments_is_still_maske
assert gemini_texts(b"".join(received.frames)) == (f"fragmented {MASK} whole",)
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"))
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 len(rig.analyzer.drain()) == 1
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:

View file

@ -2545,6 +2545,29 @@ def _gemini_texts(chunks: list[object]) -> list[str]:
]
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"),
@ -2638,7 +2661,17 @@ async def test_apply_to_output_streaming_anthropic_tool_call_arguments_only_mask
@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)]
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:
@ -2658,6 +2691,148 @@ async def test_apply_to_output_streaming_gemini_stream_without_pii_is_replayed_b
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_is_carried_on_the_masked_frame():
"""
Gemini can end a stream with an error frame after content frames, and the
client reads the error off the last frame, so the masked re-emit keeps it.
"""
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()
(frame,) = _gemini_frames(collected)
assert frame["candidates"][0]["content"]["parts"] == [{"text": "The architect was <PERSON>."}]
assert frame["error"] == {"code": 503, "message": "overloaded", "status": "UNAVAILABLE"}
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_upstream_abort_mid_stream_forwards_no_frame():
guardrail = _OPTIONAL_PresidioPIIMasking(
@ -2925,13 +3100,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]
@ -3011,15 +3186,16 @@ async def test_apply_to_output_streaming_gemini_first_frame_split_across_transpo
@pytest.mark.asyncio
@pytest.mark.parametrize(
"frame",
("frame", "reason"),
[
b"data: not json at all\n\n",
b'data: {"unexpected": true}\n\n',
b"data: " + b"x" * (70 * 1024) + b"\n",
(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", "unknown-surface", "unterminated"],
ids=["non-json", "json-scalar", "unknown-surface", "unterminated"],
)
async def test_apply_to_output_streaming_unrecognized_raw_sse_stream_is_withheld(frame: bytes):
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,
@ -3030,7 +3206,71 @@ async def test_apply_to_output_streaming_unrecognized_raw_sse_stream_is_withheld
async def mock_stream():
yield frame
with pytest.raises(GuardrailRaisedException, match="cannot read this streaming response shape"):
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 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 == []