Files
more_dots/services/tool_router.py
T
2026-03-11 23:40:39 +08:00

273 lines
8.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
工具路由器模块
支持动态注册和管理工具
"""
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)