fix(guardrails): finish stream checks before releasing tool calls

This commit is contained in:
albertbausili 2026-09-21 15:28:09 +02:00
parent e80c20aab7
commit 080b5a486a
2 changed files with 34 additions and 6 deletions

View file

@ -765,7 +765,7 @@ class UnifiedLLMGuardrails(CustomLogger):
if saw_text_content:
async for out in _round(last_chunk, is_final=False):
yield out
for tool_only in (
tool_chunks: Final = tuple(
self._tool_call_passthrough_chunk(
buffered_item,
finish_reason_per_choice=finish_reason_per_choice,
@ -773,9 +773,31 @@ class UnifiedLLMGuardrails(CustomLogger):
)
for buffered_item in responses_so_far
if self._chunk_has_tool_calls(buffered_item)
):
)
async def checked_tail() -> AsyncGenerator[object, None]:
try:
async for tail_chunk in self._emit_stream_tail(
last_chunk=last_chunk,
final_round=_round,
responses_so_far=responses_so_far,
responses_yielded=responses_yielded,
):
yield tail_chunk
except _StreamTerminated as exc:
yield exc
tail_chunks: Final = tuple([chunk async for chunk in checked_tail()])
if tail_chunks and isinstance(tail_chunks[-1], _StreamTerminated):
for error_chunk in tail_chunks[:-1]:
yield error_chunk
return
for tool_only in tool_chunks:
responses_yielded.append(tool_only)
yield tool_only
for tail_chunk in tail_chunks:
yield tail_chunk
return
async for out in self._emit_stream_tail(
last_chunk=last_chunk,

View file

@ -1375,14 +1375,15 @@ class TestStreamingTransform:
assert out[-2].choices[0].finish_reason == "stop"
@pytest.mark.asyncio
async def test_tool_call_blocking_guardrail_is_enforced(self):
@pytest.mark.parametrize(("content", "allowed_scans"), [(None, 0), ("proposal", 0), ("proposal", 1)])
async def test_tool_call_blocking_guardrail_is_enforced(self, content: str | None, allowed_scans: int):
"""A guardrail that blocks on tool calls must terminate the incremental_diff
stream: tool calls go through the block decision, not bypass it."""
from litellm.exceptions import GuardrailRaisedException
class _ToolCallBlocker(_StreamingTextGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
if input_type == "response" and inputs.get("tool_calls"):
if input_type == "response" and self.response_calls >= allowed_scans:
raise GuardrailRaisedException(
guardrail_name="tc-block",
message="blocked tool call",
@ -1395,7 +1396,7 @@ class TestStreamingTransform:
StreamingChoices(
index=0,
delta=Delta(
content=None,
content=content,
tool_calls=[
{
"index": 0,
@ -1418,8 +1419,13 @@ class TestStreamingTransform:
response=upstream(),
request_data={"guardrail_to_apply": _ToolCallBlocker(), "model": "gpt-4"},
)
async def consume_checked_stream() -> None:
async for chunk in stream:
assert isinstance(chunk, ModelResponseStream)
assert all(not choice.delta.tool_calls and choice.finish_reason is None for choice in chunk.choices)
with pytest.raises(GuardrailRaisedException):
await anext(stream)
await consume_checked_stream()
@pytest.mark.asyncio
async def test_mixed_content_and_tool_call_chunk_does_not_leak_text(self):