Files
more_dots/services/tool_router.py
T

273 lines
8.1 KiB
Python
Raw Normal View History

2026-03-11 23:40:39 +08:00
"""
工具路由器模块
支持动态注册和管理工具
"""
2026-02-26 13:43:44 +08:00
import json
2026-03-11 23:40:39 +08:00
import logging
import time
from typing import Any, Callable, Dict, List, Optional, Type
2026-02-26 13:43:44 +08:00
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
2026-03-11 23:40:39 +08:00
from core.registry import ToolRegistry, ToolMetadata
2026-02-26 13:43:44 +08:00
2026-03-11 23:40:39 +08:00
logger = logging.getLogger(__name__)
2026-02-26 13:43:44 +08:00
2026-03-11 23:40:39 +08:00
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]:
2026-02-26 13:43:44 +08:00
"""列出可用工具名称"""
return list(self._tools.keys())
2026-03-11 23:40:39 +08:00
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),
},
}
2026-02-26 13:43:44 +08:00
def call(self, tool_name: str, payload: Any) -> Dict[str, Any]:
2026-03-11 23:40:39 +08:00
"""
调用工具并返回标准化结果
Args:
tool_name: 工具名称
payload: 输入参数
Returns:
标准化结果 {ok, data, error}
"""
2026-02-26 13:43:44 +08:00
tool = self._tools.get(tool_name)
if not tool:
return {"ok": False, "data": None, "error": f"工具不存在: {tool_name}"}
2026-03-11 23:40:39 +08:00
start_time = time.time()
2026-02-26 13:43:44 +08:00
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)
2026-03-11 23:40:39 +08:00
2026-02-26 13:43:44 +08:00
result = tool.run(input_value)
2026-03-11 23:40:39 +08:00
self._record_success(tool_name, time.time() - start_time)
2026-02-26 13:43:44 +08:00
return {"ok": True, "data": result, "error": None}
2026-03-11 23:40:39 +08:00
2026-02-26 13:43:44 +08:00
except Exception as e:
2026-03-11 23:40:39 +08:00
self._record_failure(tool_name, time.time() - start_time)
2026-02-26 13:43:44 +08:00
return {"ok": False, "data": None, "error": str(e)}
2026-03-11 23:40:39 +08:00
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)