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.
This commit is contained in:
Ninad Phalak 2026-09-26 01:06:05 -05:00
parent 28697377c3
commit b9656734cd
No known key found for this signature in database
GPG key ID: 59119ED515433744
2 changed files with 89 additions and 7 deletions

View file

@ -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)

View file

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