chore(guardrails): tighten tool permission checks

This commit is contained in:
user 2026-05-01 00:55:04 -07:00
parent eab0075353
commit 150a34f2b0
2 changed files with 371 additions and 36 deletions

View file

@ -225,10 +225,10 @@ class ToolPermissionGuardrail(CustomGuardrail):
def _parse_tool_call_arguments(
self, tool_call: ChatCompletionMessageToolCall
) -> Dict[str, Any]:
) -> tuple[Optional[Dict[str, Any]], Optional[str]]:
arguments = getattr(tool_call.function, "arguments", None)
if not arguments:
return {}
return None, "missing arguments"
parsed_arguments: Any = {}
try:
@ -236,22 +236,24 @@ class ToolPermissionGuardrail(CustomGuardrail):
parsed_arguments = json.loads(arguments)
elif isinstance(arguments, dict):
parsed_arguments = arguments
except json.JSONDecodeError as exc:
else:
return None, "arguments must be a JSON object"
except (json.JSONDecodeError, TypeError) as exc:
verbose_proxy_logger.warning(
"Tool Permission Guardrail: Failed to decode arguments for tool %s: %s",
tool_call.function.name,
exc,
)
return {}
return None, "arguments could not be parsed"
if isinstance(parsed_arguments, dict):
return parsed_arguments
return parsed_arguments, None
verbose_proxy_logger.debug(
"Tool Permission Guardrail: Ignoring non-dict arguments for tool %s",
"Tool Permission Guardrail: Rejecting non-dict arguments for tool %s",
tool_call.function.name,
)
return {}
return None, "arguments must be a JSON object"
def _collect_argument_paths(
self,
@ -331,10 +333,21 @@ class ToolPermissionGuardrail(CustomGuardrail):
continue
if rule.allowed_param_patterns and should_check_params:
arguments = self._parse_tool_call_arguments(tool_call)
arguments, parse_error = self._parse_tool_call_arguments(tool_call)
if parse_error:
default_message = f"Tool '{tool_identifier}' {parse_error} required by rule '{rule.id}'"
message = self.render_violation_message(
default=default_message,
context={"tool_name": tool_identifier, "rule_id": rule.id},
)
return False, rule.id, message
if not arguments:
last_pattern_failure_msg = f"Tool '{tool_identifier}' is missing arguments required by rule '{rule.id}'"
continue
default_message = f"Tool '{tool_identifier}' is missing arguments required by rule '{rule.id}'"
message = self.render_violation_message(
default=default_message,
context={"tool_name": tool_identifier, "rule_id": rule.id},
)
return False, rule.id, message
patterns_match, failure_message = self._patterns_match_for_rule(
arguments=arguments,
@ -365,6 +378,33 @@ class ToolPermissionGuardrail(CustomGuardrail):
)
return is_allowed, None, message
@staticmethod
def _get_mapping_value(item: Any, key: str) -> Any:
if isinstance(item, dict):
return item.get(key)
return getattr(item, key, None)
@staticmethod
def _legacy_function_call_id(choice_index: int) -> str:
return f"legacy_function_call_{choice_index}"
def _legacy_function_call_to_tool_call(
self, function_call: Any, choice_index: int
) -> Optional[ChatCompletionMessageToolCall]:
if function_call is None:
return None
function_name = self._get_mapping_value(function_call, "name")
arguments = self._get_mapping_value(function_call, "arguments") or ""
if not function_name:
return None
return ChatCompletionMessageToolCall(
id=self._legacy_function_call_id(choice_index),
type="function",
function={"name": function_name, "arguments": arguments},
)
def _extract_tool_calls_from_response(
self, response: ModelResponse
) -> List[ChatCompletionMessageToolCall]:
@ -379,13 +419,72 @@ class ToolPermissionGuardrail(CustomGuardrail):
"""
tool_calls = []
for choice in response.choices:
for choice_index, choice in enumerate(response.choices):
if isinstance(choice, Choices):
for tool in choice.message.tool_calls or []:
tool_calls.append(tool)
legacy_tool_call = self._legacy_function_call_to_tool_call(
getattr(choice.message, "function_call", None), choice_index
)
if legacy_tool_call is not None:
tool_calls.append(legacy_tool_call)
return tool_calls
def _get_request_tool_name(self, tool: Any) -> tuple[Optional[str], Optional[str]]:
tool_type = self._get_mapping_value(tool, "type")
if tool_type != "function":
return None, tool_type
function = self._get_mapping_value(tool, "function")
tool_name = self._get_mapping_value(function, "name")
return tool_name, tool_type
def _get_legacy_function_name(self, function: Any) -> Optional[str]:
return self._get_mapping_value(function, "name")
def _get_named_tool_choice(self, data: dict) -> Optional[str]:
tool_choice = 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":
return None
return self._get_mapping_value(
self._get_mapping_value(tool_choice, "function"), "name"
)
def _get_named_function_call(self, data: dict) -> Optional[str]:
function_call = data.get("function_call")
if not function_call or function_call in ("auto", "none"):
return None
if isinstance(function_call, str):
return function_call
return self._get_mapping_value(function_call, "name")
def _collect_request_tools(self, data: dict) -> List[tuple[str, Optional[str]]]:
request_tools: List[tuple[str, Optional[str]]] = []
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))
for function in data.get("functions") or []:
function_name = self._get_legacy_function_name(function)
if function_name is not None:
request_tools.append((function_name, "function"))
for forced_tool_name in (
self._get_named_tool_choice(data),
self._get_named_function_call(data),
):
if forced_tool_name is not None:
request_tools.append((forced_tool_name, "function"))
return request_tools
def _modify_request_with_permission_errors(
self,
data: dict,
@ -410,19 +509,32 @@ class ToolPermissionGuardrail(CustomGuardrail):
for tool_use in denied_tool_names:
error_tool_names.add(tool_use)
# Modify the tools
tools: Optional[List[ChatCompletionToolParam]] = data.get("tools")
if tools is None:
return data
new_tools = []
for tool in tools:
if tool["type"] != "function":
continue
tool_name: str = tool["function"]["name"]
if tool_name not in error_tool_names:
if tools is not None:
new_tools = []
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"] = new_tools
functions = data.get("functions")
if functions is not None:
data["functions"] = [
function
for function in functions
if self._get_legacy_function_name(function) not in error_tool_names
]
named_tool_choice = self._get_named_tool_choice(data)
if named_tool_choice in error_tool_names:
data["tool_choice"] = "none"
named_function_call = self._get_named_function_call(data)
if named_function_call in error_tool_names:
data["function_call"] = "none"
return data
def _create_permission_error_result(
@ -472,7 +584,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
error_results[tool_use.id] = error_result
# Modify the response content
for choice in response.choices:
for choice_index, choice in enumerate(response.choices):
if isinstance(choice, Choices):
filtered_tool_calls = []
error_messages = []
@ -490,6 +602,15 @@ class ToolPermissionGuardrail(CustomGuardrail):
filtered_tool_calls if filtered_tool_calls else None
)
legacy_tool_call = self._legacy_function_call_to_tool_call(
getattr(choice.message, "function_call", None), choice_index
)
if legacy_tool_call is not None:
error_result = error_results.get(legacy_tool_call.id)
if error_result is not None:
choice.message.function_call = None
error_messages.append(error_result.content)
# Add error messages to content
if error_messages:
existing_content = choice.message.content
@ -519,21 +640,16 @@ class ToolPermissionGuardrail(CustomGuardrail):
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return data
new_tools: Optional[List[ChatCompletionToolParam]] = data.get("tools")
if new_tools is None:
new_tools = self._collect_request_tools(data)
if not new_tools:
verbose_proxy_logger.warning(
"Tool Permission Guardrail: not running guardrail. No tools in data"
"Tool Permission Guardrail: not running guardrail. No tools or functions in data"
)
return data
# Check permissions for each tool
denied_tool_names = []
for tool in new_tools:
if tool["type"] != "function":
continue
tool_name: str = tool["function"]["name"]
tool_type: Optional[str] = tool.get("type")
for tool_name, tool_type in new_tools:
is_allowed, _, message = self._check_tool_permission(tool_name, tool_type)
if not is_allowed and message is not None:

View file

@ -220,6 +220,27 @@ class TestToolPermissionGuardrail:
assert tool_calls[0].id == "call_123"
assert tool_calls[0].function.name == "Read"
def test_extract_tool_calls_legacy_function_call_format(self):
response = ModelResponse(
choices=[
Choices(
message={
"function_call": {
"name": "Read",
"arguments": '{"file_path": "/test/file.txt"}',
},
}
)
]
)
tool_calls = self.guardrail._extract_tool_calls_from_response(response)
assert len(tool_calls) == 1
assert isinstance(tool_calls[0], ChatCompletionMessageToolCall)
assert tool_calls[0].id == "legacy_function_call_0"
assert tool_calls[0].function.name == "Read"
assert tool_calls[0].function.arguments == '{"file_path": "/test/file.txt"}'
def test_extract_tool_calls_empty_response(self):
response = ModelResponse(choices=[])
tool_calls = self.guardrail._extract_tool_calls_from_response(response)
@ -271,6 +292,31 @@ class TestToolPermissionGuardrail:
data=data, user_api_key_dict=user_api_key_dict, response=response
)
@pytest.mark.asyncio
async def test_async_post_call_success_hook_with_denied_legacy_function_call_raises(
self,
):
response = ModelResponse(
choices=[
Choices(
message={
"function_call": {
"name": "Read",
"arguments": "{}",
},
}
)
]
)
user_api_key_dict = UserAPIKeyAuth()
data = {"guardrails": ["test-tool-permission"]}
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
with pytest.raises(GuardrailRaisedException):
await self.guardrail.async_post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
@pytest.mark.asyncio
async def test_async_post_call_success_hook_param_patterns_allow(self):
guardrail = ToolPermissionGuardrail(
@ -379,7 +425,9 @@ class TestToolPermissionGuardrail:
assert "berri" in choice.message.content
@pytest.mark.asyncio
async def test_async_post_call_success_hook_missing_arguments_default_allows(self):
async def test_async_post_call_success_hook_missing_arguments_blocks_param_rule(
self,
):
guardrail = ToolPermissionGuardrail(
guardrail_name="mail-guardrail",
rules=[
@ -405,9 +453,52 @@ class TestToolPermissionGuardrail:
data = {"guardrails": ["mail-guardrail"]}
with patch.object(guardrail, "should_run_guardrail", return_value=True):
await guardrail.async_post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
with pytest.raises(GuardrailRaisedException):
await guardrail.async_post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"arguments",
[
"{not-json",
'["owner@berri.ai"]',
],
)
async def test_async_post_call_success_hook_malformed_arguments_blocks_param_rule(
self, arguments
):
guardrail = ToolPermissionGuardrail(
guardrail_name="mail-guardrail",
rules=[
{
"id": "deny_gmail",
"tool_name": r"^mail_mcp-send_email$",
"decision": "deny",
"allowed_param_patterns": {"to[]": r"^.+@gmail\.com$"},
}
],
default_action="allow",
on_disallowed_action="block",
)
tool_call = {
"function": {
"name": "mail_mcp-send_email",
"arguments": arguments,
},
"type": "function",
}
response = ModelResponse(choices=[Choices(message={"tool_calls": [tool_call]})])
user_api_key_dict = UserAPIKeyAuth()
data = {"guardrails": ["mail-guardrail"]}
with patch.object(guardrail, "should_run_guardrail", return_value=True):
with pytest.raises(GuardrailRaisedException):
await guardrail.async_post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
@pytest.mark.asyncio
async def test_async_pre_call_hook_block_mode(self):
@ -430,6 +521,65 @@ class TestToolPermissionGuardrail:
)
assert excinfo.value.status_code == 400
@pytest.mark.asyncio
async def test_async_pre_call_hook_blocks_legacy_functions(self):
data = {
"functions": [
{"name": "Bash", "description": "allowed"},
{"name": "Read", "description": "denied"},
]
}
user_api_key_dict = UserAPIKeyAuth()
cache = DualCache(default_in_memory_ttl=1)
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
with pytest.raises(HTTPException) as excinfo:
await self.guardrail.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=cache,
data=data,
call_type="completion",
)
assert excinfo.value.status_code == 400
@pytest.mark.asyncio
async def test_async_pre_call_hook_blocks_named_legacy_function_call(self):
data = {
"functions": [{"name": "Bash"}],
"function_call": {"name": "Read"},
}
user_api_key_dict = UserAPIKeyAuth()
cache = DualCache(default_in_memory_ttl=1)
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
with pytest.raises(HTTPException) as excinfo:
await self.guardrail.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=cache,
data=data,
call_type="completion",
)
assert excinfo.value.status_code == 400
@pytest.mark.asyncio
async def test_async_pre_call_hook_blocks_named_tool_choice(self):
data = {
"tools": [{"type": "function", "function": {"name": "Bash"}}],
"tool_choice": {"type": "function", "function": {"name": "Read"}},
}
user_api_key_dict = UserAPIKeyAuth()
cache = DualCache(default_in_memory_ttl=1)
with patch.object(self.guardrail, "should_run_guardrail", return_value=True):
with pytest.raises(HTTPException) as excinfo:
await self.guardrail.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=cache,
data=data,
call_type="completion",
)
assert excinfo.value.status_code == 400
@pytest.mark.asyncio
async def test_async_pre_call_hook_uses_custom_template(self):
guardrail = ToolPermissionGuardrail(
@ -491,6 +641,41 @@ class TestToolPermissionGuardrail:
assert "Bash" in tool_names
assert "Read" not in tool_names
@pytest.mark.asyncio
async def test_async_pre_call_hook_rewrite_mode_filters_legacy_functions(self):
guardrail = ToolPermissionGuardrail(
guardrail_name="test-tool-permission",
rules=self.test_rules,
default_action="deny",
on_disallowed_action="rewrite",
)
data = {
"functions": [
{"name": "Bash", "description": "allowed"},
{"name": "Read", "description": "denied"},
],
"function_call": {"name": "Read"},
"tools": [
{"type": "function", "function": {"name": "Bash"}},
],
"tool_choice": {"type": "function", "function": {"name": "Read"}},
}
user_api_key_dict = UserAPIKeyAuth()
cache = DualCache(default_in_memory_ttl=1)
with patch.object(guardrail, "should_run_guardrail", return_value=True):
new_data = await guardrail.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=cache,
data=data,
call_type="completion",
)
assert isinstance(new_data, dict)
assert [function["name"] for function in new_data["functions"]] == ["Bash"]
assert new_data["function_call"] == "none"
assert new_data["tool_choice"] == "none"
def test_modify_response_with_permission_errors(self):
# Setup a response with one tool_call
tool_call = ChatCompletionMessageToolCall(
@ -522,6 +707,40 @@ class TestToolPermissionGuardrail:
assert isinstance(choice.message.content, str)
assert "Permission denied" in choice.message.content
def test_modify_response_with_permission_errors_filters_legacy_function_call(self):
response = ModelResponse(
choices=[
Choices(
message={
"function_call": {
"name": "Read",
"arguments": "{}",
},
"content": "",
}
)
]
)
tool_call = self.guardrail._extract_tool_calls_from_response(response)[0]
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.message.function_call is None
assert isinstance(choice.message.content, str)
assert "Permission denied" in choice.message.content
class TestToolPermissionGuardrailIntegration:
"""Integration tests for Tool Permission Guardrail"""