diff --git a/litellm/proxy/guardrails/anthropic_sse.py b/litellm/proxy/guardrails/anthropic_sse.py index c446cfaa815..44281f7f8f2 100644 --- a/litellm/proxy/guardrails/anthropic_sse.py +++ b/litellm/proxy/guardrails/anthropic_sse.py @@ -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 ) diff --git a/litellm/proxy/guardrails/gemini_sse.py b/litellm/proxy/guardrails/gemini_sse.py new file mode 100644 index 00000000000..177880004b4 --- /dev/null +++ b/litellm/proxy/guardrails/gemini_sse.py @@ -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})}) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 94750f08a9e..2960c3ca82e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -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: diff --git a/tests/integration/observability/test_presidio_streaming_output.py b/tests/integration/observability/test_presidio_streaming_output.py index 680617f1652..73faa6a8344 100644 --- a/tests/integration/observability/test_presidio_streaming_output.py +++ b/tests/integration/observability/test_presidio_streaming_output.py @@ -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 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 a5625e45d75..23f77089165 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -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 , 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 "" 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 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_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 ."}] + 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 ."}] + 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": ""}, + ) + 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 ."}], + [{"text": "The engineer was ."}], + ] + + +@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": ""}, ) - 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) == [""] @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": ""}, ) - 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": ""}, + ) + 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