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.
Files changed (84) hide show
  1. {agstack-2.0.0 → agstack-2.1.0}/PKG-INFO +1 -1
  2. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/__init__.py +3 -1
  3. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/agent.py +171 -91
  4. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/context.py +17 -0
  5. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/flow.py +30 -4
  6. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/registry.py +14 -1
  7. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/tool.py +90 -2
  8. {agstack-2.0.0 → agstack-2.1.0}/agstack.egg-info/PKG-INFO +1 -1
  9. {agstack-2.0.0 → agstack-2.1.0}/agstack.egg-info/SOURCES.txt +4 -1
  10. {agstack-2.0.0 → agstack-2.1.0}/pyproject.toml +1 -1
  11. agstack-2.1.0/tests/test_agent_parallel_tools.py +152 -0
  12. agstack-2.1.0/tests/test_flow_cancellation.py +185 -0
  13. agstack-2.1.0/tests/test_tool_hooks.py +218 -0
  14. {agstack-2.0.0 → agstack-2.1.0}/LICENSE +0 -0
  15. {agstack-2.0.0 → agstack-2.1.0}/README.md +0 -0
  16. {agstack-2.0.0 → agstack-2.1.0}/agstack/__init__.py +0 -0
  17. {agstack-2.0.0 → agstack-2.1.0}/agstack/cache/__init__.py +0 -0
  18. {agstack-2.0.0 → agstack-2.1.0}/agstack/cache/base.py +0 -0
  19. {agstack-2.0.0 → agstack-2.1.0}/agstack/cache/memory.py +0 -0
  20. {agstack-2.0.0 → agstack-2.1.0}/agstack/cache/redis.py +0 -0
  21. {agstack-2.0.0 → agstack-2.1.0}/agstack/config/__init__.py +0 -0
  22. {agstack-2.0.0 → agstack-2.1.0}/agstack/config/logger.py +0 -0
  23. {agstack-2.0.0 → agstack-2.1.0}/agstack/config/manager.py +0 -0
  24. {agstack-2.0.0 → agstack-2.1.0}/agstack/config/types.py +0 -0
  25. {agstack-2.0.0 → agstack-2.1.0}/agstack/contexts.py +0 -0
  26. {agstack-2.0.0 → agstack-2.1.0}/agstack/decorators.py +0 -0
  27. {agstack-2.0.0 → agstack-2.1.0}/agstack/events.py +0 -0
  28. {agstack-2.0.0 → agstack-2.1.0}/agstack/exceptions.py +0 -0
  29. {agstack-2.0.0 → agstack-2.1.0}/agstack/fastapi/__init__.py +0 -0
  30. {agstack-2.0.0 → agstack-2.1.0}/agstack/fastapi/exception.py +0 -0
  31. {agstack-2.0.0 → agstack-2.1.0}/agstack/fastapi/middleware.py +0 -0
  32. {agstack-2.0.0 → agstack-2.1.0}/agstack/fastapi/offline.py +0 -0
  33. {agstack-2.0.0 → agstack-2.1.0}/agstack/fastapi/sse.py +0 -0
  34. {agstack-2.0.0 → agstack-2.1.0}/agstack/infra/db/__init__.py +0 -0
  35. {agstack-2.0.0 → agstack-2.1.0}/agstack/infra/es/__init__.py +0 -0
  36. {agstack-2.0.0 → agstack-2.1.0}/agstack/infra/kg/__init__.py +0 -0
  37. {agstack-2.0.0 → agstack-2.1.0}/agstack/infra/mq/__init__.py +0 -0
  38. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/__init__.py +0 -0
  39. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/client.py +0 -0
  40. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/event.py +0 -0
  41. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/exceptions.py +0 -0
  42. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/factory.py +0 -0
  43. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/loader.py +0 -0
  44. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/__init__.py +0 -0
  45. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/agent_node.py +0 -0
  46. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/base.py +0 -0
  47. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/detect_node.py +0 -0
  48. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/echo_node.py +0 -0
  49. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/iterator_node.py +0 -0
  50. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/llm_chat_node.py +0 -0
  51. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/llm_embed_node.py +0 -0
  52. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/llm_rerank_node.py +0 -0
  53. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/python_node.py +0 -0
  54. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/subflow_node.py +0 -0
  55. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/switch_node.py +0 -0
  56. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/nodes/tool_node.py +0 -0
  57. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/records.py +0 -0
  58. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/sandbox.py +0 -0
  59. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/state.py +0 -0
  60. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/flow/trace.py +0 -0
  61. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/prompts.py +0 -0
  62. {agstack-2.0.0 → agstack-2.1.0}/agstack/llm/token.py +0 -0
  63. {agstack-2.0.0 → agstack-2.1.0}/agstack/messagebus/__init__.py +0 -0
  64. {agstack-2.0.0 → agstack-2.1.0}/agstack/messagebus/base.py +0 -0
  65. {agstack-2.0.0 → agstack-2.1.0}/agstack/messagebus/memory.py +0 -0
  66. {agstack-2.0.0 → agstack-2.1.0}/agstack/messagebus/redis.py +0 -0
  67. {agstack-2.0.0 → agstack-2.1.0}/agstack/schema.py +0 -0
  68. {agstack-2.0.0 → agstack-2.1.0}/agstack/security/__init__.py +0 -0
  69. {agstack-2.0.0 → agstack-2.1.0}/agstack/security/casbin.py +0 -0
  70. {agstack-2.0.0 → agstack-2.1.0}/agstack/security/crypt.py +0 -0
  71. {agstack-2.0.0 → agstack-2.1.0}/agstack/status.py +0 -0
  72. {agstack-2.0.0 → agstack-2.1.0}/agstack.egg-info/dependency_links.txt +0 -0
  73. {agstack-2.0.0 → agstack-2.1.0}/agstack.egg-info/requires.txt +0 -0
  74. {agstack-2.0.0 → agstack-2.1.0}/agstack.egg-info/top_level.txt +0 -0
  75. {agstack-2.0.0 → agstack-2.1.0}/setup.cfg +0 -0
  76. {agstack-2.0.0 → agstack-2.1.0}/tests/test_cache_memory.py +0 -0
  77. {agstack-2.0.0 → agstack-2.1.0}/tests/test_cache_redis.py +0 -0
  78. {agstack-2.0.0 → agstack-2.1.0}/tests/test_flow_error_semantics.py +0 -0
  79. {agstack-2.0.0 → agstack-2.1.0}/tests/test_flow_io.py +0 -0
  80. {agstack-2.0.0 → agstack-2.1.0}/tests/test_flow_iterator.py +0 -0
  81. {agstack-2.0.0 → agstack-2.1.0}/tests/test_flow_switch_subflow.py +0 -0
  82. {agstack-2.0.0 → agstack-2.1.0}/tests/test_llm_usage_callback.py +0 -0
  83. {agstack-2.0.0 → agstack-2.1.0}/tests/test_messagebus_memory.py +0 -0
  84. {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.0.0
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 tool_call in tool_calls:
248
- tool = self.get_tool_by_name(tool_call["name"])
249
- if not tool:
250
- error_msg = f"Tool not found: {tool_call['name']}"
251
- context.add_message(
252
- self.name,
253
- "tool",
254
- content=json.dumps({"error": error_msg}, ensure_ascii=False),
255
- tool_call_id=tool_call["id"],
256
- )
257
- # AG-UI: TOOL_CALL_RESULT (错误)
258
- yield event.tool_call_result(
259
- tool_call_id=tool_call["id"],
260
- content=json.dumps({"error": error_msg}, ensure_ascii=False),
261
- )
262
- continue
263
-
264
- # 解析 LLM 返回的工具参数;解析失败作为该次调用的失败反馈给模型,由模型自行重试
265
- try:
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 _ in self.stream(context):
290
- pass
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
- result = await self._execute(context, args)
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.0.0
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
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "agstack"
3
- version = "2.0.0"
3
+ version = "2.1.0"
4
4
  description = "Production-ready toolkit for building FastAPI and LLM applications"
5
5
  readme = "README.md"
6
6
  license = "MIT"