mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge d74f5ddb1d into b781d157d7
This commit is contained in:
commit
a7fa010ec3
5 changed files with 950 additions and 155 deletions
|
|
@ -2,7 +2,8 @@
|
|||
|
||||
`/v1/messages` streams reach a guardrail's `async_post_call_streaming_iterator_hook` as raw SSE
|
||||
frames rather than chunk objects, which `stream_chunk_builder` cannot assemble. These helpers let a
|
||||
hook scan such a stream, and re-emit it when the guardrail rewrote the response.
|
||||
hook scan such a stream, and re-emit it when the guardrail rewrote the response. The raw SSE
|
||||
parsing (`joined_sse_stream`, `parsed_sse_events`) is shared with the Gemini sibling module.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -12,6 +13,8 @@ import json
|
|||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm.types.utils import Choices, ModelResponse
|
||||
|
||||
_ANTHROPIC_EVENT_TYPES: Final = frozenset(
|
||||
|
|
@ -26,25 +29,29 @@ _ANTHROPIC_EVENT_TYPES: Final = frozenset(
|
|||
"error",
|
||||
}
|
||||
)
|
||||
_CONTENT_FREE_PAYLOADS: Final = frozenset({"", "[DONE]"})
|
||||
_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def is_raw_sse_stream(all_chunks: Sequence[object]) -> bool:
|
||||
return any(isinstance(chunk, (str, bytes)) for chunk in all_chunks)
|
||||
|
||||
|
||||
def _joined_sse_stream(all_chunks: Sequence[object]) -> str | None:
|
||||
def joined_sse_stream(all_chunks: Sequence[object]) -> str | None:
|
||||
"""The raw frames decoded as one text, with every SSE line ending (CRLF, LF, CR) folded to LF."""
|
||||
raw: Final = b"".join(
|
||||
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8")
|
||||
for chunk in all_chunks
|
||||
if isinstance(chunk, (str, bytes))
|
||||
)
|
||||
try:
|
||||
return codecs.getincrementaldecoder("utf-8")().decode(raw, final=False)
|
||||
decoded: Final = codecs.getincrementaldecoder("utf-8")().decode(raw, final=False)
|
||||
except UnicodeDecodeError:
|
||||
return None
|
||||
return decoded.replace("\r\n", "\n").replace("\r", "\n")
|
||||
|
||||
|
||||
def _parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]:
|
||||
def parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]:
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
)
|
||||
|
|
@ -52,15 +59,49 @@ def _parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]:
|
|||
return tuple(
|
||||
event_data
|
||||
for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses
|
||||
if (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing
|
||||
if isinstance(event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event), Mapping) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing
|
||||
)
|
||||
|
||||
|
||||
def _data_payload(event: str) -> str | None:
|
||||
lines: Final = tuple(line.strip() for line in event.splitlines())
|
||||
return next((line[len("data:") :].strip() for line in lines if line.startswith("data:")), None)
|
||||
|
||||
|
||||
def _is_unreadable_payload(payload: str) -> bool:
|
||||
if payload in _CONTENT_FREE_PAYLOADS:
|
||||
return False
|
||||
try:
|
||||
_JSON_OBJECT.validate_json(payload)
|
||||
except ValidationError:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def has_unreadable_sse_frames(all_chunks: Sequence[object]) -> bool:
|
||||
"""Whether a ``data:`` payload is neither blank, ``[DONE]``, nor a JSON object.
|
||||
|
||||
``parsed_sse_events`` drops such a frame silently, which suits the assemblers and not a
|
||||
masking hook: a frame it cannot read is one it cannot scan, so the hook withholds the stream
|
||||
instead of replaying it.
|
||||
"""
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
)
|
||||
|
||||
sse_stream: Final = joined_sse_stream(all_chunks)
|
||||
if sse_stream is None:
|
||||
return True
|
||||
events: Final = AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same splitter the assembler uses
|
||||
payloads: Final = tuple(_data_payload(event) for event in events)
|
||||
return any(_is_unreadable_payload(payload) for payload in payloads if payload is not None)
|
||||
|
||||
|
||||
def _anthropic_message_start(sse_stream: str) -> Mapping[str, object] | None:
|
||||
return next(
|
||||
(
|
||||
message
|
||||
for event_data in _parsed_sse_events(sse_stream)
|
||||
for event_data in parsed_sse_events(sse_stream)
|
||||
if event_data.get("type") == "message_start" and isinstance(message := event_data.get("message"), dict)
|
||||
),
|
||||
None,
|
||||
|
|
@ -75,10 +116,10 @@ def is_anthropic_sse_stream(all_chunks: Sequence[object]) -> bool:
|
|||
stream raw too. Reading its frames as Anthropic ones would refuse the response in a wire format
|
||||
its client cannot parse, so the surface is decided on the event types actually present.
|
||||
"""
|
||||
sse_stream: Final = _joined_sse_stream(all_chunks)
|
||||
sse_stream: Final = joined_sse_stream(all_chunks)
|
||||
if sse_stream is None:
|
||||
return False
|
||||
return any(event.get("type") in _ANTHROPIC_EVENT_TYPES for event in _parsed_sse_events(sse_stream))
|
||||
return any(event.get("type") in _ANTHROPIC_EVENT_TYPES for event in parsed_sse_events(sse_stream))
|
||||
|
||||
|
||||
def assemble_anthropic_sse_stream(
|
||||
|
|
@ -95,7 +136,7 @@ def assemble_anthropic_sse_stream(
|
|||
AnthropicPassthroughLoggingHandler,
|
||||
)
|
||||
|
||||
sse_stream: Final = _joined_sse_stream(all_chunks)
|
||||
sse_stream: Final = joined_sse_stream(all_chunks)
|
||||
if sse_stream is None:
|
||||
return None
|
||||
message_start: Final = _anthropic_message_start(sse_stream)
|
||||
|
|
@ -157,10 +198,10 @@ def is_sse_error_stream(all_chunks: Sequence[object]) -> bool:
|
|||
# A stream mixing typed chunks with an error frame still carries content to scan, and the
|
||||
# frames-only join below would drop exactly the part that has to be scanned
|
||||
return False
|
||||
sse_stream: Final = _joined_sse_stream(all_chunks)
|
||||
sse_stream: Final = joined_sse_stream(all_chunks)
|
||||
if sse_stream is None:
|
||||
return False
|
||||
events: Final = _parsed_sse_events(sse_stream)
|
||||
events: Final = parsed_sse_events(sse_stream)
|
||||
return len(events) > 0 and all(
|
||||
event.get("type") == "error" or isinstance(event.get("error"), Mapping) for event in events
|
||||
)
|
||||
|
|
|
|||
235
litellm/proxy/guardrails/gemini_sse.py
Normal file
235
litellm/proxy/guardrails/gemini_sse.py
Normal file
|
|
@ -0,0 +1,235 @@
|
|||
"""Gemini SSE frame masking for guardrail streaming hooks.
|
||||
|
||||
The Google ``:streamGenerateContent`` route relays the upstream ``data:`` frames as raw bytes, so a
|
||||
guardrail's ``async_post_call_streaming_iterator_hook`` cannot read them as chunk objects. These
|
||||
helpers fold such a stream into the one response body a non-streaming call would have returned,
|
||||
rewrite its text and function-call arguments in place, and re-emit it as one frame, followed by any
|
||||
upstream error frame kept as its own terminal frame. Every other field (thought signatures,
|
||||
function-call ids, model version, response id, safety ratings, usage) is carried through untouched,
|
||||
because a client echoes the model turn back on its next request and Gemini 3 rejects a function
|
||||
call whose thought signature is missing.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable, Iterable, Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from itertools import accumulate, chain, groupby
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm.proxy.guardrails.anthropic_sse import joined_sse_stream, parsed_sse_events
|
||||
|
||||
_GEMINI_RESPONSE_KEYS: Final = frozenset({"candidates", "usageMetadata", "promptFeedback"})
|
||||
_TEXT_RUN_PART_KEYS: Final = frozenset({"text", "thought", "thoughtSignature"})
|
||||
_EMPTY_UNSIGNED_TEXT_PART: Final = MappingProxyType({"text": ""})
|
||||
_UNREADABLE_ARGUMENTS: Final = "left the function call arguments unreadable as JSON after masking"
|
||||
|
||||
TextMasker: TypeAlias = Callable[[str], Awaitable[str]] # mutable-ok: Callable params
|
||||
_JsonObject: TypeAlias = Mapping[str, object]
|
||||
_JSON_OBJECT: Final = TypeAdapter(_JsonObject)
|
||||
_JSON_ARRAY: Final = TypeAdapter(tuple[object, ...])
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GeminiStreamUnchanged:
|
||||
"""Masking rewrote nothing, so the original frames are replayed as they arrived."""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GeminiStreamMasked:
|
||||
frames: tuple[bytes, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GeminiStreamUnreadable:
|
||||
reason: str
|
||||
|
||||
|
||||
GeminiStreamMaskResult: TypeAlias = GeminiStreamUnchanged | GeminiStreamMasked | GeminiStreamUnreadable
|
||||
|
||||
|
||||
def is_gemini_sse_stream(all_chunks: Sequence[object]) -> bool:
|
||||
sse_stream: Final = joined_sse_stream(all_chunks)
|
||||
if sse_stream is None:
|
||||
return False
|
||||
return any(not _GEMINI_RESPONSE_KEYS.isdisjoint(event) for event in parsed_sse_events(sse_stream))
|
||||
|
||||
|
||||
async def mask_gemini_sse_stream(all_chunks: Sequence[object], mask_text: TextMasker) -> GeminiStreamMaskResult:
|
||||
"""The stream's frames folded into one response, with every text and function-call argument masked.
|
||||
|
||||
Text arrives split across frames, so the fragments of one text part are joined before they are
|
||||
scanned; PII that only exists once the fragments meet is otherwise forwarded in halves. Top-level
|
||||
keys and candidate keys take the last frame's value, and candidates are merged by index. An
|
||||
upstream error frame stays its own terminal frame after the masked response, as it arrived.
|
||||
"""
|
||||
sse_stream: Final = joined_sse_stream(all_chunks)
|
||||
if sse_stream is None:
|
||||
return GeminiStreamUnreadable("could not decode the streaming response as UTF-8")
|
||||
events: Final = parsed_sse_events(sse_stream)
|
||||
candidates: Final = _merged_candidates(tuple(_candidates_of(events)))
|
||||
masked_candidates: Final = tuple([await _masked_candidate(candidate, mask_text) for candidate in candidates])
|
||||
unreadable: Final = next((item for item in masked_candidates if isinstance(item, GeminiStreamUnreadable)), None)
|
||||
if unreadable is not None:
|
||||
return unreadable
|
||||
if masked_candidates == candidates:
|
||||
return GeminiStreamUnchanged()
|
||||
response: Final = MappingProxyType({**_without_error(_later_wins(events)), "candidates": masked_candidates})
|
||||
return GeminiStreamMasked((_frame(response), *map(_frame, _error_frames(events))))
|
||||
|
||||
|
||||
def _frame(event: _JsonObject) -> bytes:
|
||||
return f"data: {_json_text(event)}\n\n".encode()
|
||||
|
||||
|
||||
def _without_error(event: _JsonObject) -> _JsonObject:
|
||||
return MappingProxyType({key: value for key, value in event.items() if key != "error"})
|
||||
|
||||
|
||||
def _error_frames(events: Sequence[_JsonObject]) -> Iterator[_JsonObject]:
|
||||
return (MappingProxyType({"error": event["error"]}) for event in events if "error" in event)
|
||||
|
||||
|
||||
def _json_text(value: object) -> str:
|
||||
return json.dumps(value, ensure_ascii=False, default=dict)
|
||||
|
||||
|
||||
def _json_object(value: object) -> _JsonObject | None:
|
||||
try:
|
||||
return _JSON_OBJECT.validate_python(value)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _json_objects(values: Iterable[object]) -> Iterator[_JsonObject]:
|
||||
return (parsed for parsed in map(_json_object, values) if parsed is not None)
|
||||
|
||||
|
||||
def _json_array(value: object) -> tuple[object, ...]:
|
||||
try:
|
||||
return _JSON_ARRAY.validate_python(value)
|
||||
except ValidationError:
|
||||
return ()
|
||||
|
||||
|
||||
def _later_wins(mappings: Sequence[_JsonObject]) -> _JsonObject:
|
||||
return MappingProxyType(dict(chain.from_iterable(mapping.items() for mapping in mappings)))
|
||||
|
||||
|
||||
def _candidates_of(events: Sequence[_JsonObject]) -> Iterator[_JsonObject]:
|
||||
for event in events:
|
||||
yield from _json_objects(_json_array(event.get("candidates")))
|
||||
|
||||
|
||||
def _merged_candidates(candidates: Sequence[_JsonObject]) -> tuple[_JsonObject, ...]:
|
||||
indexes: Final = tuple(dict.fromkeys(_candidate_index(candidate) for candidate in candidates))
|
||||
return tuple(
|
||||
_merged_candidate(tuple(candidate for candidate in candidates if _candidate_index(candidate) == index))
|
||||
for index in indexes
|
||||
)
|
||||
|
||||
|
||||
def _candidate_index(candidate: _JsonObject) -> int:
|
||||
index: Final = candidate.get("index")
|
||||
return index if isinstance(index, int) else 0
|
||||
|
||||
|
||||
def _merged_candidate(fragments: Sequence[_JsonObject]) -> _JsonObject:
|
||||
merged: Final = _later_wins(fragments)
|
||||
contents: Final = tuple(_json_objects(fragment.get("content") for fragment in fragments))
|
||||
if not contents:
|
||||
return merged
|
||||
content: Final = MappingProxyType({**_later_wins(contents), "parts": _merged_parts(tuple(_parts_of(contents)))})
|
||||
return MappingProxyType({**merged, "content": content})
|
||||
|
||||
|
||||
def _parts_of(contents: Sequence[_JsonObject]) -> Iterator[object]:
|
||||
for content in contents:
|
||||
yield from _json_array(content.get("parts"))
|
||||
|
||||
|
||||
def _merged_parts(parts: Sequence[object]) -> tuple[object, ...]:
|
||||
"""Adjacent fragments of one text part joined into that part; every other part kept as it came.
|
||||
|
||||
The empty unsigned text part Gemini streams on its final frame is dropped, since the
|
||||
non-streaming body carries none and Gemini rejects it when a client echoes it back.
|
||||
"""
|
||||
run_ids: Final = tuple(run_id for run_id, _ in accumulate(parts, _next_run, initial=(0, None)))[1:]
|
||||
merged: Final = tuple(
|
||||
_merged_run(tuple(part for part, _ in run)) for _, run in groupby(zip(parts, run_ids), key=lambda pair: pair[1])
|
||||
)
|
||||
return tuple(part for part in merged if part != _EMPTY_UNSIGNED_TEXT_PART)
|
||||
|
||||
|
||||
def _next_run(state: tuple[int, object], part: object) -> tuple[int, object]:
|
||||
run_id, previous = state
|
||||
return (run_id if _continues_text_run(previous, part) else run_id + 1, part)
|
||||
|
||||
|
||||
def _continues_text_run(previous: object, part: object) -> bool:
|
||||
"""Whether ``part`` is the next fragment of the text part ``previous`` started.
|
||||
|
||||
A fragment carrying a thought signature closes its run, since the signature belongs to the
|
||||
part it arrived on and one merged part can only carry one.
|
||||
"""
|
||||
previous_fragment: Final = _text_fragment(previous)
|
||||
fragment: Final = _text_fragment(part)
|
||||
if previous_fragment is None or fragment is None:
|
||||
return False
|
||||
return (
|
||||
bool(previous_fragment.get("thought")) == bool(fragment.get("thought"))
|
||||
and "thoughtSignature" not in previous_fragment
|
||||
)
|
||||
|
||||
|
||||
def _text_fragment(part: object) -> _JsonObject | None:
|
||||
fragment: Final = _json_object(part)
|
||||
if fragment is None or not _TEXT_RUN_PART_KEYS.issuperset(fragment) or not isinstance(fragment.get("text"), str):
|
||||
return None
|
||||
return fragment
|
||||
|
||||
|
||||
def _merged_run(run: Sequence[object]) -> object:
|
||||
"""A run longer than one part holds text fragments only, since nothing else continues a run."""
|
||||
if len(run) == 1:
|
||||
return run[0]
|
||||
fragments: Final = tuple(_json_objects(run))
|
||||
text: Final = "".join(str(fragment["text"]) for fragment in fragments)
|
||||
return MappingProxyType({**fragments[0], **fragments[-1], "text": text})
|
||||
|
||||
|
||||
async def _masked_candidate(candidate: _JsonObject, mask_text: TextMasker) -> _JsonObject | GeminiStreamUnreadable:
|
||||
content: Final = _json_object(candidate.get("content"))
|
||||
if content is None:
|
||||
return candidate
|
||||
parts: Final = _json_array(content.get("parts"))
|
||||
masked_parts: Final = tuple([await _masked_part(part, mask_text) for part in parts])
|
||||
unreadable: Final = next((item for item in masked_parts if isinstance(item, GeminiStreamUnreadable)), None)
|
||||
if unreadable is not None:
|
||||
return unreadable
|
||||
return MappingProxyType({**candidate, "content": MappingProxyType({**content, "parts": masked_parts})})
|
||||
|
||||
|
||||
async def _masked_part(part: object, mask_text: TextMasker) -> object | GeminiStreamUnreadable:
|
||||
fragment: Final = _json_object(part)
|
||||
if fragment is None:
|
||||
return part
|
||||
text: Final = fragment.get("text")
|
||||
if isinstance(text, str):
|
||||
return MappingProxyType({**fragment, "text": await mask_text(text)}) if text else part
|
||||
function_call: Final = _json_object(fragment.get("functionCall"))
|
||||
if function_call is None:
|
||||
return part
|
||||
arguments: Final = _json_object(function_call.get("args"))
|
||||
if arguments is None:
|
||||
return part
|
||||
masked_arguments: Final = await mask_text(_json_text(arguments))
|
||||
try:
|
||||
parsed_arguments: Final = _JSON_OBJECT.validate_json(masked_arguments)
|
||||
except ValidationError:
|
||||
return GeminiStreamUnreadable(_UNREADABLE_ARGUMENTS)
|
||||
return MappingProxyType({**fragment, "functionCall": MappingProxyType({**function_call, "args": parsed_arguments})})
|
||||
|
|
@ -19,7 +19,7 @@ from datetime import datetime
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast
|
||||
|
||||
import aiohttp
|
||||
from typing_extensions import NotRequired, ReadOnly
|
||||
from typing_extensions import NotRequired, ReadOnly, assert_never
|
||||
|
||||
import litellm
|
||||
from litellm import get_secret
|
||||
|
|
@ -45,8 +45,16 @@ from litellm.proxy.common_utils.sse_keepalive import split_complete_sse_frames
|
|||
from litellm.proxy.guardrails.anthropic_sse import (
|
||||
anthropic_sse_chunks_from_response,
|
||||
assemble_anthropic_sse_stream,
|
||||
has_unreadable_sse_frames,
|
||||
is_anthropic_sse_stream,
|
||||
model_response_text,
|
||||
is_sse_error_stream,
|
||||
)
|
||||
from litellm.proxy.guardrails.gemini_sse import (
|
||||
GeminiStreamMasked,
|
||||
GeminiStreamUnchanged,
|
||||
GeminiStreamUnreadable,
|
||||
is_gemini_sse_stream,
|
||||
mask_gemini_sse_stream,
|
||||
)
|
||||
from litellm.types.guardrails import (
|
||||
GuardrailEventHooks,
|
||||
|
|
@ -97,9 +105,6 @@ def _json_escaped_len(text: str) -> int:
|
|||
return len(json.dumps(text).encode("utf-8")) - 2 # strip the surrounding quotes
|
||||
|
||||
|
||||
_MAX_FIRST_SSE_FRAME_BYTES: Final = 64 * 1024
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _SsePreface:
|
||||
"""Complete leading SSE frames with no ``data:`` line, relayed verbatim before the stream shape is decided."""
|
||||
|
|
@ -139,8 +144,7 @@ async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGener
|
|||
``data:`` line) as they complete, and join raw ``bytes`` chunks until they
|
||||
hold one complete SSE event with a data line, so the stream shape is
|
||||
decided on a whole frame rather than a transport fragment. Everything
|
||||
after that first frame is forwarded untouched. The byte cap can only be
|
||||
reached by a single unterminated frame.
|
||||
after that first frame is forwarded untouched.
|
||||
"""
|
||||
pending = b""
|
||||
try:
|
||||
|
|
@ -154,7 +158,7 @@ async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGener
|
|||
if preface:
|
||||
yield _SsePreface(preface)
|
||||
pending = classifiable + tail
|
||||
if classifiable or len(pending) >= _MAX_FIRST_SSE_FRAME_BYTES:
|
||||
if classifiable:
|
||||
break
|
||||
else:
|
||||
if pending:
|
||||
|
|
@ -1444,18 +1448,10 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
elif isinstance(chunk, _SsePreface):
|
||||
yield chunk.raw
|
||||
elif isinstance(chunk, bytes):
|
||||
first_frame_is_anthropic = (
|
||||
not passthrough_due_to_unknown_stream_shape
|
||||
and not all_chunks
|
||||
and is_anthropic_sse_stream((chunk,))
|
||||
)
|
||||
if not first_frame_is_anthropic:
|
||||
passthrough_due_to_unknown_stream_shape = (
|
||||
passthrough_due_to_unknown_stream_shape or not all_chunks
|
||||
)
|
||||
if all_chunks or passthrough_due_to_unknown_stream_shape:
|
||||
yield chunk
|
||||
continue
|
||||
for masked_chunk in await self._mask_anthropic_sse_stream(chunk, stream, request_data):
|
||||
for masked_chunk in await self._mask_raw_sse_stream(chunk, stream, request_data):
|
||||
yield masked_chunk
|
||||
return
|
||||
else:
|
||||
|
|
@ -1467,7 +1463,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
if passthrough_due_to_unknown_stream_shape:
|
||||
verbose_proxy_logger.warning(
|
||||
"Presidio apply_to_output: streaming response was not a parsed chat completion stream "
|
||||
"(raw non-Anthropic SSE passthrough or /v1/responses events). "
|
||||
"(an error frame from an earlier guardrail, /v1/responses events, or a mixed stream). "
|
||||
"Output PII masking was skipped for this response."
|
||||
)
|
||||
return
|
||||
|
|
@ -1489,24 +1485,71 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
|
||||
async def _mask_anthropic_sse_stream(
|
||||
async def _mask_raw_sse_stream(
|
||||
self, first_chunk: bytes, rest: AsyncIterator[object], request_data: dict
|
||||
) -> tuple[object, ...]:
|
||||
rest_chunks: Final = [chunk async for chunk in rest] # mutable-ok: tuple() cannot consume an async iterator
|
||||
"""The whole raw SSE stream masked as one response, or a raised refusal when it cannot be read.
|
||||
|
||||
Raw frames are buffered to the end because PII can span frames, so no frame is forwarded
|
||||
before the joined text was scanned, and the surface is decided on the whole stream rather
|
||||
than on its first frame. A stream that is nothing but the refusal an earlier guardrail in
|
||||
the chain emitted is forwarded as it arrived. A stream carrying a frame the parser cannot
|
||||
read, whose surface is unknown, or that its surface's assembler cannot rebuild, is withheld
|
||||
rather than forwarded unmasked.
|
||||
"""
|
||||
rest_chunks: Final = tuple([chunk async for chunk in rest])
|
||||
chunks: Final = (first_chunk, *rest_chunks)
|
||||
if has_unreadable_sse_frames(chunks):
|
||||
raise self._withheld_stream_error("could not read every streamed frame")
|
||||
if is_sse_error_stream(chunks):
|
||||
return chunks
|
||||
if is_anthropic_sse_stream(chunks):
|
||||
return await self._mask_anthropic_sse_stream(chunks, request_data)
|
||||
if is_gemini_sse_stream(chunks):
|
||||
return await self._mask_gemini_sse_stream(chunks, request_data)
|
||||
raise self._withheld_stream_error("cannot read this streaming response shape")
|
||||
|
||||
async def _mask_anthropic_sse_stream(self, chunks: tuple[object, ...], request_data: dict) -> tuple[object, ...]:
|
||||
assembled: Final = assemble_anthropic_sse_stream(chunks, restore_identity=True)
|
||||
if assembled is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Presidio apply_to_output: raw SSE stream could not be assembled into a response. "
|
||||
"Output PII masking was skipped for this response."
|
||||
)
|
||||
return chunks
|
||||
original_text: Final = model_response_text(assembled)
|
||||
raise self._withheld_stream_error("could not assemble the streaming response")
|
||||
before: Final = assembled.model_dump()
|
||||
await self._process_response_for_pii(response=assembled, request_data=request_data, mode="mask")
|
||||
if model_response_text(assembled) == original_text:
|
||||
if assembled.model_dump() == before:
|
||||
return chunks
|
||||
return anthropic_sse_chunks_from_response(assembled)
|
||||
|
||||
async def _mask_gemini_sse_stream(self, chunks: tuple[object, ...], request_data: dict) -> tuple[object, ...]:
|
||||
"""Gemini frames masked in place, so thought signatures and function-call ids reach the client.
|
||||
|
||||
The frames are rewritten rather than round-tripped through a chat-completions response
|
||||
because that translation drops the fields a Gemini client has to echo on its next turn.
|
||||
"""
|
||||
presidio_config: Final = self.get_presidio_settings_from_request_data(request_data or {})
|
||||
|
||||
async def mask_text(text: str) -> str:
|
||||
return await self.check_pii(
|
||||
text=text, output_parse_pii=False, presidio_config=presidio_config, request_data=request_data
|
||||
)
|
||||
|
||||
result: Final = await mask_gemini_sse_stream(chunks, mask_text)
|
||||
match result:
|
||||
case GeminiStreamUnchanged():
|
||||
return chunks
|
||||
case GeminiStreamMasked(frames=frames):
|
||||
return frames
|
||||
case GeminiStreamUnreadable(reason=reason):
|
||||
raise self._withheld_stream_error(reason)
|
||||
case _:
|
||||
assert_never(result)
|
||||
|
||||
def _withheld_stream_error(self, reason: str) -> GuardrailRaisedException:
|
||||
return GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message=f"output PII masking {reason}, so the response was withheld instead of being forwarded unmasked",
|
||||
status_code=500,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _unmask_sse_bytes_chunk(chunk: bytes, pii_tokens: dict[str, str]) -> bytes:
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import json
|
|||
import re
|
||||
import signal
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
|
@ -257,67 +258,93 @@ def gemini_provider(reply: Reply) -> Callable[[Request], Reply]:
|
|||
return provider
|
||||
|
||||
|
||||
def test_native_gemini_first_frame_reaches_caller_before_upstream_sends_the_second(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
gate: Final = threading.Event()
|
||||
first: Final = gemini_frame("first ")
|
||||
second: Final = gemini_frame("second ")
|
||||
provider: Final = gemini_provider(
|
||||
Reply(content_type="text/event-stream", chunks=(first, second), gate_after_first=gate)
|
||||
)
|
||||
def test_native_gemini_name_split_across_frames_is_masked_as_one_frame(gateway: Gateway, tmp_path: Path) -> None:
|
||||
"""The analyzer only matches the whole name, so a per-frame scan would forward both halves."""
|
||||
first, last = PERSON.split(" ")
|
||||
frames: Final = (gemini_frame(f"{first} "), gemini_frame(f"{last} designed it."))
|
||||
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames, pause_between_chunks=0.2))
|
||||
with presidio_rig(gateway, tmp_path, provider) as rig:
|
||||
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
|
||||
assert received.status == 200, received.text
|
||||
assert PERSON not in received.text, received.text
|
||||
assert gemini_texts(b"".join(received.frames)) == (f"{MASK} designed it.",)
|
||||
analyzed: Final = rig.analyzer.drain()
|
||||
assert len(analyzed) == 1 and json.loads(analyzed[0].body)["text"] == f"{PERSON} designed it."
|
||||
assert len(rig.anonymizer.drain()) == 1
|
||||
|
||||
|
||||
def test_native_gemini_no_frame_reaches_caller_before_upstream_finishes(gateway: Gateway, tmp_path: Path) -> None:
|
||||
gate: Final = threading.Event()
|
||||
frames: Final = (gemini_frame(f"{PERSON} "), gemini_frame("designed it."))
|
||||
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames, gate_after_first=gate))
|
||||
with presidio_rig(gateway, tmp_path, provider) as rig:
|
||||
armed_at: Final = time.monotonic()
|
||||
threading.Timer(1.0, gate.set).start()
|
||||
with rig.gateway.client.stream(
|
||||
"POST", rig.gemini_path(), json=rig.gemini_body(), headers={"Authorization": f"Bearer {rig.gateway.key}"}
|
||||
) as response:
|
||||
assert response.status_code == 200, response.read().decode()
|
||||
chunks: Final = response.iter_raw()
|
||||
arrived: Final = next(chunks)
|
||||
assert gemini_texts(arrived) == ("first ",), f"first chunk while upstream is gated: {arrived!r}"
|
||||
gate.set()
|
||||
arrived_at: Final = time.monotonic()
|
||||
rest: Final = b"".join(chunks)
|
||||
assert gemini_texts(rest) == ("second ",), rest
|
||||
assert arrived_at - armed_at >= 1.0, f"a frame reached the caller while the upstream was gated: {arrived!r}"
|
||||
assert gemini_texts(arrived + rest) == (f"{MASK} designed it.",), (arrived + rest).decode()
|
||||
assert PERSON not in (arrived + rest).decode()
|
||||
assert len(rig.upstream.drain()) == 1
|
||||
assert rig.analyzer.drain() == () and rig.anonymizer.drain() == ()
|
||||
|
||||
|
||||
def test_native_gemini_frames_received_before_upstream_abort_reach_caller(gateway: Gateway, tmp_path: Path) -> None:
|
||||
def test_native_gemini_unrecognized_stream_shape_is_withheld(gateway: Gateway, tmp_path: Path) -> None:
|
||||
frames: Final = (b'data: {"unexpected": "' + PERSON.encode() + b'"}\r\n\r\n',)
|
||||
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames))
|
||||
with presidio_rig(gateway, tmp_path, provider) as rig:
|
||||
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
|
||||
assert PERSON not in received.text, received.text
|
||||
assert "cannot read this streaming response shape" in received.text, received.text
|
||||
assert rig.analyzer.drain() == ()
|
||||
|
||||
|
||||
def test_native_gemini_upstream_abort_mid_stream_returns_an_error_and_no_frame(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
"""The stream is read whole before its first byte goes out, so an upstream abort still gets an error status."""
|
||||
frames: Final = (gemini_frame(f"chunk {index} from {PERSON}. ") for index in range(3))
|
||||
provider: Final = gemini_provider(
|
||||
Reply(content_type="text/event-stream", chunks=tuple(frames), abort_after=2, pause_between_chunks=0.2)
|
||||
)
|
||||
with presidio_rig(gateway, tmp_path, provider) as rig:
|
||||
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
|
||||
assert received.status == 200, received.text
|
||||
*frames_before_abort, trailer = data_payloads(b"".join(received.frames))
|
||||
assert [gemini_text(frame) for frame in frames_before_abort] == [
|
||||
f"chunk 0 from {PERSON}. ",
|
||||
f"chunk 1 from {PERSON}. ",
|
||||
], received.text
|
||||
assert "candidates" not in trailer and json.dumps(trailer).count('"code": "500"') == 1, received.text
|
||||
assert received.status == 500, received.text
|
||||
assert PERSON not in received.text, received.text
|
||||
error: Final = json.loads(received.text)["error"]
|
||||
assert error["code"] == "500" and "candidates" not in received.text, received.text
|
||||
assert len(rig.upstream.drain()) == 1
|
||||
|
||||
|
||||
def test_native_gemini_first_frame_split_into_transport_fragments_streams_every_byte(
|
||||
def test_native_gemini_first_frame_split_into_transport_fragments_is_still_masked(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
first: Final = gemini_frame(f"fragmented {PERSON}")
|
||||
first: Final = gemini_frame(f"fragmented {PERSON} ")
|
||||
second: Final = gemini_frame("whole")
|
||||
chunks: Final = (first[:7], first[7:19], first[19:], second)
|
||||
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=chunks))
|
||||
with presidio_rig(gateway, tmp_path, provider) as rig:
|
||||
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
|
||||
assert received.status == 200, received.text
|
||||
assert gemini_texts(b"".join(received.frames)) == (f"fragmented {PERSON}", "whole")
|
||||
assert PERSON not in received.text, received.text
|
||||
assert gemini_texts(b"".join(received.frames)) == (f"fragmented {MASK} whole",)
|
||||
|
||||
|
||||
def test_native_gemini_non_json_frame_passes_through_unchanged(gateway: Gateway, tmp_path: Path) -> None:
|
||||
frames: Final = (b"data: not json at all\r\n\r\n", gemini_frame("after"))
|
||||
def test_native_gemini_stream_with_an_unreadable_frame_is_withheld(gateway: Gateway, tmp_path: Path) -> None:
|
||||
"""A frame the parser cannot read is one the guardrail cannot scan, so nothing around it is replayed either."""
|
||||
frames: Final = (gemini_frame(f"the architect was {PERSON}"), b"data: not json at all\r\n\r\n", gemini_frame("."))
|
||||
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames))
|
||||
with presidio_rig(gateway, tmp_path, provider) as rig:
|
||||
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
|
||||
assert received.status == 200, received.text
|
||||
assert received.text.replace("\r\n", "\n") == b"".join(frames).decode().replace("\r\n", "\n")
|
||||
assert received.status == 500, received.text
|
||||
assert PERSON not in received.text, received.text
|
||||
assert "could not read every streamed frame" in received.text, received.text
|
||||
assert rig.analyzer.drain() == ()
|
||||
|
||||
|
||||
def test_native_gemini_empty_stream_returns_200_with_no_body(gateway: Gateway, tmp_path: Path) -> None:
|
||||
|
|
@ -328,14 +355,14 @@ def test_native_gemini_empty_stream_returns_200_with_no_body(gateway: Gateway, t
|
|||
assert received.text == ""
|
||||
|
||||
|
||||
def test_native_gemini_streams_while_presidio_analyzer_is_down(gateway: Gateway, tmp_path: Path) -> None:
|
||||
def test_native_gemini_stream_fails_closed_when_analyzer_is_down(gateway: Gateway, tmp_path: Path) -> None:
|
||||
frames: Final = (gemini_frame(f"{PERSON} one. "), gemini_frame("two."))
|
||||
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames))
|
||||
with presidio_rig(gateway, tmp_path, provider, analyze=broken) as rig:
|
||||
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
|
||||
assert received.status == 200, received.text
|
||||
assert gemini_texts(b"".join(received.frames)) == (f"{PERSON} one. ", "two.")
|
||||
assert rig.analyzer.drain() == ()
|
||||
assert PERSON not in received.text, received.text
|
||||
assert "Presidio analyzer" in received.text, received.text
|
||||
assert rig.anonymizer.drain() == ()
|
||||
|
||||
|
||||
def test_native_gemini_unauthenticated_request_is_rejected_before_upstream(gateway: Gateway, tmp_path: Path) -> None:
|
||||
|
|
@ -526,17 +553,13 @@ def test_mixed_burst_survives_anonymizer_outage_and_recovers(gateway: Gateway, t
|
|||
|
||||
def gemini_call(index: int) -> tuple[str, str, int]:
|
||||
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
|
||||
return (
|
||||
"gemini",
|
||||
f"g{index}",
|
||||
received.status if gemini_texts(b"".join(received.frames)) == (f"{PERSON} ", "designed it.") else -1,
|
||||
)
|
||||
masked: Final = gemini_texts(b"".join(received.frames)) == (f"{MASK} designed it.",)
|
||||
return ("gemini", f"g{index}", -1 if PERSON in received.text else (1 if masked else 0))
|
||||
|
||||
def anthropic_call(index: int) -> tuple[str, str, int]:
|
||||
body: Final = {**rig.messages_body(), "messages": [{"role": "user", "content": f"a{index}"}]}
|
||||
received: Final = rig.stream("/v1/messages", body)
|
||||
leaked: Final = PERSON in received.text
|
||||
return ("anthropic", f"a{index}", -1 if leaked else (1 if MASK in received.text else 0))
|
||||
return ("anthropic", f"a{index}", -1 if PERSON in received.text else (1 if MASK in received.text else 0))
|
||||
|
||||
def phase(offset: int) -> tuple[tuple[str, str, int], ...]:
|
||||
with ThreadPoolExecutor(max_workers=12) as pool:
|
||||
|
|
@ -552,13 +575,9 @@ def test_mixed_burst_survives_anonymizer_outage_and_recovers(gateway: Gateway, t
|
|||
outage.clear()
|
||||
healthy_after: Final = phase(200)
|
||||
|
||||
for name, results in (("before", healthy_before), ("during", during), ("after", healthy_after)):
|
||||
assert all(status == 200 for kind, _, status in results if kind == "gemini"), (name, results)
|
||||
assert all(status == 1 for kind, _, status in healthy_before + healthy_after if kind == "anthropic"), (
|
||||
healthy_before,
|
||||
healthy_after,
|
||||
)
|
||||
assert all(status == 0 for kind, _, status in during if kind == "anthropic"), during
|
||||
for name, results in (("before", healthy_before), ("after", healthy_after)):
|
||||
assert all(outcome == 1 for _, _, outcome in results), (name, results)
|
||||
assert all(outcome == 0 for _, _, outcome in during), during
|
||||
identities: Final = tuple(identity for _, identity, _ in healthy_before + during + healthy_after)
|
||||
assert len(identities) == len(set(identities)) == 36
|
||||
|
||||
|
|
|
|||
|
|
@ -1449,11 +1449,10 @@ from litellm.types.utils import ModelResponseStream
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_with_bytes_chunks_does_not_crash(mock_user_api_key):
|
||||
async def test_streaming_unrecognized_raw_frame_ahead_of_typed_chunks_is_withheld(mock_user_api_key):
|
||||
"""
|
||||
Regression test: async_post_call_streaming_iterator_hook should
|
||||
gracefully handle raw bytes in the stream instead of crashing with
|
||||
'bytes' object has no attribute 'id'.
|
||||
A raw frame the hook cannot place on a known surface is not a chunk it can
|
||||
mask, so the response is refused instead of being forwarded unscanned.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
|
|
@ -1462,7 +1461,7 @@ async def test_streaming_with_bytes_chunks_does_not_crash(mock_user_api_key):
|
|||
)
|
||||
|
||||
async def mock_stream():
|
||||
yield b'data: {"id":"chatcmpl-1"}\n\n' # raw bytes
|
||||
yield b'data: {"id":"chatcmpl-1"}\n\n'
|
||||
yield ModelResponseStream(
|
||||
id="chatcmpl-1",
|
||||
choices=[],
|
||||
|
|
@ -1470,18 +1469,22 @@ async def test_streaming_with_bytes_chunks_does_not_crash(mock_user_api_key):
|
|||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
system_fingerprint=None,
|
||||
) # proper chunk
|
||||
)
|
||||
|
||||
chunks = []
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=mock_user_api_key,
|
||||
response=mock_stream(),
|
||||
request_data={},
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
# Should not crash, should produce at least one valid chunk
|
||||
assert len(chunks) >= 1
|
||||
async def collect():
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=mock_user_api_key,
|
||||
response=mock_stream(),
|
||||
request_data={},
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
with pytest.raises(GuardrailRaisedException, match="cannot read this streaming response shape"):
|
||||
await collect()
|
||||
|
||||
assert chunks == []
|
||||
|
||||
|
||||
def test_entity_deny_list_filters_detections():
|
||||
|
|
@ -2114,10 +2117,10 @@ async def test_anthropic_native_response_non_text_blocks_untouched():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_bytes_chunks_are_yielded_not_discarded():
|
||||
async def test_streaming_partial_anthropic_stream_without_message_start_is_withheld():
|
||||
"""
|
||||
Regression test: bytes chunks (Anthropic native SSE) should be yielded
|
||||
through the streaming hook, not silently discarded.
|
||||
A raw Anthropic stream that cannot be assembled (no message_start) is
|
||||
refused with a clear error rather than forwarded unmasked or dropped silently.
|
||||
"""
|
||||
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
|
|
@ -2132,15 +2135,19 @@ async def test_streaming_bytes_chunks_are_yielded_not_discarded():
|
|||
|
||||
mock_user_api_key = UserAPIKeyAuth(api_key="test-key")
|
||||
chunks = []
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=mock_user_api_key,
|
||||
response=mock_stream(),
|
||||
request_data={},
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
assert any(isinstance(c, bytes) for c in chunks), "bytes chunks must not be discarded"
|
||||
assert byte_chunk in chunks
|
||||
async def collect():
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=mock_user_api_key,
|
||||
response=mock_stream(),
|
||||
request_data={},
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
with pytest.raises(GuardrailRaisedException, match="could not assemble the streaming response"):
|
||||
await collect()
|
||||
|
||||
assert chunks == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -2528,13 +2535,397 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_without_pii_are_for
|
|||
assert collected == byte_chunks
|
||||
|
||||
|
||||
def _gemini_sse(text: str) -> bytes:
|
||||
def _gemini_sse(text: str, terminator: bytes = b"\n\n") -> bytes:
|
||||
payload = {"candidates": [{"content": {"parts": [{"text": text}], "role": "model"}, "index": 0}]}
|
||||
return f"data: {json.dumps(payload)}\n\n".encode()
|
||||
return b"data: " + json.dumps(payload).encode() + terminator
|
||||
|
||||
|
||||
def _gemini_texts(chunks: list[object]) -> list[str]:
|
||||
frames = b"".join(chunk for chunk in chunks if isinstance(chunk, bytes)).decode()
|
||||
return [
|
||||
json.loads(line[6:])["candidates"][0]["content"]["parts"][0]["text"]
|
||||
for line in frames.splitlines()
|
||||
if line.startswith("data: ")
|
||||
]
|
||||
|
||||
|
||||
def _gemini_frame(*parts: dict, finish_reason: str | None = None, **top_level: object) -> bytes:
|
||||
candidate = {"content": {"parts": list(parts), "role": "model"}, "index": 0}
|
||||
payload = {
|
||||
"candidates": [candidate if finish_reason is None else {**candidate, "finishReason": finish_reason}],
|
||||
**top_level,
|
||||
}
|
||||
return b"data: " + json.dumps(payload).encode() + b"\n\n"
|
||||
|
||||
|
||||
def _gemini_frames(chunks: list[object]) -> list[dict]:
|
||||
frames = b"".join(chunk for chunk in chunks if isinstance(chunk, bytes)).decode()
|
||||
return [json.loads(line[6:]) for line in frames.splitlines() if line.startswith("data: ")]
|
||||
|
||||
|
||||
def _gemini_fake_masking_guardrail(server: TestServer) -> _OPTIONAL_PresidioPIIMasking:
|
||||
return _OPTIONAL_PresidioPIIMasking(
|
||||
apply_to_output=True,
|
||||
presidio_analyzer_api_base=str(server.make_url("/")),
|
||||
presidio_anonymizer_api_base=str(server.make_url("/")),
|
||||
pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK},
|
||||
)
|
||||
|
||||
|
||||
async def _collect_masked_output(guardrail: _OPTIONAL_PresidioPIIMasking, stream, collected: list[object]) -> None:
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
|
||||
response=stream,
|
||||
request_data={},
|
||||
):
|
||||
collected.append(chunk)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incrementally_until_upstream_aborts():
|
||||
async def test_apply_to_output_streaming_gemini_name_split_across_frames_is_masked_as_one_response():
|
||||
"""
|
||||
The name only exists once the frames are joined, so a per-frame scan would
|
||||
forward both halves unmasked. The analyzer and anonymizer are the in-process
|
||||
fake, which finds a person only in two adjacent capitalized words.
|
||||
"""
|
||||
frames = [_gemini_sse("The architect was John"), _gemini_sse(" Smith, per the record.")]
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
collected: list[object] = []
|
||||
async with TestServer(_fake_presidio_app()) as server:
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
apply_to_output=True,
|
||||
presidio_analyzer_api_base=str(server.make_url("/")),
|
||||
presidio_anonymizer_api_base=str(server.make_url("/")),
|
||||
pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK},
|
||||
)
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
await guardrail._close_http_session()
|
||||
|
||||
assert "John Smith" not in b"".join(collected).decode()
|
||||
assert _gemini_texts(collected) == ["The architect was <PERSON>, 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 "<PERSON>" in joined, joined
|
||||
assert joined.count("event: message_start") == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("terminator", [b"\n\n", b"\r\n\r\n"])
|
||||
async def test_apply_to_output_streaming_gemini_stream_without_pii_is_replayed_byte_for_byte(terminator):
|
||||
function_call = {
|
||||
"functionCall": {"id": "call_1", "name": "lookup", "args": {"city": "Paris"}},
|
||||
"thoughtSignature": "sig",
|
||||
}
|
||||
frames = [
|
||||
_gemini_sse("nothing personal ", terminator),
|
||||
_gemini_sse("in here.", terminator),
|
||||
b"data: "
|
||||
+ json.dumps({"candidates": [{"content": {"parts": [function_call], "role": "model"}}]}).encode()
|
||||
+ terminator,
|
||||
]
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
collected: list[object] = []
|
||||
async with TestServer(_fake_presidio_app()) as server:
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
apply_to_output=True,
|
||||
presidio_analyzer_api_base=str(server.make_url("/")),
|
||||
presidio_anonymizer_api_base=str(server.make_url("/")),
|
||||
pii_entities_config={PiiEntityType.PERSON: PiiAction.MASK},
|
||||
)
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
await guardrail._close_http_session()
|
||||
|
||||
assert collected == frames
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_function_call_keeps_its_signature_and_id_while_its_arguments_are_masked():
|
||||
"""
|
||||
A Gemini client echoes the streamed model turn on its next request, and Gemini 3 rejects
|
||||
a function call whose thought signature is missing, so masking has to rewrite the frames
|
||||
in place rather than rebuild them from a chat-completions response that drops those fields.
|
||||
"""
|
||||
frames = [
|
||||
_gemini_frame(
|
||||
{"text": "emailing John"}, usageMetadata={"totalTokenCount": 1}, modelVersion="m", responseId="r1"
|
||||
),
|
||||
_gemini_frame({"text": " Smith now."}, usageMetadata={"totalTokenCount": 2}, modelVersion="m", responseId="r1"),
|
||||
_gemini_frame(
|
||||
{
|
||||
"functionCall": {"id": "call_1", "name": "send_email", "args": {"to": "John Smith", "urgent": True}},
|
||||
"thoughtSignature": "sig-fc",
|
||||
},
|
||||
usageMetadata={"totalTokenCount": 3},
|
||||
modelVersion="m",
|
||||
responseId="r1",
|
||||
),
|
||||
_gemini_frame(
|
||||
{"text": ""}, finish_reason="STOP", usageMetadata={"totalTokenCount": 9}, modelVersion="m", responseId="r1"
|
||||
),
|
||||
]
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
collected: list[object] = []
|
||||
async with TestServer(_fake_presidio_app()) as server:
|
||||
guardrail = _gemini_fake_masking_guardrail(server)
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
await guardrail._close_http_session()
|
||||
|
||||
assert "John Smith" not in b"".join(collected).decode()
|
||||
(body,) = _gemini_frames(collected)
|
||||
assert body["candidates"] == [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{"text": "emailing <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_keeps_its_own_terminal_frame():
|
||||
"""
|
||||
Gemini can end a stream with an error frame after content frames, and a client
|
||||
reads the failure off that terminal frame, so the masked re-emit keeps it separate.
|
||||
"""
|
||||
frames = [
|
||||
_gemini_frame({"text": "The architect was John Smith."}),
|
||||
b'data: {"error": {"code": 503, "message": "overloaded", "status": "UNAVAILABLE"}}\n\n',
|
||||
]
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
async with TestServer(_fake_presidio_app()) as server:
|
||||
guardrail = _gemini_fake_masking_guardrail(server)
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
await guardrail._close_http_session()
|
||||
|
||||
masked, error = _gemini_frames(collected)
|
||||
assert masked["candidates"][0]["content"]["parts"] == [{"text": "The architect was <PERSON>."}]
|
||||
assert "error" not in masked
|
||||
assert error == {"error": {"code": 503, "message": "overloaded", "status": "UNAVAILABLE"}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_content_behind_a_leading_error_frame_is_still_masked():
|
||||
"""
|
||||
The stream surface is decided on the whole buffered stream, not on its first frame, so an
|
||||
error frame ahead of content frames does not switch masking off for the content behind it.
|
||||
"""
|
||||
frames = [
|
||||
b'data: {"error": {"code": 429, "message": "quota", "status": "RESOURCE_EXHAUSTED"}}\n\n',
|
||||
_gemini_frame({"text": "The architect was John Smith."}),
|
||||
]
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
async with TestServer(_fake_presidio_app()) as server:
|
||||
guardrail = _gemini_fake_masking_guardrail(server)
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
await guardrail._close_http_session()
|
||||
|
||||
masked, error = _gemini_frames(collected)
|
||||
assert masked["candidates"][0]["content"]["parts"] == [{"text": "The architect was <PERSON>."}]
|
||||
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": "<PERSON>"},
|
||||
)
|
||||
frames = [
|
||||
b'data: {"error": {"message": "Violated guardrail policy", "type": "guardrail_error"}}\n\n',
|
||||
b"data: [DONE]\n\n",
|
||||
]
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
assert collected == frames
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_every_candidate_is_merged_by_index_and_masked():
|
||||
"""
|
||||
Gemini interleaves every candidate's fragments across frames, so each candidate
|
||||
is joined on its own index and scanned, not only the first one.
|
||||
"""
|
||||
|
||||
def frame(*texts: tuple[int, str]) -> bytes:
|
||||
candidates = [{"content": {"parts": [{"text": text}], "role": "model"}, "index": index} for index, text in texts]
|
||||
return b"data: " + json.dumps({"candidates": candidates}).encode() + b"\n\n"
|
||||
|
||||
frames = [
|
||||
frame((0, "The architect was John"), (1, "The engineer was Ada")),
|
||||
frame((0, " Smith."), (1, " Lovelace.")),
|
||||
]
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for chunk in frames:
|
||||
yield chunk
|
||||
|
||||
async with TestServer(_fake_presidio_app()) as server:
|
||||
guardrail = _gemini_fake_masking_guardrail(server)
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
await guardrail._close_http_session()
|
||||
|
||||
(masked,) = _gemini_frames(collected)
|
||||
assert [candidate["content"]["parts"] for candidate in masked["candidates"]] == [
|
||||
[{"text": "The architect was <PERSON>."}],
|
||||
[{"text": "The engineer was <PERSON>."}],
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_upstream_abort_mid_stream_forwards_no_frame():
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
|
|
@ -2548,18 +2939,31 @@ async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incremen
|
|||
yield frame
|
||||
raise ConnectionError("upstream closed mid-stream")
|
||||
|
||||
async def collect() -> None:
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
|
||||
response=mock_stream(),
|
||||
request_data={},
|
||||
):
|
||||
collected.append(chunk)
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
await collect()
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
assert collected == frames
|
||||
assert collected == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_sse_bytes_fail_closed_when_presidio_is_unreachable():
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
presidio_analyzer_api_base="http://127.0.0.1:9",
|
||||
presidio_anonymizer_api_base="http://127.0.0.1:9",
|
||||
)
|
||||
frames = [_gemini_sse("Hello John Smith"), _gemini_sse(" from the record.")]
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
with pytest.raises(Exception, match="Presidio PII analysis failed"):
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
assert collected == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -2787,13 +3191,13 @@ async def test_apply_to_output_streaming_leading_keepalive_is_forwarded_before_u
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_leading_comments_over_the_frame_cap_are_still_masked():
|
||||
async def test_apply_to_output_streaming_leading_comments_of_any_size_are_still_masked():
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
mock_redacted_text={"text": "<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]
|
||||
|
|
@ -2850,7 +3254,7 @@ async def test_apply_to_output_streaming_comment_only_stream_is_forwarded_unchan
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_first_frame_split_across_transport_chunks_streams_incrementally():
|
||||
async def test_apply_to_output_streaming_gemini_first_frame_split_across_transport_chunks_is_still_masked():
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
|
|
@ -2864,50 +3268,103 @@ async def test_apply_to_output_streaming_gemini_first_frame_split_across_transpo
|
|||
yield first[:20]
|
||||
yield first[20:]
|
||||
yield second
|
||||
raise ConnectionError("upstream closed mid-stream")
|
||||
|
||||
async def collect() -> None:
|
||||
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
|
||||
response=mock_stream(),
|
||||
request_data={},
|
||||
):
|
||||
collected.append(chunk)
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
await collect()
|
||||
|
||||
assert collected == [first, second]
|
||||
assert "John Smith" not in b"".join(collected).decode()
|
||||
assert _gemini_texts(collected) == ["<PERSON>"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_unterminated_first_frame_is_released_once_it_exceeds_the_cap():
|
||||
@pytest.mark.parametrize(
|
||||
("frame", "reason"),
|
||||
[
|
||||
(b"data: not json at all\n\n", "could not read every streamed frame"),
|
||||
(b"data: 42\n\n", "could not read every streamed frame"),
|
||||
(b'data: {"unexpected": true}\n\n', "cannot read this streaming response shape"),
|
||||
(b"data: " + b"x" * (70 * 1024) + b"\n", "could not read every streamed frame"),
|
||||
],
|
||||
ids=["non-json", "json-scalar", "unknown-surface", "unterminated"],
|
||||
)
|
||||
async def test_apply_to_output_streaming_unrecognized_raw_sse_stream_is_withheld(frame: bytes, reason: str):
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
mock_redacted_text={"text": "<PERSON>"},
|
||||
)
|
||||
piece = b"data: " + b"x" * 1023 + b"\n"
|
||||
pieces_to_cap = -(-(64 * 1024) // len(piece))
|
||||
released_at: list[int] = []
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for index in range(pieces_to_cap * 4):
|
||||
if collected:
|
||||
released_at.append(index)
|
||||
yield piece
|
||||
yield frame
|
||||
|
||||
with pytest.raises(GuardrailRaisedException, match=reason):
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
assert collected == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_gemini_stream_with_an_unreadable_frame_is_withheld():
|
||||
"""
|
||||
The parser skips a frame it cannot read, and replaying the originals around it
|
||||
would forward that frame unscanned, so the whole stream is withheld instead.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
mock_redacted_text={"text": "<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 for chunk in guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
|
||||
response=mock_stream(),
|
||||
request_data={},
|
||||
):
|
||||
collected.append(chunk)
|
||||
|
||||
assert released_at, "nothing reached the caller before the upstream finished"
|
||||
assert released_at[0] == pieces_to_cap, released_at[:3]
|
||||
assert b"".join(collected) == piece * (pieces_to_cap * 4)
|
||||
async def mock_stream():
|
||||
for frame in frames:
|
||||
yield frame
|
||||
|
||||
with pytest.raises(GuardrailRaisedException, match="could not read every streamed frame"):
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
assert collected == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_to_output_streaming_anthropic_stream_with_an_unreadable_frame_is_withheld():
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
mock_testing=True,
|
||||
apply_to_output=True,
|
||||
mock_redacted_text={"text": "Hello"},
|
||||
)
|
||||
byte_chunks = [
|
||||
_anthropic_sse(
|
||||
"message_start",
|
||||
{"type": "message_start", "message": {"id": "msg_1", "model": "claude", "content": [], "usage": {}}},
|
||||
),
|
||||
_anthropic_sse(
|
||||
"content_block_start",
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
|
||||
),
|
||||
b"event: content_block_delta\ndata: not json at all\n\n",
|
||||
_anthropic_sse(
|
||||
"content_block_delta",
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "Hello"}},
|
||||
),
|
||||
_anthropic_sse("content_block_stop", {"type": "content_block_stop", "index": 0}),
|
||||
_anthropic_sse("message_stop", {"type": "message_stop"}),
|
||||
]
|
||||
collected: list[object] = []
|
||||
|
||||
async def mock_stream():
|
||||
for chunk in byte_chunks:
|
||||
yield chunk
|
||||
|
||||
with pytest.raises(GuardrailRaisedException, match="could not read every streamed frame"):
|
||||
await _collect_masked_output(guardrail, mock_stream(), collected)
|
||||
|
||||
assert collected == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue