pipeline_core/pipeline_core/tool_registry.py

167 lines
4.9 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.

"""
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