fix(guardrails): flush every held choice, and cover tool results and suffix

Three review findings.

The trailing flush walked the last chunk's choices, so a choice that finished
earlier and stopped appearing lost whatever text was still held for it and its
answer was truncated. It is now driven by the windows themselves and emits one
chunk per choice, synthesising the choice when the terminal chunk omits it.
That was data loss, not just under-redaction.

An Anthropic tool_result carries its own content, as a string or as further
blocks, and only each part's `text` was being collected. Handled recursively;
image and audio parts still fall through untouched.

The legacy completions `suffix` is forwarded to providers that support it and
was never collected. Note the placement: it has to be gathered before the
string-prompt early return, which is what the new test pins.
This commit is contained in:
Ninad Phalak 2026-09-03 23:28:26 -05:00
parent 6536d61adf
commit 403b06ea76
No known key found for this signature in database
GPG key ID: 59119ED515433744
2 changed files with 125 additions and 18 deletions

View file

@ -100,7 +100,8 @@ def _collect_entry(entries: MutableSeq, index: int, slots: _SlotSink) -> None:
def _collect_prompt(data: MutableRequest, slots: _SlotSink) -> None:
"""The Completions API sends its text in a top-level `prompt`."""
"""The Completions API sends its text in `prompt`, and its tail in `suffix`."""
_collect(data, "suffix", slots)
prompt: Final = data.get("prompt")
if isinstance(prompt, str):
_collect(data, "prompt", slots)
@ -118,8 +119,13 @@ def _collect_content(container: MutableRequest, slots: _SlotSink) -> None:
_collect(container, "content", slots)
return
for part in content if isinstance(content, list) else ():
if isinstance(part, dict):
_collect(part, "text", slots)
if not isinstance(part, dict):
continue
_collect(part, "text", slots)
# An Anthropic tool_result carries its own content, as a string or as more
# blocks. Image and audio parts have no text and fall through untouched.
if "content" in part:
_collect_content(part, slots)
def _collect_participant_name(message: MutableRequest, slots: _SlotSink) -> None:
@ -486,8 +492,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
# A stream that ended without a finish_reason can still leave text held back.
if last_chunk is not None and any(carries.values()):
trailing: Final = last_chunk.model_copy(deep=True)
if await self._flush_trailing(trailing, carries, session_id):
async for trailing in self._flush_trailing(last_chunk, carries, session_id):
yield trailing
async def _restore_choice(self, choice: Any, carries: _CarryWindows, session_id: str) -> None:
@ -513,23 +518,50 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
carries[index] = remaining # rebind-ok: this choice's window advances.
delta.content = emitted
async def _flush_trailing(self, trailing: Any, carries: _CarryWindows, session_id: str) -> bool:
"""Empties every still-held window into a copy of the last chunk."""
emitted_any = False # rebind-ok: set once any choice contributes text.
for choice in getattr(trailing, "choices", None) or ():
delta = getattr(choice, "delta", None)
if delta is None:
continue
index = _choice_index(choice)
carry = carries.get(index, "")
async def _flush_trailing(
self, last_chunk: Any, carries: _CarryWindows, session_id: str
) -> AsyncGenerator[Any, None]:
"""Empties every window still holding text, one chunk per choice.
Driven by the windows rather than by the last chunk's choices. A choice that
finished earlier is not present in the terminal chunk, and flushing only what
that chunk carries would drop its held text and truncate its answer.
"""
for index in sorted(carries):
carry = carries[index]
if not carry:
delta.content = None
continue
text, remaining = await self._stream_step("", carry, True, session_id)
carries[index] = remaining # rebind-ok: this choice's window advances.
delta.content = text or None
emitted_any = emitted_any or bool(text)
return emitted_any
if not text:
continue
chunk = self._chunk_for_choice(last_chunk, index)
if chunk is None:
continue
chunk.choices[0].delta.content = text
yield chunk
@staticmethod
def _chunk_for_choice(last_chunk: Any, index: int) -> Any:
"""A single-choice copy of the last chunk, carrying only `index`.
Emitting one choice per chunk keeps a flush from reading as content on a
choice it does not belong to.
"""
chunk: Final = last_chunk.model_copy(deep=True)
raw_choices: Final = getattr(chunk, "choices", None)
if not raw_choices:
return None
choices: Final[tuple] = tuple(raw_choices)
matching: Final = tuple(choice for choice in choices if _choice_index(choice) == index)
kept: Final = matching[0] if matching else choices[0]
if getattr(kept, "delta", None) is None:
return None
kept.index = index
# The terminal signal, if there was one, already went out with the real chunk.
kept.finish_reason = None
chunk.choices = [kept] # mutable-ok: the chunk model requires a list.
return chunk
async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]:
"""Returns ``(text safe to emit now, window still being held)``."""

View file

@ -358,6 +358,43 @@ class TestRequestCoverage:
assert data["messages"][0]["name"] == "get_weather"
assert mock.call_args_list[0].kwargs["json"]["texts"] == ["result"]
@pytest.mark.asyncio
async def test_anthropic_tool_result_content_is_redacted(self):
"""A tool_result nests its own content, as a string or as more blocks."""
guardrail = _guardrail()
_mock_post(guardrail, {"texts": ["[EMAIL_1]", "[EMAIL_2]"]})
data = {
"messages": [
{
"role": "user",
"content": [
{"type": "tool_result", "tool_use_id": "t1", "content": "found jane.doe@example.com"},
{
"type": "tool_result",
"tool_use_id": "t2",
"content": [{"type": "text", "text": "also bob@example.com"}],
},
],
}
]
}
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion")
assert data["messages"][0]["content"][0]["content"] == "[EMAIL_1]"
assert data["messages"][0]["content"][1]["content"][0]["text"] == "[EMAIL_2]"
@pytest.mark.asyncio
async def test_completions_suffix_is_redacted(self):
"""LiteLLM forwards the legacy `suffix` to providers that support it."""
guardrail = _guardrail()
_mock_post(guardrail, {"texts": ["signed [EMAIL_1]", "write to [EMAIL_1]"]})
data = {"prompt": "write to jane.doe@example.com", "suffix": "signed jane.doe@example.com"}
await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion")
assert data["suffix"] == "signed [EMAIL_1]"
@pytest.mark.asyncio
async def test_every_shape_in_one_request_is_redacted(self):
guardrail = _guardrail()
@ -682,6 +719,44 @@ class TestStreamingRehydration:
assert sent[2]["carry"] == "A-held", "choice 0 must get its own window back"
assert sent[3]["carry"] == "B-held", "choice 1 must get its own window back"
@pytest.mark.asyncio
async def test_a_choice_missing_from_the_last_chunk_still_flushes(self):
"""Held text must not be dropped because its choice ended earlier.
Choice 1 finishes and stops appearing, then the stream ends without a
finish_reason for choice 0. Flushing only the terminal chunk's choices would
discard whatever choice 1 was still holding and truncate its answer.
"""
guardrail = _guardrail(event_hook="post_call")
_mock_post(
guardrail,
{"text": "", "carry": "held-0"},
{"text": "", "carry": "held-1"},
{"text": "zero-done", "carry": ""},
{"text": "one-done", "carry": ""},
)
async def stream():
yield ModelResponseStream(
choices=[
StreamingChoices(index=0, delta=Delta(content="a")),
StreamingChoices(index=1, delta=Delta(content="b")),
]
)
yield ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content=None))])
chunks = await _drain(
guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=None, response=stream(), request_data={"messages": []}
)
)
flushed = {
choice.index: choice.delta.content for chunk in chunks for choice in chunk.choices if choice.delta.content
}
assert flushed.get(1) == "one-done", "choice 1's held text was dropped"
assert flushed.get(0) == "zero-done"
@pytest.mark.asyncio
async def test_chunks_are_forwarded_as_they_arrive(self):
"""Restoration must not buffer the stream into a single terminal chunk."""