From 403b06ea76b43ae0eeef635daea3cc1b08db1730 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Thu, 3 Sep 2026 23:28:26 -0500 Subject: [PATCH] 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. --- .../llm_shield_proxy/llm_shield_proxy.py | 68 ++++++++++++----- .../guardrail_hooks/test_llm_shield_proxy.py | 75 +++++++++++++++++++ 2 files changed, 125 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index b1ad9563eb2..47adc39aa9e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -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)``.""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py index 372bb6ed2ab..472e2606608 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -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."""