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