mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(guardrails): hold streamed tool calls until inspection succeeds
This commit is contained in:
parent
4d88cb480d
commit
e80c20aab7
4 changed files with 34 additions and 41 deletions
|
|
@ -62,7 +62,9 @@ Under `incremental_diff` the reply is held until the end-of-stream evaluate retu
|
|||
|
||||
`incremental_diff` covers OpenAI chat completions streaming; other surfaces fall back to `block_only`.
|
||||
|
||||
Streamed tool calls are the exception in either mode: LiteLLM forwards the tool-call deltas as they arrive and only sends the assembled call to TrustGuard once the stream ends, so a blocked call can already have reached the client. `incremental_diff` narrows that window, holding back the answer text and the turn's `finish_reason` so the block lands as an error instead of trailing a stream that looks complete. Use non-streaming requests where a tool call must be vetted before the client ever sees it.
|
||||
Under `incremental_diff`, LiteLLM buffers tool-call deltas until TrustGuard inspects the assembled response. A blocking verdict releases no tool-call arguments or finish signal. Allowed tool calls retain their original deltas and order
|
||||
|
||||
The default `block_only` mode still forwards tool-call deltas before inspection. Use `incremental_diff` or non-streaming requests when tool calls must be checked before delivery. Streamed tool-call rewrites are not supported; use non-streaming requests for transformed tool arguments
|
||||
|
||||
## References
|
||||
|
||||
|
|
|
|||
|
|
@ -708,40 +708,10 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
|
||||
try:
|
||||
async for item in response:
|
||||
# v1 transforms only text. A chunk carrying tool_calls is passed
|
||||
# through raw so function-calling turns are not dropped, but ONLY
|
||||
# its tool-call fields are forwarded: content is stripped so any
|
||||
# response text (in the same delta, or in another choice of an n>1
|
||||
# chunk) can never bypass the transform. The original chunk is kept
|
||||
# in responses_so_far so its text is still accumulated + redacted +
|
||||
# emitted as synthetic deltas, and so the guardrail inspects the
|
||||
# assembled tool calls at end of stream (see the block inspection
|
||||
# below), matching block_only. finish_reason rides on the raw
|
||||
# tool-only chunk, so it is not recorded for the text flush.
|
||||
if self._chunk_has_tool_calls(item):
|
||||
saw_tool_calls = True
|
||||
responses_so_far.append(item)
|
||||
last_chunk = item
|
||||
# Fix #3 — flush accumulated text BEFORE the tool-call
|
||||
# passthrough. Without this, a stream of text chunks that
|
||||
# hasn't yet hit a sampled round can be trailed by a
|
||||
# tool-call chunk carrying finish_reason="tool_calls"; an
|
||||
# SSE-compliant client stops reading at that finish_reason
|
||||
# and drops the end-of-stream text flush that would follow.
|
||||
if saw_text_content:
|
||||
async for out in _round(item, is_final=False):
|
||||
yield out
|
||||
# Fix #1 — pass finish_reason_per_choice into the
|
||||
# passthrough so a mixed content+tool_call chunk defers its
|
||||
# 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,
|
||||
held_choices=_held_choices(held_chars_per_choice),
|
||||
)
|
||||
responses_yielded.append(tool_only)
|
||||
yield tool_only
|
||||
continue
|
||||
|
||||
if self._is_trailing_metadata_chunk(item):
|
||||
|
|
@ -759,6 +729,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
# sampled round here would guardrail the same content twice.
|
||||
if (
|
||||
not end_of_stream_only
|
||||
and not saw_tool_calls
|
||||
and not self._chunk_has_finish_reason(item)
|
||||
and chunk_counter % sampling_rate == 0
|
||||
):
|
||||
|
|
@ -791,6 +762,21 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
):
|
||||
yield out
|
||||
|
||||
if saw_text_content:
|
||||
async for out in _round(last_chunk, is_final=False):
|
||||
yield out
|
||||
for tool_only in (
|
||||
self._tool_call_passthrough_chunk(
|
||||
buffered_item,
|
||||
finish_reason_per_choice=finish_reason_per_choice,
|
||||
held_choices=_held_choices(held_chars_per_choice),
|
||||
)
|
||||
for buffered_item in responses_so_far
|
||||
if self._chunk_has_tool_calls(buffered_item)
|
||||
):
|
||||
responses_yielded.append(tool_only)
|
||||
yield tool_only
|
||||
|
||||
async for out in self._emit_stream_tail(
|
||||
last_chunk=last_chunk,
|
||||
final_round=_round,
|
||||
|
|
|
|||
|
|
@ -110,8 +110,8 @@ async def _upstream_reply() -> AsyncIterator[ModelResponseStream]:
|
|||
FORBIDDEN_TOOL = "wire_transfer"
|
||||
|
||||
|
||||
async def _upstream_tool_call() -> AsyncIterator[ModelResponseStream]:
|
||||
for chunk in REPLY_CHUNKS[:4]:
|
||||
async def _upstream_tool_call(include_text: bool = True) -> AsyncIterator[ModelResponseStream]:
|
||||
for chunk in REPLY_CHUNKS[:4] if include_text else ():
|
||||
yield _stream_chunk(chunk)
|
||||
yield ModelResponseStream(
|
||||
model="gpt-4o-mini",
|
||||
|
|
@ -1157,21 +1157,18 @@ class TestNeuralTrustGuardrail:
|
|||
assert _deltas(received) == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_block_under_incremental_diff_leaves_the_turn_unfinished(self) -> None:
|
||||
"""A streamed tool call is only scanned once the stream ends, so the turn must not look complete.
|
||||
|
||||
Until then the answer text stays withheld and no finish_reason goes out, so a client cannot treat
|
||||
the turn as done, and the block surfaces as a 400 rather than trailing a finished-looking stream.
|
||||
"""
|
||||
@pytest.mark.parametrize("include_text", [True, False])
|
||||
async def test_tool_call_block_under_incremental_diff_sends_nothing(self, include_text: bool) -> None:
|
||||
guardrail = _guardrail(event_hook="post_call", default_on=True, streaming_transform_mode="incremental_diff")
|
||||
received: list[object] = [] # mutable-ok: collects what the client saw before the block
|
||||
with patch.object(guardrail.async_handler, "post", _tool_call_blocking_trustguard()):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _drain_into(_guardrail_stream(guardrail, _upstream_tool_call()), received)
|
||||
await _drain_into(_guardrail_stream(guardrail, _upstream_tool_call(include_text)), received)
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.detail["verdict"] == "block"
|
||||
assert _deltas(received) == []
|
||||
assert _finish_reasons(received) == []
|
||||
assert received == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_streaming_mode_leaves_the_transform_off_the_wire(self) -> None:
|
||||
|
|
|
|||
|
|
@ -1410,8 +1410,16 @@ class TestStreamingTransform:
|
|||
],
|
||||
)
|
||||
|
||||
async def upstream():
|
||||
yield tool_chunk
|
||||
|
||||
stream: Final = UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"),
|
||||
response=upstream(),
|
||||
request_data={"guardrail_to_apply": _ToolCallBlocker(), "model": "gpt-4"},
|
||||
)
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await _drive_stream(UnifiedLLMGuardrails(), _ToolCallBlocker(), [tool_chunk])
|
||||
await anext(stream)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mixed_content_and_tool_call_chunk_does_not_leak_text(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue