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.
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="达到工具调用上限,请基于已有信息回答。")