Merge pull request #14519 from uc4w6c/feat/add_tools_permission_guardrail

feat: add tool-permission guardrail
This commit is contained in:
Krish Dholakia 2025-09-13 23:22:31 -07:00 • committed by GitHub
commit 11822e63f1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 1094 additions and 6 deletions

View file

@ -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
<Tabs>
<TabItem value="block" label="Block Request">
**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"
}
}
```
</TabItem>
<TabItem value="rewrite" label="Rewrite Request">
**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"
}
```
</TabItem>
</Tabs>

View file

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

View file

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

View file

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

View file

@ -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]] = {}

View file

@ -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):
@ -381,6 +379,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"
@ -469,6 +479,7 @@ class LitellmParams(
LassoGuardrailConfigModel,
PillarGuardrailConfigModel,
NomaGuardrailConfigModel,
ToolPermissionGuardrailConfigModel,
BaseLitellmParams,
):
guardrail: str = Field(description="The type of guardrail integration to use")

View file

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

View file

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