fix(guardrails): keep a Gemini upstream error frame as its own terminal frame after masking

This commit is contained in:
mateo-berri 2026-09-26 16:22:26 -07:00
parent 9e683e86f3
commit 2a9acc692e
2 changed files with 28 additions and 13 deletions

View file

@ -3,10 +3,11 @@
The Google ``:streamGenerateContent`` route relays the upstream ``data:`` frames as raw bytes, so a
guardrail's ``async_post_call_streaming_iterator_hook`` cannot read them as chunk objects. These
helpers fold such a stream into the one response body a non-streaming call would have returned,
rewrite its text and function-call arguments in place, and re-emit it as one frame. Every other
field (thought signatures, function-call ids, model version, response id, safety ratings, usage)
is carried through untouched, because a client echoes the model turn back on its next request and
Gemini 3 rejects a function call whose thought signature is missing.
rewrite its text and function-call arguments in place, and re-emit it as one frame, followed by any
upstream error frame kept as its own terminal frame. Every other field (thought signatures,
function-call ids, model version, response id, safety ratings, usage) is carried through untouched,
because a client echoes the model turn back on its next request and Gemini 3 rejects a function
call whose thought signature is missing.
"""
from __future__ import annotations
@ -63,7 +64,8 @@ async def mask_gemini_sse_stream(all_chunks: Sequence[object], mask_text: TextMa
Text arrives split across frames, so the fragments of one text part are joined before they are
scanned; PII that only exists once the fragments meet is otherwise forwarded in halves. Top-level
keys and candidate keys take the last frame's value, and candidates are merged by index.
keys and candidate keys take the last frame's value, and candidates are merged by index. An
upstream error frame stays its own terminal frame after the masked response, as it arrived.
"""
sse_stream: Final = joined_sse_stream(all_chunks)
if sse_stream is None:
@ -76,8 +78,20 @@ async def mask_gemini_sse_stream(all_chunks: Sequence[object], mask_text: TextMa
return unreadable
if masked_candidates == candidates:
return GeminiStreamUnchanged()
response: Final = MappingProxyType({**_later_wins(events), "candidates": masked_candidates})
return GeminiStreamMasked((f"data: {_json_text(response)}\n\n".encode(),))
response: Final = MappingProxyType({**_without_error(_later_wins(events)), "candidates": masked_candidates})
return GeminiStreamMasked((_frame(response), *map(_frame, _error_frames(events))))
def _frame(event: _JsonObject) -> bytes:
return f"data: {_json_text(event)}\n\n".encode()
def _without_error(event: _JsonObject) -> _JsonObject:
return MappingProxyType({key: value for key, value in event.items() if key != "error"})
def _error_frames(events: Sequence[_JsonObject]) -> Iterator[_JsonObject]:
return (MappingProxyType({"error": event["error"]}) for event in events if "error" in event)
def _json_text(value: object) -> str:

View file

@ -2808,10 +2808,10 @@ async def test_apply_to_output_streaming_gemini_masked_function_call_arguments_t
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_error_frame_mid_stream_is_carried_on_the_masked_frame():
async def test_apply_to_output_streaming_gemini_error_frame_mid_stream_keeps_its_own_terminal_frame():
"""
Gemini can end a stream with an error frame after content frames, and the
client reads the error off the last frame, so the masked re-emit keeps it.
Gemini can end a stream with an error frame after content frames, and a client
reads the failure off that terminal frame, so the masked re-emit keeps it separate.
"""
frames = [
_gemini_frame({"text": "The architect was John Smith."}),
@ -2828,9 +2828,10 @@ async def test_apply_to_output_streaming_gemini_error_frame_mid_stream_is_carrie
await _collect_masked_output(guardrail, mock_stream(), collected)
await guardrail._close_http_session()
(frame,) = _gemini_frames(collected)
assert frame["candidates"][0]["content"]["parts"] == [{"text": "The architect was <PERSON>."}]
assert frame["error"] == {"code": 503, "message": "overloaded", "status": "UNAVAILABLE"}
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