mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
ae3cebd6ed
commit
27b6f4f2ba
2 changed files with 43 additions and 6 deletions
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue