agstack 1.25.1__tar.gz → 2.0.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.0.0}/PKG-INFO +1 -1
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/agent.py +33 -4
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/context.py +4 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/flow.py +51 -116
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/tool_node.py +5 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack.egg-info/PKG-INFO +1 -1
- {agstack-1.25.1 → agstack-2.0.0}/agstack.egg-info/SOURCES.txt +1 -0
- {agstack-1.25.1 → agstack-2.0.0}/pyproject.toml +1 -1
- agstack-2.0.0/tests/test_flow_error_semantics.py +501 -0
- {agstack-1.25.1 → agstack-2.0.0}/LICENSE +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/README.md +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/cache/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/cache/base.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/cache/memory.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/cache/redis.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/config/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/config/logger.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/config/manager.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/config/types.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/contexts.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/decorators.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/events.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/exceptions.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/fastapi/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/fastapi/exception.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/fastapi/middleware.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/fastapi/offline.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/fastapi/sse.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/infra/db/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/infra/es/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/infra/kg/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/infra/mq/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/client.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/event.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/exceptions.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/factory.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/loader.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/agent_node.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/base.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/detect_node.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/echo_node.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/iterator_node.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/llm_chat_node.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/llm_embed_node.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/llm_rerank_node.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/python_node.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/subflow_node.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/nodes/switch_node.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/records.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/registry.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/sandbox.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/state.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/tool.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/flow/trace.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/prompts.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/llm/token.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/messagebus/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/messagebus/base.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/messagebus/memory.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/messagebus/redis.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/schema.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/security/__init__.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/security/casbin.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/security/crypt.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack/status.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack.egg-info/dependency_links.txt +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack.egg-info/requires.txt +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/agstack.egg-info/top_level.txt +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/setup.cfg +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/tests/test_cache_memory.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/tests/test_cache_redis.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/tests/test_flow_io.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/tests/test_flow_iterator.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/tests/test_flow_switch_subflow.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/tests/test_llm_usage_callback.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/tests/test_messagebus_memory.py +0 -0
- {agstack-1.25.1 → agstack-2.0.0}/tests/test_messagebus_redis.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: agstack
|
|
3
|
-
Version:
|
|
3
|
+
Version: 2.0.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>
|
|
@@ -10,7 +10,7 @@ from ..client import get_llm_client
|
|
|
10
10
|
from . import event
|
|
11
11
|
from .context import Usage
|
|
12
12
|
from .event import EventType
|
|
13
|
-
from .exceptions import FlowError
|
|
13
|
+
from .exceptions import AgentError, FlowError
|
|
14
14
|
|
|
15
15
|
|
|
16
16
|
if TYPE_CHECKING:
|
|
@@ -32,6 +32,7 @@ class Agent:
|
|
|
32
32
|
max_turns: int = 10,
|
|
33
33
|
*,
|
|
34
34
|
tool_choice: str = "auto",
|
|
35
|
+
on_max_turns: str = "finalize",
|
|
35
36
|
label: str | None = None,
|
|
36
37
|
echo: bool = False,
|
|
37
38
|
):
|
|
@@ -44,6 +45,7 @@ class Agent:
|
|
|
44
45
|
:param temperature: 温度参数
|
|
45
46
|
:param max_tokens: 最大 token 数
|
|
46
47
|
:param max_turns: 最大轮次
|
|
48
|
+
:param on_max_turns: max_turns 耗尽时的行为,"finalize"(降级输出并标记 truncated)或 "error"(抛出异常)
|
|
47
49
|
:param label: 面向用户的展示名称(控制 STEP 进度事件可见性)
|
|
48
50
|
:param echo: 是否转发 TEXT_MESSAGE 给用户
|
|
49
51
|
"""
|
|
@@ -55,6 +57,7 @@ class Agent:
|
|
|
55
57
|
self.max_tokens = max_tokens
|
|
56
58
|
self.max_turns = max_turns
|
|
57
59
|
self.tool_choice = tool_choice
|
|
60
|
+
self.on_max_turns = on_max_turns
|
|
58
61
|
self.label = label
|
|
59
62
|
self.echo = echo
|
|
60
63
|
|
|
@@ -117,6 +120,7 @@ class Agent:
|
|
|
117
120
|
client = get_llm_client()
|
|
118
121
|
|
|
119
122
|
# Agent 循环
|
|
123
|
+
assistant_content = ""
|
|
120
124
|
for _ in range(self.max_turns):
|
|
121
125
|
context.increment_turn()
|
|
122
126
|
|
|
@@ -257,11 +261,21 @@ class Agent:
|
|
|
257
261
|
)
|
|
258
262
|
continue
|
|
259
263
|
|
|
260
|
-
# 解析 LLM
|
|
264
|
+
# 解析 LLM 返回的工具参数;解析失败作为该次调用的失败反馈给模型,由模型自行重试
|
|
261
265
|
try:
|
|
262
266
|
tool_args = json.loads(tool_call["arguments"]) if tool_call["arguments"] else {}
|
|
263
|
-
except json.JSONDecodeError:
|
|
264
|
-
|
|
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
|
|
265
279
|
|
|
266
280
|
# 执行前进度事件
|
|
267
281
|
progress_label = tool.get_progress_label(tool_args)
|
|
@@ -323,3 +337,18 @@ class Agent:
|
|
|
323
337
|
|
|
324
338
|
# 更新消息列表,继续下一轮
|
|
325
339
|
messages = [self.get_system_message()] + context.history + context.get_messages(self.name)
|
|
340
|
+
|
|
341
|
+
# max_turns 耗尽:必须显式收尾,禁止静默截断
|
|
342
|
+
if self.on_max_turns == "error":
|
|
343
|
+
error_msg = f"Agent {self.name} exceeded max_turns={self.max_turns}"
|
|
344
|
+
yield event.run_error(message=error_msg, code="AGENT_MAX_TURNS_EXCEEDED")
|
|
345
|
+
raise AgentError("AGENT_MAX_TURNS_EXCEEDED", 500, {"agent": self.name})
|
|
346
|
+
|
|
347
|
+
# finalize:最后一轮已生成的部分文本作为降级输出,带截断标记
|
|
348
|
+
yield event.custom(
|
|
349
|
+
name="agent_max_turns",
|
|
350
|
+
value={"agentName": self.name, "maxTurns": self.max_turns},
|
|
351
|
+
)
|
|
352
|
+
context.set_output(self.name, {"result": assistant_content, "truncated": True})
|
|
353
|
+
yield event.text_message_end(message_id=msg_id)
|
|
354
|
+
context.set_variable("_agent_call_id", None)
|
|
@@ -80,6 +80,10 @@ class FlowContext:
|
|
|
80
80
|
"""设置变量值"""
|
|
81
81
|
self.variables[key] = value
|
|
82
82
|
|
|
83
|
+
def pop_variable(self, key: str, default: Any = None) -> Any:
|
|
84
|
+
"""取出并移除变量"""
|
|
85
|
+
return self.variables.pop(key, default)
|
|
86
|
+
|
|
83
87
|
def update_variables(self, updates: dict[str, Any]) -> None:
|
|
84
88
|
"""批量更新变量"""
|
|
85
89
|
self.variables.update(updates)
|
|
@@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, AsyncIterator
|
|
|
9
9
|
from uuid import uuid4
|
|
10
10
|
|
|
11
11
|
from . import event
|
|
12
|
+
from .context import Usage
|
|
12
13
|
from .exceptions import NodeExecutionError
|
|
13
14
|
|
|
14
15
|
|
|
@@ -39,6 +40,23 @@ def _parse_literal(s: str) -> Any:
|
|
|
39
40
|
return s
|
|
40
41
|
|
|
41
42
|
|
|
43
|
+
def _usage_snapshot(usage: Usage) -> tuple[int, int, int]:
|
|
44
|
+
"""记录节点执行前的用量快照,用于差值归因"""
|
|
45
|
+
return (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens)
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def _usage_delta(usage: Usage, before: tuple[int, int, int]) -> Usage | None:
|
|
49
|
+
"""节点执行前后的用量差值;全零返回 None,避免 trace 体积膨胀"""
|
|
50
|
+
delta = Usage(
|
|
51
|
+
prompt_tokens=usage.prompt_tokens - before[0],
|
|
52
|
+
completion_tokens=usage.completion_tokens - before[1],
|
|
53
|
+
total_tokens=usage.total_tokens - before[2],
|
|
54
|
+
)
|
|
55
|
+
if delta.prompt_tokens or delta.completion_tokens or delta.total_tokens:
|
|
56
|
+
return delta
|
|
57
|
+
return None
|
|
58
|
+
|
|
59
|
+
|
|
42
60
|
@dataclass
|
|
43
61
|
class RetryPolicy:
|
|
44
62
|
"""节点重试策略"""
|
|
@@ -262,120 +280,14 @@ class Flow:
|
|
|
262
280
|
return False
|
|
263
281
|
|
|
264
282
|
async def run(self, context: "FlowContext") -> dict[str, Any]:
|
|
265
|
-
"""执行 Flow
|
|
266
|
-
if not self.edges:
|
|
267
|
-
for node in self.nodes:
|
|
268
|
-
node_id = node.get("id")
|
|
269
|
-
if not node_id:
|
|
270
|
-
continue
|
|
271
|
-
context.current_node = node_id
|
|
272
|
-
node_type: str = node.get("type", "")
|
|
273
|
-
handler = self._node_handlers.get(node_type)
|
|
274
|
-
if handler:
|
|
275
|
-
result = await handler.execute(node, context)
|
|
276
|
-
context.set_output(node_id, result)
|
|
277
|
-
else:
|
|
278
|
-
raise NodeExecutionError("UNKNOWN_NODE_TYPE", args={"node_type": node_type})
|
|
279
|
-
else:
|
|
280
|
-
current_node_id: str | None = self.nodes[0]["id"] if self.nodes else None
|
|
281
|
-
visit_count: dict[str, int] = {}
|
|
282
|
-
|
|
283
|
-
while current_node_id:
|
|
284
|
-
node = self.get_node_config(current_node_id)
|
|
285
|
-
if not node:
|
|
286
|
-
break
|
|
287
|
-
|
|
288
|
-
# 循环计数与超限检测
|
|
289
|
-
visit_count[current_node_id] = visit_count.get(current_node_id, 0) + 1
|
|
290
|
-
force_fallback = self._check_cycle_limit(current_node_id, visit_count)
|
|
291
|
-
if force_fallback:
|
|
292
|
-
current_node_id = self._resolve_next_node(current_node_id, context, force_fallback=True)
|
|
293
|
-
continue
|
|
294
|
-
|
|
295
|
-
context.current_node = current_node_id
|
|
296
|
-
node_type: str = node.get("type", "")
|
|
297
|
-
|
|
298
|
-
if node_type == "message":
|
|
299
|
-
config = node.get("config", {})
|
|
300
|
-
template = config.get("content", "")
|
|
301
|
-
text = template.format_map(_SafeFormatDict(context.variables))
|
|
302
|
-
context.set_output(current_node_id, {"result": text})
|
|
303
|
-
current_node_id = self._resolve_next_node(current_node_id, context)
|
|
304
|
-
|
|
305
|
-
elif node_type == "parallel":
|
|
306
|
-
config = node.get("config", {})
|
|
307
|
-
branches: list[str] = config.get("branches", [])
|
|
308
|
-
|
|
309
|
-
async def _run_branch(branch_id: str) -> None:
|
|
310
|
-
branch_node = self.get_node_config(branch_id)
|
|
311
|
-
if not branch_node:
|
|
312
|
-
return
|
|
313
|
-
context.current_node = branch_id
|
|
314
|
-
branch_type: str = branch_node.get("type", "")
|
|
315
|
-
branch_handler = self._node_handlers.get(branch_type)
|
|
316
|
-
if branch_handler:
|
|
317
|
-
result = await branch_handler.execute(branch_node, context)
|
|
318
|
-
context.set_output(branch_id, result)
|
|
319
|
-
|
|
320
|
-
await asyncio.gather(*[_run_branch(bid) for bid in branches])
|
|
321
|
-
merged: dict[str, Any] = {}
|
|
322
|
-
for bid in branches:
|
|
323
|
-
branch_result = context.outputs.get(bid, {})
|
|
324
|
-
if isinstance(branch_result, dict):
|
|
325
|
-
merged.update(branch_result)
|
|
326
|
-
context.set_output(current_node_id, merged)
|
|
327
|
-
current_node_id = self._resolve_next_node(current_node_id, context)
|
|
328
|
-
|
|
329
|
-
elif node_type == "iteration":
|
|
330
|
-
config = node.get("config", {})
|
|
331
|
-
items_ref = config.get("items", "")
|
|
332
|
-
items = context.resolve_reference(items_ref) if isinstance(items_ref, str) else items_ref
|
|
333
|
-
if not isinstance(items, list):
|
|
334
|
-
items = [items]
|
|
335
|
-
|
|
336
|
-
item_var = config.get("item_variable", "item")
|
|
337
|
-
index_var = config.get("index_variable", "index")
|
|
338
|
-
body_nodes: list[str] = config.get("body", [])
|
|
339
|
-
results: list[Any] = []
|
|
340
|
-
|
|
341
|
-
for idx, item in enumerate(items):
|
|
342
|
-
context.set_variable(item_var, item)
|
|
343
|
-
context.set_variable(index_var, idx)
|
|
344
|
-
for body_node_id in body_nodes:
|
|
345
|
-
body_node = self.get_node_config(body_node_id)
|
|
346
|
-
if not body_node:
|
|
347
|
-
continue
|
|
348
|
-
body_type: str = body_node.get("type", "")
|
|
349
|
-
body_handler = self._node_handlers.get(body_type)
|
|
350
|
-
if body_handler:
|
|
351
|
-
body_result = await body_handler.execute(body_node, context)
|
|
352
|
-
context.set_output(body_node_id, body_result)
|
|
353
|
-
if body_nodes:
|
|
354
|
-
results.append(context.outputs.get(body_nodes[-1]))
|
|
355
|
-
|
|
356
|
-
context.set_output(current_node_id, {"results": results})
|
|
357
|
-
current_node_id = self._resolve_next_node(current_node_id, context)
|
|
358
|
-
|
|
359
|
-
elif node_type in self._node_handlers:
|
|
360
|
-
handler = self._node_handlers[node_type]
|
|
361
|
-
try:
|
|
362
|
-
result = await handler.execute(node, context)
|
|
363
|
-
except Exception as e:
|
|
364
|
-
iter_target = self._find_iterator_fallback(current_node_id)
|
|
365
|
-
if iter_target:
|
|
366
|
-
context.set_variable(f"_iter_{iter_target}_error", str(e))
|
|
367
|
-
context.set_output(current_node_id, {"error": str(e)})
|
|
368
|
-
context.set_variable("_prev_node_id", current_node_id)
|
|
369
|
-
current_node_id = iter_target
|
|
370
|
-
continue
|
|
371
|
-
raise
|
|
372
|
-
context.set_output(current_node_id, result)
|
|
373
|
-
context.set_variable("_prev_node_id", current_node_id)
|
|
374
|
-
current_node_id = self._resolve_next_node(current_node_id, context)
|
|
375
|
-
|
|
376
|
-
else:
|
|
377
|
-
raise NodeExecutionError("UNKNOWN_NODE_TYPE", args={"node_type": node_type})
|
|
283
|
+
"""执行 Flow(非流式)
|
|
378
284
|
|
|
285
|
+
stream() 的消费包装:两条路径共享同一执行引擎,重试策略、
|
|
286
|
+
FlowTrace、output_mode、iterator 状态清理等行为完全一致。
|
|
287
|
+
节点失败抛 NodeExecutionError(包装原始异常)。
|
|
288
|
+
"""
|
|
289
|
+
async for _ in self.stream(context):
|
|
290
|
+
pass
|
|
379
291
|
return context.outputs
|
|
380
292
|
|
|
381
293
|
async def stream(self, context: "FlowContext") -> AsyncIterator[dict[str, Any]]:
|
|
@@ -459,6 +371,7 @@ class Flow:
|
|
|
459
371
|
if node_type == "message":
|
|
460
372
|
msg_config = node.get("config", {})
|
|
461
373
|
context.trace.record_node_start(current_node_id, "message", inputs=msg_config)
|
|
374
|
+
usage_before = _usage_snapshot(context.usage)
|
|
462
375
|
|
|
463
376
|
# message 节点增加 STEP 事件
|
|
464
377
|
msg_sid = str(uuid4())
|
|
@@ -485,7 +398,11 @@ class Flow:
|
|
|
485
398
|
fin_evt["_echo"] = msg_config.get("echo", True)
|
|
486
399
|
yield fin_evt
|
|
487
400
|
|
|
488
|
-
context.trace.record_node_end(
|
|
401
|
+
context.trace.record_node_end(
|
|
402
|
+
current_node_id,
|
|
403
|
+
outputs={"result": text},
|
|
404
|
+
usage=_usage_delta(context.usage, usage_before),
|
|
405
|
+
)
|
|
489
406
|
current_node_id = self._resolve_next_node(current_node_id, context)
|
|
490
407
|
|
|
491
408
|
elif node_type == "parallel":
|
|
@@ -493,6 +410,7 @@ class Flow:
|
|
|
493
410
|
branches = config.get("branches", [])
|
|
494
411
|
|
|
495
412
|
context.trace.record_node_start(current_node_id, "parallel", inputs=config)
|
|
413
|
+
usage_before = _usage_snapshot(context.usage)
|
|
496
414
|
|
|
497
415
|
parallel_sid = str(uuid4())
|
|
498
416
|
step_evt = event.step_started(step_name=f"parallel:{current_node_id}", step_id=parallel_sid)
|
|
@@ -523,7 +441,12 @@ class Flow:
|
|
|
523
441
|
try:
|
|
524
442
|
result = await branch_handler.execute(branch_node, context)
|
|
525
443
|
context.set_output(branch_id, result)
|
|
526
|
-
context
|
|
444
|
+
# 分支并发共享 context,差值无法按分支切分——分支不记 usage,整体归因容器节点
|
|
445
|
+
context.trace.record_node_end(
|
|
446
|
+
branch_id,
|
|
447
|
+
outputs=result,
|
|
448
|
+
error=context.pop_variable("_last_node_error"),
|
|
449
|
+
)
|
|
527
450
|
except Exception as e:
|
|
528
451
|
context.trace.record_node_end(branch_id, error=str(e))
|
|
529
452
|
raise
|
|
@@ -542,7 +465,11 @@ class Flow:
|
|
|
542
465
|
fin_evt["_echo"] = False
|
|
543
466
|
yield fin_evt
|
|
544
467
|
|
|
545
|
-
context.trace.record_node_end(
|
|
468
|
+
context.trace.record_node_end(
|
|
469
|
+
current_node_id,
|
|
470
|
+
outputs=merged,
|
|
471
|
+
usage=_usage_delta(context.usage, usage_before),
|
|
472
|
+
)
|
|
546
473
|
current_node_id = self._resolve_next_node(current_node_id, context)
|
|
547
474
|
|
|
548
475
|
elif node_type == "iteration":
|
|
@@ -586,6 +513,7 @@ class Flow:
|
|
|
586
513
|
parent_id=context.trace._qualify_id(current_node_id),
|
|
587
514
|
iteration_index=idx,
|
|
588
515
|
)
|
|
516
|
+
body_usage_before = _usage_snapshot(context.usage)
|
|
589
517
|
body_result = await body_handler.execute(body_node, context)
|
|
590
518
|
context.set_output(body_node_id, body_result)
|
|
591
519
|
# 收集 body 节点产生的 execution_records
|
|
@@ -593,6 +521,8 @@ class Flow:
|
|
|
593
521
|
context.trace.record_node_end(
|
|
594
522
|
body_node_id,
|
|
595
523
|
outputs=body_result,
|
|
524
|
+
error=context.pop_variable("_last_node_error"),
|
|
525
|
+
usage=_usage_delta(context.usage, body_usage_before),
|
|
596
526
|
tool_calls=body_tool_calls if body_tool_calls else None,
|
|
597
527
|
)
|
|
598
528
|
if body_nodes:
|
|
@@ -607,6 +537,7 @@ class Flow:
|
|
|
607
537
|
fin_evt["_echo"] = False
|
|
608
538
|
yield fin_evt
|
|
609
539
|
|
|
540
|
+
# body 串行执行已按差值归因 usage,容器不重复归因
|
|
610
541
|
context.trace.record_node_end(current_node_id, outputs=iteration_output)
|
|
611
542
|
current_node_id = self._resolve_next_node(current_node_id, context)
|
|
612
543
|
|
|
@@ -622,6 +553,7 @@ class Flow:
|
|
|
622
553
|
inputs=resolved_inputs,
|
|
623
554
|
label=config.get("label"),
|
|
624
555
|
)
|
|
556
|
+
usage_before = _usage_snapshot(context.usage)
|
|
625
557
|
|
|
626
558
|
# output_mode: "append" — 保存执行前的历史
|
|
627
559
|
append_mode = config.get("output_mode") == "append"
|
|
@@ -666,6 +598,9 @@ class Flow:
|
|
|
666
598
|
context.trace.record_node_end(
|
|
667
599
|
current_node_id,
|
|
668
600
|
outputs=context.outputs.get(current_node_id),
|
|
601
|
+
# 节点内部容错的失败(如 tool 节点 on_error: "continue")经 context 传递
|
|
602
|
+
error=context.pop_variable("_last_node_error"),
|
|
603
|
+
usage=_usage_delta(context.usage, usage_before),
|
|
669
604
|
tool_calls=tool_calls if tool_calls else None,
|
|
670
605
|
messages=messages,
|
|
671
606
|
)
|
|
@@ -33,5 +33,10 @@ class ToolNodeHandler(NodeHandler):
|
|
|
33
33
|
tool = self._create_tool(config)
|
|
34
34
|
result = await tool.execute_async(context, inputs=resolved)
|
|
35
35
|
if not result.success:
|
|
36
|
+
# on_error: "continue" — 失败降级为节点输出,flow 继续走边路由,
|
|
37
|
+
# 条件边可用 $o.<node>.success == false 分流;默认 "raise" 保持原语义
|
|
38
|
+
if config.get("on_error") == "continue":
|
|
39
|
+
context.set_variable("_last_node_error", result.error)
|
|
40
|
+
return {"success": False, "error": result.error}
|
|
36
41
|
raise ToolExecutionError("TOOL_EXECUTION_FAILED", args={"tool_name": tool.name, "error": result.error})
|
|
37
42
|
return result.result
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: agstack
|
|
3
|
-
Version:
|
|
3
|
+
Version: 2.0.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>
|
|
@@ -0,0 +1,501 @@
|
|
|
1
|
+
# Copyright (c) 2020-2026 XtraVisions, All rights reserved.
|
|
2
|
+
|
|
3
|
+
"""Flow 系统失败语义测试 — 1.26.0 修复的五项缺陷(D1-D5)验收用例
|
|
4
|
+
|
|
5
|
+
D1: Agent max_turns 耗尽显式收尾
|
|
6
|
+
D2: 工具参数 JSON 解析失败反馈给模型
|
|
7
|
+
D3: tool 节点 on_error 开关
|
|
8
|
+
D4: NodeTrace.usage 按节点归因
|
|
9
|
+
D5: run() 收敛为 stream() 消费者
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
import asyncio
|
|
13
|
+
import json
|
|
14
|
+
from types import SimpleNamespace
|
|
15
|
+
from unittest.mock import AsyncMock, MagicMock, patch
|
|
16
|
+
|
|
17
|
+
import pytest
|
|
18
|
+
|
|
19
|
+
from agstack.llm.flow.agent import Agent
|
|
20
|
+
from agstack.llm.flow.context import FlowContext
|
|
21
|
+
from agstack.llm.flow.event import EventType
|
|
22
|
+
from agstack.llm.flow.exceptions import AgentError, NodeExecutionError
|
|
23
|
+
from agstack.llm.flow.flow import Flow
|
|
24
|
+
from agstack.llm.flow.registry import registry
|
|
25
|
+
from agstack.llm.flow.tool import Tool
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def _run(coro):
|
|
29
|
+
return asyncio.get_event_loop().run_until_complete(coro)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
# ── 可编程 LLM 流式客户端桩 ──
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _text_chunk(text: str):
|
|
36
|
+
return SimpleNamespace(
|
|
37
|
+
choices=[SimpleNamespace(delta=SimpleNamespace(content=text, tool_calls=None), finish_reason=None)]
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _tool_call_chunk(call_id: str, name: str, arguments: str):
|
|
42
|
+
tc = SimpleNamespace(index=0, id=call_id, function=SimpleNamespace(name=name, arguments=arguments))
|
|
43
|
+
return SimpleNamespace(
|
|
44
|
+
choices=[SimpleNamespace(delta=SimpleNamespace(content=None, tool_calls=[tc]), finish_reason=None)]
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def _finish_chunk(reason: str = "stop"):
|
|
49
|
+
return SimpleNamespace(
|
|
50
|
+
choices=[SimpleNamespace(delta=SimpleNamespace(content=None, tool_calls=None), finish_reason=reason)]
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
class FakeStreamClient:
|
|
55
|
+
"""按轮次返回预编程 chunk 序列;超出编程轮次时重复最后一轮"""
|
|
56
|
+
|
|
57
|
+
def __init__(self, turns: list[list]):
|
|
58
|
+
self.turns = turns
|
|
59
|
+
self.requests: list[dict] = []
|
|
60
|
+
|
|
61
|
+
async def chat(self, stream: bool = True, **kwargs):
|
|
62
|
+
self.requests.append(kwargs)
|
|
63
|
+
idx = min(len(self.requests) - 1, len(self.turns) - 1)
|
|
64
|
+
chunks = self.turns[idx]
|
|
65
|
+
|
|
66
|
+
async def _gen():
|
|
67
|
+
for c in chunks:
|
|
68
|
+
yield c
|
|
69
|
+
|
|
70
|
+
return _gen()
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
async def _collect(aiter):
|
|
74
|
+
return [evt async for evt in aiter]
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
# ── D1: Agent max_turns 耗尽显式收尾 ──
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
class TestAgentMaxTurns:
|
|
81
|
+
def _looping_tool(self, counter: dict) -> Tool:
|
|
82
|
+
def fn(context, inputs):
|
|
83
|
+
counter["n"] = counter.get("n", 0) + 1
|
|
84
|
+
return {"note": "need more tool calls"}
|
|
85
|
+
|
|
86
|
+
return Tool(name="looper", description="always asks for more", function=fn)
|
|
87
|
+
|
|
88
|
+
@patch("agstack.llm.flow.agent.get_llm_client")
|
|
89
|
+
def test_exhaustion_finalizes_explicitly(self, mock_get_client):
|
|
90
|
+
"""打满场景:END 事件闭合、agent_max_turns CUSTOM 事件、输出含 truncated 标记"""
|
|
91
|
+
counter: dict = {}
|
|
92
|
+
turn = [
|
|
93
|
+
_text_chunk("thinking"),
|
|
94
|
+
_tool_call_chunk("tc1", "looper", "{}"),
|
|
95
|
+
_finish_chunk("tool_calls"),
|
|
96
|
+
]
|
|
97
|
+
mock_get_client.return_value = FakeStreamClient([turn])
|
|
98
|
+
|
|
99
|
+
agent = Agent(name="worker", tools=[self._looping_tool(counter)], max_turns=2)
|
|
100
|
+
ctx = FlowContext(variables={"input": "go"})
|
|
101
|
+
events = _run(_collect(agent.stream(ctx)))
|
|
102
|
+
|
|
103
|
+
types = [e["type"] for e in events]
|
|
104
|
+
assert types.count(EventType.TEXT_MESSAGE_END) == 1
|
|
105
|
+
customs = [e for e in events if e["type"] == EventType.CUSTOM and e.get("name") == "agent_max_turns"]
|
|
106
|
+
assert len(customs) == 1
|
|
107
|
+
assert customs[0]["value"] == {"agentName": "worker", "maxTurns": 2}
|
|
108
|
+
assert ctx.outputs["worker"]["truncated"] is True
|
|
109
|
+
assert ctx.outputs["worker"]["result"] == "thinking"
|
|
110
|
+
assert ctx.get_variable("_agent_call_id") is None
|
|
111
|
+
assert counter["n"] == 2 # 两轮各执行一次工具
|
|
112
|
+
|
|
113
|
+
@patch("agstack.llm.flow.agent.get_llm_client")
|
|
114
|
+
def test_normal_exit_unchanged(self, mock_get_client):
|
|
115
|
+
"""正常场景:无 truncated 键、无 agent_max_turns 事件"""
|
|
116
|
+
mock_get_client.return_value = FakeStreamClient([[_text_chunk("hello"), _finish_chunk()]])
|
|
117
|
+
|
|
118
|
+
agent = Agent(name="worker", max_turns=2)
|
|
119
|
+
ctx = FlowContext(variables={"input": "hi"})
|
|
120
|
+
events = _run(_collect(agent.stream(ctx)))
|
|
121
|
+
|
|
122
|
+
assert ctx.outputs["worker"] == {"result": "hello"}
|
|
123
|
+
assert "truncated" not in ctx.outputs["worker"]
|
|
124
|
+
assert [e["type"] for e in events].count(EventType.TEXT_MESSAGE_END) == 1
|
|
125
|
+
assert not any(e["type"] == EventType.CUSTOM and e.get("name") == "agent_max_turns" for e in events)
|
|
126
|
+
|
|
127
|
+
@patch("agstack.llm.flow.agent.get_llm_client")
|
|
128
|
+
def test_run_returns_partial_text_on_exhaustion(self, mock_get_client):
|
|
129
|
+
counter: dict = {}
|
|
130
|
+
turn = [
|
|
131
|
+
_text_chunk("partial"),
|
|
132
|
+
_tool_call_chunk("tc1", "looper", "{}"),
|
|
133
|
+
_finish_chunk("tool_calls"),
|
|
134
|
+
]
|
|
135
|
+
mock_get_client.return_value = FakeStreamClient([turn])
|
|
136
|
+
|
|
137
|
+
agent = Agent(name="worker", tools=[self._looping_tool(counter)], max_turns=2)
|
|
138
|
+
result = _run(agent.run(FlowContext(variables={"input": "go"})))
|
|
139
|
+
assert "partial" in result["result"]
|
|
140
|
+
|
|
141
|
+
@patch("agstack.llm.flow.agent.get_llm_client")
|
|
142
|
+
def test_on_max_turns_error_mode(self, mock_get_client):
|
|
143
|
+
"""严格模式:yield RUN_ERROR 后抛 AgentError"""
|
|
144
|
+
counter: dict = {}
|
|
145
|
+
turn = [_tool_call_chunk("tc1", "looper", "{}"), _finish_chunk("tool_calls")]
|
|
146
|
+
mock_get_client.return_value = FakeStreamClient([turn])
|
|
147
|
+
|
|
148
|
+
agent = Agent(name="worker", tools=[self._looping_tool(counter)], max_turns=1, on_max_turns="error")
|
|
149
|
+
ctx = FlowContext(variables={"input": "go"})
|
|
150
|
+
|
|
151
|
+
async def _consume():
|
|
152
|
+
events = []
|
|
153
|
+
with pytest.raises(AgentError):
|
|
154
|
+
async for evt in agent.stream(ctx):
|
|
155
|
+
events.append(evt)
|
|
156
|
+
return events
|
|
157
|
+
|
|
158
|
+
events = _run(_consume())
|
|
159
|
+
assert any(e["type"] == EventType.RUN_ERROR and e.get("code") == "AGENT_MAX_TURNS_EXCEEDED" for e in events)
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
# ── D2: 工具参数 JSON 解析失败反馈给模型 ──
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
class TestToolArgsParseFailure:
|
|
166
|
+
@patch("agstack.llm.flow.agent.get_llm_client")
|
|
167
|
+
def test_invalid_json_feeds_error_back_to_model(self, mock_get_client):
|
|
168
|
+
"""解析失败:工具不被调用,错误进 tool message 与 TOOL_CALL_RESULT,下一轮模型可见"""
|
|
169
|
+
counter: dict = {"n": 0}
|
|
170
|
+
|
|
171
|
+
def fn(context, inputs):
|
|
172
|
+
counter["n"] += 1
|
|
173
|
+
return {"docs": []}
|
|
174
|
+
|
|
175
|
+
tool = Tool(name="search", description="retrieval", function=fn)
|
|
176
|
+
bad_args = '{"query": "未闭合'
|
|
177
|
+
client = FakeStreamClient(
|
|
178
|
+
[
|
|
179
|
+
[_tool_call_chunk("tc1", "search", bad_args), _finish_chunk("tool_calls")],
|
|
180
|
+
[_text_chunk("done"), _finish_chunk()],
|
|
181
|
+
]
|
|
182
|
+
)
|
|
183
|
+
mock_get_client.return_value = client
|
|
184
|
+
|
|
185
|
+
agent = Agent(name="worker", tools=[tool])
|
|
186
|
+
ctx = FlowContext(variables={"input": "find it"})
|
|
187
|
+
events = _run(_collect(agent.stream(ctx)))
|
|
188
|
+
|
|
189
|
+
# 工具函数不被调用
|
|
190
|
+
assert counter["n"] == 0
|
|
191
|
+
|
|
192
|
+
# 事件流出现含 parse failed 信息的 TOOL_CALL_RESULT
|
|
193
|
+
results = [e for e in events if e["type"] == EventType.TOOL_CALL_RESULT]
|
|
194
|
+
assert len(results) == 1
|
|
195
|
+
payload = json.loads(results[0]["content"])
|
|
196
|
+
assert "JSON parse failed" in payload["error"]
|
|
197
|
+
assert payload["raw_arguments"] == bad_args
|
|
198
|
+
|
|
199
|
+
# 下一轮模型请求的 messages 中含该 tool 角色错误消息
|
|
200
|
+
second_request_messages = client.requests[1]["messages"]
|
|
201
|
+
tool_msgs = [m for m in second_request_messages if m.get("role") == "tool"]
|
|
202
|
+
assert len(tool_msgs) == 1
|
|
203
|
+
assert "JSON parse failed" in tool_msgs[0]["content"]
|
|
204
|
+
assert tool_msgs[0]["tool_call_id"] == "tc1"
|
|
205
|
+
|
|
206
|
+
# 循环正常继续并结束
|
|
207
|
+
assert ctx.outputs["worker"] == {"result": "done"}
|
|
208
|
+
|
|
209
|
+
@patch("agstack.llm.flow.agent.get_llm_client")
|
|
210
|
+
def test_empty_arguments_still_calls_tool(self, mock_get_client):
|
|
211
|
+
"""arguments 为空字符串维持现行为:无参工具正常调用"""
|
|
212
|
+
captured: dict = {"n": 0}
|
|
213
|
+
|
|
214
|
+
def fn(context, inputs):
|
|
215
|
+
captured["n"] += 1
|
|
216
|
+
captured["inputs"] = inputs
|
|
217
|
+
return {"ok": True}
|
|
218
|
+
|
|
219
|
+
tool = Tool(name="noargs", description="no-arg tool", function=fn)
|
|
220
|
+
client = FakeStreamClient(
|
|
221
|
+
[
|
|
222
|
+
[_tool_call_chunk("tc1", "noargs", ""), _finish_chunk("tool_calls")],
|
|
223
|
+
[_text_chunk("done"), _finish_chunk()],
|
|
224
|
+
]
|
|
225
|
+
)
|
|
226
|
+
mock_get_client.return_value = client
|
|
227
|
+
|
|
228
|
+
agent = Agent(name="worker", tools=[tool])
|
|
229
|
+
_run(_collect(agent.stream(FlowContext(variables={"input": "go"}))))
|
|
230
|
+
assert captured["n"] == 1
|
|
231
|
+
assert captured["inputs"] == {}
|
|
232
|
+
|
|
233
|
+
|
|
234
|
+
# ── D3: tool 节点 on_error 开关 ──
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def _fail_fn(context, inputs):
|
|
238
|
+
raise ValueError("boom")
|
|
239
|
+
|
|
240
|
+
|
|
241
|
+
class TestToolNodeOnError:
|
|
242
|
+
def test_default_raise_preserved(self):
|
|
243
|
+
"""不写 on_error:失败仍中断 flow(NodeExecutionError 包装)"""
|
|
244
|
+
registry.register_tool("d3_fail_raise", Tool(name="d3_fail_raise", description="", function=_fail_fn))
|
|
245
|
+
flow = Flow(
|
|
246
|
+
flow_id="t",
|
|
247
|
+
name="t",
|
|
248
|
+
nodes=[{"id": "search", "type": "tool", "config": {"tool_name": "d3_fail_raise"}}],
|
|
249
|
+
)
|
|
250
|
+
with pytest.raises(NodeExecutionError):
|
|
251
|
+
_run(flow.run(FlowContext()))
|
|
252
|
+
|
|
253
|
+
def test_on_error_continue_routes_condition_edge(self):
|
|
254
|
+
"""on_error: continue:flow 不中断,条件边按 success == false 分流,trace 记录 error"""
|
|
255
|
+
registry.register_tool("d3_fail_cont", Tool(name="d3_fail_cont", description="", function=_fail_fn))
|
|
256
|
+
flow = Flow(
|
|
257
|
+
flow_id="t",
|
|
258
|
+
name="t",
|
|
259
|
+
nodes=[
|
|
260
|
+
{
|
|
261
|
+
"id": "search",
|
|
262
|
+
"type": "tool",
|
|
263
|
+
"config": {"tool_name": "d3_fail_cont", "on_error": "continue"},
|
|
264
|
+
},
|
|
265
|
+
{
|
|
266
|
+
"id": "fallback",
|
|
267
|
+
"type": "python",
|
|
268
|
+
"config": {"code": "def main(**kwargs):\n return {'handled': True}"},
|
|
269
|
+
},
|
|
270
|
+
{
|
|
271
|
+
"id": "happy",
|
|
272
|
+
"type": "python",
|
|
273
|
+
"config": {"code": "def main(**kwargs):\n return {'happy': True}"},
|
|
274
|
+
},
|
|
275
|
+
],
|
|
276
|
+
edges=[
|
|
277
|
+
{"source": "search", "condition": "$o.search.success == false", "target": "fallback"},
|
|
278
|
+
{"source": "search", "target": "happy"},
|
|
279
|
+
],
|
|
280
|
+
)
|
|
281
|
+
ctx = FlowContext()
|
|
282
|
+
_run(flow.run(ctx))
|
|
283
|
+
|
|
284
|
+
assert ctx.outputs["search"]["success"] is False
|
|
285
|
+
assert "boom" in ctx.outputs["search"]["error"]
|
|
286
|
+
assert ctx.outputs["fallback"] == {"handled": True}
|
|
287
|
+
assert "happy" not in ctx.outputs
|
|
288
|
+
|
|
289
|
+
search_trace = next(n for n in ctx.trace.nodes if n.node_id == "search")
|
|
290
|
+
assert search_trace.error is not None and "boom" in search_trace.error
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
# ── D4: NodeTrace.usage 按节点归因 ──
|
|
294
|
+
|
|
295
|
+
|
|
296
|
+
def _chat_response(text: str, prompt: int, completion: int):
|
|
297
|
+
resp = MagicMock()
|
|
298
|
+
choice = MagicMock()
|
|
299
|
+
choice.message.content = text
|
|
300
|
+
resp.choices = [choice]
|
|
301
|
+
resp.usage = MagicMock(prompt_tokens=prompt, completion_tokens=completion, total_tokens=prompt + completion)
|
|
302
|
+
return resp
|
|
303
|
+
|
|
304
|
+
|
|
305
|
+
_NEVER_EDGE = {"condition": "$v._never == yes"} # 恒不满足且无 fallback:驱动 edge-driven 路径后自然结束
|
|
306
|
+
|
|
307
|
+
|
|
308
|
+
class TestNodeUsageAttribution:
|
|
309
|
+
@patch("agstack.llm.flow.nodes.llm_chat_node.get_llm_client")
|
|
310
|
+
def test_single_llm_node(self, mock_get_client):
|
|
311
|
+
mock_client = AsyncMock()
|
|
312
|
+
mock_client.chat = AsyncMock(return_value=_chat_response("hi", 10, 5))
|
|
313
|
+
mock_get_client.return_value = mock_client
|
|
314
|
+
|
|
315
|
+
flow = Flow(
|
|
316
|
+
flow_id="t",
|
|
317
|
+
name="t",
|
|
318
|
+
nodes=[{"id": "chat1", "type": "llm_chat", "config": {"prompt": "hello"}}],
|
|
319
|
+
edges=[{"source": "chat1", "target": "chat1", **_NEVER_EDGE}],
|
|
320
|
+
)
|
|
321
|
+
ctx = FlowContext()
|
|
322
|
+
_run(flow.run(ctx))
|
|
323
|
+
|
|
324
|
+
node = next(n for n in ctx.trace.nodes if n.node_id == "chat1")
|
|
325
|
+
assert node.usage is not None
|
|
326
|
+
assert node.usage.total_tokens == 15
|
|
327
|
+
assert node.usage.total_tokens == ctx.trace.total_usage.total_tokens
|
|
328
|
+
|
|
329
|
+
@patch("agstack.llm.flow.nodes.llm_chat_node.get_llm_client")
|
|
330
|
+
def test_two_serial_llm_nodes_sum_to_total(self, mock_get_client):
|
|
331
|
+
mock_client = AsyncMock()
|
|
332
|
+
mock_client.chat = AsyncMock(side_effect=[_chat_response("a", 10, 5), _chat_response("b", 6, 2)])
|
|
333
|
+
mock_get_client.return_value = mock_client
|
|
334
|
+
|
|
335
|
+
flow = Flow(
|
|
336
|
+
flow_id="t",
|
|
337
|
+
name="t",
|
|
338
|
+
nodes=[
|
|
339
|
+
{"id": "chat1", "type": "llm_chat", "config": {"prompt": "one"}},
|
|
340
|
+
{"id": "chat2", "type": "llm_chat", "config": {"prompt": "two"}},
|
|
341
|
+
],
|
|
342
|
+
edges=[{"source": "chat1", "target": "chat2"}],
|
|
343
|
+
)
|
|
344
|
+
ctx = FlowContext()
|
|
345
|
+
_run(flow.run(ctx))
|
|
346
|
+
|
|
347
|
+
n1 = next(n for n in ctx.trace.nodes if n.node_id == "chat1")
|
|
348
|
+
n2 = next(n for n in ctx.trace.nodes if n.node_id == "chat2")
|
|
349
|
+
assert n1.usage is not None and n2.usage is not None
|
|
350
|
+
assert n1.usage.total_tokens == 15
|
|
351
|
+
assert n2.usage.total_tokens == 8
|
|
352
|
+
assert n1.usage.total_tokens + n2.usage.total_tokens == ctx.trace.total_usage.total_tokens
|
|
353
|
+
|
|
354
|
+
def test_non_llm_node_usage_is_none(self):
|
|
355
|
+
"""无 LLM 调用的节点:usage 为 None 而非零值 Usage"""
|
|
356
|
+
flow = Flow(
|
|
357
|
+
flow_id="t",
|
|
358
|
+
name="t",
|
|
359
|
+
nodes=[
|
|
360
|
+
{
|
|
361
|
+
"id": "py1",
|
|
362
|
+
"type": "python",
|
|
363
|
+
"config": {"code": "def main(**kwargs):\n return {'ok': True}"},
|
|
364
|
+
}
|
|
365
|
+
],
|
|
366
|
+
edges=[{"source": "py1", "target": "py1", **_NEVER_EDGE}],
|
|
367
|
+
)
|
|
368
|
+
ctx = FlowContext()
|
|
369
|
+
_run(flow.run(ctx))
|
|
370
|
+
node = next(n for n in ctx.trace.nodes if n.node_id == "py1")
|
|
371
|
+
assert node.usage is None
|
|
372
|
+
|
|
373
|
+
@patch("agstack.llm.flow.nodes.llm_chat_node.get_llm_client")
|
|
374
|
+
def test_parallel_container_owns_branch_usage(self, mock_get_client):
|
|
375
|
+
"""parallel 容器 usage = 分支用量总和,分支节点 usage 为 None"""
|
|
376
|
+
mock_client = AsyncMock()
|
|
377
|
+
mock_client.chat = AsyncMock(side_effect=[_chat_response("a", 10, 5), _chat_response("b", 6, 2)])
|
|
378
|
+
mock_get_client.return_value = mock_client
|
|
379
|
+
|
|
380
|
+
flow = Flow(
|
|
381
|
+
flow_id="t",
|
|
382
|
+
name="t",
|
|
383
|
+
nodes=[
|
|
384
|
+
{"id": "par", "type": "parallel", "config": {"branches": ["b1", "b2"]}},
|
|
385
|
+
{"id": "b1", "type": "llm_chat", "config": {"prompt": "one"}},
|
|
386
|
+
{"id": "b2", "type": "llm_chat", "config": {"prompt": "two"}},
|
|
387
|
+
],
|
|
388
|
+
edges=[{"source": "par", "target": "par", **_NEVER_EDGE}],
|
|
389
|
+
)
|
|
390
|
+
ctx = FlowContext()
|
|
391
|
+
_run(flow.run(ctx))
|
|
392
|
+
|
|
393
|
+
par = next(n for n in ctx.trace.nodes if n.node_id == "par")
|
|
394
|
+
assert par.usage is not None
|
|
395
|
+
assert par.usage.total_tokens == 23
|
|
396
|
+
for branch_id in ("b1", "b2"):
|
|
397
|
+
branch = next(n for n in ctx.trace.nodes if n.node_id == branch_id)
|
|
398
|
+
assert branch.usage is None
|
|
399
|
+
|
|
400
|
+
|
|
401
|
+
# ── D5: run() 收敛为 stream() 消费者 ──
|
|
402
|
+
|
|
403
|
+
|
|
404
|
+
class TestRunStreamConvergence:
|
|
405
|
+
def test_run_applies_retry_policy(self):
|
|
406
|
+
"""首次必失败、二次成功 + retry 配置:run() 成功返回"""
|
|
407
|
+
calls = {"n": 0}
|
|
408
|
+
|
|
409
|
+
def flaky(context, inputs):
|
|
410
|
+
calls["n"] += 1
|
|
411
|
+
if calls["n"] == 1:
|
|
412
|
+
raise ValueError("first attempt fails")
|
|
413
|
+
return {"ok": True}
|
|
414
|
+
|
|
415
|
+
registry.register_tool("d5_flaky", Tool(name="d5_flaky", description="", function=flaky))
|
|
416
|
+
flow = Flow(
|
|
417
|
+
flow_id="t",
|
|
418
|
+
name="t",
|
|
419
|
+
nodes=[
|
|
420
|
+
{
|
|
421
|
+
"id": "job",
|
|
422
|
+
"type": "tool",
|
|
423
|
+
"config": {
|
|
424
|
+
"tool_name": "d5_flaky",
|
|
425
|
+
"retry": {"max_retries": 2, "delay": 0.001, "backoff": 1.0},
|
|
426
|
+
},
|
|
427
|
+
}
|
|
428
|
+
],
|
|
429
|
+
)
|
|
430
|
+
ctx = FlowContext()
|
|
431
|
+
outputs = _run(flow.run(ctx))
|
|
432
|
+
assert calls["n"] == 2
|
|
433
|
+
assert outputs["job"] == {"ok": True}
|
|
434
|
+
|
|
435
|
+
def test_run_populates_trace(self):
|
|
436
|
+
flow = Flow(
|
|
437
|
+
flow_id="t",
|
|
438
|
+
name="t",
|
|
439
|
+
nodes=[
|
|
440
|
+
{
|
|
441
|
+
"id": "py1",
|
|
442
|
+
"type": "python",
|
|
443
|
+
"config": {"code": "def main(**kwargs):\n return {'ok': True}"},
|
|
444
|
+
}
|
|
445
|
+
],
|
|
446
|
+
edges=[{"source": "py1", "target": "py1", **_NEVER_EDGE}],
|
|
447
|
+
)
|
|
448
|
+
ctx = FlowContext()
|
|
449
|
+
_run(flow.run(ctx))
|
|
450
|
+
assert len(ctx.trace.nodes) == 1
|
|
451
|
+
assert ctx.trace.started_at is not None
|
|
452
|
+
assert ctx.trace.finished_at is not None
|
|
453
|
+
assert ctx.trace.total_usage is ctx.usage
|
|
454
|
+
|
|
455
|
+
def test_run_output_mode_append(self):
|
|
456
|
+
"""output_mode: append 的节点被访问两次后输出是长度 2 的 list"""
|
|
457
|
+
flow = Flow(
|
|
458
|
+
flow_id="t",
|
|
459
|
+
name="t",
|
|
460
|
+
nodes=[
|
|
461
|
+
{
|
|
462
|
+
"id": "gen",
|
|
463
|
+
"type": "python",
|
|
464
|
+
"config": {
|
|
465
|
+
"code": "def main(**kwargs):\n return {'tick': 1}",
|
|
466
|
+
"output_mode": "append",
|
|
467
|
+
},
|
|
468
|
+
},
|
|
469
|
+
{
|
|
470
|
+
"id": "end",
|
|
471
|
+
"type": "python",
|
|
472
|
+
"config": {"code": "def main(**kwargs):\n return {'end': True}"},
|
|
473
|
+
},
|
|
474
|
+
],
|
|
475
|
+
edges=[
|
|
476
|
+
{"source": "gen", "condition": "$v.always == yes", "target": "gen"},
|
|
477
|
+
{"source": "gen", "target": "end"},
|
|
478
|
+
],
|
|
479
|
+
cycle_limits={"gen": 2},
|
|
480
|
+
)
|
|
481
|
+
ctx = FlowContext(variables={"always": "yes"})
|
|
482
|
+
_run(flow.run(ctx))
|
|
483
|
+
assert isinstance(ctx.outputs["gen"], list)
|
|
484
|
+
assert len(ctx.outputs["gen"]) == 2
|
|
485
|
+
assert ctx.outputs["end"] == {"end": True}
|
|
486
|
+
|
|
487
|
+
def test_run_wraps_errors_as_node_execution_error(self):
|
|
488
|
+
"""异常类型收窄:run() 节点失败抛 NodeExecutionError(1.25.1 抛原始异常)"""
|
|
489
|
+
flow = Flow(
|
|
490
|
+
flow_id="t",
|
|
491
|
+
name="t",
|
|
492
|
+
nodes=[
|
|
493
|
+
{
|
|
494
|
+
"id": "bad",
|
|
495
|
+
"type": "python",
|
|
496
|
+
"config": {"code": "def main(**kwargs):\n raise RuntimeError('inner')"},
|
|
497
|
+
}
|
|
498
|
+
],
|
|
499
|
+
)
|
|
500
|
+
with pytest.raises(NodeExecutionError):
|
|
501
|
+
_run(flow.run(FlowContext()))
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|