diff --git a/litellm/proxy/guardrails/anthropic_sse.py b/litellm/proxy/guardrails/anthropic_sse.py index d8c785a498b..f4466dcf93d 100644 --- a/litellm/proxy/guardrails/anthropic_sse.py +++ b/litellm/proxy/guardrails/anthropic_sse.py @@ -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( ( diff --git a/litellm/proxy/guardrails/gemini_sse.py b/litellm/proxy/guardrails/gemini_sse.py index 7ecbc40b3b1..ff41df3c9b0 100644 --- a/litellm/proxy/guardrails/gemini_sse.py +++ b/litellm/proxy/guardrails/gemini_sse.py @@ -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})}) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 51119d2c4cf..ee10dc91c1d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -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: diff --git a/tests/integration/observability/test_presidio_streaming_output.py b/tests/integration/observability/test_presidio_streaming_output.py index d5190664347..73faa6a8344 100644 --- a/tests/integration/observability/test_presidio_streaming_output.py +++ b/tests/integration/observability/test_presidio_streaming_output.py @@ -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: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index 17bec2ecf36..9da3048a728 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -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 now."}, + { + "functionCall": { + "id": "call_1", + "name": "send_email", + "args": {"to": "", "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 ", "thought": True}, + {"text": "hello ,", "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": ""}, + ) + 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 ."}] + 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": ""}, ) - 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": ""}, + ) + 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 == []