mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(guardrails): check Anthropic-format request tools in tool_permission pre_call
The pre_call hook only collected tools shaped as {"type": "function", "function": {...}},
so flat tools from /v1/messages ({"name", "input_schema"}) and /v1/responses
({"type": "function", "name"}) were never checked and deny rules silently passed.
Reuse the shared anthropic_tool_names helper so every request tool name is checked,
and handle Anthropic tool_choice {"type": "tool"} in both block and rewrite modes.
Fixes #40583
This commit is contained in:
parent
9496f16f12
commit
739ff2f2da
2 changed files with 69 additions and 20 deletions
|
|
@ -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,
|
||||
|
|
@ -605,28 +606,32 @@ 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) -> 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 = "function" if tool_type in (None, "custom") else tool_type
|
||||
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":
|
||||
choice_type: Final[object] = self._get_mapping_value(tool_choice, "type")
|
||||
if choice_type == "tool":
|
||||
return self._get_mapping_value(tool_choice, "name")
|
||||
if choice_type != "function":
|
||||
return None
|
||||
return self._get_mapping_value(self._get_mapping_value(tool_choice, "function"), "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")
|
||||
if not function_call or function_call in ("auto", "none"):
|
||||
|
|
@ -639,9 +644,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
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))
|
||||
|
||||
for function in data.get("functions") or []:
|
||||
function_name = self._get_legacy_function_name(function)
|
||||
|
|
@ -681,13 +684,11 @@ 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 self._get_request_tool_targets(tool))
|
||||
]
|
||||
|
||||
functions: Final = data.get("functions")
|
||||
if functions is not None:
|
||||
|
|
@ -697,7 +698,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:
|
||||
|
|
|
|||
|
|
@ -1177,6 +1177,54 @@ class TestToolPermissionGuardrailAnthropicMessages:
|
|||
def _tool_use(self, name, tool_id="tu_1"):
|
||||
return {"type": "tool_use", "id": tool_id, "name": name, "input": {"command": "ls"}}
|
||||
|
||||
@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):
|
||||
data = {"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "hi"}], "tools": [denied_tool]}
|
||||
|
||||
with patch.object(self.blocking, "should_run_guardrail", return_value=True):
|
||||
with pytest.raises(HTTPException) as excinfo:
|
||||
await self.blocking.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
|
||||
async def test_pre_call_rewrite_strips_denied_anthropic_tool_and_forced_choice(self):
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"tools": [
|
||||
{"name": "Bash", "input_schema": {"type": "object", "properties": {}}},
|
||||
{"name": "Read", "input_schema": {"type": "object", "properties": {}}},
|
||||
],
|
||||
"tool_choice": {"type": "tool", "name": "Read"},
|
||||
}
|
||||
|
||||
with patch.object(self.rewriting, "should_run_guardrail", return_value=True):
|
||||
result = await self.rewriting.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(default_in_memory_ttl=1),
|
||||
data=data,
|
||||
call_type="anthropic_messages",
|
||||
)
|
||||
|
||||
assert [tool["name"] for tool in result["tools"]] == ["Bash"]
|
||||
assert result["tool_choice"] == {"type": "none"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_denied_anthropic_tool_use_is_blocked(self):
|
||||
response = self._response({"type": "text", "text": "reading"}, self._tool_use("Read"))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue