agentchat-task-agent 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- agentchat_task_agent-0.1.0.dist-info/METADATA +251 -0
- agentchat_task_agent-0.1.0.dist-info/RECORD +21 -0
- agentchat_task_agent-0.1.0.dist-info/WHEEL +5 -0
- agentchat_task_agent-0.1.0.dist-info/entry_points.txt +3 -0
- agentchat_task_agent-0.1.0.dist-info/licenses/LICENSE +21 -0
- agentchat_task_agent-0.1.0.dist-info/top_level.txt +1 -0
- task_agent/__init__.py +44 -0
- task_agent/cli.py +98 -0
- task_agent/config.py +32 -0
- task_agent/demo.py +132 -0
- task_agent/executor.py +42 -0
- task_agent/graph.py +314 -0
- task_agent/judge.py +48 -0
- task_agent/llm.py +19 -0
- task_agent/memory.py +45 -0
- task_agent/nodes.py +394 -0
- task_agent/prompts.py +120 -0
- task_agent/py.typed +1 -0
- task_agent/state.py +36 -0
- task_agent/telemetry.py +45 -0
- task_agent/tools.py +155 -0
task_agent/tools.py
ADDED
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
"""工具调用执行器与内置纯计算工具(零第三方依赖)。"""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import ast
|
|
5
|
+
import inspect
|
|
6
|
+
import operator
|
|
7
|
+
import random
|
|
8
|
+
from dataclasses import dataclass, field
|
|
9
|
+
from datetime import datetime, timezone
|
|
10
|
+
from typing import Any, Callable
|
|
11
|
+
|
|
12
|
+
from task_agent.executor import ExecuteRequest, StepResult
|
|
13
|
+
from task_agent.llm import LLMFactory, llm_text
|
|
14
|
+
from task_agent.nodes import _jump_json
|
|
15
|
+
from task_agent.prompts import TOOLCALL_PROMPT
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@dataclass(frozen=True)
|
|
19
|
+
class Tool:
|
|
20
|
+
"""工具声明:name / description / 参数说明(JSON Schema 子集,零依赖)。"""
|
|
21
|
+
|
|
22
|
+
name: str
|
|
23
|
+
description: str
|
|
24
|
+
parameters: dict[str, dict[str, Any]] = field(default_factory=dict)
|
|
25
|
+
func: Callable[..., Any] | None = None
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def _safe_calc(expression: str) -> str:
|
|
29
|
+
"""AST 白名单安全求值(禁 eval / 属性访问 / IO)。"""
|
|
30
|
+
_OPS = {
|
|
31
|
+
ast.Add: operator.add,
|
|
32
|
+
ast.Sub: operator.sub,
|
|
33
|
+
ast.Mult: operator.mul,
|
|
34
|
+
ast.Div: operator.truediv,
|
|
35
|
+
ast.FloorDiv: operator.floordiv,
|
|
36
|
+
ast.Mod: operator.mod,
|
|
37
|
+
ast.Pow: operator.pow,
|
|
38
|
+
ast.USub: operator.neg,
|
|
39
|
+
ast.UAdd: operator.pos,
|
|
40
|
+
}
|
|
41
|
+
|
|
42
|
+
def _eval(node: ast.AST):
|
|
43
|
+
if isinstance(node, ast.Expression):
|
|
44
|
+
return _eval(node.body)
|
|
45
|
+
if isinstance(node, ast.Constant) and isinstance(node.value, (int, float)):
|
|
46
|
+
return node.value
|
|
47
|
+
if isinstance(node, ast.BinOp) and type(node.op) in _OPS:
|
|
48
|
+
return _OPS[type(node.op)](_eval(node.left), _eval(node.right))
|
|
49
|
+
if isinstance(node, ast.UnaryOp) and type(node.op) in _OPS:
|
|
50
|
+
return _OPS[type(node.op)](_eval(node.operand))
|
|
51
|
+
raise ValueError("不支持的表达式")
|
|
52
|
+
|
|
53
|
+
try:
|
|
54
|
+
return str(_eval(ast.parse(expression, mode="eval")))
|
|
55
|
+
except Exception as exc:
|
|
56
|
+
return f"计算失败:{exc}"
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def _current_time(timezone_name: str = "Asia/Shanghai") -> str:
|
|
60
|
+
try:
|
|
61
|
+
from zoneinfo import ZoneInfo
|
|
62
|
+
|
|
63
|
+
return datetime.now(ZoneInfo(timezone_name)).strftime("%Y-%m-%d %H:%M:%S %Z")
|
|
64
|
+
except Exception:
|
|
65
|
+
return datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S UTC")
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _random_number(lo: float = 1, hi: float = 100) -> str:
|
|
69
|
+
return str(random.randint(int(lo), int(hi)))
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
builtin_tools: list[Tool] = [
|
|
73
|
+
Tool(
|
|
74
|
+
"calculator",
|
|
75
|
+
"安全计算数学表达式,如 '1 + 2 * 3'",
|
|
76
|
+
{"expression": {"type": "string", "description": "数学表达式"}},
|
|
77
|
+
_safe_calc,
|
|
78
|
+
),
|
|
79
|
+
Tool(
|
|
80
|
+
"current_time",
|
|
81
|
+
"获取指定时区的当前时间",
|
|
82
|
+
{"timezone_name": {"type": "string", "description": "IANA 时区名,默认 Asia/Shanghai"}},
|
|
83
|
+
_current_time,
|
|
84
|
+
),
|
|
85
|
+
Tool(
|
|
86
|
+
"random_number",
|
|
87
|
+
"生成指定范围内的随机整数",
|
|
88
|
+
{
|
|
89
|
+
"lo": {"type": "number", "description": "下限,默认 1"},
|
|
90
|
+
"hi": {"type": "number", "description": "上限,默认 100"},
|
|
91
|
+
},
|
|
92
|
+
_random_number,
|
|
93
|
+
),
|
|
94
|
+
]
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
class ToolCallingExecutor:
|
|
98
|
+
"""工具调用执行器:LLM 决定调工具或直答,最多 max_tool_calls 轮。
|
|
99
|
+
|
|
100
|
+
作为 task-agent 的 Executor 注入;把 LangChain 的"工具调用循环"概念
|
|
101
|
+
以零依赖方式实现(引擎不感知工具,保持接口缝设计)。
|
|
102
|
+
"""
|
|
103
|
+
|
|
104
|
+
def __init__(
|
|
105
|
+
self,
|
|
106
|
+
llm_factory: LLMFactory,
|
|
107
|
+
tools: list[Tool] | None = None,
|
|
108
|
+
max_tool_calls: int = 3,
|
|
109
|
+
) -> None:
|
|
110
|
+
self._llm_factory = llm_factory
|
|
111
|
+
self._tools = {t.name: t for t in (tools or [])}
|
|
112
|
+
self.max_tool_calls = max(1, max_tool_calls)
|
|
113
|
+
|
|
114
|
+
def _tool_desc(self) -> str:
|
|
115
|
+
lines = []
|
|
116
|
+
for t in self._tools.values():
|
|
117
|
+
params = ", ".join(f"{k}: {v.get('type', 'any')}" for k, v in (t.parameters or {}).items())
|
|
118
|
+
lines.append(f"- {t.name}({params}): {t.description}")
|
|
119
|
+
return "\n".join(lines) or "(无工具)"
|
|
120
|
+
|
|
121
|
+
async def __call__(self, request: ExecuteRequest) -> StepResult:
|
|
122
|
+
action = request.action
|
|
123
|
+
trace: list[str] = []
|
|
124
|
+
for _ in range(self.max_tool_calls):
|
|
125
|
+
resp = (
|
|
126
|
+
await llm_text(
|
|
127
|
+
self._llm_factory(),
|
|
128
|
+
TOOLCALL_PROMPT.format(tools=self._tool_desc(), action=action),
|
|
129
|
+
)
|
|
130
|
+
).strip()
|
|
131
|
+
data = _jump_json(resp)
|
|
132
|
+
if not data:
|
|
133
|
+
# LLM 未按 JSON 格式输出 → 视为直接回答(避免误报工具错误)
|
|
134
|
+
return StepResult(answer=resp or "(无输出)")
|
|
135
|
+
if data.get("answer") is not None:
|
|
136
|
+
text = str(data["answer"]).strip()
|
|
137
|
+
if trace:
|
|
138
|
+
text = f"{text}\n[已调用] {' | '.join(trace)}"
|
|
139
|
+
return StepResult(answer=text or "(无输出)")
|
|
140
|
+
name = str(data.get("tool") or "")
|
|
141
|
+
tool = self._tools.get(name)
|
|
142
|
+
if tool is None or tool.func is None:
|
|
143
|
+
return StepResult(
|
|
144
|
+
answer=f"工具不存在:{name or '空'}(可用:{', '.join(self._tools)})"
|
|
145
|
+
)
|
|
146
|
+
args = data.get("args") if isinstance(data.get("args"), dict) else {}
|
|
147
|
+
try:
|
|
148
|
+
result = tool.func(**args)
|
|
149
|
+
if inspect.isawaitable(result):
|
|
150
|
+
result = await result
|
|
151
|
+
trace.append(f"{name}={result}")
|
|
152
|
+
action = f"已调用 {name} 得到结果:{result}。请据此给出最终回答。"
|
|
153
|
+
except Exception as exc: # noqa: BLE001 - 工具失败返回友好错误
|
|
154
|
+
return StepResult(answer=f"工具 {name} 执行失败:{exc}")
|
|
155
|
+
return StepResult(answer="达到工具调用上限,请基于已有信息回答。")
|