Coverage for agentos/tools/registry.py: 50%
40 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-10 01:30 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-10 01:30 +0800
1"""
2统一工具注册表 — 核心循环不关心具体实现。
3"""
5from __future__ import annotations
7import asyncio
8import uuid
10from agentos.tools.base import BaseTool, ToolCall, ToolResult
13class ToolRegistry:
14 """统一工具注册表。所有工具在这里注册,核心循环不关心具体实现。"""
16 def __init__(self):
17 self._tools: dict[str, BaseTool] = {}
19 def register(self, tool: BaseTool):
20 self._tools[tool.name] = tool
22 def register_many(self, tools: list[BaseTool]):
23 for tool in tools:
24 self.register(tool)
26 def get(self, name: str) -> BaseTool | None:
27 return self._tools.get(name)
29 def list_names(self) -> list[str]:
30 return list(self._tools.keys())
32 def get_schemas_for_model(self, model_type: str) -> list[dict]:
33 """根据模型类型生成工具schema。"""
34 if model_type in ("openai", "deepseek", "kimi", "qwen", "glm", "minimax"):
35 return [t.to_openai_schema() for t in self._tools.values()]
36 elif model_type == "anthropic":
37 return [t.to_anthropic_schema() for t in self._tools.values()]
38 else:
39 return [t.to_openai_schema() for t in self._tools.values()]
41 async def execute_batch(self, calls: list[ToolCall], sandbox=None) -> list[ToolResult]:
42 """并行执行一组工具调用。"""
43 tasks = []
44 for call in calls:
45 tool = self._tools.get(call.name)
46 if not tool:
47 tasks.append(self._unknown_tool_result(call))
48 else:
49 tasks.append(self._execute_one(tool, call, sandbox))
50 return await asyncio.gather(*tasks)
52 async def _execute_one(self, tool: BaseTool, call: ToolCall, sandbox=None) -> ToolResult:
53 try:
54 return await tool.execute(call.arguments, sandbox=sandbox)
55 except Exception as e:
56 return ToolResult(call_id=call.id, error=str(e))
58 async def _unknown_tool_result(self, call: ToolCall) -> ToolResult:
59 return ToolResult(
60 call_id=call.id,
61 error=f"Unknown tool: {call.name}. Available: {self.list_names()}",
62 )
64 @staticmethod
65 def make_call_id() -> str:
66 return f"call_{uuid.uuid4().hex[:12]}"