From f16f3e23cdaaff327a6bb4b30856c371e4c1ece7 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 5 Aug 2026 15:06:51 -0700 Subject: [PATCH] fix(tool_permission): fail closed on unverifiable SSE streams and end the turn when every tool call is denied An SSE stream that cannot be positively identified as Anthropic (no parseable message_start event) now blocks instead of passing through unscanned, closing the bypass where any raw-SSE backend skipped tool permission checks entirely. Buffered chunks are joined back into one stream before parsing, so events split across network chunk boundaries assemble correctly instead of being silently dropped. Rewrite mode now resets finish_reason to stop when no tool call survives, so the re-encoded Anthropic stream reports stop_reason end_turn and clients do not wait for a tool result that never comes --- .../guardrail_hooks/tool_permission.py | 45 ++++++++++++- .../guardrail_hooks/test_tool_permission.py | 65 +++++++++++++++++++ 2 files changed, 107 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py index 148e86ee0f0..5710af8ff3d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py @@ -734,6 +734,13 @@ class ToolPermissionGuardrail(CustomGuardrail): else: choice.message.content = "\n".join(error_messages) + if ( + not choice.message.tool_calls + and getattr(choice.message, "function_call", None) is None + and choice.finish_reason in ("tool_calls", "function_call") + ): + choice.finish_reason = "stop" + @log_guardrail_information async def async_pre_call_hook( self, @@ -878,6 +885,14 @@ class ToolPermissionGuardrail(CustomGuardrail): anthropic_response: Final = self._assemble_anthropic_stream(all_chunks) if anthropic_response is None: + if self._is_raw_sse_stream(all_chunks): + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=( + "Streamed response could not be verified for tool permissions " + "(not a parseable Anthropic SSE stream), blocking it" + ), + ) for chunk in all_chunks: yield chunk return @@ -910,18 +925,42 @@ class ToolPermissionGuardrail(CustomGuardrail): verbose_proxy_logger.debug("Tool Permission Guardrail Post-Call Hook: All tools allowed") return denied_tools + @staticmethod + def _joined_sse_stream(all_chunks: Sequence[Any]) -> str | None: + raw: Final = b"".join( + chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") + for chunk in all_chunks + if isinstance(chunk, (str, bytes)) + ) + try: + return raw.decode("utf-8") + except UnicodeDecodeError: + return None + + @staticmethod + def _has_anthropic_message_start(sse_stream: str) -> bool: + from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( + AnthropicPassthroughLoggingHandler, + ) + + return any( + (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing + and event_data.get("type") == "message_start" + for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses + ) + @staticmethod def _assemble_anthropic_stream(all_chunks: Sequence[Any]) -> ModelResponse | None: from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( AnthropicPassthroughLoggingHandler, ) - sse_chunks: Final = tuple(chunk for chunk in all_chunks if isinstance(chunk, (str, bytes))) - if not sse_chunks: + sse_stream: Final = ToolPermissionGuardrail._joined_sse_stream(all_chunks) + if sse_stream is None or not ToolPermissionGuardrail._has_anthropic_message_start(sse_stream): return None try: assembled = AnthropicPassthroughLoggingHandler._build_complete_streaming_response( # pyright: ignore[reportPrivateUsage] # the only SSE-to-ModelResponse assembler; reimplementing it here would fork the parser - all_chunks=sse_chunks, + all_chunks=(sse_stream,), litellm_logging_obj=None, # pyright: ignore[reportArgumentType] # only forwarded to stream_chunk_builder, which accepts None model="", ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py index 3fe3d8db5eb..6cfd0dde2f8 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py @@ -752,6 +752,26 @@ class TestToolPermissionGuardrail: assert isinstance(choice.message.content, str) assert "Permission denied" in choice.message.content + def test_modify_response_resets_finish_reason_when_every_tool_call_is_denied(self): + tool_call = ChatCompletionMessageToolCall(function={"name": "Read", "arguments": "{}"}, id="call_123") + response = ModelResponse( + choices=[Choices(finish_reason="tool_calls", message={"tool_calls": [tool_call], "content": ""})] + ) + denied_tools = [ + ( + tool_call, + PermissionError(tool_name="Read", rule_id="deny_read", message="Tool 'Read' denied by rule 'deny_read'"), + ) + ] + + self.guardrail._modify_response_with_permission_errors(response, denied_tools) + + choice = response.choices[0] + assert isinstance(choice, Choices) + assert choice.finish_reason == "stop", ( + "keeping finish_reason tool_calls with no surviving tool calls leaves the client waiting on a tool" + ) + def test_modify_response_with_permission_errors_filters_legacy_function_call(self): response = ModelResponse( choices=[ @@ -1190,3 +1210,48 @@ class TestToolPermissionGuardrailAnthropicMessages: body = b"".join(c if isinstance(c, bytes) else str(c).encode() for c in out).decode() assert '"type": "tool_use"' not in body, "denied tool_use must not survive into the rewritten stream" assert "Permission denied" in body + assert '"stop_reason": "end_turn"' in body, ( + "dropping every tool_use must end the turn, or the client waits for a tool result that never comes" + ) + assert '"stop_reason": "tool_use"' not in body + + def _resplit(self, chunks, size=7): + joined = b"".join(chunks) + return [joined[i : i + size] for i in range(0, len(joined), size)] + + @pytest.mark.asyncio + async def test_denied_tool_use_is_caught_when_sse_events_are_split_across_chunk_boundaries(self): + with patch.object(self.blocking, "should_run_guardrail", return_value=True): + with pytest.raises(GuardrailRaisedException) as exc_info: + await self._drain(self.blocking, self._resplit(self._sse_chunks("Read"))) + + assert "deny_read" in str(exc_info.value), ( + "a stream split mid-event must still assemble and hit the rule, not fail as unparseable" + ) + + @pytest.mark.asyncio + async def test_allowed_stream_split_across_chunk_boundaries_is_passed_through_verbatim(self): + chunks = self._resplit(self._sse_chunks("Bash")) + + with patch.object(self.blocking, "should_run_guardrail", return_value=True): + out = await self._drain(self.blocking, chunks) + + assert out == chunks + + @pytest.mark.asyncio + async def test_non_anthropic_sse_stream_fails_closed(self): + gemini_chunks = [ + b'data: {"candidates": [{"content": {"parts": [{"functionCall": ' + b'{"name": "run_shell", "args": {"command": "ls"}}}], "role": "model"}}]}\n\n', + b'data: {"candidates": [{"content": {"parts": [{"text": "done"}]}, "finishReason": "STOP"}]}\n\n', + ] + + with patch.object(self.blocking, "should_run_guardrail", return_value=True): + with pytest.raises(GuardrailRaisedException): + await self._drain(self.blocking, gemini_chunks) + + @pytest.mark.asyncio + async def test_unparseable_sse_stream_fails_closed(self): + with patch.object(self.blocking, "should_run_guardrail", return_value=True): + with pytest.raises(GuardrailRaisedException): + await self._drain(self.blocking, [b"data: not-json\n\n", b"event: weird\n\n"])