diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py index 64753d9fa85..3c7ec967396 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py @@ -1,6 +1,6 @@ import json import re -from typing import Any, AsyncGenerator, Dict, List, Literal, Optional, Union +from typing import Any, AsyncGenerator, Dict, List, Literal, Optional, Union, cast from fastapi import HTTPException @@ -17,6 +17,7 @@ from litellm.proxy.common_utils.callback_utils import ( add_guardrail_to_applied_guardrails_header, ) from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.llms.openai import ChatCompletionToolCallChunk from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( PermissionError, ToolPermissionRule, @@ -26,6 +27,7 @@ from litellm.types.utils import ( CallTypesLiteral, ChatCompletionMessageToolCall, Choices, + GenericGuardrailAPIInputs, LLMResponseTypes, ModelResponse, ModelResponseStream, @@ -34,6 +36,7 @@ from litellm.types.utils import ( GUARDRAIL_NAME = "tool_permission" + class ToolPermissionGuardrail(CustomGuardrail): def __init__( self, @@ -164,7 +167,7 @@ class ToolPermissionGuardrail(CustomGuardrail): self, tool_name: Optional[str], tool_type: Optional[str] = None, - ) -> tuple[bool, Optional[str], Optional[str]]: + ) -> tuple[bool, Optional[PermissionError]]: """ Check if a tool is allowed based on the configured rules @@ -173,7 +176,7 @@ class ToolPermissionGuardrail(CustomGuardrail): tool_type: Type of the tool to check Returns: - Tuple of (is_allowed, rule_id, message) + Tuple of (is_allowed, PermissionError) """ verbose_proxy_logger.debug( f"Checking permission for tool: {tool_name or tool_type}" @@ -198,7 +201,8 @@ class ToolPermissionGuardrail(CustomGuardrail): }, ) verbose_proxy_logger.debug(message) - return is_allowed, rule.id, message + error = None if is_allowed else PermissionError(tool_name=tool_identifier, rule_id=rule.id, message=message) + return is_allowed, error # No rule matched, use default action is_allowed = self.default_action == "allow" @@ -212,7 +216,8 @@ class ToolPermissionGuardrail(CustomGuardrail): }, ) verbose_proxy_logger.debug(message) - return is_allowed, None, message + error = None if is_allowed else PermissionError(tool_name=tool_identifier, rule_id=None, message=message) + return is_allowed, error def _parse_tool_call_arguments( self, tool_call: ChatCompletionMessageToolCall @@ -298,11 +303,21 @@ class ToolPermissionGuardrail(CustomGuardrail): def _get_permission_for_tool_call( self, tool_call: ChatCompletionMessageToolCall - ) -> tuple[bool, Optional[str], Optional[str]]: + ) -> tuple[bool, Optional[PermissionError]]: + """ + Check if a tool call is allowed based on the configured rules + + Args: + tool_call: ChatCompletionMessageToolCall + + Returns: + Tuple of (is_allowed, PermissionError) + """ + tool_name = tool_call.function.name if tool_call.function else None tool_type = getattr(tool_call, "type", None) if not tool_name and not tool_type: - return self.default_action == "allow", None, None + return self.default_action == "allow", None tool_identifier = tool_name or tool_type or "unknown_tool" @@ -338,7 +353,8 @@ class ToolPermissionGuardrail(CustomGuardrail): default=default_message, context={"tool_name": tool_identifier, "rule_id": rule.id}, ) - return is_allowed, rule.id, message + error = None if is_allowed else PermissionError(tool_name=tool_identifier,rule_id=rule.id, message=message) + return is_allowed, error is_allowed = self.default_action == "allow" default_message = ( @@ -350,7 +366,8 @@ class ToolPermissionGuardrail(CustomGuardrail): default=default_message, context={"tool_name": tool_identifier, "rule_id": None}, ) - return is_allowed, None, message + error = None if is_allowed else PermissionError(tool_name=tool_identifier,rule_id=None, message=message) + return is_allowed, error def _extract_tool_calls_from_response( self, response: ModelResponse @@ -404,8 +421,6 @@ class ToolPermissionGuardrail(CustomGuardrail): 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: new_tools.append(tool) @@ -487,6 +502,119 @@ class ToolPermissionGuardrail(CustomGuardrail): else: choice.message.content = "\n".join(error_messages) + def _normalize_tool_call_input( + self, tool_call_entry: Union[ChatCompletionToolCallChunk, ChatCompletionMessageToolCall] + ) -> Optional[ChatCompletionMessageToolCall]: + if isinstance(tool_call_entry, ChatCompletionMessageToolCall): + return tool_call_entry + if isinstance(tool_call_entry, dict): + function_payload = tool_call_entry.get("function") or {} + if isinstance(function_payload, dict): + function_payload = cast(Dict[str, Any], function_payload) + else: + function_payload = dict(function_payload) + try: + return ChatCompletionMessageToolCall( + id=tool_call_entry.get("id"), + type=tool_call_entry.get("type"), + function=function_payload, + ) + except Exception as exc: + verbose_proxy_logger.warning( + "Tool Permission Guardrail: Failed to normalize tool call %s: %s", + tool_call_entry, + exc, + ) + return None + return None + + def _sanitize_tool_definitions( + self, inputs: GenericGuardrailAPIInputs + ) -> GenericGuardrailAPIInputs: + tools = inputs.get("tools") + if not tools: + return inputs + + allowed_tools: List[ChatCompletionToolParam] = [] + + for tool in tools: + tool_name: str = tool["function"]["name"] + tool_type: Optional[str] = tool.get("type") + + is_allowed, error = self._check_tool_permission( + tool_name, tool_type + ) + if is_allowed or error is None: + allowed_tools.append(tool) + continue + + verbose_proxy_logger.warning( + "Tool Permission Guardrail: %s", error.message + ) + if self.on_disallowed_action == "block": + raise HTTPException( + status_code=400, + detail={ + "error": "Violated guardrail policy", + "detection_message": error.message, + }, + ) + + if len(tools) == len(allowed_tools): + verbose_proxy_logger.debug( + "Tool Permission Guardrail: All tool definitions allowed" + ) + inputs["tools"] = allowed_tools + return inputs + + def _sanitize_tool_calls( + self, inputs: GenericGuardrailAPIInputs, *, is_request: bool + ) -> GenericGuardrailAPIInputs: + if not inputs.get("tool_calls"): + return inputs + + filtered_tool_calls = [] + error_messages = [] + for tool_call_entry in inputs.get("tool_calls", []): + normalized_tool_call = self._normalize_tool_call_input(tool_call_entry) + if normalized_tool_call is None: + filtered_tool_calls.append(tool_call_entry) + continue + + is_allowed, error = self._get_permission_for_tool_call( + normalized_tool_call + ) + if is_allowed or error is None: + filtered_tool_calls.append(tool_call_entry) + continue + + verbose_proxy_logger.warning( + "Tool Permission Guardrail: %s", error.message + ) + + if self.on_disallowed_action == "block": + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=error.message, + ) + error_result = self._create_permission_error_result(normalized_tool_call, error) + error_messages.append(error_result.content) + + replaced_text = inputs.get("texts", []) + if error_messages: + if replaced_text: + replaced_text[-1] = replaced_text[-1] + "\n\n" + "\n".join(error_messages) + else: + replaced_text.append("\n".join(error_messages)) + else: + verbose_proxy_logger.debug( + "Tool Permission Guardrail: All tool calls allowed" + ) + + inputs["texts"] = replaced_text + inputs["tool_calls"] = filtered_tool_calls + return inputs + @log_guardrail_information async def async_pre_call_hook( self, @@ -516,21 +644,19 @@ class ToolPermissionGuardrail(CustomGuardrail): # 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") - is_allowed, _, message = self._check_tool_permission(tool_name, tool_type) + is_allowed, error = self._check_tool_permission(tool_name, tool_type) - if not is_allowed and message is not None: - verbose_proxy_logger.warning(f"Tool Permission Guardrail: {message}") + if not is_allowed and error is not None: + verbose_proxy_logger.warning(f"Tool Permission Guardrail: {error.message}") if self.on_disallowed_action == "block": raise HTTPException( status_code=400, detail={ "error": "Violated guardrail policy", - "detection_message": message, + "detection_message": error.message, }, ) denied_tool_names.append(tool_name) @@ -591,28 +717,20 @@ class ToolPermissionGuardrail(CustomGuardrail): # Check permissions for each tool use denied_tools = [] for tool_call in tool_calls: - is_allowed, rule_id, message = self._get_permission_for_tool_call(tool_call) + is_allowed, error = self._get_permission_for_tool_call(tool_call) - if not is_allowed and message is not None: - verbose_proxy_logger.warning(f"Tool Permission Guardrail: {message}") + if not is_allowed and error is not None: + verbose_proxy_logger.warning(f"Tool Permission Guardrail: {error.message}") if self.on_disallowed_action == "block": raise GuardrailRaisedException( guardrail_name=self.guardrail_name, - message=message, + message=error.message, ) denied_tools.append( ( tool_call, - PermissionError( - tool_name=( - tool_call.function.name - if tool_call.function and tool_call.function.name - else "unknown_tool" - ), - rule_id=rule_id, - message=message, - ), + error, ) ) @@ -679,32 +797,24 @@ class ToolPermissionGuardrail(CustomGuardrail): # Check permissions for each tool use denied_tools = [] for tool_call in tool_calls: - is_allowed, rule_id, message = self._get_permission_for_tool_call( + is_allowed, error = self._get_permission_for_tool_call( tool_call ) - if not is_allowed and message is not None: + if not is_allowed and error is not None: verbose_proxy_logger.warning( - f"Tool Permission Guardrail: {message}" + f"Tool Permission Guardrail: {error.message}" ) if self.on_disallowed_action == "block": raise GuardrailRaisedException( guardrail_name=self.guardrail_name, - message=message, + message=error.message, ) denied_tools.append( ( tool_call, - PermissionError( - tool_name=( - tool_call.function.name - if tool_call.function and tool_call.function.name - else "unknown_tool" - ), - rule_id=rule_id, - message=message, - ), + error, ) ) @@ -726,3 +836,23 @@ class ToolPermissionGuardrail(CustomGuardrail): else: for chunk in all_chunks: yield chunk + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + """ + Apply the Tool Permission guardrail to structured inputs used by guardrail translation handlers. + """ + verbose_proxy_logger.debug( + "Tool Permission Guardrail: apply_guardrail invoked for %s", + input_type, + ) + + if input_type == "request": + return self._sanitize_tool_definitions(inputs) + else: + return self._sanitize_tool_calls(inputs, is_request=False) 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 a7fd1c64955..8a57ce3265f 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 @@ -27,6 +27,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( from litellm.types.utils import ( ChatCompletionMessageToolCall, Choices, + GenericGuardrailAPIInputs, ModelResponse, ) @@ -90,13 +91,14 @@ class TestToolPermissionGuardrail: on_disallowed_action="block", ) - is_allowed, rule_id, _ = guardrail._check_tool_permission("AnyTool", "function") + is_allowed, error = guardrail._check_tool_permission("AnyTool", "function") assert is_allowed is True - assert rule_id == "allow_functions" + assert error is None - is_allowed, rule_id, _ = guardrail._check_tool_permission("AnyTool", "custom") + is_allowed, error = guardrail._check_tool_permission("AnyTool", "custom") assert is_allowed is False - assert rule_id is None + assert error + assert error.rule_id is None def test_rule_matches_tool_with_name_and_type(self): guardrail = ToolPermissionGuardrail( @@ -113,13 +115,14 @@ class TestToolPermissionGuardrail: on_disallowed_action="block", ) - is_allowed, rule_id, _ = guardrail._check_tool_permission("Bash", "function") + is_allowed, error = guardrail._check_tool_permission("Bash", "function") assert is_allowed is True - assert rule_id == "allow_specific" + assert error is None - is_allowed, rule_id, _ = guardrail._check_tool_permission("Bash", "custom") + is_allowed, error = guardrail._check_tool_permission("Bash", "custom") assert is_allowed is False - assert rule_id is None + assert error + assert error.rule_id is None def test_rule_requires_name_or_type(self): with pytest.raises(ValueError): @@ -150,32 +153,33 @@ class TestToolPermissionGuardrail: type="function", ) - is_allowed, rule_id, _ = guardrail._get_permission_for_tool_call(tool_call) + is_allowed, error = guardrail._get_permission_for_tool_call(tool_call) assert is_allowed is True - assert rule_id == "allow_type_only" + assert error is None def test_check_tool_permission_allow(self): - is_allowed, rule_id, msg = self.guardrail._check_tool_permission("Bash") + is_allowed, error = self.guardrail._check_tool_permission("Bash") assert is_allowed is True - assert rule_id == "allow_bash" - assert "allowed" in (msg or "") + assert error is None - is_allowed, rule_id, _ = self.guardrail._check_tool_permission( + is_allowed, error = self.guardrail._check_tool_permission( "mcp__github_add_issue_comment" ) assert is_allowed is True - assert rule_id == "allow_github" + assert error is None def test_check_tool_permission_deny(self): - is_allowed, rule_id, msg = self.guardrail._check_tool_permission("Read") + is_allowed, error = self.guardrail._check_tool_permission("Read") assert is_allowed is False - assert rule_id == "deny_read" - assert "denied" in (msg or "") + assert error + assert error.rule_id == "deny_read" + assert "denied" in error.message - is_allowed, rule_id, msg = self.guardrail._check_tool_permission("UnknownTool") + is_allowed, error = self.guardrail._check_tool_permission("UnknownTool") assert is_allowed is False - assert rule_id is None - assert "default" in (msg or "") + assert error + assert error.rule_id is None + assert "default" in error.message def test_check_tool_permission_custom_template(self): guardrail = ToolPermissionGuardrail( @@ -185,15 +189,19 @@ class TestToolPermissionGuardrail: violation_message_template="custom {tool_name} {rule_id} :: {default_message}", ) - _, rule_id, message = guardrail._check_tool_permission("Read") - assert rule_id == "deny_read" - assert message.startswith("custom Read deny_read") - assert "Tool 'Read' denied" in message + is_allowed, error = guardrail._check_tool_permission("Read") + assert is_allowed is False + assert error + assert error.rule_id == "deny_read" + assert error.message.startswith("custom Read") + assert "Tool 'Read' denied by rule 'deny_read'" in error.message - _, rule_id, message = guardrail._check_tool_permission("UnknownTool") - assert rule_id is None - assert message.startswith("custom UnknownTool None") - assert "Tool 'UnknownTool' denied by default action" in message + is_allowed, error = guardrail._check_tool_permission("UnknownTool") + assert is_allowed is False + assert error + assert error.rule_id is None + assert error.message.startswith("custom UnknownTool") + assert "Tool 'UnknownTool' denied by default action" in error.message def test_extract_tool_calls_openai_format(self): tool_call = { @@ -525,6 +533,89 @@ class TestToolPermissionGuardrail: assert isinstance(choice.message.content, str) assert "Permission denied" in choice.message.content + @pytest.mark.asyncio + async def test_apply_guardrail_request_rewrite_filters_tools(self): + guardrail = ToolPermissionGuardrail( + guardrail_name="rewrite-permissions", + rules=self.test_rules, + default_action="deny", + on_disallowed_action="rewrite", + ) + inputs = GenericGuardrailAPIInputs( + tools=[ + {"type": "function", "function": {"name": "Bash"}}, + {"type": "function", "function": {"name": "Read"}}, + ] + ) + + updated_inputs = await guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="request", + ) + + assert "tools" in updated_inputs + tool_names = [tool["function"]["name"] for tool in updated_inputs["tools"]] + assert tool_names == ["Bash"] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_block_raises(self): + inputs = GenericGuardrailAPIInputs( + tools=[ + {"type": "function", "function": {"name": "Read"}}, + ] + ) + with pytest.raises(HTTPException): + await self.guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="request", + ) + + @pytest.mark.asyncio + async def test_apply_guardrail_response_rewrite_mutates_tool_calls(self): + guardrail = ToolPermissionGuardrail( + guardrail_name="rewrite-response", + rules=self.test_rules, + default_action="deny", + on_disallowed_action="rewrite", + ) + tool_call = { + "id": "call_123", + "type": "function", + "function": {"name": "Read", "arguments": "{}"}, + } + inputs = GenericGuardrailAPIInputs(tool_calls=[tool_call]) + + updated_inputs = await guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="response", + ) + + assert updated_inputs.get("tool_calls") == [] + assert updated_inputs.get("texts") + assert "Permission denied" in updated_inputs["texts"][0] + + @pytest.mark.asyncio + async def test_apply_guardrail_response_block_raises(self): + inputs = GenericGuardrailAPIInputs( + tool_calls=[ + { + "id": "call_234", + "type": "function", + "function": {"name": "Read", "arguments": "{}"}, + } + ] + ) + + with pytest.raises(GuardrailRaisedException): + await self.guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="response", + ) + class TestToolPermissionGuardrailIntegration: """Integration tests for Tool Permission Guardrail""" @@ -538,14 +629,14 @@ class TestToolPermissionGuardrailIntegration: default_action="allow", ) - is_allowed, rule_id, message = guardrail._check_tool_permission("UnknownTool") + is_allowed, error = guardrail._check_tool_permission("UnknownTool") assert is_allowed is True - assert rule_id is None - assert "default" in (message or "") + assert error is None - is_allowed, rule_id, _ = guardrail._check_tool_permission("Read") + is_allowed, error = guardrail._check_tool_permission("Read") assert is_allowed is False - assert rule_id == "deny_read" + assert error + assert error.rule_id == "deny_read" def test_empty_rules(self): guardrail = ToolPermissionGuardrail( @@ -554,7 +645,6 @@ class TestToolPermissionGuardrailIntegration: default_action="allow", ) - is_allowed, rule_id, message = guardrail._check_tool_permission("AnyTool") + is_allowed, error = guardrail._check_tool_permission("AnyTool") assert is_allowed is True - assert rule_id is None - assert "default" in (message or "") + assert error is None \ No newline at end of file