From fc3a95b152b9349e74c42d5f5a4d58a9203dcb06 Mon Sep 17 00:00:00 2001 From: oudecheng <13802883547@139.com> Date: Thu, 13 Aug 2026 22:12:20 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=BC=82=E6=AD=A5=E5=AD=90=E4=BB=A3?= =?UTF-8?q?=E7=90=86=20v8=20=E5=AE=9E=E7=8E=B0=E4=B8=8E=20/stop=20?= =?UTF-8?q?=E5=8F=96=E6=B6=88=E6=B5=81=E7=A8=8B=E4=BF=AE=E5=A4=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 后端核心实现: - 新增 wait_coordinator:释放/重获取 serial_lock 的 in-tool waiting 模式,替代旧的 break-exit 方案 - 新增 wait_for_subagents 工具:select! 等待子代理完成/用户消息/超时,支持 try_drain 批量消费 - task 工具异步 spawn 路径:CancellationToken 注册表 + RAII guard + Semaphore 并发限流 - process exit 安全检查:pending 子代理存在时抑制 ExecutionCompleted - 崩溃恢复:启动时标记 running->interrupted,history 加载时对账占位 /stop 取消流程修复: - wait_coordinator select! 添加 cancel 分支,完整清理状态 - ToolContext 注入 cancel_rx(watch::Receiver clone) - agent_loop 工具执行 select! 对 wait 工具跳过竞速,防止 drop coordinator 清理逻辑 - 子代理完成状态通过 execution_completed metadata 传播 存储层: - pending_subagents 表 + 条件 UPDATE - mark_all_running_as_interrupted 崩溃恢复 --- src/agent/agent_loop.rs | 28 +- src/command/handlers/stop_execution.rs | 33 +- src/config/mod.rs | 16 + src/gateway/agent_factory.rs | 22 +- src/gateway/execution.rs | 72 +++- src/gateway/mod.rs | 31 +- src/gateway/processor.rs | 158 ++++++-- src/gateway/runtime.rs | 14 +- src/gateway/session.rs | 239 +++++++++++- src/gateway/session_history.rs | 113 ++++++ src/gateway/tool_registry_factory.rs | 7 +- src/gateway/wait_coordinator.rs | 291 +++++++++++++++ src/gateway/ws.rs | 1 + src/protocol/mod.rs | 6 + src/protocol/ws_adapter.rs | 2 + src/storage/migrations.rs | 25 ++ src/storage/mod.rs | 143 +++++++- src/storage/records.rs | 25 ++ src/tools/mod.rs | 4 +- src/tools/task/error.rs | 4 + src/tools/task/mod.rs | 4 +- src/tools/task/runtime.rs | 488 ++++++++++++++++++++++--- src/tools/task/types.rs | 47 ++- src/tools/todo_read.rs | 1 + src/tools/todo_write.rs | 1 + src/tools/traits.rs | 98 ++++- src/tools/wait_tool.rs | 250 +++++++++++++ 27 files changed, 2010 insertions(+), 113 deletions(-) create mode 100644 src/gateway/wait_coordinator.rs create mode 100644 src/tools/wait_tool.rs diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index 48a8301..01f5f93 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -1447,10 +1447,29 @@ impl AgentLoop { .await; // Execute tools and add results to messages - // 工具执行与取消信号竞速:取消时 drop join_all 或 sequential future, - // 未完成的工具调用被丢弃。 - let tool_results = if self.cancel_token.is_some() { + // + // 取消竞速策略: + // - 包含 wait_for_subagents 时:不使用 select!,直接 await execute_tools。 + // 原因:coordinator.wait() 在 select! 返回后有不可中断的清理逻辑 + // (步骤 6-8:重获取 serial_lock、回填 guard_slot、清除 is_waiting)。 + // 若 agent_loop 的 select! 在清理期间 drop execute_tools, + // 会导致 is_waiting=true 永久残留、guard_slot 为空、serial_lock 未持有, + // 后续所有用户消息走注入路径但 wakeup 无接收者 → 系统永久卡死。 + // cancel 由 coordinator 内部 select! 的 cancel 分支处理,清理不会被打断。 + // + // - 不含 wait_for_subagents 时:保留 select! 竞速,允许 /stop 中断 + // 长时间运行的工具(如 MCP HTTP 请求)。 + let has_wait_tool = response + .tool_calls + .iter() + .any(|tc| tc.name == "wait_for_subagents"); + + let tool_results = if self.cancel_token.is_some() && !has_wait_tool { tokio::select! { + biased; + results = self.execute_tools(&response.tool_calls) => { + results + } _ = self.cancel_signal() => { // 为所有 tool_calls 补充取消结果,避免孤立 assistant(tool_calls) for tool_call in &response.tool_calls { @@ -1468,9 +1487,6 @@ impl AgentLoop { self.emit_live_tool_call_message(cancel.final_response.clone()).await; return Ok(cancel); } - results = self.execute_tools(&response.tool_calls) => { - results - } } } else { self.execute_tools(&response.tool_calls).await diff --git a/src/command/handlers/stop_execution.rs b/src/command/handlers/stop_execution.rs index a191de5..b79451d 100644 --- a/src/command/handlers/stop_execution.rs +++ b/src/command/handlers/stop_execution.rs @@ -1,4 +1,5 @@ use async_trait::async_trait; +use std::sync::Arc; use crate::command::Command; use crate::command::context::CommandContext; @@ -6,18 +7,27 @@ use crate::command::handler::{CommandHandler, CommandMetadata}; use crate::command::response::{CommandError, CommandResponse, MessageKind}; use crate::gateway::cancel_manager::CancelManager; use crate::gateway::session::SessionManager; +use crate::tools::SubAgentRuntime; /// 处理 StopExecution 命令:按话题取消当前正在执行的 Agent。 +/// +/// 取消传播:同时取消该 topic 下所有正在运行的异步子代理(通过 CancellationToken)。 pub struct StopExecutionCommandHandler { cancel_manager: CancelManager, session_manager: SessionManager, + subagent_executor: Option>, } impl StopExecutionCommandHandler { - pub fn new(cancel_manager: CancelManager, session_manager: SessionManager) -> Self { + pub fn new( + cancel_manager: CancelManager, + session_manager: SessionManager, + subagent_executor: Option>, + ) -> Self { Self { cancel_manager, session_manager, + subagent_executor, } } } @@ -89,9 +99,26 @@ impl CommandHandler for StopExecutionCommandHandler { let cancelled = self.cancel_manager.cancel_by_topic(&topic_id).await; - if cancelled { + // 取消传播:同时取消该 topic 下所有正在运行的异步子代理 + let cancelled_subagents = if let Some(ref executor) = self.subagent_executor { + executor.cancel_pending_for_topic(&topic_id).await + } else { + 0 + }; + + if cancelled || cancelled_subagents > 0 { + let msg = if cancelled && cancelled_subagents > 0 { + format!( + "正在停止当前任务及 {} 个后台子代理...", + cancelled_subagents + ) + } else if cancelled { + "正在停止当前任务...".to_string() + } else { + format!("正在停止 {} 个后台子代理...", cancelled_subagents) + }; Ok(CommandResponse::success(ctx.request_id) - .with_message(MessageKind::Notification, "正在停止当前任务...")) + .with_message(MessageKind::Notification, msg)) } else { Ok(CommandResponse::success(ctx.request_id) .with_message(MessageKind::Notification, "当前没有正在执行的任务")) diff --git a/src/config/mod.rs b/src/config/mod.rs index 8d231e8..9a2a834 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -337,6 +337,12 @@ pub struct TaskConfig { pub allowed_tools: Vec, #[serde(default = "default_task_max_nesting_depth")] pub max_nesting_depth: u32, + /// 异步子代理最大并发数(Semaphore 限流,仅主 agent 顶层 spawn 生效) + #[serde(default = "default_task_max_concurrent")] + pub max_concurrent: usize, + /// wait_for_subagents 工具默认超时(秒),LLM 可通过参数覆盖 + #[serde(default = "default_task_wait_default_timeout_secs")] + pub wait_default_timeout_secs: u64, } fn default_task_enabled() -> bool { @@ -355,6 +361,14 @@ fn default_task_max_nesting_depth() -> u32 { 2 } +fn default_task_max_concurrent() -> usize { + 8 +} + +fn default_task_wait_default_timeout_secs() -> u64 { + 60 +} + fn default_task_allowed_tools() -> Vec { vec![ "read".to_string(), @@ -380,6 +394,8 @@ impl Default for TaskConfig { ttl_hours: default_task_ttl_hours(), allowed_tools: default_task_allowed_tools(), max_nesting_depth: default_task_max_nesting_depth(), + max_concurrent: default_task_max_concurrent(), + wait_default_timeout_secs: default_task_wait_default_timeout_secs(), } } } diff --git a/src/gateway/agent_factory.rs b/src/gateway/agent_factory.rs index 123e8b9..1991290 100644 --- a/src/gateway/agent_factory.rs +++ b/src/gateway/agent_factory.rs @@ -1,5 +1,7 @@ use std::sync::Arc; +use tokio::sync::mpsc; + use crate::agent::context_compressor::ContextCompressor; use crate::agent::{AgentError, AgentLoop, AgentRuntimeConfig, CompositeSystemPromptProvider, SystemPromptProvider}; use crate::config::{CompactionConfig, LLMProviderConfig, ModelResolver}; @@ -14,7 +16,8 @@ use crate::skills::{SkillPromptProvider, SkillRuntime}; use crate::storage::PromptInjectionRepository; use crate::storage::persistent_session_id; use crate::tools::task::runtime::{SubagentPromptProvider, SubagentRuntime}; -use crate::tools::{ToolContext, ToolRegistry}; +use crate::tools::task::SubagentResult; +use crate::tools::{ToolContext, ToolRegistry, WaitCoordinator}; /// 构建与 Agent 实际使用的完全一致的组合系统提示词 Provider。 /// @@ -76,6 +79,11 @@ pub(crate) struct AgentBuildRequest<'a> { pub(crate) cancel_token: Option>, /// 端到端追踪 ID(从 InboundMessage 继承,注入 ToolContext 供 tool 执行路径日志关联) pub(crate) trace_id: Option, + /// 异步子代理完成队列的 sender(按 topic 隔离)。 + /// 仅主 agent 有值:TaskTool 据此在子代理完成时发送 SubagentResult。 + pub(crate) sub_done_sender: Option>, + /// wait_for_subagents 工具的协调器(仅主 agent 有值)。 + pub(crate) wait_coordinator: Option>, } impl AgentFactory { @@ -234,6 +242,12 @@ impl AgentFactory { .notification_chat_id .unwrap_or(request.session_chat_id); // 构建上下文压缩器(参数内聚到 ContextCompressor,CompactionConfig 注入) + // 注入取消信号 receiver 的 clone 到 ToolContext, + // 供 wait_for_subagents 工具传递给 coordinator.wait() 的 select!。 + // watch::Receiver::clone() 创建共享同一 sender 的新 receiver, + // 各 receiver 的 has_changed()/changed() 状态独立,互不影响。 + let cancel_rx_for_context = request.cancel_token.as_ref().map(|rx| rx.clone()); + let runtime_config = AgentRuntimeConfig::from(effective_provider_config.clone()); let compressor = Arc::new(self.build_compressor(&runtime_config)); let mut agent = agent @@ -253,6 +267,12 @@ impl AgentFactory { // 注入专家 capability,TaskTool 据此强制校验子代理白/黑名单 parent_capability: expert_capability.clone(), trace_id: request.trace_id.clone(), + // 注入异步子代理完成队列 sender(按 topic 隔离) + sub_done_sender: request.sub_done_sender.clone(), + // 注入 wait 协调器(封装释放/重获取 serial_lock 逻辑) + wait_coordinator: request.wait_coordinator.clone(), + // 注入取消信号 receiver clone(供 wait 工具的 cancel 检查) + cancel_rx: cancel_rx_for_context, }) .with_compressor(Some(compressor)); // 注入观测器(依赖注入,agent_loop 只认 Observer trait) diff --git a/src/gateway/execution.rs b/src/gateway/execution.rs index 2e2010a..ff3ea4c 100644 --- a/src/gateway/execution.rs +++ b/src/gateway/execution.rs @@ -15,6 +15,8 @@ use tokio::sync::Mutex; use super::compaction::schedule_background_history_compaction; use super::message_prepare::enrich_user_content_with_media_refs; use super::session::Session; +use super::wait_coordinator::SessionWaitCoordinator; +use crate::tools::WaitCoordinator; /// 空的 EmittedMessageHandler,不转发消息,仅配合 PersistingEmittedMessageHandler 做持久化。 struct NoOpEmittedMessageHandler; @@ -284,15 +286,32 @@ impl AgentExecutionService { // 获取该 topic 的串行锁(通过短暂获取 session 锁) // 同一 topic 的消息处理必须串行执行,防止并发 loop 操作同一历史的不同快照 // 不同 topic 之间互不阻塞,支持多话题并发执行 - let serial_lock = { + let (serial_lock, store, lock_key) = { let mut session_guard = request.session.lock().await; - let lock_key = request.topic_id.as_deref().unwrap_or(request.chat_id); - session_guard.topic_serial_lock(lock_key) + let lock_key = request + .topic_id + .as_deref() + .unwrap_or(request.chat_id) + .to_string(); + session_guard.ensure_sub_done_channel(&lock_key); + ( + session_guard.topic_serial_lock(&lock_key), + session_guard.session_store(), + lock_key, + ) }; // 等待该 topic 的前一条消息处理完成(含压缩) // await 串行锁时不持有 session 锁,其他 topic 的消息可以正常处理 - let _serial_guard = serial_lock.lock().await; + // 使用 lock_owned 获取 OwnedMutexGuard,存入 guard_slot 供 wait_coordinator 释放/重获取 + // 注意:lock_owned 消费 Arc,需 clone 保留 serial_lock 供 coordinator 使用 + let serial_guard = serial_lock.clone().lock_owned().await; + + // guard_slot:wait_coordinator 通过此 slot 释放/重获取 serial_lock。 + // 正常执行时 guard 留在 slot 中(锁持有);wait 工具调用时 take guard 释放锁, + // select! 等待结束后重获取锁并回填新 guard。 + // guard_slot 作为 Arc 共享于执行路径与 coordinator,二者全部 drop 时 guard 才释放锁。 + let guard_slot = Arc::new(Mutex::new(Some(serial_guard))); let (history, agent, user_message, user_message_count, original_topic_id) = { let mut session_guard = request.session.lock().await; @@ -340,12 +359,27 @@ impl AgentExecutionService { let history = session_guard.get_or_create_history(history_key).clone(); session_guard.record_skill_offer(request.chat_id)?; + // 创建 wait 协调器(封装释放/重获取 serial_lock + select! 等待逻辑)。 + // 仅主 agent 注入;coordinator 通过 guard_slot 释放/重获取 serial_lock, + // 使 wait_for_subagents 工具能在等待期间让 process_one 注入用户消息。 + let wait_coordinator: Option> = { + let coordinator = SessionWaitCoordinator::new( + request.session.clone(), + guard_slot.clone(), + serial_lock.clone(), + store.clone(), + lock_key.clone(), + ); + Some(Arc::new(coordinator)) + }; + let mut agent = session_guard.create_agent( request.chat_id, Some(request.sender_id), Some(&user_message.id), original_topic_id.as_deref(), request.trace_id, + wait_coordinator, )?; if let Some(handler) = request.live_emitter.clone() { agent = agent.with_emitted_message_handler(handler); @@ -408,17 +442,26 @@ impl AgentExecutionService { // 获取该 topic 的串行锁(与普通消息路径共享,保证串行执行) // 定时任务由调度器触发,无用户消息竞态;在锁前一次性捕获 topic_id, // 锁后复用同一值作为 original_topic_id,保证锁键与写入目标一致。 - let (serial_lock, lock_time_topic_id) = { + let (serial_lock, session_store, lock_key, lock_time_topic_id) = { let mut session_guard = request.session.lock().await; let tid = session_guard .current_topic(request.chat_id) .map(|s| s.to_string()); - let lock_key = tid.as_deref().unwrap_or(request.chat_id); - (session_guard.topic_serial_lock(lock_key), tid) + let lock_key = tid.as_deref().unwrap_or(request.chat_id).to_string(); + session_guard.ensure_sub_done_channel(&lock_key); + ( + session_guard.topic_serial_lock(&lock_key), + session_guard.session_store(), + lock_key, + tid, + ) }; // 等待该 topic 的前一条消息处理完成(含压缩) - let _serial_guard = serial_lock.lock().await; + // 使用 lock_owned 获取 OwnedMutexGuard,存入 guard_slot 供 wait_coordinator 释放/重获取 + // 注意:lock_owned 消费 Arc,需 clone 保留 serial_lock 供 coordinator 使用 + let serial_guard = serial_lock.clone().lock_owned().await; + let guard_slot = Arc::new(Mutex::new(Some(serial_guard))); let ( history, @@ -475,6 +518,18 @@ impl AgentExecutionService { let history = session_guard.get_or_create_history(history_key).clone(); session_guard.record_skill_offer(request.chat_id)?; + // 创建 wait 协调器(与普通消息路径一致,支持定时任务中 spawn 异步子代理) + let wait_coordinator: Option> = { + let coordinator = SessionWaitCoordinator::new( + request.session.clone(), + guard_slot.clone(), + serial_lock.clone(), + session_store.clone(), + lock_key.clone(), + ); + Some(Arc::new(coordinator)) + }; + let agent = session_guard.create_agent_with_provider_config( request.chat_id, request.notification_chat_id, // 传入真实 chat_id @@ -483,6 +538,7 @@ impl AgentExecutionService { request.provider_config.clone(), original_topic_id.as_deref(), &request.trace_id, + wait_coordinator, )?; // 获取 store 和 session_id,用于构造消息持久化 handler diff --git a/src/gateway/mod.rs b/src/gateway/mod.rs index 5d86bdb..3d912c8 100644 --- a/src/gateway/mod.rs +++ b/src/gateway/mod.rs @@ -28,6 +28,7 @@ pub mod session_pool; pub mod static_files; pub mod tool_prompt_provider; pub mod tool_registry_factory; +pub mod wait_coordinator; pub mod ws; use axum::{Router, middleware, routing}; @@ -70,6 +71,8 @@ pub struct GatewayState { pub skills: Arc, pub experts: Arc, pub subagent_runtime: Arc, + /// 异步子代理执行器(DefaultSubAgentRuntime),用于取消传播等操作 + pub subagent_executor: Option>, /// per-session 的用户模型选择(覆盖专家配置) pub model_selections: Arc, /// Prometheus metrics handle(/metrics 端点渲染用)。 @@ -105,7 +108,7 @@ impl GatewayState { mcp_servers: config.mcp_servers.clone(), }; - let (session_manager, task_repository, mcp_manager, subagent_runtime, model_selections) = + let (session_manager, task_repository, mcp_manager, subagent_runtime, model_selections, subagent_executor) = build_session_manager_with_sender( agent_prompt_reinject_every, show_tool_results, @@ -150,6 +153,7 @@ impl GatewayState { skills, experts, subagent_runtime, + subagent_executor, model_selections, prometheus_handle, }) @@ -178,6 +182,7 @@ impl GatewayState { semaphore, provider_config, self.cancel_manager.clone(), + self.subagent_executor.clone(), ); tokio::spawn(inbound_processor.run()); @@ -213,6 +218,30 @@ pub async fn run( let state = Arc::new(GatewayState::from_config(config, restart_tx)?); + // ── 崩溃恢复:标记中断的异步子代理 ── + // 服务器重启后,之前 spawn 的异步子代理进程已不存在, + // 将 pending_subagents 表中所有 status='running' 的记录标记为 'interrupted'。 + // 下次 wait_for_subagents 调用时,这些 task_id 不会出现在 pending 列表中, + // agent 可据此判断子代理未正常完成。 + match state.session_manager.store().mark_all_running_as_interrupted() { + Ok(0) => { + tracing::info!("Crash recovery: no interrupted subagents to recover"); + } + Ok(n) => { + tracing::info!( + recovered_count = n, + "Crash recovery: marked {} running subagents as interrupted (server restarted)", + n + ); + } + Err(e) => { + tracing::error!( + error = %e, + "Crash recovery: failed to mark interrupted subagents" + ); + } + } + // Get provider config for channels let cfg = state.config.read().await; let provider_config = cfg.get_provider_config("default")?; diff --git a/src/gateway/processor.rs b/src/gateway/processor.rs index b445175..40e91c4 100644 --- a/src/gateway/processor.rs +++ b/src/gateway/processor.rs @@ -29,6 +29,7 @@ use crate::storage::persistent_session_id; use crate::topic_description::generate_topic_description; use super::session::{BusToolCallEmitter, SessionManager}; +use super::message_prepare::enrich_user_content_with_media_refs; #[derive(Clone)] pub struct InboundProcessor { @@ -48,6 +49,7 @@ impl InboundProcessor { semaphore: Arc, provider_config: LLMProviderConfig, cancel_manager: CancelManager, + subagent_executor: Option>, ) -> Self { // 创建命令路由器并注册处理器 let mut command_router = CommandRouter::new(); @@ -118,6 +120,7 @@ impl InboundProcessor { command_router.register(Box::new(StopExecutionCommandHandler::new( cancel_manager.clone(), session_manager.clone(), + subagent_executor, ))); Self { @@ -312,6 +315,86 @@ impl InboundProcessor { if let Some(ref topic_id) = current_topic { emitter_metadata.insert("topic_id".to_string(), topic_id.clone()); } + + // ── 异步子代理等待注入路径 ── + // 当主 agent 正在 wait_for_subagents 中等待(已释放 serial_lock、is_waiting=true)时, + // 新用户消息不应启动新的 agent loop,而应注入 history 并唤醒等待中的 agent。 + // + // 流程: + // 1. 获取 serial_lock(若 agent 正常运行则阻塞;若 agent 在 wait 中则立即获取) + // 2. 检查 is_waiting:true → 注入 + wakeup + return;false → 释放锁走正常路径 + // + // 安全性:is_waiting 在持锁状态下检查,wait_coordinator 清除 is_waiting 需先重获取锁, + // 两者互斥,无 TOCTOU。 + if let Some(ref topic_id) = current_topic { + if let Some(session) = self.session_manager.get(&inbound.channel).await { + let lock_key = topic_id.clone(); + + // 获取 serial_lock Arc(短暂持有 session 锁) + let serial_lock = { + let mut g = session.lock().await; + g.ensure_sub_done_channel(&lock_key); + g.topic_serial_lock(&lock_key) + }; + + // 阻塞获取 serial_lock + // - agent 正常运行:阻塞至其完成(天然串行化) + // - agent 在 wait 中:wait 已释放锁,可立即获取 + let _inject_guard = serial_lock.clone().lock_owned().await; + + // 检查 is_waiting(持锁状态下安全) + let is_waiting = { + let g = session.lock().await; + g.is_waiting(&lock_key) + }; + + if is_waiting { + // Agent 正在 wait_for_subagents 中等待 → 注入用户消息 + 唤醒 + tracing::info!( + topic_id = %lock_key, + "Topic is in waiting state, injecting user message and waking up agent" + ); + + let wakeup = { + let mut g = session.lock().await; + // 确保 session 和 chat 已加载 + g.ensure_persistent_session(&inbound.chat_id)?; + g.ensure_chat_loaded(&inbound.chat_id, Some(&lock_key))?; + + // 构造用户消息(与 prepare_and_execute_message 一致的处理流程) + let media_refs: Vec = inbound + .media + .iter() + .map(|m| m.path.clone()) + .collect(); + let enriched_content = + enrich_user_content_with_media_refs(&inbound.content, &media_refs)?; + let user_message = + g.create_user_message(&enriched_content, media_refs); + g.append_persisted_message( + &inbound.chat_id, + Some(&lock_key), + user_message, + )?; + + // 获取 wakeup 信号 + g.wait_wakeup(&lock_key) + }; + + // 唤醒等待中的 agent(wait_coordinator 的 select! 会捕获此通知) + wakeup.notify_one(); + + // _inject_guard 在此处 drop → 释放 serial_lock + // wait_coordinator 重获取锁后继续处理(history 已包含新用户消息) + // + // 跳过 handle_message / cancel 注册 / execution_completed, + // 因为等待中的 agent 会处理这条消息。 + return Ok(()); + } + // is_waiting=false:_inject_guard drop 释放锁,走正常 handle_message 路径 + } + } + let live_emitter = Arc::new(PersistingEmittedMessageHandler::new( BusToolCallEmitter::new( self.bus.clone(), @@ -497,31 +580,58 @@ impl InboundProcessor { self.cancel_manager.remove_by_topic(topic_id).await; } - // 发送执行完成信号,通知前端可以停止 loading 状态 - // 无论成功还是失败都发送,确保前端状态正确 - let mut completion_metadata = inbound.forwarded_metadata.clone(); - if let Some(ref topic_id) = current_topic { - completion_metadata.insert("topic_id".to_string(), topic_id.clone()); - } - if let Err(error) = self - .bus - .publish_outbound( - OutboundMessage::execution_completed( - channel, - chat_id, - Some(session_id), - completion_metadata, + // 发送执行完成信号,通知前端可以停止 loading 状态。 + // + // 退出兜底(safety net):如果当前 topic 仍有 running 状态的子代理, + // 不发送 ExecutionCompleted。这防止 LLM 未调用 wait_for_subagents 就退出时, + // 前端过早停止 loading 导致子代理结果"丢失"的观感。 + // 恢复路径:下一条用户消息触发新的 process_one → 加载 history → + // LLM 看到 "running" 占位 → 调用 wait_for_subagents → 消费 sub_done_q 结果。 + let has_pending_subagents = if let Some(ref topic_id) = current_topic { + let pending = self + .session_manager + .store() + .list_pending_subagents(topic_id, Some("running")) + .unwrap_or_default(); + if !pending.is_empty() { + tracing::info!( + topic_id = %topic_id, + pending_count = pending.len(), + "Skipping ExecutionCompleted: pending subagents still running" + ); + true + } else { + false + } + } else { + false + }; + + if !has_pending_subagents { + let mut completion_metadata = inbound.forwarded_metadata.clone(); + if let Some(ref topic_id) = current_topic { + completion_metadata.insert("topic_id".to_string(), topic_id.clone()); + } + if let Err(error) = self + .bus + .publish_outbound( + OutboundMessage::execution_completed( + channel, + chat_id, + Some(session_id), + completion_metadata, + ) + .with_trace_id(&inbound.trace_id), ) - .with_trace_id(&inbound.trace_id), - ) - .await - { - match error { - crate::bus::BusError::Dropped => { - tracing::warn!(error = %error, "Outbound dropped (bus full)"); - } - crate::bus::BusError::Closed => { - tracing::error!(error = %error, "Failed to publish execution_completed"); + .await + { + match error { + crate::bus::BusError::Dropped => { + tracing::warn!(error = %error, "Outbound dropped (bus full)"); + } + crate::bus::BusError::Closed => { + tracing::error!(error = %error, "Failed to publish execution_completed"); + } } } } diff --git a/src/gateway/runtime.rs b/src/gateway/runtime.rs index df3315f..0b1ee18 100644 --- a/src/gateway/runtime.rs +++ b/src/gateway/runtime.rs @@ -26,7 +26,7 @@ use crate::tools::task::runtime::SubagentRuntime; use crate::tools::todo_write::TodoItem; use crate::tools::{ DefaultSubAgentRuntime, InMemoryTaskRepository, NoopSessionMessageSender, SessionMessageSender, - SubAgentRuntimeConfig, SubagentCatalog, TaskTool, ToolRegistry, + SubAgentRuntime, SubAgentRuntimeConfig, SubagentCatalog, TaskTool, ToolRegistry, }; use super::agent_factory::AgentFactory; @@ -67,6 +67,7 @@ pub(crate) fn build_session_manager( Option>, Arc, Arc, + Option>, ), AgentError, > { @@ -119,6 +120,7 @@ pub(crate) fn build_session_manager_with_sender( Option>, Arc, Arc, + Option>, ), AgentError, > { @@ -209,10 +211,11 @@ pub(crate) fn build_session_manager_with_sender( } // Create SubAgentRuntime (if task tool is enabled) - let (factory, task_repository, subagent_runtime): ( + let (factory, task_repository, subagent_runtime, subagent_executor): ( _, Arc, Arc, + Option>, ) = if task_config.enabled { let task_repository = Arc::new(InMemoryTaskRepository::new()); // Build subagent tools with MCP tools (task tool registered separately below) @@ -237,6 +240,7 @@ pub(crate) fn build_session_manager_with_sender( default_max_execution_secs: task_config.max_execution_secs, ttl_hours: task_config.ttl_hours, max_nesting_depth: task_config.max_nesting_depth, + max_concurrent: task_config.max_concurrent, }; let default_subagent_runtime = Arc::new(DefaultSubAgentRuntime::new( @@ -260,10 +264,14 @@ pub(crate) fn build_session_manager_with_sender( )); } + let subagent_executor: Option> = + Some(default_subagent_runtime.clone()); + ( factory.with_subagent_runtime(default_subagent_runtime), task_repository, subagent_runtime, + subagent_executor, ) } else { // task_config 未启用时仍创建 subagent_runtime(供 API 使用) @@ -272,6 +280,7 @@ pub(crate) fn build_session_manager_with_sender( factory, Arc::new(InMemoryTaskRepository::new()), subagent_runtime, + None, ) }; @@ -367,5 +376,6 @@ pub(crate) fn build_session_manager_with_sender( mcp_manager, subagent_runtime, model_selections, + subagent_executor, )) } diff --git a/src/gateway/session.rs b/src/gateway/session.rs index 67fcd0e..4e3d859 100644 --- a/src/gateway/session.rs +++ b/src/gateway/session.rs @@ -14,10 +14,12 @@ use crate::storage::{ use crate::tools::ToolRegistry; use crate::tools::task::repository::TaskRepository; use crate::tools::task::runtime::SubagentRuntime; +use crate::tools::task::SubagentResult; +use crate::tools::WaitCoordinator; use async_trait::async_trait; use std::collections::HashMap; use std::sync::Arc; -use tokio::sync::{Mutex, mpsc}; +use tokio::sync::{Mutex, Notify, mpsc}; use uuid::Uuid; use super::agent_factory::{AgentBuildRequest, AgentFactory}; @@ -413,12 +415,17 @@ impl Session { /// 确保指定 topic 的历史已加载到内存。 /// 按 topic_id 键化查找,已存在则直接返回,否则从 DB 加载。 + /// 加载后对账 pending_subagents DB,替换过时的 "running" 占位。 pub fn ensure_chat_loaded( &mut self, chat_id: &str, topic_id: Option<&str>, ) -> Result<(), AgentError> { - self.history.ensure_chat_loaded(chat_id, topic_id) + self.history.ensure_chat_loaded(chat_id, topic_id)?; + if let Some(tid) = topic_id { + self.reconcile_running_placeholders(tid); + } + Ok(()) } pub fn ensure_agent_prompt_before_user_message( @@ -587,19 +594,147 @@ impl Session { self.history.topic_serial_lock(topic_id) } + /// 确保该 topic 的 sub_done 队列已创建(与 topic_serial_lock 同步初始化)。 + pub(crate) fn ensure_sub_done_channel(&mut self, topic_id: &str) { + self.history.ensure_sub_done_channel(topic_id); + } + + /// 获取该 topic 的 sub_done 队列 sender(用于后台子代理发送结果)。 + #[allow(dead_code)] + pub(crate) fn sub_done_sender( + &mut self, + topic_id: &str, + ) -> Option> { + self.history.sub_done_sender(topic_id) + } + + /// 取出该 topic 的 sub_done 队列 receiver(wait 工具进入 select! 前调用)。 + pub(crate) fn take_sub_done_receiver( + &mut self, + topic_id: &str, + ) -> Option> { + self.history.take_sub_done_receiver(topic_id) + } + + /// 归还该 topic 的 sub_done 队列 receiver(wait 工具 select! 结束后调用)。 + pub(crate) fn restore_sub_done_receiver( + &mut self, + topic_id: &str, + rx: mpsc::Receiver, + ) { + self.history.restore_sub_done_receiver(topic_id, rx); + } + + /// 获取或创建该 topic 的 wait 唤醒信号。 + pub(crate) fn wait_wakeup(&mut self, topic_id: &str) -> Arc { + self.history.wait_wakeup(topic_id) + } + + /// 设置该 topic 的等待状态。 + pub(crate) fn set_waiting(&mut self, topic_id: &str, waiting: bool) { + self.history.set_waiting(topic_id, waiting); + } + + /// 检查该 topic 是否处于等待状态。 + pub(crate) fn is_waiting(&self, topic_id: &str) -> bool { + self.history.is_waiting(topic_id) + } + /// 按 topic_id 从 DB 重新加载历史到内存 pub(crate) fn reload_topic_history( &mut self, chat_id: &str, topic_id: &str, ) -> Result<(), AgentError> { - self.history.reload_topic_history(chat_id, topic_id) + self.history.reload_topic_history(chat_id, topic_id)?; + self.reconcile_running_placeholders(topic_id); + Ok(()) + } + + /// 对账 pending_subagents DB,将内存 history 中过时的 "running" 占位 + /// 替换为 DB 中的实际状态。 + /// + /// 场景:服务器崩溃重启后,history 中仍保留 "running" 占位, + /// 但 DB 中该子代理状态可能已被启动扫描标记为 "interrupted"。 + /// 若不对账,LLM 会看到 "running" → 调用 wait_for_subagents → + /// query_pending_task_ids 返回空(DB 已非 running)→ 返回 "no pending" → + /// LLM 困惑:history 说 running 但 wait 说无 pending。 + /// + /// 支持两种 content 格式: + /// - JSON: `{"status":"running","task_id":"xxx",...}`(当前 task 工具返回格式) + /// - 纯文本: `running, task_id=xxx. ...`(旧格式,向后兼容) + /// + /// 仅修改内存缓存,不持久化到 DB(每次从 DB 加载时重新对账,幂等)。 + fn reconcile_running_placeholders(&mut self, topic_id: &str) { + let pending = match self.store.list_pending_subagents(topic_id, None) { + Ok(records) => records, + Err(e) => { + tracing::warn!( + error = %e, + topic_id = %topic_id, + "Failed to query pending_subagents for reconciliation" + ); + return; + } + }; + + if pending.is_empty() { + return; + } + + let status_map: std::collections::HashMap<&str, &str> = pending + .iter() + .map(|r| (r.task_id.as_str(), r.status.as_str())) + .collect(); + + let history = self.history.get_or_create_history(topic_id); + let mut reconciled = 0; + + for msg in history.iter_mut() { + if msg.role != "tool" { + continue; + } + + // 尝试提取 task_id(同时支持 JSON 和纯文本格式) + let (task_id, is_json) = match extract_task_id_from_content(&msg.content) { + Some(id) => id, + None => continue, + }; + + let actual_status = match status_map.get(task_id.as_str()) { + Some(s) => *s, + None => continue, // 记录不存在(已清理),保留原占位 + }; + + if actual_status == "running" { + continue; // 仍在运行,保留占位 + } + + // 替换为实际状态 + msg.content = format_reconciled_content(&task_id, actual_status, is_json); + reconciled += 1; + } + + if reconciled > 0 { + tracing::info!( + topic_id = %topic_id, + reconciled_count = reconciled, + "Reconciled stale 'running' placeholders with DB status" + ); + } } pub(crate) fn store(&self) -> Arc { self.history.conversations() } + /// 获取底层 SessionStore(用于 pending_subagents 查询等)。 + /// 与 `store()` 不同:后者返回 ConversationRepository trait object, + /// 此方法返回具体的 SessionStore 类型,暴露 pending_subagents 等 CRUD。 + pub(crate) fn session_store(&self) -> Arc { + self.store.clone() + } + pub fn record_skill_offer(&self, chat_id: &str) -> Result<(), AgentError> { if self.skills.is_empty() { return Ok(()); @@ -621,6 +756,7 @@ impl Session { message_id: Option<&str>, explicit_topic_id: Option<&str>, trace_id: &str, + wait_coordinator: Option>, ) -> Result { self.create_agent_with_provider_config( chat_id, @@ -630,6 +766,7 @@ impl Session { self.provider_config.clone(), explicit_topic_id, trace_id, + wait_coordinator, ) } @@ -642,6 +779,7 @@ impl Session { provider_config: LLMProviderConfig, explicit_topic_id: Option<&str>, trace_id: &str, + wait_coordinator: Option>, ) -> Result { // 优先使用显式传入的 topic_id;回退到当前 chat 的活跃 topic let topic_id = explicit_topic_id @@ -658,6 +796,12 @@ impl Session { None => self.pending_cancel_tokens.remove(session_chat_id), }; + // 获取该 topic 的 sub_done 队列 sender(用于异步子代理完成回调) + // 仅主 agent 注入;无 topic 时为 None(走同步路径) + let sub_done_sender = topic_id + .as_deref() + .and_then(|tid| self.history.sub_done_sender(tid)); + self.agent_factory.create(AgentBuildRequest { channel_name: &self.channel_name, session_chat_id, @@ -668,10 +812,97 @@ impl Session { provider_config, cancel_token, trace_id: Some(trace_id.to_string()), + sub_done_sender, + wait_coordinator, }) } } +/// 从 tool result content 中提取 task_id,同时支持 JSON 和纯文本格式。 +/// +/// JSON 格式: `{"status":"running","task_id":"xxx",...}` +/// 纯文本格式: `running, task_id=xxx. ...` +/// +/// 返回 (task_id, is_json)。仅当 status=="running" 时才提取(只对账 running 占位)。 +fn extract_task_id_from_content(content: &str) -> Option<(String, bool)> { + // 尝试 JSON 格式 + let json_start = content.find('{')?; + if let Ok(parsed) = serde_json::from_str::(&content[json_start..]) { + let status = parsed.get("status").and_then(|v| v.as_str())?; + if status != "running" { + return None; + } + let task_id = parsed.get("task_id").and_then(|v| v.as_str())?; + if task_id.is_empty() { + return None; + } + return Some((task_id.to_string(), true)); + } + + // 回退到纯文本格式(向后兼容) + let prefix = "running, task_id="; + let rest = content.strip_prefix(prefix)?; + let end = rest.find('.').unwrap_or(rest.len()); + let task_id = &rest[..end]; + if task_id.is_empty() { + None + } else { + Some((task_id.to_string(), false)) + } +} + +/// 根据子代理的实际状态生成替换内容。 +/// +/// JSON 格式:更新 JSON 中的 status 字段(保持前端 parseTaskResult 兼容)。 +/// 纯文本格式:替换为描述性文本(向后兼容)。 +fn format_reconciled_content(task_id: &str, status: &str, is_json: bool) -> String { + if is_json { + // 更新 JSON 中的 status 字段,保持前端 parseTaskResult 能正确解析 + let placeholder = match status { + "interrupted" => format!( + "Subagent {} was interrupted (server restart). Result is unavailable.", + task_id + ), + "completed" => format!( + "Subagent {} has completed. Call wait_for_subagents to retrieve the result.", + task_id + ), + "failed" => format!( + "Subagent {} has failed. Call wait_for_subagents to retrieve error details.", + task_id + ), + "timeout" => format!("Subagent {} timed out.", task_id), + "cancelled" => format!("Subagent {} was cancelled.", task_id), + other => format!("Subagent {} status: {}.", task_id, other), + }; + serde_json::json!({ + "status": status, + "summary": placeholder, + "output": placeholder, + "task_id": task_id, + }) + .to_string() + } else { + match status { + "interrupted" => format!( + "Subagent {} was interrupted (server restart). Result is unavailable.", + task_id + ), + "completed" => format!( + "Subagent {} has completed. Call wait_for_subagents to retrieve the result.", + task_id + ), + "failed" => format!( + "Subagent {} has failed. Call wait_for_subagents to retrieve error details.", + task_id + ), + "timeout" => format!("Subagent {} timed out.", task_id), + "cancelled" => format!("Subagent {} was cancelled.", task_id), + other => format!("Subagent {} status: {}.", task_id, other), + } + } +} + /// SessionManager 管理所有 Session,按 channel_name 路由 #[derive(Clone)] pub struct SessionManager { @@ -763,7 +994,7 @@ impl SessionManager { model_resolver, crate::config::CompactionConfig::default(), ) - .map(|(session_manager, _, _, _, _)| session_manager) + .map(|(session_manager, _, _, _, _, _)| session_manager) } pub fn tools(&self) -> Arc { diff --git a/src/gateway/session_history.rs b/src/gateway/session_history.rs index b6a4c62..aa9600c 100644 --- a/src/gateway/session_history.rs +++ b/src/gateway/session_history.rs @@ -1,11 +1,14 @@ use std::collections::{HashMap, HashSet}; use std::sync::Arc; +use tokio::sync::{Notify, mpsc}; + use crate::agent::AgentError; use crate::bus::ChatMessage; use crate::storage::{ ConversationRepository, SessionRecord, SkillEventRepository, persistent_session_id, }; +use crate::tools::task::SubagentResult; /// 内存中缓存的 topic 历史上限。 /// 超过此值时,驱逐非活跃 topic(不在 chat_topic_ids 当前引用中的 topic)。 @@ -35,6 +38,23 @@ pub(crate) struct SessionHistory { /// 防止并发 loop 操作同一历史的不同快照产生交错序列。 /// 不同 topic 之间互不阻塞,支持多话题并发执行。 topic_serial_locks: HashMap>>, + /// per-topic 子代理完成队列的 sender 端。 + /// 后台子代理完成后 clone 此 sender 发送 SubagentResult, + /// 由 wait_for_subagents 工具的 receiver 端消费。 + /// 生命周期与 topic_serial_locks 完全一致。 + sub_done_senders: HashMap>, + /// per-topic 子代理完成队列的 receiver 端。 + /// 使用 Option 包装以便 wait_for_subagents 工具在 select! 期间 + /// 将其取出(不持有 Session 锁),结束后归还。 + sub_done_receivers: HashMap>>, + /// per-topic wait 唤醒信号。 + /// wait 释放锁进入 select! 后,process_one 在注入用户消息到 history 后 + /// 调用 notify_one() 唤醒 wait。 + wait_wakeups: HashMap>, + /// per-topic 等待状态标志。 + /// true 表示该 topic 的 agent 正在 wait_for_subagents 中等待(已释放 serial_lock)。 + /// process_one 持锁后检查此标志:true 则注入用户消息 + wakeup,false 则正常处理。 + waiting_flags: HashMap, conversations: Arc, skill_events: Arc, } @@ -67,6 +87,20 @@ impl SessionHistory { if active.contains(tid.as_str()) || self.compression_in_flight.contains(*tid) { return false; } + // 不变量 2:lock 实例单射性 — 驱逐会移除 topic_serial_locks/wakeup/sender 等 entry, + // 下次访问会创建新实例,破坏原 Arc 持有者(wait_coordinator)与新建者(process_one) + // 的串行化。因此必须充分覆盖所有"活跃"语义: + // a) waiting_flag=true → wait_coordinator 正在 select! 等待 + // b) sub_done_receivers[tid]=None → receiver 被 take 走(wait_coordinator 持有) + // c) topic_serial_lock 被持有 → 有 agent 任务正在处理 + // 任一为真都不能驱逐。 + if self.waiting_flags.get(*tid).copied().unwrap_or(false) { + return false; + } + if matches!(self.sub_done_receivers.get(*tid), Some(None)) { + // receiver 被 take 走 = wait_coordinator 正在 select! 中 + return false; + } // 检查是否有活跃 agent 任务(serial lock 被持有) // try_lock 成功 = 锁空闲 = 无活跃任务 = 可驱逐 // try_lock 失败 = 锁被持有 = 有活跃任务 = 不驱逐 @@ -81,6 +115,11 @@ impl SessionHistory { if let Some(tid) = to_evict.cloned() { let msg_count = self.topic_histories.get(&tid).map(|h| h.len()).unwrap_or(0); self.topic_histories.remove(&tid); + self.topic_serial_locks.remove(&tid); + self.sub_done_senders.remove(&tid); + self.sub_done_receivers.remove(&tid); + self.wait_wakeups.remove(&tid); + self.waiting_flags.remove(&tid); tracing::info!( topic_id = %tid, evicted_messages = msg_count, @@ -101,6 +140,10 @@ impl SessionHistory { chat_topic_ids: HashMap::new(), compression_in_flight: HashSet::new(), topic_serial_locks: HashMap::new(), + sub_done_senders: HashMap::new(), + sub_done_receivers: HashMap::new(), + wait_wakeups: HashMap::new(), + waiting_flags: HashMap::new(), conversations, skill_events, } @@ -109,6 +152,9 @@ impl SessionHistory { /// 获取或创建该 topic 的串行化锁。 /// 同一 topic 的所有消息处理共享同一个锁,保证串行执行; /// 不同 topic 之间互不阻塞,支持多话题并发执行。 + /// + /// 同时同步创建该 topic 的 sub_done 队列、wait 唤醒信号和等待标志, + /// 生命周期与 serial_lock 完全一致。 pub(crate) fn topic_serial_lock(&mut self, topic_id: &str) -> Arc> { self.topic_serial_locks .entry(topic_id.to_string()) @@ -116,6 +162,68 @@ impl SessionHistory { .clone() } + /// 确保该 topic 的 sub_done 队列已创建。 + /// 应在 topic 初始化时(与 topic_serial_lock 同步)调用。 + /// 队列容量为 32,足够缓存多个子代理同时完成的结果。 + pub(crate) fn ensure_sub_done_channel(&mut self, topic_id: &str) { + if !self.sub_done_senders.contains_key(topic_id) { + let (tx, rx) = mpsc::channel::(32); + self.sub_done_senders.insert(topic_id.to_string(), tx); + self.sub_done_receivers + .insert(topic_id.to_string(), Some(rx)); + } + } + + /// 获取该 topic 的 sub_done 队列 sender(用于后台子代理发送结果)。 + /// 调用前应已通过 `ensure_sub_done_channel` 创建队列。 + pub(crate) fn sub_done_sender(&mut self, topic_id: &str) -> Option> { + self.ensure_sub_done_channel(topic_id); + self.sub_done_senders.get(topic_id).cloned() + } + + /// 取出该 topic 的 sub_done 队列 receiver。 + /// 由 wait_for_subagents 工具在进入 select! 前调用(需释放 Session 锁)。 + /// 调用后 receiver 不在 map 中,需通过 `restore_sub_done_receiver` 归还。 + pub(crate) fn take_sub_done_receiver( + &mut self, + topic_id: &str, + ) -> Option> { + self.sub_done_receivers + .get_mut(topic_id) + .and_then(|opt| opt.take()) + } + + /// 归还该 topic 的 sub_done 队列 receiver。 + /// 由 wait_for_subagents 工具在 select! 结束后调用。 + pub(crate) fn restore_sub_done_receiver( + &mut self, + topic_id: &str, + rx: mpsc::Receiver, + ) { + self.sub_done_receivers + .insert(topic_id.to_string(), Some(rx)); + } + + /// 获取或创建该 topic 的 wait 唤醒信号。 + pub(crate) fn wait_wakeup(&mut self, topic_id: &str) -> Arc { + self.wait_wakeups + .entry(topic_id.to_string()) + .or_insert_with(|| Arc::new(Notify::new())) + .clone() + } + + /// 设置该 topic 的等待状态。 + /// true = agent 正在 wait_for_subagents 中等待(已释放 serial_lock)。 + pub(crate) fn set_waiting(&mut self, topic_id: &str, waiting: bool) { + self.waiting_flags.insert(topic_id.to_string(), waiting); + } + + /// 检查该 topic 是否处于等待状态。 + /// process_one 持锁后调用:true 则走注入路径,false 则正常处理。 + pub(crate) fn is_waiting(&self, topic_id: &str) -> bool { + self.waiting_flags.get(topic_id).copied().unwrap_or(false) + } + pub(crate) fn persistent_session_id(&self, chat_id: &str) -> String { persistent_session_id(&self.channel_name, chat_id) } @@ -212,6 +320,11 @@ impl SessionHistory { // (仅在无活跃任务时安全移除;有活跃任务时 lock 被 Arc clone 持有, // 移除 HashMap entry 不影响正在使用 lock 的任务) self.topic_serial_locks.remove(topic_id); + // 同步清理 sub_done 队列、wait 唤醒信号和等待标志 + self.sub_done_senders.remove(topic_id); + self.sub_done_receivers.remove(topic_id); + self.wait_wakeups.remove(topic_id); + self.waiting_flags.remove(topic_id); } /// 清空指定 chat/topic 的内存历史和 DB 消息。 diff --git a/src/gateway/tool_registry_factory.rs b/src/gateway/tool_registry_factory.rs index eeec384..2d660ae 100644 --- a/src/gateway/tool_registry_factory.rs +++ b/src/gateway/tool_registry_factory.rs @@ -14,7 +14,7 @@ use crate::tools::{ BashTool, CalculatorTool, FileEditTool, FileReadTool, FileWriteTool, HttpRequestTool, MemoryManageTool, MemorySearchTool, SchedulerManageTool, SessionMessageSender, SessionSendTool, ShellSessionManager, SkillActivateTool, SkillManageTool, SubAgentRuntime, TaskTool, TimeTool, - TodoReadTool, TodoWriteTool, ToolRegistry, WebFetchTool, + TodoReadTool, TodoWriteTool, ToolRegistry, WaitForSubagentsTool, WebFetchTool, }; pub(crate) struct ToolRegistryFactory { @@ -160,6 +160,11 @@ impl ToolRegistryFactory { if self.is_enabled("task") && self.task_config.enabled { if let Some(runtime) = &self.subagent_runtime { registry.register(TaskTool::new(runtime.clone(), None)); + // 注册 wait_for_subagents 工具(仅主 agent,用于等待异步子代理完成) + // 默认超时从配置读取,LLM 可通过 timeout_secs 参数覆盖 + registry.register(WaitForSubagentsTool::new( + self.task_config.wait_default_timeout_secs, + )); } } diff --git a/src/gateway/wait_coordinator.rs b/src/gateway/wait_coordinator.rs new file mode 100644 index 0000000..c1cfd4c --- /dev/null +++ b/src/gateway/wait_coordinator.rs @@ -0,0 +1,291 @@ +use std::sync::Arc; +use std::time::Duration; + +use async_trait::async_trait; +use tokio::sync::{Mutex, OwnedMutexGuard, watch}; +use tokio::time::sleep; + +use crate::gateway::session::Session; +use crate::storage::SessionStore; +use crate::tools::task::SubagentResult; +use crate::tools::{WaitCoordinator, WaitEvent}; + +/// 基于 Session 的 wait 协调器实现。 +/// +/// 持有 serial_lock 的 guard slot(可通过 take/drop 释放锁,通过 put 回填新 guard), +/// 以及 Session 引用(用于管理 sub_done_receiver、wait_wakeup、waiting_flag)。 +/// +/// wait() 流程: +/// 1. 设置 waiting=true(让 process_one 走注入路径) +/// 2. 取出 sub_done_receiver(从 Session 中 take,select! 期间不持有 Session 锁) +/// 3. 获取 wait_wakeup(Arc,clone 后不持有 Session 锁) +/// 4. 释放 serial_lock(从 guard_slot take 并 drop guard) +/// 5. select! { sub_done_q.recv(), wakeup.notified(), timeout } +/// 6. 重新获取 serial_lock(serial_lock.lock_owned().await) +/// 7. 回填 guard 到 guard_slot +/// 8. 设置 waiting=false(先获取锁后清除,避免 TOCTOU) +/// 9. 归还 sub_done_receiver +pub struct SessionWaitCoordinator { + /// Session 引用(Arc>),用于访问 SessionHistory 的队列和状态 + session: Arc>, + /// serial_lock guard 的存储槽。 + /// 执行路径(execution.rs)获取锁后将 guard 存入此槽; + /// wait() 取出并 drop 以释放锁,重获取后回填新 guard。 + guard_slot: Arc>>>, + /// serial_lock 本体(Arc>),用于重获取锁 + serial_lock: Arc>, + /// SessionStore 引用,用于查询 pending_subagents + store: Arc, + /// 当前 topic_id + topic_id: String, +} + +impl SessionWaitCoordinator { + pub fn new( + session: Arc>, + guard_slot: Arc>>>, + serial_lock: Arc>, + store: Arc, + topic_id: String, + ) -> Self { + Self { + session, + guard_slot, + serial_lock, + store, + topic_id, + } + } +} + +#[async_trait] +impl WaitCoordinator for SessionWaitCoordinator { + fn query_pending_task_ids(&self) -> Vec { + self.store + .list_pending_subagents(&self.topic_id, Some("running")) + .map(|records| records.into_iter().map(|r| r.task_id).collect()) + .unwrap_or_default() + } + + async fn try_drain_queued_results(&self) -> Vec { + // 取出 receiver → try_recv 排空 → 归还 receiver + // 安全性:此方法在执行路径中被调用(serial_lock 已持有), + // 无其他代码并发访问 receiver。 + let rx = { + let mut session = self.session.lock().await; + session.take_sub_done_receiver(&self.topic_id) + }; + + let mut results = Vec::new(); + if let Some(mut rx) = rx { + while let Ok(result) = rx.try_recv() { + results.push(result); + } + // 归还 receiver(即使已排空,仍需放回供后续 wait() 使用) + let mut session = self.session.lock().await; + session.restore_sub_done_receiver(&self.topic_id, rx); + } + + if !results.is_empty() { + tracing::debug!( + topic_id = %self.topic_id, + drained_count = results.len(), + "Drained buffered subagent results from sub_done_q" + ); + } + + results + } + + async fn wait( + &self, + timeout: Duration, + cancel_rx: Option>, + ) -> WaitEvent { + // 1. 设置 waiting=true + { + let mut session = self.session.lock().await; + session.set_waiting(&self.topic_id, true); + } + + // 2. 取出 sub_done_receiver(select! 期间不持有 Session 锁) + let receiver = { + let mut session = self.session.lock().await; + session.take_sub_done_receiver(&self.topic_id) + }; + + // 3. 获取 wait_wakeup(Arc,clone 后不持有 Session 锁) + let wakeup = { + let mut session = self.session.lock().await; + session.wait_wakeup(&self.topic_id) + }; + + // 3.5. 记录等待前的用户消息数量(用于 wakeup 后提取新注入的消息) + // 直接从 SQLite 读取,不持有任何锁 + let user_msg_count_before = self + .store + .load_messages_for_topic(&self.topic_id, None) + .map(|msgs| msgs.iter().filter(|m| m.role == "user").count()) + .unwrap_or(0); + + // 4. 释放 serial_lock(取出 guard 并 drop) + { + let mut slot = self.guard_slot.lock().await; + let _ = slot.take(); // drop guard → 释放 serial_lock + } + + tracing::debug!( + topic_id = %self.topic_id, + timeout_secs = timeout.as_secs(), + user_msg_count_before, + has_cancel_rx = cancel_rx.is_some(), + "SessionWaitCoordinator: lock released, entering select!" + ); + + // 5. select! 等待(不持有任何锁) + // + // cancel 分支放在最后(biased 排序中最后被 poll), + // 确保子代理结果和用户消息优先于取消信号被处理。 + // 场景:/stop 后用户立即发消息 → process_one 注入消息 + wakeup, + // select! 优先消费 wakeup(UserMessage),而非 cancel(Cancelled), + // 使已注入的用户消息能被 Agent 处理而非丢失。 + // + // 但 cancel_rx.changed() 不会无限阻塞——若无子代理结果、无用户消息, + // cancel 仍是唯一就绪分支,等待被优雅终止。 + let event = if let Some(mut rx) = receiver { + let mut cancel_rx = cancel_rx; + tokio::select! { + biased; + result = rx.recv() => { + match result { + Some(subagent_result) => { + let event = WaitEvent::SubagentResult(subagent_result); + let mut session = self.session.lock().await; + session.restore_sub_done_receiver(&self.topic_id, rx); + event + } + None => { + // sender 全部 drop(所有 sub_done_sender 被释放) + WaitEvent::Timeout + } + } + } + _ = wakeup.notified() => { + // 用户消息到达(process_one 已注入 history 并 wakeup) + let mut session = self.session.lock().await; + session.restore_sub_done_receiver(&self.topic_id, rx); + // 提取等待期间新注入的用户消息内容 + let new_messages = self.fetch_new_user_messages(user_msg_count_before); + tracing::info!( + topic_id = %self.topic_id, + new_msg_count = new_messages.len(), + "SessionWaitCoordinator: woke up by user message" + ); + WaitEvent::UserMessage(new_messages) + } + _ = async { + if let Some(ref mut crx) = cancel_rx { + let _ = crx.changed().await; + } else { + std::future::pending::<()>().await; + } + } => { + // 取消信号到达(/stop)→ 归还 receiver,返回 Cancelled + tracing::info!( + topic_id = %self.topic_id, + "SessionWaitCoordinator: cancelled by /stop during wait" + ); + let mut session = self.session.lock().await; + session.restore_sub_done_receiver(&self.topic_id, rx); + WaitEvent::Cancelled + } + _ = sleep(timeout) => { + let mut session = self.session.lock().await; + session.restore_sub_done_receiver(&self.topic_id, rx); + WaitEvent::Timeout + } + } + } else { + // 无 receiver(topic 无 sub_done 队列),直接等待 timeout 或 wakeup + let mut cancel_rx = cancel_rx; + tokio::select! { + biased; + _ = wakeup.notified() => { + let new_messages = self.fetch_new_user_messages(user_msg_count_before); + tracing::info!( + topic_id = %self.topic_id, + new_msg_count = new_messages.len(), + "SessionWaitCoordinator: woke up by user message (no receiver)" + ); + WaitEvent::UserMessage(new_messages) + } + _ = async { + if let Some(ref mut crx) = cancel_rx { + let _ = crx.changed().await; + } else { + std::future::pending::<()>().await; + } + } => { + tracing::info!( + topic_id = %self.topic_id, + "SessionWaitCoordinator: cancelled by /stop during wait (no receiver)" + ); + WaitEvent::Cancelled + } + _ = sleep(timeout) => WaitEvent::Timeout, + } + }; + + // 6. 重新获取 serial_lock + // lock_owned 消费 Arc,需 clone 保留 self.serial_lock 供后续可能的重入 + // + // 取消场景下此处可能阻塞——如果 process_one 正持有锁注入用户消息, + // 需等其释放后才能重获取。这是正确行为:确保 is_waiting 清除与 + // process_one 的注入互斥,避免 TOCTOU。 + let new_guard = self.serial_lock.clone().lock_owned().await; + + // 7. 回填 guard 到 guard_slot + { + let mut slot = self.guard_slot.lock().await; + *slot = Some(new_guard); + } + + // 8. 设置 waiting=false(先获取锁后清除,避免 TOCTOU) + // 注意:serial_lock 已在步骤 6 获取,此时 process_one 无法获取锁, + // 所以清除 waiting 是安全的。 + { + let mut session = self.session.lock().await; + session.set_waiting(&self.topic_id, false); + } + + tracing::debug!( + topic_id = %self.topic_id, + "SessionWaitCoordinator: lock reacquired, waiting cleared" + ); + + event + } +} + +impl SessionWaitCoordinator { + /// 提取等待期间新注入的用户消息内容。 + /// 通过对比等待前的用户消息数量,从 SQLite 中取出新增的用户消息。 + fn fetch_new_user_messages(&self, count_before: usize) -> Vec { + match self.store.load_messages_for_topic(&self.topic_id, None) { + Ok(msgs) => msgs + .iter() + .filter(|m| m.role == "user") + .skip(count_before) + .map(|m| m.content.clone()) + .collect(), + Err(e) => { + tracing::warn!( + error = %e, + topic_id = %self.topic_id, + "Failed to load messages for fetching new user messages" + ); + Vec::new() + } + } + } +} diff --git a/src/gateway/ws.rs b/src/gateway/ws.rs index 44782d2..49f6996 100644 --- a/src/gateway/ws.rs +++ b/src/gateway/ws.rs @@ -535,6 +535,7 @@ async fn handle_inbound( router.register(Box::new(StopExecutionCommandHandler::new( state.cancel_manager.clone(), state.session_manager.clone(), + state.subagent_executor.clone(), ))); // 构建命令上下文 diff --git a/src/protocol/mod.rs b/src/protocol/mod.rs index 86966da..9610640 100644 --- a/src/protocol/mod.rs +++ b/src/protocol/mod.rs @@ -330,6 +330,12 @@ pub enum WsOutbound { timestamp: Option, #[serde(default, skip_serializing_if = "Option::is_none")] subagent_task_id: Option, + /// 子代理最终状态(completed/failed/timeout/cancelled/interrupted), + /// 供前端更新主视图中 task tool result 占位消息的显示状态。 + #[serde(default, skip_serializing_if = "Option::is_none")] + subagent_status: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + subagent_summary: Option, }, #[serde(rename = "todo_list")] TodoList { diff --git a/src/protocol/ws_adapter.rs b/src/protocol/ws_adapter.rs index 7b84551..56c2560 100644 --- a/src/protocol/ws_adapter.rs +++ b/src/protocol/ws_adapter.rs @@ -221,6 +221,8 @@ pub(crate) fn ws_outbound_from_outbound_message(message: &OutboundMessage) -> Ve topic_id: message.metadata.get("topic_id").cloned(), timestamp: Some(crate::protocol::now_timestamp()), subagent_task_id: message.metadata.get("subagent_task_id").cloned(), + subagent_status: message.metadata.get("subagent_status").cloned(), + subagent_summary: message.metadata.get("subagent_summary").cloned(), }], } } diff --git a/src/storage/migrations.rs b/src/storage/migrations.rs index 971886a..544baae 100644 --- a/src/storage/migrations.rs +++ b/src/storage/migrations.rs @@ -323,3 +323,28 @@ pub(super) fn add_column_if_missing(conn: &Connection, sql: &str) -> Result<(), Err(error) => Err(StorageError::Database(error)), } } + +/// pending_subagents 表:跟踪异步子代理执行状态,用于崩溃恢复。 +pub(super) fn ensure_pending_subagents_schema(conn: &Connection) -> Result<(), StorageError> { + conn.execute_batch( + " + CREATE TABLE IF NOT EXISTS pending_subagents ( + task_id TEXT PRIMARY KEY, + parent_session_id TEXT NOT NULL, + parent_topic_id TEXT NOT NULL, + parent_chat_id TEXT NOT NULL, + parent_channel TEXT NOT NULL, + def_name TEXT, + spawned_at INTEGER NOT NULL, + status TEXT NOT NULL DEFAULT 'running' + ); + + CREATE INDEX IF NOT EXISTS idx_pending_subagents_topic + ON pending_subagents(parent_topic_id, status); + + CREATE INDEX IF NOT EXISTS idx_pending_subagents_session + ON pending_subagents(parent_session_id); + ", + )?; + Ok(()) +} diff --git a/src/storage/mod.rs b/src/storage/mod.rs index fa6f4ee..ecba5e7 100644 --- a/src/storage/mod.rs +++ b/src/storage/mod.rs @@ -28,10 +28,10 @@ pub use ports::{ SkillEventRepository, TodoRepository, }; pub use records::{ - ALLOWED_MEMORY_NAMESPACES, GLOBAL_SCOPE_KEY, MemoryRecord, MemoryUpsert, SchedulerJobRecord, - SchedulerJobState, SchedulerJobStatus, SchedulerJobUpsert, SessionRecord, SessionTokenStats, - SkillEventRecord, TodoRecord, TopicRecord, allowed_namespace_names, get_namespace_description, - is_valid_namespace, + ALLOWED_MEMORY_NAMESPACES, GLOBAL_SCOPE_KEY, MemoryRecord, MemoryUpsert, PendingSubagentRecord, + SchedulerJobRecord, SchedulerJobState, SchedulerJobStatus, SchedulerJobUpsert, SessionRecord, + SessionTokenStats, SkillEventRecord, TodoRecord, TopicRecord, allowed_namespace_names, + get_namespace_description, is_valid_namespace, }; #[derive(Clone)] @@ -234,6 +234,7 @@ impl SessionStore { ensure_scheduler_schema(&conn)?; ensure_memory_scope_key_migration(&conn)?; ensure_todos_schema(&conn)?; + ensure_pending_subagents_schema(&conn)?; drop(conn); @@ -2052,6 +2053,140 @@ impl SessionStore { } Ok(todos) } + + // ==================== pending_subagents ==================== + + /// 插入一条 pending_subagent 记录(task 工具 spawn 时调用)。 + pub fn insert_pending_subagent( + &self, + record: &PendingSubagentRecord, + ) -> Result<(), StorageError> { + let conn = self.pool.get()?; + conn.execute( + "INSERT OR REPLACE INTO pending_subagents + (task_id, parent_session_id, parent_topic_id, parent_chat_id, parent_channel, def_name, spawned_at, status) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", + params![ + record.task_id, + record.parent_session_id, + record.parent_topic_id, + record.parent_chat_id, + record.parent_channel, + record.def_name, + record.spawned_at, + record.status, + ], + )?; + Ok(()) + } + + /// 查询指定 topic 下匹配状态的 pending_subagent 记录。 + /// `status` 为 None 时查询所有状态。 + pub fn list_pending_subagents( + &self, + topic_id: &str, + status: Option<&str>, + ) -> Result, StorageError> { + let conn = self.pool.get()?; + let sql = if status.is_some() { + "SELECT task_id, parent_session_id, parent_topic_id, parent_chat_id, parent_channel, def_name, spawned_at, status + FROM pending_subagents + WHERE parent_topic_id = ?1 AND status = ?2 + ORDER BY spawned_at ASC" + } else { + "SELECT task_id, parent_session_id, parent_topic_id, parent_chat_id, parent_channel, def_name, spawned_at, status + FROM pending_subagents + WHERE parent_topic_id = ?1 + ORDER BY spawned_at ASC" + }; + let mut stmt = conn.prepare(sql)?; + let rows = if let Some(s) = status { + stmt.query_map(params![topic_id, s], map_pending_subagent_record)? + } else { + stmt.query_map(params![topic_id], map_pending_subagent_record)? + }; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + + /// 获取指定 task_id 的 pending_subagent 记录。 + pub fn get_pending_subagent( + &self, + task_id: &str, + ) -> Result, StorageError> { + let conn = self.pool.get()?; + let mut stmt = conn.prepare( + "SELECT task_id, parent_session_id, parent_topic_id, parent_chat_id, parent_channel, def_name, spawned_at, status + FROM pending_subagents + WHERE task_id = ?1", + )?; + let mut rows = stmt.query_map(params![task_id], map_pending_subagent_record)?; + match rows.next() { + Some(row) => Ok(Some(row?)), + None => Ok(None), + } + } + + /// 更新指定 task_id 的状态(子代理完成或取消时调用)。 + pub fn update_pending_subagent_status( + &self, + task_id: &str, + new_status: &str, + ) -> Result<(), StorageError> { + let conn = self.pool.get()?; + conn.execute( + "UPDATE pending_subagents SET status = ?1 WHERE task_id = ?2", + params![new_status, task_id], + )?; + Ok(()) + } + + /// 条件更新状态:仅在当前状态为 `expected_current` 时才更新为 `new_status`。 + /// + /// 实现状态机不可逆性不变量:避免 cancel 路径覆盖 spawn 已写入的终态 + /// (completed → cancelled 是非法转换)。 + /// + /// 返回是否实际更新(affected rows > 0)。false 表示状态已被其他路径更新, + /// 调用方应跳过后续基于该假设的操作。 + pub fn try_update_pending_subagent_status( + &self, + task_id: &str, + expected_current: &str, + new_status: &str, + ) -> Result { + let conn = self.pool.get()?; + let affected = conn.execute( + "UPDATE pending_subagents SET status = ?1 WHERE task_id = ?2 AND status = ?3", + params![new_status, task_id, expected_current], + )?; + Ok(affected > 0) + } + + /// 将所有 running 状态的 pending_subagent 标记为 interrupted(启动时崩溃恢复调用)。 + pub fn mark_all_running_as_interrupted(&self) -> Result { + let conn = self.pool.get()?; + let affected = conn.execute( + "UPDATE pending_subagents SET status = 'interrupted' WHERE status = 'running'", + [], + )?; + Ok(affected) + } +} + +fn map_pending_subagent_record(row: &rusqlite::Row<'_>) -> rusqlite::Result { + Ok(PendingSubagentRecord { + task_id: row.get(0)?, + parent_session_id: row.get(1)?, + parent_topic_id: row.get(2)?, + parent_chat_id: row.get(3)?, + parent_channel: row.get(4)?, + def_name: row.get(5)?, + spawned_at: row.get(6)?, + status: row.get(7)?, + }) } pub fn persistent_session_id(channel_name: &str, chat_id: &str) -> String { diff --git a/src/storage/records.rs b/src/storage/records.rs index 5087353..5085f9e 100644 --- a/src/storage/records.rs +++ b/src/storage/records.rs @@ -111,6 +111,31 @@ pub struct TopicRecord { pub message_count: i64, } +/// pending_subagents 表的记录,跟踪异步子代理执行状态。 +/// +/// 生命周期:task 工具 spawn 时插入(status=running), +/// 子代理完成时更新为 completed/failed/timeout, +/// 进程重启时 running 状态被标记为 interrupted。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PendingSubagentRecord { + /// 子代理 task_id(主键,与 TaskSession.id 一致) + pub task_id: String, + /// 父会话 session_id + pub parent_session_id: String, + /// 父会话 topic_id(用于按 topic 查询未完成子代理) + pub parent_topic_id: String, + /// 父会话 chat_id + pub parent_chat_id: String, + /// 父会话 channel_name + pub parent_channel: String, + /// 子代理定义名称(可选,用于诊断) + pub def_name: Option, + /// 子代理启动时间戳 + pub spawned_at: i64, + /// 执行状态:running / completed / failed / interrupted / cancelled / timeout + pub status: String, +} + /// 单个 session 的 token 用量统计(聚合结果)。 #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct SessionTokenStats { diff --git a/src/tools/mod.rs b/src/tools/mod.rs index 6f97921..c874e88 100644 --- a/src/tools/mod.rs +++ b/src/tools/mod.rs @@ -18,6 +18,7 @@ pub mod time; pub mod todo_read; pub mod todo_write; pub mod traits; +pub mod wait_tool; pub mod web_fetch; pub use bash::BashTool; @@ -45,7 +46,8 @@ pub use task::{ pub use time::TimeTool; pub use todo_read::TodoReadTool; pub use todo_write::TodoWriteTool; -pub use traits::{Tool, ToolContext, ToolResult}; +pub use traits::{Tool, ToolContext, ToolResult, WaitCoordinator, WaitEvent}; +pub use wait_tool::WaitForSubagentsTool; pub use web_fetch::WebFetchTool; /// Extract a string parameter from JSON args. diff --git a/src/tools/task/error.rs b/src/tools/task/error.rs index 07d4ac6..01d537a 100644 --- a/src/tools/task/error.rs +++ b/src/tools/task/error.rs @@ -18,6 +18,9 @@ pub enum TaskError { #[error("Task execution timed out")] Timeout, + #[error("Task cancelled by user")] + Cancelled, + #[error("Repository error: {0}")] RepositoryError(#[from] StorageError), @@ -35,6 +38,7 @@ impl TaskError { pub fn as_status(&self) -> &'static str { match self { Self::Timeout => "timeout", + Self::Cancelled => "cancelled", Self::SessionNotFound(_) => "failed", Self::InvalidParentSession => "failed", Self::AgentCreationFailed(_) => "failed", diff --git a/src/tools/task/mod.rs b/src/tools/task/mod.rs index bdd2be2..69c29fc 100644 --- a/src/tools/task/mod.rs +++ b/src/tools/task/mod.rs @@ -14,6 +14,6 @@ pub use runtime::{ }; pub use tool::TaskTool; pub use types::{ - SubagentDef, SubagentSource, SubagentType, TaskDefinition, TaskHandle, TaskSession, - TaskSessionState, TaskToolArgs, TaskToolResult, + SubagentDef, SubagentResult, SubagentSource, SubagentStatus, SubagentType, TaskDefinition, + TaskHandle, TaskSession, TaskSessionState, TaskToolArgs, TaskToolResult, }; diff --git a/src/tools/task/runtime.rs b/src/tools/task/runtime.rs index 72f46f0..2676c3c 100644 --- a/src/tools/task/runtime.rs +++ b/src/tools/task/runtime.rs @@ -8,6 +8,49 @@ use std::time::Duration; use async_trait::async_trait; use serde::Deserialize; +/// RAII guard:spawn 任务退出时(正常/early return/panic)确定性清理 cancel_registry 条目。 +/// +/// 实现不变量 3(资源生命周期与作用域严格绑定): +/// spawn 块末行的手动清理是脆弱的,panic/early return 会绕过; +/// 用 Drop impl 把清理封进作用域语义,由编译器保证执行。 +struct CancelRegistryGuard { + task_id: String, + registry: Arc>>, + /// 标记是否已显式释放(例如 spawn 块成功路径末尾主动 disarm)。 + /// 默认 false,drop 时执行清理。 + disarmed: bool, +} + +impl CancelRegistryGuard { + fn new( + task_id: String, + registry: Arc>>, + ) -> Self { + Self { + task_id, + registry, + disarmed: false, + } + } + + /// 显式释放:成功路径末尾调用,避免重复清理。 + /// (实际 drop 也会幂等移除,但 disarm 让语义更清晰。) + #[allow(dead_code)] + fn disarm(&mut self) { + self.disarmed = true; + } +} + +impl Drop for CancelRegistryGuard { + fn drop(&mut self) { + if self.disarmed { + return; + } + // 幂等:条目可能已被 cancel_pending_for_topic 移除或先前已 drop + self.registry.lock().remove(&self.task_id); + } +} + use crate::agent::{ AgentLoop, AgentRuntimeConfig, EmittedMessageHandler, PersistingEmittedMessageHandler, SystemPrompt, SystemPromptContext, SystemPromptProvider, @@ -20,14 +63,18 @@ use crate::domain::CapabilityPolicy; use crate::experts::ExpertRuntime; use crate::providers::StreamDelta; use crate::skills::SkillRuntime; -use crate::storage::{ConversationRepository, SessionStore}; +use crate::storage::{ConversationRepository, PendingSubagentRecord, SessionStore}; use crate::tools::{ToolContext, ToolRegistry}; +use crate::utils::current_timestamp; use super::error::TaskError; use super::prompt::{SubagentPromptBuilder, extract_summary}; use super::repository::TaskRepository; use super::tool::TaskTool; -use super::types::{SubagentDef, SubagentSource, TaskDefinition, TaskSession, TaskToolResult}; +use super::types::{ + SubagentDef, SubagentResult, SubagentSource, SubagentStatus, TaskDefinition, TaskSession, + TaskToolResult, +}; /// 子代理运行时配置 #[derive(Debug, Clone)] @@ -40,6 +87,8 @@ pub struct SubAgentRuntimeConfig { pub ttl_hours: u64, /// 子代理最大嵌套深度(0 = 禁止嵌套,1 = 允许 1 层孙代理) pub max_nesting_depth: u32, + /// 异步子代理最大并发数(Semaphore 限流) + pub max_concurrent: usize, } impl Default for SubAgentRuntimeConfig { @@ -62,6 +111,7 @@ impl Default for SubAgentRuntimeConfig { default_max_execution_secs: 3600, // 60分钟 ttl_hours: 24, max_nesting_depth: 1, + max_concurrent: 8, } } } @@ -92,6 +142,12 @@ pub trait SubAgentRuntime: Send + Sync + 'static { /// 获取可用的子代理类型列表 fn available_subagent_names(&self) -> Vec; + + /// 取消指定 topic 下所有正在运行的异步子代理。 + /// + /// 用于 /stop 命令传播:用户取消主 agent 时,同步取消其后台子代理。 + /// 返回被触发取消的子代理数量。 + async fn cancel_pending_for_topic(&self, topic_id: &str) -> usize; } /// 静态系统提示词提供者(用于子代理) @@ -312,6 +368,15 @@ fn build_subagent_event_metadata(session: &TaskSession) -> HashMap, /// 技能运行时(实时计算技能索引,替代冻结快照) skills: Arc, + /// 异步子代理并发限流(按 config.max_concurrent 初始化) + semaphore: Arc, + /// task_id → CancellationToken 映射,用于取消传播 + /// Arc 包装以便 spawned task 完成后清理自身条目 + cancel_registry: Arc>>, } impl DefaultSubAgentRuntime { @@ -407,6 +477,8 @@ impl DefaultSubAgentRuntime { store: Arc, skills: Arc, ) -> Self { + let max_concurrent = config.max_concurrent.max(1); + let semaphore = Arc::new(tokio::sync::Semaphore::new(max_concurrent)); Self { config, task_repository, @@ -418,6 +490,8 @@ impl DefaultSubAgentRuntime { bus, store, skills, + semaphore, + cancel_registry: Arc::new(parking_lot::Mutex::new(HashMap::new())), } } @@ -553,6 +627,12 @@ impl DefaultSubAgentRuntime { parent_capability: def.map(|d| d.capability.clone()), // 从父 ToolContext 继承 trace_id,保持端到端追踪贯通子代理 trace_id: trace_id.clone(), + // 子代理不注入 sub_done_sender:嵌套层不支持异步,走同步路径 + sub_done_sender: None, + // 子代理不注入 wait_coordinator:嵌套层不支持异步 wait + wait_coordinator: None, + // 子代理不注入 cancel_rx:嵌套层不支持异步 wait,无需 cancel 检查 + cancel_rx: None, }); // 如果有 MessageBus,附加实时广播 emitter @@ -596,6 +676,18 @@ impl DefaultSubAgentRuntime { session: &TaskSession, def: &SubagentDef, prompt: String, + ) -> Result { + let max_secs = self.effective_max_execution_secs(def); + Self::execute_task_static(agent, session, def, prompt, max_secs).await + } + + /// 静态执行任务(供 tokio::spawn 调用,不依赖 &self) + async fn execute_task_static( + agent: AgentLoop, + session: &TaskSession, + _def: &SubagentDef, + prompt: String, + max_secs: u64, ) -> Result { // 构建初始消息 let history = vec![ChatMessage::user(prompt)]; @@ -606,7 +698,6 @@ impl DefaultSubAgentRuntime { }; // 设置超时 - let max_secs = self.effective_max_execution_secs(def); let timeout_duration = Duration::from_secs(max_secs); let result = tokio::time::timeout( @@ -825,62 +916,282 @@ impl SubAgentRuntime for DefaultSubAgentRuntime { } } - // 6-8. 构建提示词、创建子代理、执行任务 - // 统一为单个 Result 表达式:model_resolver / create_subagent / execute_task - // 的任何失败都流入下方 match 的 Err 分支,经 handle_task_failure 返回结构化结果。 - let result: Result = { - // 6. 构建子代理系统提示词 - // 实时按 def.capability 过滤技能索引(替代冻结快照,反映运行时技能增删) - let skills_index = if def.capability.has_skill_policy() { - self.skills.system_index_prompt_filtered( - def.capability.allowed_skills.as_deref(), - &def.capability.denied_skills, + // 6. 构建子代理系统提示词 + // 实时按 def.capability 过滤技能索引(替代冻结快照,反映运行时技能增删) + let skills_index = if def.capability.has_skill_policy() { + self.skills.system_index_prompt_filtered( + def.capability.allowed_skills.as_deref(), + &def.capability.denied_skills, + ) + } else { + self.skills.system_index_prompt() + }; + // 同步解析 def 中的 provider/model 覆盖,保证环境提示中的模型名与实际使用的模型一致 + let effective_provider_config = match (def.provider.is_some(), def.model.is_some()) { + (true, _) | (_, true) => self + .model_resolver + .resolve( + def.provider.as_deref(), + def.model.as_deref(), + &self.provider_config, ) - } else { - self.skills.system_index_prompt() - }; - // 同步解析 def 中的 provider/model 覆盖,保证环境提示中的模型名与实际使用的模型一致 - let effective_provider_config = match (def.provider.is_some(), def.model.is_some()) { - (true, _) | (_, true) => self - .model_resolver - .resolve( - def.provider.as_deref(), - def.model.as_deref(), - &self.provider_config, - ) - .map_err(|e| { - TaskError::AgentCreationFailed(format!( - "subagent '{}' model resolution failed: {}", - def.name, e - )) - })?, - _ => self.provider_config.clone(), - }; - let system_prompt = SubagentPromptBuilder::build( - &def, - &task.description, - &task.prompt, - &effective_provider_config, - skills_index.as_deref(), - ); + .map_err(|e| { + TaskError::AgentCreationFailed(format!( + "subagent '{}' model resolution failed: {}", + def.name, e + )) + })?, + _ => self.provider_config.clone(), + }; + let system_prompt = SubagentPromptBuilder::build( + &def, + &task.description, + &task.prompt, + &effective_provider_config, + skills_index.as_deref(), + ); - // 7. 创建子代理 - let agent = self.create_subagent( - &session, - system_prompt, - Some(&def), - parent_context.nesting_depth, - parent_context.task_id.clone(), - parent_context.trace_id.clone(), - )?; - - // 8. 执行任务 - self.execute_task(agent, &session, &def, task.prompt.clone()) - .await + // 7. 创建子代理 + let agent = match self.create_subagent( + &session, + system_prompt, + Some(&def), + parent_context.nesting_depth, + parent_context.task_id.clone(), + parent_context.trace_id.clone(), + ) { + Ok(agent) => agent, + Err(e) => { + let trace_id = parent_context.trace_id.as_deref().unwrap_or(""); + return self.handle_task_failure(session, e, trace_id).await; + } }; - // 9. 更新会话状态并保存 let trace_id = parent_context.trace_id.as_deref().unwrap_or(""); + + // 8. 判断执行模式:异步(主 agent + 有 sub_done_sender)或同步(子代理/无 sender) + let is_async_mode = parent_context.nesting_depth == 0 + && parent_context.sub_done_sender.is_some() + && session.parent_topic_id.is_some(); + + if is_async_mode { + // ===== 异步路径 ===== + let topic_id = session.parent_topic_id.clone().unwrap_or_default(); + let task_id = session.id.clone(); + + // 8a. INSERT pending_subagents 记录 + let pending_record = PendingSubagentRecord { + task_id: task_id.clone(), + parent_session_id: session.parent_session_id.clone(), + parent_topic_id: topic_id.clone(), + parent_chat_id: session.parent_chat_id.clone(), + parent_channel: session.parent_channel_name.clone(), + def_name: Some(def.name.clone()), + spawned_at: current_timestamp(), + status: "running".to_string(), + }; + if let Err(e) = self.store.insert_pending_subagent(&pending_record) { + tracing::warn!( + error = %e, + task_id = %task_id, + "Failed to insert pending_subagent record" + ); + } + + // 8b. tokio::spawn 后台执行子代理 + let store = self.store.clone(); + let task_repository = self.task_repository.clone(); + let bus = self.bus.clone(); + let sub_done_sender = parent_context.sub_done_sender.clone().unwrap(); + let session_clone = session.clone(); + let def_clone = def.clone(); + let prompt = task.prompt.clone(); + let trace_id_owned = trace_id.to_string(); + let max_secs = self.effective_max_execution_secs(&def); + let task_id_for_spawn = task_id.clone(); + let semaphore = self.semaphore.clone(); + let cancel_registry = self.cancel_registry.clone(); + + // 创建 CancellationToken 并注册到 registry(供 /stop 取消传播) + let cancel_token = tokio_util::sync::CancellationToken::new(); + cancel_registry + .lock() + .insert(task_id_for_spawn.clone(), cancel_token.clone()); + + // RAII guard:spawn 任务退出时(正常/early return/panic)确定性清理 registry + // 不变量 3:清理与作用域绑定,避免末行清理被 panic 绕过 + let registry_guard = + CancelRegistryGuard::new(task_id_for_spawn.clone(), cancel_registry.clone()); + + tokio::spawn(async move { + // guard 在闭包退出时 drop,确定性清理 cancel_registry 条目 + let _registry_guard = registry_guard; + + // 获取并发许可(Semaphore 限流) + let _permit = match semaphore.acquire_owned().await { + Ok(p) => p, + Err(e) => { + // 不变量 1:状态机收敛性 — early return 也必须收敛终态 + // 发送 Failed 结果让 wait 收到,更新 DB 状态,否则系统出现悬空记录 + tracing::warn!( + error = %e, + task_id = %task_id_for_spawn, + "Semaphore closed, subagent cannot start; converging state machine to Failed" + ); + let result = SubagentResult { + task_id: task_id_for_spawn.clone(), + status: SubagentStatus::Failed, + output: String::new(), + pending_task_ids: store + .list_pending_subagents(&topic_id, Some("running")) + .map(|records| { + records + .into_iter() + .map(|r| r.task_id) + .filter(|id| id != &task_id_for_spawn) + .collect::>() + }) + .unwrap_or_default(), + }; + let _ = sub_done_sender.send(result).await; + let _ = + store.update_pending_subagent_status(&task_id_for_spawn, "failed"); + // _registry_guard drop 时清理 registry 条目 + return; + } + }; + + // select! 等待执行完成或取消信号 + let exec_result = tokio::select! { + biased; + _ = cancel_token.cancelled() => { + tracing::info!( + task_id = %task_id_for_spawn, + "Subagent cancelled by user" + ); + Err(TaskError::Cancelled) + } + r = Self::execute_task_static( + agent, + &session_clone, + &def_clone, + prompt, + max_secs, + ) => r, + }; + + // 完成回调:查询未完成 → send SubagentResult → UPDATE status + let (status, output, _summary) = match &exec_result { + Ok(tool_result) => ( + SubagentStatus::Completed, + serde_json::to_string(&tool_result).unwrap_or_default(), + tool_result.summary.clone(), + ), + Err(TaskError::Timeout) => ( + SubagentStatus::Timeout, + String::new(), + "timeout".to_string(), + ), + Err(TaskError::Cancelled) => ( + SubagentStatus::Cancelled, + String::new(), + "cancelled".to_string(), + ), + Err(e) => ( + SubagentStatus::Failed, + String::new(), + e.to_string(), + ), + }; + + // 查询同 topic 下仍未完成的子代理列表 + let pending_task_ids = store + .list_pending_subagents(&topic_id, Some("running")) + .map(|records| { + records + .into_iter() + .map(|r| r.task_id) + .filter(|id| id != &task_id_for_spawn) + .collect::>() + }) + .unwrap_or_default(); + + // 发送 SubagentResult 到 sub_done_q + let result = SubagentResult { + task_id: task_id_for_spawn.clone(), + status, + output, + pending_task_ids, + }; + if let Err(e) = sub_done_sender.send(result).await { + tracing::warn!( + error = %e, + task_id = %task_id_for_spawn, + "Failed to send SubagentResult to sub_done_q (receiver dropped?)" + ); + } + + // UPDATE pending_subagents 状态 + let status_str = match status { + SubagentStatus::Completed => "completed", + SubagentStatus::Failed => "failed", + SubagentStatus::Timeout => "timeout", + SubagentStatus::Cancelled => "cancelled", + }; + if let Err(e) = store.update_pending_subagent_status(&task_id_for_spawn, status_str) { + tracing::warn!( + error = %e, + task_id = %task_id_for_spawn, + "Failed to update pending_subagent status" + ); + } + + // 更新 TaskSession 状态并发布完成事件 + let mut session_done = session_clone; + match exec_result { + Ok(tool_result) => { + session_done.mark_completed(tool_result.summary); + if let Err(e) = task_repository.save_task_session(&session_done).await { + tracing::warn!(error = %e, task_id = %task_id_for_spawn, "Failed to save completed session"); + } + publish_subagent_completion(&bus, &session_done, &trace_id_owned).await; + } + Err(e) => { + let err_str = e.to_string(); + if matches!(e, TaskError::Timeout) { + session_done.mark_timeout(); + } else if matches!(e, TaskError::Cancelled) { + session_done.mark_cancelled(); + } else { + session_done.mark_failed(err_str); + } + if let Err(e) = task_repository.save_task_session(&session_done).await { + tracing::warn!(error = %e, task_id = %task_id_for_spawn, "Failed to save failed session"); + } + publish_subagent_error(&bus, &session_done, &e.to_string(), &trace_id_owned).await; + } + } + // _registry_guard 在此 drop,确定性清理 cancel_registry 条目 + // (替代原末行手动 remove,覆盖 panic/early return 全路径) + }); + + // 8c. 立即返回 running 占位结果 + // 注意: summary 留空,output 只含引导信息。LLM 看到 running 后应调 wait_for_subagents。 + return Ok(TaskToolResult { + status: "running".to_string(), + summary: format!("Task {} spawned asynchronously", task_id), + output: format!( + "running, task_id={}. Call wait_for_subagents(timeout_secs) to wait for subagent completion and get results.", + task_id + ), + task_id, + }); + } + + // ===== 同步路径(子代理嵌套或无 sub_done_sender) ===== + // 9. 执行任务并处理结果 + let result = self.execute_task(agent, &session, &def, task.prompt.clone()).await; + match result { Ok(tool_result) => { let mut session = session; @@ -1010,6 +1321,75 @@ impl SubAgentRuntime for DefaultSubAgentRuntime { fn available_subagent_names(&self) -> Vec { self.subagent_runtime.available_names() } + + async fn cancel_pending_for_topic(&self, topic_id: &str) -> usize { + // 查询该 topic 下所有 running 的子代理 + let running = match self.store.list_pending_subagents(topic_id, Some("running")) { + Ok(records) => records, + Err(e) => { + tracing::warn!( + error = %e, + topic_id = %topic_id, + "Failed to list pending subagents for cancellation" + ); + return 0; + } + }; + + let count = running.len(); + if count == 0 { + return 0; + } + + tracing::info!( + topic_id = %topic_id, + count, + "Cancelling pending subagents for topic" + ); + + // 触发每个子代理的 CancellationToken + let registry = self.cancel_registry.lock(); + for record in &running { + if let Some(token) = registry.get(&record.task_id) { + token.cancel(); + tracing::info!( + task_id = %record.task_id, + "Cancelled subagent token" + ); + } else { + // token 不在 registry 中(可能已完成但 DB 状态未更新,或进程重启后丢失) + // 不变量 1:条件 UPDATE,仅在 status='running' 时转为 cancelled, + // 避免 spawn 已完成的终态被覆盖(completed → cancelled 是非法转换) + match self + .store + .try_update_pending_subagent_status(&record.task_id, "running", "cancelled") + { + Ok(true) => { + tracing::info!( + task_id = %record.task_id, + "Marked subagent as cancelled in DB (token not in registry)" + ); + } + Ok(false) => { + tracing::info!( + task_id = %record.task_id, + "Subagent status already updated by another path, skip cancel" + ); + } + Err(e) => { + tracing::warn!( + error = %e, + task_id = %record.task_id, + "Failed to mark subagent as cancelled in DB" + ); + } + } + } + } + drop(registry); + + count + } } /// 子代理定义目录 diff --git a/src/tools/task/types.rs b/src/tools/task/types.rs index 877de98..74c5643 100644 --- a/src/tools/task/types.rs +++ b/src/tools/task/types.rs @@ -17,6 +17,8 @@ pub enum TaskSessionState { Failed, /// 已超时 Timeout, + /// 已取消(用户 /stop 传播) + Cancelled, /// 状态未知(如重启后从 DB 重建时无法可靠推断原状态) Unknown, } @@ -27,6 +29,19 @@ impl Default for TaskSessionState { } } +impl TaskSessionState { + pub fn as_str(&self) -> &'static str { + match self { + Self::Running => "running", + Self::Completed => "completed", + Self::Failed => "failed", + Self::Timeout => "timeout", + Self::Cancelled => "cancelled", + Self::Unknown => "unknown", + } + } +} + /// 子代理来源 #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] @@ -223,6 +238,13 @@ impl TaskSession { self.error = Some("Task execution timed out".to_string()); self.updated_at = current_timestamp(); } + + /// 标记取消 + pub fn mark_cancelled(&mut self) { + self.state = TaskSessionState::Cancelled; + self.error = Some("Task cancelled by user".to_string()); + self.updated_at = current_timestamp(); + } } /// 任务工具参数 @@ -271,7 +293,7 @@ pub struct TaskHandle { /// 任务执行结果 #[derive(Debug, Clone, Serialize)] pub struct TaskToolResult { - /// 状态: success/failed/timeout + /// 状态: success/failed/timeout/running pub status: String, /// 任务完成总结 pub summary: String, @@ -280,3 +302,26 @@ pub struct TaskToolResult { /// 会话 ID(用于恢复) pub task_id: String, } + +/// 异步子代理完成状态 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum SubagentStatus { + Completed, + Failed, + Timeout, + Cancelled, +} + +/// 异步子代理完成回调内容(通过 sub_done_q 传递给 wait 工具) +#[derive(Debug, Clone)] +pub struct SubagentResult { + /// 完成的子代理 task_id + pub task_id: String, + /// 完成状态 + pub status: SubagentStatus, + /// 子代理输出(与 TaskToolResult.output 格式一致) + pub output: String, + /// 仍未完成的子代理 task_id 列表(供 LLM 判断全局进度) + pub pending_task_ids: Vec, +} diff --git a/src/tools/todo_read.rs b/src/tools/todo_read.rs index 6892dbb..7623826 100644 --- a/src/tools/todo_read.rs +++ b/src/tools/todo_read.rs @@ -185,6 +185,7 @@ mod tests { tool_call_id: None, parent_capability: None, trace_id: None, + ..Default::default() } } diff --git a/src/tools/todo_write.rs b/src/tools/todo_write.rs index 25d6f70..d873c14 100644 --- a/src/tools/todo_write.rs +++ b/src/tools/todo_write.rs @@ -487,6 +487,7 @@ mod tests { tool_call_id: None, parent_capability: None, trace_id: None, + ..Default::default() } } diff --git a/src/tools/traits.rs b/src/tools/traits.rs index 630f6df..af98442 100644 --- a/src/tools/traits.rs +++ b/src/tools/traits.rs @@ -1,6 +1,11 @@ +use std::time::Duration; +use std::sync::Arc; + use async_trait::async_trait; +use tokio::sync::{mpsc, watch}; use crate::domain::CapabilityPolicy; +use crate::tools::task::SubagentResult; #[derive(Debug, Clone)] pub struct ToolResult { @@ -9,7 +14,58 @@ pub struct ToolResult { pub error: Option, } -#[derive(Debug, Clone, Default)] +/// wait_for_subagents 工具等待期间的事件。 +#[derive(Debug, Clone)] +pub enum WaitEvent { + /// 一个子代理完成,携带其结果 + SubagentResult(SubagentResult), + /// wait 期间有新用户消息到达(已注入 history),携带新用户消息的内容列表 + UserMessage(Vec), + /// 等待超时 + Timeout, + /// 收到取消信号(/stop),等待已优雅终止,状态已清理完毕 + Cancelled, +} + +/// wait_for_subagents 工具的协调器接口。 +/// +/// 封装了释放/重获取 serial_lock + select! 等待 + 管理等待状态的逻辑, +/// 使 wait 工具不直接依赖 gateway 模块(避免循环依赖)。 +/// +/// 具体实现 `SessionWaitCoordinator` 在 gateway 模块中,持有 Session 引用。 +#[async_trait] +pub trait WaitCoordinator: Send + Sync + 'static { + /// 查询当前 topic 下仍处于 running 状态的子代理 task_id 列表。 + fn query_pending_task_ids(&self) -> Vec; + + /// 尝试排空 sub_done_q 中已缓冲的子代理结果(非阻塞 drain,但方法本身是 async + /// 因为需要获取 Session 锁)。 + /// + /// 用于两种场景: + /// 1. wait_for_subagents 入口处:即使 DB 中无 running 子代理, + /// 队列可能仍缓冲了已完成子代理的结果(子代理完成 → send 到队列 → + /// DB 更新为 completed,但 LLM 上轮未消费队列)。 + /// 2. wait() 进入 select! 前:多个子代理同时完成时,批量消费避免 + /// 每个结果各触发一次 LLM 调用。 + async fn try_drain_queued_results(&self) -> Vec; + + /// 进入等待状态:释放 serial_lock → select! → 重获取 serial_lock。 + /// + /// 调用前提:serial_lock 已被执行路径获取,guard 存于 coordinator 内部。 + /// 返回后:serial_lock 已被重新获取,waiting 标志已清除。 + /// + /// `cancel_rx`:可选的取消信号接收端。当收到信号时(/stop 命令), + /// select! 立即返回 `WaitEvent::Cancelled`,并完成完整的状态清理 + ///(重获取锁、回填 guard、清除 is_waiting、归还 receiver)。 + /// 为 None 时退化为不检查取消(向后兼容,子代理场景)。 + async fn wait( + &self, + timeout: Duration, + cancel_rx: Option>, + ) -> WaitEvent; +} + +#[derive(Clone, Default)] pub struct ToolContext { pub channel_name: Option, pub sender_id: Option, @@ -35,6 +91,46 @@ pub struct ToolContext { /// 端到端追踪 ID(从 InboundMessage 继承,用于 tool 执行路径的日志关联)。 /// None 表示无追踪上下文(如子代理独立执行或测试环境)。 pub trace_id: Option, + /// 异步子代理完成队列的 sender(按 topic 隔离)。 + /// 仅主 agent(nesting_depth=0)有值:agent_factory 构建时从 SessionHistory 注入。 + /// TaskTool spawn 异步子代理后,子代理完成时通过此 sender 发送 SubagentResult, + /// 由 wait_for_subagents 工具的 receiver 端消费。 + /// 子代理自身(nesting_depth>0)为 None:嵌套层不支持异步,走同步路径。 + pub sub_done_sender: Option>, + /// wait_for_subagents 工具的协调器(仅主 agent 有值)。 + /// 封装了释放/重获取 serial_lock + select! 等待逻辑。 + /// wait 工具通过此接口实现真等待(释放锁让 process_one 注入用户消息)。 + pub wait_coordinator: Option>, + /// 取消信号接收端(仅主 agent 有值,由 agent_factory 从 cancel_token clone 注入)。 + /// wait_for_subagents 工具将其传给 coordinator.wait() 的 select!, + /// 使 /stop 命令能立即中断等待并完成状态清理。 + /// watch::Receiver 可安全 clone,多个 receiver 共享同一 sender, + /// 互不影响各自的 has_changed() / changed() 状态。 + pub cancel_rx: Option>, +} + +impl std::fmt::Debug for ToolContext { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ToolContext") + .field("channel_name", &self.channel_name) + .field("sender_id", &self.sender_id) + .field("chat_id", &self.chat_id) + .field("session_id", &self.session_id) + .field("topic_id", &self.topic_id) + .field("message_id", &self.message_id) + .field("message_seq", &self.message_seq) + .field("subagent_description", &self.subagent_description) + .field("nesting_depth", &self.nesting_depth) + .field("task_id", &self.task_id) + .field("parent_task_id", &self.parent_task_id) + .field("tool_call_id", &self.tool_call_id) + .field("parent_capability", &self.parent_capability) + .field("trace_id", &self.trace_id) + .field("sub_done_sender", &self.sub_done_sender) + .field("wait_coordinator", &self.wait_coordinator.is_some()) + .field("cancel_rx", &self.cancel_rx.is_some()) + .finish() + } } #[async_trait] diff --git a/src/tools/wait_tool.rs b/src/tools/wait_tool.rs new file mode 100644 index 0000000..54eca29 --- /dev/null +++ b/src/tools/wait_tool.rs @@ -0,0 +1,250 @@ +use std::time::Duration; + +use async_trait::async_trait; +use serde_json::json; + +use crate::tools::{Tool, ToolContext, ToolResult, WaitEvent}; + +/// wait_for_subagents 工具 — 等待异步子代理完成或用户消息到达。 +/// +/// 调用后释放 serial_lock,进入 select! 等待: +/// - 子代理完成 → 返回结果 + 未完成列表 +/// - 用户消息到达 → 返回 "有新用户消息"(消息已注入 history) +/// - 超时 → 返回超时 + 未完成列表 +/// +/// 等待结束后重新获取 serial_lock,保证后续工具调用串行。 +pub struct WaitForSubagentsTool { + /// 默认超时(秒),LLM 未指定时使用 + default_timeout_secs: u64, +} + +impl WaitForSubagentsTool { + pub const TOOL_NAME: &'static str = "wait_for_subagents"; + + pub fn new(default_timeout_secs: u64) -> Self { + Self { + default_timeout_secs, + } + } +} + +#[async_trait] +impl Tool for WaitForSubagentsTool { + fn name(&self) -> &str { + Self::TOOL_NAME + } + + fn description(&self) -> &str { + "Wait for asynchronous subagents to complete, or for new user messages to arrive. \ + Use this after launching subagents via the task tool to receive their results. \ + Returns the first completed subagent's result and a list of still-pending task IDs. \ + If pending_task_ids is non-empty, call this tool again to wait for the next one." + } + + fn parameters_schema(&self) -> serde_json::Value { + json!({ + "type": "object", + "properties": { + "timeout_secs": { + "type": "integer", + "description": "Maximum seconds to wait. Default 60. The tool returns early if a subagent completes or a user message arrives.", + "default": 60 + } + }, + "required": [] + }) + } + + fn read_only(&self) -> bool { + false + } + + fn exclusive(&self) -> bool { + // wait 工具释放/重获取 serial_lock,不应与其他工具并发 + true + } + + async fn execute(&self, _args: serde_json::Value) -> anyhow::Result { + Ok(ToolResult { + success: false, + output: String::new(), + error: Some( + "wait_for_subagents requires tool context with wait_coordinator".to_string(), + ), + }) + } + + async fn execute_with_context( + &self, + context: &ToolContext, + args: serde_json::Value, + ) -> anyhow::Result { + // 1. 获取 wait_coordinator(仅主 agent 有值) + let coordinator = match &context.wait_coordinator { + Some(c) => c.clone(), + None => { + return Ok(ToolResult { + success: false, + output: String::new(), + error: Some( + "wait_for_subagents is not available in this context (no wait_coordinator)" + .to_string(), + ), + }); + } + }; + + // 2. 解析超时参数 + let timeout_secs = args + .get("timeout_secs") + .and_then(|v| v.as_u64()) + .unwrap_or(self.default_timeout_secs); + let timeout = Duration::from_secs(timeout_secs); + + // 3. 先尝试排空队列中已缓冲的结果 + // 场景:子代理已完成 → 结果 send 到队列 + DB 更新为 completed, + // 但 LLM 上轮未调用 wait(或 wait 超时未消费)→ 结果缓冲在队列中。 + // 若不排空,query_pending_task_ids 返回空(DB 已非 running)→ + // 返回 "No pending" → 缓冲结果永远丢失。 + let drained = coordinator.try_drain_queued_results().await; + if !drained.is_empty() { + let pending_after = coordinator.query_pending_task_ids(); + let pending_str = if pending_after.is_empty() { + "none".to_string() + } else { + pending_after.join(", ") + }; + let formatted: Vec = drained + .iter() + .map(|r| { + format!( + "Subagent {} completed (status: {:?}). Output: {}", + r.task_id, r.status, r.output + ) + }) + .collect(); + return Ok(ToolResult { + success: true, + output: format!( + "Retrieved {} buffered subagent result(s):\n{}\nStill pending: [{}]", + drained.len(), + formatted.join("\n"), + pending_str + ), + error: None, + }); + } + + // 4. 查询 pending 子代理 + let pending = coordinator.query_pending_task_ids(); + if pending.is_empty() { + return Ok(ToolResult { + success: true, + output: "No pending subagents to wait for.".to_string(), + error: None, + }); + } + + tracing::info!( + topic_id = ?context.topic_id, + pending_count = pending.len(), + timeout_secs, + "wait_for_subagents: entering wait" + ); + + // 4. 进入等待(coordinator 内部:释放锁 → select! → 重获取锁) + // 传入 cancel_rx 使 /stop 命令能立即中断等待。 + // coordinator 在 select! 中以 biased 优先级处理: + // 子代理结果 > 用户消息 > 取消信号 > 超时 + let event = coordinator + .wait(timeout, context.cancel_rx.clone()) + .await; + + // 5. 格式化返回结果 + let output = match event { + WaitEvent::SubagentResult(result) => { + let pending_str = if result.pending_task_ids.is_empty() { + "none".to_string() + } else { + result.pending_task_ids.join(", ") + }; + format!( + "Subagent {} completed (status: {:?}). Output: {}\nStill pending: [{}]", + result.task_id, result.status, result.output, pending_str + ) + } + WaitEvent::UserMessage(messages) => { + let pending = coordinator.query_pending_task_ids(); + let pending_str = if pending.is_empty() { + "none".to_string() + } else { + pending.join(", ") + }; + if messages.is_empty() { + format!( + "A new user message arrived while waiting (content could not be retrieved). \ + Still pending subagents: [{}]", + pending_str + ) + } else { + let formatted_msgs: Vec = messages + .iter() + .enumerate() + .map(|(i, msg)| format!(" [{}] {}", i + 1, msg)) + .collect(); + format!( + "New user message(s) arrived while waiting:\n{}\n\ + These messages have been added to the conversation history. \ + Still pending subagents: [{}]", + formatted_msgs.join("\n"), + pending_str + ) + } + } + WaitEvent::Timeout => { + let pending = coordinator.query_pending_task_ids(); + format!( + "Wait timed out after {}s. Still pending subagents: [{}]", + timeout_secs, + if pending.is_empty() { + "none".to_string() + } else { + pending.join(", ") + } + ) + } + WaitEvent::Cancelled => { + // /stop 命令中断了等待。coordinator 已完成全部状态清理 + //(重获取 serial_lock、回填 guard、清除 is_waiting、归还 receiver)。 + // 返回提示性输出,Agent 下一轮迭代会检测到 cancel 并退出。 + let pending = coordinator.query_pending_task_ids(); + tracing::info!( + topic_id = ?context.topic_id, + pending_count = pending.len(), + "wait_for_subagents: cancelled by /stop" + ); + format!( + "Wait was cancelled by /stop command. \ + Pending subagents (if any) have been cancelled separately. \ + Still pending in DB: [{}]", + if pending.is_empty() { + "none".to_string() + } else { + pending.join(", ") + } + ) + } + }; + + tracing::info!( + topic_id = ?context.topic_id, + "wait_for_subagents: returning result" + ); + + Ok(ToolResult { + success: true, + output, + error: None, + }) + } +}