167 lines
4.9 KiB
Python
167 lines
4.9 KiB
Python
"""
|
||
pipeline-core: Tool Registry — 工具注册与发现
|
||
|
||
对照 Hermes Agent 的 tools/registry.py:
|
||
- 集中注册所有可用工具
|
||
- 生成 LLM function-calling schema
|
||
- 按产线启用/禁用工具
|
||
- 工具分类管理
|
||
"""
|
||
|
||
import json
|
||
import logging
|
||
from typing import Callable, Dict, List, Optional
|
||
|
||
from .agent_config import ToolDefinition
|
||
|
||
logger = logging.getLogger("pipeline.tool_registry")
|
||
|
||
|
||
class ToolRegistry:
|
||
"""全局工具注册表。每个工具包含定义(schema)和执行函数。"""
|
||
|
||
def __init__(self):
|
||
self._tools: Dict[str, ToolDefinition] = {}
|
||
self._handlers: Dict[str, Callable] = {}
|
||
|
||
def register(
|
||
self,
|
||
tool: ToolDefinition,
|
||
handler: Callable = None,
|
||
):
|
||
"""注册一个工具。
|
||
|
||
Args:
|
||
tool: 工具定义(名称、描述、参数 schema)
|
||
handler: 执行函数 async def handler(params, context) -> str
|
||
"""
|
||
self._tools[tool.name] = tool
|
||
if handler:
|
||
self._handlers[tool.name] = handler
|
||
logger.debug(f"ToolRegistry: registered {tool.name}")
|
||
|
||
def unregister(self, name: str):
|
||
"""移除工具"""
|
||
self._tools.pop(name, None)
|
||
self._handlers.pop(name, None)
|
||
|
||
def enable(self, name: str):
|
||
"""启用工具"""
|
||
if name in self._tools:
|
||
self._tools[name].enabled = True
|
||
|
||
def disable(self, name: str):
|
||
"""禁用工具"""
|
||
if name in self._tools:
|
||
self._tools[name].enabled = False
|
||
|
||
def get(self, name: str) -> Optional[ToolDefinition]:
|
||
"""获取工具定义"""
|
||
return self._tools.get(name)
|
||
|
||
def get_handler(self, name: str) -> Optional[Callable]:
|
||
"""获取工具执行函数"""
|
||
return self._handlers.get(name)
|
||
|
||
def list_enabled(self, category: str = None) -> List[ToolDefinition]:
|
||
"""列出启用的工具,可按分类筛选"""
|
||
tools = [t for t in self._tools.values() if t.enabled]
|
||
if category:
|
||
tools = [t for t in tools if t.category == category]
|
||
return tools
|
||
|
||
def to_openai_schema(self, category: str = None) -> List[dict]:
|
||
"""生成 OpenAI function-calling 格式的 tool 定义。
|
||
|
||
等价于 HA 在 system prompt 中注入的 tool descriptions。
|
||
"""
|
||
tools = self.list_enabled(category)
|
||
return [
|
||
{
|
||
"type": "function",
|
||
"function": {
|
||
"name": t.name,
|
||
"description": t.description,
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": {
|
||
k: {"type": "string", "description": v}
|
||
for k, v in (t.parameters or {}).items()
|
||
},
|
||
"required": list(t.parameters.keys()) if t.parameters else [],
|
||
},
|
||
},
|
||
}
|
||
for t in tools
|
||
]
|
||
|
||
def to_text_description(self, category: str = None) -> str:
|
||
"""生成纯文本工具描述(用于不支持 native function calling 的模型)。
|
||
|
||
等价于 HA 的 tool description 注入到 system prompt。
|
||
"""
|
||
tools = self.list_enabled(category)
|
||
lines = []
|
||
for t in tools:
|
||
params = ", ".join(
|
||
f"{k}: {v}" for k, v in (t.parameters or {}).items()
|
||
)
|
||
lines.append(f"- {t.name}({params}): {t.description}")
|
||
return "\n".join(lines)
|
||
|
||
def get_tool_names(self) -> List[str]:
|
||
"""获取所有已注册工具名称"""
|
||
return list(self._tools.keys())
|
||
|
||
def get_by_category(self) -> Dict[str, List[str]]:
|
||
"""按分类分组工具名"""
|
||
groups: Dict[str, List[str]] = {}
|
||
for t in self._tools.values():
|
||
groups.setdefault(t.category, []).append(t.name)
|
||
return groups
|
||
|
||
|
||
# ── 全局单例 ──
|
||
|
||
_registry: Optional[ToolRegistry] = None
|
||
|
||
|
||
def get_tool_registry() -> ToolRegistry:
|
||
"""获取全局工具注册表单例"""
|
||
global _registry
|
||
if _registry is None:
|
||
_registry = ToolRegistry()
|
||
return _registry
|
||
|
||
|
||
# ── 便捷函数 ──
|
||
|
||
def register_tool(
|
||
name: str,
|
||
description: str,
|
||
parameters: dict = None,
|
||
category: str = "general",
|
||
requires_confirmation: bool = False,
|
||
):
|
||
"""装饰器:注册工具。
|
||
|
||
用法:
|
||
@register_tool("search_web", "搜索网页", {"query": "搜索关键词"})
|
||
async def search_web(params, context):
|
||
...
|
||
"""
|
||
|
||
def decorator(handler: Callable):
|
||
tool = ToolDefinition(
|
||
name=name,
|
||
description=description,
|
||
parameters=parameters or {},
|
||
category=category,
|
||
requires_confirmation=requires_confirmation,
|
||
)
|
||
registry = get_tool_registry()
|
||
registry.register(tool, handler)
|
||
return handler
|
||
|
||
return decorator
|