273 lines
8.1 KiB
Python
273 lines
8.1 KiB
Python
"""
|
||
工具路由器模块
|
||
|
||
支持动态注册和管理工具
|
||
"""
|
||
|
||
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)
|