From a2efc1c321c9ce69d1940c80592675b7cf58de0d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 13:50:41 -0700 Subject: [PATCH 1/5] fix(guardrails): buffer and mask raw Gemini SSE streams in Presidio output masking and fail closed on unknown stream shapes With presidio_filter_scope output or both, a Google :streamGenerateContent stream reached the Presidio streaming hook as raw SSE bytes that only the Anthropic path knew how to read, so the Gemini frames were forwarded with the PII the guardrail was configured to mask, while the same request without streaming came back masked. Raw Gemini streams are now assembled, masked, and re-emitted as one provider-shaped frame the way Anthropic streams already were, and any raw stream whose shape neither assembler can read is withheld with a 500 instead of being passed through unmasked. Whether masking changed the response is decided on the whole assembled response, so a masked tool-call argument is re-emitted too. --- litellm/proxy/guardrails/anthropic_sse.py | 23 +- litellm/proxy/guardrails/gemini_sse.py | 65 +++++ .../guardrails/guardrail_hooks/presidio.py | 81 ++++-- .../test_presidio_streaming_output.py | 105 ++++--- .../guardrail_hooks/test_presidio.py | 268 +++++++++++++----- 5 files changed, 393 insertions(+), 149 deletions(-) create mode 100644 litellm/proxy/guardrails/gemini_sse.py diff --git a/litellm/proxy/guardrails/anthropic_sse.py b/litellm/proxy/guardrails/anthropic_sse.py index 26dd2a95dc4..d8c785a498b 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 @@ -32,19 +33,21 @@ def is_raw_sse_stream(all_chunks: Sequence[object]) -> bool: return any(isinstance(chunk, (str, bytes)) for chunk in all_chunks) -def _joined_sse_stream(all_chunks: Sequence[object]) -> str | None: +def joined_sse_stream(all_chunks: Sequence[object]) -> str | None: + """The raw frames decoded as one text, with every SSE line ending (CRLF, LF, CR) folded to LF.""" raw: Final = b"".join( chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in all_chunks if isinstance(chunk, (str, bytes)) ) try: - return codecs.getincrementaldecoder("utf-8")().decode(raw, final=False) + decoded: Final = codecs.getincrementaldecoder("utf-8")().decode(raw, final=False) except UnicodeDecodeError: return None + return decoded.replace("\r\n", "\n").replace("\r", "\n") -def _parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]: +def parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]: from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( AnthropicPassthroughLoggingHandler, ) @@ -60,7 +63,7 @@ def _anthropic_message_start(sse_stream: str) -> Mapping[str, object] | None: return next( ( message - for event_data in _parsed_sse_events(sse_stream) + for event_data in parsed_sse_events(sse_stream) if event_data.get("type") == "message_start" and isinstance(message := event_data.get("message"), dict) ), None, @@ -75,10 +78,10 @@ def is_anthropic_sse_stream(all_chunks: Sequence[object]) -> bool: stream raw too. Reading its frames as Anthropic ones would refuse the response in a wire format its client cannot parse, so the surface is decided on the event types actually present. """ - sse_stream: Final = _joined_sse_stream(all_chunks) + sse_stream: Final = joined_sse_stream(all_chunks) if sse_stream is None: return False - return any(event.get("type") in _ANTHROPIC_EVENT_TYPES for event in _parsed_sse_events(sse_stream)) + return any(event.get("type") in _ANTHROPIC_EVENT_TYPES for event in parsed_sse_events(sse_stream)) def assemble_anthropic_sse_stream( @@ -95,7 +98,7 @@ def assemble_anthropic_sse_stream( AnthropicPassthroughLoggingHandler, ) - sse_stream: Final = _joined_sse_stream(all_chunks) + sse_stream: Final = joined_sse_stream(all_chunks) if sse_stream is None: return None message_start: Final = _anthropic_message_start(sse_stream) @@ -157,10 +160,10 @@ def is_sse_error_stream(all_chunks: Sequence[object]) -> bool: # A stream mixing typed chunks with an error frame still carries content to scan, and the # frames-only join below would drop exactly the part that has to be scanned return False - sse_stream: Final = _joined_sse_stream(all_chunks) + sse_stream: Final = joined_sse_stream(all_chunks) if sse_stream is None: return False - events: Final = _parsed_sse_events(sse_stream) + events: Final = parsed_sse_events(sse_stream) return len(events) > 0 and all( event.get("type") == "error" or isinstance(event.get("error"), Mapping) for event in events ) diff --git a/litellm/proxy/guardrails/gemini_sse.py b/litellm/proxy/guardrails/gemini_sse.py new file mode 100644 index 00000000000..7ecbc40b3b1 --- /dev/null +++ b/litellm/proxy/guardrails/gemini_sse.py @@ -0,0 +1,65 @@ +"""Gemini SSE <-> ModelResponse conversion for guardrail streaming hooks. + +The Google ``:streamGenerateContent`` route relays the upstream ``data:`` frames as raw bytes, so a +guardrail's ``async_post_call_streaming_iterator_hook`` cannot read them as chunk objects. These +helpers assemble such a stream into a ModelResponse and re-emit it once the guardrail rewrote it. +""" + +from __future__ import annotations + +import json +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import Final + +from litellm.proxy.guardrails.anthropic_sse import joined_sse_stream, parsed_sse_events +from litellm.types.utils import ModelResponse + +_GEMINI_RESPONSE_KEYS: Final = frozenset({"candidates", "usageMetadata", "promptFeedback"}) + + +@dataclass(frozen=True, slots=True) +class _StreamParserLogging: + """The one attribute the Gemini stream parser reads off its logging object.""" + + optional_params: Mapping[str, object] + + +def is_gemini_sse_stream(all_chunks: Sequence[object]) -> bool: + sse_stream: Final = joined_sse_stream(all_chunks) + if sse_stream is None: + return False + return any(not _GEMINI_RESPONSE_KEYS.isdisjoint(event) for event in parsed_sse_events(sse_stream)) + + +def assemble_gemini_sse_stream(all_chunks: Sequence[object]) -> ModelResponse | None: + """Assemble raw Gemini SSE frames into a ModelResponse. + + A frame the Gemini parser rejects (a mid-stream ``error`` payload, say) raises the parser's own + error so the caller reports the upstream failure rather than a guardrail one. + """ + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ModelResponseIterator + from litellm.main import stream_chunk_builder + + sse_stream: Final = joined_sse_stream(all_chunks) + if sse_stream is None: + return None + parser: Final = ModelResponseIterator( + streaming_response=None, + sync_stream=False, + logging_obj=_StreamParserLogging(optional_params={}), # pyright: ignore[reportArgumentType] # the parser reads only optional_params, and a native Gemini request carries no legacy `functions` + ) + chunks: Final = tuple( + parsed for event in parsed_sse_events(sse_stream) if (parsed := parser.chunk_parser(dict(event))) is not None + ) + if not chunks: + return None + assembled: Final = stream_chunk_builder(chunks=list(chunks)) + return assembled if isinstance(assembled, ModelResponse) else None + + +def gemini_sse_chunks_from_response(assembled: ModelResponse) -> tuple[bytes, ...]: + from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter + + body: Final = GoogleGenAIAdapter().translate_completion_to_generate_content(response=assembled) + return (f"data: {json.dumps(body)}\n\n".encode(),) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index f5e24c501f1..51119d2c4cf 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -12,11 +12,11 @@ import asyncio import json import re import threading -from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaitable, Iterator, Sequence +from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaitable, Callable, Iterator, Sequence from contextlib import asynccontextmanager from dataclasses import dataclass from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast +from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeAlias, TypedDict, cast import aiohttp from typing_extensions import NotRequired, ReadOnly @@ -46,7 +46,12 @@ from litellm.proxy.guardrails.anthropic_sse import ( anthropic_sse_chunks_from_response, assemble_anthropic_sse_stream, is_anthropic_sse_stream, - model_response_text, + is_sse_error_stream, +) +from litellm.proxy.guardrails.gemini_sse import ( + assemble_gemini_sse_stream, + gemini_sse_chunks_from_response, + is_gemini_sse_stream, ) from litellm.types.guardrails import ( GuardrailEventHooks, @@ -97,9 +102,6 @@ def _json_escaped_len(text: str) -> int: return len(json.dumps(text).encode("utf-8")) - 2 # strip the surrounding quotes -_MAX_FIRST_SSE_FRAME_BYTES: Final = 64 * 1024 - - @dataclass(frozen=True, slots=True) class _SsePreface: """Complete leading SSE frames with no ``data:`` line, relayed verbatim before the stream shape is decided.""" @@ -139,8 +141,7 @@ async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGener ``data:`` line) as they complete, and join raw ``bytes`` chunks until they hold one complete SSE event with a data line, so the stream shape is decided on a whole frame rather than a transport fragment. Everything - after that first frame is forwarded untouched. The byte cap can only be - reached by a single unterminated frame. + after that first frame is forwarded untouched. """ pending = b"" try: @@ -154,7 +155,7 @@ async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGener if preface: yield _SsePreface(preface) pending = classifiable + tail - if classifiable or len(pending) >= _MAX_FIRST_SSE_FRAME_BYTES: + if classifiable: break else: if pending: @@ -169,6 +170,18 @@ async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGener yield chunk +_RawSseReEmit: TypeAlias = Callable[[ModelResponse], tuple[bytes, ...]] + + +def _assemble_raw_sse_stream(chunks: Sequence[object]) -> tuple[ModelResponse | None, _RawSseReEmit] | None: + """The stream assembled by its surface, with that surface's re-emitter; None for an unknown surface.""" + if is_anthropic_sse_stream(chunks): + return assemble_anthropic_sse_stream(chunks, restore_identity=True), anthropic_sse_chunks_from_response + if is_gemini_sse_stream(chunks): + return assemble_gemini_sse_stream(chunks), gemini_sse_chunks_from_response + return None + + class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): user_api_key_cache = None ad_hoc_recognizers: list[str] | None = None @@ -1442,18 +1455,13 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): elif isinstance(chunk, _SsePreface): yield chunk.raw elif isinstance(chunk, bytes): - first_frame_is_anthropic = ( - not passthrough_due_to_unknown_stream_shape - and not all_chunks - and is_anthropic_sse_stream((chunk,)) - ) - if not first_frame_is_anthropic: + if all_chunks or passthrough_due_to_unknown_stream_shape or is_sse_error_stream((chunk,)): passthrough_due_to_unknown_stream_shape = ( passthrough_due_to_unknown_stream_shape or not all_chunks ) yield chunk continue - for masked_chunk in await self._mask_anthropic_sse_stream(chunk, stream, request_data): + for masked_chunk in await self._mask_raw_sse_stream(chunk, stream, request_data): yield masked_chunk return else: @@ -1465,7 +1473,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if passthrough_due_to_unknown_stream_shape: verbose_proxy_logger.warning( "Presidio apply_to_output: streaming response was not a parsed chat completion stream " - "(raw non-Anthropic SSE passthrough or /v1/responses events). " + "(an error frame from an earlier guardrail, /v1/responses events, or a mixed stream). " "Output PII masking was skipped for this response." ) return @@ -1487,23 +1495,42 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): for chunk in all_chunks: yield chunk - async def _mask_anthropic_sse_stream( + async def _mask_raw_sse_stream( self, first_chunk: bytes, rest: AsyncIterator[object], request_data: dict ) -> tuple[object, ...]: + """The whole raw SSE stream masked as one response, or a raised refusal when it cannot be read. + + Raw frames are buffered to the end because PII can span frames, so no frame is forwarded + before the joined text was scanned. A stream whose surface is unknown, or that its surface's + assembler cannot rebuild, is withheld rather than forwarded unmasked. + """ rest_chunks: Final = [chunk async for chunk in rest] # mutable-ok: tuple() cannot consume an async iterator chunks: Final = (first_chunk, *rest_chunks) - assembled: Final = assemble_anthropic_sse_stream(chunks, restore_identity=True) - if assembled is None: - verbose_proxy_logger.warning( - "Presidio apply_to_output: raw SSE stream could not be assembled into a response. " - "Output PII masking was skipped for this response." + surface: Final = _assemble_raw_sse_stream(chunks) + if surface is None: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=( + "output PII masking cannot read this streaming response shape, " + "so the response was withheld instead of being forwarded unmasked" + ), + status_code=500, ) - return chunks - original_text: Final = model_response_text(assembled) + assembled, re_emit = surface + if assembled is None: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=( + "output PII masking could not assemble the streaming response, " + "so the response was withheld instead of being forwarded unmasked" + ), + status_code=500, + ) + before: Final = assembled.model_dump() await self._process_response_for_pii(response=assembled, request_data=request_data, mode="mask") - if model_response_text(assembled) == original_text: + if assembled.model_dump() == before: return chunks - return anthropic_sse_chunks_from_response(assembled) + return re_emit(assembled) @staticmethod def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: dict[str, str]) -> bytes: diff --git a/tests/integration/observability/test_presidio_streaming_output.py b/tests/integration/observability/test_presidio_streaming_output.py index 680617f1652..d5190664347 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: +def test_native_gemini_non_json_frame_in_a_stream_without_pii_is_replayed_unchanged( + gateway: Gateway, tmp_path: Path +) -> None: frames: Final = (b"data: not json at all\r\n\r\n", gemini_frame("after")) provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames)) with presidio_rig(gateway, tmp_path, provider) as rig: received: Final = rig.stream(rig.gemini_path(), rig.gemini_body()) assert received.status == 200, received.text assert received.text.replace("\r\n", "\n") == b"".join(frames).decode().replace("\r\n", "\n") + assert len(rig.analyzer.drain()) == 1 def test_native_gemini_empty_stream_returns_200_with_no_body(gateway: Gateway, tmp_path: Path) -> None: @@ -328,14 +355,14 @@ def test_native_gemini_empty_stream_returns_200_with_no_body(gateway: Gateway, t assert received.text == "" -def test_native_gemini_streams_while_presidio_analyzer_is_down(gateway: Gateway, tmp_path: Path) -> None: +def test_native_gemini_stream_fails_closed_when_analyzer_is_down(gateway: Gateway, tmp_path: Path) -> None: frames: Final = (gemini_frame(f"{PERSON} one. "), gemini_frame("two.")) provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames)) with presidio_rig(gateway, tmp_path, provider, analyze=broken) as rig: received: Final = rig.stream(rig.gemini_path(), rig.gemini_body()) - assert received.status == 200, received.text - assert gemini_texts(b"".join(received.frames)) == (f"{PERSON} one. ", "two.") - assert rig.analyzer.drain() == () + assert PERSON not in received.text, received.text + assert "Presidio analyzer" in received.text, received.text + assert rig.anonymizer.drain() == () def test_native_gemini_unauthenticated_request_is_rejected_before_upstream(gateway: Gateway, tmp_path: Path) -> None: @@ -526,17 +553,13 @@ def test_mixed_burst_survives_anonymizer_outage_and_recovers(gateway: Gateway, t def gemini_call(index: int) -> tuple[str, str, int]: received: Final = rig.stream(rig.gemini_path(), rig.gemini_body()) - return ( - "gemini", - f"g{index}", - received.status if gemini_texts(b"".join(received.frames)) == (f"{PERSON} ", "designed it.") else -1, - ) + masked: Final = gemini_texts(b"".join(received.frames)) == (f"{MASK} designed it.",) + return ("gemini", f"g{index}", -1 if PERSON in received.text else (1 if masked else 0)) def anthropic_call(index: int) -> tuple[str, str, int]: body: Final = {**rig.messages_body(), "messages": [{"role": "user", "content": f"a{index}"}]} received: Final = rig.stream("/v1/messages", body) - leaked: Final = PERSON in received.text - return ("anthropic", f"a{index}", -1 if leaked else (1 if MASK in received.text else 0)) + return ("anthropic", f"a{index}", -1 if PERSON in received.text else (1 if MASK in received.text else 0)) def phase(offset: int) -> tuple[tuple[str, str, int], ...]: with ThreadPoolExecutor(max_workers=12) as pool: @@ -552,13 +575,9 @@ def test_mixed_burst_survives_anonymizer_outage_and_recovers(gateway: Gateway, t outage.clear() healthy_after: Final = phase(200) - for name, results in (("before", healthy_before), ("during", during), ("after", healthy_after)): - assert all(status == 200 for kind, _, status in results if kind == "gemini"), (name, results) - assert all(status == 1 for kind, _, status in healthy_before + healthy_after if kind == "anthropic"), ( - healthy_before, - healthy_after, - ) - assert all(status == 0 for kind, _, status in during if kind == "anthropic"), during + for name, results in (("before", healthy_before), ("after", healthy_after)): + assert all(outcome == 1 for _, _, outcome in results), (name, results) + assert all(outcome == 0 for _, _, outcome in during), during identities: Final = tuple(identity for _, identity, _ in healthy_before + during + healthy_after) assert len(identities) == len(set(identities)) == 36 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 0a4ffbaef26..17bec2ecf36 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 @@ -2524,13 +2531,135 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_without_pii_are_for assert collected == byte_chunks -def _gemini_sse(text: str) -> bytes: +def _gemini_sse(text: str, terminator: bytes = b"\n\n") -> bytes: payload = {"candidates": [{"content": {"parts": [{"text": text}], "role": "model"}, "index": 0}]} - return f"data: {json.dumps(payload)}\n\n".encode() + return b"data: " + json.dumps(payload).encode() + terminator + + +def _gemini_texts(chunks: list[object]) -> list[str]: + frames = b"".join(chunk for chunk in chunks if isinstance(chunk, bytes)).decode() + return [ + json.loads(line[6:])["candidates"][0]["content"]["parts"][0]["text"] + for line in frames.splitlines() + if line.startswith("data: ") + ] + + +async def _collect_masked_output(guardrail: _OPTIONAL_PresidioPIIMasking, stream, collected: list[object]) -> None: + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + response=stream, + request_data={}, + ): + collected.append(chunk) @pytest.mark.asyncio -async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incrementally_until_upstream_aborts(): +async def test_apply_to_output_streaming_gemini_name_split_across_frames_is_masked_as_one_response(): + """ + The name only exists once the frames are joined, so a per-frame scan would + forward both halves unmasked. The analyzer and anonymizer are the in-process + fake, which finds a person only in two adjacent capitalized words. + """ + frames = [_gemini_sse("The architect was John"), _gemini_sse(" Smith, per the record.")] + + async def mock_stream(): + for frame in frames: + yield frame + + collected: list[object] = [] + async with TestServer(_fake_presidio_app()) as server: + guardrail = _OPTIONAL_PresidioPIIMasking( + apply_to_output=True, + presidio_analyzer_api_base=str(server.make_url("/")), + presidio_anonymizer_api_base=str(server.make_url("/")), + pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK}, + ) + await _collect_masked_output(guardrail, mock_stream(), collected) + await guardrail._close_http_session() + + assert "John Smith" not in b"".join(collected).decode() + assert _gemini_texts(collected) == ["The architect was , 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): + frames = [_gemini_sse("nothing personal ", terminator), _gemini_sse("in here.", terminator)] + + async def mock_stream(): + for frame in frames: + yield frame + + collected: list[object] = [] + async with TestServer(_fake_presidio_app()) as server: + guardrail = _OPTIONAL_PresidioPIIMasking( + apply_to_output=True, + presidio_analyzer_api_base=str(server.make_url("/")), + presidio_anonymizer_api_base=str(server.make_url("/")), + pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK}, + ) + await _collect_masked_output(guardrail, mock_stream(), collected) + await guardrail._close_http_session() + + assert collected == frames + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_gemini_upstream_abort_mid_stream_forwards_no_frame(): guardrail = _OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, @@ -2544,18 +2673,31 @@ async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incremen yield frame raise ConnectionError("upstream closed mid-stream") - async def collect() -> None: - async for chunk in guardrail.async_post_call_streaming_iterator_hook( - user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), - response=mock_stream(), - request_data={}, - ): - collected.append(chunk) - with pytest.raises(ConnectionError): - await collect() + await _collect_masked_output(guardrail, mock_stream(), collected) - assert collected == frames + assert collected == [] + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_gemini_sse_bytes_fail_closed_when_presidio_is_unreachable(): + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + presidio_analyzer_api_base="http://127.0.0.1:9", + presidio_anonymizer_api_base="http://127.0.0.1:9", + ) + frames = [_gemini_sse("Hello John Smith"), _gemini_sse(" from the record.")] + collected: list[object] = [] + + async def mock_stream(): + for frame in frames: + yield frame + + with pytest.raises(Exception, match="Presidio PII analysis failed"): + await _collect_masked_output(guardrail, mock_stream(), collected) + + assert collected == [] @pytest.mark.asyncio @@ -2846,7 +2988,7 @@ async def test_apply_to_output_streaming_comment_only_stream_is_forwarded_unchan @pytest.mark.asyncio -async def test_apply_to_output_streaming_gemini_first_frame_split_across_transport_chunks_streams_incrementally(): +async def test_apply_to_output_streaming_gemini_first_frame_split_across_transport_chunks_is_still_masked(): guardrail = _OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, @@ -2860,50 +3002,38 @@ async def test_apply_to_output_streaming_gemini_first_frame_split_across_transpo yield first[:20] yield first[20:] yield second - raise ConnectionError("upstream closed mid-stream") - async def collect() -> None: - async for chunk in guardrail.async_post_call_streaming_iterator_hook( - user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), - response=mock_stream(), - request_data={}, - ): - collected.append(chunk) + await _collect_masked_output(guardrail, mock_stream(), collected) - with pytest.raises(ConnectionError): - await collect() - - assert collected == [first, second] + assert "John Smith" not in b"".join(collected).decode() + assert _gemini_texts(collected) == [""] @pytest.mark.asyncio -async def test_apply_to_output_streaming_unterminated_first_frame_is_released_once_it_exceeds_the_cap(): +@pytest.mark.parametrize( + "frame", + [ + b"data: not json at all\n\n", + b'data: {"unexpected": true}\n\n', + b"data: " + b"x" * (70 * 1024) + b"\n", + ], + ids=["non-json", "unknown-surface", "unterminated"], +) +async def test_apply_to_output_streaming_unrecognized_raw_sse_stream_is_withheld(frame: bytes): guardrail = _OPTIONAL_PresidioPIIMasking( mock_testing=True, apply_to_output=True, mock_redacted_text={"text": ""}, ) - piece = b"data: " + b"x" * 1023 + b"\n" - pieces_to_cap = -(-(64 * 1024) // len(piece)) - released_at: list[int] = [] + collected: list[object] = [] async def mock_stream(): - for index in range(pieces_to_cap * 4): - if collected: - released_at.append(index) - yield piece + yield frame - collected: list[object] = [] - async for chunk in guardrail.async_post_call_streaming_iterator_hook( - user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), - response=mock_stream(), - request_data={}, - ): - collected.append(chunk) + with pytest.raises(GuardrailRaisedException, match="cannot read this streaming response shape"): + await _collect_masked_output(guardrail, mock_stream(), collected) - assert released_at, "nothing reached the caller before the upstream finished" - assert released_at[0] == pieces_to_cap, released_at[:3] - assert b"".join(collected) == piece * (pieces_to_cap * 4) + assert collected == [] @pytest.mark.asyncio From 95cf66096138ccece15f9f4135f481f2ac5cd6bc Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 15:52:07 -0700 Subject: [PATCH 2/5] fix(guardrails): rewrite masked Gemini SSE frames in place and withhold streams with unreadable frames --- litellm/proxy/guardrails/anthropic_sse.py | 40 ++- litellm/proxy/guardrails/gemini_sse.py | 212 ++++++++++++-- .../guardrails/guardrail_hooks/presidio.py | 95 ++++--- .../test_presidio_streaming_output.py | 14 +- .../guardrail_hooks/test_presidio.py | 260 +++++++++++++++++- 5 files changed, 535 insertions(+), 86 deletions(-) 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 == [] From 9e683e86f312f61eb0dc4ed3a565743948f7dd31 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 15:54:45 -0700 Subject: [PATCH 3/5] test(guardrails): cover masking of every Gemini candidate merged by index --- .../guardrail_hooks/test_presidio.py | 33 +++++++++++++++++++ 1 file changed, 33 insertions(+) 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 9da3048a728..41fcf6806a0 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -2833,6 +2833,39 @@ async def test_apply_to_output_streaming_gemini_error_frame_mid_stream_is_carrie assert frame["error"] == {"code": 503, "message": "overloaded", "status": "UNAVAILABLE"} +@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( From 2a9acc692e9f3c311b49557d0637d007e795c403 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 16:22:26 -0700 Subject: [PATCH 4/5] fix(guardrails): keep a Gemini upstream error frame as its own terminal frame after masking --- litellm/proxy/guardrails/gemini_sse.py | 28 ++++++++++++++----- .../guardrail_hooks/test_presidio.py | 13 +++++---- 2 files changed, 28 insertions(+), 13 deletions(-) diff --git a/litellm/proxy/guardrails/gemini_sse.py b/litellm/proxy/guardrails/gemini_sse.py index ff41df3c9b0..177880004b4 100644 --- a/litellm/proxy/guardrails/gemini_sse.py +++ b/litellm/proxy/guardrails/gemini_sse.py @@ -3,10 +3,11 @@ 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. 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. +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 @@ -63,7 +64,8 @@ async def mask_gemini_sse_stream(all_chunks: Sequence[object], mask_text: TextMa 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. + 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: @@ -76,8 +78,20 @@ async def mask_gemini_sse_stream(all_chunks: Sequence[object], mask_text: TextMa 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(),)) + 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: 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 41fcf6806a0..2dc4ff53bce 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -2808,10 +2808,10 @@ async def test_apply_to_output_streaming_gemini_masked_function_call_arguments_t @pytest.mark.asyncio -async def test_apply_to_output_streaming_gemini_error_frame_mid_stream_is_carried_on_the_masked_frame(): +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 the - client reads the error off the last frame, so the masked re-emit keeps it. + 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."}), @@ -2828,9 +2828,10 @@ async def test_apply_to_output_streaming_gemini_error_frame_mid_stream_is_carrie 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"} + 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 From d74f5ddb1d3cc209f77b28b5a22ac7db95c7bf95 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 17:19:49 -0700 Subject: [PATCH 5/5] fix(guardrails): classify a raw SSE stream as error-only on the whole drained stream, not its first frame An error frame ahead of Gemini or Anthropic content frames switched masking off for every frame behind it. The refusal check now runs on the drained stream, so an error-only stream is still forwarded as it arrived and content behind a leading error frame is masked --- .../guardrails/guardrail_hooks/presidio.py | 15 +++--- .../guardrail_hooks/test_presidio.py | 53 +++++++++++++++++++ 2 files changed, 61 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index ee10dc91c1d..9cc5bfcf5b6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -1446,10 +1446,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): elif isinstance(chunk, _SsePreface): yield chunk.raw elif isinstance(chunk, bytes): - if all_chunks or passthrough_due_to_unknown_stream_shape or is_sse_error_stream((chunk,)): - passthrough_due_to_unknown_stream_shape = ( - passthrough_due_to_unknown_stream_shape or not all_chunks - ) + if all_chunks or passthrough_due_to_unknown_stream_shape: yield chunk continue for masked_chunk in await self._mask_raw_sse_stream(chunk, stream, request_data): @@ -1492,14 +1489,18 @@ 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 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. + 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): 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 2dc4ff53bce..a23eb0805bf 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -2834,6 +2834,59 @@ async def test_apply_to_output_streaming_gemini_error_frame_mid_stream_keeps_its 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(): """