From b9656734cdf549fa7c442570a2f4c9049544c9c4 Mon Sep 17 00:00:00 2001 From: Ninad Phalak Date: Sat, 26 Sep 2026 01:06:05 -0500 Subject: [PATCH] fix(guardrails): build llm_shield_proxy stream deltas without new mutable literals The lint job's LIT002 budget gate failed on this PR: the file added 11 mutable-collection constructions and the tree sits at its limit. Build the index-only tool-call continuation in one helper, keep read-only inputs as tuples, and annotate the lists the delta and texts fields require. Adds tests for the two tool-call flush paths the refactor touches, which had no coverage: held arguments landing in the finish_reason chunk next to that chunk's own fragment, and the trailing flush of a stream that ends without a finish_reason. --- .../llm_shield_proxy/llm_shield_proxy.py | 22 ++++-- .../guardrail_hooks/test_llm_shield_proxy.py | 74 +++++++++++++++++++ 2 files changed, 89 insertions(+), 7 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 9877f746730..5ae8a434793 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 @@ -290,6 +290,14 @@ def _carry_sort_key(key: tuple) -> tuple: return (choice_index, -1 if tool_index is None else tool_index) +def _continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]: + """A `tool_calls` delta carrying `text` as an index-only continuation. + + Clients concatenate tool-call fragments by index, so no id or name is needed. + """ + return [{"index": tool_index, "function": {"arguments": text}}] # mutable-ok: delta.tool_calls is a list. + + class LLMShieldProxyGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -761,10 +769,10 @@ class LLMShieldProxyGuardrail(CustomGuardrail): if tool_index is None: delta.content = text else: - continuations.append({"index": tool_index, "function": {"arguments": text}}) + continuations.extend(_continuation_delta(tool_index, text)) if continuations: - existing: Final[list] = list(getattr(delta, "tool_calls", None) or []) - delta.tool_calls = existing + continuations + existing: Final = tuple(getattr(delta, "tool_calls", None) or ()) + delta.tool_calls = [*existing, *continuations] # mutable-ok: delta.tool_calls is a list. async def _flush_trailing( self, last_chunk: Any, carries: _CarryWindows, session_id: str @@ -799,7 +807,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): # delivered. Replace rather than append, and drop the content, or the # client sees them twice. chunk.choices[0].delta.content = None - chunk.choices[0].delta.tool_calls = [{"index": tool_index, "function": {"arguments": text}}] + chunk.choices[0].delta.tool_calls = _continuation_delta(tool_index, text) yield chunk @staticmethod @@ -860,8 +868,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail): whichever path ran. The request side is left to the native pre-call hook, because redacting it here as well would redact it twice. """ - text_list: Final[list] = list(inputs.get("texts") or ()) - tool_calls: Final[list] = list(inputs.get("tool_calls") or ()) if input_type == "response" else [] + text_list: Final = tuple(inputs.get("texts") or ()) + tool_calls: Final = tuple(inputs.get("tool_calls") or ()) if input_type == "response" else () if not text_list and not tool_calls: return inputs @@ -882,7 +890,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): if input_type == "request" else await self._rehydrate(tuple(spans), self._session_id(request_data)) ) - restored_values: Final[list] = list(replaced) + restored_values: Final[list] = list(replaced) # mutable-ok: sliced into the texts list. for write, replacement in zip(writers, restored_values[len(text_list) :]): write(replacement) 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 bfa8c04c603..0197b83c9f0 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 @@ -53,6 +53,19 @@ async def _drain(generator) -> list: return [chunk async for chunk in generator] +def _tool_chunk(arguments: str, finish_reason: str | None = None) -> ModelResponseStream: + """One streamed fragment of tool call 0's arguments.""" + tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "send", "arguments": arguments}} + return ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[tool_call]), finish_reason=finish_reason)] + ) + + +def _field(holder: object, name: str) -> object: + """Reads a field from a dict or a model; the guardrail emits both shapes.""" + return holder.get(name) if isinstance(holder, dict) else getattr(holder, name) + + def test_llm_shield_guardrail_config(monkeypatch: pytest.MonkeyPatch): """Should register through init_guardrails_v2 like any other provider.""" monkeypatch.setattr(litellm, "guardrail_name_config_map", {}) @@ -893,6 +906,67 @@ class TestStreamingRehydration: 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_held_tool_arguments_land_in_the_finishing_chunk(self): + """A client parses tool arguments on finish_reason, so the flush must ride that chunk. + + The finishing chunk also carries its own fragment for the same tool call. That + entry has to survive, with the held text appended after it as a continuation. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": '{"to": "', "carry": "[EMA"}, + {"text": "", "carry": '[EMAIL_1]"}'}, + {"text": 'a@example.com"}', "carry": ""}, + ) + + async def stream(): + yield _tool_chunk('{"to": "[EMA') + yield _tool_chunk('IL_1]"}', finish_reason="tool_calls") + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 2, "the flush must not arrive after the finish_reason chunk" + final_calls = chunks[1].choices[0].delta.tool_calls + assert len(final_calls) == 2, "the finishing chunk's own fragment was dropped" + assert _field(final_calls[1], "index") == 0 + arguments = "".join( + _field(_field(call, "function"), "arguments") or "" + for chunk in chunks + for call in chunk.choices[0].delta.tool_calls + ) + assert json.loads(arguments) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_held_tool_arguments_flush_when_the_stream_ends_unfinished(self): + """No finish_reason at all: a trailing chunk carries the held arguments alone.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": '{"to": "', "carry": "[EMA"}, + {"text": 'a@example.com"}', "carry": ""}, + ) + + async def stream(): + yield _tool_chunk('{"to": "[EMA') + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 2 + trailing = chunks[1].choices[0].delta.tool_calls + assert trailing == [{"index": 0, "function": {"arguments": 'a@example.com"}'}}], ( + "the copied chunk's own fragment was already delivered and must not repeat" + ) + @pytest.mark.asyncio async def test_chunks_are_forwarded_as_they_arrive(self): """Restoration must not buffer the stream into a single terminal chunk."""