// src/agent/runtime/executor/mod.rs // // 工具调用验证与并行执行器。 mod helpers; // Re-export 公共类型 pub use helpers::{PreparedCall, ToolExecutionResult, ToolResultMessage}; use futures_util::stream::FuturesUnordered; use futures_util::StreamExt; use sqlx::SqlitePool; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; use tokio::sync::{mpsc, oneshot}; use tracing::{info, warn}; use crate::api::{AppState, PendingPermission}; use crate::clients::llm::{ChatMessage, ToolCall}; use super::checkpoint::CheckpointManager; use super::denial_tracker::DenialTracker; use super::file_cache::FileStateCache; use super::hardline; use super::partitioner::ToolPartitioner; use super::permission::{PermissionChecker, PermissionResult}; use super::permission_explainer::explain_permission; use super::{AgentStreamEvent, DuplicateDetector}; use crate::agent::hooks::{event_label, HookRegistry, PreToolUseContext}; use crate::agent::tools::{ToolContext, ToolRegistry}; use helpers::{execute_single_tool, process_single_result, save_tool_message_sync}; /// 验证工具调用:死循环检测 + 参数解析。 /// /// 返回 (prepared_calls, has_duplicate)。 /// 死循环或参数无效时,错误消息直接注入到 messages。 #[allow(clippy::too_many_arguments)] pub fn validate_and_prepare( tool_calls: &[ToolCall], duplicate_detector: &mut DuplicateDetector, duplicate_threshold: usize, messages: &mut Vec, tx: &mpsc::UnboundedSender, db: &SqlitePool, session_id: &str, turn_index: i32, step: usize, ) -> (Vec, bool) { let mut prepared_calls: Vec = Vec::new(); let mut has_duplicate = false; for tool_call in tool_calls { let tool_name = &tool_call.function.name; let tool_args_str = &tool_call.function.arguments; // 确保每个工具调用有唯一 ID(LLM 可能不返回 id) let call_id = if tool_call.id.is_empty() { format!("call_{}", &uuid::Uuid::new_v4().to_string()[..8]) } else { tool_call.id.clone() }; // 死循环检测 if duplicate_detector.record(tool_name, tool_args_str, duplicate_threshold) { warn!( "[Executor] 检测到死循环:{} 连续调用 {} 次", tool_name, duplicate_threshold ); let _ = tx.send(AgentStreamEvent::Error { message: format!("检测到工具 {} 的重复调用,已自动终止循环。", tool_name), }); let error_msg = ChatMessage::tool_result( &call_id, format!( "错误:工具 {} 被连续重复调用 {} 次,参数完全相同。\ 请停止重复调用并直接给出目前收集到的答案。", tool_name, duplicate_threshold ), ); messages.push(error_msg); has_duplicate = true; continue; } // 解析参数 let args: serde_json::Value = match serde_json::from_str(tool_args_str) { Ok(v) => v, Err(e) => { let error_output = format!("工具参数 JSON 解析失败: {}", e); let _ = tx.send(AgentStreamEvent::ToolResult { tool_call_id: call_id.clone(), name: tool_name.clone(), output: error_output.clone(), is_error: true, metadata: serde_json::json!({}), step, }); let tool_msg = ChatMessage::tool_result(&call_id, &error_output); save_tool_message_sync(db, session_id, turn_index, step, &tool_msg); messages.push(tool_msg); continue; } }; prepared_calls.push(PreparedCall { tool_call_id: call_id.clone(), tool_name: tool_name.clone(), args, }); } (prepared_calls, has_duplicate) } /// 并行执行所有准备好的工具调用。 /// /// 流程: /// 1. 权限检查(deny 规则阻止不可执行工具) /// 2. 发送 ToolCall SSE 事件 /// 3. 运行 PreToolUse hooks /// 4. 工具分区 + 并行执行(并发安全工具一批并行,不安全工具单独串行) /// 5. 收集结果、发送 ToolResult SSE、运行 PostToolUse hooks /// 6. 返回 ToolResultMessage 列表供调用方推入 messages #[allow(clippy::too_many_arguments)] pub async fn execute_parallel( prepared_calls: &[PreparedCall], tool_registry: &ToolRegistry, app_state: Arc, hook_registry: &HookRegistry, permission_checker: Option<&PermissionChecker>, session_permission_checker: Option<&PermissionChecker>, denial_tracker: Option<&std::sync::Mutex>, checkpoint_manager: Option<&std::sync::Arc>, tx: &mpsc::UnboundedSender, db: &SqlitePool, session_id: &str, agent_name: &str, turn_index: i32, step: usize, tool_timeout_secs: u64, max_output_chars: usize, read_file_state: Arc>, enable_thinking: bool, additional_allowed_dirs: Vec, ) -> ToolExecutionResult { if prepared_calls.is_empty() { return ToolExecutionResult { tool_messages: Vec::new(), was_cancelled: false, had_duplicate: false, hook_contexts: Vec::new(), blocking_errors: Vec::new(), }; } let sid = session_id.to_string(); // Phase 1: 发送 ToolCall SSE 事件 for prep in prepared_calls { let _ = tx.send(AgentStreamEvent::ToolCall { id: prep.tool_call_id.clone(), name: prep.tool_name.clone(), arguments: prep.args.clone(), step, }); } // Phase 2: PreToolUse hooks — 收集修改后的参数和附加上下文 let exec_start = std::time::Instant::now(); let mut mutated_args: Vec = Vec::new(); let mut additional_contexts: Vec = Vec::new(); let mut hook_permission_info: Vec> = Vec::new(); // ^^^ (permission_desc, hook_tool_name) let mut hook_blocking_errors: Vec = Vec::new(); for prep in prepared_calls { let hook_ctx = PreToolUseContext { session_id: sid.clone(), tool_name: prep.tool_name.clone(), tool_args: prep.args.clone(), step, }; let result = hook_registry.run_pre_tool_use(&hook_ctx).await; if result.action.is_blocked() { let reason = result.action.block_reason().unwrap_or("unknown"); warn!( "[Executor] PreToolUse hook 阻止了 {} 的执行: {}", prep.tool_name, reason ); } // 收集所有阻塞错误详情(含多个 hook 同时 block 的情况) for be in &result.blocking_errors { hook_blocking_errors.push(format!( "[{}] 阻止 {}: {}", be.hook_name, prep.tool_name, be.reason )); } // 收集 hook 的权限请求(保留完整信息用于 AskUser prompt) if let Some((permission, tool_name)) = result.permission_info() { info!( "[Executor] Hook 请求了工具 {} 的权限确认: {}", prep.tool_name, permission ); hook_permission_info.push(Some((permission.to_string(), tool_name.to_string()))); } else { hook_permission_info.push(None); } // 使用 hook 可能修改后的参数 mutated_args.push(result.final_args); // 收集所有 hook 注入的上下文(优先使用带来源标记的 tagged_contexts) if !result.tagged_contexts.is_empty() { for tc in &result.tagged_contexts { additional_contexts.push(format!( "[Hook: {} | {}] {}", tc.hook_name, event_label(tc.source_event), tc.content, )); } } else { for ctx in &result.additional_contexts { additional_contexts.push(ctx.clone()); } } } // Phase 2.5: 权限检查 — PermissionChecker 规则引擎拦截被拒绝的工具。 // 被拒绝的工具直接注入错误 result,不进入执行队列。 let mut tool_messages: Vec = Vec::new(); let mut denied_indices: std::collections::HashSet = std::collections::HashSet::new(); // ── Hardline 预检查(在任何模式下都不可绕过)── // 在 PermissionChecker 之前执行,确保 hardline 规则始终生效。 for (i, prep) in prepared_calls.iter().enumerate() { let hardline_result = match prep.tool_name.as_str() { "run_bash" => { if let Some(cmd) = prep.args.get("command").and_then(|v| v.as_str()) { hardline::check_command(cmd) } else { hardline::HardlineResult::allowed() } } "file_write" | "file_edit" => { if let Some(path) = prep.args.get("file_path").and_then(|v| v.as_str()) { hardline::check_dangerous_path(path) } else if let Some(path) = prep.args.get("path").and_then(|v| v.as_str()) { hardline::check_dangerous_path(path) } else { hardline::HardlineResult::allowed() } } _ => hardline::HardlineResult::allowed(), }; if hardline_result.blocked { warn!( "[Executor] Hardline 阻止了工具 {} (category={}): {}", prep.tool_name, hardline_result.category.as_deref().unwrap_or("unknown"), hardline_result.reason ); let err_output = hardline_result.reason.clone(); let _ = tx.send(AgentStreamEvent::ToolResult { tool_call_id: prep.tool_call_id.clone(), name: prep.tool_name.clone(), output: err_output.clone(), is_error: true, metadata: serde_json::json!({ "hardline_blocked": true, "hardline_category": hardline_result.category, }), step, }); let err_msg = ChatMessage::tool_result(&prep.tool_call_id, &err_output); save_tool_message_sync(db, &sid, turn_index, step, &err_msg); tool_messages.push(ToolResultMessage { chat_message: err_msg, was_error: true, }); // 记录拒绝追踪 if let Some(dt) = denial_tracker { if let Ok(mut tracker) = dt.lock() { tracker.record_denial(); } } denied_indices.insert(i); } } if let Some(checker) = permission_checker { for (i, prep) in prepared_calls.iter().enumerate() { let mut perm_result = checker.check(&prep.tool_name, Some(&prep.args)); perm_result = checker.apply_mode(perm_result, &prep.tool_name); // Hook PermissionRequired — 若 Checker 返回 Allowed,升级为 Ask // 使用 hook 提供的具体权限描述替换泛型消息 if let Some(Some((ref perm_desc, _))) = hook_permission_info.get(i) { if perm_result.is_allowed() { perm_result = PermissionResult::AskUser { message: format!( "[Hook 权限请求] {}\n\n工具: {}\n参数: {}", perm_desc, prep.tool_name, serde_json::to_string_pretty(&prep.args).unwrap_or_default(), ), }; } } // 工具级 check_permissions() — 在 PermissionChecker 结果基础上叠加 // PermissionChecker Deny/Ask 优先,工具级规则在 Allow 时可升级为 Ask if let Some(tool) = tool_registry.get(&prep.tool_name) { let tool_rules = tool.check_permissions(&prep.args); for tool_rule in &tool_rules { match tool_rule { crate::agent::tools::PermissionRule::Deny { reason, .. } => { // 工具级 Deny 仅在 PermissionChecker 未 Deny 时生效 if !perm_result.is_denied() { perm_result = PermissionResult::Denied { reason: reason.clone(), }; } } crate::agent::tools::PermissionRule::Ask { message, .. } => { // 工具级 Ask:若 PermissionChecker 返回 Allowed,升级为 Ask if perm_result.is_allowed() { perm_result = PermissionResult::AskUser { message: message.clone(), }; } } _ => {} } } } // 会话级权限检查(API 动态添加的规则,优先级高于环境变量规则) if let Some(session_checker) = session_permission_checker { let session_result = session_checker.check(&prep.tool_name, Some(&prep.args)); // 会话规则结果覆盖或升级 match session_result { PermissionResult::Denied { reason } => { // 会话 Deny 强制覆盖 perm_result = PermissionResult::Denied { reason }; } PermissionResult::AskUser { message } => { // 会话 Ask 在 Allow 时升级 if perm_result.is_allowed() { perm_result = PermissionResult::AskUser { message }; } } PermissionResult::Allowed => { // 会话 Allow 仅覆盖 Allowed,保持 Deny/AskUser 不变 // 避免覆盖工具级 check_permissions() 升级的 AskUser } } } match perm_result { PermissionResult::Denied { reason } => { warn!( "[Executor] PermissionChecker 拒绝了工具 {}: {}", prep.tool_name, reason ); let err_output = format!("工具 {} 被权限规则拒绝执行: {}", prep.tool_name, reason); let _ = tx.send(AgentStreamEvent::ToolResult { tool_call_id: prep.tool_call_id.clone(), name: prep.tool_name.clone(), output: err_output.clone(), is_error: true, metadata: serde_json::json!({}), step, }); let err_msg = ChatMessage::tool_result(&prep.tool_call_id, &err_output); save_tool_message_sync(db, &sid, turn_index, step, &err_msg); tool_messages.push(ToolResultMessage { chat_message: err_msg, was_error: true, }); // 记录拒绝追踪 if let Some(dt) = denial_tracker { if let Ok(mut tracker) = dt.lock() { tracker.record_denial(); } } denied_indices.insert(i); } PermissionResult::AskUser { message } => { info!( "[Executor] PermissionChecker 请求用户确认工具 {}: {}", prep.tool_name, message ); // 生成权限风险解释 let permission_exp = explain_permission(&prep.tool_name, &prep.args); let explanation_json = serde_json::to_value(&permission_exp).ok(); // 发送权限请求 SSE 事件 let _ = tx.send(AgentStreamEvent::PermissionRequest { tool_call_id: prep.tool_call_id.clone(), tool_name: prep.tool_name.clone(), message: message.clone(), arguments: prep.args.clone(), explanation: explanation_json, }); // 创建 oneshot 通道等待用户响应 let (resp_tx, resp_rx) = oneshot::channel(); let perm_id = uuid::Uuid::new_v4().to_string(); let tc_id = prep.tool_call_id.clone(); let t_name = prep.tool_name.clone(); // 存储待处理的权限请求 { let mut perms = match app_state.pending_permissions.lock() { Ok(p) => p, Err(e) => { warn!("[Executor] 权限系统内部错误: {}", e); let err_output = format!("权限系统内部错误,工具 {} 被拒绝", prep.tool_name); let _ = tx.send(AgentStreamEvent::ToolResult { tool_call_id: prep.tool_call_id.clone(), name: prep.tool_name.clone(), output: err_output.clone(), is_error: true, metadata: serde_json::json!({}), step, }); let err_msg = ChatMessage::tool_result(&prep.tool_call_id, &err_output); save_tool_message_sync(db, &sid, turn_index, step, &err_msg); tool_messages.push(ToolResultMessage { chat_message: err_msg, was_error: true, }); // 内部错误 → 记录拒绝追踪 if let Some(dt) = denial_tracker { if let Ok(mut tracker) = dt.lock() { tracker.record_denial(); } } denied_indices.insert(i); continue; } }; perms.insert( perm_id.clone(), PendingPermission { tool_call_id: tc_id.clone(), tool_name: t_name.clone(), message: message.clone(), arguments: prep.args.clone(), response_tx: resp_tx, }, ); } // 等待用户响应(120 秒超时) let timeout_dur = std::time::Duration::from_secs(120); let perm_result = tokio::time::timeout(timeout_dur, resp_rx).await; // 清理待处理的权限请求 if let Ok(mut perms) = app_state.pending_permissions.lock() { perms.remove(&perm_id); } match perm_result { Ok(Ok(response)) if response.allowed => { info!("[Executor] 用户允许了工具 {} 的执行", prep.tool_name); let _ = tx.send(AgentStreamEvent::PermissionResponse { tool_call_id: prep.tool_call_id.clone(), allowed: true, }); // 用户允许 → 重置连续拒绝计数 if let Some(dt) = denial_tracker { if let Ok(mut tracker) = dt.lock() { tracker.record_success(); } } } Ok(Ok(_response)) => { // 用户拒绝 info!("[Executor] 用户拒绝了工具 {}", prep.tool_name); let err_output = format!("用户拒绝了工具 {} 的执行", prep.tool_name); let _ = tx.send(AgentStreamEvent::ToolResult { tool_call_id: prep.tool_call_id.clone(), name: prep.tool_name.clone(), output: err_output.clone(), is_error: true, metadata: serde_json::json!({}), step, }); let err_msg = ChatMessage::tool_result(&prep.tool_call_id, &err_output); save_tool_message_sync(db, &sid, turn_index, step, &err_msg); tool_messages.push(ToolResultMessage { chat_message: err_msg, was_error: true, }); // 用户拒绝 → 记录拒绝追踪 if let Some(dt) = denial_tracker { if let Ok(mut tracker) = dt.lock() { tracker.record_denial(); } } denied_indices.insert(i); let _ = tx.send(AgentStreamEvent::PermissionResponse { tool_call_id: prep.tool_call_id.clone(), allowed: false, }); } _ => { // 超时或通道关闭 warn!("[Executor] 权限请求超时或取消: {}", prep.tool_name); let err_output = format!("权限请求超时 (120s): {} 未获得用户确认", prep.tool_name); let _ = tx.send(AgentStreamEvent::ToolResult { tool_call_id: prep.tool_call_id.clone(), name: prep.tool_name.clone(), output: err_output.clone(), is_error: true, metadata: serde_json::json!({}), step, }); let err_msg = ChatMessage::tool_result(&prep.tool_call_id, &err_output); save_tool_message_sync(db, &sid, turn_index, step, &err_msg); tool_messages.push(ToolResultMessage { chat_message: err_msg, was_error: true, }); // 超时 → 记录拒绝追踪 if let Some(dt) = denial_tracker { if let Ok(mut tracker) = dt.lock() { tracker.record_denial(); } } denied_indices.insert(i); let _ = tx.send(AgentStreamEvent::PermissionResponse { tool_call_id: prep.tool_call_id.clone(), allowed: false, }); } } } PermissionResult::Allowed => { // 工具被允许 → 重置连续拒绝计数 if let Some(dt) = denial_tracker { if let Ok(mut tracker) = dt.lock() { tracker.record_success(); } } } } } } // if let Some(checker) // Phase 3: 分区并行执行(参考 Claude Code partitionToolCalls + runTools)。 // // 改进:原实现将所有非拒绝工具放入单个 FuturesUnordered 无差别并发, // 可能导致非并发安全工具(如 run_bash)错误地并行执行。 // 新实现使用 ToolPartitioner 将工具按并发安全性分批: // - 连续的并发安全工具放入同一个并行批次(FuturesUnordered) // - 非并发安全工具独占一个串行批次(逐次执行) // 批次内工具执行完成后立即推送 SSE 事件,不等待整个批次完成。 let cancelled = Arc::new(AtomicBool::new(false)); let cancel_flag = cancelled.clone(); let app_state_ref = app_state.clone(); let sid_ref = sid.clone(); let cancel_handle = tokio::spawn(async move { loop { tokio::time::sleep(std::time::Duration::from_millis(250)).await; if let Ok(locked) = app_state_ref.cancelled_runs.lock() { if locked.contains(&sid_ref) { cancel_flag.store(true, Ordering::SeqCst); return; } } } }); let timeout_dur = std::time::Duration::from_secs(tool_timeout_secs); // ── Checkpoint 预触发:对文件变更类工具在执行前创建快照 ── if let Some(ckpt) = checkpoint_manager { let cwd = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from(".")); for prep in prepared_calls .iter() .filter(|p| CheckpointManager::should_checkpoint(&p.tool_name)) { ckpt.ensure_checkpoint(&cwd, &format!("pre-{}", prep.tool_name)); } } // ── Phase 3a: 构建不包含被拒绝工具的 (原索引, PreparedCall) 映射 ── let non_denied: Vec<(usize, &PreparedCall)> = prepared_calls .iter() .enumerate() .filter(|(i, _)| !denied_indices.contains(i)) .collect(); // ── Phase 3b: 分区 ── let non_denied_calls: Vec = non_denied.iter().map(|(_, p)| (*p).clone()).collect(); let partitioner = ToolPartitioner::new(10); let batches = partitioner.partition(&non_denied_calls, tool_registry); // 预设非拒绝工具中哪些原索引属于已拒绝列表(不会有,但安全起见) let original_index_of: std::collections::HashMap = non_denied .iter() .map(|(orig_idx, prep)| (prep.tool_call_id.clone(), *orig_idx)) .collect(); info!( "[Executor] 工具分区完成: {} 工具 → {} 批次 ({} 串行 + {} 并行)", non_denied.len(), batches.len(), batches.iter().filter(|b| !b.is_parallel).count(), batches.iter().filter(|b| b.is_parallel).count(), ); let mut was_cancelled = false; // ── Phase 3c: 逐批次执行 ── // 批次之间串行;并行批次内工具并发执行;串行批次内工具逐个执行。 for batch in &batches { if was_cancelled { break; } if batch.is_parallel { // ── 并行批次:FuturesUnordered 并发执行 ── let mut exec_futs: FuturesUnordered<_> = batch .calls .iter() .map(|prep| { let orig_idx = original_index_of .get(&prep.tool_call_id) .copied() .unwrap_or(0); let tool_name = prep.tool_name.clone(); let args = mutated_args .get(orig_idx) .cloned() .unwrap_or_else(|| prep.args.clone()); let tool_ctx = ToolContext::with_file_cache(app_state.clone(), read_file_state.clone()) .with_sse_tx(tx.clone()) .with_session_id(session_id.to_string()) .with_thinking(enable_thinking) .with_additional_dirs(additional_allowed_dirs.clone()) .with_tool_call_id(prep.tool_call_id.clone()) .with_max_output_chars(max_output_chars); let cancelled = cancelled.clone(); let tool_opt = tool_registry.get(&tool_name); Box::pin(async move { let output = execute_single_tool( tool_opt, args, &tool_ctx, &cancelled, timeout_dur, &tool_name, ) .await; let was_cancelled = cancelled.load(Ordering::SeqCst); ( prep.tool_call_id.clone(), prep.tool_name.clone(), prep.args.clone(), output, was_cancelled, ) }) }) .collect(); // 渐进式处理:每个工具一完成就处理 while let Some((tool_call_id, tool_name, tool_args, output, cancelled_flag)) = exec_futs.next().await { if cancelled_flag { was_cancelled = true; } process_single_result( &tool_call_id, &tool_name, &tool_args, &output, cancelled_flag, exec_start, tx, hook_registry, &app_state.config.library_dir, &sid, agent_name, step, max_output_chars, &mut tool_messages, &mut additional_contexts, db, turn_index, ) .await; } } else { // ── 串行批次:逐个执行 ── for prep in &batch.calls { let orig_idx = original_index_of .get(&prep.tool_call_id) .copied() .unwrap_or(0); let tool_name = prep.tool_name.clone(); let args = mutated_args .get(orig_idx) .cloned() .unwrap_or_else(|| prep.args.clone()); let tool_ctx = ToolContext::with_file_cache(app_state.clone(), read_file_state.clone()) .with_sse_tx(tx.clone()) .with_session_id(session_id.to_string()) .with_thinking(enable_thinking) .with_additional_dirs(additional_allowed_dirs.clone()); let tool_opt = tool_registry.get(&tool_name); let output = execute_single_tool( tool_opt, args, &tool_ctx, &cancelled, timeout_dur, &tool_name, ) .await; let cancelled_flag = cancelled.load(Ordering::SeqCst); if cancelled_flag { was_cancelled = true; } process_single_result( &prep.tool_call_id, &tool_name, &prep.args, &output, cancelled_flag, exec_start, tx, hook_registry, &app_state.config.library_dir, &sid, agent_name, step, max_output_chars, &mut tool_messages, &mut additional_contexts, db, turn_index, ) .await; if was_cancelled { break; } } } } cancel_handle.abort(); ToolExecutionResult { tool_messages, was_cancelled, had_duplicate: false, hook_contexts: additional_contexts, blocking_errors: hook_blocking_errors, } }