Coverage for merco/tools/registry.py: 95%
37 statements
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-07 14:04 +0800
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-07 14:04 +0800
1"""工具注册中心 — 管理所有可用工具,支持 toolset 过滤、可用性检查、动态描述"""
3from typing import Optional
4from .base import BaseTool
5from merco.tools.middleware import ToolContext, ToolMiddlewareChain
8class ToolRegistry:
9 """中央工具注册表
11 支持:
12 - toolset 分组:通过 set_enabled_toolsets() 控制启用哪些分组
13 - check_fn 过滤:工具不可用时自动从 LLM 列表中移除
14 - 动态描述:get_definitions(context) 可传入运行时上下文增强工具描述
15 - 安全守卫:所有工具执行前通过 ToolGuard 检查
16 """
18 def __init__(self):
19 self._tools: dict[str, BaseTool] = {}
20 self._enabled_toolsets: set | None = None # None = 全部启用
21 self._middleware = ToolMiddlewareChain()
23 def use(self, middleware) -> "ToolRegistry":
24 """挂载中间件"""
25 self._middleware.use(middleware)
26 return self
28 def register(self, tool: BaseTool):
29 """注册一个工具"""
30 self._tools[tool.name] = tool
32 def unregister(self, name: str):
33 """注销一个工具"""
34 self._tools.pop(name, None)
36 def get(self, name: str) -> Optional[BaseTool]:
37 """获取工具"""
38 return self._tools.get(name)
40 def list_tools(self) -> list[BaseTool]:
41 """列出所有已注册工具(不区分 toolset)"""
42 return list(self._tools.values())
44 def set_enabled_toolsets(self, toolsets: list[str] | None):
45 """设置启用的 toolset 分组。None = 全部启用。
47 工具注册后调用此方法限制可见范围。
48 例如 set_enabled_toolsets(["file", "bash"]) 只暴露文件操作和终端工具。
49 """
50 self._enabled_toolsets = set(toolsets) if toolsets is not None else None
52 def get_definitions(self, context: dict | None = None) -> list[dict]:
53 """获取可用工具定义(用于 LLM function calling)
55 过滤规则:
56 1. check() 返回 False 的工具被排除
57 2. 如果设置了 enabled_toolsets,只包含匹配分组的工具
59 context 传给 tool.describe() 用于动态描述增强。
60 """
61 definitions = []
62 for tool in self._tools.values():
63 # 可用性检查
64 if not tool.check():
65 continue
66 # toolset 过滤
67 if self._enabled_toolsets is not None and tool.toolset not in self._enabled_toolsets:
68 continue
69 definitions.append(tool.get_definition(context))
70 return definitions
72 async def execute(self, tool_name: str, **kwargs) -> dict:
73 """执行指定工具。中间件链处理安全检查和错误处理。"""
74 tool = self.get(tool_name)
75 if tool is None:
76 return {"error": f"工具 '{tool_name}' 不存在"}
78 ctx = ToolContext(tool_name=tool_name, arguments=kwargs, tool=tool)
79 return await self._middleware.execute(ctx, lambda: tool.execute(**kwargs))
82# 模块级全局单例 — 工具模块在 import 时通过此实例自注册
83tool_registry = ToolRegistry()