This commit is contained in:
Meryem Sakin 2026-09-27 22:16:40 +00:00 • committed by GitHub
commit 163914ca23
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 138 additions and 23 deletions

View file

@ -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"

View file

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