diff --git a/tests/tools/test_delegate_composite_toolsets.py b/tests/tools/test_delegate_composite_toolsets.py index b5d8aae398..fbb984ea4b 100644 --- a/tests/tools/test_delegate_composite_toolsets.py +++ b/tests/tools/test_delegate_composite_toolsets.py @@ -2,7 +2,7 @@ import unittest -from tools.delegate_tool import _expand_parent_toolsets +from tools.delegate_tool import _expand_parent_toolsets, _strip_blocked_tools class TestExpandParentToolsets(unittest.TestCase): @@ -26,6 +26,14 @@ class TestExpandParentToolsets(unittest.TestCase): child_toolsets = [t for t in toolsets if t in expanded] self.assertEqual(child_toolsets, ["web"]) + def test_composites_with_allowed_included_tools_are_not_stripped(self): + toolsets = ["safe", "hermes-gateway", "hermes-cli", "delegation", "kanban"] + + self.assertEqual( + _strip_blocked_tools(toolsets), + ["safe", "hermes-gateway", "hermes-cli"], + ) + if __name__ == "__main__": unittest.main() diff --git a/tools/delegate_tool_toolsets.py b/tools/delegate_tool_toolsets.py index be75cf95a5..0252dd9b65 100644 --- a/tools/delegate_tool_toolsets.py +++ b/tools/delegate_tool_toolsets.py @@ -5,7 +5,7 @@ from __future__ import annotations import logging from typing import List, Optional -from toolsets import TOOLSETS +from toolsets import TOOLSETS, resolve_toolset from tools.delegate_tool_config import _get_inherit_mcp_toolsets logger = logging.getLogger("tools.delegate_tool") # log-record parity with the origin module @@ -51,7 +51,10 @@ def _strip_blocked_tools(toolsets: List[str]) -> List[str]: """Remove toolsets whose tools are ALL blocked (derived from DELEGATE_BLOCKED_TOOLS so the two can't drift) plus composite toolsets children must never get (``delegation``, ``kanban``).""" blocked_toolset_names = {"delegation", "kanban"} | { - name for name, defn in TOOLSETS.items() if all(t in DELEGATE_BLOCKED_TOOLS for t in defn.get("tools", [])) + name + for name in TOOLSETS + if (resolved_tools := resolve_toolset(name, include_registry=False)) + and all(tool in DELEGATE_BLOCKED_TOOLS for tool in resolved_tools) } return [t for t in toolsets if t not in blocked_toolset_names]