agstack 2.0.0__tar.gz → 2.2.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {agstack-2.0.0 → agstack-2.2.0}/PKG-INFO +1 -1
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/__init__.py +3 -1
- agstack-2.2.0/agstack/llm/flow/agent.py +468 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/context.py +17 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/flow.py +30 -4
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/registry.py +14 -1
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/tool.py +90 -2
- {agstack-2.0.0 → agstack-2.2.0}/agstack.egg-info/PKG-INFO +1 -1
- {agstack-2.0.0 → agstack-2.2.0}/agstack.egg-info/SOURCES.txt +5 -1
- {agstack-2.0.0 → agstack-2.2.0}/pyproject.toml +1 -1
- agstack-2.2.0/tests/test_agent_parallel_tools.py +152 -0
- agstack-2.2.0/tests/test_agent_request_overrides.py +50 -0
- agstack-2.2.0/tests/test_flow_cancellation.py +185 -0
- agstack-2.2.0/tests/test_tool_hooks.py +218 -0
- agstack-2.0.0/agstack/llm/flow/agent.py +0 -354
- {agstack-2.0.0 → agstack-2.2.0}/LICENSE +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/README.md +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/cache/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/cache/base.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/cache/memory.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/cache/redis.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/config/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/config/logger.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/config/manager.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/config/types.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/contexts.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/decorators.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/events.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/exceptions.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/fastapi/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/fastapi/exception.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/fastapi/middleware.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/fastapi/offline.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/fastapi/sse.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/infra/db/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/infra/es/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/infra/kg/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/infra/mq/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/client.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/event.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/exceptions.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/factory.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/loader.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/agent_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/base.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/detect_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/echo_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/iterator_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/llm_chat_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/llm_embed_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/llm_rerank_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/python_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/subflow_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/switch_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/nodes/tool_node.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/records.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/sandbox.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/state.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/flow/trace.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/prompts.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/llm/token.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/messagebus/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/messagebus/base.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/messagebus/memory.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/messagebus/redis.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/schema.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/security/__init__.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/security/casbin.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/security/crypt.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack/status.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack.egg-info/dependency_links.txt +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack.egg-info/requires.txt +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/agstack.egg-info/top_level.txt +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/setup.cfg +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/tests/test_cache_memory.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/tests/test_cache_redis.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/tests/test_flow_error_semantics.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/tests/test_flow_io.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/tests/test_flow_iterator.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/tests/test_flow_switch_subflow.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/tests/test_llm_usage_callback.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/tests/test_messagebus_memory.py +0 -0
- {agstack-2.0.0 → agstack-2.2.0}/tests/test_messagebus_redis.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: agstack
|
|
3
|
-
Version: 2.
|
|
3
|
+
Version: 2.2.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",
|
|
@@ -0,0 +1,468 @@
|
|
|
1
|
+
# Copyright (c) 2020-2026 XtraVisions, All rights reserved.
|
|
2
|
+
|
|
3
|
+
"""Agent 定义和执行"""
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import json
|
|
7
|
+
from typing import TYPE_CHECKING, Any, AsyncIterator
|
|
8
|
+
from uuid import uuid4
|
|
9
|
+
|
|
10
|
+
from ..client import get_llm_client
|
|
11
|
+
from . import event
|
|
12
|
+
from .context import Usage
|
|
13
|
+
from .event import EventType
|
|
14
|
+
from .exceptions import AgentError, FlowError
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
if TYPE_CHECKING:
|
|
18
|
+
from .context import FlowContext
|
|
19
|
+
from .tool import Tool
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class Agent:
|
|
23
|
+
"""Agent 定义"""
|
|
24
|
+
|
|
25
|
+
def __init__(
|
|
26
|
+
self,
|
|
27
|
+
name: str,
|
|
28
|
+
instructions: str = "",
|
|
29
|
+
tools: list["Tool"] | None = None,
|
|
30
|
+
model: str = "gpt-4o",
|
|
31
|
+
temperature: float = 0.7,
|
|
32
|
+
max_tokens: int | None = None,
|
|
33
|
+
max_turns: int = 10,
|
|
34
|
+
*,
|
|
35
|
+
tool_choice: str = "auto",
|
|
36
|
+
on_max_turns: str = "finalize",
|
|
37
|
+
retry_empty_response: bool = False,
|
|
38
|
+
label: str | None = None,
|
|
39
|
+
echo: bool = False,
|
|
40
|
+
):
|
|
41
|
+
"""初始化 Agent
|
|
42
|
+
|
|
43
|
+
:param name: Agent 名称
|
|
44
|
+
:param instructions: 系统指令
|
|
45
|
+
:param tools: 可用工具列表
|
|
46
|
+
:param model: 模型名称
|
|
47
|
+
:param temperature: 温度参数
|
|
48
|
+
:param max_tokens: 最大 token 数
|
|
49
|
+
:param max_turns: 最大轮次
|
|
50
|
+
:param on_max_turns: max_turns 耗尽时的行为,"finalize"(降级输出并标记 truncated)或 "error"(抛出异常)
|
|
51
|
+
:param retry_empty_response: 一轮既无文字也无 tool_calls 时(如推理模型把输出预算耗尽在 reasoning 上)
|
|
52
|
+
以 ``request_overrides(..., retry=True)`` 的覆盖参数重试一次
|
|
53
|
+
:param label: 面向用户的展示名称(控制 STEP 进度事件可见性)
|
|
54
|
+
:param echo: 是否转发 TEXT_MESSAGE 给用户
|
|
55
|
+
"""
|
|
56
|
+
self.name = name
|
|
57
|
+
self.instructions = instructions or f"You are {name}, a helpful AI assistant."
|
|
58
|
+
self.tools = tools or []
|
|
59
|
+
self.model = model
|
|
60
|
+
self.temperature = temperature
|
|
61
|
+
self.max_tokens = max_tokens
|
|
62
|
+
self.max_turns = max_turns
|
|
63
|
+
self.tool_choice = tool_choice
|
|
64
|
+
self.on_max_turns = on_max_turns
|
|
65
|
+
self.retry_empty_response = retry_empty_response
|
|
66
|
+
self.label = label
|
|
67
|
+
self.echo = echo
|
|
68
|
+
|
|
69
|
+
def get_system_message(self) -> dict[str, Any]:
|
|
70
|
+
"""获取系统消息"""
|
|
71
|
+
return {"role": "system", "content": self.instructions}
|
|
72
|
+
|
|
73
|
+
def get_tools_schema(self) -> list[dict[str, Any]]:
|
|
74
|
+
"""获取工具 schema"""
|
|
75
|
+
return [tool.to_openai_tool() for tool in self.tools]
|
|
76
|
+
|
|
77
|
+
def request_overrides(self, context: "FlowContext", turn: int, *, retry: bool = False) -> dict[str, Any]:
|
|
78
|
+
"""按轮覆盖本次模型请求参数的钩子(子类实现,默认不覆盖)
|
|
79
|
+
|
|
80
|
+
返回值合并进 ``client.chat`` 的 kwargs:``extra_body`` 按键合并,其余键直接覆盖。
|
|
81
|
+
典型用法:决策轮 / 作答轮分别设置 ``extra_body={"enable_thinking": ...}`` 与 ``max_tokens``;
|
|
82
|
+
``retry=True`` 表示上一次请求空响应后的重试。
|
|
83
|
+
|
|
84
|
+
:param turn: 本 agent 本次运行内的轮次,从 1 起
|
|
85
|
+
"""
|
|
86
|
+
return {}
|
|
87
|
+
|
|
88
|
+
@staticmethod
|
|
89
|
+
def _apply_overrides(kwargs: dict[str, Any], overrides: dict[str, Any]) -> None:
|
|
90
|
+
for key, value in overrides.items():
|
|
91
|
+
if key == "extra_body" and isinstance(value, dict):
|
|
92
|
+
merged = dict(kwargs.get("extra_body") or {})
|
|
93
|
+
merged.update(value)
|
|
94
|
+
kwargs["extra_body"] = merged
|
|
95
|
+
else:
|
|
96
|
+
kwargs[key] = value
|
|
97
|
+
|
|
98
|
+
def get_tool_by_name(self, name: str) -> "Tool | None":
|
|
99
|
+
"""根据名称获取工具"""
|
|
100
|
+
for tool in self.tools:
|
|
101
|
+
if tool.name == name:
|
|
102
|
+
return tool
|
|
103
|
+
return None
|
|
104
|
+
|
|
105
|
+
def _group_tool_calls(self, tool_calls: list[dict[str, Any]]) -> list[list[dict[str, Any]]]:
|
|
106
|
+
"""按声明分组:连续的 concurrency_safe 调用聚为一组并发执行,其余单独成组串行
|
|
107
|
+
|
|
108
|
+
未注册/未声明的工具一律按不安全处理(fail closed),
|
|
109
|
+
默认全 False 时每组恰好一个调用,行为与串行完全一致。
|
|
110
|
+
"""
|
|
111
|
+
groups: list[list[dict[str, Any]]] = []
|
|
112
|
+
prev_safe = False
|
|
113
|
+
for tc in tool_calls:
|
|
114
|
+
tool = self.get_tool_by_name(tc["name"])
|
|
115
|
+
safe = bool(tool and tool.concurrency_safe)
|
|
116
|
+
if safe and prev_safe:
|
|
117
|
+
groups[-1].append(tc)
|
|
118
|
+
else:
|
|
119
|
+
groups.append([tc])
|
|
120
|
+
prev_safe = safe
|
|
121
|
+
return groups
|
|
122
|
+
|
|
123
|
+
async def _stream_tool_call(
|
|
124
|
+
self,
|
|
125
|
+
context: "FlowContext",
|
|
126
|
+
tool_call: dict[str, Any],
|
|
127
|
+
message_sink: list[dict[str, Any]] | None = None,
|
|
128
|
+
) -> AsyncIterator[dict[str, Any]]:
|
|
129
|
+
"""执行单个 tool_call,yield 其 AG-UI 事件
|
|
130
|
+
|
|
131
|
+
tool 消息默认即时写回 context;并发分组执行时传入 message_sink 收集,
|
|
132
|
+
由调用方在组完成后按 tool_call 原始顺序统一写回
|
|
133
|
+
(OpenAI 协议要求 tool 消息与 assistant.tool_calls 顺序对应)。
|
|
134
|
+
"""
|
|
135
|
+
|
|
136
|
+
def _emit_message(**kwargs: Any) -> None:
|
|
137
|
+
if message_sink is None:
|
|
138
|
+
context.add_message(self.name, "tool", **kwargs)
|
|
139
|
+
else:
|
|
140
|
+
message_sink.append(kwargs)
|
|
141
|
+
|
|
142
|
+
tool = self.get_tool_by_name(tool_call["name"])
|
|
143
|
+
if not tool:
|
|
144
|
+
error_content = json.dumps({"error": f"Tool not found: {tool_call['name']}"}, ensure_ascii=False)
|
|
145
|
+
_emit_message(content=error_content, tool_call_id=tool_call["id"])
|
|
146
|
+
# AG-UI: TOOL_CALL_RESULT (错误)
|
|
147
|
+
yield event.tool_call_result(tool_call_id=tool_call["id"], content=error_content)
|
|
148
|
+
return
|
|
149
|
+
|
|
150
|
+
# 解析 LLM 返回的工具参数;解析失败作为该次调用的失败反馈给模型,由模型自行重试
|
|
151
|
+
try:
|
|
152
|
+
tool_args = json.loads(tool_call["arguments"]) if tool_call["arguments"] else {}
|
|
153
|
+
except json.JSONDecodeError as e:
|
|
154
|
+
error_content = json.dumps(
|
|
155
|
+
{
|
|
156
|
+
"error": f"Invalid tool arguments (JSON parse failed): {e}",
|
|
157
|
+
"raw_arguments": tool_call["arguments"][:500],
|
|
158
|
+
},
|
|
159
|
+
ensure_ascii=False,
|
|
160
|
+
)
|
|
161
|
+
_emit_message(content=error_content, tool_call_id=tool_call["id"])
|
|
162
|
+
# AG-UI: TOOL_CALL_RESULT (错误)
|
|
163
|
+
yield event.tool_call_result(tool_call_id=tool_call["id"], content=error_content)
|
|
164
|
+
return
|
|
165
|
+
|
|
166
|
+
# 执行前进度事件
|
|
167
|
+
progress_label = tool.get_progress_label(tool_args)
|
|
168
|
+
if progress_label:
|
|
169
|
+
yield event.custom(
|
|
170
|
+
name="skill_progress",
|
|
171
|
+
value={
|
|
172
|
+
"progressId": tool_call["id"],
|
|
173
|
+
"description": progress_label,
|
|
174
|
+
"status": "running",
|
|
175
|
+
},
|
|
176
|
+
)
|
|
177
|
+
|
|
178
|
+
# 执行工具(传入 LLM 解析的参数作为 inputs)
|
|
179
|
+
result = await tool.execute_async(context, tool_args)
|
|
180
|
+
|
|
181
|
+
# 执行后进度事件
|
|
182
|
+
if progress_label:
|
|
183
|
+
yield event.custom(
|
|
184
|
+
name="skill_progress",
|
|
185
|
+
value={
|
|
186
|
+
"progressId": tool_call["id"],
|
|
187
|
+
"status": "completed" if result.success else "failed",
|
|
188
|
+
},
|
|
189
|
+
)
|
|
190
|
+
|
|
191
|
+
# 使用 result.content 作为 LLM 上下文(Tool 已计算好)
|
|
192
|
+
result_content = result.content or (
|
|
193
|
+
json.dumps(result.result, ensure_ascii=False)
|
|
194
|
+
if result.success
|
|
195
|
+
else json.dumps({"error": result.error}, ensure_ascii=False)
|
|
196
|
+
)
|
|
197
|
+
_emit_message(content=result_content, tool_call_id=tool_call["id"], summary=result.summary)
|
|
198
|
+
|
|
199
|
+
# AG-UI: TOOL_CALL_RESULT
|
|
200
|
+
yield event.tool_call_result(tool_call_id=tool_call["id"], content=result_content)
|
|
201
|
+
|
|
202
|
+
# 实时用户进度 — 有 summary 时告知前端
|
|
203
|
+
if result.summary:
|
|
204
|
+
yield event.custom(
|
|
205
|
+
name="tool_progress",
|
|
206
|
+
value={
|
|
207
|
+
"tool_call_id": tool_call["id"],
|
|
208
|
+
"tool_name": result.name,
|
|
209
|
+
"success": result.success,
|
|
210
|
+
"summary": result.summary,
|
|
211
|
+
},
|
|
212
|
+
)
|
|
213
|
+
|
|
214
|
+
# 业务自定义事件 — flush pending
|
|
215
|
+
for pending_evt in context.pop_pending_custom_events():
|
|
216
|
+
yield pending_evt
|
|
217
|
+
|
|
218
|
+
async def _gather_tool_calls(
|
|
219
|
+
self, context: "FlowContext", group: list[dict[str, Any]]
|
|
220
|
+
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
|
221
|
+
"""并发执行一组 concurrency_safe 的 tool_calls
|
|
222
|
+
|
|
223
|
+
返回 (按完成顺序展平的事件列表, 按原始顺序展平的 tool 消息 kwargs 列表)。
|
|
224
|
+
组内事件先缓冲、gather 结束后统一交调用方 yield,保证单 generator 语义;
|
|
225
|
+
单个调用的意外异常兜底转失败结果,不影响组内其它调用。
|
|
226
|
+
注意:并发期间 execution_records / pending_custom_events 的追加顺序不再确定。
|
|
227
|
+
"""
|
|
228
|
+
event_buffers: list[list[dict[str, Any]]] = [[] for _ in group]
|
|
229
|
+
message_buffers: list[list[dict[str, Any]]] = [[] for _ in group]
|
|
230
|
+
finish_order: list[int] = []
|
|
231
|
+
|
|
232
|
+
async def _run_one(i: int, tc: dict[str, Any]) -> None:
|
|
233
|
+
try:
|
|
234
|
+
async for evt in self._stream_tool_call(context, tc, message_sink=message_buffers[i]):
|
|
235
|
+
event_buffers[i].append(evt)
|
|
236
|
+
except Exception as e:
|
|
237
|
+
error_content = json.dumps({"error": str(e)}, ensure_ascii=False)
|
|
238
|
+
message_buffers[i].append({"content": error_content, "tool_call_id": tc["id"]})
|
|
239
|
+
event_buffers[i].append(event.tool_call_result(tool_call_id=tc["id"], content=error_content))
|
|
240
|
+
finally:
|
|
241
|
+
finish_order.append(i)
|
|
242
|
+
|
|
243
|
+
await asyncio.gather(*[_run_one(i, tc) for i, tc in enumerate(group)])
|
|
244
|
+
events = [evt for i in finish_order for evt in event_buffers[i]]
|
|
245
|
+
messages = [msg for msgs in message_buffers for msg in msgs]
|
|
246
|
+
return events, messages
|
|
247
|
+
|
|
248
|
+
async def run(self, context: "FlowContext", inputs: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
249
|
+
"""执行 Agent 逻辑"""
|
|
250
|
+
content_parts = []
|
|
251
|
+
async for evt in self.stream(context, inputs):
|
|
252
|
+
# AG-UI 事件格式
|
|
253
|
+
if isinstance(evt, dict):
|
|
254
|
+
if evt.get("type") == EventType.TEXT_MESSAGE_CONTENT:
|
|
255
|
+
content_parts.append(evt.get("delta", ""))
|
|
256
|
+
elif evt.get("type") == EventType.RUN_ERROR:
|
|
257
|
+
raise FlowError("AGENT_EXECUTION_FAILED", 500, {"error": evt.get("message")})
|
|
258
|
+
return {"result": "".join(content_parts)}
|
|
259
|
+
|
|
260
|
+
async def stream(
|
|
261
|
+
self, context: "FlowContext", inputs: dict[str, Any] | None = None
|
|
262
|
+
) -> AsyncIterator[dict[str, Any]]:
|
|
263
|
+
"""流式执行 Agent,输出 AG-UI 标准事件"""
|
|
264
|
+
|
|
265
|
+
# 注入 agent_call_id 供 tool 审计关联
|
|
266
|
+
agent_call_id = str(uuid4())
|
|
267
|
+
context.set_variable("_agent_call_id", agent_call_id)
|
|
268
|
+
|
|
269
|
+
# 输入来源:优先 inputs 参数,回退到 context.variables
|
|
270
|
+
user_input = ""
|
|
271
|
+
if inputs:
|
|
272
|
+
user_input = inputs.get("input", "")
|
|
273
|
+
if not user_input:
|
|
274
|
+
user_input = context.get_variable("input") or context.get_variable("query", "")
|
|
275
|
+
msg_id = context.message_id or str(uuid4())
|
|
276
|
+
|
|
277
|
+
# 添加用户消息(scoped by agent name)
|
|
278
|
+
context.add_message(self.name, "user", user_input)
|
|
279
|
+
context.last_agent = self.name
|
|
280
|
+
|
|
281
|
+
# AG-UI: TEXT_MESSAGE_START
|
|
282
|
+
yield event.text_message_start(message_id=msg_id, role="assistant")
|
|
283
|
+
|
|
284
|
+
# 构建消息列表:system + 共享历史 + 当前 agent 的隔离消息
|
|
285
|
+
messages = [self.get_system_message()] + context.history + context.get_messages(self.name)
|
|
286
|
+
tools_schema = self.get_tools_schema() if self.tools else None
|
|
287
|
+
|
|
288
|
+
# 获取 LLM 客户端
|
|
289
|
+
client = get_llm_client()
|
|
290
|
+
|
|
291
|
+
# Agent 循环
|
|
292
|
+
assistant_content = ""
|
|
293
|
+
for turn in range(1, self.max_turns + 1):
|
|
294
|
+
# 协作式取消检查点:不再开始新的 LLM 轮次
|
|
295
|
+
if context.is_cancelled:
|
|
296
|
+
if not context.get_variable("_cancel_emitted"):
|
|
297
|
+
context.set_variable("_cancel_emitted", True)
|
|
298
|
+
yield event.run_error(message="FLOW_CANCELLED", code="CANCELLED")
|
|
299
|
+
return
|
|
300
|
+
|
|
301
|
+
context.increment_turn()
|
|
302
|
+
|
|
303
|
+
# 调用模型;空响应(无文字无 tool_calls)且开启 retry_empty_response 时以重试覆盖参数再请求一次
|
|
304
|
+
attempt = 0
|
|
305
|
+
while True:
|
|
306
|
+
assistant_content = ""
|
|
307
|
+
tool_calls: list[dict[str, Any]] = []
|
|
308
|
+
tool_calls_buffer: dict[int, dict[str, Any]] = {}
|
|
309
|
+
|
|
310
|
+
try:
|
|
311
|
+
kwargs: dict[str, Any] = {
|
|
312
|
+
"messages": messages,
|
|
313
|
+
"model": self.model,
|
|
314
|
+
"temperature": self.temperature,
|
|
315
|
+
}
|
|
316
|
+
|
|
317
|
+
if self.max_tokens:
|
|
318
|
+
kwargs["max_tokens"] = self.max_tokens
|
|
319
|
+
|
|
320
|
+
if tools_schema:
|
|
321
|
+
kwargs["tools"] = tools_schema
|
|
322
|
+
kwargs["tool_choice"] = self.tool_choice
|
|
323
|
+
|
|
324
|
+
self._apply_overrides(kwargs, self.request_overrides(context, turn, retry=attempt > 0) or {})
|
|
325
|
+
|
|
326
|
+
stream = await client.chat(stream=True, **kwargs)
|
|
327
|
+
|
|
328
|
+
async for chunk in stream:
|
|
329
|
+
if not chunk.choices:
|
|
330
|
+
continue
|
|
331
|
+
|
|
332
|
+
choice = chunk.choices[0]
|
|
333
|
+
delta = choice.delta
|
|
334
|
+
|
|
335
|
+
# 内容增量 - AG-UI: TEXT_MESSAGE_CONTENT
|
|
336
|
+
if delta.content:
|
|
337
|
+
assistant_content += delta.content
|
|
338
|
+
yield event.text_message_content(
|
|
339
|
+
message_id=msg_id,
|
|
340
|
+
delta=delta.content,
|
|
341
|
+
)
|
|
342
|
+
|
|
343
|
+
# 工具调用
|
|
344
|
+
if delta.tool_calls:
|
|
345
|
+
for tool_call_delta in delta.tool_calls:
|
|
346
|
+
idx = tool_call_delta.index # noqa
|
|
347
|
+
if idx not in tool_calls_buffer:
|
|
348
|
+
tool_calls_buffer[idx] = {
|
|
349
|
+
"id": tool_call_delta.id or "", # noqa
|
|
350
|
+
"name": "",
|
|
351
|
+
"arguments": "",
|
|
352
|
+
}
|
|
353
|
+
|
|
354
|
+
if tool_call_delta.id: # noqa
|
|
355
|
+
tool_calls_buffer[idx]["id"] = tool_call_delta.id # noqa
|
|
356
|
+
if tool_call_delta.function and tool_call_delta.function.name: # noqa
|
|
357
|
+
tool_calls_buffer[idx]["name"] = tool_call_delta.function.name # noqa
|
|
358
|
+
if tool_call_delta.function and tool_call_delta.function.arguments: # noqa
|
|
359
|
+
tool_calls_buffer[idx]["arguments"] += tool_call_delta.function.arguments # noqa
|
|
360
|
+
|
|
361
|
+
# 完成
|
|
362
|
+
if choice.finish_reason:
|
|
363
|
+
# AG-UI: 工具调用事件
|
|
364
|
+
for tool_call_data in tool_calls_buffer.values():
|
|
365
|
+
tool_calls.append(tool_call_data)
|
|
366
|
+
|
|
367
|
+
# TOOL_CALL_START
|
|
368
|
+
yield event.tool_call_start(
|
|
369
|
+
tool_call_id=tool_call_data["id"],
|
|
370
|
+
tool_call_name=tool_call_data["name"],
|
|
371
|
+
)
|
|
372
|
+
|
|
373
|
+
# TOOL_CALL_ARGS
|
|
374
|
+
yield event.tool_call_args(
|
|
375
|
+
tool_call_id=tool_call_data["id"],
|
|
376
|
+
delta=tool_call_data["arguments"],
|
|
377
|
+
)
|
|
378
|
+
|
|
379
|
+
# TOOL_CALL_END
|
|
380
|
+
yield event.tool_call_end(tool_call_id=tool_call_data["id"])
|
|
381
|
+
|
|
382
|
+
# 更新 usage
|
|
383
|
+
if hasattr(chunk, "usage") and chunk.usage:
|
|
384
|
+
context.add_usage(
|
|
385
|
+
Usage(
|
|
386
|
+
prompt_tokens=chunk.usage.prompt_tokens or 0,
|
|
387
|
+
completion_tokens=chunk.usage.completion_tokens or 0,
|
|
388
|
+
total_tokens=chunk.usage.total_tokens or 0,
|
|
389
|
+
)
|
|
390
|
+
)
|
|
391
|
+
|
|
392
|
+
except Exception as e:
|
|
393
|
+
error_msg = str(e)
|
|
394
|
+
# AG-UI: RUN_ERROR
|
|
395
|
+
yield event.run_error(message=error_msg)
|
|
396
|
+
raise FlowError("AGENT_EXECUTION_FAILED", 500, {"error": error_msg}) from e
|
|
397
|
+
|
|
398
|
+
if self.retry_empty_response and attempt == 0 and not tool_calls and not assistant_content.strip():
|
|
399
|
+
attempt = 1
|
|
400
|
+
continue
|
|
401
|
+
break
|
|
402
|
+
|
|
403
|
+
# 保存 assistant 消息(tool_calls 转为 OpenAI 标准格式)
|
|
404
|
+
if tool_calls:
|
|
405
|
+
openai_tool_calls = [
|
|
406
|
+
{
|
|
407
|
+
"id": tc["id"],
|
|
408
|
+
"type": "function",
|
|
409
|
+
"function": {"name": tc["name"], "arguments": tc["arguments"]},
|
|
410
|
+
}
|
|
411
|
+
for tc in tool_calls
|
|
412
|
+
]
|
|
413
|
+
context.add_message(
|
|
414
|
+
self.name,
|
|
415
|
+
"assistant",
|
|
416
|
+
content=assistant_content or None,
|
|
417
|
+
tool_calls=openai_tool_calls,
|
|
418
|
+
)
|
|
419
|
+
else:
|
|
420
|
+
context.add_message(self.name, "assistant", assistant_content)
|
|
421
|
+
|
|
422
|
+
# 如果没有工具调用,结束循环
|
|
423
|
+
if not tool_calls:
|
|
424
|
+
# 存储结果供 Flow/A2A 使用
|
|
425
|
+
context.set_output(self.name, {"result": assistant_content})
|
|
426
|
+
# AG-UI: TEXT_MESSAGE_END
|
|
427
|
+
yield event.text_message_end(message_id=msg_id)
|
|
428
|
+
context.set_variable("_agent_call_id", None)
|
|
429
|
+
return
|
|
430
|
+
|
|
431
|
+
# 执行工具调用:连续的 concurrency_safe 调用聚组并发,其余保持串行(默认全串行)
|
|
432
|
+
for group in self._group_tool_calls(tool_calls):
|
|
433
|
+
# 协作式取消检查点:不再开始新的工具执行(不中断在途工具)
|
|
434
|
+
if context.is_cancelled:
|
|
435
|
+
if not context.get_variable("_cancel_emitted"):
|
|
436
|
+
context.set_variable("_cancel_emitted", True)
|
|
437
|
+
yield event.run_error(message="FLOW_CANCELLED", code="CANCELLED")
|
|
438
|
+
return
|
|
439
|
+
|
|
440
|
+
if len(group) == 1:
|
|
441
|
+
# 串行路径:事件实时 yield,tool 消息即时写回
|
|
442
|
+
async for evt in self._stream_tool_call(context, group[0]):
|
|
443
|
+
yield evt
|
|
444
|
+
else:
|
|
445
|
+
# 并发组:事件按完成顺序 yield,tool 消息按原始顺序写回
|
|
446
|
+
group_events, group_messages = await self._gather_tool_calls(context, group)
|
|
447
|
+
for evt in group_events:
|
|
448
|
+
yield evt
|
|
449
|
+
for msg in group_messages:
|
|
450
|
+
context.add_message(self.name, "tool", **msg)
|
|
451
|
+
|
|
452
|
+
# 更新消息列表,继续下一轮
|
|
453
|
+
messages = [self.get_system_message()] + context.history + context.get_messages(self.name)
|
|
454
|
+
|
|
455
|
+
# max_turns 耗尽:必须显式收尾,禁止静默截断
|
|
456
|
+
if self.on_max_turns == "error":
|
|
457
|
+
error_msg = f"Agent {self.name} exceeded max_turns={self.max_turns}"
|
|
458
|
+
yield event.run_error(message=error_msg, code="AGENT_MAX_TURNS_EXCEEDED")
|
|
459
|
+
raise AgentError("AGENT_MAX_TURNS_EXCEEDED", 500, {"agent": self.name})
|
|
460
|
+
|
|
461
|
+
# finalize:最后一轮已生成的部分文本作为降级输出,带截断标记
|
|
462
|
+
yield event.custom(
|
|
463
|
+
name="agent_max_turns",
|
|
464
|
+
value={"agentName": self.name, "maxTurns": self.max_turns},
|
|
465
|
+
)
|
|
466
|
+
context.set_output(self.name, {"result": assistant_content, "truncated": True})
|
|
467
|
+
yield event.text_message_end(message_id=msg_id)
|
|
468
|
+
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
|
|
@@ -84,6 +88,19 @@ class FlowContext:
|
|
|
84
88
|
"""取出并移除变量"""
|
|
85
89
|
return self.variables.pop(key, default)
|
|
86
90
|
|
|
91
|
+
def cancel(self) -> None:
|
|
92
|
+
"""请求取消执行(幂等)
|
|
93
|
+
|
|
94
|
+
协作式取消:引擎在下一个检查点(节点执行前、agent 轮次开始、
|
|
95
|
+
tool_call 执行前)停止,不强杀在途的工具或 LLM 调用。
|
|
96
|
+
"""
|
|
97
|
+
self._cancel_event.set()
|
|
98
|
+
|
|
99
|
+
@property
|
|
100
|
+
def is_cancelled(self) -> bool:
|
|
101
|
+
"""是否已请求取消(长 I/O 工具可自查以提前返回)"""
|
|
102
|
+
return self._cancel_event.is_set()
|
|
103
|
+
|
|
87
104
|
def update_variables(self, updates: dict[str, Any]) -> None:
|
|
88
105
|
"""批量更新变量"""
|
|
89
106
|
self.variables.update(updates)
|
|
@@ -10,7 +10,7 @@ from uuid import uuid4
|
|
|
10
10
|
|
|
11
11
|
from . import event
|
|
12
12
|
from .context import Usage
|
|
13
|
-
from .exceptions import NodeExecutionError
|
|
13
|
+
from .exceptions import FlowExecutionError, NodeExecutionError
|
|
14
14
|
|
|
15
15
|
|
|
16
16
|
if TYPE_CHECKING:
|
|
@@ -234,6 +234,12 @@ class Flow:
|
|
|
234
234
|
last_error: Exception | None = None
|
|
235
235
|
|
|
236
236
|
for attempt in range(policy.max_retries + 1):
|
|
237
|
+
# 协作式取消检查点:取消后不再开始新的重试
|
|
238
|
+
if attempt > 0 and context.is_cancelled:
|
|
239
|
+
if not context.get_variable("_cancel_emitted"):
|
|
240
|
+
context.set_variable("_cancel_emitted", True)
|
|
241
|
+
yield event.run_error(message="FLOW_CANCELLED", code="CANCELLED")
|
|
242
|
+
return
|
|
237
243
|
try:
|
|
238
244
|
if attempt > 0:
|
|
239
245
|
wait = policy.delay * (policy.backoff ** (attempt - 1))
|
|
@@ -284,10 +290,12 @@ class Flow:
|
|
|
284
290
|
|
|
285
291
|
stream() 的消费包装:两条路径共享同一执行引擎,重试策略、
|
|
286
292
|
FlowTrace、output_mode、iterator 状态清理等行为完全一致。
|
|
287
|
-
节点失败抛 NodeExecutionError
|
|
293
|
+
节点失败抛 NodeExecutionError(包装原始异常);
|
|
294
|
+
取消(context.cancel())抛 FlowExecutionError("FLOW_CANCELLED")。
|
|
288
295
|
"""
|
|
289
|
-
async for
|
|
290
|
-
|
|
296
|
+
async for evt in self.stream(context):
|
|
297
|
+
if evt.get("type") == event.EventType.RUN_ERROR and evt.get("code") == "CANCELLED":
|
|
298
|
+
raise FlowExecutionError("FLOW_CANCELLED", args={"flow": self.name})
|
|
291
299
|
return context.outputs
|
|
292
300
|
|
|
293
301
|
async def stream(self, context: "FlowContext") -> AsyncIterator[dict[str, Any]]:
|
|
@@ -310,11 +318,22 @@ class Flow:
|
|
|
310
318
|
context.trace.finished_at = time.time()
|
|
311
319
|
context.trace.total_usage = context.usage
|
|
312
320
|
|
|
321
|
+
# 取消停止:RUN_ERROR(code=CANCELLED) 是终止事件,事件流以其结束
|
|
322
|
+
if context.get_variable("_cancel_emitted"):
|
|
323
|
+
return
|
|
324
|
+
|
|
313
325
|
yield event.step_finished(step_name=f"flow:{self.name}", step_id=flow_sid)
|
|
314
326
|
|
|
315
327
|
async def _stream_sequential(self, context: "FlowContext") -> AsyncIterator[dict[str, Any]]:
|
|
316
328
|
"""顺序流式执行"""
|
|
317
329
|
for node in self.nodes:
|
|
330
|
+
# 协作式取消检查点:不再开始新的节点执行
|
|
331
|
+
if context.is_cancelled:
|
|
332
|
+
if not context.get_variable("_cancel_emitted"):
|
|
333
|
+
context.set_variable("_cancel_emitted", True)
|
|
334
|
+
yield event.run_error(message="FLOW_CANCELLED", code="CANCELLED")
|
|
335
|
+
return
|
|
336
|
+
|
|
318
337
|
node_id = node.get("id")
|
|
319
338
|
if not node_id:
|
|
320
339
|
continue
|
|
@@ -350,6 +369,13 @@ class Flow:
|
|
|
350
369
|
visit_count: dict[str, int] = {}
|
|
351
370
|
|
|
352
371
|
while current_node_id:
|
|
372
|
+
# 协作式取消检查点:不再开始新的节点执行(粒度是节点边界,不中断在途节点)
|
|
373
|
+
if context.is_cancelled:
|
|
374
|
+
if not context.get_variable("_cancel_emitted"):
|
|
375
|
+
context.set_variable("_cancel_emitted", True)
|
|
376
|
+
yield event.run_error(message="FLOW_CANCELLED", code="CANCELLED")
|
|
377
|
+
return
|
|
378
|
+
|
|
353
379
|
node = self.get_node_config(current_node_id)
|
|
354
380
|
if not node:
|
|
355
381
|
yield event.run_error(
|
|
@@ -8,7 +8,7 @@ import copy
|
|
|
8
8
|
from typing import Any, cast
|
|
9
9
|
|
|
10
10
|
from .agent import Agent
|
|
11
|
-
from .tool import Tool
|
|
11
|
+
from .tool import Tool, ToolHook, clear_tool_hooks, register_tool_hook
|
|
12
12
|
|
|
13
13
|
|
|
14
14
|
class FlowRegistry:
|
|
@@ -45,6 +45,19 @@ class FlowRegistry:
|
|
|
45
45
|
if echo:
|
|
46
46
|
self._tool_echo[name] = True
|
|
47
47
|
|
|
48
|
+
def register_tool_hook(self, hook: ToolHook, *, prepend: bool = False) -> None:
|
|
49
|
+
"""注册全局工具执行钩子,对所有 Tool.execute_async 生效
|
|
50
|
+
|
|
51
|
+
pre_execute 按注册顺序、post_execute 按逆序执行;
|
|
52
|
+
prepend=True 抢占链头(其 post_execute 成为最外层,适合 spill/截断类钩子)。
|
|
53
|
+
钩子链存于 tool 模块(保持 registry→tool 单向导入),此处仅转发注册。
|
|
54
|
+
"""
|
|
55
|
+
register_tool_hook(hook, prepend=prepend)
|
|
56
|
+
|
|
57
|
+
def clear_tool_hooks(self) -> None:
|
|
58
|
+
"""清空全局工具钩子(测试隔离用)"""
|
|
59
|
+
clear_tool_hooks()
|
|
60
|
+
|
|
48
61
|
def register_agent(
|
|
49
62
|
self, name: str, agent_class: type[Agent], *, label: str | None = None, echo: bool = False
|
|
50
63
|
) -> None:
|