registry.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591
  1. """
  2. Tool Registry - 工具注册表和装饰器
  3. 职责:
  4. 1. @tool 装饰器:自动注册工具并生成 Schema
  5. 2. 管理所有工具的 Schema 和实现
  6. 3. 路由工具调用到具体实现
  7. 4. 支持域名过滤、敏感数据处理、工具统计
  8. 从 Resonote/llm/tools/registry.py 抽取并扩展
  9. """
  10. import json
  11. import inspect
  12. import logging
  13. import time
  14. from typing import Any, Callable, Dict, List, Optional, Sequence
  15. from agent.tools.models import ToolCapability
  16. from agent.tools.url_matcher import filter_by_url
  17. logger = logging.getLogger(__name__)
  18. class ToolStats:
  19. """工具使用统计"""
  20. def __init__(self):
  21. self.call_count: int = 0
  22. self.success_count: int = 0
  23. self.failure_count: int = 0
  24. self.total_duration: float = 0.0
  25. self.last_called: Optional[float] = None
  26. @property
  27. def average_duration(self) -> float:
  28. """平均执行时间(秒)"""
  29. return self.total_duration / self.call_count if self.call_count > 0 else 0.0
  30. @property
  31. def success_rate(self) -> float:
  32. """成功率"""
  33. return self.success_count / self.call_count if self.call_count > 0 else 0.0
  34. def to_dict(self) -> Dict[str, Any]:
  35. return {
  36. "call_count": self.call_count,
  37. "success_count": self.success_count,
  38. "failure_count": self.failure_count,
  39. "average_duration": self.average_duration,
  40. "success_rate": self.success_rate,
  41. "last_called": self.last_called
  42. }
  43. class ToolRegistry:
  44. """工具注册表"""
  45. def __init__(self):
  46. self._tools: Dict[str, Dict[str, Any]] = {}
  47. self._stats: Dict[str, ToolStats] = {}
  48. def register(
  49. self,
  50. func: Callable,
  51. schema: Optional[Dict] = None,
  52. requires_confirmation: bool = False,
  53. editable_params: Optional[List[str]] = None,
  54. display: Optional[Dict[str, Dict[str, Any]]] = None,
  55. url_patterns: Optional[List[str]] = None,
  56. hidden_params: Optional[List[str]] = None,
  57. inject_params: Optional[Dict[str, Any]] = None,
  58. groups: Optional[List[str]] = None,
  59. capabilities: Optional[Sequence[ToolCapability | str]] = None,
  60. ):
  61. """
  62. 注册工具
  63. Args:
  64. func: 工具函数
  65. schema: 工具 Schema(如果为 None,自动生成)
  66. requires_confirmation: 是否需要用户确认
  67. editable_params: 允许用户编辑的参数列表
  68. display: i18n 展示信息 {"zh": {"name": "xx", "params": {...}}, "en": {...}}
  69. url_patterns: URL 模式列表(如 ["*.google.com"],None = 无限制)
  70. hidden_params: 隐藏参数列表(不生成 schema,LLM 看不到)
  71. inject_params: 注入参数规则 {param_name: injector_func}
  72. groups: 工具分组标签(如 ["core"]、["browser"]),用于 RunConfig.tool_groups 过滤
  73. capabilities: 安全副作用标签;explicit_validation 工具必须显式声明
  74. """
  75. func_name = func.__name__
  76. # 如果没有提供 Schema,自动生成
  77. if schema is None:
  78. try:
  79. from agent.tools.schema import SchemaGenerator
  80. schema = SchemaGenerator.generate(func, hidden_params=hidden_params or [])
  81. except Exception as e:
  82. logger.error(f"Failed to generate schema for {func_name}: {e}")
  83. raise
  84. self._tools[func_name] = {
  85. "func": func,
  86. "schema": schema,
  87. "url_patterns": url_patterns,
  88. "hidden_params": hidden_params or [],
  89. "inject_params": inject_params or {},
  90. "groups": groups or [],
  91. "capabilities": frozenset(ToolCapability(item) for item in capabilities or []),
  92. "ui_metadata": {
  93. "requires_confirmation": requires_confirmation,
  94. "editable_params": editable_params or [],
  95. "display": display or {}
  96. }
  97. }
  98. # 初始化统计
  99. self._stats[func_name] = ToolStats()
  100. logger.debug(
  101. f"[ToolRegistry] Registered: {func_name} "
  102. f"(requires_confirmation={requires_confirmation}, "
  103. f"editable_params={editable_params or []}, "
  104. f"url_patterns={url_patterns or 'none'})"
  105. )
  106. @staticmethod
  107. def _resolve_key_path(context: Dict[str, Any], key_path: str) -> Any:
  108. """
  109. 从 context 中按路径取值。
  110. 支持 "obj.field" 格式:第一段从 context dict 取值,后续段用 getattr。
  111. 例如 "knowledge_config.default_tags" → context["knowledge_config"].default_tags
  112. Args:
  113. context: 上下文字典
  114. key_path: 取值路径
  115. Returns:
  116. 取到的值,路径无效返回 None
  117. """
  118. parts = key_path.split(".")
  119. value = context.get(parts[0])
  120. for part in parts[1:]:
  121. if value is None:
  122. return None
  123. value = getattr(value, part, None)
  124. return value
  125. def is_registered(self, tool_name: str) -> bool:
  126. """检查工具是否已注册"""
  127. return tool_name in self._tools
  128. def get_schemas(self, tool_names: Optional[List[str]] = None) -> List[Dict]:
  129. """
  130. 获取工具 Schema
  131. Args:
  132. tool_names: 工具名称列表(None = 所有工具)
  133. Returns:
  134. OpenAI Tool Schema 列表
  135. """
  136. if tool_names is None:
  137. tool_names = list(self._tools.keys())
  138. schemas = []
  139. for name in tool_names:
  140. if name in self._tools:
  141. schemas.append(self._tools[name]["schema"])
  142. else:
  143. logger.warning(f"[ToolRegistry] Tool not found: {name}")
  144. return schemas
  145. def get_tool_names(self, current_url: Optional[str] = None, groups: Optional[List[str]] = None) -> List[str]:
  146. """
  147. 获取工具名称列表(可选 URL 过滤 + group 过滤)
  148. Args:
  149. current_url: 当前 URL(None = 不过滤 URL)
  150. groups: 工具分组白名单(None = 不过滤 group,返回所有工具)
  151. Returns:
  152. 工具名称列表
  153. """
  154. # 1. group 过滤
  155. if groups is not None:
  156. group_set = set(groups)
  157. candidates = {
  158. name for name, tool in self._tools.items()
  159. if group_set & set(tool.get("groups", []))
  160. }
  161. else:
  162. candidates = set(self._tools.keys())
  163. # 2. URL 过滤
  164. if current_url is None:
  165. return list(candidates)
  166. tool_items = [
  167. {"name": name, "url_patterns": self._tools[name]["url_patterns"]}
  168. for name in candidates
  169. ]
  170. filtered = filter_by_url(tool_items, current_url, url_field="url_patterns")
  171. return [item["name"] for item in filtered]
  172. def get_available_groups(self) -> List[str]:
  173. """获取所有已注册的工具分组"""
  174. groups = set()
  175. for tool in self._tools.values():
  176. groups.update(tool.get("groups", []))
  177. return sorted(groups)
  178. def get_capabilities(self, tool_name: str) -> frozenset[ToolCapability]:
  179. """Return declared effects; an empty set means unclassified."""
  180. tool = self._tools.get(tool_name)
  181. if not tool:
  182. return frozenset()
  183. return frozenset(tool.get("capabilities", ()))
  184. def get_schemas_for_url(self, current_url: Optional[str] = None) -> List[Dict]:
  185. """
  186. 根据当前 URL 获取匹配的工具 Schema
  187. Args:
  188. current_url: 当前 URL(None = 返回无 URL 限制的工具)
  189. Returns:
  190. 过滤后的工具 Schema 列表
  191. """
  192. tool_names = self.get_tool_names(current_url)
  193. return self.get_schemas(tool_names)
  194. async def execute(
  195. self,
  196. name: str,
  197. arguments: Dict[str, Any],
  198. uid: str = "",
  199. context: Optional[Dict[str, Any]] = None,
  200. sensitive_data: Optional[Dict[str, Any]] = None,
  201. inject_values: Optional[Dict[str, Any]] = None
  202. ) -> str:
  203. """
  204. 执行工具调用
  205. Args:
  206. name: 工具名称
  207. arguments: 工具参数
  208. uid: 用户ID(自动注入)
  209. context: 额外上下文
  210. sensitive_data: 敏感数据字典(用于替换 <secret> 占位符)
  211. Returns:
  212. JSON 字符串格式的结果
  213. """
  214. if name not in self._tools:
  215. error_msg = f"Unknown tool: {name}"
  216. logger.error(f"[ToolRegistry] {error_msg}")
  217. return json.dumps({"error": error_msg}, ensure_ascii=False)
  218. start_time = time.time()
  219. stats = self._stats[name]
  220. stats.call_count += 1
  221. stats.last_called = start_time
  222. try:
  223. func = self._tools[name]["func"]
  224. tool_info = self._tools[name]
  225. # 处理敏感数据占位符
  226. if sensitive_data:
  227. from agent.tools.sensitive import replace_sensitive_data
  228. current_url = context.get("page_url") if context else None
  229. arguments = replace_sensitive_data(arguments, sensitive_data, current_url)
  230. # 准备参数:只注入函数需要的参数
  231. sig = inspect.signature(func)
  232. # 过滤掉函数签名中不存在的参数(如 Claude SDK 发送的 {"_": true} 占位符)
  233. valid_params = set(sig.parameters.keys())
  234. kwargs = {k: v for k, v in arguments.items() if k in valid_params}
  235. # 注入隐藏参数(hidden_params)
  236. hidden_params = tool_info.get("hidden_params", [])
  237. if "uid" in hidden_params and "uid" in sig.parameters:
  238. kwargs["uid"] = uid
  239. if "context" in hidden_params and "context" in sig.parameters:
  240. kwargs["context"] = context
  241. # 注入参数(inject_params)
  242. inject_params = tool_info.get("inject_params", {})
  243. for param_name, rule in inject_params.items():
  244. if param_name not in sig.parameters:
  245. continue
  246. if not isinstance(rule, dict) or "mode" not in rule:
  247. # 兼容旧格式:直接值或 callable
  248. if param_name not in kwargs or kwargs[param_name] is None:
  249. kwargs[param_name] = rule() if callable(rule) else rule
  250. continue
  251. mode = rule["mode"]
  252. key_path = rule.get("key")
  253. # 从 context 中按路径取值
  254. value = self._resolve_key_path(context, key_path) if key_path and context else None
  255. if value is None:
  256. continue
  257. if mode == "default":
  258. # 默认值模式:LLM 未提供则注入
  259. if param_name not in kwargs or kwargs[param_name] is None:
  260. kwargs[param_name] = value
  261. elif mode == "merge":
  262. # 合并模式:框架值始终保留,LLM 可追加新内容
  263. llm_value = kwargs.get(param_name)
  264. if isinstance(value, dict):
  265. # dict: LLM 追加新 key,同名 key 以框架值为准
  266. kwargs[param_name] = {**(llm_value or {}), **value}
  267. elif isinstance(value, list):
  268. # list: 合并去重
  269. kwargs[param_name] = list(set((llm_value or []) + value))
  270. else:
  271. kwargs[param_name] = value
  272. # 执行函数
  273. if inspect.iscoroutinefunction(func):
  274. result = await func(**kwargs)
  275. else:
  276. result = func(**kwargs)
  277. # 记录成功
  278. stats.success_count += 1
  279. duration = time.time() - start_time
  280. stats.total_duration += duration
  281. # 返回结果:ToolResult 转为可序列化格式
  282. if isinstance(result, str):
  283. return result
  284. # 处理 ToolResult 对象
  285. from agent.tools.models import ToolResult
  286. if isinstance(result, ToolResult):
  287. ret = {"text": result.to_llm_message()}
  288. # Runner 消费的控制字段与业务文本分离,不能靠解析 output 猜测终止。
  289. if result.terminate_run:
  290. ret["_control"] = {
  291. "terminate_run": True,
  292. "result_summary": result.result_summary or result.long_term_memory or result.output,
  293. }
  294. # 保留images
  295. if result.images:
  296. ret["images"] = result.images
  297. # 保留tool_usage
  298. if result.tool_usage:
  299. ret["tool_usage"] = result.tool_usage
  300. # 向后兼容:只有 text 时返回字符串;有控制信息时必须保留 dict。
  301. if len(ret) == 1:
  302. return ret["text"]
  303. return ret
  304. return json.dumps(result, ensure_ascii=False, indent=2)
  305. except Exception as e:
  306. # 记录失败
  307. stats.failure_count += 1
  308. duration = time.time() - start_time
  309. stats.total_duration += duration
  310. error_msg = f"Error executing tool '{name}': {str(e)}"
  311. logger.error(f"[ToolRegistry] {error_msg}")
  312. import traceback
  313. logger.error(traceback.format_exc())
  314. return json.dumps({"error": error_msg}, ensure_ascii=False)
  315. def get_stats(self, tool_name: Optional[str] = None) -> Dict[str, Dict[str, Any]]:
  316. """
  317. 获取工具统计信息
  318. Args:
  319. tool_name: 工具名称(None = 所有工具)
  320. Returns:
  321. 统计信息字典
  322. """
  323. if tool_name:
  324. if tool_name in self._stats:
  325. return {tool_name: self._stats[tool_name].to_dict()}
  326. return {}
  327. return {name: stats.to_dict() for name, stats in self._stats.items()}
  328. def get_top_tools(self, limit: int = 10, by: str = "call_count") -> List[str]:
  329. """
  330. 获取排名靠前的工具
  331. Args:
  332. limit: 返回数量
  333. by: 排序依据(call_count, success_rate, average_duration)
  334. Returns:
  335. 工具名称列表
  336. """
  337. if by == "call_count":
  338. sorted_tools = sorted(
  339. self._stats.items(),
  340. key=lambda x: x[1].call_count,
  341. reverse=True
  342. )
  343. elif by == "success_rate":
  344. sorted_tools = sorted(
  345. self._stats.items(),
  346. key=lambda x: x[1].success_rate,
  347. reverse=True
  348. )
  349. elif by == "average_duration":
  350. sorted_tools = sorted(
  351. self._stats.items(),
  352. key=lambda x: x[1].average_duration,
  353. reverse=False # 越快越好
  354. )
  355. else:
  356. raise ValueError(f"Invalid sort by: {by}")
  357. return [name for name, _ in sorted_tools[:limit]]
  358. def check_confirmation_required(self, tool_calls: List[Dict]) -> bool:
  359. """检查是否有工具需要用户确认"""
  360. for tc in tool_calls:
  361. tool_name = tc.get("function", {}).get("name")
  362. if tool_name and tool_name in self._tools:
  363. if self._tools[tool_name]["ui_metadata"].get("requires_confirmation", False):
  364. return True
  365. return False
  366. def get_confirmation_flags(self, tool_calls: List[Dict]) -> List[bool]:
  367. """返回每个工具是否需要确认"""
  368. flags = []
  369. for tc in tool_calls:
  370. tool_name = tc.get("function", {}).get("name")
  371. if tool_name and tool_name in self._tools:
  372. flags.append(self._tools[tool_name]["ui_metadata"].get("requires_confirmation", False))
  373. else:
  374. flags.append(False)
  375. return flags
  376. def check_any_param_editable(self, tool_calls: List[Dict]) -> bool:
  377. """检查是否有任何工具允许参数编辑"""
  378. for tc in tool_calls:
  379. tool_name = tc.get("function", {}).get("name")
  380. if tool_name and tool_name in self._tools:
  381. editable_params = self._tools[tool_name]["ui_metadata"].get("editable_params", [])
  382. if editable_params:
  383. return True
  384. return False
  385. def get_editable_params_map(self, tool_calls: List[Dict]) -> Dict[str, List[str]]:
  386. """返回每个工具调用的可编辑参数列表"""
  387. params_map = {}
  388. for tc in tool_calls:
  389. tool_call_id = tc.get("id")
  390. tool_name = tc.get("function", {}).get("name")
  391. if tool_name and tool_name in self._tools:
  392. editable_params = self._tools[tool_name]["ui_metadata"].get("editable_params", [])
  393. params_map[tool_call_id] = editable_params
  394. else:
  395. params_map[tool_call_id] = []
  396. return params_map
  397. def get_ui_metadata(
  398. self,
  399. locale: str = "zh",
  400. tool_names: Optional[List[str]] = None
  401. ) -> Dict[str, Dict[str, Any]]:
  402. """
  403. 获取工具的UI元数据(用于前端展示)
  404. Returns:
  405. {
  406. "tool_name": {
  407. "display_name": "搜索笔记",
  408. "param_display_names": {"query": "搜索关键词"},
  409. "requires_confirmation": false,
  410. "editable_params": ["query"]
  411. }
  412. }
  413. """
  414. if tool_names is None:
  415. tool_names = list(self._tools.keys())
  416. metadata = {}
  417. for name in tool_names:
  418. if name not in self._tools:
  419. continue
  420. ui_meta = self._tools[name]["ui_metadata"]
  421. display = ui_meta.get("display", {}).get(locale, {})
  422. metadata[name] = {
  423. "display_name": display.get("name", name),
  424. "param_display_names": display.get("params", {}),
  425. "requires_confirmation": ui_meta.get("requires_confirmation", False),
  426. "editable_params": ui_meta.get("editable_params", [])
  427. }
  428. return metadata
  429. # 全局单例
  430. _global_registry = ToolRegistry()
  431. def tool(
  432. description: Optional[str] = None,
  433. param_descriptions: Optional[Dict[str, str]] = None,
  434. requires_confirmation: bool = False,
  435. editable_params: Optional[List[str]] = None,
  436. display: Optional[Dict[str, Dict[str, Any]]] = None,
  437. url_patterns: Optional[List[str]] = None,
  438. hidden_params: Optional[List[str]] = None,
  439. inject_params: Optional[Dict[str, Any]] = None,
  440. groups: Optional[List[str]] = None,
  441. capabilities: Optional[Sequence[ToolCapability | str]] = None,
  442. ):
  443. """
  444. 工具装饰器 - 自动注册工具并生成 Schema
  445. Args:
  446. description: 函数描述(可选,从 docstring 提取)
  447. param_descriptions: 参数描述(可选,从 docstring 提取)
  448. requires_confirmation: 是否需要用户确认(默认 False)
  449. editable_params: 允许用户编辑的参数列表
  450. display: i18n 展示信息
  451. url_patterns: URL 模式列表(如 ["*.google.com"],None = 无限制)
  452. hidden_params: 隐藏参数列表(不生成 schema,LLM 看不到)
  453. inject_params: 注入参数规则 {param_name: injector_func}
  454. groups: 工具分组标签(如 ["core"]、["browser"]),用于 RunConfig.tool_groups 过滤
  455. capabilities: 安全副作用标签(read/write/agent_spawn/external_send 等)
  456. Example:
  457. @tool(
  458. hidden_params=["context", "uid"],
  459. inject_params={
  460. "owner": lambda ctx: ctx.config.knowledge.get_owner(),
  461. },
  462. editable_params=["query"],
  463. url_patterns=["*.google.com"],
  464. display={
  465. "zh": {"name": "搜索笔记", "params": {"query": "搜索关键词"}},
  466. "en": {"name": "Search Notes", "params": {"query": "Query"}}
  467. }
  468. )
  469. async def search_blocks(
  470. query: str,
  471. limit: int = 10,
  472. owner: Optional[str] = None,
  473. context: Optional[ToolContext] = None,
  474. uid: str = ""
  475. ) -> str:
  476. '''搜索用户的笔记块'''
  477. ...
  478. """
  479. def decorator(func: Callable) -> Callable:
  480. # 注册到全局 registry
  481. _global_registry.register(
  482. func,
  483. requires_confirmation=requires_confirmation,
  484. editable_params=editable_params,
  485. display=display,
  486. url_patterns=url_patterns,
  487. hidden_params=hidden_params,
  488. inject_params=inject_params,
  489. groups=groups,
  490. capabilities=capabilities,
  491. )
  492. return func
  493. return decorator
  494. def get_tool_registry() -> ToolRegistry:
  495. """获取全局工具注册表"""
  496. return _global_registry