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."""