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"])