mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(guardrails): keep the Responses API custom tool type distinct
Anthropic's user-defined tools carry type "custom", so pre_call mapped that to "function" for tool_type rules. The Responses API has its own custom tool type, and mapping it let a custom tool slip past a tool_type deny rule when the default action is allow. Normalize only for the formats where custom means a plain function, and read tool names directly in the rewrite path, which never needed a type.
This commit is contained in:
parent
c25cee4e92
commit
b4b66e6bf5
2 changed files with 45 additions and 8 deletions
|
|
@ -41,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:
|
||||
|
|
@ -606,9 +607,13 @@ 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_targets(self, tool: object) -> tuple[tuple[str, 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")
|
||||
normalized_type: Final = "function" if tool_type in (None, "custom") else 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:
|
||||
|
|
@ -638,11 +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 []:
|
||||
request_tools.extend(self._get_request_tool_targets(tool))
|
||||
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)
|
||||
|
|
@ -683,9 +688,7 @@ class ToolPermissionGuardrail(CustomGuardrail):
|
|||
tools: Final[list[ChatCompletionToolParam] | None] = data.get("tools")
|
||||
if tools is not None:
|
||||
data["tools"] = [
|
||||
tool
|
||||
for tool in tools
|
||||
if not any(name in error_tool_names for name, _ in self._get_request_tool_targets(tool))
|
||||
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")
|
||||
|
|
@ -806,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"
|
||||
|
|
|
|||
|
|
@ -1213,6 +1213,40 @@ class TestToolPermissionGuardrailAnthropicMessages:
|
|||
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"),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue