From 080b5a486ace0984a0bd1c98fc122d393d8ab3dc Mon Sep 17 00:00:00 2001 From: albertbausili Date: Mon, 21 Sep 2026 15:28:09 +0200 Subject: [PATCH] fix(guardrails): finish stream checks before releasing tool calls --- .../unified_guardrail/unified_guardrail.py | 26 +++++++++++++++++-- .../test_unified_guardrail.py | 14 +++++++--- 2 files changed, 34 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index e5cb8292f3f..6be65c2b686 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -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, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index ca507db791d..c953dea1fb4 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -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):