diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py index e20f0b320b9..819e743c209 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py @@ -14,6 +14,7 @@ from litellm.integrations.custom_guardrail import ( CustomGuardrail, log_guardrail_information, ) +from litellm.llms.base_llm.guardrail_translation.utils import anthropic_tool_names from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.callback_utils import ( add_guardrail_to_applied_guardrails_header, @@ -40,6 +41,7 @@ from litellm.types.utils import ( ) GUARDRAIL_NAME: Final = "tool_permission" +_RESPONSES_CALL_TYPES: Final = frozenset({"responses", "aresponses", "_aresponses_websocket"}) def _object_mapping(value: object) -> Mapping[str, object] | None: @@ -605,27 +607,33 @@ class ToolPermissionGuardrail(CustomGuardrail): if not any(_is_tool_use_block(block) for block in kept_blocks): response["stop_reason"] = "end_turn" # rebind-ok: dropping every tool_use ends the turn - def _get_request_tool_name(self, tool: object) -> tuple[str | None, str | None]: + def _get_request_tool_targets( + self, tool: object, call_type: CallTypesLiteral + ) -> tuple[tuple[str, str | None], ...]: tool_type: Final = self._get_mapping_value(tool, "type") - if tool_type != "function": - return None, tool_type - - function: Final = self._get_mapping_value(tool, "function") - tool_name: Final = self._get_mapping_value(function, "name") - return tool_name, tool_type + normalized_type: Final = ( + tool_type if call_type in _RESPONSES_CALL_TYPES or tool_type not in (None, "custom") else "function" + ) + return tuple((name, normalized_type) for name in anthropic_tool_names(tool)) def _get_legacy_function_name(self, function: object) -> str | None: return self._get_mapping_value(function, "name") - def _get_named_tool_choice(self, data: dict) -> str | None: + def _get_named_tool_choice(self, data: Mapping[str, object]) -> str | None: tool_choice: Final = data.get("tool_choice") if not tool_choice or tool_choice in ("auto", "none", "required"): return None if isinstance(tool_choice, str): return tool_choice - if self._get_mapping_value(tool_choice, "type") != "function": + if self._get_mapping_value(tool_choice, "type") not in ("tool", "function"): return None - return self._get_mapping_value(self._get_mapping_value(tool_choice, "function"), "name") + function_name: Final = self._get_mapping_value(self._get_mapping_value(tool_choice, "function"), "name") + return function_name or self._get_mapping_value(tool_choice, "name") + + @staticmethod + def _is_anthropic_tool_choice(data: Mapping[str, object]) -> bool: + tool_choice: Final = _object_mapping(data.get("tool_choice")) + return tool_choice is not None and tool_choice.get("type") == "tool" def _get_named_function_call(self, data: dict) -> str | None: function_call: Final = data.get("function_call") @@ -635,13 +643,11 @@ class ToolPermissionGuardrail(CustomGuardrail): return function_call return self._get_mapping_value(function_call, "name") - def _collect_request_tools(self, data: dict) -> list[tuple[str, str | None]]: + def _collect_request_tools(self, data: dict, call_type: CallTypesLiteral) -> list[tuple[str, str | None]]: request_tools: Final[list[tuple[str, str | None]]] = [] for tool in data.get("tools") or []: - tool_name, tool_type = self._get_request_tool_name(tool) - if tool_name is not None: - request_tools.append((tool_name, tool_type)) + request_tools.extend(self._get_request_tool_targets(tool, call_type)) for function in data.get("functions") or []: function_name = self._get_legacy_function_name(function) @@ -681,13 +687,9 @@ class ToolPermissionGuardrail(CustomGuardrail): tools: Final[list[ChatCompletionToolParam] | None] = data.get("tools") if tools is not None: - new_tools: Final = [] - for tool in tools: - tool_name, tool_type = self._get_request_tool_name(tool) - if tool_type == "function" and tool_name in error_tool_names: - continue - new_tools.append(tool) - data["tools"] = new_tools + data["tools"] = [ + tool for tool in tools if not any(name in error_tool_names for name in anthropic_tool_names(tool)) + ] functions: Final = data.get("functions") if functions is not None: @@ -697,7 +699,7 @@ class ToolPermissionGuardrail(CustomGuardrail): named_tool_choice: Final = self._get_named_tool_choice(data) if named_tool_choice in error_tool_names: - data["tool_choice"] = "none" + data["tool_choice"] = {"type": "none"} if self._is_anthropic_tool_choice(data) else "none" named_function_call: Final = self._get_named_function_call(data) if named_function_call in error_tool_names: @@ -807,7 +809,7 @@ class ToolPermissionGuardrail(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=event_type) is not True: return data - new_tools: Final = self._collect_request_tools(data) + new_tools: Final = self._collect_request_tools(data, call_type) if not new_tools: verbose_proxy_logger.debug( "Tool Permission Guardrail: not running guardrail. No tools or functions in data" 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 427fa43ffd5..ec5f6efefba 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 @@ -5,6 +5,7 @@ Unit tests for Tool Permission Guardrail (OpenAI tool_calls semantics) import json import logging import re +from typing import Literal from unittest.mock import patch import pytest @@ -24,6 +25,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( PermissionError, ) from litellm.types.utils import ( + CallTypesLiteral, ChatCompletionMessageToolCall, Choices, ModelResponse, @@ -1177,6 +1179,117 @@ class TestToolPermissionGuardrailAnthropicMessages: def _tool_use(self, name, tool_id="tu_1"): return {"type": "tool_use", "id": tool_id, "name": name, "input": {"command": "ls"}} + def _always_on_pre_call(self, on_disallowed_action: Literal["block", "rewrite"]) -> ToolPermissionGuardrail: + return ToolPermissionGuardrail( + guardrail_name=f"anthropic-pre-call-{on_disallowed_action}", + rules=self.rules, + default_action="deny", + on_disallowed_action=on_disallowed_action, + event_hook=GuardrailEventHooks.pre_call, + default_on=True, + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "denied_tool", + [ + {"name": "Read", "input_schema": {"type": "object", "properties": {}}}, + {"type": "custom", "name": "Read", "input_schema": {"type": "object", "properties": {}}}, + {"type": "function", "name": "Read", "parameters": {"type": "object", "properties": {}}}, + ], + ids=["anthropic", "anthropic_custom_type", "responses_api_flat_function"], + ) + async def test_pre_call_blocks_denied_request_tool_in_flat_format(self, denied_tool: dict[str, object]) -> None: + data = {"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "hi"}], "tools": [denied_tool]} + + with pytest.raises(HTTPException) as excinfo: + await self._always_on_pre_call("block").async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=DualCache(default_in_memory_ttl=1), + data=data, + call_type="anthropic_messages", + ) + + assert excinfo.value.status_code == 400 + assert excinfo.value.detail["detection_message"] == "Tool 'Read' denied by rule 'deny_read'" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("rule_tool_type", "call_type"), + [(r"^custom$", "responses"), (r"^function$", "anthropic_messages")], + ids=["responses_custom_stays_custom", "anthropic_custom_reads_as_function"], + ) + async def test_pre_call_tool_type_rules_follow_the_request_format( + self, rule_tool_type: str, call_type: CallTypesLiteral + ) -> None: + guardrail = ToolPermissionGuardrail( + guardrail_name=f"tool-type-{call_type}", + rules=[{"id": "deny_type", "tool_type": rule_tool_type, "decision": "deny"}], + default_action="allow", + on_disallowed_action="block", + event_hook=GuardrailEventHooks.pre_call, + default_on=True, + ) + data = { + "model": "claude-sonnet-4-5", + "messages": [{"role": "user", "content": "hi"}], + "tools": [{"type": "custom", "name": "apply_patch", "input_schema": {"type": "object", "properties": {}}}], + } + + with pytest.raises(HTTPException) as excinfo: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=DualCache(default_in_memory_ttl=1), + data=data, + call_type=call_type, + ) + + assert excinfo.value.status_code == 400 + assert excinfo.value.detail["detection_message"] == "Tool 'apply_patch' denied by rule 'deny_type'" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("tool_shape", "tool_choice", "call_type", "expected_tool_choice"), + [ + ( + {"input_schema": {"type": "object", "properties": {}}}, + {"type": "tool", "name": "Read"}, + "anthropic_messages", + {"type": "none"}, + ), + ( + {"type": "function", "parameters": {"type": "object", "properties": {}}}, + {"type": "function", "name": "Read"}, + "responses", + "none", + ), + ], + ids=["anthropic", "responses_api"], + ) + async def test_pre_call_rewrite_strips_denied_flat_tool_and_forced_choice( + self, + tool_shape: dict[str, object], + tool_choice: dict[str, str], + call_type: CallTypesLiteral, + expected_tool_choice: dict[str, str] | str, + ) -> None: + data = { + "model": "claude-sonnet-4-5", + "messages": [{"role": "user", "content": "hi"}], + "tools": [{"name": "Bash", **tool_shape}, {"name": "Read", **tool_shape}], + "tool_choice": tool_choice, + } + + result = await self._always_on_pre_call("rewrite").async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=DualCache(default_in_memory_ttl=1), + data=data, + call_type=call_type, + ) + + assert [tool["name"] for tool in result["tools"]] == ["Bash"] + assert result["tool_choice"] == expected_tool_choice + @pytest.mark.asyncio async def test_denied_anthropic_tool_use_is_blocked(self): response = self._response({"type": "text", "text": "reading"}, self._tool_use("Read"))