mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): rewrite masked Gemini SSE frames in place and withhold streams with unreadable frames
This commit is contained in:
parent
a2efc1c321
commit
95cf660961
5 changed files with 535 additions and 86 deletions
|
|
@ -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(
|
||||
(
|
||||
|
|
|
|||
|
|
@ -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})})
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -2545,6 +2545,29 @@ def _gemini_texts(chunks: list[object]) -> list[str]:
|
|||
]
|
||||
|
||||
|
||||
def _gemini_frame(*parts: dict, finish_reason: str | None = None, **top_level: object) -> bytes:
|
||||
candidate = {"content": {"parts": list(parts), "role": "model"}, "index": 0}
|
||||
payload = {
|
||||
"candidates": [candidate if finish_reason is None else {**candidate, "finishReason": finish_reason}],
|
||||
**top_level,
|
||||
}
|
||||
return b"data: " + json.dumps(payload).encode() + b"\n\n"
|
||||
|
||||
|
||||
def _gemini_frames(chunks: list[object]) -> list[dict]:
|
||||
frames = b"".join(chunk for chunk in chunks if isinstance(chunk, bytes)).decode()
|
||||
return [json.loads(line[6:]) for line in frames.splitlines() if line.startswith("data: ")]
|
||||
|
||||
|
||||
def _gemini_fake_masking_guardrail(server: TestServer) -> _OPTIONAL_PresidioPIIMasking:
|
||||
return _OPTIONAL_PresidioPIIMasking(
|
||||
apply_to_output=True,
|
||||
presidio_analyzer_api_base=str(server.make_url("/")),
|
||||
presidio_anonymizer_api_base=str(server.make_url("/")),
|
||||
pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK},
|
||||
)
|
||||
|
||||
|
||||
async def _collect_masked_output(guardrail: _OPTIONAL_PresidioPIIMasking, stream, collected: list[object]) -> None:
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
|
||||
|
|
@ -2638,7 +2661,17 @@ async def test_apply_to_output_streaming_anthropic_tool_call_arguments_only_mask
|
|||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("terminator", [b"\n\n", b"\r\n\r\n"])
|
||||
async def test_apply_to_output_streaming_gemini_stream_without_pii_is_replayed_byte_for_byte(terminator):
|
||||
frames = [_gemini_sse("nothing personal ", terminator), _gemini_sse("in here.", terminator)]
|
||||
function_call = {
|
||||
"functionCall": {"id": "call_1", "name": "lookup", "args": {"city": "Paris"}},
|
||||
"thoughtSignature": "sig",
|
||||
}
|
||||
frames = [
|
||||
_gemini_sse("nothing personal ", terminator),
|
||||
_gemini_sse("in here.", terminator),
|
||||
b"data: "
|
||||
+ json.dumps({"candidates": [{"content": {"parts": [function_call], "role": "model"}}]}).encode()
|
||||
+ terminator,
|
||||
]
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
|
|
@ -2658,6 +2691,148 @@ async def test_apply_to_output_streaming_gemini_stream_without_pii_is_replayed_b
|
|||
assert collected == frames
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_function_call_keeps_its_signature_and_id_while_its_arguments_are_masked():
|
||||
"""
|
||||
A Gemini client echoes the streamed model turn on its next request, and Gemini 3 rejects
|
||||
a function call whose thought signature is missing, so masking has to rewrite the frames
|
||||
in place rather than rebuild them from a chat-completions response that drops those fields.
|
||||
"""
|
||||
frames = [
|
||||
_gemini_frame(
|
||||
{"text": "emailing John"}, usageMetadata={"totalTokenCount": 1}, modelVersion="m", responseId="r1"
|
||||
),
|
||||
_gemini_frame({"text": " Smith now."}, usageMetadata={"totalTokenCount": 2}, modelVersion="m", responseId="r1"),
|
||||
_gemini_frame(
|
||||
{
|
||||
"functionCall": {"id": "call_1", "name": "send_email", "args": {"to": "John Smith", "urgent": True}},
|
||||
"thoughtSignature": "sig-fc",
|
||||
},
|
||||
usageMetadata={"totalTokenCount": 3},
|
||||
modelVersion="m",
|
||||
responseId="r1",
|
||||
),
|
||||
_gemini_frame(
|
||||
{"text": ""}, finish_reason="STOP", usageMetadata={"totalTokenCount": 9}, modelVersion="m", responseId="r1"
|
||||
),
|
||||
]
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
collected: list[object] = []
|
||||
async with TestServer(_fake_presidio_app()) as server:
|
||||
guardrail = _gemini_fake_masking_guardrail(server)
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
await guardrail._close_http_session()
|
||||
|
||||
assert "John Smith" not in b"".join(collected).decode()
|
||||
(body,) = _gemini_frames(collected)
|
||||
assert body["candidates"] == [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{"text": "emailing <PERSON> now."},
|
||||
{
|
||||
"functionCall": {
|
||||
"id": "call_1",
|
||||
"name": "send_email",
|
||||
"args": {"to": "<PERSON>", "urgent": True},
|
||||
},
|
||||
"thoughtSignature": "sig-fc",
|
||||
},
|
||||
],
|
||||
"role": "model",
|
||||
},
|
||||
"index": 0,
|
||||
"finishReason": "STOP",
|
||||
}
|
||||
]
|
||||
assert body["usageMetadata"] == {"totalTokenCount": 9}
|
||||
assert body["modelVersion"] == "m"
|
||||
assert body["responseId"] == "r1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_thought_part_and_trailing_signature_survive_the_merge():
|
||||
"""
|
||||
A thought part is never merged into the visible text that follows it, and the signature
|
||||
Gemini sends on a trailing empty text part lands on the text it signs.
|
||||
"""
|
||||
frames = [
|
||||
_gemini_frame({"text": "thinking about John Smith", "thought": True}),
|
||||
_gemini_frame({"text": "hello John"}),
|
||||
_gemini_frame({"text": " Smith,"}),
|
||||
_gemini_frame({"text": "", "thoughtSignature": "sig-text"}, finish_reason="STOP"),
|
||||
]
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
collected: list[object] = []
|
||||
async with TestServer(_fake_presidio_app()) as server:
|
||||
guardrail = _gemini_fake_masking_guardrail(server)
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
await guardrail._close_http_session()
|
||||
|
||||
(body,) = _gemini_frames(collected)
|
||||
assert body["candidates"][0]["content"]["parts"] == [
|
||||
{"text": "thinking about <PERSON>", "thought": True},
|
||||
{"text": "hello <PERSON>,", "thoughtSignature": "sig-text"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_masked_function_call_arguments_that_no_longer_parse_are_withheld():
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
mock_redacted_text={"text": "<PERSON>"},
|
||||
)
|
||||
frames = [
|
||||
_gemini_frame({"functionCall": {"name": "lookup", "args": {"person": "John Smith"}}}, finish_reason="STOP")
|
||||
]
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
with pytest.raises(GuardrailRaisedException) as raised:
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
assert raised.value.status_code == 500
|
||||
assert collected == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_error_frame_mid_stream_is_carried_on_the_masked_frame():
|
||||
"""
|
||||
Gemini can end a stream with an error frame after content frames, and the
|
||||
client reads the error off the last frame, so the masked re-emit keeps it.
|
||||
"""
|
||||
frames = [
|
||||
_gemini_frame({"text": "The architect was John Smith."}),
|
||||
b'data: {"error": {"code": 503, "message": "overloaded", "status": "UNAVAILABLE"}}\n\n',
|
||||
]
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
async with TestServer(_fake_presidio_app()) as server:
|
||||
guardrail = _gemini_fake_masking_guardrail(server)
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
await guardrail._close_http_session()
|
||||
|
||||
(frame,) = _gemini_frames(collected)
|
||||
assert frame["candidates"][0]["content"]["parts"] == [{"text": "The architect was <PERSON>."}]
|
||||
assert frame["error"] == {"code": 503, "message": "overloaded", "status": "UNAVAILABLE"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_upstream_abort_mid_stream_forwards_no_frame():
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
|
|
@ -2925,13 +3100,13 @@ async def test_apply_to_output_streaming_leading_keepalive_is_forwarded_before_u
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_leading_comments_over_the_frame_cap_are_still_masked():
|
||||
async def test_apply_to_output_streaming_leading_comments_of_any_size_are_still_masked():
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
mock_redacted_text={"text": "<PERSON>"},
|
||||
)
|
||||
keepalives = [b": keepalive\n\n" * 512] * 12 # ~72 KiB of complete comment frames, over the 64 KiB cap
|
||||
keepalives = [b": keepalive\n\n" * 512] * 12 # ~72 KiB of complete comment frames
|
||||
byte_chunks = [
|
||||
*keepalives[:-1],
|
||||
keepalives[-1]
|
||||
|
|
@ -3011,15 +3186,16 @@ async def test_apply_to_output_streaming_gemini_first_frame_split_across_transpo
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"frame",
|
||||
("frame", "reason"),
|
||||
[
|
||||
b"data: not json at all\n\n",
|
||||
b'data: {"unexpected": true}\n\n',
|
||||
b"data: " + b"x" * (70 * 1024) + b"\n",
|
||||
(b"data: not json at all\n\n", "could not read every streamed frame"),
|
||||
(b"data: 42\n\n", "could not read every streamed frame"),
|
||||
(b'data: {"unexpected": true}\n\n', "cannot read this streaming response shape"),
|
||||
(b"data: " + b"x" * (70 * 1024) + b"\n", "could not read every streamed frame"),
|
||||
],
|
||||
ids=["non-json", "unknown-surface", "unterminated"],
|
||||
ids=["non-json", "json-scalar", "unknown-surface", "unterminated"],
|
||||
)
|
||||
async def test_apply_to_output_streaming_unrecognized_raw_sse_stream_is_withheld(frame: bytes):
|
||||
async def test_apply_to_output_streaming_unrecognized_raw_sse_stream_is_withheld(frame: bytes, reason: str):
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
|
|
@ -3030,7 +3206,71 @@ async def test_apply_to_output_streaming_unrecognized_raw_sse_stream_is_withheld
|
|||
async def mock_stream():
|
||||
yield frame
|
||||
|
||||
with pytest.raises(GuardrailRaisedException, match="cannot read this streaming response shape"):
|
||||
with pytest.raises(GuardrailRaisedException, match=reason):
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
assert collected == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_stream_with_an_unreadable_frame_is_withheld():
|
||||
"""
|
||||
The parser skips a frame it cannot read, and replaying the originals around it
|
||||
would forward that frame unscanned, so the whole stream is withheld instead.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
mock_redacted_text={"text": "<PERSON>"},
|
||||
)
|
||||
frames = [
|
||||
_gemini_frame({"text": "The architect was"}),
|
||||
b"data: not json at all\n\n",
|
||||
_gemini_frame({"text": " Ada."}, finish_reason="STOP"),
|
||||
]
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
with pytest.raises(GuardrailRaisedException, match="could not read every streamed frame"):
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
assert collected == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_anthropic_stream_with_an_unreadable_frame_is_withheld():
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
mock_redacted_text={"text": "Hello"},
|
||||
)
|
||||
byte_chunks = [
|
||||
_anthropic_sse(
|
||||
"message_start",
|
||||
{"type": "message_start", "message": {"id": "msg_1", "model": "claude", "content": [], "usage": {}}},
|
||||
),
|
||||
_anthropic_sse(
|
||||
"content_block_start",
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
|
||||
),
|
||||
b"event: content_block_delta\ndata: not json at all\n\n",
|
||||
_anthropic_sse(
|
||||
"content_block_delta",
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "Hello"}},
|
||||
),
|
||||
_anthropic_sse("content_block_stop", {"type": "content_block_stop", "index": 0}),
|
||||
_anthropic_sse("message_stop", {"type": "message_stop"}),
|
||||
]
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for chunk in byte_chunks:
|
||||
yield chunk
|
||||
|
||||
with pytest.raises(GuardrailRaisedException, match="could not read every streamed frame"):
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
assert collected == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue