agstack 1.25.1__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-1.25.1 → agstack-2.1.0}/PKG-INFO +1 -1
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/__init__.py +3 -1
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/agent.py +191 -82
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/context.py +21 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/flow.py +78 -117
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/tool_node.py +5 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/registry.py +14 -1
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/tool.py +90 -2
- {agstack-1.25.1 → agstack-2.1.0}/agstack.egg-info/PKG-INFO +1 -1
- {agstack-1.25.1 → agstack-2.1.0}/agstack.egg-info/SOURCES.txt +5 -1
- {agstack-1.25.1 → 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_flow_error_semantics.py +501 -0
- agstack-2.1.0/tests/test_tool_hooks.py +218 -0
- {agstack-1.25.1 → agstack-2.1.0}/LICENSE +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/README.md +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/cache/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/cache/base.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/cache/memory.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/cache/redis.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/config/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/config/logger.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/config/manager.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/config/types.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/contexts.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/decorators.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/events.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/exceptions.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/fastapi/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/fastapi/exception.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/fastapi/middleware.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/fastapi/offline.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/fastapi/sse.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/infra/db/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/infra/es/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/infra/kg/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/infra/mq/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/client.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/event.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/exceptions.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/factory.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/loader.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/agent_node.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/base.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/detect_node.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/echo_node.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/iterator_node.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/llm_chat_node.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/llm_embed_node.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/llm_rerank_node.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/python_node.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/subflow_node.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/nodes/switch_node.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/records.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/sandbox.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/state.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/flow/trace.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/prompts.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/llm/token.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/messagebus/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/messagebus/base.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/messagebus/memory.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/messagebus/redis.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/schema.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/security/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/security/casbin.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/security/crypt.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack/status.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack.egg-info/dependency_links.txt +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack.egg-info/requires.txt +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/agstack.egg-info/top_level.txt +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/setup.cfg +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/tests/test_cache_memory.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/tests/test_cache_redis.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/tests/test_flow_io.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/tests/test_flow_iterator.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/tests/test_flow_switch_subflow.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/tests/test_llm_usage_callback.py +0 -0
- {agstack-1.25.1 → agstack-2.1.0}/tests/test_messagebus_memory.py +0 -0
- {agstack-1.25.1 → 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: 1.
|
|
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
|
|
@@ -10,7 +11,7 @@ from ..client import get_llm_client
|
|
|
10
11
|
from . import event
|
|
11
12
|
from .context import Usage
|
|
12
13
|
from .event import EventType
|
|
13
|
-
from .exceptions import FlowError
|
|
14
|
+
from .exceptions import AgentError, FlowError
|
|
14
15
|
|
|
15
16
|
|
|
16
17
|
if TYPE_CHECKING:
|
|
@@ -32,6 +33,7 @@ class Agent:
|
|
|
32
33
|
max_turns: int = 10,
|
|
33
34
|
*,
|
|
34
35
|
tool_choice: str = "auto",
|
|
36
|
+
on_max_turns: str = "finalize",
|
|
35
37
|
label: str | None = None,
|
|
36
38
|
echo: bool = False,
|
|
37
39
|
):
|
|
@@ -44,6 +46,7 @@ class Agent:
|
|
|
44
46
|
:param temperature: 温度参数
|
|
45
47
|
:param max_tokens: 最大 token 数
|
|
46
48
|
:param max_turns: 最大轮次
|
|
49
|
+
:param on_max_turns: max_turns 耗尽时的行为,"finalize"(降级输出并标记 truncated)或 "error"(抛出异常)
|
|
47
50
|
:param label: 面向用户的展示名称(控制 STEP 进度事件可见性)
|
|
48
51
|
:param echo: 是否转发 TEXT_MESSAGE 给用户
|
|
49
52
|
"""
|
|
@@ -55,6 +58,7 @@ class Agent:
|
|
|
55
58
|
self.max_tokens = max_tokens
|
|
56
59
|
self.max_turns = max_turns
|
|
57
60
|
self.tool_choice = tool_choice
|
|
61
|
+
self.on_max_turns = on_max_turns
|
|
58
62
|
self.label = label
|
|
59
63
|
self.echo = echo
|
|
60
64
|
|
|
@@ -73,6 +77,149 @@ class Agent:
|
|
|
73
77
|
return tool
|
|
74
78
|
return None
|
|
75
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
|
+
|
|
76
223
|
async def run(self, context: "FlowContext", inputs: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
77
224
|
"""执行 Agent 逻辑"""
|
|
78
225
|
content_parts = []
|
|
@@ -117,7 +264,15 @@ class Agent:
|
|
|
117
264
|
client = get_llm_client()
|
|
118
265
|
|
|
119
266
|
# Agent 循环
|
|
267
|
+
assistant_content = ""
|
|
120
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
|
+
|
|
121
276
|
context.increment_turn()
|
|
122
277
|
|
|
123
278
|
# 调用模型
|
|
@@ -239,87 +394,41 @@ class Agent:
|
|
|
239
394
|
context.set_variable("_agent_call_id", None)
|
|
240
395
|
return
|
|
241
396
|
|
|
242
|
-
#
|
|
243
|
-
for
|
|
244
|
-
|
|
245
|
-
if
|
|
246
|
-
|
|
247
|
-
|
|
248
|
-
|
|
249
|
-
|
|
250
|
-
|
|
251
|
-
|
|
252
|
-
|
|
253
|
-
|
|
254
|
-
|
|
255
|
-
|
|
256
|
-
|
|
257
|
-
)
|
|
258
|
-
|
|
259
|
-
|
|
260
|
-
|
|
261
|
-
|
|
262
|
-
tool_args = json.loads(tool_call["arguments"]) if tool_call["arguments"] else {}
|
|
263
|
-
except json.JSONDecodeError:
|
|
264
|
-
tool_args = {}
|
|
265
|
-
|
|
266
|
-
# 执行前进度事件
|
|
267
|
-
progress_label = tool.get_progress_label(tool_args)
|
|
268
|
-
if progress_label:
|
|
269
|
-
yield event.custom(
|
|
270
|
-
name="skill_progress",
|
|
271
|
-
value={
|
|
272
|
-
"progressId": tool_call["id"],
|
|
273
|
-
"description": progress_label,
|
|
274
|
-
"status": "running",
|
|
275
|
-
},
|
|
276
|
-
)
|
|
277
|
-
|
|
278
|
-
# 执行工具(传入 LLM 解析的参数作为 inputs)
|
|
279
|
-
result = await tool.execute_async(context, tool_args)
|
|
280
|
-
|
|
281
|
-
# 执行后进度事件
|
|
282
|
-
if progress_label:
|
|
283
|
-
yield event.custom(
|
|
284
|
-
name="skill_progress",
|
|
285
|
-
value={
|
|
286
|
-
"progressId": tool_call["id"],
|
|
287
|
-
"status": "completed" if result.success else "failed",
|
|
288
|
-
},
|
|
289
|
-
)
|
|
290
|
-
|
|
291
|
-
# 使用 result.content 作为 LLM 上下文(Tool 已计算好)
|
|
292
|
-
result_content = result.content or (
|
|
293
|
-
json.dumps(result.result, ensure_ascii=False)
|
|
294
|
-
if result.success
|
|
295
|
-
else json.dumps({"error": result.error}, ensure_ascii=False)
|
|
296
|
-
)
|
|
297
|
-
context.add_message(
|
|
298
|
-
self.name,
|
|
299
|
-
"tool",
|
|
300
|
-
content=result_content,
|
|
301
|
-
tool_call_id=tool_call["id"],
|
|
302
|
-
summary=result.summary,
|
|
303
|
-
)
|
|
304
|
-
|
|
305
|
-
# AG-UI: TOOL_CALL_RESULT
|
|
306
|
-
yield event.tool_call_result(tool_call_id=tool_call["id"], content=result_content)
|
|
307
|
-
|
|
308
|
-
# 实时用户进度 — 有 summary 时告知前端
|
|
309
|
-
if result.summary:
|
|
310
|
-
yield event.custom(
|
|
311
|
-
name="tool_progress",
|
|
312
|
-
value={
|
|
313
|
-
"tool_call_id": tool_call["id"],
|
|
314
|
-
"tool_name": result.name,
|
|
315
|
-
"success": result.success,
|
|
316
|
-
"summary": result.summary,
|
|
317
|
-
},
|
|
318
|
-
)
|
|
319
|
-
|
|
320
|
-
# 业务自定义事件 — flush pending
|
|
321
|
-
for pending_evt in context.pop_pending_custom_events():
|
|
322
|
-
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)
|
|
323
417
|
|
|
324
418
|
# 更新消息列表,继续下一轮
|
|
325
419
|
messages = [self.get_system_message()] + context.history + context.get_messages(self.name)
|
|
420
|
+
|
|
421
|
+
# max_turns 耗尽:必须显式收尾,禁止静默截断
|
|
422
|
+
if self.on_max_turns == "error":
|
|
423
|
+
error_msg = f"Agent {self.name} exceeded max_turns={self.max_turns}"
|
|
424
|
+
yield event.run_error(message=error_msg, code="AGENT_MAX_TURNS_EXCEEDED")
|
|
425
|
+
raise AgentError("AGENT_MAX_TURNS_EXCEEDED", 500, {"agent": self.name})
|
|
426
|
+
|
|
427
|
+
# finalize:最后一轮已生成的部分文本作为降级输出,带截断标记
|
|
428
|
+
yield event.custom(
|
|
429
|
+
name="agent_max_turns",
|
|
430
|
+
value={"agentName": self.name, "maxTurns": self.max_turns},
|
|
431
|
+
)
|
|
432
|
+
context.set_output(self.name, {"result": assistant_content, "truncated": True})
|
|
433
|
+
yield event.text_message_end(message_id=msg_id)
|
|
434
|
+
context.set_variable("_agent_call_id", None)
|
|
@@ -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
|
|
@@ -80,6 +84,23 @@ class FlowContext:
|
|
|
80
84
|
"""设置变量值"""
|
|
81
85
|
self.variables[key] = value
|
|
82
86
|
|
|
87
|
+
def pop_variable(self, key: str, default: Any = None) -> Any:
|
|
88
|
+
"""取出并移除变量"""
|
|
89
|
+
return self.variables.pop(key, default)
|
|
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
|
+
|
|
83
104
|
def update_variables(self, updates: dict[str, Any]) -> None:
|
|
84
105
|
"""批量更新变量"""
|
|
85
106
|
self.variables.update(updates)
|