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