mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
chore(guardrails): tighten tool permission checks
This commit is contained in:
parent
eab0075353
commit
150a34f2b0
2 changed files with 371 additions and 36 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue