fix(guardrails): don't repeat usage in llm_shield_proxy end-of-stream flush chunks

With n>=2 and stream_options.include_usage, the end-of-stream flush copies the last
chunk the stream carried, which is the one holding usage, so each synthetic flush
chunk repeated it and a consumer summing usage chunks counted the request twice.
The copy now drops `usage`, matching a normal mid-stream chunk. Reported by
@yucheng-berri.
This commit is contained in:
Ninad Phalak 2026-10-05 08:53:32 +00:00
parent ae3cebd6ed
commit 27b6f4f2ba
No known key found for this signature in database
2 changed files with 43 additions and 6 deletions

View file

@ -568,7 +568,7 @@ def _read_list(holder: object, name: str) -> Sequence[object]:
return tuple(_as_array(value) or ())
def _write_field(holder: object, name: str, value: str) -> None:
def _write_field(holder: object, name: str, value: object) -> None:
"""Writes one string field back into a dict or an object. Pairs with _read_field."""
if isinstance(holder, dict):
holder[name] = value
@ -1600,16 +1600,18 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
chunk = self._chunk_for_choice(last_chunk, choice_index)
if chunk is None:
continue
if isinstance(chunk.choices[0], TextChoices):
chunk.choices[0].text = text
choice: object = chunk.choices[0]
delta = _read_field(choice, "delta")
if isinstance(choice, TextChoices):
choice.text = text
elif tool_index is None:
chunk.choices[0].delta.content = text
_write_field(delta, "content", text)
else:
# The copy carried this chunk's own content and tool calls, both already
# 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 = _continuation_delta(tool_index, text)
_write_field(delta, "content", None)
_write_field(delta, "tool_calls", _continuation_delta(tool_index, text))
yield chunk
@staticmethod
@ -1632,6 +1634,11 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
# The terminal signal, if there was one, already went out with the real chunk.
kept.finish_reason = None
chunk.choices = [kept]
# So did the usage, which `stream_options.include_usage` puts on that last chunk. A
# client that sums usage across chunks would count the request twice; a mid-stream
# chunk carries no `usage` attribute at all, so the copy drops it.
if hasattr(chunk, "usage"):
del chunk.usage
return chunk
async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]:

View file

@ -31,6 +31,7 @@ from litellm.types.utils import (
StreamingChoices,
TextChoices,
TextCompletionResponse,
Usage,
)
@ -1936,3 +1937,32 @@ class TestProxyWiring:
assert reply.choices[0].message.content == "Repeat alice@example.com"
assert cache.cache_dict == {}, "the redacted request's reply must not be cached"
class TestStreamUsage:
@pytest.mark.asyncio
async def test_trailing_flush_does_not_repeat_usage(self):
"""With n>=2 and include_usage, the last chunk carries usage; a flush copied from it must not.
Any consumer that sums usage chunks would otherwise count the request twice.
"""
guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"})
both = ModelResponseStream(
choices=[
StreamingChoices(index=0, delta=Delta(content="Mail [EMAI")),
StreamingChoices(index=1, delta=Delta(content="Call [EMAI")),
]
)
usage_chunk = ModelResponseStream(
choices=[StreamingChoices(index=1, delta=Delta(content=None))],
usage=Usage(prompt_tokens=5, completion_tokens=7, total_tokens=12),
)
out = await _restore_stream(guardrail, [both, usage_chunk])
with_usage = [chunk for chunk in out if getattr(chunk, "usage", None) is not None]
assert with_usage == [out[1]], "only the provider's own usage chunk carries usage"
flushed = out[2:]
assert flushed, "the held-back text is flushed at end of stream"
assert sorted(chunk.choices[0].index for chunk in flushed) == [0, 1]
assert [chunk.choices[0].delta.content for chunk in flushed] == ["[EMAI", "[EMAI"]