mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
28697377c3
commit
b9656734cd
2 changed files with 89 additions and 7 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue