""" 工具路由器模块 支持动态注册和管理工具 """ import json import logging import time from typing import Any, Callable, Dict, List, Optional, Type from langchain_core.tools import BaseTool from tools.calculator import CalculatorTool from tools.web_search import WebSearchTool from tools.rest_api_tool import RestApiTool from tools.sr_api_tool import SrApiQueryTool from core.registry import ToolRegistry, ToolMetadata logger = logging.getLogger(__name__) class ToolRouter: """ 工具路由器:统一调用入口 支持特性: - 动态注册工具 - 工具元数据管理 - 执行监控 """ def __init__(self, tools: Optional[List[BaseTool]] = None): self._tools: Dict[str, BaseTool] = {} self._tool_metadata: Dict[str, ToolMetadata] = {} self._execution_stats: Dict[str, Dict[str, Any]] = {} if tools is not None: for tool in tools: self.register_tool(tool) else: self._register_default_tools() def _register_default_tools(self) -> None: """注册默认工具""" default_tools = [ CalculatorTool(), WebSearchTool(), RestApiTool(), SrApiQueryTool(), ] for tool in default_tools: self.register_tool(tool) def register_tool( self, tool: BaseTool, description: str = "", version: str = "1.0.0", timeout: int = 30, retry: int = 0, tags: Optional[List[str]] = None, ) -> None: """ 注册工具 Args: tool: 工具实例 description: 描述(默认使用 tool.description) version: 版本 timeout: 超时时间 retry: 重试次数 tags: 标签 """ name = tool.name metadata = ToolMetadata( name=name, description=description or tool.description, version=version, timeout=timeout, retry=retry, tags=tags or [], ) self._tools[name] = tool self._tool_metadata[name] = metadata self._execution_stats[name] = { "total_calls": 0, "success_calls": 0, "failed_calls": 0, "total_time_ms": 0, } ToolRegistry._entries[name] = type( "RegistryEntry", (), {"instance": tool, "metadata": {"tool_metadata": metadata}} )() logger.info(f"Registered tool: {name} (v{version})") def unregister_tool(self, name: str) -> bool: """ 注销工具 Args: name: 工具名称 Returns: 是否成功注销 """ if name in self._tools: del self._tools[name] del self._tool_metadata[name] del self._execution_stats[name] ToolRegistry.unregister(name) logger.info(f"Unregistered tool: {name}") return True return False def get_tool(self, name: str) -> Optional[BaseTool]: """获取工具实例""" return self._tools.get(name) def get_tool_metadata(self, name: str) -> Optional[ToolMetadata]: """获取工具元数据""" return self._tool_metadata.get(name) def list_tools(self) -> List[str]: """列出可用工具名称""" return list(self._tools.keys()) def get_tool_info(self, name: str) -> Optional[Dict[str, Any]]: """获取工具详细信息""" if name not in self._tools: return None tool = self._tools[name] metadata = self._tool_metadata.get(name) stats = self._execution_stats.get(name, {}) return { "name": name, "description": metadata.description if metadata else tool.description, "version": metadata.version if metadata else "unknown", "timeout": metadata.timeout if metadata else 30, "tags": metadata.tags if metadata else [], "stats": { "total_calls": stats.get("total_calls", 0), "success_rate": self._calculate_success_rate(name), }, } def call(self, tool_name: str, payload: Any) -> Dict[str, Any]: """ 调用工具并返回标准化结果 Args: tool_name: 工具名称 payload: 输入参数 Returns: 标准化结果 {ok, data, error} """ tool = self._tools.get(tool_name) if not tool: return {"ok": False, "data": None, "error": f"工具不存在: {tool_name}"} start_time = time.time() try: if isinstance(payload, (dict, list)): input_value = json.dumps(payload, ensure_ascii=False) elif payload is None: input_value = "" else: input_value = str(payload) result = tool.run(input_value) self._record_success(tool_name, time.time() - start_time) return {"ok": True, "data": result, "error": None} except Exception as e: self._record_failure(tool_name, time.time() - start_time) return {"ok": False, "data": None, "error": str(e)} def call_with_metadata( self, tool_name: str, payload: Any, ) -> Dict[str, Any]: """ 调用工具并返回包含元数据的结果 Args: tool_name: 工具名称 payload: 输入参数 Returns: 包含元数据的结果 """ result = self.call(tool_name, payload) metadata = self.get_tool_metadata(tool_name) return { **result, "tool_name": tool_name, "tool_version": metadata.version if metadata else "unknown", "execution_time_ms": self._execution_stats.get(tool_name, {}).get("last_time_ms", 0), } def _record_success(self, tool_name: str, elapsed: float) -> None: """记录成功执行""" if tool_name in self._execution_stats: stats = self._execution_stats[tool_name] stats["total_calls"] += 1 stats["success_calls"] += 1 stats["total_time_ms"] += elapsed * 1000 stats["last_time_ms"] = elapsed * 1000 def _record_failure(self, tool_name: str, elapsed: float) -> None: """记录失败执行""" if tool_name in self._execution_stats: stats = self._execution_stats[tool_name] stats["total_calls"] += 1 stats["failed_calls"] += 1 stats["total_time_ms"] += elapsed * 1000 stats["last_time_ms"] = elapsed * 1000 def _calculate_success_rate(self, tool_name: str) -> float: """计算成功率""" stats = self._execution_stats.get(tool_name) if not stats or stats["total_calls"] == 0: return 0.0 return stats["success_calls"] / stats["total_calls"] def get_all_stats(self) -> Dict[str, Dict[str, Any]]: """获取所有工具的执行统计""" result = {} for name in self._tools: result[name] = { **self._execution_stats.get(name, {}), "success_rate": self._calculate_success_rate(name), } return result def register_function( self, name: str, func: Callable, description: str = "", timeout: int = 30, ) -> None: """ 将普通函数注册为工具 Args: name: 工具名称 func: 函数 description: 描述 timeout: 超时时间 """ from langchain_core.tools import Tool tool = Tool( name=name, description=description, func=func, ) self.register_tool(tool, description=description, timeout=timeout)