mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
Merge a21d5b49c2 into b94f5bdbed
This commit is contained in:
commit
163914ca23
2 changed files with 138 additions and 23 deletions
|
|
@ -14,6 +14,7 @@ from litellm.integrations.custom_guardrail import (
|
|||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import anthropic_tool_names
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
|
|
@ -40,6 +41,7 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
GUARDRAIL_NAME: Final = "tool_permission"
|
||||
_RESPONSES_CALL_TYPES: Final = frozenset({"responses", "aresponses", "_aresponses_websocket"})
|
||||
|
||||
|
||||
def _object_mapping(value: object) -> Mapping[str, object] | None:
|
||||
|
|
@ -605,27 +607,33 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
if not any(_is_tool_use_block(block) for block in kept_blocks):
|
||||
response["stop_reason"] = "end_turn" # rebind-ok: dropping every tool_use ends the turn
|
||||
|
||||
def _get_request_tool_name(self, tool: object) -> tuple[str | None, str | None]:
|
||||
def _get_request_tool_targets(
|
||||
self, tool: object, call_type: CallTypesLiteral
|
||||
) -> tuple[tuple[str, str | None], ...]:
|
||||
tool_type: Final = self._get_mapping_value(tool, "type")
|
||||
if tool_type != "function":
|
||||
return None, tool_type
|
||||
|
||||
function: Final = self._get_mapping_value(tool, "function")
|
||||
tool_name: Final = self._get_mapping_value(function, "name")
|
||||
return tool_name, tool_type
|
||||
normalized_type: Final = (
|
||||
tool_type if call_type in _RESPONSES_CALL_TYPES or tool_type not in (None, "custom") else "function"
|
||||
)
|
||||
return tuple((name, normalized_type) for name in anthropic_tool_names(tool))
|
||||
|
||||
def _get_legacy_function_name(self, function: object) -> str | None:
|
||||
return self._get_mapping_value(function, "name")
|
||||
|
||||
def _get_named_tool_choice(self, data: dict) -> str | None:
|
||||
def _get_named_tool_choice(self, data: Mapping[str, object]) -> str | None:
|
||||
tool_choice: Final = data.get("tool_choice")
|
||||
if not tool_choice or tool_choice in ("auto", "none", "required"):
|
||||
return None
|
||||
if isinstance(tool_choice, str):
|
||||
return tool_choice
|
||||
if self._get_mapping_value(tool_choice, "type") != "function":
|
||||
if self._get_mapping_value(tool_choice, "type") not in ("tool", "function"):
|
||||
return None
|
||||
return self._get_mapping_value(self._get_mapping_value(tool_choice, "function"), "name")
|
||||
function_name: Final = self._get_mapping_value(self._get_mapping_value(tool_choice, "function"), "name")
|
||||
return function_name or self._get_mapping_value(tool_choice, "name")
|
||||
|
||||
@staticmethod
|
||||
def _is_anthropic_tool_choice(data: Mapping[str, object]) -> bool:
|
||||
tool_choice: Final = _object_mapping(data.get("tool_choice"))
|
||||
return tool_choice is not None and tool_choice.get("type") == "tool"
|
||||
|
||||
def _get_named_function_call(self, data: dict) -> str | None:
|
||||
function_call: Final = data.get("function_call")
|
||||
|
|
@ -635,13 +643,11 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
return function_call
|
||||
return self._get_mapping_value(function_call, "name")
|
||||
|
||||
def _collect_request_tools(self, data: dict) -> list[tuple[str, str | None]]:
|
||||
def _collect_request_tools(self, data: dict, call_type: CallTypesLiteral) -> list[tuple[str, str | None]]:
|
||||
request_tools: Final[list[tuple[str, str | None]]] = []
|
||||
|
||||
for tool in data.get("tools") or []:
|
||||
tool_name, tool_type = self._get_request_tool_name(tool)
|
||||
if tool_name is not None:
|
||||
request_tools.append((tool_name, tool_type))
|
||||
request_tools.extend(self._get_request_tool_targets(tool, call_type))
|
||||
|
||||
for function in data.get("functions") or []:
|
||||
function_name = self._get_legacy_function_name(function)
|
||||
|
|
@ -681,13 +687,9 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
|
||||
tools: Final[list[ChatCompletionToolParam] | None] = data.get("tools")
|
||||
if tools is not None:
|
||||
new_tools: Final = []
|
||||
for tool in tools:
|
||||
tool_name, tool_type = self._get_request_tool_name(tool)
|
||||
if tool_type == "function" and tool_name in error_tool_names:
|
||||
continue
|
||||
new_tools.append(tool)
|
||||
data["tools"] = new_tools
|
||||
data["tools"] = [
|
||||
tool for tool in tools if not any(name in error_tool_names for name in anthropic_tool_names(tool))
|
||||
]
|
||||
|
||||
functions: Final = data.get("functions")
|
||||
if functions is not None:
|
||||
|
|
@ -697,7 +699,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
|
||||
named_tool_choice: Final = self._get_named_tool_choice(data)
|
||||
if named_tool_choice in error_tool_names:
|
||||
data["tool_choice"] = "none"
|
||||
data["tool_choice"] = {"type": "none"} if self._is_anthropic_tool_choice(data) else "none"
|
||||
|
||||
named_function_call: Final = self._get_named_function_call(data)
|
||||
if named_function_call in error_tool_names:
|
||||
|
|
@ -807,7 +809,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
|
||||
return data
|
||||
|
||||
new_tools: Final = self._collect_request_tools(data)
|
||||
new_tools: Final = self._collect_request_tools(data, call_type)
|
||||
if not new_tools:
|
||||
verbose_proxy_logger.debug(
|
||||
"Tool Permission Guardrail: not running guardrail. No tools or functions in data"
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Unit tests for Tool Permission Guardrail (OpenAI tool_calls semantics)
|
|||
import json
|
||||
import logging
|
||||
import re
|
||||
from typing import Literal
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -24,6 +25,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import (
|
|||
PermissionError,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
CallTypesLiteral,
|
||||
ChatCompletionMessageToolCall,
|
||||
Choices,
|
||||
ModelResponse,
|
||||
|
|
@ -1177,6 +1179,117 @@ class TestToolPermissionGuardrailAnthropicMessages:
|
|||
def _tool_use(self, name, tool_id="tu_1"):
|
||||
return {"type": "tool_use", "id": tool_id, "name": name, "input": {"command": "ls"}}
|
||||
|
||||
def _always_on_pre_call(self, on_disallowed_action: Literal["block", "rewrite"]) -> ToolPermissionGuardrail:
|
||||
return ToolPermissionGuardrail(
|
||||
guardrail_name=f"anthropic-pre-call-{on_disallowed_action}",
|
||||
rules=self.rules,
|
||||
default_action="deny",
|
||||
on_disallowed_action=on_disallowed_action,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"denied_tool",
|
||||
[
|
||||
{"name": "Read", "input_schema": {"type": "object", "properties": {}}},
|
||||
{"type": "custom", "name": "Read", "input_schema": {"type": "object", "properties": {}}},
|
||||
{"type": "function", "name": "Read", "parameters": {"type": "object", "properties": {}}},
|
||||
],
|
||||
ids=["anthropic", "anthropic_custom_type", "responses_api_flat_function"],
|
||||
)
|
||||
async def test_pre_call_blocks_denied_request_tool_in_flat_format(self, denied_tool: dict[str, object]) -> None:
|
||||
data = {"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "hi"}], "tools": [denied_tool]}
|
||||
|
||||
with pytest.raises(HTTPException) as excinfo:
|
||||
await self._always_on_pre_call("block").async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(default_in_memory_ttl=1),
|
||||
data=data,
|
||||
call_type="anthropic_messages",
|
||||
)
|
||||
|
||||
assert excinfo.value.status_code == 400
|
||||
assert excinfo.value.detail["detection_message"] == "Tool 'Read' denied by rule 'deny_read'"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("rule_tool_type", "call_type"),
|
||||
[(r"^custom$", "responses"), (r"^function$", "anthropic_messages")],
|
||||
ids=["responses_custom_stays_custom", "anthropic_custom_reads_as_function"],
|
||||
)
|
||||
async def test_pre_call_tool_type_rules_follow_the_request_format(
|
||||
self, rule_tool_type: str, call_type: CallTypesLiteral
|
||||
) -> None:
|
||||
guardrail = ToolPermissionGuardrail(
|
||||
guardrail_name=f"tool-type-{call_type}",
|
||||
rules=[{"id": "deny_type", "tool_type": rule_tool_type, "decision": "deny"}],
|
||||
default_action="allow",
|
||||
on_disallowed_action="block",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
default_on=True,
|
||||
)
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"tools": [{"type": "custom", "name": "apply_patch", "input_schema": {"type": "object", "properties": {}}}],
|
||||
}
|
||||
|
||||
with pytest.raises(HTTPException) as excinfo:
|
||||
await guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(default_in_memory_ttl=1),
|
||||
data=data,
|
||||
call_type=call_type,
|
||||
)
|
||||
|
||||
assert excinfo.value.status_code == 400
|
||||
assert excinfo.value.detail["detection_message"] == "Tool 'apply_patch' denied by rule 'deny_type'"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("tool_shape", "tool_choice", "call_type", "expected_tool_choice"),
|
||||
[
|
||||
(
|
||||
{"input_schema": {"type": "object", "properties": {}}},
|
||||
{"type": "tool", "name": "Read"},
|
||||
"anthropic_messages",
|
||||
{"type": "none"},
|
||||
),
|
||||
(
|
||||
{"type": "function", "parameters": {"type": "object", "properties": {}}},
|
||||
{"type": "function", "name": "Read"},
|
||||
"responses",
|
||||
"none",
|
||||
),
|
||||
],
|
||||
ids=["anthropic", "responses_api"],
|
||||
)
|
||||
async def test_pre_call_rewrite_strips_denied_flat_tool_and_forced_choice(
|
||||
self,
|
||||
tool_shape: dict[str, object],
|
||||
tool_choice: dict[str, str],
|
||||
call_type: CallTypesLiteral,
|
||||
expected_tool_choice: dict[str, str] | str,
|
||||
) -> None:
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"tools": [{"name": "Bash", **tool_shape}, {"name": "Read", **tool_shape}],
|
||||
"tool_choice": tool_choice,
|
||||
}
|
||||
|
||||
result = await self._always_on_pre_call("rewrite").async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(default_in_memory_ttl=1),
|
||||
data=data,
|
||||
call_type=call_type,
|
||||
)
|
||||
|
||||
assert [tool["name"] for tool in result["tools"]] == ["Bash"]
|
||||
assert result["tool_choice"] == expected_tool_choice
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_denied_anthropic_tool_use_is_blocked(self):
|
||||
response = self._response({"type": "text", "text": "reading"}, self._tool_use("Read"))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue