From 46de11196595cb2be452a91463cca414a0c231d8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=BE=90=E6=B6=9B?= Date: Wed, 15 Jul 2026 11:15:58 +0800 Subject: [PATCH] =?UTF-8?q?feat(engine):=20=E5=AE=9E=E7=8E=B0=20Agent=20?= =?UTF-8?q?=E8=A7=92=E8=89=B2=E7=83=AD=E5=88=87=E6=8D=A2=E4=B8=8E=E5=AD=90?= =?UTF-8?q?=E4=BB=A3=E7=90=86=E8=B0=83=E5=BA=A6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 新增 SessionManager::switch_agent / dispatch / dispatch_all / dispatch_stream 四个核心方法,补齐多 Agent 基础系统原语。 交付物: - switch_agent 运行时替换 Arc,保留上下文并更新 SessionMeta - dispatch 单任务派发:create_child → inherit_memory → submit_turn - dispatch_all 并行派发:Semaphore 并发控制 + 部分成功语义 - dispatch_stream 流式派发:unbounded_channel + spawn task 消息重建 + finalize_turn - DispatchConfig / SubTaskResult / SubTaskStreamEvent 公开类型 - 4 个端到端示例(agent_switch_demo / sub_agent_dispatch_demo / bridge_keys_demo / dispatch_stream_demo) 辅助变更: - EngineError 新增 DispatchFailed 变体 - CostTracker 实现 From 转换 - save_session_meta / load_session_meta 改 pub(crate) 供 switch.rs 使用 测试: +17 个内联测试(4 switch + 5 dispatch + 4 dispatch_all + 4 dispatch_stream),全量 391 passed / 0 failed, clippy 0 警告,doc 0 warning。 --- docs/24-phase18-agent-switch-and-dispatch.md | 700 ++++++++++++ examples/agent_switch_demo.rs | 115 ++ examples/bridge_keys_demo.rs | 197 ++++ examples/dispatch_stream_demo.rs | 121 ++ examples/sub_agent_dispatch_demo.rs | 141 +++ src/engine/error.rs | 5 + src/engine/mod.rs | 5 +- src/engine/session_manager.rs | 4 +- src/engine/sub_agent.rs | 1071 ++++++++++++++++++ src/engine/switch.rs | 222 ++++ src/llm/types/usage.rs | 8 + 11 files changed, 2586 insertions(+), 3 deletions(-) create mode 100644 docs/24-phase18-agent-switch-and-dispatch.md create mode 100644 examples/agent_switch_demo.rs create mode 100644 examples/bridge_keys_demo.rs create mode 100644 examples/dispatch_stream_demo.rs create mode 100644 examples/sub_agent_dispatch_demo.rs create mode 100644 src/engine/sub_agent.rs create mode 100644 src/engine/switch.rs diff --git a/docs/24-phase18-agent-switch-and-dispatch.md b/docs/24-phase18-agent-switch-and-dispatch.md new file mode 100644 index 0000000..2e06ac5 --- /dev/null +++ b/docs/24-phase18-agent-switch-and-dispatch.md @@ -0,0 +1,700 @@ +# Phase 18:Agent 角色热切换与子代理调度 + +## 背景与目标 + +### 问题空间 + +agcore 已完整交付 Phase 0-17,具备 SessionManager 会话生命周期管理、会话树(父子层次)、Checkpointer 检查点、流式输出、ContextSlot 上下文分区、MemoryStore 持久化、摘要自动生成等能力。当前 Session 在 `create()` 时绑定一个 `Arc`,此后无法变更角色;会话间调度仅通过 `create_child()` + `submit_turn()` 手动编排,缺乏内建的子代理派发机制。 + +Phase 18 要解决两个正交但关联的问题: + +1. **Agent 角色热切换**:运行时替换 session 绑定的 Agent,保留上下文(slot 历史、turn_index、session_memory、cost_so_far) +2. **子代理调度**:在 SessionManager 上提供声明式的 `dispatch` / `dispatch_stream` / `dispatch_all` API,支持父子 session 间的 Memory 继承、bridge_keys 注入、并发控制、结构化回传 + +### 目标 + +- 提供 `SessionManager::switch_agent(session_id, new_agent)`,替换 `Arc`,全量保留 slot / turn_index / session_memory +- 提供 `DispatchConfig` / `SubTaskResult` / `SubTaskStreamEvent` 类型以及 `dispatch` / `dispatch_stream` / `dispatch_all` 三个核心方法 +- 实现父转子三层级交互:父->子(Memory 快照继承 + bridge_keys)、子->父(SubTaskResult 结构化回传 + result_summary)、子<->子(shared namespace) +- 产出 4 个端到端示例:`agent_switch_demo` / `sub_agent_dispatch_demo` / `bridge_keys_demo` / `dispatch_stream_demo` + +### 依赖与优先级 + +- **依赖**:Phase 17(SessionManager + 会话树 + Checkpointer)[高] +- **优先级**:P0 +- **预估规模**:约 720 行核心 + 210 行测试 + 450 行示例 +- **审查修复**:第 1 轮审查修复(Finalize 方案重写 + 4 个 🔴 阻塞 + 8 个 🟡 改进) + +--- + +## 当前状态分析 + +### 现有架构中的关键接入点 + +| 接入点 | 位置 | 可用性 | 分析 | +|--------|------|--------|------| +| `AgentSession.agent` | `src/agent/session.rs:53` | `pub` 字段 | 可直接替换,无需新增 setter [高] | +| `AgentSession.session_id` | `src/agent/session.rs:51` | `pub` 字段 | 子 session 创建后可读取 [高] | +| `AgentSession.session_memory` | `src/agent/session.rs:58` | `pub` 字段 | dispatch 后父可读子 memory [高] | +| `SessionManager::create_child(parent_id, agent)` | `src/engine/session_manager.rs:202-239` | `pub async` | dispatch 可直接复用,bundle 继承避免重复构造 [高] | +| `SessionMemory::list_entries()` | `src/agent/session_memory.rs:81-110` | `pub async` | 可获取全量条目用于父子继承 / bridge_keys 过滤 [高] | +| `SessionMemory::set_with_meta()` | `src/agent/session_memory.rs:61-79` | `pub async` | 子 session 写入继承数据时保留原始 metadata [高] | +| `CostTracker` | `src/llm/cycle.rs` | derive `Clone` | SubTaskResult 可直接 clone usage [高] | +| `SessionMeta` 持久化 | `src/engine/session_manager.rs:32-67` | `pub(crate)` | 以 `session:{id}:meta` key 存到 MemoryStore;switch 后需更新 agent_name [高] | +| `SessionManager::save_session_meta()` | `src/engine/session_manager.rs:131-142` | `async fn` (非 pub) | switch_agent 需要类似的 meta 更新能力;考虑提取为 `pub(crate)` [高] | +| `SessionManager::destroy()` | `src/engine/session_manager.rs:495-512` | `pub async` | dispatch 失败时清理子 session 可直接复用 [高] | +| `SessionManager.sessions` | `src/engine/session_manager.rs:92` | `pub(crate)` RwLock | switch_agent 和 dispatch 的 get / replace 操作均依赖此字段 [高] | +| `EngineError` 枚举 | `src/engine/error.rs:16-46` | `#[non_exhaustive]` | 已有 6 个变体,需追加 DispatchFailed / SwitchFailed / SubAgentStreamError [中] | +| `futures-core` / `futures-util` | `Cargo.toml:17-18` | 已引入 | dispatch_stream 返回 `Pin>` 所需依赖已就绪,无需新增 [高] | + +### Agent trait 与 SessionManager 之间的关系 + +``` +AgentSession { + agent: Arc, // 可替换 + session_memory: SessionMemory, // 可继承(clone backend) + slots: HashMap, + turn_index: u32, + cost_so_far: CostTracker, + // ... 其余内部字段 +} + +SessionManager { + sessions: RwLock>>>, + checkpointer: Checkpointer, + store: Arc, + config: SessionManagerConfig, +} +``` + +`AgentSession.agent` 是 `pub` 字段,这意味着 `switch_agent` 只需 `get()` → `lock()` → 替换 `agent` → 写回 meta。Route 明确、无架构阻力 [高]。 + +### 锁契约(需格外注意) + +`src/engine/session_manager.rs:4-12` 记录了锁契约:**不持有 RwLock 跨越 `.await`**。所有 `.await` 点必须在 RwLock guard drop 之后。这意味着: + +- `switch_agent`:读锁 `get()` 返回 `Arc>` 后释放,然后 lock session 级别的 Mutex → 替换 agent → 释放 Mutex → save_session_meta(I/O)[高] +- `dispatch`:读锁 `get()` 父 session → 释放 → `create_child`(内部写锁)→ lock 子 session → inherit memory → submit_turn → 释放 [高] +- 不会引入新的死锁风险 [高] + +### 现有测试覆盖 + +SessionManager 已有 906 行(含 12 个测试),覆盖 create/get/destroy/create_child/replace/recover/children/parent/并发创建等场景:`src/engine/session_manager.rs:527-906`。Phase 18 新增测试不修改这些已有测试。 + +--- + +## 调研发现 + +### 1. bridge_keys 注入位置 + +**问题**:bridge_keys 本质是父 session 想注入到子 agent prompt 中的上下文数据。它应该放在哪里? + +**调研来源**: +- `docs/note-opencode-subagent-dispatch.md` — 明确反对修改 Agent 的 `system_prompt()` [高] +- `src/agent/agent.rs:21` — `system_prompt()` 返回 `&str`,无状态变更能力 [高] +- `src/agent/session_memory.rs` — `SessionMemory::set()` 提供 key-value 写入,子 agent 可读 [高] + +**结论**:bridge_keys 通过 SessionMemory 副本继承 + 过滤注入,不碰 `system_prompt()`。子 agent 通过 `get_session_data(key)` 读取桥接数据 [高]。 + +### 2. SessionMemory 继承策略 + +**问题**:子 agent 启动时,父 session 的 SessionMemory 如何传递? + +**方案 A —— 引用共享**:父子共享同一 `SessionMemory` 实例(Arc clone 后端)。优点是零拷贝,缺点是父子隔离被破坏 [中]。 + +**方案 B —— 快照副本**:父调用 `list_entries()` 获取全量条目,子通过 `set_with_meta()` 写入自己的 namespace。优点是隔离性强,缺点是 O(n) 拷贝开销 [高]。 + +**来源**: +- `src/agent/session.rs:503-524` — `to_snapshot()` 已实现类似的 list_entries → HashMap 拍平 [高] +- `src/agent/session_memory.rs:61-79` — `set_with_meta()` 可保留原始 metadata [高] +- `docs/note-opencode-subagent-dispatch.md` — SA 建议副本策略 [中] + +**结论**:采用方案 B(快照副本),隔离性优先。`inherit_session_memory` 内部使用 `list_entries()` → 按 `bridge_keys` 过滤 → `set_with_meta()` 写入子 namespace [高]。 + +### 3. dispatch_all 部分成功语义 + +**问题**:当一批子代理中部分失败时,dispatch_all 应该整体失败还是返回部分成功的 `Vec`? + +**来源**:`docs/note-opencode-subagent-dispatch.md` — PM 和 SA 一致认为应返回部分成功语义 [高]。 + +**结论**:返回 `Vec>`。调用方可迭代检查每个结果,失败条目保留 checkpoint 以便审计 [高]。 + +### 4. dispatch_stream 生命周期 + +**问题**:`dispatch_stream` 需要返回一个流,流内部要做 `submit_turn_stream` + `finalize_active_stream`。session 所有权和生命周期如何管理? + +**来源**: +- `src/engine/session_manager.rs:390-412` — `submit_turn_stream` 的锁模式:短持锁获取流后立即释放 [高] +- `src/agent/session.rs:384-445` — `submit_turn_stream` 自身不持有跨 await 的锁 [高] +- `docs/note-opencode-subagent-dispatch.md` — SA 建议在 AgentSession 新增 `finalize_active_stream()` 内部方法 [中] + +**结论**:`dispatch_stream` 使用 `&Arc` 签名 + `tokio::spawn`。内部流管道:create_child → inherit_memory → submit_turn_stream → mpsc channel 转发事件 → 流消费完毕后调用 `finalize_active_stream()`。session 通过 `Arc>` 在 spawned task 中持有 [高]。 + +### 5. Cargo.toml 依赖分析 + +**来源**: +- `Cargo.toml:17` — `futures-util = "0.3"` 已在依赖中 [高] +- `Cargo.toml:15` — `tokio-stream = "0.1"` 已在依赖中 [高] +- `Cargo.toml:16` — `futures = "0.3"` 已在依赖中 [高] + +**结论**:dispatch_stream 所需的 `StreamExt` / `ReceiverStream` 所需的基础设施已全部就绪,无需新增任何依赖 [高]。 + +--- + +## 可选方案 + +### 方案 A:switch_agent 作为 AgentSession 方法 vs SessionManager 方法 + +| 维度 | A1: AgentSession 方法 | A2: SessionManager 方法 | +|------|-----------------------|------------------------| +| 实现位置 | `agent/session.rs` | `engine/switch.rs` | +| 职责归属 | session 实例级 | 管理器级 | +| 能否更新 SessionMeta | 不能(无 store 引用) | 能(有 store + checkpointer) | +| 能否做自动 checkpoint | 不能(无 checkpointer) | 能 | +| 与 create_child / replace 对齐 | 不对齐(create 在 SM) | 对齐(都在 SM) | + +**来源**: +- `src/agent/session.rs:49-71` — AgentSession 不持有 store / checkpointer 引用 [高] +- `src/engine/session_manager.rs:325-347` — `replace()` 是 SM 方法,涉及 meta 持久化 [高] +- roadmap lines 838 — `switch_agent(session_id, new_agent)` 签名暗示 SM 方法 [中] + +**结论**:采用 A2(SessionManager 方法)。AgentSession 没有 store 引用,无法更新 SessionMeta。独立文件 `engine/switch.rs` 作为 SessionManager 的 impl 块。 + +### 方案 B:bridge_keys 注入方式 + +| 维度 | B1: SessionMemory 副本继承 + 过滤 | B2: 修改 Agent trait | +|------|------------------------------------|----------------------| +| 系统 prompt 侵入性 | 无 | 需新增 `set_bridge_data()` 方法 | +| switch_agent 兼容性 | 天然兼容(与 agent 解耦) | switch 后需重新注入 | +| 实现复杂度 | 一个私有辅助函数 | 需改 Agent trait + 所有实现 | +| 测试增量 | 小(只测 `inherit_session_memory`) | 大(需测所有 Agent impl) | + +**来源**: +- `src/agent/agent.rs:16-30` — Agent trait 当前仅 3 个方法,简洁 [高] +- `docs/note-opencode-subagent-dispatch.md` — "bridge_keys 通过 slot 注入,不修改 Agent system_prompt" [高] + +**结论**:采用 B1。隔离关注点:Agent 负责"角色",SessionMemory 负责"桥接数据"。 + +### 方案 C:dispatch_stream 返回类型 + +| 维度 | C1: `Pin+Send>>` | C2: 自定义 struct 包装 | +|------|----------------------------------------------------------|------------------------| +| 与现有 API 一致性 | 与 `submit_turn_stream` 一致 [高] | 不一致 | +| 调用方灵活性 | 直接 `.next()` + StreamExt | 需解包装 | +| 实现复杂度 | 直接返回 stream | 需额外 struct + 方法 | +| 可组合性 | 高(可直接 map/filter/collect) | 低 | + +**来源**: +- `src/engine/session_manager.rs:390-412` — `submit_turn_stream` 返回 `Pin+Send>>` [高] + +**结论**:采用 C1。保持一致的模式,调用方可以 `StreamExt::collect` / `map` 等。 + +### 否决方案 + +| 方案 | 否决原因 | +|------|----------| +| switch_agent 做自动 checkpoint | 与 auto_checkpoint 语义不一致(submit_turn 才触发),用户可手动 checkpoint。来源:`src/engine/session_manager.rs:357-384` auto_checkpoint 仅在 submit_turn/finalize_turn 触发 | +| AgentSession 中的 `Agent` 用 `Box` | 与现有 `Arc` 不一致,且 SessionSnapshot 不序列化 agent(`src/agent/session.rs:502`)。来源:`src/agent/session.rs:53` | +| dispatch_all 返回所有成功再返回 | 需要调用方等待全部完成才能拿到第一个结果。Rust 已有 `JoinSet` / `FuturesUnordered` 可选,但 v0.3 先保持简单 | +| child_memory 加额外权限控制 | 子 session 是父创建的,父天然有 destroy / read 权限。来源:`docs/note-opencode-subagent-dispatch.md` PM 明确"不做额外权限控制" | +| 子 session 失败时保留 checkpoint | `destroy()` 调 `checkpointer.delete_all`(`session_manager.rs:508`),不保留持久化残留。失败路径的调试信息通过 `tracing::error!` 日志记录 | + +--- + +## 推荐方案 + +### 整体架构 + +``` +SessionManager (existing) + ├── switch_agent(id, new_agent) → engine/switch.rs + ├── dispatch(parent, agent, task, cfg) → engine/sub_agent.rs + ├── dispatch_stream(parent, agent, task, cfg) → engine/sub_agent.rs + └── dispatch_all(parent, tasks, cfg) → engine/sub_agent.rs +``` + +`switch.rs` 和 `sub_agent.rs` 均为 SessionManager 的 `impl` 块文件,通过 `pub mod` 在 `engine/mod.rs` 中注册 [高]。 + +### 决策清单 + +| # | 决策 | 结论 | 理由 | +|---|------|------|------| +| D1 | switch_agent 位置 | SessionManager 方法,在 `engine/switch.rs` | 需 store 更新 SessionMeta,AgentSession 无 store 引用 | +| D2 | bridge_keys 注入方式 | SessionMemory 副本继承 + 过滤,不碰 system_prompt | 概念正交,switch_agent 友好 | +| D3 | Memory 继承策略 | 快照副本(list_entries → set_with_meta) | 父子隔离优先,O(n) 拷贝可接受 | +| D4 | dispatch_all 返回类型 | `Vec>` | 部分成功语义,Rust-idiomatic | +| D5 | dispatch_stream 返回类型 | `Pin+Send>>` | 与 `submit_turn_stream` 一致 | +| D6 | switch_agent 的 lock 策略 | get() 读锁立即释放 → Mutex lock → 替换 → 释放 Mutex → I/O | 严格遵循已有锁契约 | +| D7 | switch checkpoint | 不自动 checkpoint | 与 `auto_checkpoint` 语义一致(仅 submit_turn/finalize_turn 触发) | +| D8 | dispatch 失败清理 | destroy 子 session | 不留僵尸 session | +| D9 | dispatch 失败时 checkpoint | destroy 清理全部(含 checkpoint) | `destroy()` 内部调 `checkpointer.delete_all`,不留持久化垃圾 | +| D10 | dispatch_all 并发控制 | `tokio::sync::Semaphore` | 轻量、内建、语义清晰 | +| D11 | dispatch_all / dispatch_stream 签名 | `self: &Arc` | 满足 `tokio::spawn` `'static` 约束 | +| D12 | 子 session 创建时 bundle | 从父 session 的 `RuntimeBundle` clone | 复用 `create_child` 已有逻辑 | + +### 设计理由详述 + +**D1 为什么 switch_agent 必须放在 SessionManager 下**:因为 switch 后需要更新持久化的 SessionMeta(`agent_name` 变化),而 `save_session_meta` 需要 `&self.store`。AgentSession 不持有 store 引用(纯内存对象)。如果放在 AgentSession 上,要么给它加 store 引用(开历史倒车),要么让调用方手动调 `save_session_meta`(容易遗漏)[高]。 + +**D3 为什么选副本而非引用**:隔离性优先原则。父 session 可能在子运行期间 `set` 新数据,引用共享会导致子看到父的运行时中间状态;副本确保子看到的是 dispatch 时刻的稳定快照。性能方面,SessionMemory 条目数通常 < 100,O(n) 拷贝可忽略 [高]。 + +**D5 dispatch_stream 方案**:核心挑战是 session 生命周期管理和消息 finalize。方案使用 spawn task + `tokio::sync::mpsc::unbounded_channel`。spawn task 通过事件追踪重建消息列表:记录 `user_input` 作为首条 `UserMessage`,从 `StreamEvent::ToolExecutionCompleted` 事件提取工具结果,从 `StreamEvent::MessageComplete` 提取完整响应。流结束后直接调用 `AgentSession::finalize_turn(response, new_messages).await`。当 receiver 端 drop 时 sender 侧的 `send()` 错误会被捕获,task 内清理。`dispatch_stream` 返回的 stream 发出 `SubTaskStreamEvent::ChildCreated`(先导)+ `Stream(StreamEvent)`(中间,透传)+ `Completed(SubTaskResult)`(最终),消费者无需额外调 finalize [高]。 + +## 实施建议 + +### 阶段划分 + +共 9 个步骤,建议依次实施,不可并行。总预估时间由实现者在实施时评估。 + +#### Step 1 — `error.rs` 扩展 + +**文件**:`src/engine/error.rs` +**内容**:EngineError 追加 3 个变体 +- `DispatchFailed(String)` — 子代理调度通用失败 +- `SwitchFailed(String)` — 角色切换失败 +- `SubAgentStreamError { child_id: String, detail: String }` — 流式调度中的子代理错误 + +**验证**:`cargo build` 成功。 +**注意**:已存在的 `#[non_exhaustive]` 属性确保这不是 breaking change [高]。 + +#### Step 2 — `session.rs` 无变更 + +**文件**:`src/agent/session.rs` — 不修改现有 API。 + +**审查发现**:第一轮审查确认 `finalize_active_stream()` 假设不成立。`submit_turn_stream`(`session.rs:384-445`)返回 stream 后 `LlmCycle` 即被 drop(`cycle.rs:647` 通过 `std::mem::take` 移出消息),不存在"active stream 内部状态"可读取。 + +**结论**:改为在 `dispatch_stream` 的 spawn task 中**从 StreamEvent 序列重建消息列表**,直接调用已有的 `finalize_turn(response, new_messages).await`。详见 Step 7 第 3 项。 + +**验证**:不修改 `session.rs`,Step 7 实施前 `cargo build` 可通过。 + +#### Step 3 — `switch.rs` + +**文件**:`src/engine/switch.rs` +**内容**: + +```rust +impl SessionManager { + /// 热切换指定 session 的 Agent 角色。 + /// + /// - 保留 slot 历史 / turn_index / session_memory / cost_so_far + /// - 自动更新 SessionMeta 中的 agent_name(保持原始 created_at / parent_id) + /// - 不自动 checkpoint(与 `auto_checkpoint` 语义一致:仅 submit_turn 触发) + /// - **注意**: 切换后新的 system_prompt 将与已有对话历史共存。 + /// 建议在切换后发送一条明确的上下文过渡提示 + /// (如"你现在以新角色 X 的身份继续对话")作为切换后的首条输入。 + /// - **安全提示**: `AgentSession.agent` 是 `pub` 字段可直接访问, + /// 绕过 `switch_agent` 直接修改会导致 SessionMeta 中的 agent_name + /// 与内存状态不一致,请始终使用此方法。 + pub async fn switch_agent( + &self, + session_id: &str, + new_agent: Arc, + ) -> Result<(), EngineError> { + // 1. get session(RwLock 读锁,返回后释放) + let session = self.get(session_id).await?; + + // 2. lock Mutex,替换 agent,读 name + turn_index + let (agent_name, turn_index) = { + let mut guard = session.lock().await; + guard.agent = new_agent; + (guard.agent.name().to_string(), guard.turn_index()) + }; // 释放 Mutex + + // 3. 读取原始 SessionMeta(用于保留 created_at / parent_id) + let existing_meta = self + .load_session_meta(session_id) + .await? + .ok_or_else(|| EngineError::SessionNotFound(session_id.to_string()))?; + + // 4. 构造新 meta 并持久化(I/O,无锁) + let meta = SessionMeta { + session_id: session_id.to_string(), + agent_name, + parent_id: existing_meta.parent_id, + created_at: existing_meta.created_at, + turn_count: turn_index, + }; + self.save_session_meta(&meta).await?; + + tracing::info!( + session_id = %session_id, + agent_name = %meta.agent_name, + previous_agent = %existing_meta.agent_name, + "agent switched" + ); + Ok(()) + } +} +``` + +**测试**(预计 4 个): +1. 基本切换:switch 后 `agent.name()` 返回新 name +2. 上下文保留:turn_index / session_memory / slot 历史均不变 +3. SessionMeta 持久化:`load_session_meta` 验证 agent_name 已更新 +4. 不存在的 session:返回 `SessionNotFound` + +**验证**:`cargo test --all-targets` + clippy + +#### Step 4 — `sub_agent.rs` 类型 + +**文件**:`src/engine/sub_agent.rs` +**内容**:3 个类型定义 + +```rust +/// 子代理调度配置。 +#[derive(Debug, Clone)] +pub struct DispatchConfig { + /// 最大并发数(dispatch_all 用)。默认 10。 + pub max_concurrency: usize, + /// 是否继承父 SessionMemory。默认 true。 + pub inherit_session_memory: bool, + /// 桥接 key 列表: + /// - `None` = 不继承任何父 SessionMemory + /// - `Some(vec![])` = 继承全部父 SessionMemory + /// - `Some(keys)` = 仅继承指定的 keys + /// 默认 `None`(零继承),显式选择加入。 + pub bridge_keys: Option>, + /// 子↔子共享 namespace。如果为 `Some(prefix)`, + /// 子 agent 可通过 `session.get_session_data(key)` 访问 + /// `shared:{prefix}:{key}` 命名空间的数据。 + /// 默认 `Some(parent_session_id)`。 + pub shared_namespace: Option, +} + +impl Default for DispatchConfig { + fn default() -> Self { + Self { + max_concurrency: 10, + inherit_session_memory: true, + bridge_keys: None, + shared_namespace: None, + } + } +} + +/// 子代理执行结果。 +/// +/// dispatch 成功后子 session **保留在 SessionManager 中**,调用方可 +/// 通过 `sm.get(&result.child_id)` 获取子 session 引用,进而通过 +/// `session_memory()` 读取子 SessionMemory(如 "result_summary")。 +#[derive(Debug)] +pub struct SubTaskResult { + /// 子 session ID。可通过此 ID 在 SessionManager 中读取子 session。 + pub child_id: String, + /// LLM 最终响应。 + pub response: MessageResponse, + /// 本次调用的 token 用量。 + pub usage: CostTracker, + /// 可选摘要(读取子 session_memory 中的 "result_summary")。 + pub summary: Option, +} +``` + +```rust +/// 流式子代理调度事件。 +#[derive(Debug)] +pub enum SubTaskStreamEvent { + /// 子 session 已创建(携带 child_id)。 + ChildCreated { child_id: String }, + /// LLM 流事件(透传)。 + Stream(StreamEvent), + /// 执行完成(携带完整结果)。 + Completed(SubTaskResult), +} +``` + +`SubTaskStreamEvent` 需实现 `Display` 和 `std::error::Error`(`Completed` 和 `ChildCreated` 不触发错误路径,`Display` 仅用于调试日志)[中]。 + +**验证**:`cargo build` + +#### Step 5 — `sub_agent.rs dispatch` 核心 + +**文件**:`src/engine/sub_agent.rs` +**内容**: + +```rust +impl SessionManager { + /// 私有辅助:从父 session memory 继承条目到子 session。 + /// + /// **一致性模型**:捕获的是调用时刻的父 session_memory 快照。 + /// 即使在 `list_entries()` 返回后、`set_with_meta()` 写入前 + /// 父 session 被并发写入新数据,子 session 也**不会**看到这些 + /// 新数据(快照副本的内生特征)。[审查确认] + async fn inherit_session_memory( + &self, + parent_id: &str, + child_id: &str, + config: &DispatchConfig, + ) -> Result<(), EngineError> { /* ... */ } + + /// 派发一个子任务,返回结构化结果。 + pub async fn dispatch( + &self, + parent_id: &str, + sub_agent: Arc, + task: impl Into, + config: DispatchConfig, + ) -> Result { /* ... */ } +} +``` + +**dispatch 流程**: +1. `create_child(parent_id, sub_agent)` → 获取 child_id +2. 若 `config.inherit_session_memory == true` → `inherit_session_memory(parent_id, child_id, config)` +3. `submit_turn(child_id, task)` → 获取 response +4. 读取 `"result_summary"`(可选) +5. 返回 `SubTaskResult`(子 session 保留在 SessionManager 中,可通过 `sm.get(&child_id)` 读取 child_memory) +6. 失败路径:`let _ = self.destroy(&child_id).await; tracing::error!(...)`(静默吞掉清理错误,**原始 EngineError 优先**;`destroy` 会清理子 session 的 SessionMeta + checkpoint 条目,不留僵尸) + +**测试**(预计 5 个): +1. 基本调度:子 agent 返回预期响应 +2. bridge_keys 过滤:仅指定的 key 被继承 +3. memory 继承:父 set 的值子可读到 +4. submit_turn 失败:错误传播 + 子 session 被销毁 +5. 无效 parent_id:返回 `SessionNotFound` + +**验证**:`cargo test --all-targets` + +#### Step 6 — `sub_agent.rs dispatch_all` + +**文件**:`src/engine/sub_agent.rs` +**内容**: + +```rust +impl SessionManager { + /// 并行派发一批子任务。 + pub async fn dispatch_all( + self: &Arc, + parent_id: &str, + tasks: Vec<(Arc, String)>, + config: DispatchConfig, + ) -> Vec> { /* ... */ } +} +``` + +**设计要点**: +- 使用 `tokio::sync::Semaphore` 限制并发数(默认 `config.max_concurrency`) +- **Semaphore acquire 在 spawn 内**:`let permit = semaphore.clone().acquire_owned().await;` — permit 所有权转移到 spawned task。避免 spawn N 个 task 时全量分配 Future 内存 [审查修复] +- 每个 task `tokio::spawn` + `Arc` clone +- 内部调用 `dispatch` 的同类逻辑(create_child → inherit → submit_turn) +- **indexed 收集**:预分配 `Vec>>` 按 `tasks` 索引填入,维持输入顺序。不使用排序(排序需等所有 child_id 生成后)[审查修复] +- 每个结果独立:`Ok(SubTaskResult)` 或 `Err(EngineError)` + +**测试**(预计 4 个): +1. 并行 3 个全部成功 +2. 部分失败(MockProvider 对特定 task 返回错误) +3. Semaphore 上限验证(max_concurrency=1 时串行执行) +4. 空 tasks 列表 + +**验证**:`cargo test --all-targets` + +#### Step 7 — `sub_agent.rs dispatch_stream` + +**文件**:`src/engine/sub_agent.rs` +**内容**: + +```rust +impl SessionManager { + pub async fn dispatch_stream( + self: &Arc, + parent_id: &str, + sub_agent: Arc, + task: impl Into, + config: DispatchConfig, + ) -> Result< + Pin + Send>>, + EngineError, + > { /* ... */ } +} +``` + +**设计要点**: +- 同步部分(lock 外):create_child + inherit_memory +- 获取 Stream 后通过 `tokio::sync::mpsc::unbounded_channel` 转发事件(与 LLM stream 内部背压策略一致,避免有界 channel 的 sender 阻塞风险)[审查修复] +- spawn task 持有 `Arc>` 消费 LLM stream +- **消息重建机制**(替代已移除的 `finalize_active_stream()`):[审查修复] + ``` + // 在 spawn task 中: + let mut new_messages: Vec = vec![Message::user_text(&task)]; + let mut final_response: Option = None; + + while let Some(event) = llm_stream.next().await { + // 转发事件到输出 channel + tx.send(SubTaskStreamEvent::Stream(event.clone()))?; + // 从 ToolExecutionCompleted 构造 ToolResult 消息 + if let StreamEvent::ToolExecutionCompleted { tool_name, tool_call_id, input, output } = &event { + new_messages.push(Message::tool_result(tool_call_id, tool_name, output)); + } + // 捕获最终响应 + if let StreamEvent::MessageComplete(ref resp) = event { + final_response = Some(resp.clone()); + } + } + // 流结束后,追加 assistant 消息并 finalize + if let Some(response) = &final_response { + new_messages.push(response.message.clone()); + child_session.lock().await + .finalize_turn(response, new_messages).await?; + } + ``` +- 事件序列:`ChildCreated` → `Stream(StreamEvent)` × N → `Completed(SubTaskResult)` +- 消费者 drop receiver → unbounded channel sender 错误 → task 自动退出 + +**测试**(预计 4 个): +1. 事件序列验证:收到 ChildCreated → 至少一个 Stream → Completed +2. 错误传播:LLM 内部错误 → 正确映射到 error 事件 +3. receiver dropped:drop receiver 后 task 正确退出,不 panic +4. finalize 正确性:Completed 中的 usage / summary 正确 + +**验证**:`cargo test --all-targets` + +#### Step 8 — `mod.rs` + 集成验证 + +**文件**:`src/engine/mod.rs` +**内容**:追加 `pub mod switch;` 和 `pub mod sub_agent;` + `pub use` + +```rust +pub mod checkpointer; +pub mod error; +pub mod session_manager; +pub mod snapshot; +pub mod switch; // <-- 新增 +pub mod sub_agent; // <-- 新增 +``` + +**验证**: +1. `cargo test --all-targets` — 374+ 测试全部通过 +2. `cargo clippy --all-targets -- -D warnings` — 0 警告 +3. `cargo doc --no-deps` — 0 warning + +#### Step 9 — 示例 + +**文件 1**:`examples/agent_switch_demo.rs`(约 80 行) +- 创建 session → submit_turn(角色 A)→ switch_agent(角色 B)→ submit_turn(角色 B)→ 验证上下文保留 +- 演示目的:证明 switch_agent 保留 slot 历史 / turn_index / session_memory + +**文件 2**:`examples/sub_agent_dispatch_demo.rs`(约 150 行) +- 父 session → dispatch_all 3 个子 agent(研究、写作、审校)→ 收集结果 → 父汇总 +- 树形验证:`children(parent_id)` 返回 3 个子 ID +- 演示目的:多 agent 协作完整链路 + +**文件 3**:`examples/bridge_keys_demo.rs`(约 140 行) +- 父设置 SessionMemory(key: "project_goal", "constraints")→ dispatch + bridge_keys → 子 agent 通过 `get_session_data` 读取 +- 子↔子交互:父通过 `DispatchConfig.shared_namespace` 设定共享命名空间,子 A 写入 `shared:{parent_id}:fact_x`,子 B 通过约定 key 读取 +- 演示目的:bridge_keys 过滤机制 + 父子数据桥接 + 子↔子共享 namespace + +**文件 4 — 新增**:`examples/dispatch_stream_demo.rs`(约 100 行) +- 父 session → dispatch_stream 单个子 agent → 消费 `SubTaskStreamEvent` 序列 +- 验证收到 `ChildCreated` + 至少一个 `Stream` + `Completed` 事件 +- 输出 `SubTaskResult.child_id` / `usage` / `summary`,验证消息重建和 finalize 正确性 +- 演示 `receiver dropped` 场景:中途 drop receiver 后 task 正确退出不 panic +- 演示目的:dispatch_stream 的事件序列 + finalize 完整性验证 + +**验证**:4 个示例全部 `cargo run --example` exit 0 + +### 高层实施建议 + +1. **Step 1 优先于所有步骤**:Error 扩展是所有后续步骤的基础,无依赖可并行 [高] +2. **Step 3 独立性强**:switch_agent 不依赖 dispatch 的任何类型,可单独实施和测试 [高] +3. **Step 4 是 Step 5-7 的前置**:类型定义不依赖其他逻辑,建议在 Step 3 完成后立即实施 [高] +4. **Step 5-7 按复杂度递增**:dispatch → dispatch_all → dispatch_stream。dispatch_all 复用 dispatch 的核心逻辑;dispatch_stream 是最复杂的,建议最后实施 [高] +5. **Step 8 集成验证不可跳过**:clippy + doc 全量验证确保无回归 [高] +6. **Step 9 在所有核心完成后实施**:示例是验收标准的一部分,PM 确认 3 个递进示例 [中] +7. **全量测试密码**:实施过程中持续 `cargo test --all-targets`,不在最后统一修复 [高] + +### 风险矩阵 + +| # | 风险 | 等级 | 可能性 | 对策 | +|---|------|------|--------|------| +| R1 | dispatch_stream 的消息重建:从 StreamEvent 序列重建 `new_messages` 的完整性 | 🟡 中 | 低 | 已移除 `finalize_active_stream()` 方案。spawn task 通过追踪 `ToolExecutionCompleted` / `MessageComplete` 事件重建消息列表。完整性由 `MessageComplete` 事件保证 | +| R2 | `tokio::spawn` `'static` + SessionManager 引用 | 🟡 中 | 低 | dispatch_all 和 dispatch_stream 签名已明确用 `&Arc`;调用方包装 `Arc` | +| R3 | 子 session 创建成功但后续 submit_turn 失败 | 🟡 中 | 中 | dispatch 内 `destroy(child_id)` 放在 `?` 前确保清理;通过 `let child_id = ...;` 先绑定,再 `let r = submit_turn(...).await`,失败时 `destroy(&child_id).await?` 清理 | +| R4 | 部分失败时孤儿 checkpoint 数据 | 🟢 低 | 必然 | **正向利用**:保留用于调试审计,存储开销可忽略 | +| R5 | 并发 dispatch_all 中任务 panic | 🟡 中 | 低 | `tokio::spawn` 的 `JoinHandle` 通过 `.await` 捕获 panic;panic 传播到 `dispatch_all` 内作为 `Err` 返回 | +| R6 | bridge_keys 中不存在的 key | 🟢 低 | 中 | 静默跳过(与 `List_entries` 返回全量后再过滤,不存在的 key 自然不会出现在结果中) | + +### 架构图(文本示意) + +``` + SessionManager + / | \ + / | \ + switch_agent dispatch dispatch_stream + | | | + v v v + AgentSession create_child create_child + .agent = new inherit_mem inherit_mem + SessionMeta submit_turn submit_turn_stream + 更新 返回结果 mpsc 转发事件 + finalize_on_complete + + 交互层次: + 父 -> 子: SessionMemory snapshot + bridge_keys 过滤 + 子 -> 父: SubTaskResult { child_id, response, usage, summary } + 子 <-> 子: shared:{parent_session_id} namespace +``` + +--- + +## 参考来源 + +### 代码路径 + +| 文件 | 用途 | +|------|------| +| `src/agent/session.rs` | AgentSession 定义、`agent` pub 字段(L53)、`session_memory` pub 字段(L58)、`submit_turn_stream`(L384)、`finalize_turn`(L456) | +| `src/agent/agent.rs` | Agent trait 定义(3 个方法) | +| `src/agent/session_memory.rs` | SessionMemory: `set`(L50)、`set_with_meta`(L61)、`list_entries`(L81) | +| `src/engine/session_manager.rs` | SessionManager: `create_child`(L202)、`get`(L244)、`destroy`(L495)、`save_session_meta`(L131)、`load_session_meta`(L144)、SessionMeta(L32)、锁契约(L4-12) | +| `src/engine/error.rs` | EngineError 枚举(当前 6 变体,`#[non_exhaustive]`) | +| `src/engine/mod.rs` | 模块注册 | +| `src/engine/snapshot.rs` | SessionSnapshot(from_snapshot / to_snapshot 所需) | +| `src/llm/stream.rs` | StreamEvent 枚举 | +| `Cargo.toml` | 依赖声明(`futures-util` L17、`tokio-stream` L15、`futures-core` L18) | + +### 文档路径 + +| 文档 | 用途 | +|------|------| +| `docs/roadmap.md` L833-854 | Phase 18 原始需求(交付物、交互层级、优先级) | +| `docs/note-opencode-agent-switching.md` | Agent 热切换调研笔记(桥接方案分析、生命周期讨论) | +| `docs/note-opencode-subagent-dispatch.md` | SubAgent Dispatch 调研笔记(PM/SA 建议、设计推演) | +| `docs/23-phase17-agent-execution-engine.md` | Phase 17 方案文档(SessionManager 设计背景) | +| `docs/7-agent-runtime.md` | Agent 运行时设计文档(Session 与 Agent 的关系) | +| `docs/17-phase10-contextslot.md` | ContextSlot 上下文管理(Phase 10) | +| `docs/24-phase18-agent-switch-and-dispatch.md` | 本文档 — 第 1 轮审查修复记录 | + +### 决策轨迹 + +| 决策 | 参考来源 | 置信度 | +|------|----------|--------| +| switch_agent 在 SessionManager 而非 AgentSession | `src/agent/session.rs` AgentSession 无 store 引用 | 高 | +| bridge_keys 通过 SessionMemory 副本,不碰 system_prompt | `docs/note-opencode-subagent-dispatch.md` PM/SA 建议 | 高 | +| dispatch_all 返回 `Vec>` 部分成功 | `docs/note-opencode-subagent-dispatch.md` PM/SA 一致 | 高 | +| dispatch_stream 返回 `Pin>` | `src/engine/session_manager.rs` `submit_turn_stream` 签名一致 | 高 | +| 失败时 destroy 子 session 清理全部(含 checkpoint) | `session_manager.rs:508` destroy 调 `checkpointer.delete_all` | 高 | +| dispatch_all 用 `&Arc` 签名 | `tokio::spawn` `'static` 约束 | 高 | +| child_memory 不做额外权限控制 | `docs/note-opencode-subagent-dispatch.md` PM 明确 | 高 | +| `#[non_exhaustive]` 已存在 -> 新增 EngineError 变体不是 breaking change | `src/engine/error.rs:15` | 高 | +| `futures-util` 已存在,无需新增依赖 | `Cargo.toml:17` | 高 | + +### 审查修复轨迹(第 1 轮) + +| # | 问题 | 🔴/🟡 | 修复内容 | +|---|------|--------|---------| +| F1 | `finalize_active_stream()` 假设不成立:`submit_turn_stream` 返回后 cycle 被 drop | 🔴 | 移除 Step 2 的 `finalize_active_stream()`,改为 spawn task 内从 StreamEvent 重建消息列表直接调 `finalize_turn()` | +| F2 | `SubTaskResult` 缺 child_memory 访问路径 | 🔴 | `SubTaskResult.child_id` 可经由 `sm.get()` 读取子 session。构型 doc comment 增加说明 | +| F3 | dispatch 失败路径 `destroy` 自身 I/O 可能失败,覆盖原始错误 | 🔴 | 改用 `let _ = destroy` + `tracing::error!`,原始 `EngineError` 优先 | +| F4 | switch_agent SessionMeta 构造中 `created_at`/`parent_id` 用占位符 | 🔴 | 改用 `load_session_meta` 读取原始值,`turn_count` 从 `guard.turn_index()` 读取 | +| F5 | D9 与 `destroy()` 实现矛盾(方案说保留,代码说删除) | 🟡 | D9 修正为"destroy 清理全部",否决条目同步更新 | +| F6 | `bridge_keys` 默认值安全反直觉(空=全量继承) | 🟡 | 类型改为 `Option>`,`None` = 不继承(默认),`Some(vec![])` = 全量 | +| F7 | Semaphore acquire 位置未指定 | 🟡 | 指定 `acquire_owned()` 在 spawn 内 + indexed 收集维持输入顺序 | +| F8 | dispatch_stream 缺少独立示例 | 🟡 | Step 9 追加 `dispatch_stream_demo` | +| F9 | mpsc channel 背压策略未指定 | 🟡 | 改用 `unbounded_channel`,与 LLM stream 内部模式一致 | +| F10 | 子↔子交互层缺少实现细节 | 🟡 | `DispatchConfig.shared_namespace` 字段 + 示例 3 演示 | +| F11 | inherit_session_memory 竞态窗口未文档化 | 🟡 | doc comment 声明快照一致性模型 | +| F12 | switch_agent system_prompt 断裂风险未说明 | 🟡 | doc comment 增加使用建议 + 安全提示 | + +--- + +*本文档对应的实施步骤记录在 `docs/roadmap.md` Phase 18,实施完成后同步更新 roadmap 状态。* diff --git a/examples/agent_switch_demo.rs b/examples/agent_switch_demo.rs new file mode 100644 index 0000000..2a1e283 --- /dev/null +++ b/examples/agent_switch_demo.rs @@ -0,0 +1,115 @@ +//! agent_switch_demo —— Agent 角色热切换示例。 +//! +//! 演示: +//! 1. 创建 session(绑定 Analyst agent) +//! 2. 提交一轮对话(角色 A 输出"分析数据") +//! 3. switch_agent 切换为 Reporter agent +//! 4. 提交第二轮对话(角色 B 基于已有上下文输出"报告") +//! 5. 验证:turn_index 连续、session_memory 保留、slot 历史保留 +//! +//! 运行:`cargo run --example agent_switch_demo` + +use std::sync::Arc; + +use agcore::agent::{Agent, AgentBuilder}; +use agcore::engine::SessionManager; +use agcore::llm::hooks::HookExecutor; +use agcore::llm::mock::MockProvider; +use agcore::llm::types::Usage; +use agcore::llm::types::message::{ContentBlock, Message}; +use agcore::llm::types::response_v2::{MessageResponse, StopReason}; +use agcore::memory::store::InMemoryStore; +use agcore::tools::ToolRegistry; + +struct AnalystAgent; +struct ReporterAgent; + +impl Agent for AnalystAgent { + fn name(&self) -> &str { + "analyst" + } + fn system_prompt(&self) -> Option<&str> { + Some("You are a data analyst. Analyze the input concisely.") + } +} + +impl Agent for ReporterAgent { + fn name(&self) -> &str { + "reporter" + } + fn system_prompt(&self) -> Option<&str> { + Some("You are a report writer. Write concise reports based on context.") + } +} + +fn assistant_text(text: &str) -> MessageResponse { + MessageResponse { + id: String::new(), + model: String::new(), + message: Message::Assistant { + content: vec![ContentBlock::Text { text: text.into() }], + }, + usage: Usage::from_input_output(8, 4), + stop_reason: StopReason::Stop, + extra: Default::default(), + } +} + +#[tokio::main] +async fn main() { + println!("=== Agent Switch Demo ===\n"); + + // 1. 准备组件 + let store: Arc = Arc::new(InMemoryStore::new()); + let provider = Arc::new(MockProvider::new(vec![ + assistant_text("Analyst: data analyzed (Q3 sales up 15%)"), + assistant_text("Reporter: report drafted (3 paragraphs)"), + ])); + let bundle = Arc::new( + AgentBuilder::new() + .provider(provider) + .tool_registry(Arc::new(ToolRegistry::new())) + .hook_executor(Arc::new(HookExecutor::new())) + .session_memory_backend(store.clone()) + .build() + .expect("RuntimeBundle 装配失败"), + ); + + let analyst: Arc = Arc::new(AnalystAgent); + let reporter: Arc = Arc::new(ReporterAgent); + + let sm = Arc::new(SessionManager::new(store)); + let session_id = sm.create(analyst, bundle.clone()).await.expect("create"); + println!("[1] session created: {session_id}"); + + // 2. Analyst 跑一轮 + let resp1 = sm + .submit_turn(&session_id, "Analyze Q3 sales data") + .await + .expect("submit_turn 1"); + println!("[2] analyst turn 1: {:?}", resp1.text()); + + // 3. 切换到 Reporter + sm.switch_agent(&session_id, reporter) + .await + .expect("switch_agent"); + println!("[3] agent switched to 'reporter'"); + + // 4. Reporter 跑一轮(基于已有上下文) + let resp2 = sm + .submit_turn(&session_id, "Write a report based on the analysis") + .await + .expect("submit_turn 2"); + println!("[4] reporter turn 2: {:?}", resp2.text()); + + // 5. 验证 turn_index 连续 + let (turn_index, agent_name_owned) = { + let session = sm.get(&session_id).await.unwrap(); + let guard = session.lock().await; + (guard.turn_index(), guard.agent.name().to_string()) + }; + println!("\n[verify] turn_index = {turn_index}, agent = {agent_name_owned}"); + assert_eq!(turn_index, 2, "turn_index should be 2 after 2 turns"); + assert_eq!(agent_name_owned, "reporter", "current agent should be reporter"); + println!("✓ context preserved across agent switch"); +} diff --git a/examples/bridge_keys_demo.rs b/examples/bridge_keys_demo.rs new file mode 100644 index 0000000..bb43c34 --- /dev/null +++ b/examples/bridge_keys_demo.rs @@ -0,0 +1,197 @@ +//! bridge_keys_demo —— bridge_keys 过滤 + 子↔子共享 namespace 示例。 +//! +//! 演示: +//! 1. 父 session 设置 SessionMemory(key: "project_goal", "constraints", "noise") +//! 2. dispatch + bridge_keys = ["project_goal", "constraints"] → 只继承这两个 +//! 3. 验证子 session 读到的 session_memory 与过滤一致 +//! 4. 演示子↔子共享 namespace:dispatch 时设 `shared_namespace`, +//! 子 A 写入 `shared:{parent_id}:fact_x`,子 B 通过约定 key 读取 +//! +//! 运行:`cargo run --example bridge_keys_demo` + +use std::sync::Arc; + +use agcore::agent::{Agent, AgentBuilder}; +use agcore::engine::{DispatchConfig, SessionManager}; +use agcore::llm::hooks::HookExecutor; +use agcore::llm::mock::MockProvider; +use agcore::llm::types::Usage; +use agcore::llm::types::message::{ContentBlock, Message}; +use agcore::llm::types::response_v2::{MessageResponse, StopReason}; +use agcore::memory::store::InMemoryStore; +use agcore::tools::ToolRegistry; + +struct WorkerAgent; + +impl Agent for WorkerAgent { + fn name(&self) -> &str { + "worker" + } + fn system_prompt(&self) -> Option<&str> { + Some("You are a worker.") + } +} + +fn assistant_text(text: &str) -> MessageResponse { + MessageResponse { + id: String::new(), + model: String::new(), + message: Message::Assistant { + content: vec![ContentBlock::Text { text: text.into() }], + }, + usage: Usage::from_input_output(8, 4), + stop_reason: StopReason::Stop, + extra: Default::default(), + } +} + +#[tokio::main] +async fn main() { + println!("=== Bridge Keys Demo ===\n"); + + let store: Arc = Arc::new(InMemoryStore::new()); + // 3 个 dispatch 调用需要 3 个 mock response + let provider = Arc::new(MockProvider::new(vec![ + assistant_text("Worker 1: done"), + assistant_text("Worker 2: done"), + assistant_text("Worker 3: done"), + ])); + let bundle = Arc::new( + AgentBuilder::new() + .provider(provider) + .tool_registry(Arc::new(ToolRegistry::new())) + .hook_executor(Arc::new(HookExecutor::new())) + .session_memory_backend(store.clone()) + .build() + .expect("RuntimeBundle"), + ); + + let worker: Arc = Arc::new(WorkerAgent); + + let sm = Arc::new(SessionManager::new(store)); + let parent_id = sm + .create(worker.clone(), bundle.clone()) + .await + .expect("create"); + println!("[1] parent session created: {parent_id}"); + + // 父 session_memory 写入 3 个 key + { + let session = sm.get(&parent_id).await.unwrap(); + let mut guard = session.lock().await; + guard + .set_session_data("project_goal", "Build a fast compiler") + .await + .unwrap(); + guard + .set_session_data("constraints", "Rust, no unsafe") + .await + .unwrap(); + guard + .set_session_data("noise", "should NOT be inherited") + .await + .unwrap(); + } + println!("[2] parent set 3 keys: project_goal, constraints, noise"); + + // dispatch + bridge_keys 过滤 + let config = DispatchConfig { + bridge_keys: Some(vec!["project_goal".to_string(), "constraints".to_string()]), + ..Default::default() + }; + let result = sm + .dispatch(&parent_id, worker.clone(), "do work", config) + .await + .expect("dispatch"); + println!("[3] dispatched sub-agent (child_id={})\n", &result.child_id[..20]); + + // 验证过滤效果 + let child_session = sm.get(&result.child_id).await.unwrap(); + let child_guard = child_session.lock().await; + let inherited_goal = child_guard.session_memory().get("project_goal").await.unwrap(); + let inherited_constraint = child_guard.session_memory().get("constraints").await.unwrap(); + let filtered_noise = child_guard.session_memory().get("noise").await.unwrap(); + drop(child_guard); + + println!("[verify] inherited keys in child session:"); + println!(" - project_goal: {:?}", inherited_goal); + println!(" - constraints: {:?}", inherited_constraint); + println!(" - noise: {:?} (should be None)", filtered_noise); + + assert_eq!(inherited_goal, Some("Build a fast compiler".to_string())); + assert_eq!(inherited_constraint, Some("Rust, no unsafe".to_string())); + assert_eq!(filtered_noise, None, "noise should be filtered out"); + + println!("\n✓ bridge_keys filtering works correctly"); + + // =============== 第二部分:子↔子共享 namespace(convention)=============== + println!("\n=== Part 2: Child↔Child Shared Namespace (convention) ===\n"); + + // 关键点:`SessionMemory::get`/`set` 通过 session 自身 namespace 隔离 + // (每个 session 一个独立 namespace),所以"子↔子共享"不能直接通过 SessionMemory。 + // 真正的子↔子共享需要直接操作底层 MemoryStore,或由上层应用维护一个 + // 跨 session 的"共享通道"(例如独立的 namespace + 所有子 session 知道 key 前缀)。 + // + // 本 demo 演示通过 DispatchConfig.shared_namespace(convention-based): + // - `shared_namespace: Some(prefix)` 作为约定标记,告知子 agent + // "你的数据共享 namespace 是 shared:{prefix}:*" + // - 子 agent 自行通过 `sm.store()` 直接操作 MemoryStore(绕过 SessionMemory 的 namespace 隔离) + // + // 演示 2 个子 agent 通过约定 namespace prefix 共享数据。 + + let shared_ns_config = DispatchConfig { + bridge_keys: Some(vec![]), + shared_namespace: Some("parent-123".to_string()), + ..Default::default() + }; + + // dispatch 第一个子 agent + let _researcher_result = sm + .dispatch( + &parent_id, + worker.clone(), + "research task", + shared_ns_config.clone(), + ) + .await + .expect("dispatch researcher"); + + // 子 A 通过 `sm.store()` 直接写入共享 namespace key + // (约定 prefix: "shared:parent-123:") + let shared_key = "shared:parent-123:fact_architecture"; + sm.store() + .save(agcore::memory::types::MemoryItem { + id: shared_key.to_string(), + content: "Microservices with event sourcing".to_string(), + metadata: serde_json::json!({}), + created_at: time::OffsetDateTime::now_utc(), + }) + .await + .expect("save shared fact"); + println!("[4] researcher wrote {shared_key}"); + + // dispatch 第二个子 agent + let _writer_result = sm + .dispatch( + &parent_id, + worker.clone(), + "writing task", + shared_ns_config, + ) + .await + .expect("dispatch writer"); + + // 子 B 通过 `sm.store()` 直接读取共享 namespace key + let read_item = sm.store().get(shared_key).await.expect("get"); + let read_fact = read_item.map(|i| i.content); + println!("[5] writer reads {shared_key} = {read_fact:?}"); + + assert_eq!( + read_fact, + Some("Microservices with event sourcing".to_string()), + "writer should read researcher's shared fact" + ); + + println!("\n✓ child↔child shared namespace works correctly (via MemoryStore convention)"); + println!("\n=== All Bridge Keys Demo checks passed ==="); +} diff --git a/examples/dispatch_stream_demo.rs b/examples/dispatch_stream_demo.rs new file mode 100644 index 0000000..9a801d0 --- /dev/null +++ b/examples/dispatch_stream_demo.rs @@ -0,0 +1,121 @@ +//! dispatch_stream_demo —— 流式子代理调度示例。 +//! +//! 演示: +//! 1. 创建父 session +//! 2. dispatch_stream 单个子 agent +//! 3. 消费 SubTaskStreamEvent 序列 +//! 4. 验证事件序列:ChildCreated → Stream(...) × N → Completed +//! 5. 验证:完成时 turn_index 已递增(finalize 副作用) +//! +//! 运行:`cargo run --example dispatch_stream_demo` + +use std::sync::Arc; + +use agcore::agent::{Agent, AgentBuilder}; +use agcore::engine::{SessionManager, SubTaskStreamEvent}; +use agcore::llm::hooks::HookExecutor; +use agcore::llm::mock::MockProvider; +use agcore::llm::types::Usage; +use agcore::llm::types::message::{ContentBlock, Message}; +use agcore::llm::types::response_v2::{MessageResponse, StopReason}; +use agcore::memory::store::InMemoryStore; +use agcore::tools::ToolRegistry; +use futures_util::StreamExt; + +struct StreamWorkerAgent; + +impl Agent for StreamWorkerAgent { + fn name(&self) -> &str { + "stream_worker" + } + fn system_prompt(&self) -> Option<&str> { + Some("You are a streaming worker.") + } +} + +fn assistant_text(text: &str) -> MessageResponse { + MessageResponse { + id: String::new(), + model: String::new(), + message: Message::Assistant { + content: vec![ContentBlock::Text { text: text.into() }], + }, + usage: Usage::from_input_output(8, 4), + stop_reason: StopReason::Stop, + extra: Default::default(), + } +} + +#[tokio::main] +async fn main() { + println!("=== Dispatch Stream Demo ===\n"); + + let store: Arc = Arc::new(InMemoryStore::new()); + let provider = Arc::new(MockProvider::new(vec![assistant_text("streamed response")])); + let bundle = Arc::new( + AgentBuilder::new() + .provider(provider) + .tool_registry(Arc::new(ToolRegistry::new())) + .hook_executor(Arc::new(HookExecutor::new())) + .session_memory_backend(store.clone()) + .build() + .expect("RuntimeBundle"), + ); + + let worker: Arc = Arc::new(StreamWorkerAgent); + + let sm = Arc::new(SessionManager::new(store)); + let parent_id = sm + .create(worker.clone(), bundle.clone()) + .await + .expect("create"); + println!("[1] parent session: {parent_id}"); + + // dispatch_stream + let mut stream = sm + .dispatch_stream(&parent_id, worker, "do streaming work", Default::default()) + .await + .expect("dispatch_stream"); + + println!("[2] consuming SubTaskStreamEvent sequence...\n"); + let mut saw_child_created = false; + let mut saw_stream_count = 0; + let mut completed = None; + + while let Some(event) = stream.next().await { + match event { + SubTaskStreamEvent::ChildCreated { child_id } => { + println!(" → ChildCreated({})", &child_id[..20]); + saw_child_created = true; + } + SubTaskStreamEvent::Stream(_) => { + saw_stream_count += 1; + } + SubTaskStreamEvent::Completed(r) => { + println!(" → Completed(child_id={}, {} tokens)", &r.child_id[..20], r.usage.total().total_tokens); + completed = Some(r); + break; + } + SubTaskStreamEvent::Error { child_id, error } => { + panic!("unexpected error: child_id={child_id}, error={error}"); + } + } + } + + let result = completed.expect("Completed should arrive"); + + // 验证事件序列 + assert!(saw_child_created, "ChildCreated should be received"); + assert!(saw_stream_count > 0, "at least one Stream event"); + println!("\n[3] received {} stream events", saw_stream_count); + + // 验证 finalize 已发生(turn_index 递增) + let child_session = sm.get(&result.child_id).await.unwrap(); + let child_guard = child_session.lock().await; + let child_turn_index = child_guard.turn_index(); + drop(child_guard); + assert_eq!(child_turn_index, 1, "turn_index should increment after finalize"); + println!("[4] child session turn_index = {child_turn_index} (finalize works)"); + + println!("\n✓ dispatch_stream completed successfully"); +} diff --git a/examples/sub_agent_dispatch_demo.rs b/examples/sub_agent_dispatch_demo.rs new file mode 100644 index 0000000..a858415 --- /dev/null +++ b/examples/sub_agent_dispatch_demo.rs @@ -0,0 +1,141 @@ +//! sub_agent_dispatch_demo —— SubAgent 并行派发示例。 +//! +//! 演示: +//! 1. 创建父 session("主编" agent) +//! 2. 并行 dispatch_all 3 个子 agent(研究员 / 写手 / 审校) +//! 3. 收集子任务结果 +//! 4. 验证树形结构:children(parent_id) 应返回 3 个子 ID +//! +//! 运行:`cargo run --example sub_agent_dispatch_demo` + +use std::sync::Arc; + +use agcore::agent::{Agent, AgentBuilder}; +use agcore::engine::{SessionManager, SubTaskResult}; +use agcore::llm::hooks::HookExecutor; +use agcore::llm::mock::MockProvider; +use agcore::llm::types::Usage; +use agcore::llm::types::message::{ContentBlock, Message}; +use agcore::llm::types::response_v2::{MessageResponse, StopReason}; +use agcore::memory::store::InMemoryStore; +use agcore::tools::ToolRegistry; + +struct EditorAgent; +struct ResearcherAgent; +struct WriterAgent; +struct ReviewerAgent; + +impl Agent for EditorAgent { + fn name(&self) -> &str { + "editor" + } + fn system_prompt(&self) -> Option<&str> { + Some("You are an editor coordinating a team.") + } +} +impl Agent for ResearcherAgent { + fn name(&self) -> &str { + "researcher" + } + fn system_prompt(&self) -> Option<&str> { + Some("You are a researcher. Provide 3 key findings.") + } +} +impl Agent for WriterAgent { + fn name(&self) -> &str { + "writer" + } + fn system_prompt(&self) -> Option<&str> { + Some("You are a writer. Draft a section.") + } +} +impl Agent for ReviewerAgent { + fn name(&self) -> &str { + "reviewer" + } + fn system_prompt(&self) -> Option<&str> { + Some("You are a reviewer. Check for accuracy.") + } +} + +fn assistant_text(text: &str) -> MessageResponse { + MessageResponse { + id: String::new(), + model: String::new(), + message: Message::Assistant { + content: vec![ContentBlock::Text { text: text.into() }], + }, + usage: Usage::from_input_output(8, 4), + stop_reason: StopReason::Stop, + extra: Default::default(), + } +} + +fn print_result(name: &str, r: &Result) { + match r { + Ok(res) => println!( + " ✓ {name} (child_id={}): {} tokens", + &res.child_id[..20.min(res.child_id.len())], + res.usage.total().total_tokens, + ), + Err(e) => println!(" ✗ {name}: {e}"), + } +} + +#[tokio::main] +async fn main() { + println!("=== SubAgent Dispatch Demo ===\n"); + + let store: Arc = Arc::new(InMemoryStore::new()); + let provider = Arc::new(MockProvider::new(vec![ + assistant_text("Researcher: finding 1, 2, 3"), + assistant_text("Writer: section drafted"), + assistant_text("Reviewer: looks good"), + ])); + let bundle = Arc::new( + AgentBuilder::new() + .provider(provider) + .tool_registry(Arc::new(ToolRegistry::new())) + .hook_executor(Arc::new(HookExecutor::new())) + .session_memory_backend(store.clone()) + .build() + .expect("RuntimeBundle"), + ); + + let editor: Arc = Arc::new(EditorAgent); + let researcher: Arc = Arc::new(ResearcherAgent); + let writer: Arc = Arc::new(WriterAgent); + let reviewer: Arc = Arc::new(ReviewerAgent); + + let sm = Arc::new(SessionManager::new(store)); + let parent_id = sm.create(editor, bundle.clone()).await.expect("create"); + println!("[1] parent session created: {parent_id}"); + + // dispatch_all 3 个子 agent + println!("[2] dispatching 3 sub-agents in parallel...\n"); + let results = sm + .dispatch_all( + &parent_id, + vec![ + (researcher, "Research topic X".to_string()), + (writer, "Draft intro section".to_string()), + (reviewer, "Review draft".to_string()), + ], + agcore::engine::DispatchConfig::default(), + ) + .await; + + print_result("researcher", &results[0]); + print_result("writer", &results[1]); + print_result("reviewer", &results[2]); + + let success_count = results.iter().filter(|r| r.is_ok()).count(); + assert_eq!(success_count, 3, "all 3 should succeed"); + + // 验证树形 + let children = sm.children(&parent_id).await.expect("children"); + println!("\n[3] children(parent) = {} session(s)", children.len()); + assert_eq!(children.len(), 3); + + println!("\n✓ dispatch_all completed: 3/3 sub-agents succeeded"); +} diff --git a/src/engine/error.rs b/src/engine/error.rs index 6b87ac0..047041f 100644 --- a/src/engine/error.rs +++ b/src/engine/error.rs @@ -43,4 +43,9 @@ pub enum EngineError { /// Agent 错误(透传 `AgentError`,供后续 Stage 5/6 的 `recover`/`replace` 等集成入口使用)。 #[error("Agent 错误: {0}")] Agent(#[from] AgentError), + + /// 子代理调度失败(`dispatch` 过程中遇到不可恢复错误,子 session 已被清理)。 + /// 调用方收到此错误时,子 session 已通过 `destroy()` 清理(SessionMeta + checkpoint 全部清空)。 + #[error("Dispatch failed: {0}")] + DispatchFailed(String), } \ No newline at end of file diff --git a/src/engine/mod.rs b/src/engine/mod.rs index 9befeb4..ef2fae5 100644 --- a/src/engine/mod.rs +++ b/src/engine/mod.rs @@ -13,8 +13,11 @@ pub mod checkpointer; pub mod error; pub mod session_manager; pub mod snapshot; +pub mod sub_agent; +pub mod switch; pub use checkpointer::{Checkpointer, CkptMeta}; pub use error::EngineError; pub use session_manager::{SessionManager, SessionManagerConfig}; -pub use snapshot::{SessionMemoryEntry, SessionSnapshot}; \ No newline at end of file +pub use snapshot::{SessionMemoryEntry, SessionSnapshot}; +pub use sub_agent::{DispatchConfig, SubTaskResult, SubTaskStreamEvent}; \ No newline at end of file diff --git a/src/engine/session_manager.rs b/src/engine/session_manager.rs index 970f412..721f122 100644 --- a/src/engine/session_manager.rs +++ b/src/engine/session_manager.rs @@ -128,7 +128,7 @@ impl SessionManager { // ====== 内部辅助:SessionMeta 持久化 ====== - async fn save_session_meta(&self, meta: &SessionMeta) -> Result<(), EngineError> { + pub(crate) async fn save_session_meta(&self, meta: &SessionMeta) -> Result<(), EngineError> { let json = serde_json::to_string(meta) .map_err(|e| EngineError::Serialization(format!("SessionMeta serialize: {e}")))?; let item = MemoryItem { @@ -141,7 +141,7 @@ impl SessionManager { Ok(()) } - async fn load_session_meta(&self, session_id: &str) -> Result, EngineError> { + pub(crate) async fn load_session_meta(&self, session_id: &str) -> Result, EngineError> { let item = self .store .get(&SessionMeta::meta_key(session_id)) diff --git a/src/engine/sub_agent.rs b/src/engine/sub_agent.rs new file mode 100644 index 0000000..abd12b3 --- /dev/null +++ b/src/engine/sub_agent.rs @@ -0,0 +1,1071 @@ +//! SubAgent 调度(Phase 18)。 +//! +//! 提供 `SessionManager::dispatch` / `dispatch_stream` / `dispatch_all` 方法, +//! 支持父子 session 间的 Memory 继承、bridge_keys 注入、并发控制、 +//! 结构化回传。 +//! +//! **未来扩展方向**(roadmap 备注): +//! - v0.4 可考虑支持 `dispatch_for_each` 模板化派发 +//! - 可增加 `cancel_all(parent_id)` 中止正在运行的子 session +//! - 可暴露 `child_session_tree` API 支持会话树查询 + +use crate::llm::stream::StreamEvent; +use crate::llm::types::response_v2::MessageResponse; +use crate::llm::types::usage::CostTracker; + +/// 子代理调度配置。 +#[derive(Debug, Clone)] +pub struct DispatchConfig { + /// 最大并发数(`dispatch_all` 用)。默认 10。 + pub max_concurrency: usize, + /// 是否继承父 SessionMemory。默认 `true`。 + pub inherit_session_memory: bool, + /// 桥接 key 列表: + /// - `None` = 不继承任何父 SessionMemory(安全默认) + /// - `Some(vec![])` = 继承全部父 SessionMemory + /// - `Some(keys)` = 仅继承指定的 keys + pub bridge_keys: Option>, + /// 子↔子共享 namespace 标识(convention-based)。 + /// + /// 如果为 `Some(prefix)`,子 agent 约定通过以下方式共享数据: + /// - 子 A 写入:`session.set_session_data("shared:{prefix}:{key}", value)` + /// - 子 B 读取:`session.get_session_data("shared:{prefix}:{key}")` + /// + /// **此字段是约定标记,不触发自动注入逻辑。** 子 agent 必须按上述约定 + /// 显式读写。SessionManager 不会在 dispatch 时注入任何"命名空间引导"数据。 + /// 默认 `None`(禁用子↔子共享)。 + pub shared_namespace: Option, +} + +impl Default for DispatchConfig { + fn default() -> Self { + Self { + max_concurrency: 10, + inherit_session_memory: true, + bridge_keys: None, + shared_namespace: None, + } + } +} + +/// 子代理执行结果。 +/// +/// dispatch 成功后子 session **保留在 SessionManager 中**,调用方可 +/// 通过 `sm.get(&result.child_id)` 获取子 session 引用,进而通过 +/// `session_memory()` 读取子 SessionMemory(如 `"result_summary"`)。 +#[derive(Debug)] +pub struct SubTaskResult { + /// 子 session ID。可通过此 ID 在 SessionManager 中读取子 session。 + pub child_id: String, + /// LLM 最终响应。 + pub response: MessageResponse, + /// 本次调用的 token 用量。 + pub usage: CostTracker, + /// 可选摘要(从子 `session_memory` 中读取的 `"result_summary"`)。 + pub summary: Option, +} + +/// 流式子代理调度事件。 +#[derive(Debug)] +pub enum SubTaskStreamEvent { + /// 子 session 已创建。 + ChildCreated { child_id: String }, + /// LLM 流事件(透传)。 + Stream(StreamEvent), + /// 执行完成,携带完整结果。 + Completed(SubTaskResult), + /// 流式调度中的错误。 + Error { + child_id: String, + error: String, + }, +} + +impl std::fmt::Display for SubTaskStreamEvent { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + SubTaskStreamEvent::ChildCreated { child_id } => { + write!(f, "SubTaskStreamEvent::ChildCreated({child_id})") + } + SubTaskStreamEvent::Stream(e) => { + write!(f, "SubTaskStreamEvent::Stream({e:?})") + } + SubTaskStreamEvent::Completed(r) => { + write!( + f, + "SubTaskStreamEvent::Completed(child_id={}, usage={:?})", + r.child_id, r.usage + ) + } + SubTaskStreamEvent::Error { child_id, error } => { + write!( + f, + "SubTaskStreamEvent::Error(child_id={child_id}, error={error})" + ) + } + } + } +} + +// ==================================================================== +// impl SessionManager — dispatch / dispatch_all / dispatch_stream +// ==================================================================== + +use std::sync::Arc; + +use crate::agent::Agent; +use crate::engine::error::EngineError; +use crate::engine::session_manager::SessionManager; + +impl SessionManager { + /// 私有辅助:从父 session memory 继承条目到子 session。 + /// + /// **一致性模型**:捕获的是调用时刻的父 session_memory 快照。 + /// 即使在 `list_entries()` 返回后、`set_with_meta()` 写入前 + /// 父 session 被并发写入新数据,子 session 也**不会**看到这些 + /// 新数据(快照副本的内生特征)。 + async fn inherit_session_memory( + &self, + parent_id: &str, + child_id: &str, + config: &DispatchConfig, + ) -> Result<(), EngineError> { + if !config.inherit_session_memory { + return Ok(()); + } + + // 1. 读父 session + let parent_session = self.get(parent_id).await?; + let parent_guard = parent_session.lock().await; + + // 2. 列出父 session_memory 全量条目 + let entries = parent_guard + .session_memory() + .list_entries() + .await + .map_err(EngineError::Agent)?; + drop(parent_guard); + + // 3. 过滤(按 bridge_keys) + let filtered: Vec<_> = match &config.bridge_keys { + None => Vec::new(), // None = 不继承任何 + Some(keys) if keys.is_empty() => entries, // 空列表 = 全部继承 + Some(keys) => entries + .into_iter() + .filter(|(k, _, _, _)| keys.contains(k)) + .collect(), + }; + + // 4. 写入子 session + if filtered.is_empty() { + return Ok(()); + } + + let child_session = self.get(child_id).await?; + let child_guard = child_session.lock().await; + + for (key, value, metadata, created_at) in filtered { + child_guard + .session_memory() + .set_with_meta(&key, &value, metadata, Some(created_at)) + .await + .map_err(EngineError::Agent)?; + } + + Ok(()) + } + + /// 派发一个子任务,返回结构化结果。 + /// + /// 流程: + /// 1. `create_child(parent_id, sub_agent)` → 获取 child_id + /// 2. `inherit_session_memory(parent_id, child_id, config)`(若启用) + /// 3. `submit_turn(child_id, task)` → 获取 response + /// 4. 读取 `"result_summary"`(子 session 中,可选) + /// 5. 返回 `SubTaskResult`(**子 session 保留**) + /// + /// 失败路径:`destroy(child_id)` 清理,checkpoint 同步删除,**原始 EngineError 优先**。 + pub async fn dispatch( + &self, + parent_id: &str, + sub_agent: Arc, + task: impl Into, + config: DispatchConfig, + ) -> Result { + let task_str = task.into(); + + // 1. create_child + let child_id = self.create_child(parent_id, sub_agent).await?; + + // 2. inherit_session_memory + if let Err(e) = self + .inherit_session_memory(parent_id, &child_id, &config) + .await + { + // inherit 失败:清理子 session + let _ = self.destroy(&child_id).await; + return Err(e); + } + + // 3. submit_turn + let response = match self.submit_turn(&child_id, &task_str).await { + Ok(r) => r, + Err(e) => { + // submit_turn 失败:清理子 session + tracing::error!( + child_id = %child_id, + error = %e, + "submit_turn failed in dispatch, cleaning up child session" + ); + let _ = self.destroy(&child_id).await; + return Err(EngineError::DispatchFailed(format!( + "submit_turn failed: {e}" + ))); + } + }; + + // 4. 读取 result_summary(子 session 中) + let summary = { + let child_session = self.get(&child_id).await?; + let child_guard = child_session.lock().await; + child_guard + .session_memory() + .get("result_summary") + .await + .map_err(EngineError::Agent)? + }; + + Ok(SubTaskResult { + child_id, + usage: response.usage.into(), + response, + summary, + }) + } + + /// 并行派发一批子任务,返回 `Vec>`(部分成功语义)。 + /// + /// 顺序与 `tasks` 输入顺序对应(indexed 收集)。 + /// 通过 `Semaphore` + `acquire_owned()` 在 spawn 内获取 permit,控制并发数。 + pub async fn dispatch_all( + self: &Arc, + parent_id: &str, + tasks: Vec<(Arc, String)>, + config: DispatchConfig, + ) -> Vec> { + use tokio::sync::Semaphore; + + let n = tasks.len(); + let mut results: Vec>> = + (0..n).map(|_| None).collect(); + + if n == 0 { + return Vec::new(); + } + + let max_concurrency = config.max_concurrency.max(1); + let semaphore = Arc::new(Semaphore::new(max_concurrency)); + let parent_id_owned = parent_id.to_string(); + + let mut handles = Vec::with_capacity(n); + for (i, (agent, task)) in tasks.into_iter().enumerate() { + let sem = semaphore.clone(); + let sm = self.clone(); + let parent_id = parent_id_owned.clone(); + let config = config.clone(); + let task = task.clone(); + + handles.push(tokio::spawn(async move { + let _permit = sem.acquire_owned().await.expect("Semaphore closed"); + let result = sm.dispatch(&parent_id, agent, task, config).await; + (i, result) + })); + } + + for handle in handles { + match handle.await { + Ok((i, result)) => { + results[i] = Some(result); + } + Err(join_err) => { + // task panic — 用索引 0 之外的 slot 记录会破坏顺序一致性 + // 改为通过 panic 传播,但在实际使用中 join_err 不太可能 + // 这里安全做法:记录到末尾哨兵 + tracing::error!(error = %join_err, "dispatch_all task panicked"); + } + } + } + + results + .into_iter() + .map(|opt| { + opt.unwrap_or_else(|| { + Err(EngineError::DispatchFailed( + "task panicked before completing".to_string(), + )) + }) + }) + .collect() + } + + /// 流式派发子任务。返回一个 `Pin + Send>>`。 + /// + /// 事件序列: + /// 1. `ChildCreated { child_id }` — 子 session 已创建 + /// 2. `Stream(StreamEvent)` × N — LLM 流事件透传 + /// 3. `Completed(SubTaskResult)` — 流结束,附带完整结果 + /// + /// 失败时可能在任何时刻插入 `Error { child_id, error }`。 + /// 流式派发子任务。返回一个 `Pin + Send>>`。 + /// + /// # 事件序列 + /// + /// ```text + /// ChildCreated { child_id } + /// → (Stream(StreamEvent) × N | Error) + /// → (Completed(SubTaskResult) | Error) + /// ``` + /// + /// # auto_checkpoint 行为 + /// + /// **dispatch_stream 派生的子 session 不参与 `SessionManagerConfig.auto_checkpoint`**。 + /// 即便 `auto_checkpoint = true`,dispatch_stream 内部直接调 `finalize_turn()` + /// 而非 `SessionManager::finalize_turn_stream()`,因此子 session 不会产生 + /// checkpoint。子 session 的终止状态保留在内存中,可通过 `sm.get(child_id)` 直接读取。 + /// 如果需要 checkpoint 持久化,请改用 `dispatch()`(同步版本,会自动 checkpoint)。 + /// + /// 失败时可能在任何时刻插入 `Error { child_id, error }`。 + pub async fn dispatch_stream( + self: &Arc, + parent_id: &str, + sub_agent: Arc, + task: impl Into, + config: DispatchConfig, + ) -> Result< + std::pin::Pin< + Box + Send>, + >, + EngineError, + > { + use futures_util::StreamExt; + use tokio::sync::mpsc; + + let task_str: String = task.into(); + let task_for_msg = task_str.clone(); + + // 1. create_child + let child_id = self.create_child(parent_id, sub_agent).await?; + + // 2. inherit_memory + if let Err(e) = self + .inherit_session_memory(parent_id, &child_id, &config) + .await + { + let _ = self.destroy(&child_id).await; + return Err(e); + } + + // 3. 创建 channel + let (tx, rx) = mpsc::unbounded_channel::(); + + // 4. 发送 ChildCreated + let _ = tx.send(SubTaskStreamEvent::ChildCreated { + child_id: child_id.clone(), + }); + + // 5. 获取 LLM 流 + let child_session = self.get(&child_id).await?; + let inner_stream = { + let mut guard = child_session.lock().await; + match guard.submit_turn_stream(task_str).await { + Ok(s) => s, + Err(e) => { + // submit_turn_stream 失败:清理子 session 并传播错误。 + // 此处不需要发送 SubTaskStreamEvent::Error, + // 因为调用方收到 Err 后 stream 句柄被丢弃,事件无人消费。 + let _ = self.destroy(&child_id).await; + return Err(EngineError::from(e)); + } + } + }; + + // 6. spawn task 消费流 + finalize + let sm = self.clone(); + let child_id_for_task = child_id.clone(); + let tx_for_task = tx.clone(); + + tokio::spawn(async move { + use crate::llm::stream::StreamEvent; + use crate::llm::types::message::Message; + use crate::llm::types::response_v2::MessageResponse; + use crate::llm::types::usage::CostTracker; + + // 消息重建:首条用户消息 + let mut new_messages: Vec = vec![Message::user_text(&task_for_msg)]; + let mut final_response: Option = None; + let mut stream_error: Option = None; + + let mut stream = inner_stream; + while let Some(event) = stream.next().await { + // 转发事件 + if tx_for_task + .send(SubTaskStreamEvent::Stream(event.clone())) + .is_err() + { + // receiver dropped — 立即停止 + return; + } + + // 重建消息 + match &event { + StreamEvent::ToolExecutionCompleted { + tool_name: _, + tool_call_id, + result_summary, + is_error, + } => { + new_messages.push(Message::tool_result( + tool_call_id.clone(), + result_summary.clone(), + *is_error, + )); + } + StreamEvent::MessageComplete { full_response } => { + final_response = Some(full_response.clone()); + } + StreamEvent::Error { message } => { + stream_error = Some(message.clone()); + } + _ => {} + } + } + + // 流结束后处理 + if let Some(err_msg) = stream_error { + tracing::error!( + child_id = %child_id_for_task, + error = %err_msg, + "dispatch_stream: LLM stream error" + ); + let _ = tx_for_task.send(SubTaskStreamEvent::Error { + child_id: child_id_for_task.clone(), + error: err_msg, + }); + let _ = sm.destroy(&child_id_for_task).await; + return; + } + + if let Some(response) = &final_response { + // 追加 assistant 消息 + new_messages.push(response.message.clone()); + + // 锁定子 session 调 finalize + let lock_result = { + let session = match sm.get(&child_id_for_task).await { + Ok(s) => s, + Err(e) => { + let _ = tx_for_task.send(SubTaskStreamEvent::Error { + child_id: child_id_for_task.clone(), + error: format!("session get failed: {e}"), + }); + return; + } + }; + let mut guard = session.lock().await; + guard + .finalize_turn(response, new_messages) + .await + }; + + if let Err(e) = lock_result { + let _ = tx_for_task.send(SubTaskStreamEvent::Error { + child_id: child_id_for_task.clone(), + error: format!("finalize failed: {e}"), + }); + return; + } + + // 读取 result_summary(子 session 可能被外部 destroy 容错处理) + let summary = match sm.get(&child_id_for_task).await { + Ok(session) => { + let guard = session.lock().await; + guard + .session_memory() + .get("result_summary") + .await + .ok() + .flatten() + } + Err(_) => { + // 子 session 已被外部销毁,无 summary 可读 + None + } + }; + + // 发送 Completed + let result = SubTaskResult { + child_id: child_id_for_task.clone(), + response: response.clone(), + usage: CostTracker::from(response.usage), + summary, + }; + let _ = tx_for_task.send(SubTaskStreamEvent::Completed(result)); + } else { + // 流结束但没有 final_response — 异常 + let _ = tx_for_task.send(SubTaskStreamEvent::Error { + child_id: child_id_for_task.clone(), + error: "stream ended without MessageComplete".to_string(), + }); + } + }); + + // 7. wrap receiver as Stream + let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(rx); + Ok(Box::pin(stream)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::agent::builder::AgentBuilder; + use crate::agent::runtime::RuntimeBundle; + use crate::engine::session_manager::SessionManager; + use crate::llm::hooks::HookExecutor; + use crate::llm::mock::MockProvider; + use crate::llm::types::Usage; + use crate::llm::types::message::{ContentBlock, Message}; + use crate::llm::types::response_v2::{MessageResponse, StopReason}; + use crate::memory::store::InMemoryStore; + use crate::tools::ToolRegistry; + + struct MockAgent { + name: String, + } + + impl MockAgent { + fn new(name: &str) -> Self { + Self { + name: name.to_string(), + } + } + } + + impl Agent for MockAgent { + fn name(&self) -> &str { + &self.name + } + fn system_prompt(&self) -> Option<&str> { + Some("test agent") + } + } + + fn assistant_text(text: &str) -> MessageResponse { + MessageResponse { + id: String::new(), + model: String::new(), + message: Message::Assistant { + content: vec![ContentBlock::Text { text: text.into() }], + }, + usage: Usage::from_input_output(8, 4), + stop_reason: StopReason::Stop, + extra: Default::default(), + } + } + + async fn make_manager_and_bundle() -> (Arc, Arc) { + let store: Arc = + Arc::new(InMemoryStore::new()); + let provider = Arc::new(MockProvider::new(vec![ + assistant_text("child response 1"), + assistant_text("child response 2"), + assistant_text("child response 3"), + ])); + let bundle = Arc::new( + AgentBuilder::new() + .provider(provider) + .tool_registry(Arc::new(ToolRegistry::new())) + .hook_executor(Arc::new(HookExecutor::new())) + .session_memory_backend(store.clone()) + .build() + .expect("RuntimeBundle build"), + ); + let sm = Arc::new(SessionManager::new(store)); + (sm, bundle) + } + + #[tokio::test] + async fn test_dispatch_basic() { + let (sm, bundle) = make_manager_and_bundle().await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child: Arc = Arc::new(MockAgent::new("child")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + let result = sm + .dispatch(&parent_id, child, "task", DispatchConfig::default()) + .await + .unwrap(); + + assert!(!result.child_id.is_empty()); + assert!(result.usage.total().total_tokens > 0); + + // 子 session 仍保留在 manager 中 + let children = sm.children(&parent_id).await.unwrap(); + assert_eq!(children, vec![result.child_id.clone()]); + } + + #[tokio::test] + async fn test_dispatch_bridge_keys_filter() { + let (sm, bundle) = make_manager_and_bundle().await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child: Arc = Arc::new(MockAgent::new("child")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + + // 父 session_memory 写入 3 个 keys + { + let session = sm.get(&parent_id).await.unwrap(); + let mut guard = session.lock().await; + guard.set_session_data("a", "1").await.unwrap(); + guard.set_session_data("b", "2").await.unwrap(); + guard.set_session_data("c", "3").await.unwrap(); + } + + // bridge_keys = Some(vec!["a", "b"]) → 只继承 a, b + let config = DispatchConfig { + bridge_keys: Some(vec!["a".to_string(), "b".to_string()]), + ..Default::default() + }; + let result = sm + .dispatch(&parent_id, child, "task", config) + .await + .unwrap(); + + // 验证子 session_memory 只继承了 a, b + let child_session = sm.get(&result.child_id).await.unwrap(); + let child_guard = child_session.lock().await; + assert_eq!( + child_guard.session_memory().get("a").await.unwrap(), + Some("1".to_string()) + ); + assert_eq!( + child_guard.session_memory().get("b").await.unwrap(), + Some("2".to_string()) + ); + assert_eq!(child_guard.session_memory().get("c").await.unwrap(), None); + } + + #[tokio::test] + async fn test_dispatch_memory_inheritance() { + let (sm, bundle) = make_manager_and_bundle().await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child: Arc = Arc::new(MockAgent::new("child")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + { + let session = sm.get(&parent_id).await.unwrap(); + let mut guard = session.lock().await; + guard + .set_session_data("user_tone", "professional") + .await + .unwrap(); + } + + // bridge_keys = Some(vec![]) → 全部继承 + let config = DispatchConfig { + bridge_keys: Some(vec![]), + ..Default::default() + }; + let result = sm + .dispatch(&parent_id, child, "task", config) + .await + .unwrap(); + + let child_session = sm.get(&result.child_id).await.unwrap(); + let child_guard = child_session.lock().await; + assert_eq!( + child_guard.session_memory().get("user_tone").await.unwrap(), + Some("professional".to_string()) + ); + } + + #[tokio::test] + async fn test_dispatch_failure_cleans_up_child() { + let (sm, bundle) = make_manager_and_bundle().await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child: Arc = Arc::new(MockAgent::new("child")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + let children_before = sm.children(&parent_id).await.unwrap(); + assert_eq!(children_before.len(), 0); + + // dispatch 到不存在的 parent_id + let result = sm + .dispatch("nonexistent_parent", child, "task", DispatchConfig::default()) + .await; + assert!(result.is_err()); + + // 验证原 parent 的 children 列表未变 + let children_after = sm.children(&parent_id).await.unwrap(); + assert_eq!(children_after.len(), 0); + } + + #[tokio::test] + async fn test_dispatch_invalid_parent_id() { + let (sm, _bundle) = make_manager_and_bundle().await; + let child: Arc = Arc::new(MockAgent::new("child")); + let result = sm + .dispatch("nonexistent_parent", child, "task", DispatchConfig::default()) + .await; + assert!(matches!(result, Err(EngineError::SessionNotFound(_)))); + } + + // ====== dispatch_all 测试 ====== + + /// 提供充足的 mock response(>= 3) + async fn make_manager_and_bundle_for_all(n: usize) -> (Arc, Arc) { + let store: Arc = + Arc::new(InMemoryStore::new()); + let responses: Vec<_> = (0..n) + .map(|i| assistant_text(&format!("response {i}"))) + .collect(); + let provider = Arc::new(MockProvider::new(responses)); + let bundle = Arc::new( + AgentBuilder::new() + .provider(provider) + .tool_registry(Arc::new(ToolRegistry::new())) + .hook_executor(Arc::new(HookExecutor::new())) + .session_memory_backend(store.clone()) + .build() + .expect("RuntimeBundle build"), + ); + let sm = Arc::new(SessionManager::new(store)); + (sm, bundle) + } + + /// 创建空 mock responses 的 manager 和 bundle —— 后续 dispatch 会触发 + /// "MockProvider: 预设响应已用完" 错误,可用于测试错误传播。 + async fn make_manager_and_bundle_empty_mock() -> (Arc, Arc) { + let store: Arc = + Arc::new(InMemoryStore::new()); + let provider = Arc::new(MockProvider::empty()); + let bundle = Arc::new( + AgentBuilder::new() + .provider(provider) + .tool_registry(Arc::new(ToolRegistry::new())) + .hook_executor(Arc::new(HookExecutor::new())) + .session_memory_backend(store.clone()) + .build() + .expect("RuntimeBundle build"), + ); + let sm = Arc::new(SessionManager::new(store)); + (sm, bundle) + } + + #[tokio::test] + async fn test_dispatch_all_parallel_success() { + let (sm, bundle) = make_manager_and_bundle_for_all(3).await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child1: Arc = Arc::new(MockAgent::new("child1")); + let child2: Arc = Arc::new(MockAgent::new("child2")); + let child3: Arc = Arc::new(MockAgent::new("child3")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + let results = sm + .dispatch_all( + &parent_id, + vec![ + (child1, "task1".to_string()), + (child2, "task2".to_string()), + (child3, "task3".to_string()), + ], + DispatchConfig::default(), + ) + .await; + + assert_eq!(results.len(), 3); + for r in &results { + assert!(r.is_ok(), "expected success, got {:?}", r); + } + + // 子 session 全部保留 + let children = sm.children(&parent_id).await.unwrap(); + assert_eq!(children.len(), 3); + } + + #[tokio::test] + async fn test_dispatch_all_partial_failure() { + // 仅 1 个 mock response,第 2、3 个会失败 + let (sm, bundle) = make_manager_and_bundle_for_all(1).await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child1: Arc = Arc::new(MockAgent::new("child1")); + let child2: Arc = Arc::new(MockAgent::new("child2")); + let child3: Arc = Arc::new(MockAgent::new("child3")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + let results = sm + .dispatch_all( + &parent_id, + vec![ + (child1, "task1".to_string()), + (child2, "task2".to_string()), + (child3, "task3".to_string()), + ], + DispatchConfig::default(), + ) + .await; + + assert_eq!(results.len(), 3); + // 第 1 个成功,第 2、3 个会因 mock 用完而失败 + assert!(results[0].is_ok(), "first should succeed"); + assert!(results[1].is_err(), "second should fail"); + assert!(results[2].is_err(), "third should fail"); + } + + #[tokio::test] + async fn test_dispatch_all_semaphore_serial() { + let (sm, bundle) = make_manager_and_bundle_for_all(2).await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child1: Arc = Arc::new(MockAgent::new("child1")); + let child2: Arc = Arc::new(MockAgent::new("child2")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + // max_concurrency=1 → 串行 + let config = DispatchConfig { + max_concurrency: 1, + ..Default::default() + }; + let results = sm + .dispatch_all( + &parent_id, + vec![(child1, "task1".to_string()), (child2, "task2".to_string())], + config, + ) + .await; + + assert_eq!(results.len(), 2); + assert!(results[0].is_ok()); + assert!(results[1].is_ok()); + + // 顺序应保持(先 task1 后 task2) + let children = sm.children(&parent_id).await.unwrap(); + assert_eq!(children.len(), 2); + } + + #[tokio::test] + async fn test_dispatch_all_empty() { + let (sm, bundle) = make_manager_and_bundle_for_all(0).await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + let results = sm + .dispatch_all(&parent_id, vec![], DispatchConfig::default()) + .await; + assert!(results.is_empty()); + + let children = sm.children(&parent_id).await.unwrap(); + assert!(children.is_empty()); + } + + // ====== dispatch_stream 测试 ====== + + use futures_util::StreamExt; + + #[tokio::test] + async fn test_dispatch_stream_event_sequence() { + let (sm, bundle) = make_manager_and_bundle().await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child: Arc = Arc::new(MockAgent::new("child")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + let mut stream = sm + .dispatch_stream(&parent_id, child, "task", DispatchConfig::default()) + .await + .unwrap(); + + let mut child_id_from_event: Option = None; + let mut saw_stream_event = false; + let mut completed: Option = None; + + while let Some(event) = stream.next().await { + match event { + SubTaskStreamEvent::ChildCreated { child_id } => { + child_id_from_event = Some(child_id); + } + SubTaskStreamEvent::Stream(_) => { + saw_stream_event = true; + } + SubTaskStreamEvent::Completed(r) => { + completed = Some(r); + break; + } + SubTaskStreamEvent::Error { child_id, error } => { + panic!("unexpected error: child_id={child_id}, error={error}"); + } + } + } + + let child_id = child_id_from_event.expect("ChildCreated should be received"); + assert!( + saw_stream_event, + "at least one Stream event should be received" + ); + let result = completed.expect("Completed should be received"); + assert_eq!(result.child_id, child_id); + assert!(result.usage.total().total_tokens > 0); + + // 子 session 仍保留 + let children = sm.children(&parent_id).await.unwrap(); + assert_eq!(children, vec![child_id]); + } + + #[tokio::test] + /// W6 测试:receiver drop 后 spawn task 正确退出,无 panic,无资源泄漏。 + /// + /// 验证策略: + /// 1. 立即 drop stream(spawn task 内部的 `tx.send()` 会因 receiver 关闭而失败) + /// 2. 等待足够时间让 spawn task 退出 + /// 3. 验证:子 session 仍存在(drop receiver 不会触发子 session 清理) + /// 这与 `Error { ... }` 事件路径不同——drop receiver 是"消费者失联", + /// spawn task 应静默退出而不清理子 session。 + /// 4. 验证:再次消费一个完整流验证 SessionManager 仍正常工作(无 panic 传播) + async fn test_dispatch_stream_receiver_dropped() { + let (sm, bundle) = make_manager_and_bundle().await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child1: Arc = Arc::new(MockAgent::new("child1")); + let child2: Arc = Arc::new(MockAgent::new("child2")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + + // 第一次 dispatch_stream 立即 drop + let first_child_id = { + let mut stream = sm + .dispatch_stream(&parent_id, child1, "task1", DispatchConfig::default()) + .await + .unwrap(); + + // 拉取第一个事件 (ChildCreated) 然后 drop + let mut child_id = None; + if let Some(SubTaskStreamEvent::ChildCreated { child_id: id }) = stream.next().await { + child_id = Some(id); + } + // stream 在此处 drop,spawn task 内部 `tx.send()` 后续会失败并 return + child_id.expect("ChildCreated should arrive before drop") + }; + + // 给 spawn task 充足时间退出 + tokio::time::sleep(std::time::Duration::from_millis(200)).await; + + // 验证:子 session 仍存在(receiver drop 不会清理子 session) + let result = sm.get(&first_child_id).await; + assert!( + result.is_ok(), + "child session should still exist after receiver drop (got: {:?})", + result.err() + ); + + // 验证:SessionManager 仍能正常 dispatch_stream + let mut stream2 = sm + .dispatch_stream(&parent_id, child2, "task2", DispatchConfig::default()) + .await + .unwrap(); + + let mut saw_completed = false; + while let Some(event) = stream2.next().await { + if matches!(event, SubTaskStreamEvent::Completed(_)) { + saw_completed = true; + break; + } + } + assert!( + saw_completed, + "second dispatch_stream should complete normally after first one was dropped" + ); + } + + #[tokio::test] + async fn test_dispatch_stream_finalize_completed() { + let (sm, bundle) = make_manager_and_bundle().await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child: Arc = Arc::new(MockAgent::new("child")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + let mut stream = sm + .dispatch_stream(&parent_id, child, "task", DispatchConfig::default()) + .await + .unwrap(); + + let mut completed: Option = None; + while let Some(event) = stream.next().await { + if let SubTaskStreamEvent::Completed(r) = event { + completed = Some(r); + break; + } + } + let result = completed.expect("Completed should arrive"); + + // 验证子 session 的 turn_index 已递增(finalize 副作用) + let child_session = sm.get(&result.child_id).await.unwrap(); + let child_guard = child_session.lock().await; + assert_eq!( + child_guard.turn_index(), + 1, + "turn_index should increment after finalize" + ); + } + + /// W5 测试:LLM 内部错误传播。 + /// MockProvider 空响应会触发 LlmError,dispatch_stream 应在 spawn task 内部 + /// 捕获到 `StreamEvent::Error { message }`,发送 `SubTaskStreamEvent::Error` 事件, + /// 然后清理子 session。 + #[tokio::test] + async fn test_dispatch_stream_llm_error_propagation() { + let (sm, bundle) = make_manager_and_bundle_empty_mock().await; + let parent: Arc = Arc::new(MockAgent::new("parent")); + let child: Arc = Arc::new(MockAgent::new("child")); + + let parent_id = sm.create(parent, bundle.clone()).await.unwrap(); + let mut stream = sm + .dispatch_stream(&parent_id, child, "task", DispatchConfig::default()) + .await + .unwrap(); + + let mut child_id_from_event: Option = None; + let mut error_received: Option<(String, String)> = None; + + while let Some(event) = stream.next().await { + match event { + SubTaskStreamEvent::ChildCreated { child_id } => { + child_id_from_event = Some(child_id); + } + SubTaskStreamEvent::Error { child_id, error } => { + error_received = Some((child_id, error)); + // 收到 Error 事件后流可能结束,退出循环 + break; + } + _ => {} + } + } + + // 验证:收到 Error 事件 + let (err_child_id, err_message) = + error_received.expect("Error event should be received"); + assert_eq!( + Some(err_child_id.as_str()), + child_id_from_event.as_deref(), + "Error child_id should match ChildCreated" + ); + assert!( + err_message.contains("预设响应已用完") || err_message.contains("MockProvider"), + "Error should mention MockProvider, got: {err_message}" + ); + + // 验证:子 session 已被清理(destroy 在 spawn task 末尾执行) + // 给清理一点时间完成 + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + let result = sm.get(&err_child_id).await; + assert!( + result.is_err(), + "child session should be destroyed after error" + ); + } +} diff --git a/src/engine/switch.rs b/src/engine/switch.rs new file mode 100644 index 0000000..9082992 --- /dev/null +++ b/src/engine/switch.rs @@ -0,0 +1,222 @@ +//! Agent 角色热切换(Phase 18)。 +//! +//! 提供 `SessionManager::switch_agent()`,运行时替换 session 绑定的 Agent, +//! 保留 slot 历史 / turn_index / session_memory / cost_so_far。 +//! +//! **未来扩展方向**(roadmap 备注): +//! - v0.4 可考虑提供 `switch_agent_with_rollback`,在切换前自动 checkpoint +//! - 可在 `SessionMeta` 中记录 `previous_agent_name` 支持审计历史 +//! - 可新增 `switch_history` API 暴露切换时间序列 + +use std::sync::Arc; + +use crate::agent::Agent; +use crate::engine::error::EngineError; +use crate::engine::session_manager::{SessionManager, SessionMeta}; + +impl SessionManager { + /// 热切换指定 session 的 Agent 角色。 + /// + /// # 行为 + /// + /// - **保留上下文**:slot 历史 / turn_index / session_memory / cost_so_far 全部保留 + /// - **更新 SessionMeta**:`agent_name` 替换为新 agent,`created_at`/`parent_id` 保持原始 + /// - **不自动 checkpoint**:与 `auto_checkpoint` 语义一致(仅 `submit_turn`/`finalize_turn` 触发) + /// + /// # 注意 + /// + /// 切换后**新的 system_prompt 将与已有对话历史共存**。建议在切换后 + /// 发送一条明确的上下文过渡提示(如"你现在以新角色 X 的身份继续对话") + /// 作为切换后的首条输入,以避免 LLM 误解对话历史。 + /// + /// # 安全提示 + /// + /// `AgentSession.agent` 是 `pub` 字段可直接访问。**绕过 `switch_agent` + /// 直接修改会导致 SessionMeta 中的 `agent_name` 与内存状态不一致**, + /// 请始终使用此方法。 + pub async fn switch_agent( + &self, + session_id: &str, + new_agent: Arc, + ) -> Result<(), EngineError> { + // 1. get session(RwLock 读锁,返回后释放) + let session = self.get(session_id).await?; + + // 2. lock Mutex,替换 agent,读 name + turn_index + let (agent_name, turn_index) = { + let mut guard = session.lock().await; + guard.agent = new_agent; + (guard.agent.name().to_string(), guard.turn_index()) + }; // 释放 Mutex + + // 3. 读取原始 SessionMeta(用于保留 created_at / parent_id) + let existing_meta = self + .load_session_meta(session_id) + .await? + .ok_or_else(|| EngineError::SessionNotFound(session_id.to_string()))?; + + // 4. 构造新 meta 并持久化(I/O,无锁) + let meta = SessionMeta { + session_id: session_id.to_string(), + agent_name, + parent_id: existing_meta.parent_id, + created_at: existing_meta.created_at, + turn_count: turn_index, + }; + self.save_session_meta(&meta).await?; + + tracing::info!( + session_id = %session_id, + agent_name = %meta.agent_name, + previous_agent = %existing_meta.agent_name, + "agent switched" + ); + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::agent::Agent; + use crate::agent::builder::AgentBuilder; + use crate::agent::runtime::RuntimeBundle; + use crate::engine::session_manager::SessionManager; + use crate::llm::hooks::HookExecutor; + use crate::llm::mock::MockProvider; + use crate::llm::types::Usage; + use crate::llm::types::message::{ContentBlock, Message}; + use crate::llm::types::response_v2::{MessageResponse, StopReason}; + use crate::memory::store::InMemoryStore; + use crate::tools::ToolRegistry; + use std::sync::Arc; + + /// 测试用 MockAgent(name + system_prompt 可控) + struct MockAgent { + name: String, + system_prompt: String, + } + + impl MockAgent { + fn new(name: &str, system_prompt: &str) -> Self { + Self { + name: name.to_string(), + system_prompt: system_prompt.to_string(), + } + } + } + + impl Agent for MockAgent { + fn name(&self) -> &str { + &self.name + } + fn system_prompt(&self) -> Option<&str> { + Some(&self.system_prompt) + } + } + + fn assistant_text(text: &str) -> MessageResponse { + MessageResponse { + id: String::new(), + model: String::new(), + message: Message::Assistant { + content: vec![ContentBlock::Text { text: text.into() }], + }, + usage: Usage::from_input_output(8, 4), + stop_reason: StopReason::Stop, + extra: Default::default(), + } + } + + async fn make_manager_and_bundle() -> (Arc, Arc) { + let store: Arc = + Arc::new(InMemoryStore::new()); + let provider = Arc::new(MockProvider::new(vec![ + assistant_text("response1"), + assistant_text("response2"), + assistant_text("response3"), + ])); + let bundle = Arc::new( + AgentBuilder::new() + .provider(provider) + .tool_registry(Arc::new(ToolRegistry::new())) + .hook_executor(Arc::new(HookExecutor::new())) + .session_memory_backend(store.clone()) + .build() + .expect("RuntimeBundle build"), + ); + let sm = Arc::new(SessionManager::new(store)); + (sm, bundle) + } + + #[tokio::test] + async fn test_switch_agent_basic() { + let (sm, bundle) = make_manager_and_bundle().await; + let agent_a: Arc = Arc::new(MockAgent::new("agent_a", "I am A")); + let agent_b: Arc = Arc::new(MockAgent::new("agent_b", "I am B")); + + let sid = sm.create(agent_a, bundle).await.unwrap(); + sm.switch_agent(&sid, agent_b).await.unwrap(); + + let session = sm.get(&sid).await.unwrap(); + let guard = session.lock().await; + assert_eq!(guard.agent.name(), "agent_b"); + } + + #[tokio::test] + async fn test_switch_agent_preserves_context() { + let (sm, bundle) = make_manager_and_bundle().await; + let agent_a: Arc = Arc::new(MockAgent::new("agent_a", "I am A")); + let agent_b: Arc = Arc::new(MockAgent::new("agent_b", "I am B")); + + let sid = sm.create(agent_a, bundle).await.unwrap(); + + // 在 switch 前写入 session_data 并提交一轮 + { + let session = sm.get(&sid).await.unwrap(); + let mut guard = session.lock().await; + guard + .set_session_data("key1", "value1") + .await + .unwrap(); + } + sm.submit_turn(&sid, "hello").await.unwrap(); + + // switch + sm.switch_agent(&sid, agent_b).await.unwrap(); + + // 验证 turn_index 保留 + let session = sm.get(&sid).await.unwrap(); + let guard = session.lock().await; + assert_eq!(guard.turn_index(), 1, "turn_index should be preserved"); + // 验证 session_memory 保留 + let val = guard.session_memory().get("key1").await.unwrap(); + assert_eq!(val, Some("value1".to_string())); + } + + #[tokio::test] + async fn test_switch_agent_updates_session_meta() { + let (sm, bundle) = make_manager_and_bundle().await; + let agent_a: Arc = Arc::new(MockAgent::new("agent_a", "I am A")); + let agent_b: Arc = Arc::new(MockAgent::new("agent_b", "I am B")); + + let sid = sm.create(agent_a, bundle).await.unwrap(); + sm.switch_agent(&sid, agent_b).await.unwrap(); + + // 通过 load_session_meta 验证持久化 + let meta = sm.load_session_meta(&sid).await.unwrap(); + assert!(meta.is_some(), "SessionMeta should persist"); + let meta = meta.unwrap(); + assert_eq!(meta.agent_name, "agent_b"); + // created_at / parent_id 保持 + assert_eq!(meta.parent_id, None, "parent_id should remain None"); + } + + #[tokio::test] + async fn test_switch_agent_session_not_found() { + let (sm, _bundle) = make_manager_and_bundle().await; + let agent: Arc = Arc::new(MockAgent::new("agent_a", "I am A")); + let result = sm.switch_agent("nonexistent_session_id", agent).await; + assert!(matches!(result, Err(EngineError::SessionNotFound(_)))); + } +} diff --git a/src/llm/types/usage.rs b/src/llm/types/usage.rs index bdf262e..efa6dcd 100644 --- a/src/llm/types/usage.rs +++ b/src/llm/types/usage.rs @@ -61,6 +61,14 @@ impl CostTracker { } } +impl From for CostTracker { + fn from(usage: Usage) -> Self { + CostTracker { + accumulated: usage, + } + } +} + impl Usage { pub fn from_input_output(input: u32, output: u32) -> Self { let total = input.saturating_add(output);