mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #14519 from uc4w6c/feat/add_tools_permission_guardrail
feat: add tool-permission guardrail
This commit is contained in:
commit
11822e63f1
8 changed files with 1094 additions and 6 deletions
153
docs/my-website/docs/proxy/guardrails/tool_permission.md
Normal file
153
docs/my-website/docs/proxy/guardrails/tool_permission.md
Normal 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>
|
||||
|
|
@ -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"]
|
||||
511
litellm/proxy/guardrails/guardrail_hooks/tool_permission.py
Normal file
511
litellm/proxy/guardrails/guardrail_hooks/tool_permission.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]] = {}
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
@ -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 "")
|
||||
Loading…
Add table
Reference in a new issue