mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): keep usage chunk and defer tool_calls finish_reason behind held text in incremental_diff
A stream_options.include_usage usage chunk (empty delta plus usage) was folded into the final transform round and rebuilt without its usage, so token counts and cost vanished from clients. Metadata-only chunks are now replayed after the final text flush. A terminal tool-call chunk arriving while earlier text was still held back carried finish_reason=tool_calls ahead of that text. The finish_reason is now deferred to the final text chunk whenever the choice has held text. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
7815719de7
commit
060abd263e
2 changed files with 129 additions and 8 deletions
|
|
@ -104,6 +104,10 @@ def _chunk_choices(item: object) -> Sequence[object]:
|
|||
return choices
|
||||
|
||||
|
||||
def _held_choices(held_chars_per_choice: Mapping[int, int]) -> frozenset[int]:
|
||||
return frozenset(idx for idx, held in held_chars_per_choice.items() if held > 0)
|
||||
|
||||
|
||||
def _is_redundant_scan(scan_key: "StreamingScanKey | None", last_scan_key: "StreamingScanKey | None") -> bool:
|
||||
if scan_key is None:
|
||||
return False
|
||||
|
|
@ -472,6 +476,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
emitted_text_per_choice: dict[int, str],
|
||||
holdback_per_choice: dict[int, int],
|
||||
finish_reason_per_choice: dict[int, str | None],
|
||||
held_chars_per_choice: dict[int, int],
|
||||
is_final: bool,
|
||||
) -> ModelResponseStream | None:
|
||||
"""Build the synthetic chunk carrying the newly-guardrailed deltas.
|
||||
|
|
@ -479,7 +484,9 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
For each choice, the new delta is the mutated accumulated text past what
|
||||
has already been emitted, minus a trailing holdback (forced to 0 on the
|
||||
final flush). ``emitted_text_per_choice`` holds the exact bytes already
|
||||
sent per choice and is extended in place. Returns None when there is no
|
||||
sent per choice and is extended in place; ``held_chars_per_choice`` is
|
||||
updated in place with how many mutated chars per choice are still withheld
|
||||
after this round. Returns None when there is no
|
||||
text to emit (e.g. a tool-call-only turn) or nothing new and this is not
|
||||
the final chunk.
|
||||
|
||||
|
|
@ -536,6 +543,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
holdback = 0 if is_final else max(0, holdback_per_choice.get(choice_idx, 0))
|
||||
end = max(len(already), len(text) - holdback)
|
||||
deltas[choice_idx] = text[len(already) : end]
|
||||
held_chars_per_choice[choice_idx] = len(text) - end
|
||||
|
||||
# Iterate the mutated choices (not just those in reference_chunk) so a
|
||||
# choice with pending text is never dropped for n > 1. finish_reason is
|
||||
|
|
@ -590,6 +598,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
responses_yielded: list[object],
|
||||
emitted_text_per_choice: dict[int, str],
|
||||
finish_reason_per_choice: dict[int, str | None],
|
||||
held_chars_per_choice: dict[int, int],
|
||||
is_final: bool,
|
||||
) -> AsyncGenerator[object, None]:
|
||||
"""Run one guardrail processing round and emit the resulting diff chunk.
|
||||
|
|
@ -618,6 +627,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
emitted_text_per_choice=emitted_text_per_choice,
|
||||
holdback_per_choice=sink.holdback_per_choice,
|
||||
finish_reason_per_choice=finish_reason_per_choice,
|
||||
held_chars_per_choice=held_chars_per_choice,
|
||||
is_final=is_final,
|
||||
)
|
||||
except ModifyResponseException as e:
|
||||
|
|
@ -673,6 +683,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
responses_yielded: Final[list[object]] = []
|
||||
emitted_text_per_choice: Final[dict[int, str]] = {}
|
||||
finish_reason_per_choice: Final[dict[int, str | None]] = {}
|
||||
held_chars_per_choice: Final[dict[int, int]] = {}
|
||||
chunk_counter = 0
|
||||
last_chunk: object | None = None
|
||||
|
||||
|
|
@ -688,6 +699,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
responses_yielded=responses_yielded,
|
||||
emitted_text_per_choice=emitted_text_per_choice,
|
||||
finish_reason_per_choice=finish_reason_per_choice,
|
||||
held_chars_per_choice=held_chars_per_choice,
|
||||
is_final=is_final,
|
||||
)
|
||||
|
||||
|
|
@ -724,12 +736,18 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
# finish_reason to the final text terminator (see the
|
||||
# _tool_call_passthrough_chunk docstring).
|
||||
tool_only = self._tool_call_passthrough_chunk(
|
||||
item, finish_reason_per_choice=finish_reason_per_choice
|
||||
item,
|
||||
finish_reason_per_choice=finish_reason_per_choice,
|
||||
held_choices=_held_choices(held_chars_per_choice),
|
||||
)
|
||||
responses_yielded.append(tool_only)
|
||||
yield tool_only
|
||||
continue
|
||||
|
||||
if self._is_trailing_metadata_chunk(item):
|
||||
responses_so_far.append(item)
|
||||
continue
|
||||
|
||||
chunk_counter += 1
|
||||
responses_so_far.append(item)
|
||||
last_chunk = item
|
||||
|
|
@ -773,12 +791,33 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
):
|
||||
yield out
|
||||
|
||||
if last_chunk is not None:
|
||||
async for out in _round(last_chunk, is_final=True):
|
||||
yield out
|
||||
async for out in self._emit_stream_tail(
|
||||
last_chunk=last_chunk,
|
||||
final_round=_round,
|
||||
responses_so_far=responses_so_far,
|
||||
responses_yielded=responses_yielded,
|
||||
):
|
||||
yield out
|
||||
except _StreamTerminated:
|
||||
return
|
||||
|
||||
async def _emit_stream_tail(
|
||||
self,
|
||||
*,
|
||||
last_chunk: object | None,
|
||||
final_round: Callable[[object, bool], AsyncGenerator[object, None]],
|
||||
responses_so_far: Sequence[object],
|
||||
responses_yielded: list[object],
|
||||
) -> AsyncGenerator[object, None]:
|
||||
"""Flush the held text with holdback 0, then replay metadata-only chunks
|
||||
(usage) so they land after the text and its finish_reason, as upstream sent them."""
|
||||
if last_chunk is not None:
|
||||
async for out in final_round(last_chunk, True):
|
||||
yield out
|
||||
for trailing in self._trailing_metadata_chunks(responses_so_far):
|
||||
responses_yielded.append(trailing)
|
||||
yield trailing
|
||||
|
||||
async def _inspect_full_response_for_block(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -829,6 +868,23 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
return True
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
def _is_trailing_metadata_chunk(cls, item: object) -> bool:
|
||||
"""True for a chunk that carries only stream metadata (no choices, or a
|
||||
``usage`` chunk whose deltas are empty); such chunks are replayed after
|
||||
the final text flush instead of being folded into the transform."""
|
||||
if not _chunk_choices(item):
|
||||
return True
|
||||
return (
|
||||
getattr(item, "usage", None) is not None
|
||||
and not cls._chunk_carries_text(item)
|
||||
and not cls._chunk_has_finish_reason(item)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _trailing_metadata_chunks(cls, items: Sequence[object]) -> tuple[object, ...]:
|
||||
return tuple(item for item in items if cls._is_trailing_metadata_chunk(item))
|
||||
|
||||
@staticmethod
|
||||
def _chunk_carries_text(item: object) -> bool:
|
||||
"""True if any choice in this chunk has non-empty string ``delta.content``."""
|
||||
|
|
@ -843,6 +899,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
def _tool_call_passthrough_chunk(
|
||||
item: object,
|
||||
finish_reason_per_choice: "dict[int, str | None] | None" = None,
|
||||
held_choices: frozenset[int] = frozenset(),
|
||||
) -> ModelResponseStream:
|
||||
"""Copy of a chunk carrying tool calls with all text content stripped.
|
||||
|
||||
|
|
@ -851,8 +908,9 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
transform instead). Applies per choice so an n>1 chunk mixing a text
|
||||
choice and a tool-call choice does not leak the text choice.
|
||||
|
||||
For a choice that carries BOTH text content AND tool_calls, ``finish_reason``
|
||||
is suppressed on the passthrough and recorded on
|
||||
For a choice that carries BOTH text content AND tool_calls, or whose earlier
|
||||
text is still withheld (``held_choices``), ``finish_reason`` is suppressed on
|
||||
the passthrough and recorded on
|
||||
``finish_reason_per_choice`` (when provided) so the final synthetic text
|
||||
chunk delivers it. Emitting the passthrough's ``finish_reason`` before the
|
||||
text flush would let a spec-compliant SSE client stop reading at
|
||||
|
|
@ -865,7 +923,8 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
idx = getattr(choice, "index", 0) or 0
|
||||
original_finish = getattr(choice, "finish_reason", None)
|
||||
has_text = isinstance(getattr(delta, "content", None), str) and getattr(delta, "content", "") != ""
|
||||
if has_text and original_finish is not None and finish_reason_per_choice is not None:
|
||||
text_pending = has_text or idx in held_choices
|
||||
if text_pending and original_finish is not None and finish_reason_per_choice is not None:
|
||||
finish_reason_per_choice[idx] = original_finish
|
||||
passthrough_finish: str | None = None
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -1119,6 +1119,7 @@ class TestStreamingTransform:
|
|||
emitted_text_per_choice={},
|
||||
holdback_per_choice={},
|
||||
finish_reason_per_choice={0: "stop", 1: "length"},
|
||||
held_chars_per_choice={},
|
||||
is_final=True,
|
||||
)
|
||||
|
||||
|
|
@ -1157,6 +1158,7 @@ class TestStreamingTransform:
|
|||
emitted_text_per_choice={},
|
||||
holdback_per_choice={},
|
||||
finish_reason_per_choice={},
|
||||
held_chars_per_choice={},
|
||||
is_final=False,
|
||||
)
|
||||
|
||||
|
|
@ -1179,6 +1181,7 @@ class TestStreamingTransform:
|
|||
emitted_text_per_choice={0: "My SSN is 123"},
|
||||
holdback_per_choice={},
|
||||
finish_reason_per_choice={},
|
||||
held_chars_per_choice={},
|
||||
is_final=False,
|
||||
)
|
||||
|
||||
|
|
@ -1312,6 +1315,65 @@ class TestStreamingTransform:
|
|||
assert out[1].choices[0].delta.tool_calls
|
||||
assert out[1].choices[0].finish_reason == "tool_calls"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_held_text_flushes_before_tool_call_finish_reason(self):
|
||||
"""Text still held back when a separate terminal tool-call chunk arrives is
|
||||
delivered before the stream's finish_reason, not after it."""
|
||||
guardrail = _StreamingTextGuardrail(holdback_schedule=[100, 100, 100])
|
||||
|
||||
tool_chunk = ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(
|
||||
content=None,
|
||||
tool_calls=[
|
||||
{
|
||||
"index": 0,
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
),
|
||||
finish_reason="tool_calls",
|
||||
)
|
||||
],
|
||||
)
|
||||
chunks = [_stream_chunk("let me check "), tool_chunk]
|
||||
|
||||
out = await _drive_stream(UnifiedLLMGuardrails(), guardrail, chunks)
|
||||
|
||||
finished_at = [i for i, item in enumerate(out) if item.choices[0].finish_reason is not None]
|
||||
assert finished_at == [len(out) - 1]
|
||||
assert out[-1].choices[0].finish_reason == "tool_calls"
|
||||
assert "".join(_delta_text(i) for i in out) == "LET ME CHECK "
|
||||
assert any(item.choices[0].delta.tool_calls for item in out)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"usage_choices",
|
||||
[[], [StreamingChoices(index=0, delta=Delta(), finish_reason=None)]],
|
||||
ids=["choiceless", "empty-delta"],
|
||||
)
|
||||
async def test_usage_chunk_is_forwarded_after_final_text(self, usage_choices):
|
||||
"""A trailing usage chunk (stream_options.include_usage) is delivered after
|
||||
the transformed text instead of being swallowed, whether it arrives with
|
||||
no choices or, as CustomStreamWrapper emits it, with one empty delta."""
|
||||
guardrail = _StreamingTextGuardrail(holdback_schedule=[100, 100, 100])
|
||||
usage_chunk = ModelResponseStream(
|
||||
choices=usage_choices,
|
||||
usage={"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
|
||||
)
|
||||
chunks = [_stream_chunk("hello "), _stream_chunk("world", finish_reason="stop"), usage_chunk]
|
||||
|
||||
out = await _drive_stream(UnifiedLLMGuardrails(), guardrail, chunks)
|
||||
|
||||
assert "".join(_delta_text(i) for i in out) == "HELLO WORLD"
|
||||
assert out[-1].usage.total_tokens == 5
|
||||
assert not _delta_text(out[-1])
|
||||
assert out[-2].choices[0].finish_reason == "stop"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_blocking_guardrail_is_enforced(self):
|
||||
"""A guardrail that blocks on tool calls must terminate the incremental_diff
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue