agstack 2.0.0__tar.gz → 2.1.0__tar.gz
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.
- {agstack-2.0.0 → agstack-2.1.0}/PKG-INFO +1 -1
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/__init__.py +3 -1
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/agent.py +171 -91
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/context.py +17 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/flow.py +30 -4
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/registry.py +14 -1
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/tool.py +90 -2
- {agstack-2.0.0 → agstack-2.1.0}/agstack.egg-info/PKG-INFO +1 -1
- {agstack-2.0.0 → agstack-2.1.0}/agstack.egg-info/SOURCES.txt +4 -1
- {agstack-2.0.0 → agstack-2.1.0}/pyproject.toml +1 -1
- agstack-2.1.0/tests/test_agent_parallel_tools.py +152 -0
- agstack-2.1.0/tests/test_flow_cancellation.py +185 -0
- agstack-2.1.0/tests/test_tool_hooks.py +218 -0
- {agstack-2.0.0 → agstack-2.1.0}/LICENSE +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/README.md +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/cache/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/cache/base.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/cache/memory.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/cache/redis.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/config/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/config/logger.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/config/manager.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/config/types.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/contexts.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/decorators.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/events.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/exceptions.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/fastapi/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/fastapi/exception.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/fastapi/middleware.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/fastapi/offline.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/fastapi/sse.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/infra/db/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/infra/es/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/infra/kg/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/infra/mq/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/client.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/event.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/exceptions.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/factory.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/loader.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/agent_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/base.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/detect_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/echo_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/iterator_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/llm_chat_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/llm_embed_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/llm_rerank_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/python_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/subflow_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/switch_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/tool_node.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/records.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/sandbox.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/state.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/trace.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/prompts.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/token.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/messagebus/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/messagebus/base.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/messagebus/memory.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/messagebus/redis.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/schema.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/security/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/security/casbin.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/security/crypt.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack/status.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack.egg-info/dependency_links.txt +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack.egg-info/requires.txt +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/agstack.egg-info/top_level.txt +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/setup.cfg +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/tests/test_cache_memory.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/tests/test_cache_redis.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/tests/test_flow_error_semantics.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/tests/test_flow_io.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/tests/test_flow_iterator.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/tests/test_flow_switch_subflow.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/tests/test_llm_usage_callback.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/tests/test_messagebus_memory.py +0 -0
- {agstack-2.0.0 → agstack-2.1.0}/tests/test_messagebus_redis.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: agstack
|
|
3
|
-
Version: 2.
|
|
3
|
+
Version: 2.1.0
|
|
4
4
|
Summary: Production-ready toolkit for building FastAPI and LLM applications
|
|
5
5
|
Author-email: XtraVisions <gitadmin@xtravisions.com>, Chen Hao <chenhao@xtravisions.com>
|
|
6
6
|
Maintainer-email: XtraVisions <gitadmin@xtravisions.com>, Chen Hao <chenhao@xtravisions.com>
|
|
@@ -22,7 +22,7 @@ from .nodes import NodeHandler
|
|
|
22
22
|
from .records import Record, Status
|
|
23
23
|
from .registry import registry
|
|
24
24
|
from .state import FlowState
|
|
25
|
-
from .tool import Tool, ToolResult
|
|
25
|
+
from .tool import Deny, Tool, ToolHook, ToolResult
|
|
26
26
|
from .trace import EdgeTrace, FlowTrace, NodeTrace
|
|
27
27
|
|
|
28
28
|
|
|
@@ -32,6 +32,8 @@ __all__ = [
|
|
|
32
32
|
# 核心抽象
|
|
33
33
|
"Tool",
|
|
34
34
|
"ToolResult",
|
|
35
|
+
"ToolHook",
|
|
36
|
+
"Deny",
|
|
35
37
|
"Agent",
|
|
36
38
|
"Flow",
|
|
37
39
|
"FlowContext",
|
|
@@ -2,6 +2,7 @@
|
|
|
2
2
|
|
|
3
3
|
"""Agent 定义和执行"""
|
|
4
4
|
|
|
5
|
+
import asyncio
|
|
5
6
|
import json
|
|
6
7
|
from typing import TYPE_CHECKING, Any, AsyncIterator
|
|
7
8
|
from uuid import uuid4
|
|
@@ -76,6 +77,149 @@ class Agent:
|
|
|
76
77
|
return tool
|
|
77
78
|
return None
|
|
78
79
|
|
|
80
|
+
def _group_tool_calls(self, tool_calls: list[dict[str, Any]]) -> list[list[dict[str, Any]]]:
|
|
81
|
+
"""按声明分组:连续的 concurrency_safe 调用聚为一组并发执行,其余单独成组串行
|
|
82
|
+
|
|
83
|
+
未注册/未声明的工具一律按不安全处理(fail closed),
|
|
84
|
+
默认全 False 时每组恰好一个调用,行为与串行完全一致。
|
|
85
|
+
"""
|
|
86
|
+
groups: list[list[dict[str, Any]]] = []
|
|
87
|
+
prev_safe = False
|
|
88
|
+
for tc in tool_calls:
|
|
89
|
+
tool = self.get_tool_by_name(tc["name"])
|
|
90
|
+
safe = bool(tool and tool.concurrency_safe)
|
|
91
|
+
if safe and prev_safe:
|
|
92
|
+
groups[-1].append(tc)
|
|
93
|
+
else:
|
|
94
|
+
groups.append([tc])
|
|
95
|
+
prev_safe = safe
|
|
96
|
+
return groups
|
|
97
|
+
|
|
98
|
+
async def _stream_tool_call(
|
|
99
|
+
self,
|
|
100
|
+
context: "FlowContext",
|
|
101
|
+
tool_call: dict[str, Any],
|
|
102
|
+
message_sink: list[dict[str, Any]] | None = None,
|
|
103
|
+
) -> AsyncIterator[dict[str, Any]]:
|
|
104
|
+
"""执行单个 tool_call,yield 其 AG-UI 事件
|
|
105
|
+
|
|
106
|
+
tool 消息默认即时写回 context;并发分组执行时传入 message_sink 收集,
|
|
107
|
+
由调用方在组完成后按 tool_call 原始顺序统一写回
|
|
108
|
+
(OpenAI 协议要求 tool 消息与 assistant.tool_calls 顺序对应)。
|
|
109
|
+
"""
|
|
110
|
+
|
|
111
|
+
def _emit_message(**kwargs: Any) -> None:
|
|
112
|
+
if message_sink is None:
|
|
113
|
+
context.add_message(self.name, "tool", **kwargs)
|
|
114
|
+
else:
|
|
115
|
+
message_sink.append(kwargs)
|
|
116
|
+
|
|
117
|
+
tool = self.get_tool_by_name(tool_call["name"])
|
|
118
|
+
if not tool:
|
|
119
|
+
error_content = json.dumps({"error": f"Tool not found: {tool_call['name']}"}, ensure_ascii=False)
|
|
120
|
+
_emit_message(content=error_content, tool_call_id=tool_call["id"])
|
|
121
|
+
# AG-UI: TOOL_CALL_RESULT (错误)
|
|
122
|
+
yield event.tool_call_result(tool_call_id=tool_call["id"], content=error_content)
|
|
123
|
+
return
|
|
124
|
+
|
|
125
|
+
# 解析 LLM 返回的工具参数;解析失败作为该次调用的失败反馈给模型,由模型自行重试
|
|
126
|
+
try:
|
|
127
|
+
tool_args = json.loads(tool_call["arguments"]) if tool_call["arguments"] else {}
|
|
128
|
+
except json.JSONDecodeError as e:
|
|
129
|
+
error_content = json.dumps(
|
|
130
|
+
{
|
|
131
|
+
"error": f"Invalid tool arguments (JSON parse failed): {e}",
|
|
132
|
+
"raw_arguments": tool_call["arguments"][:500],
|
|
133
|
+
},
|
|
134
|
+
ensure_ascii=False,
|
|
135
|
+
)
|
|
136
|
+
_emit_message(content=error_content, tool_call_id=tool_call["id"])
|
|
137
|
+
# AG-UI: TOOL_CALL_RESULT (错误)
|
|
138
|
+
yield event.tool_call_result(tool_call_id=tool_call["id"], content=error_content)
|
|
139
|
+
return
|
|
140
|
+
|
|
141
|
+
# 执行前进度事件
|
|
142
|
+
progress_label = tool.get_progress_label(tool_args)
|
|
143
|
+
if progress_label:
|
|
144
|
+
yield event.custom(
|
|
145
|
+
name="skill_progress",
|
|
146
|
+
value={
|
|
147
|
+
"progressId": tool_call["id"],
|
|
148
|
+
"description": progress_label,
|
|
149
|
+
"status": "running",
|
|
150
|
+
},
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
# 执行工具(传入 LLM 解析的参数作为 inputs)
|
|
154
|
+
result = await tool.execute_async(context, tool_args)
|
|
155
|
+
|
|
156
|
+
# 执行后进度事件
|
|
157
|
+
if progress_label:
|
|
158
|
+
yield event.custom(
|
|
159
|
+
name="skill_progress",
|
|
160
|
+
value={
|
|
161
|
+
"progressId": tool_call["id"],
|
|
162
|
+
"status": "completed" if result.success else "failed",
|
|
163
|
+
},
|
|
164
|
+
)
|
|
165
|
+
|
|
166
|
+
# 使用 result.content 作为 LLM 上下文(Tool 已计算好)
|
|
167
|
+
result_content = result.content or (
|
|
168
|
+
json.dumps(result.result, ensure_ascii=False)
|
|
169
|
+
if result.success
|
|
170
|
+
else json.dumps({"error": result.error}, ensure_ascii=False)
|
|
171
|
+
)
|
|
172
|
+
_emit_message(content=result_content, tool_call_id=tool_call["id"], summary=result.summary)
|
|
173
|
+
|
|
174
|
+
# AG-UI: TOOL_CALL_RESULT
|
|
175
|
+
yield event.tool_call_result(tool_call_id=tool_call["id"], content=result_content)
|
|
176
|
+
|
|
177
|
+
# 实时用户进度 — 有 summary 时告知前端
|
|
178
|
+
if result.summary:
|
|
179
|
+
yield event.custom(
|
|
180
|
+
name="tool_progress",
|
|
181
|
+
value={
|
|
182
|
+
"tool_call_id": tool_call["id"],
|
|
183
|
+
"tool_name": result.name,
|
|
184
|
+
"success": result.success,
|
|
185
|
+
"summary": result.summary,
|
|
186
|
+
},
|
|
187
|
+
)
|
|
188
|
+
|
|
189
|
+
# 业务自定义事件 — flush pending
|
|
190
|
+
for pending_evt in context.pop_pending_custom_events():
|
|
191
|
+
yield pending_evt
|
|
192
|
+
|
|
193
|
+
async def _gather_tool_calls(
|
|
194
|
+
self, context: "FlowContext", group: list[dict[str, Any]]
|
|
195
|
+
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
|
196
|
+
"""并发执行一组 concurrency_safe 的 tool_calls
|
|
197
|
+
|
|
198
|
+
返回 (按完成顺序展平的事件列表, 按原始顺序展平的 tool 消息 kwargs 列表)。
|
|
199
|
+
组内事件先缓冲、gather 结束后统一交调用方 yield,保证单 generator 语义;
|
|
200
|
+
单个调用的意外异常兜底转失败结果,不影响组内其它调用。
|
|
201
|
+
注意:并发期间 execution_records / pending_custom_events 的追加顺序不再确定。
|
|
202
|
+
"""
|
|
203
|
+
event_buffers: list[list[dict[str, Any]]] = [[] for _ in group]
|
|
204
|
+
message_buffers: list[list[dict[str, Any]]] = [[] for _ in group]
|
|
205
|
+
finish_order: list[int] = []
|
|
206
|
+
|
|
207
|
+
async def _run_one(i: int, tc: dict[str, Any]) -> None:
|
|
208
|
+
try:
|
|
209
|
+
async for evt in self._stream_tool_call(context, tc, message_sink=message_buffers[i]):
|
|
210
|
+
event_buffers[i].append(evt)
|
|
211
|
+
except Exception as e:
|
|
212
|
+
error_content = json.dumps({"error": str(e)}, ensure_ascii=False)
|
|
213
|
+
message_buffers[i].append({"content": error_content, "tool_call_id": tc["id"]})
|
|
214
|
+
event_buffers[i].append(event.tool_call_result(tool_call_id=tc["id"], content=error_content))
|
|
215
|
+
finally:
|
|
216
|
+
finish_order.append(i)
|
|
217
|
+
|
|
218
|
+
await asyncio.gather(*[_run_one(i, tc) for i, tc in enumerate(group)])
|
|
219
|
+
events = [evt for i in finish_order for evt in event_buffers[i]]
|
|
220
|
+
messages = [msg for msgs in message_buffers for msg in msgs]
|
|
221
|
+
return events, messages
|
|
222
|
+
|
|
79
223
|
async def run(self, context: "FlowContext", inputs: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
80
224
|
"""执行 Agent 逻辑"""
|
|
81
225
|
content_parts = []
|
|
@@ -122,6 +266,13 @@ class Agent:
|
|
|
122
266
|
# Agent 循环
|
|
123
267
|
assistant_content = ""
|
|
124
268
|
for _ in range(self.max_turns):
|
|
269
|
+
# 协作式取消检查点:不再开始新的 LLM 轮次
|
|
270
|
+
if context.is_cancelled:
|
|
271
|
+
if not context.get_variable("_cancel_emitted"):
|
|
272
|
+
context.set_variable("_cancel_emitted", True)
|
|
273
|
+
yield event.run_error(message="FLOW_CANCELLED", code="CANCELLED")
|
|
274
|
+
return
|
|
275
|
+
|
|
125
276
|
context.increment_turn()
|
|
126
277
|
|
|
127
278
|
# 调用模型
|
|
@@ -243,97 +394,26 @@ class Agent:
|
|
|
243
394
|
context.set_variable("_agent_call_id", None)
|
|
244
395
|
return
|
|
245
396
|
|
|
246
|
-
#
|
|
247
|
-
for
|
|
248
|
-
|
|
249
|
-
if
|
|
250
|
-
|
|
251
|
-
|
|
252
|
-
|
|
253
|
-
|
|
254
|
-
|
|
255
|
-
|
|
256
|
-
|
|
257
|
-
|
|
258
|
-
|
|
259
|
-
|
|
260
|
-
|
|
261
|
-
)
|
|
262
|
-
|
|
263
|
-
|
|
264
|
-
|
|
265
|
-
|
|
266
|
-
tool_args = json.loads(tool_call["arguments"]) if tool_call["arguments"] else {}
|
|
267
|
-
except json.JSONDecodeError as e:
|
|
268
|
-
error_content = json.dumps(
|
|
269
|
-
{
|
|
270
|
-
"error": f"Invalid tool arguments (JSON parse failed): {e}",
|
|
271
|
-
"raw_arguments": tool_call["arguments"][:500],
|
|
272
|
-
},
|
|
273
|
-
ensure_ascii=False,
|
|
274
|
-
)
|
|
275
|
-
context.add_message(self.name, "tool", content=error_content, tool_call_id=tool_call["id"])
|
|
276
|
-
# AG-UI: TOOL_CALL_RESULT (错误)
|
|
277
|
-
yield event.tool_call_result(tool_call_id=tool_call["id"], content=error_content)
|
|
278
|
-
continue
|
|
279
|
-
|
|
280
|
-
# 执行前进度事件
|
|
281
|
-
progress_label = tool.get_progress_label(tool_args)
|
|
282
|
-
if progress_label:
|
|
283
|
-
yield event.custom(
|
|
284
|
-
name="skill_progress",
|
|
285
|
-
value={
|
|
286
|
-
"progressId": tool_call["id"],
|
|
287
|
-
"description": progress_label,
|
|
288
|
-
"status": "running",
|
|
289
|
-
},
|
|
290
|
-
)
|
|
291
|
-
|
|
292
|
-
# 执行工具(传入 LLM 解析的参数作为 inputs)
|
|
293
|
-
result = await tool.execute_async(context, tool_args)
|
|
294
|
-
|
|
295
|
-
# 执行后进度事件
|
|
296
|
-
if progress_label:
|
|
297
|
-
yield event.custom(
|
|
298
|
-
name="skill_progress",
|
|
299
|
-
value={
|
|
300
|
-
"progressId": tool_call["id"],
|
|
301
|
-
"status": "completed" if result.success else "failed",
|
|
302
|
-
},
|
|
303
|
-
)
|
|
304
|
-
|
|
305
|
-
# 使用 result.content 作为 LLM 上下文(Tool 已计算好)
|
|
306
|
-
result_content = result.content or (
|
|
307
|
-
json.dumps(result.result, ensure_ascii=False)
|
|
308
|
-
if result.success
|
|
309
|
-
else json.dumps({"error": result.error}, ensure_ascii=False)
|
|
310
|
-
)
|
|
311
|
-
context.add_message(
|
|
312
|
-
self.name,
|
|
313
|
-
"tool",
|
|
314
|
-
content=result_content,
|
|
315
|
-
tool_call_id=tool_call["id"],
|
|
316
|
-
summary=result.summary,
|
|
317
|
-
)
|
|
318
|
-
|
|
319
|
-
# AG-UI: TOOL_CALL_RESULT
|
|
320
|
-
yield event.tool_call_result(tool_call_id=tool_call["id"], content=result_content)
|
|
321
|
-
|
|
322
|
-
# 实时用户进度 — 有 summary 时告知前端
|
|
323
|
-
if result.summary:
|
|
324
|
-
yield event.custom(
|
|
325
|
-
name="tool_progress",
|
|
326
|
-
value={
|
|
327
|
-
"tool_call_id": tool_call["id"],
|
|
328
|
-
"tool_name": result.name,
|
|
329
|
-
"success": result.success,
|
|
330
|
-
"summary": result.summary,
|
|
331
|
-
},
|
|
332
|
-
)
|
|
333
|
-
|
|
334
|
-
# 业务自定义事件 — flush pending
|
|
335
|
-
for pending_evt in context.pop_pending_custom_events():
|
|
336
|
-
yield pending_evt
|
|
397
|
+
# 执行工具调用:连续的 concurrency_safe 调用聚组并发,其余保持串行(默认全串行)
|
|
398
|
+
for group in self._group_tool_calls(tool_calls):
|
|
399
|
+
# 协作式取消检查点:不再开始新的工具执行(不中断在途工具)
|
|
400
|
+
if context.is_cancelled:
|
|
401
|
+
if not context.get_variable("_cancel_emitted"):
|
|
402
|
+
context.set_variable("_cancel_emitted", True)
|
|
403
|
+
yield event.run_error(message="FLOW_CANCELLED", code="CANCELLED")
|
|
404
|
+
return
|
|
405
|
+
|
|
406
|
+
if len(group) == 1:
|
|
407
|
+
# 串行路径:事件实时 yield,tool 消息即时写回
|
|
408
|
+
async for evt in self._stream_tool_call(context, group[0]):
|
|
409
|
+
yield evt
|
|
410
|
+
else:
|
|
411
|
+
# 并发组:事件按完成顺序 yield,tool 消息按原始顺序写回
|
|
412
|
+
group_events, group_messages = await self._gather_tool_calls(context, group)
|
|
413
|
+
for evt in group_events:
|
|
414
|
+
yield evt
|
|
415
|
+
for msg in group_messages:
|
|
416
|
+
context.add_message(self.name, "tool", **msg)
|
|
337
417
|
|
|
338
418
|
# 更新消息列表,继续下一轮
|
|
339
419
|
messages = [self.get_system_message()] + context.history + context.get_messages(self.name)
|
|
@@ -2,6 +2,7 @@
|
|
|
2
2
|
|
|
3
3
|
"""统一的执行上下文"""
|
|
4
4
|
|
|
5
|
+
import asyncio
|
|
5
6
|
import uuid
|
|
6
7
|
from dataclasses import dataclass, field
|
|
7
8
|
from datetime import datetime
|
|
@@ -66,6 +67,9 @@ class FlowContext:
|
|
|
66
67
|
# 业务自定义事件缓冲(tool 内部通过 context 注入,agent 循环中 flush)
|
|
67
68
|
pending_custom_events: list[dict[str, Any]] = field(default_factory=list)
|
|
68
69
|
|
|
70
|
+
# 协作式取消信号(不参与序列化,恢复即未取消)
|
|
71
|
+
_cancel_event: asyncio.Event = field(default_factory=asyncio.Event, repr=False)
|
|
72
|
+
|
|
69
73
|
def __post_init__(self) -> None:
|
|
70
74
|
if self.trace is None:
|
|
71
75
|
from .trace import FlowTrace
|
|
@@ -84,6 +88,19 @@ class FlowContext:
|
|
|
84
88
|
"""取出并移除变量"""
|
|
85
89
|
return self.variables.pop(key, default)
|
|
86
90
|
|
|
91
|
+
def cancel(self) -> None:
|
|
92
|
+
"""请求取消执行(幂等)
|
|
93
|
+
|
|
94
|
+
协作式取消:引擎在下一个检查点(节点执行前、agent 轮次开始、
|
|
95
|
+
tool_call 执行前)停止,不强杀在途的工具或 LLM 调用。
|
|
96
|
+
"""
|
|
97
|
+
self._cancel_event.set()
|
|
98
|
+
|
|
99
|
+
@property
|
|
100
|
+
def is_cancelled(self) -> bool:
|
|
101
|
+
"""是否已请求取消(长 I/O 工具可自查以提前返回)"""
|
|
102
|
+
return self._cancel_event.is_set()
|
|
103
|
+
|
|
87
104
|
def update_variables(self, updates: dict[str, Any]) -> None:
|
|
88
105
|
"""批量更新变量"""
|
|
89
106
|
self.variables.update(updates)
|
|
@@ -10,7 +10,7 @@ from uuid import uuid4
|
|
|
10
10
|
|
|
11
11
|
from . import event
|
|
12
12
|
from .context import Usage
|
|
13
|
-
from .exceptions import NodeExecutionError
|
|
13
|
+
from .exceptions import FlowExecutionError, NodeExecutionError
|
|
14
14
|
|
|
15
15
|
|
|
16
16
|
if TYPE_CHECKING:
|
|
@@ -234,6 +234,12 @@ class Flow:
|
|
|
234
234
|
last_error: Exception | None = None
|
|
235
235
|
|
|
236
236
|
for attempt in range(policy.max_retries + 1):
|
|
237
|
+
# 协作式取消检查点:取消后不再开始新的重试
|
|
238
|
+
if attempt > 0 and context.is_cancelled:
|
|
239
|
+
if not context.get_variable("_cancel_emitted"):
|
|
240
|
+
context.set_variable("_cancel_emitted", True)
|
|
241
|
+
yield event.run_error(message="FLOW_CANCELLED", code="CANCELLED")
|
|
242
|
+
return
|
|
237
243
|
try:
|
|
238
244
|
if attempt > 0:
|
|
239
245
|
wait = policy.delay * (policy.backoff ** (attempt - 1))
|
|
@@ -284,10 +290,12 @@ class Flow:
|
|
|
284
290
|
|
|
285
291
|
stream() 的消费包装:两条路径共享同一执行引擎,重试策略、
|
|
286
292
|
FlowTrace、output_mode、iterator 状态清理等行为完全一致。
|
|
287
|
-
节点失败抛 NodeExecutionError
|
|
293
|
+
节点失败抛 NodeExecutionError(包装原始异常);
|
|
294
|
+
取消(context.cancel())抛 FlowExecutionError("FLOW_CANCELLED")。
|
|
288
295
|
"""
|
|
289
|
-
async for
|
|
290
|
-
|
|
296
|
+
async for evt in self.stream(context):
|
|
297
|
+
if evt.get("type") == event.EventType.RUN_ERROR and evt.get("code") == "CANCELLED":
|
|
298
|
+
raise FlowExecutionError("FLOW_CANCELLED", args={"flow": self.name})
|
|
291
299
|
return context.outputs
|
|
292
300
|
|
|
293
301
|
async def stream(self, context: "FlowContext") -> AsyncIterator[dict[str, Any]]:
|
|
@@ -310,11 +318,22 @@ class Flow:
|
|
|
310
318
|
context.trace.finished_at = time.time()
|
|
311
319
|
context.trace.total_usage = context.usage
|
|
312
320
|
|
|
321
|
+
# 取消停止:RUN_ERROR(code=CANCELLED) 是终止事件,事件流以其结束
|
|
322
|
+
if context.get_variable("_cancel_emitted"):
|
|
323
|
+
return
|
|
324
|
+
|
|
313
325
|
yield event.step_finished(step_name=f"flow:{self.name}", step_id=flow_sid)
|
|
314
326
|
|
|
315
327
|
async def _stream_sequential(self, context: "FlowContext") -> AsyncIterator[dict[str, Any]]:
|
|
316
328
|
"""顺序流式执行"""
|
|
317
329
|
for node in self.nodes:
|
|
330
|
+
# 协作式取消检查点:不再开始新的节点执行
|
|
331
|
+
if context.is_cancelled:
|
|
332
|
+
if not context.get_variable("_cancel_emitted"):
|
|
333
|
+
context.set_variable("_cancel_emitted", True)
|
|
334
|
+
yield event.run_error(message="FLOW_CANCELLED", code="CANCELLED")
|
|
335
|
+
return
|
|
336
|
+
|
|
318
337
|
node_id = node.get("id")
|
|
319
338
|
if not node_id:
|
|
320
339
|
continue
|
|
@@ -350,6 +369,13 @@ class Flow:
|
|
|
350
369
|
visit_count: dict[str, int] = {}
|
|
351
370
|
|
|
352
371
|
while current_node_id:
|
|
372
|
+
# 协作式取消检查点:不再开始新的节点执行(粒度是节点边界,不中断在途节点)
|
|
373
|
+
if context.is_cancelled:
|
|
374
|
+
if not context.get_variable("_cancel_emitted"):
|
|
375
|
+
context.set_variable("_cancel_emitted", True)
|
|
376
|
+
yield event.run_error(message="FLOW_CANCELLED", code="CANCELLED")
|
|
377
|
+
return
|
|
378
|
+
|
|
353
379
|
node = self.get_node_config(current_node_id)
|
|
354
380
|
if not node:
|
|
355
381
|
yield event.run_error(
|
|
@@ -8,7 +8,7 @@ import copy
|
|
|
8
8
|
from typing import Any, cast
|
|
9
9
|
|
|
10
10
|
from .agent import Agent
|
|
11
|
-
from .tool import Tool
|
|
11
|
+
from .tool import Tool, ToolHook, clear_tool_hooks, register_tool_hook
|
|
12
12
|
|
|
13
13
|
|
|
14
14
|
class FlowRegistry:
|
|
@@ -45,6 +45,19 @@ class FlowRegistry:
|
|
|
45
45
|
if echo:
|
|
46
46
|
self._tool_echo[name] = True
|
|
47
47
|
|
|
48
|
+
def register_tool_hook(self, hook: ToolHook, *, prepend: bool = False) -> None:
|
|
49
|
+
"""注册全局工具执行钩子,对所有 Tool.execute_async 生效
|
|
50
|
+
|
|
51
|
+
pre_execute 按注册顺序、post_execute 按逆序执行;
|
|
52
|
+
prepend=True 抢占链头(其 post_execute 成为最外层,适合 spill/截断类钩子)。
|
|
53
|
+
钩子链存于 tool 模块(保持 registry→tool 单向导入),此处仅转发注册。
|
|
54
|
+
"""
|
|
55
|
+
register_tool_hook(hook, prepend=prepend)
|
|
56
|
+
|
|
57
|
+
def clear_tool_hooks(self) -> None:
|
|
58
|
+
"""清空全局工具钩子(测试隔离用)"""
|
|
59
|
+
clear_tool_hooks()
|
|
60
|
+
|
|
48
61
|
def register_agent(
|
|
49
62
|
self, name: str, agent_class: type[Agent], *, label: str | None = None, echo: bool = False
|
|
50
63
|
) -> None:
|
|
@@ -28,6 +28,62 @@ class ToolResult:
|
|
|
28
28
|
summary: str | None = None
|
|
29
29
|
|
|
30
30
|
|
|
31
|
+
class Deny:
|
|
32
|
+
"""pre_execute 的拒绝决策:工具本体不执行,reason 作为失败结果反馈给模型"""
|
|
33
|
+
|
|
34
|
+
def __init__(self, reason: str):
|
|
35
|
+
self.reason = reason
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class ToolHook:
|
|
39
|
+
"""工具执行钩子基类(全局链,经 registry.register_tool_hook 注册)
|
|
40
|
+
|
|
41
|
+
子类覆写其一或两者;方法必须是 async def。
|
|
42
|
+
pre_execute 按注册顺序、post_execute 按逆序执行(洋葱模型)。
|
|
43
|
+
F3 并行工具落地后钩子会被并发调用:实现必须无状态或自行同步
|
|
44
|
+
(与 Tool 单例的既有纪律同构)。
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
async def pre_execute(
|
|
48
|
+
self, context: "FlowContext", tool: "Tool", inputs: dict[str, Any]
|
|
49
|
+
) -> "dict[str, Any] | Deny":
|
|
50
|
+
"""工具入参进入工具函数前调用
|
|
51
|
+
|
|
52
|
+
返回(可改写的)inputs 继续执行;返回 Deny 拒绝执行;
|
|
53
|
+
抛异常按 Deny 处理(fail closed:权限门自己出错时不放行)。
|
|
54
|
+
"""
|
|
55
|
+
return inputs
|
|
56
|
+
|
|
57
|
+
async def post_execute(self, context: "FlowContext", tool: "Tool", result: "ToolResult") -> "ToolResult":
|
|
58
|
+
"""结果落入上下文前调用,原样返回=纯观察
|
|
59
|
+
|
|
60
|
+
可替换/截断/落盘换 locator;改写对 LLM 消费内容、用户摘要、
|
|
61
|
+
execution_records 三个出口同时生效。抛异常记日志并放行原结果
|
|
62
|
+
(fail open:审计钩子的 bug 不毁掉主流程)。
|
|
63
|
+
Deny 产生的失败结果同样穿过 post 链(审计要看到被拒绝的调用)。
|
|
64
|
+
"""
|
|
65
|
+
return result
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
_TOOL_HOOKS: list[ToolHook] = []
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def register_tool_hook(hook: ToolHook, *, prepend: bool = False) -> None:
|
|
72
|
+
"""注册全局工具执行钩子,对所有 Tool.execute_async 生效
|
|
73
|
+
|
|
74
|
+
prepend=True 抢占链头(其 post_execute 成为最外层,适合 spill/截断类钩子)。
|
|
75
|
+
"""
|
|
76
|
+
if prepend:
|
|
77
|
+
_TOOL_HOOKS.insert(0, hook)
|
|
78
|
+
else:
|
|
79
|
+
_TOOL_HOOKS.append(hook)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def clear_tool_hooks() -> None:
|
|
83
|
+
"""清空全局工具钩子(测试隔离用)"""
|
|
84
|
+
_TOOL_HOOKS.clear()
|
|
85
|
+
|
|
86
|
+
|
|
31
87
|
class Tool:
|
|
32
88
|
"""工具定义"""
|
|
33
89
|
|
|
@@ -41,6 +97,7 @@ class Tool:
|
|
|
41
97
|
label: str | None = None,
|
|
42
98
|
echo: bool = False,
|
|
43
99
|
category: str | None = None,
|
|
100
|
+
concurrency_safe: bool = False,
|
|
44
101
|
summary_fn: Callable[["ToolResult"], str | None] | None = None,
|
|
45
102
|
result_formatter: Callable[["ToolResult"], str] | None = None,
|
|
46
103
|
progress_label_fn: Callable[[dict[str, Any]], str] | None = None,
|
|
@@ -54,6 +111,7 @@ class Tool:
|
|
|
54
111
|
:param label: 面向用户的展示名称(控制 STEP/TOOL_CALL 进度事件可见性)
|
|
55
112
|
:param echo: 是否转发 TEXT_MESSAGE 给用户
|
|
56
113
|
:param category: 工具分类(retrieval / analysis / action / utility)
|
|
114
|
+
:param concurrency_safe: 声明该工具可与其它 concurrency_safe 工具并发执行(默认串行)
|
|
57
115
|
:param summary_fn: 生成面向用户摘要的函数 (ToolResult) -> str | None
|
|
58
116
|
:param result_formatter: 自定义 LLM 内容格式化函数 (ToolResult) -> str
|
|
59
117
|
:param progress_label_fn: 基于调用参数生成动态进度描述 (args) -> str
|
|
@@ -65,6 +123,7 @@ class Tool:
|
|
|
65
123
|
self.label = label
|
|
66
124
|
self.echo = echo
|
|
67
125
|
self.category = category
|
|
126
|
+
self.concurrency_safe = concurrency_safe
|
|
68
127
|
self.summary_fn = summary_fn
|
|
69
128
|
self.result_formatter = result_formatter
|
|
70
129
|
self.progress_label_fn = progress_label_fn
|
|
@@ -82,10 +141,39 @@ class Tool:
|
|
|
82
141
|
return self.label
|
|
83
142
|
|
|
84
143
|
async def execute_async(self, context: "FlowContext", inputs: dict[str, Any] | None = None) -> ToolResult:
|
|
85
|
-
"""
|
|
144
|
+
"""异步执行工具(包含钩子链、计时、摘要生成、结果格式化、可观测性记录)"""
|
|
86
145
|
args = inputs or {}
|
|
87
146
|
_t0 = time.perf_counter()
|
|
88
|
-
|
|
147
|
+
|
|
148
|
+
# pre 钩子链(注册顺序):可改写入参;返回 Deny 或抛异常=拒绝执行(fail closed)
|
|
149
|
+
result: ToolResult | None = None
|
|
150
|
+
for hook in _TOOL_HOOKS:
|
|
151
|
+
try:
|
|
152
|
+
outcome = await hook.pre_execute(context, self, args)
|
|
153
|
+
except Exception as e:
|
|
154
|
+
logger.warning("Tool hook pre_execute failed for %s: %s", self.name, e, exc_info=True)
|
|
155
|
+
outcome = Deny(f"tool hook error: {e}")
|
|
156
|
+
if isinstance(outcome, Deny):
|
|
157
|
+
result = ToolResult(name=self.name, arguments=args, result={}, success=False, error=outcome.reason)
|
|
158
|
+
break
|
|
159
|
+
args = outcome
|
|
160
|
+
|
|
161
|
+
if result is None:
|
|
162
|
+
result = await self._execute(context, args)
|
|
163
|
+
|
|
164
|
+
# post 钩子链(逆序):可改写结果;抛异常=放行原结果(fail open)。
|
|
165
|
+
# Deny 的失败结果同样穿过 post 链,审计钩子能看到被拒绝的调用。
|
|
166
|
+
for hook in reversed(_TOOL_HOOKS):
|
|
167
|
+
try:
|
|
168
|
+
revised = await hook.post_execute(context, self, result)
|
|
169
|
+
except Exception as e:
|
|
170
|
+
logger.warning("Tool hook post_execute failed for %s: %s", self.name, e, exc_info=True)
|
|
171
|
+
continue
|
|
172
|
+
if isinstance(revised, ToolResult):
|
|
173
|
+
result = revised
|
|
174
|
+
else:
|
|
175
|
+
logger.warning("Tool hook post_execute for %s returned %r, ignored", self.name, type(revised))
|
|
176
|
+
|
|
89
177
|
_duration_ms = int((time.perf_counter() - _t0) * 1000)
|
|
90
178
|
|
|
91
179
|
# 计算 LLM 消费内容
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: agstack
|
|
3
|
-
Version: 2.
|
|
3
|
+
Version: 2.1.0
|
|
4
4
|
Summary: Production-ready toolkit for building FastAPI and LLM applications
|
|
5
5
|
Author-email: XtraVisions <gitadmin@xtravisions.com>, Chen Hao <chenhao@xtravisions.com>
|
|
6
6
|
Maintainer-email: XtraVisions <gitadmin@xtravisions.com>, Chen Hao <chenhao@xtravisions.com>
|
|
@@ -68,12 +68,15 @@ agstack/messagebus/redis.py
|
|
|
68
68
|
agstack/security/__init__.py
|
|
69
69
|
agstack/security/casbin.py
|
|
70
70
|
agstack/security/crypt.py
|
|
71
|
+
tests/test_agent_parallel_tools.py
|
|
71
72
|
tests/test_cache_memory.py
|
|
72
73
|
tests/test_cache_redis.py
|
|
74
|
+
tests/test_flow_cancellation.py
|
|
73
75
|
tests/test_flow_error_semantics.py
|
|
74
76
|
tests/test_flow_io.py
|
|
75
77
|
tests/test_flow_iterator.py
|
|
76
78
|
tests/test_flow_switch_subflow.py
|
|
77
79
|
tests/test_llm_usage_callback.py
|
|
78
80
|
tests/test_messagebus_memory.py
|
|
79
|
-
tests/test_messagebus_redis.py
|
|
81
|
+
tests/test_messagebus_redis.py
|
|
82
|
+
tests/test_tool_hooks.py
|