fix(guardrails): hold streamed tool calls until inspection succeeds

This commit is contained in:
albertbausili 2026-09-21 14:48:45 +02:00
parent 4d88cb480d
commit e80c20aab7
4 changed files with 34 additions and 41 deletions

View file

@ -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

View file

@ -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,

View file

@ -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:

View file

@ -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):