From 1de7f076ac14fa9047d3e0b714115c99d6809429 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Fri, 29 Aug 2025 07:29:49 +0900 Subject: [PATCH] feat: add tool-permission guardrail --- .../docs/proxy/guardrails/tool_permission.md | 153 ++++++ .../tool_permission_example.yaml | 36 ++ .../guardrail_hooks/tool_permission.py | 511 ++++++++++++++++++ .../guardrails/guardrail_initializers.py | 15 + .../proxy/guardrails/guardrail_registry.py | 2 + litellm/types/guardrails.py | 29 +- .../guardrail_hooks/tool_permission.py | 41 ++ .../guardrail_hooks/test_tool_permission.py | 319 +++++++++++ 8 files changed, 1097 insertions(+), 9 deletions(-) create mode 100644 docs/my-website/docs/proxy/guardrails/tool_permission.md create mode 100644 litellm/proxy/example_config_yaml/tool_permission_example.yaml create mode 100644 litellm/proxy/guardrails/guardrail_hooks/tool_permission.py create mode 100644 litellm/types/proxy/guardrails/guardrail_hooks/tool_permission.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py diff --git a/docs/my-website/docs/proxy/guardrails/tool_permission.md b/docs/my-website/docs/proxy/guardrails/tool_permission.md new file mode 100644 index 00000000000..9ed05ed46a8 --- /dev/null +++ b/docs/my-website/docs/proxy/guardrails/tool_permission.md @@ -0,0 +1,153 @@ +import Image from '@theme/IdealImage'; +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Tool Permission Guardrail + +LiteLLM provides a Tool Permission Guardrail that lets you control which **tool calls** a model is allowed to invoke, using configurable allow/deny rules. This offers fine-grained, provider-agnostic control over tool execution (e.g., OpenAI Chat Completions `tool_calls`, Anthropic Messages `tool_use`, MCP tools). + +## Quick Start +### 1. Define Guardrails on your LiteLLM config.yaml + +Define your guardrails under the `guardrails` section +```yaml +guardrails: + - guardrail_name: "tool-permission-guardrail" + litellm_params: + guardrail: tool_permission + mode: "post_call" + rules: + - id: "allow_bash" + tool_name: "Bash" + decision: "allow" + - id: "allow_github_mcp" + tool_name: "mcp__github_*" + decision: "allow" + - id: "allow_aws_documentation" + tool_name: "mcp__aws-documentation_*_documentation" + decision: "allow" + - id: "deny_read_commands" + tool_name: "Read" + decision: "Deny" + default_action: "deny" # Fallback when no rule matches: "allow" or "deny" + on_disallowed_action: "block" # How to handle disallowed tools: "block" or "rewrite" +``` + +#### Rule Structure + +```yaml +- id: "unique_rule_id" # Unique identifier for the rule + tool_name: "pattern" # Tool name or pattern to match + decision: "allow" # "allow" or "deny" +``` + +#### Supported values for `mode` + +- `pre_call` Run **before** LLM call, on **input** +- `post_call` Run **after** LLM call, on **input & output** + +### 2. Start the Proxy + +```shell +litellm --config config.yaml --port 4000 +``` + +## Examples + + + + +**Block requset** + +```bash +# Test +curl -X POST "http://localhost:4000/v1/chat/completions" \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer your-master-key-here" \ + -d '{ + "model": "gpt-5-mini", + "messages": [{"role": "user","content": "What is the weather like in Tokyo today?"}], + "tools": [ + { + "type":"function", + "function": { + "name":"get_current_weather", + "description": "Get the current weather in a given location" + } + } + ] + }' +``` + +**Expected response (Denied):** + +```json +{ + "error": + { + "message": "Guardrail raised an exception, Guardrail: tool-permission-guardrail, Message: Tool 'get_current_weather' denied by default action", + "type": "None", + "param": "None", + "code": "500" + } +} +``` + + + + +**Rewrite requset** + +```bash +# Test +curl -X POST "http://localhost:4000/v1/chat/completions" \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer your-master-key-here" \ + -d '{ + "model": "gpt-5-mini", + "messages": [{"role": "user","content": "What is the weather like in Tokyo today?"}], + "tools": [ + { + "type":"function", + "function": { + "name":"get_current_weather", + "description": "Get the current weather in a given location" + } + } + ] + }' +``` + +**Expected response:** + +```json +{ + "id": "chatcmpl-xxxxxxxxxxxxxxx", + "created": 1757716050, + "model": "gpt-5-mini-2025-08-07", + "object": "chat.completion", + "choices": [ + { + "finish_reason": "stop", + "index": 0, + "message": { + "content": "I can’t fetch live weather — I don’t have real‑time internet access.", + "role": "assistant", + "annotations": [] + }, + "provider_specific_fields": {} + } + ], + "usage": { + "prompt_tokens": 112, + "total_tokens": 735, + "completion_tokens_details": { + "reasoning_tokens": 384, + }, + }, + "service_tier": "default" +} +``` + + + diff --git a/litellm/proxy/example_config_yaml/tool_permission_example.yaml b/litellm/proxy/example_config_yaml/tool_permission_example.yaml new file mode 100644 index 00000000000..e18425ba383 --- /dev/null +++ b/litellm/proxy/example_config_yaml/tool_permission_example.yaml @@ -0,0 +1,36 @@ +model_list: + - model_name: claude-3-5-sonnet + litellm_params: + model: anthropic/claude-3-5-sonnet-20241022 + api_key: os.environ/ANTHROPIC_API_KEY + +guardrails: + - guardrail_name: "tool-permission-guardrail" + litellm_params: + guardrail: tool_permission + mode: "post_call" + default_on: true # Apply to all requests by default + rules: + - id: "allow_bash" + tool_name: "Bash" + decision: "allow" + - id: "allow_github_mcp" + tool_name: "mcp__github_*" + decision: "allow" + - id: "allow_aws_documentation" + tool_name: "mcp__aws-documentation_*_documentation" + decision: "allow" + - id: "deny_read_commands" + tool_name: "Read" + decision: "Deny" + default_action: "deny" # deny by default if no rule matches + on_disallowed_action: "block" # block by default if no rule matches + +# Optional: Configure general settings +general_settings: + master_key: sk-1234 + +# Optional: Add logging configuration +litellm_settings: + success_callback: ["langfuse"] + failure_callback: ["langfuse"] \ No newline at end of file diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py new file mode 100644 index 00000000000..7519f0fa479 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py @@ -0,0 +1,511 @@ +from fastapi import HTTPException + +import re +from typing import Any, AsyncGenerator, Dict, List, Literal, Optional, Union + +from litellm import ChatCompletionToolParam + +from litellm._logging import verbose_proxy_logger +from litellm.caching.dual_cache import DualCache +from litellm.exceptions import GuardrailRaisedException +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, +) +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( + PermissionError, + ToolPermissionRule, + ToolResult, +) +from litellm.types.utils import ( + ModelResponse, + ModelResponseStream, + LLMResponseTypes, + Choices, + ChatCompletionMessageToolCall, +) + +GUARDRAIL_NAME = "tool_permission" + + +class ToolPermissionGuardrail(CustomGuardrail): + def __init__( + self, + rules: Optional[List[Dict]] = None, + default_action: Literal["deny", "allow"] = "deny", + on_disallowed_action: Literal["block", "rewrite"] = "block", + **kwargs, + ): + """ + Initialize the Tool Permission Guardrail + + Args: + rules: List of permission rules + default_action: Default action when no rule matches ("allow" or "deny") + on_disallowed_action: + **kwargs: Additional arguments passed to CustomGuardrail + """ + # Set supported event hooks - this guardrail only works on post_call + if "supported_event_hooks" not in kwargs: + kwargs["supported_event_hooks"] = [ + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + ] + + super().__init__(**kwargs) + + self.rules: List[ToolPermissionRule] = [] + if rules: + for rule_dict in rules: + self.rules.append(ToolPermissionRule(**rule_dict)) + + self.default_action = default_action + self.on_disallowed_action = on_disallowed_action + + verbose_proxy_logger.debug( + "Tool Permission Guardrail initialized with %d rules, default_action: %s", + len(self.rules), + self.default_action, + ) + + def _matches_pattern(self, tool_name: str, pattern: str) -> bool: + """ + Check if a tool name matches a pattern + + Supports patterns like: + - "Bash" - exact match + - "mcp__*" - prefix pattern (matches names starting wich "mcp__") + - "*_read" - suffix wildcard (matches names ending with "_read") + - "mcp__github_*_read" - infix wildcard (matches names like "mcp__github_mark_all_notifications_read") + + Args: + tool_name: Name of the tool to check + pattern: Pattern to match against + + Returns: + True if the tool name matches the pattern + """ + # Handle exact matches + if tool_name == pattern: + return True + + if "*" in pattern: + # Escape regex special chars except '*' + escaped_pattern = re.escape(pattern) + # Turn \* into .* + regex_pattern = escaped_pattern.replace(r"\*", ".*") + return bool(re.fullmatch(regex_pattern, tool_name)) + + return False + + def _check_tool_permission( + self, tool_name: str + ) -> tuple[bool, Optional[str], Optional[str]]: + """ + Check if a tool is allowed based on the configured rules + + Args: + tool_name: Name of the tool to check + + Returns: + Tuple of (is_allowed, rule_id, message) + """ + verbose_proxy_logger.debug(f"Checking permission for tool: {tool_name}") + + # Check each rule in order + for rule in self.rules: + if self._matches_pattern(tool_name, rule.tool_name): + is_allowed = rule.decision == "allow" + message = f"Tool '{tool_name}' {'allowed' if is_allowed else 'denied'} by rule '{rule.id}'" + verbose_proxy_logger.debug(message) + return is_allowed, rule.id, message + + # No rule matched, use default action + is_allowed = self.default_action == "allow" + message = f"Tool '{tool_name}' {'allowed' if is_allowed else 'denied'} by default action" + verbose_proxy_logger.debug(message) + return is_allowed, None, message + + def _extract_tool_calls_from_response( + self, response: ModelResponse + ) -> List[ChatCompletionMessageToolCall]: + """ + Extract tool_calls from all choices in a model response. + + Args: + response: The model response to analyze + + Returns: + List of tool_calls blocks found in the response + """ + tool_calls = [] + + for choice in response.choices: + if isinstance(choice, Choices): + for tool in choice.message.tool_calls or []: + tool_calls.append(tool) + + return tool_calls + + def _modify_request_with_permission_errors( + self, + data: dict, + denied_tool_names: List[str], + ): + """ + Modify the request to replace denied tool_calls blocks with error results + + Args: + data: The model request to modify + denied_tools: List of (tool_use, error) tuples for denied tools + """ + if not denied_tool_names: + return data + + verbose_proxy_logger.info( + f"Blocking {len(denied_tool_names)} unauthorized tool uses" + ) + + # Create a mapping of tool_use_id to error result + error_tool_names = set() + 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: + new_tools.append(tool) + data["tools"] = new_tools + return data + + def _create_permission_error_result( + self, tool_call: ChatCompletionMessageToolCall, error: PermissionError + ) -> ToolResult: + """ + Create a tool_result block for a permission error + + Args: + tool_use: The tool use that was denied + error: The permission error details + + Returns: + A tool_result block with the error message + """ + error_message = f"Permission denied: {error.message}" + if error.rule_id: + error_message += f" (Rule: {error.rule_id})" + + return ToolResult( + tool_use_id=tool_call.id, content=error_message, is_error=True + ) + + def _modify_response_with_permission_errors( + self, + response: ModelResponse, + denied_tools: List[tuple[ChatCompletionMessageToolCall, PermissionError]], + ) -> None: + """ + Modify the response to replace denied tool_calls blocks with error results + + Args: + response: The model response to modify + denied_tools: List of (tool_use, error) tuples for denied tools + """ + if not denied_tools: + return + + verbose_proxy_logger.info( + f"Blocking {len(denied_tools)} unauthorized tool uses" + ) + + # Create a mapping of tool_use_id to error result + error_results = {} + for tool_use, error in denied_tools: + error_result = self._create_permission_error_result(tool_use, error) + error_results[tool_use.id] = error_result + + # Modify the response content + for choice in response.choices: + if isinstance(choice, Choices): + filtered_tool_calls = [] + error_messages = [] + + # Rewrite tool_calls + for tool_call in choice.message.tool_calls or []: + tool_call_id = tool_call.id + if tool_call_id in error_results: + error_result = error_results[tool_call_id] + error_messages.append(error_result.content) + else: + filtered_tool_calls.append(tool_call) + + choice.message.tool_calls = ( + filtered_tool_calls if filtered_tool_calls else None + ) + + # Add error messages to content + if error_messages: + existing_content = choice.message.content + if existing_content: + choice.message.content = ( + existing_content + "\n\n" + "\n".join(error_messages) + ) + else: + choice.message.content = "\n".join(error_messages) + + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: Literal[ + "completion", + "text_completion", + "embeddings", + "image_generation", + "moderation", + "audio_transcription", + "pass_through_endpoint", + "rerank", + "mcp_call", + ], + ) -> Union[Exception, str, dict, None]: + """ """ + verbose_proxy_logger.debug("Tool Permission Guardrail Pre-Call Hook") + + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) + + event_type: GuardrailEventHooks = GuardrailEventHooks.pre_call + 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: + verbose_proxy_logger.warning( + "Tool Permission Guardrail: not running guardrail. No tools 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"] + + is_allowed, _, message = self._check_tool_permission(tool_name) + + if not is_allowed and message is not None: + verbose_proxy_logger.warning(f"Tool Permission Guardrail: {message}") + if self.on_disallowed_action == "block": + raise HTTPException( + status_code=400, + detail={ + "error": "Violated guardrail policy", + "detection_message": message, + }, + ) + denied_tool_names.append(tool_name) + + if denied_tool_names: + data = self._modify_request_with_permission_errors(data, denied_tool_names) + + verbose_proxy_logger.debug( + "Tool Permission Guardrail Pre-Call Hook: All tools allowed" + ) + + add_guardrail_to_applied_guardrails_header( + request_data=data, guardrail_name=self.guardrail_name + ) + return data + + @log_guardrail_information + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response: LLMResponseTypes, + ): + """ + Check tool usage permissions after the LLM call + + Args: + data: Request data + user_api_key_dict: User API key information (unused but required by interface) + response: The model response to check + """ + if not isinstance(response, ModelResponse): + return + + verbose_proxy_logger.debug( + "Tool Permission Guardrail Post-Call Hook: Checking response" + ) + + if not self.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.post_call + ): + verbose_proxy_logger.debug( + "Tool Permission Guardrail: Skipping check (not enabled)" + ) + return + + # Extract tool_calls from the response + tool_calls = self._extract_tool_calls_from_response(response) + + if not tool_calls: + verbose_proxy_logger.debug("Tool Permission Guardrail: No tool uses found") + return + + verbose_proxy_logger.debug( + f"Tool Permission Guardrail: Found {len(tool_calls)} tool calls" + ) + + # Check permissions for each tool use + denied_tools = [] + for tool_call in tool_calls: + if tool_call.function.name is None: + continue + is_allowed, rule_id, message = self._check_tool_permission( + tool_call.function.name + ) + + if not is_allowed and message is not None: + verbose_proxy_logger.warning(f"Tool Permission Guardrail: {message}") + + if self.on_disallowed_action == "block": + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=message, + ) + denied_tools.append( + ( + tool_call, + PermissionError( + tool_name=tool_call.function.name, + rule_id=rule_id, + message=message, + ), + ) + ) + + if denied_tools: + self._modify_response_with_permission_errors(response, denied_tools) + + verbose_proxy_logger.debug( + "Tool Permission Guardrail Post-Call Hook: All tools allowed" + ) + + add_guardrail_to_applied_guardrails_header( + request_data=data, guardrail_name=self.guardrail_name + ) + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + request_data: dict, + ) -> AsyncGenerator[ModelResponseStream, None]: + """ + Check tool usage permissions after the LLM stream call + + Args: + user_api_key_dict: User API key information (unused but required by interface) + response: The model response to check + request_data: The model request (unused but required by interface) + """ + + # Import here to avoid circular imports + from litellm.llms.base_llm.base_model_iterator import MockResponseIterator + from litellm.main import stream_chunk_builder + from litellm.types.utils import TextCompletionResponse + + # Collect all chunks to process them together + all_chunks: List[ModelResponseStream] = [] + async for chunk in response: + all_chunks.append(chunk) + + assembled_model_response: Optional[ + Union[ModelResponse, TextCompletionResponse] + ] = stream_chunk_builder( + chunks=all_chunks, + ) + if isinstance(assembled_model_response, ModelResponse): + verbose_proxy_logger.debug("Tool Permission Guardrail: Checking response") + + # Extract tool_calls from the response + tool_calls = self._extract_tool_calls_from_response(assembled_model_response) + + if not tool_calls: + verbose_proxy_logger.debug( + "Tool Permission Guardrail: No tool uses found" + ) + return + + verbose_proxy_logger.debug( + f"Tool Permission Guardrail: Found {len(tool_calls)} tool calls" + ) + + # Check permissions for each tool use + denied_tools = [] + for tool_call in tool_calls: + if tool_call.function.name is None: + continue + is_allowed, rule_id, message = self._check_tool_permission( + tool_call.function.name + ) + + if not is_allowed and message is not None: + verbose_proxy_logger.warning( + f"Tool Permission Guardrail: {message}" + ) + + if self.on_disallowed_action == "block": + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=message, + ) + denied_tools.append( + ( + tool_call, + PermissionError( + tool_name=tool_call.function.name, + rule_id=rule_id, + message=message, + ), + ) + ) + + verbose_proxy_logger.debug( + "Tool Permission Guardrail Post-Call Hook: All tools allowed" + ) + + if denied_tools: + self._modify_response_with_permission_errors( + assembled_model_response, denied_tools + ) + + mock_response = MockResponseIterator( + model_response=assembled_model_response + ) + # Return the reconstructed stream + async for chunk in mock_response: + yield chunk + else: + for chunk in all_chunks: + yield chunk diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 1cf0b15ab5c..23731528d7b 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -123,3 +123,18 @@ def initialize_hide_secrets(litellm_params: LitellmParams, guardrail: Guardrail) return _secret_detection_object +def initialize_tool_permission(litellm_params: LitellmParams, guardrail: Guardrail): + from litellm.proxy.guardrails.guardrail_hooks.tool_permission import ( + ToolPermissionGuardrail, + ) + + _tool_permission_callback = ToolPermissionGuardrail( + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + rules=litellm_params.rules, + default_action=getattr(litellm_params, "default_action", "deny"), + on_disallowed_action=getattr(litellm_params, "on_disallowed_action", "block"), + default_on=litellm_params.default_on, + ) + litellm.logging_callback_manager.add_litellm_callback(_tool_permission_callback) + return _tool_permission_callback diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 21429f462d4..69c8fd8084e 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -26,6 +26,7 @@ from .guardrail_initializers import ( initialize_lakera, initialize_lakera_v2, initialize_presidio, + initialize_tool_permission, ) guardrail_initializer_registry = { @@ -34,6 +35,7 @@ guardrail_initializer_registry = { SupportedGuardrailIntegrations.LAKERA_V2.value: initialize_lakera_v2, SupportedGuardrailIntegrations.PRESIDIO.value: initialize_presidio, SupportedGuardrailIntegrations.HIDE_SECRETS.value: initialize_hide_secrets, + SupportedGuardrailIntegrations.TOOL_PERMISSION.value: initialize_tool_permission, } guardrail_class_registry: Dict[str, Type[CustomGuardrail]] = {} diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index f31f304bda9..03cfa42a33d 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -2,12 +2,8 @@ from datetime import datetime from enum import Enum from typing import Any, Dict, List, Literal, Optional, TypedDict, Union -from pydantic import BaseModel, ConfigDict, Field, SecretStr -from typing_extensions import Required, TypedDict - -from litellm.types.proxy.guardrails.guardrail_hooks.openai.openai_moderation import ( - OpenAIModerationGuardrailConfigModel, -) +from pydantic import BaseModel, ConfigDict, Field +from typing_extensions import Required """ Pydantic object defining how to set guardrails on litellm proxy @@ -41,6 +37,8 @@ class SupportedGuardrailIntegrations(Enum): MODEL_ARMOR = "model_armor" OPENAI_MODERATION = "openai_moderation" NOMA = "noma" + TOOL_PERMISSION = "tool_permission" + class Role(Enum): SYSTEM = "system" @@ -312,7 +310,6 @@ class BedrockGuardrailConfigModel(BaseModel): ) - class LakeraV2GuardrailConfigModel(BaseModel): """Configuration parameters for the Lakera AI v2 guardrail""" @@ -377,6 +374,18 @@ class NomaGuardrailConfigModel(BaseModel): ) +class ToolPermissionGuardrailConfigModel(BaseModel): + """Configuration parameters for the Tool Permission guardrail""" + + rules: Optional[List[Dict]] = Field( + default=None, description="List of permission rules for tool usage" + ) + default_action: Optional[str] = Field( + default="Deny", + description="Default action when no rule matches (Allow or Deny)", + ) + + class BaseLitellmParams(BaseModel): # works for new and patch update guardrails api_key: Optional[str] = Field( default=None, description="API key for the guardrail service" @@ -425,7 +434,8 @@ class BaseLitellmParams(BaseModel): # works for new and patch update guardrails ) model: Optional[str] = Field( - default=None, description="Optional field if guardrail requires a 'model' parameter" + default=None, + description="Optional field if guardrail requires a 'model' parameter", ) # Model Armor params @@ -446,7 +456,7 @@ class BaseLitellmParams(BaseModel): # works for new and patch update guardrails default=True, description="Whether to fail the request if Model Armor encounters an error", ) - + model_config = ConfigDict(extra="allow", protected_namespaces=()) @@ -464,6 +474,7 @@ class LitellmParams( LassoGuardrailConfigModel, PillarGuardrailConfigModel, NomaGuardrailConfigModel, + ToolPermissionGuardrailConfigModel, BaseLitellmParams, ): guardrail: str = Field(description="The type of guardrail integration to use") diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/tool_permission.py b/litellm/types/proxy/guardrails/guardrail_hooks/tool_permission.py new file mode 100644 index 00000000000..dd4d63d75f5 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/tool_permission.py @@ -0,0 +1,41 @@ +# Tool Permission Guardrail Type Definitions +from typing import Literal, Optional + +from pydantic import BaseModel, Field + + +class ToolPermissionRule(BaseModel): + """ + A rule defining permission for a specific tool or tool pattern + """ + + id: str = Field(description="Unique identifier for the rule") + tool_name: str = Field( + description="Tool name or pattern (e.g., 'Bash', 'mcp__github_*', 'mcp__github_*_read', '*_read')" + ) + decision: Literal["allow", "deny"] = Field( + description="Whether to allow or deny this tool usage" + ) + + +class ToolResult(BaseModel): + """ + Represents a tool_result block to be added to the response + """ + + type: str = Field(default="tool_result", description="Should be 'tool_result'") + tool_use_id: str = Field( + description="ID of the tool use this result corresponds to" + ) + content: str = Field(description="Result content") + is_error: bool = Field(default=True, description="Whether this is an error result") + + +class PermissionError(BaseModel): + """ + Error information for permission denial + """ + + tool_name: str = Field(description="Name of the denied tool") + rule_id: Optional[str] = Field(description="ID of the rule that caused denial") + message: str = Field(description="Error message") 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 new file mode 100644 index 00000000000..a9d87398217 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_permission.py @@ -0,0 +1,319 @@ +""" +Unit tests for Tool Permission Guardrail (OpenAI tool_calls semantics) +""" + +import os +import sys +from unittest.mock import patch + +from litellm.caching.dual_cache import DualCache +import pytest + +sys.path.insert(0, os.path.abspath("../../../../../..")) + +from fastapi import HTTPException +from litellm.exceptions import GuardrailRaisedException +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.tool_permission import ( + ToolPermissionGuardrail, +) +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( + PermissionError, +) +from litellm.types.utils import ( + ModelResponse, + Choices, + ChatCompletionMessageToolCall, +) + + +class TestToolPermissionGuardrail: + """Test class for Tool Permission Guardrail functionality""" + + def setup_method(self): + """Set up test fixtures""" + self.test_rules = [ + {"id": "allow_bash", "tool_name": "Bash", "decision": "allow"}, + {"id": "allow_github", "tool_name": "mcp__github_*", "decision": "allow"}, + { + "id": "allow_documentation", + "tool_name": "mcp__aws-documentation_*_documentation", + "decision": "allow", + }, + {"id": "deny_read", "tool_name": "Read", "decision": "deny"}, + {"id": "deny_get", "tool_name": "*_get", "decision": "deny"}, + ] + + self.guardrail = ToolPermissionGuardrail( + guardrail_name="test-tool-permission", + rules=self.test_rules, + default_action="deny", + on_disallowed_action="block", + ) + + def test_initialization(self): + """Test guardrail initialization""" + assert self.guardrail.guardrail_name == "test-tool-permission" + assert len(self.guardrail.rules) == 5 + assert self.guardrail.default_action == "deny" + assert self.guardrail.on_disallowed_action == "block" + assert GuardrailEventHooks.post_call in ( + self.guardrail.supported_event_hooks or [] + ) + + def test_pattern_matching_exact(self): + """Test exact pattern matching""" + assert self.guardrail._matches_pattern("Read", "Read") is True + assert self.guardrail._matches_pattern("Write", "Read") is False + + def test_pattern_matching_wildcards(self): + """Test wildcard pattern matching""" + assert ( + self.guardrail._matches_pattern( + "mcp__github_add_issue_comment", "mcp__github_*" + ) + is True + ) + assert ( + self.guardrail._matches_pattern( + "mcp__github_add_issue_comment", "mcp__github_*_comment" + ) + is True + ) + assert ( + self.guardrail._matches_pattern( + "mcp__github_add_issue_comment", "*_comment" + ) + is True + ) + assert ( + self.guardrail._matches_pattern( + "mcp__git_add_issue_comment", "mcp__github_*" + ) + is False + ) + assert ( + self.guardrail._matches_pattern( + "mcp__github_assign_copilot_to_issue", "mcp__github_*_comment" + ) + is False + ) + assert ( + self.guardrail._matches_pattern( + "mcp__github_assign_copilot_to_issue", "*_comment" + ) + is False + ) + + def test_check_tool_permission_allow(self): + is_allowed, rule_id, msg = self.guardrail._check_tool_permission("Bash") + assert is_allowed is True + assert rule_id == "allow_bash" + assert "allowed" in (msg or "") + + is_allowed, rule_id, _ = self.guardrail._check_tool_permission( + "mcp__github_add_issue_comment" + ) + assert is_allowed is True + assert rule_id == "allow_github" + + def test_check_tool_permission_deny(self): + is_allowed, rule_id, msg = self.guardrail._check_tool_permission("Read") + assert is_allowed is False + assert rule_id == "deny_read" + assert "denied" in (msg or "") + + is_allowed, rule_id, msg = self.guardrail._check_tool_permission("UnknownTool") + assert is_allowed is False + assert rule_id is None + assert "default" in (msg or "") + + def test_extract_tool_calls_openai_format(self): + tool_call = { + "id": "call_123", + "function": { + "name": "Read", + "arguments": '{"file_path": "/test/file.txt"}', + }, + "type": "function", + } + response = ModelResponse( + choices=[ + Choices( + message={ + "tool_calls": [tool_call], + } + ) + ] + ) + + 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 == "call_123" + assert tool_calls[0].function.name == "Read" + + def test_extract_tool_calls_empty_response(self): + response = ModelResponse(choices=[]) + tool_calls = self.guardrail._extract_tool_calls_from_response(response) + assert len(tool_calls) == 0 + + @pytest.mark.asyncio + async def test_async_post_call_success_hook_no_tools(self): + response = ModelResponse(choices=[Choices(message={})]) + user_api_key_dict = UserAPIKeyAuth() + data = {"guardrails": ["test-tool-permission"]} + + with patch.object(self.guardrail, "should_run_guardrail", return_value=True): + result = await self.guardrail.async_post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response + ) + assert result is None + + @pytest.mark.asyncio + async def test_async_post_call_success_hook_with_allowed_tools(self): + tool_call = { + "function": {"name": "Bash", "arguments": "{}"}, + "type": "function", + } + response = ModelResponse(choices=[Choices(message={"tool_calls": [tool_call]})]) + user_api_key_dict = UserAPIKeyAuth() + data = {"guardrails": ["test-tool-permission"]} + + with patch.object(self.guardrail, "should_run_guardrail", return_value=True): + result = await self.guardrail.async_post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response + ) + assert result is None + + @pytest.mark.asyncio + async def test_async_post_call_success_hook_with_denied_tools_raises(self): + tool_call = { + "function": {"name": "Read", "arguments": "{}"}, + "type": "function", + } + response = ModelResponse(choices=[Choices(message={"tool_calls": [tool_call]})]) + 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_pre_call_hook_block_mode(self): + data = { + "tools": [ + {"type": "function", "function": {"name": "Bash"}}, + {"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_rewrite_mode(self): + guardrail = ToolPermissionGuardrail( + guardrail_name="test-tool-permission", + rules=self.test_rules, + default_action="deny", + on_disallowed_action="rewrite", + ) + data = { + "tools": [ + {"type": "function", "function": {"name": "Bash"}}, + {"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 "tools" in new_data + tool_names = [t["function"]["name"] for t in new_data["tools"]] + assert "Bash" in tool_names + assert "Read" not in tool_names + + def test_modify_response_with_permission_errors(self): + # Setup a response with one tool_call + tool_call = ChatCompletionMessageToolCall( + function={"name": "Read", "arguments": "{}"}, id="call_123" + ) + response = ModelResponse( + choices=[Choices(message={"tool_calls": [tool_call], "content": ""})] + ) + + # Denied tools tuple of (tool_call, PermissionError) + denied_tools = [ + ( + tool_call, + PermissionError( + tool_name="Read", + rule_id="deny_read", + message="Tool 'Bash' denied by rule 'deny_read'", + ), + ) + ] + + # Apply modifications + self.guardrail._modify_response_with_permission_errors(response, denied_tools) + + # Verify: tool_calls removed and content contains error message + choice = response.choices[0] + assert isinstance(choice, Choices) + assert choice.message.tool_calls is None or choice.message.tool_calls == [] + assert isinstance(choice.message.content, str) + assert "Permission denied" in choice.message.content + + +class TestToolPermissionGuardrailIntegration: + """Integration tests for Tool Permission Guardrail""" + + def test_default_action_allow(self): + guardrail = ToolPermissionGuardrail( + guardrail_name="test-allow-default", + rules=[{"id": "deny_read", "tool_name": "Read", "decision": "deny"}], + default_action="allow", + ) + + is_allowed, rule_id, message = guardrail._check_tool_permission("UnknownTool") + assert is_allowed is True + assert rule_id is None + assert "default" in (message or "") + + is_allowed, rule_id, _ = guardrail._check_tool_permission("Read") + assert is_allowed is False + assert rule_id == "deny_read" + + def test_empty_rules(self): + guardrail = ToolPermissionGuardrail( + guardrail_name="test-no-rules", + rules=[], + default_action="allow", + ) + + is_allowed, rule_id, message = guardrail._check_tool_permission("AnyTool") + assert is_allowed is True + assert rule_id is None + assert "default" in (message or "")