diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index 39c0930d7cf..f84f0f429aa 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -77,6 +77,9 @@ _Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's p # _locate_request_texts, which freezes it into a tuple before returning. _SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. +# Sliding windows keyed by streaming choice index, threaded through one stream. +_CarryWindows: TypeAlias = dict # mutable-ok: per-choice windows advanced in place. + # A caller-owned list whose entries are rewritten in place, such as a Completions # `prompt` sent as an array of strings. MutableSeq: TypeAlias = list # mutable-ok: the request payload's own list. @@ -163,6 +166,12 @@ def _collect_responses_fields(data: MutableRequest, slots: _SlotSink) -> None: _collect(item, "output", slots) +def _choice_index(choice: object) -> int: + """Streaming choices are matched across chunks by their index.""" + index: Final = getattr(choice, "index", 0) + return index if isinstance(index, int) else 0 + + class LLMShieldGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -441,10 +450,10 @@ class LLMShieldGuardrail(CustomGuardrail): ) -> AsyncGenerator[Any, None]: """Restores original values incrementally, without buffering the stream. - The carry-over window is a local of this generator, so it is scoped to one - stream and cannot leak between concurrent requests. LLM Shield returns the - text that is safe to emit now plus the trailing characters it is still - holding, which are sent back with the next delta. + Each choice is its own token stream, so the sliding window is tracked per + choice index. One shared window would splice the characters held back for + one choice onto the next. The windows are locals of this generator, so they + are scoped to a single stream and cannot leak between concurrent requests. """ if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True: async for chunk in response: @@ -452,39 +461,61 @@ class LLMShieldGuardrail(CustomGuardrail): return session_id: Final = self._session_id(request_data) - carry = "" # rebind-ok: the sliding window advances with every delta. + carries: Final[dict] = {} # mutable-ok: per-choice windows, local to this stream. last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. async for chunk in response: last_chunk = chunk - delta = self._stream_delta(chunk) - text = getattr(delta, "content", None) if delta is not None else None - is_final = self._is_final_chunk(chunk) - - if delta is None or not isinstance(text, str) or not text: - # Nothing to restore in this chunk, but a final chunk still has to - # flush whatever the window is holding. - if is_final and carry and delta is not None: - emitted, carry = await self._stream_step("", carry, True, session_id) - if emitted: - delta.content = emitted - yield chunk - continue - - emitted, carry = await self._stream_step(text, carry, is_final, session_id) - delta.content = emitted + for choice in getattr(chunk, "choices", None) or (): + await self._restore_choice(choice, carries, session_id) yield chunk # A stream that ended without a finish_reason can still leave text held back. - if carry and last_chunk is not None: - flushed: Final = await self._stream_step("", carry, True, session_id) - trailing_text, carry = flushed # rebind-ok: window advances. - if trailing_text: - trailing: Final = last_chunk.model_copy(deep=True) - trailing_delta: Final = self._stream_delta(trailing) - if trailing_delta is not None: - trailing_delta.content = trailing_text - yield trailing + 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): + yield trailing + + async def _restore_choice(self, choice: Any, carries: _CarryWindows, session_id: str) -> None: + """Restores one choice's delta, advancing that choice's own window.""" + delta: Final = getattr(choice, "delta", None) + if delta is None: + return + index: Final = _choice_index(choice) + carry: Final = carries.get(index, "") + text: Final = getattr(delta, "content", None) + is_final: Final = bool(getattr(choice, "finish_reason", None)) + + if not isinstance(text, str) or not text: + # Nothing to restore here, but a final chunk still has to flush the window. + if is_final and carry: + flushed, flushed_carry = await self._stream_step("", carry, True, session_id) + carries[index] = flushed_carry # rebind-ok: this choice's window advances. + if flushed: + delta.content = flushed + return + + emitted, remaining = await self._stream_step(text, carry, is_final, session_id) + 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, "") + 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 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)``.""" @@ -503,20 +534,6 @@ class LLMShieldGuardrail(CustomGuardrail): ) return emitted, remaining - @staticmethod - def _stream_delta(chunk: object) -> Any: - choices: Final = getattr(chunk, "choices", None) - if not choices: - return None - return getattr(choices[0], "delta", None) - - @staticmethod - def _is_final_chunk(chunk: object) -> bool: - choices: Final = getattr(chunk, "choices", None) - if not choices: - return False - return bool(getattr(choices[0], "finish_reason", None)) - # --- unified API (powers the UI "Test guardrail" button) ----------------------- @log_guardrail_information diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py index f16584fbd1d..e9cc1a36b07 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -587,6 +587,73 @@ class TestStreamingRehydration: assert mock.call_args_list[1].kwargs["json"]["carry"] == "hold" assert mock.call_args_list[1].kwargs["json"]["final"] is True + @pytest.mark.asyncio + async def test_every_choice_is_restored(self): + """With n>1 a later choice must not be handed back still holding a placeholder.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": "first@example.com", "carry": ""}, + {"text": "second@example.com", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="[EMAIL_1]"), finish_reason="stop"), + StreamingChoices(index=1, delta=Delta(content="[EMAIL_2]"), finish_reason="stop"), + ] + ) + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + restored = [choice.delta.content for choice in chunks[0].choices] + assert restored == ["first@example.com", "second@example.com"] + + @pytest.mark.asyncio + async def test_choice_windows_do_not_cross_contaminate(self): + """Each choice is its own token stream, so each carries its own window. + + One shared window would send the characters held back for choice 0 up + against choice 1's next delta and splice the two streams together. + """ + guardrail = _guardrail(event_hook="post_call") + mock = _mock_post( + guardrail, + {"text": "", "carry": "A-held"}, + {"text": "", "carry": "B-held"}, + {"text": "a-done", "carry": ""}, + {"text": "b-done", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a1")), + StreamingChoices(index=1, delta=Delta(content="b1")), + ] + ) + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a2"), finish_reason="stop"), + StreamingChoices(index=1, delta=Delta(content="b2"), finish_reason="stop"), + ] + ) + + await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + sent = [call.kwargs["json"] for call in mock.call_args_list] + 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_chunks_are_forwarded_as_they_arrive(self): """Restoration must not buffer the stream into a single terminal chunk."""