x
This commit is contained in:
+244
-13
@@ -1,5 +1,13 @@
|
||||
"""
|
||||
工具路由器模块
|
||||
|
||||
支持动态注册和管理工具
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
import logging
|
||||
import time
|
||||
from typing import Any, Callable, Dict, List, Optional, Type
|
||||
|
||||
from langchain_core.tools import BaseTool
|
||||
|
||||
@@ -7,26 +15,159 @@ 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):
|
||||
if tools is None:
|
||||
tools = [CalculatorTool(), WebSearchTool(), RestApiTool(), SrApiQueryTool()]
|
||||
self._tools: Dict[str, BaseTool] = {tool.name: tool for tool in tools}
|
||||
|
||||
def list_tools(self) -> list[str]:
|
||||
"""
|
||||
工具路由器:统一调用入口
|
||||
|
||||
支持特性:
|
||||
- 动态注册工具
|
||||
- 工具元数据管理
|
||||
- 执行监控
|
||||
"""
|
||||
|
||||
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)
|
||||
@@ -34,8 +175,98 @@ class ToolRouter:
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user