fix(guardrails): classify a raw SSE stream as error-only on the whole drained stream, not its first frame

An error frame ahead of Gemini or Anthropic content frames switched masking off
for every frame behind it. The refusal check now runs on the drained stream, so
an error-only stream is still forwarded as it arrived and content behind a
leading error frame is masked
This commit is contained in:
mateo-berri 2026-09-26 17:19:49 -07:00
parent 2a9acc692e
commit d74f5ddb1d
2 changed files with 61 additions and 7 deletions

View file

@ -1446,10 +1446,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
elif isinstance(chunk, _SsePreface):
yield chunk.raw
elif isinstance(chunk, bytes):
if all_chunks or passthrough_due_to_unknown_stream_shape or is_sse_error_stream((chunk,)):
passthrough_due_to_unknown_stream_shape = (
passthrough_due_to_unknown_stream_shape or not all_chunks
)
if all_chunks or passthrough_due_to_unknown_stream_shape:
yield chunk
continue
for masked_chunk in await self._mask_raw_sse_stream(chunk, stream, request_data):
@ -1492,14 +1489,18 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
"""The whole raw SSE stream masked as one response, or a raised refusal when it cannot be read.
Raw frames are buffered to the end because PII can span frames, so no frame is forwarded
before the joined text was scanned. A stream carrying a frame the parser cannot read, whose
surface is unknown, or that its surface's assembler cannot rebuild, is withheld rather than
forwarded unmasked.
before the joined text was scanned, and the surface is decided on the whole stream rather
than on its first frame. A stream that is nothing but the refusal an earlier guardrail in
the chain emitted is forwarded as it arrived. A stream carrying a frame the parser cannot
read, whose surface is unknown, or that its surface's assembler cannot rebuild, is withheld
rather than forwarded unmasked.
"""
rest_chunks: Final = tuple([chunk async for chunk in rest])
chunks: Final = (first_chunk, *rest_chunks)
if has_unreadable_sse_frames(chunks):
raise self._withheld_stream_error("could not read every streamed frame")
if is_sse_error_stream(chunks):
return chunks
if is_anthropic_sse_stream(chunks):
return await self._mask_anthropic_sse_stream(chunks, request_data)
if is_gemini_sse_stream(chunks):

View file

@ -2834,6 +2834,59 @@ async def test_apply_to_output_streaming_gemini_error_frame_mid_stream_keeps_its
assert error == {"error": {"code": 503, "message": "overloaded", "status": "UNAVAILABLE"}}
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_content_behind_a_leading_error_frame_is_still_masked():
"""
The stream surface is decided on the whole buffered stream, not on its first frame, so an
error frame ahead of content frames does not switch masking off for the content behind it.
"""
frames = [
b'data: {"error": {"code": 429, "message": "quota", "status": "RESOURCE_EXHAUSTED"}}\n\n',
_gemini_frame({"text": "The architect was John Smith."}),
]
collected: list[object] = []
async def mock_stream():
for frame in frames:
yield frame
async with TestServer(_fake_presidio_app()) as server:
guardrail = _gemini_fake_masking_guardrail(server)
await _collect_masked_output(guardrail, mock_stream(), collected)
await guardrail._close_http_session()
masked, error = _gemini_frames(collected)
assert masked["candidates"][0]["content"]["parts"] == [{"text": "The architect was <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():
"""