From 602fd073847d6a903ef88e33323c6d33b259beea Mon Sep 17 00:00:00 2001 From: coso Date: Wed, 31 Dec 2025 19:35:02 +0800 Subject: [PATCH] feat: add agent tool calling frontend integration (v0.25.0) - Add StreamEvent types (TextDelta, ToolStart, ToolEnd, Done, Error) - Add TokenUsage and ToolExecutionResult types - Add ToolCallState for UI state management - Create ToolCallDisplay component for tool execution status - Create StreamingRenderer for real-time markdown rendering - Create TokenUsageDisplay for token usage statistics - Update MessageList to integrate tool calls and streaming - Fix lint errors and format code - Bump version to 0.25.0 Requirements: 9.1, 9.2, 9.3, 9.4, 9.5 --- package.json | 2 +- src-tauri/Cargo.lock | 2 +- src-tauri/Cargo.toml | 2 +- .../agent/native_agent.txt | 7 + .../proptest-regressions/agent/tools/bash.txt | 7 + .../agent/tools/read_file.txt | 8 + src-tauri/src/agent/README.md | 33 +- src-tauri/src/agent/mod.rs | 4 + src-tauri/src/agent/native_agent.rs | 933 +++++++++- src-tauri/src/agent/tool_loop.rs | 1124 ++++++++++++ src-tauri/src/agent/tools/README.md | 284 +++ src-tauri/src/agent/tools/bash.rs | 1062 +++++++++++ src-tauri/src/agent/tools/edit_file.rs | 1628 +++++++++++++++++ src-tauri/src/agent/tools/mod.rs | 32 + src-tauri/src/agent/tools/prompt.rs | 613 +++++++ src-tauri/src/agent/tools/read_file.rs | 1015 ++++++++++ src-tauri/src/agent/tools/registry.rs | 390 ++++ src-tauri/src/agent/tools/security.rs | 594 ++++++ src-tauri/src/agent/tools/types.rs | 554 ++++++ src-tauri/src/agent/tools/write_file.rs | 802 ++++++++ src-tauri/src/agent/types.rs | 141 +- src-tauri/tauri.conf.json | 2 +- src/assets/models/claude.png | Bin 0 -> 930 bytes src/assets/models/deepseek.png | Bin 0 -> 1190 bytes src/assets/models/doubao.png | Bin 0 -> 8854 bytes src/assets/models/ernie.png | Bin 0 -> 1256 bytes src/assets/models/gemini.png | Bin 0 -> 807 bytes src/assets/models/glm.png | Bin 0 -> 1756 bytes src/assets/models/jimeng.png | Bin 0 -> 1187 bytes src/assets/models/kimi.png | Bin 0 -> 890 bytes src/assets/platforms/csdn.png | Bin 0 -> 4353 bytes src/assets/platforms/gzh.png | Bin 0 -> 32374 bytes src/assets/platforms/juejin.png | Bin 0 -> 5973 bytes src/assets/platforms/toutiao.png | Bin 0 -> 7995 bytes src/assets/platforms/xhs.png | Bin 0 -> 4622 bytes src/assets/platforms/zhihu.png | Bin 0 -> 4073 bytes .../agent/chat/components/ChatNavbar.tsx | 10 +- .../agent/chat/components/EmptyState.tsx | 789 +++++++- .../agent/chat/components/MessageList.tsx | 15 + .../chat/components/StreamingRenderer.tsx | 106 ++ .../chat/components/TokenUsageDisplay.tsx | 69 + .../agent/chat/components/ToolCallDisplay.tsx | 288 +++ src/components/agent/chat/index.tsx | 28 +- src/components/agent/chat/types.ts | 6 + src/components/ui/select.tsx | 32 +- src/lib/api/agent.ts | 154 ++ 46 files changed, 10614 insertions(+), 122 deletions(-) create mode 100644 src-tauri/proptest-regressions/agent/native_agent.txt create mode 100644 src-tauri/proptest-regressions/agent/tools/bash.txt create mode 100644 src-tauri/proptest-regressions/agent/tools/read_file.txt create mode 100644 src-tauri/src/agent/tool_loop.rs create mode 100644 src-tauri/src/agent/tools/README.md create mode 100644 src-tauri/src/agent/tools/bash.rs create mode 100644 src-tauri/src/agent/tools/edit_file.rs create mode 100644 src-tauri/src/agent/tools/mod.rs create mode 100644 src-tauri/src/agent/tools/prompt.rs create mode 100644 src-tauri/src/agent/tools/read_file.rs create mode 100644 src-tauri/src/agent/tools/registry.rs create mode 100644 src-tauri/src/agent/tools/security.rs create mode 100644 src-tauri/src/agent/tools/types.rs create mode 100644 src-tauri/src/agent/tools/write_file.rs create mode 100644 src/assets/models/claude.png create mode 100644 src/assets/models/deepseek.png create mode 100644 src/assets/models/doubao.png create mode 100644 src/assets/models/ernie.png create mode 100644 src/assets/models/gemini.png create mode 100644 src/assets/models/glm.png create mode 100644 src/assets/models/jimeng.png create mode 100644 src/assets/models/kimi.png create mode 100644 src/assets/platforms/csdn.png create mode 100644 src/assets/platforms/gzh.png create mode 100644 src/assets/platforms/juejin.png create mode 100644 src/assets/platforms/toutiao.png create mode 100644 src/assets/platforms/xhs.png create mode 100644 src/assets/platforms/zhihu.png create mode 100644 src/components/agent/chat/components/StreamingRenderer.tsx create mode 100644 src/components/agent/chat/components/TokenUsageDisplay.tsx create mode 100644 src/components/agent/chat/components/ToolCallDisplay.tsx diff --git a/package.json b/package.json index 2c933511c..33e7acca7 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.24.0", + "version": "0.25.0", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 5c147956f..3d5f6975a 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -3674,7 +3674,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.24.0" +version = "0.25.0" dependencies = [ "anyhow", "arboard", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 05cdf4bd7..f22a2b387 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "proxycast" -version = "0.24.0" +version = "0.25.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" diff --git a/src-tauri/proptest-regressions/agent/native_agent.txt b/src-tauri/proptest-regressions/agent/native_agent.txt new file mode 100644 index 000000000..f1875cabc --- /dev/null +++ b/src-tauri/proptest-regressions/agent/native_agent.txt @@ -0,0 +1,7 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc 558c6d57a7ec605e06f771be388a08ad009c754723b45232756487915045fcb9 # shrinks to tool_name = "bash", tool_id = "call_00Aaaa0A", arg_key = "aaa", arg_value = "a" diff --git a/src-tauri/proptest-regressions/agent/tools/bash.txt b/src-tauri/proptest-regressions/agent/tools/bash.txt new file mode 100644 index 000000000..e89b77e13 --- /dev/null +++ b/src-tauri/proptest-regressions/agent/tools/bash.txt @@ -0,0 +1,7 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc d1da341c69acca3c4fcf45adfc0ca70a0835f955f5478539f7c24b847ca80c45 # shrinks to content = "-" diff --git a/src-tauri/proptest-regressions/agent/tools/read_file.txt b/src-tauri/proptest-regressions/agent/tools/read_file.txt new file mode 100644 index 000000000..e082bf1cc --- /dev/null +++ b/src-tauri/proptest-regressions/agent/tools/read_file.txt @@ -0,0 +1,8 @@ +# Seeds for failure cases proptest has generated in the past. It is +# automatically read and these particular cases re-run before any +# novel cases are generated. +# +# It is recommended to check this file in to source control so that +# everyone who runs the test benefits from these saved cases. +cc 7191f893879cb1f039c1686c2b2538314a7a1382e6b9cf1403e2a522d8c47c6c # shrinks to lines = [" C iG 0c0m", " FMf", "9 2 SD1GJ2X6zWI Mn", "4 PG2u6c M4e 3c U2gNc2DwEM1Aith", "CjGFRg F9L", "0y0aKdrF", " ", "jD9CWhlaRf", "SB urH38uh2qk ubPF 4 h1", "G Io fV633I3EViAtUk ekP2 C RJ", "oBNMH ls 45O0VU", "zgpqmoHl JA C5hq JMfE AVS", "qi4 ZnR f0w kv546Hi UNV8 ui k4bx 8mPr 13m", " 5 kZT n 7y1Eo zTf 8 ", "b N6 4Hu HFwkJ15", "3Tj 9WYz0 45 GYtb3GV LVx0 xpwh RYlu8Wo", " Wft 10nYZ Q40", "FD9dF1u bmlWx mdgJ VakP7v23Q", "595taG8u9x4IB 6Hu hv jy 2 0rfDcRf UL", "ql95v A", "pzbaL t7lV t mz n8 s890x 9OIHP p", "RpMgr n6y5 Zw7pt 34aV7m7cK H uhK Z6 o x ", "6LdW7FkF 8EdNr qcCav qp 0US69SCdJdTr", " Kz3 z vd8yI58Q kH 3nf4gJlnU", "JfR35hUSE SgZ C WH Xc5Ud 7f xlesr", "s1Mfb25X7 XBix V", "7BXzRq09oe brNC oUIzWMtwJkDZ8I7", "O", "6m2ai", "O 3Ft671 rM R6laKqS4ef rN20Qr", "Y82AvuHqrm6l Nfk6j a3 J4 0 2IfheC yz62 0", "R9OxTMZU67y 3 p HZ2EUMyZx468UIVB4gn", " m qg noz", "7TuXaK H yGvdYm4i rE9 ", "TE4CSD8H1KXNX4 24 o Iw448QnR c ", "72T", "642tY3FG Pl89X6 oq6iW9Z3UoaP N2M gs8tM7 8nQ6G6", "y 7WQguyJ E8D2 CZ", "O9 Xfp", "VIVXE SNN25D7 x9mJ 3TKdzIZA5aA0", "4T79 XYGsV0wAxU 3 1UG RZM", " PdN JA0R 4zPQ7 Q CBDKXjp4gnxZ 3", "7f j0ahlBI4tn4SS", "8v2lyDgaafHSGQb2lc4Q 6L0TKP7s yqC1 8P2", " XndE 4AA4eht9bIaoAO838 yginQ2CR3 Zh ", " p q8 GvYx2c507XrKCd2U97 73", "7 1Rnek2 y 02 1", "tFR", "5 bQ5", "01yM9Uo3KrMJ 08Jqd N 1Lm2q 05 7eT", "5DnH 6i70GsUE5Gcidwjd0 05Xg3yMiJnLl4g", "5hqDP0YC w17 AG31 U XV mN36d02YBkB8GEM14 AId1", "9GoS ", "T M oR B4Qb Uv0Mk7VsD2Ei 3 ", "2RI14wxg d3 3MtXJo IJ W y ", "EEs UFknE049Y5n", " wh ZirEFtZ67qquw ", "6t x3Yqz 7 b SjqXz1w k 8xH5ycwq WPdgR 73j", "Y KUL40NPzWud", "xjnx2Ow b XEGoVdePIBRwVmv srfK OI8 4P7Fh1", "Dh5a N5D1 rFXz0hKt4t7 4fr a NDDz", "iHuv4H ezvpK P 1EDX 2MT9EPb7hYR7v4", " 76FC41f1 K9B7 ea E5a 6K6 d1", "42 BkFMSr87 uixTdNnd115sSbr4c ", "ewc1 vIf", "DCBwDi8J d0UID OP 4l", "x DtFa6 O7O6stg3", "rDo Dlzn7 6 y g", "mjzRp7QiKy", "C Y0Z U9 28C QNUUWe4g cFqK8hjTHA9GbEq5Jn5Lx2i", "BFY7KGi42X67", "4 KX eO CJ8z f2x12TC HP6gD", "lj3PfjIZOvf MA93FheQo8HgmTV OPjL51q7IE 8X 4c", "hI TJ utF FVNHdkz w9JGBtTCae n4p2mvpz8H w8", " rc 5gQ19wGT k OYpCfDv S8TRN9JlF uU", "NBS t5D9 K fWPT6e3dBjIGsL81r", "8y 29182j3F2 Ty5 mrhvVZ6uo2iQMZ Plt e 1D03L 55 ", "K c aJ999JO XG 55Q dO8KVtyqu3e3jamJ K c5 ib3 P ", " P9OM8zVt0 Uk rxc2QBADHy k7h6b0 zQ8VXWvAJ ", " YyyIOm0", " zw ynSJ v08d IPY7F3", "P 5m eXGbx", "3gwH7 b18D 5SQU W", "0GND44b Eum5t88epJ8nrpc f8eiwvJ", "W h n5iP XgC", "nsU6u7qsa6dLB6 66l lELD49 F4z0 4SeYCDHgFS 6 4vek7 ", "SQ", "Wuk0IEfOW p4Sqc d jZP6C3210i3D6b u hM 7sKE"] +cc 22df1900e88351d89c99fcbe46488f3238532ecb4de326a180fa8e9adfc33350 # shrinks to lines = [" ", "H1SiOq hfWM1 4Rn3gXwADK", "g", "ubEzO2l", "L3dIUqqayi20Wy61PXnszI V 1Itm 4a 1 Y2m7", "l 3 YiL2lWHy nY984S0eH bkcIyJj Hs OEy83A2Y7NbnIaf", "38eOI nEF qx6AvWI5ZKKu5 j64Vg9eSTpIG", "a 44ICR8 p5fWunU co0 mZO4 xQq", "UD 8iEGV v1GVR4X ", "Smu7Hwde6DxdEu5 iGq n8Q 5 i3zpaBMF1b3LB ipR", " g7q 9vug3 ma ByY 6Y ", "9D rMa21kZmZ mWvY5w560HU N2J5fIjp QG040IJ4", "0r d342PHch hxT8 P20Eeng36mI4xrc4l", " Uv3Xhc i 1gm eWw", "s8 ", "o50vB F5N cf7 3 G dy C L9l Hs7F 3tz INB", " k9 45f Po Z BI njV jt ", " QRTc7 YsqG3mjj2 C8 ZD0Bg r6 gu vL", "6A3lB5lq 3 L ZDSWt go56cLAi X3Mu qm9t", "263", "2Qv06h1OF3 U EP1Jb EXhvyZJgw QtFHI4", "C3", "y sHbK4op QQcwH2k t 4y Aap 1E49 v", " 0LqE4 G3KumS EsvU Y Mykc7AG", " lja c", "aA cu8IGU k PBr3IdaNl75", "yN4n mgNln60 bn7 0", " 53jpSbNjC BScZ 7 a DCVfxJi s4pjL WCb p4s", "9DZ 2cmDI2vn", "mut Cr 0 J9 y5 G otM3b qd 6LId1o BwC61HH", " 0O 7SiAE2q x7QYax7H", "16 3XtzIK6z 16 t3 0CO9jdnPQshTR U", " OZkB9niv 5cs ", " 42E gZ ", "Y alH6x34J917 6 7tkpeo3xYq Y f DwQ aJO ", "y4 7rW yrX23GWC bF4G OCVyV4q"] diff --git a/src-tauri/src/agent/README.md b/src-tauri/src/agent/README.md index f806eff28..f290bde5c 100644 --- a/src-tauri/src/agent/README.md +++ b/src-tauri/src/agent/README.md @@ -4,7 +4,7 @@ ## 架构说明 -AI Agent 集成模块,提供原生 Rust Agent 功能,支持**连续对话**和**工具调用**。 +AI Agent 集成模块,提供原生 Rust Agent 功能,支持**连续对话**和**工具调用循环**。 ### 设计决策 @@ -12,14 +12,18 @@ AI Agent 集成模块,提供原生 Rust Agent 功能,支持**连续对话** - **会话管理**:支持多会话,每个会话独立维护消息历史和系统提示词 - **连续对话**:每次请求携带 session_id,自动包含历史消息 - **流式响应**:通过 Tauri 事件系统向前端推送流式内容 +- **工具系统**:可扩展的工具定义和执行框架,支持 Bash、文件操作等 +- **工具调用循环**:自动执行工具调用并继续对话,直到产生最终响应 ## 文件索引 -| 文件 | 说明 | +| 文件/目录 | 说明 | |------|------| | `mod.rs` | 模块入口,导出公共类型 | | `types.rs` | Agent 相关类型定义(会话、消息、工具、配置) | | `native_agent.rs` | 原生 Rust Agent 实现(NativeAgent、NativeAgentState) | +| `tool_loop.rs` | 工具调用循环引擎(ToolLoopEngine、ToolLoopConfig) | +| `tools/` | 工具系统子模块(类型定义、注册表、具体工具实现) | ## 核心类型 @@ -31,10 +35,21 @@ AI Agent 集成模块,提供原生 Rust Agent 功能,支持**连续对话** - `MessageContent`: 消息内容(文本或多部分) - `ContentPart`: 内容部分(文本/图片) -### 工具支持(预留) +### 工具系统 +- `ToolDefinition`: 工具定义(名称、描述、参数 Schema) +- `JsonSchema`: JSON Schema 参数定义 +- `PropertySchema`: 属性 Schema(类型、描述、默认值) - `ToolCall`: 工具调用请求 -- `ToolDefinition`: 工具定义 -- `FunctionDefinition`: 函数定义 +- `ToolResult`: 工具执行结果 +- `ToolError`: 工具错误类型 +- `Tool` trait: 工具接口(definition + execute) +- `ToolRegistry`: 工具注册表(注册、查找、验证、执行) + +### 工具调用循环 +- `ToolLoopEngine`: 工具循环引擎,执行工具调用并继续对话 +- `ToolLoopConfig`: 循环配置(最大迭代次数等) +- `ToolLoopState`: 循环状态跟踪 +- `ToolCallResult`: 工具调用结果 ### Agent 实现 - `NativeAgent`: Agent 核心实现 @@ -58,6 +73,14 @@ let request = NativeChatRequest { stream: false, }; let response = agent_state.chat(request).await?; + +// 使用工具调用循环 +let registry = Arc::new(ToolRegistry::new()); +registry.register(BashTool::new(security.clone()))?; +let engine = ToolLoopEngine::new(registry); + +let (tx, rx) = mpsc::channel(100); +let result = agent_state.chat_stream_with_tools(request, tx, &engine).await?; ``` ## 更新提醒 diff --git a/src-tauri/src/agent/mod.rs b/src-tauri/src/agent/mod.rs index b777e0021..16768426f 100644 --- a/src-tauri/src/agent/mod.rs +++ b/src-tauri/src/agent/mod.rs @@ -1,9 +1,13 @@ //! AI Agent 集成模块 //! //! 提供基于 OpenAI 兼容 API 的 Agent 实现 +//! 包含工具系统、流式处理和工具调用循环 pub mod native_agent; +pub mod tool_loop; +pub mod tools; pub mod types; pub use native_agent::{NativeAgent, NativeAgentState}; +pub use tool_loop::{ToolCallResult, ToolLoopConfig, ToolLoopEngine, ToolLoopError, ToolLoopState}; pub use types::*; diff --git a/src-tauri/src/agent/native_agent.rs b/src-tauri/src/agent/native_agent.rs index 271c74037..fee201654 100644 --- a/src-tauri/src/agent/native_agent.rs +++ b/src-tauri/src/agent/native_agent.rs @@ -2,7 +2,17 @@ //! //! 支持连续对话(Conversation History)和工具调用(Tools) //! 参考 goose 项目的 Agent 设计 +//! +//! ## 流式处理 +//! - 实现 SSE 流解析,支持 text_delta 和 tool_calls 解析 +//! - Requirements: 1.1, 1.3, 1.4 +//! +//! ## 工具调用循环 +//! - 实现工具调用检测、执行和结果收集 +//! - Requirements: 7.1, 7.2, 7.3, 7.4, 7.5, 7.6 +use crate::agent::tool_loop::{ToolCallResult, ToolLoopConfig, ToolLoopEngine, ToolLoopState}; +use crate::agent::tools::ToolRegistry; use crate::agent::types::*; use crate::models::openai::{ ChatCompletionRequest, ChatCompletionResponse, ChatMessage, ContentPart as OpenAIContentPart, @@ -16,7 +26,188 @@ use std::collections::HashMap; use std::sync::Arc; use std::time::Duration; use tokio::sync::mpsc; -use tracing::{debug, error, info}; +use tracing::{debug, error, info, warn}; + +/// SSE 流解析器 +/// +/// 解析 Server-Sent Events 流,提取 text_delta 和 tool_calls +/// Requirements: 1.1, 1.3, 1.4 +#[derive(Debug, Default)] +struct SSEParser { + /// 累积的完整内容 + full_content: String, + /// 累积的工具调用 + tool_calls: Vec, + /// 当前正在构建的工具调用索引 + current_tool_indices: HashMap, +} + +/// 工具调用增量数据 +#[derive(Debug, Clone, Default)] +struct ToolCallDelta { + /// 工具调用索引 + index: usize, + /// 工具调用 ID + id: String, + /// 工具类型 + call_type: String, + /// 函数名 + function_name: String, + /// 函数参数(累积的 JSON 字符串) + function_arguments: String, +} + +impl SSEParser { + fn new() -> Self { + Self::default() + } + + /// 解析 SSE 数据行 + /// + /// 返回 (text_delta, is_done, usage) + fn parse_data(&mut self, data: &str) -> (Option, bool, Option) { + if data.trim() == "[DONE]" { + return (None, true, None); + } + + let json: Value = match serde_json::from_str(data) { + Ok(v) => v, + Err(e) => { + warn!("[SSEParser] 解析 JSON 失败: {} - data: {}", e, data); + return (None, false, None); + } + }; + + // 提取 usage 信息(如果存在) + let usage = json.get("usage").and_then(|u| { + let input = u.get("prompt_tokens").and_then(|v| v.as_u64()).unwrap_or(0) as u32; + let output = u + .get("completion_tokens") + .and_then(|v| v.as_u64()) + .unwrap_or(0) as u32; + if input > 0 || output > 0 { + Some(TokenUsage::new(input, output)) + } else { + None + } + }); + + // 检查是否有 choices + let choices = match json.get("choices").and_then(|c| c.as_array()) { + Some(c) => c, + None => return (None, false, usage), + }; + + if choices.is_empty() { + return (None, false, usage); + } + + let choice = &choices[0]; + let delta = match choice.get("delta") { + Some(d) => d, + None => return (None, false, usage), + }; + + // 检查 finish_reason + let finish_reason = choice + .get("finish_reason") + .and_then(|f| f.as_str()) + .unwrap_or(""); + let is_done = finish_reason == "stop" || finish_reason == "tool_calls"; + + // 提取文本内容 + let text_delta = delta + .get("content") + .and_then(|c| c.as_str()) + .filter(|s| !s.is_empty()) + .map(|s| { + self.full_content.push_str(s); + s.to_string() + }); + + // 提取工具调用 + if let Some(tool_calls) = delta.get("tool_calls").and_then(|tc| tc.as_array()) { + for tc in tool_calls { + self.parse_tool_call_delta(tc); + } + } + + (text_delta, is_done, usage) + } + + /// 解析工具调用增量 + fn parse_tool_call_delta(&mut self, tc: &Value) { + let index = tc.get("index").and_then(|i| i.as_u64()).unwrap_or(0) as usize; + + // 获取或创建工具调用 + let tool_call = self + .current_tool_indices + .entry(index) + .or_insert_with(|| ToolCallDelta { + index, + ..Default::default() + }); + + // 更新 ID + if let Some(id) = tc.get("id").and_then(|i| i.as_str()) { + tool_call.id = id.to_string(); + } + + // 更新类型 + if let Some(t) = tc.get("type").and_then(|t| t.as_str()) { + tool_call.call_type = t.to_string(); + } + + // 更新函数信息 + if let Some(function) = tc.get("function") { + if let Some(name) = function.get("name").and_then(|n| n.as_str()) { + tool_call.function_name = name.to_string(); + } + if let Some(args) = function.get("arguments").and_then(|a| a.as_str()) { + tool_call.function_arguments.push_str(args); + } + } + } + + /// 完成解析,返回最终的工具调用列表 + fn finalize_tool_calls(&mut self) -> Vec { + // 按索引排序并转换为 ToolCall + let mut indices: Vec<_> = self.current_tool_indices.keys().cloned().collect(); + indices.sort(); + + indices + .into_iter() + .filter_map(|idx| { + let delta = self.current_tool_indices.get(&idx)?; + if delta.id.is_empty() || delta.function_name.is_empty() { + return None; + } + Some(ToolCall { + id: delta.id.clone(), + call_type: if delta.call_type.is_empty() { + "function".to_string() + } else { + delta.call_type.clone() + }, + function: FunctionCall { + name: delta.function_name.clone(), + arguments: delta.function_arguments.clone(), + }, + }) + }) + .collect() + } + + /// 获取完整内容 + fn get_full_content(&self) -> String { + self.full_content.clone() + } + + /// 是否有工具调用 + fn has_tool_calls(&self) -> bool { + !self.current_tool_indices.is_empty() + } +} /// 原生 Agent 实现 pub struct NativeAgent { @@ -378,11 +569,14 @@ impl NativeAgent { } /// 流式聊天(支持连续对话) + /// + /// 实现 SSE 流解析,支持 text_delta 和 tool_calls 解析 + /// Requirements: 1.1, 1.3, 1.4 pub async fn chat_stream( &self, request: NativeChatRequest, tx: mpsc::Sender, - ) -> Result<(), String> { + ) -> Result { let model = request.model.unwrap_or_else(|| self.config.model.clone()); let session_id = request.session_id.clone(); @@ -442,7 +636,8 @@ impl NativeAgent { let mut stream = response.bytes_stream(); let mut buffer = String::new(); - let mut full_content = String::new(); + let mut parser = SSEParser::new(); + let mut final_usage: Option = None; while let Some(chunk) = stream.next().await { match chunk { @@ -450,13 +645,35 @@ impl NativeAgent { let text = String::from_utf8_lossy(&bytes); buffer.push_str(&text); + // 处理完整的 SSE 事件(以 \n\n 分隔) while let Some(pos) = buffer.find("\n\n") { let event = buffer[..pos].to_string(); buffer = buffer[pos + 2..].to_string(); for line in event.lines() { if let Some(data) = line.strip_prefix("data: ") { - if data.trim() == "[DONE]" { + let (text_delta, is_done, usage) = parser.parse_data(data); + + // 更新 usage + if usage.is_some() { + final_usage = usage; + } + + // 发送文本增量 + if let Some(text) = text_delta { + let _ = tx.send(StreamEvent::TextDelta { text }).await; + } + + // 检查是否完成 + if is_done { + // 获取最终结果 + let full_content = parser.get_full_content(); + let tool_calls = if parser.has_tool_calls() { + Some(parser.finalize_tool_calls()) + } else { + None + }; + // 更新会话历史 if let Some(sid) = &session_id { self.add_message_to_session( @@ -465,34 +682,25 @@ impl NativeAgent { MessageContent::Text(request.message.clone()), request.images.as_deref(), ); - self.add_message_to_session( + self.add_assistant_message_to_session( sid, - "assistant", MessageContent::Text(full_content.clone()), - None, + tool_calls.clone(), ); } - let _ = tx.send(StreamEvent::Done { usage: None }).await; - return Ok(()); - } - if let Ok(json) = serde_json::from_str::(data) { - if let Some(delta) = json - .get("choices") - .and_then(|c| c.get(0)) - .and_then(|c| c.get("delta")) - .and_then(|d| d.get("content")) - .and_then(|c| c.as_str()) - { - if !delta.is_empty() { - full_content.push_str(delta); - let _ = tx - .send(StreamEvent::TextDelta { - text: delta.to_string(), - }) - .await; - } - } + // 发送完成事件 + let _ = tx + .send(StreamEvent::Done { + usage: final_usage.clone(), + }) + .await; + + return Ok(StreamResult { + content: full_content, + tool_calls, + usage: final_usage, + }); } } } @@ -510,6 +718,14 @@ impl NativeAgent { } } + // 流正常结束但没有收到 [DONE] + let full_content = parser.get_full_content(); + let tool_calls = if parser.has_tool_calls() { + Some(parser.finalize_tool_calls()) + } else { + None + }; + // 更新会话历史 if let Some(sid) = &session_id { self.add_message_to_session( @@ -518,11 +734,323 @@ impl NativeAgent { MessageContent::Text(request.message.clone()), request.images.as_deref(), ); - self.add_message_to_session(sid, "assistant", MessageContent::Text(full_content), None); + self.add_assistant_message_to_session( + sid, + MessageContent::Text(full_content.clone()), + tool_calls.clone(), + ); } - let _ = tx.send(StreamEvent::Done { usage: None }).await; - Ok(()) + let _ = tx + .send(StreamEvent::Done { + usage: final_usage.clone(), + }) + .await; + + Ok(StreamResult { + content: full_content, + tool_calls, + usage: final_usage, + }) + } + + /// 添加 assistant 消息到会话(支持工具调用) + fn add_assistant_message_to_session( + &self, + session_id: &str, + content: MessageContent, + tool_calls: Option>, + ) { + let mut sessions = self.sessions.write(); + if let Some(session) = sessions.get_mut(session_id) { + session.messages.push(AgentMessage { + role: "assistant".to_string(), + content, + timestamp: chrono::Utc::now().to_rfc3339(), + tool_calls, + tool_call_id: None, + }); + session.updated_at = chrono::Utc::now().to_rfc3339(); + } + } + + /// 添加工具结果消息到会话 + /// + /// Requirements: 7.2 - THE Tool_Loop SHALL send tool results back to the Agent as tool role messages + fn add_tool_result_to_session(&self, session_id: &str, tool_result: &ToolCallResult) { + let mut sessions = self.sessions.write(); + if let Some(session) = sessions.get_mut(session_id) { + session.messages.push(tool_result.to_agent_message()); + session.updated_at = chrono::Utc::now().to_rfc3339(); + } + } + + /// 流式聊天(支持工具调用循环) + /// + /// 实现完整的工具调用循环: + /// 1. 发送请求到 LLM + /// 2. 如果响应包含工具调用,执行工具 + /// 3. 将工具结果发送回 LLM + /// 4. 重复直到 LLM 产生最终响应或达到最大迭代次数 + /// + /// Requirements: 7.1, 7.2, 7.3, 7.4, 7.5, 7.6 + pub async fn chat_stream_with_tools( + &self, + request: NativeChatRequest, + tx: mpsc::Sender, + tool_loop_engine: &ToolLoopEngine, + ) -> Result { + let session_id = request.session_id.clone(); + let mut state = ToolLoopState::new(); + + // 首次请求 + let mut current_result = self.chat_stream(request.clone(), tx.clone()).await?; + + // 工具调用循环 + // Requirements: 7.3 - THE Tool_Loop SHALL continue until the Agent produces a final response without tool_calls + while tool_loop_engine.should_continue(¤t_result, state.iteration) { + state.increment_iteration(); + + let tool_calls = current_result.tool_calls.as_ref().unwrap(); + state.add_tool_calls(tool_calls.len()); + + info!( + "[NativeAgent] 工具循环迭代 {}: 执行 {} 个工具调用", + state.iteration, + tool_calls.len() + ); + + // 执行所有工具调用 + // Requirements: 7.1 - THE Tool_Loop SHALL execute each tool and collect results + // Requirements: 7.6 - WHILE the Tool_Loop is executing, THE Frontend SHALL display the current tool + let tool_results = tool_loop_engine + .execute_all_tool_calls(tool_calls, Some(&tx)) + .await; + + // 将工具结果添加到会话 + // Requirements: 7.2 - THE Tool_Loop SHALL send tool results back to the Agent as tool role messages + if let Some(sid) = &session_id { + for result in &tool_results { + self.add_tool_result_to_session(sid, result); + } + } + + // 构建继续对话的请求 + let continue_request = NativeChatRequest { + session_id: session_id.clone(), + message: String::new(), // 空消息,因为我们使用会话历史 + model: request.model.clone(), + images: None, + stream: true, + }; + + // 继续对话 + current_result = self + .chat_stream_continue(continue_request, tx.clone()) + .await?; + } + + // 检查是否因为达到最大迭代次数而停止 + // Requirements: 7.5 - THE Tool_Loop SHALL enforce a maximum iteration limit + if state.iteration >= tool_loop_engine.max_iterations() && current_result.has_tool_calls() { + warn!( + "[NativeAgent] 达到最大迭代次数 {},强制停止工具循环", + tool_loop_engine.max_iterations() + ); + let _ = tx + .send(StreamEvent::Error { + message: format!( + "达到最大工具调用迭代次数限制 ({})", + tool_loop_engine.max_iterations() + ), + }) + .await; + } + + state.mark_completed(current_result.content.clone()); + + info!( + "[NativeAgent] 工具循环完成: {} 次迭代, {} 个工具调用", + state.iteration, state.total_tool_calls + ); + + Ok(current_result) + } + + /// 继续流式对话(使用会话历史) + /// + /// 用于工具调用循环中继续对话 + async fn chat_stream_continue( + &self, + request: NativeChatRequest, + tx: mpsc::Sender, + ) -> Result { + let model = request.model.unwrap_or_else(|| self.config.model.clone()); + let session_id = request.session_id.as_ref().ok_or("需要 session_id")?; + + debug!( + "[NativeAgent] 继续流式对话: model={}, session={}", + model, session_id + ); + + // 获取会话 + let session = self + .sessions + .read() + .get(session_id) + .cloned() + .ok_or_else(|| format!("会话不存在: {}", session_id))?; + + // 构建消息(使用会话历史,不添加新的用户消息) + let messages = self.build_messages_from_session(&session); + + let chat_request = ChatCompletionRequest { + model: model.clone(), + messages, + stream: true, + temperature: self.config.temperature, + max_tokens: self.config.max_tokens, + top_p: None, + tools: None, // TODO: 添加工具定义 + tool_choice: None, + reasoning_effort: None, + }; + + let url = format!("{}/v1/chat/completions", self.base_url); + + let response = self + .client + .post(&url) + .header("Authorization", format!("Bearer {}", self.api_key)) + .header("Content-Type", "application/json") + .json(&chat_request) + .send() + .await + .map_err(|e| format!("请求失败: {}", e))?; + + let status = response.status(); + if !status.is_success() { + let body = response.text().await.unwrap_or_default(); + error!("[NativeAgent] 流式请求失败: {} - {}", status, body); + let _ = tx + .send(StreamEvent::Error { + message: format!("API 错误 ({}): {}", status, body), + }) + .await; + return Err(format!("API 错误: {}", status)); + } + + let mut stream = response.bytes_stream(); + let mut buffer = String::new(); + let mut parser = SSEParser::new(); + let mut final_usage: Option = None; + + while let Some(chunk) = stream.next().await { + match chunk { + Ok(bytes) => { + let text = String::from_utf8_lossy(&bytes); + buffer.push_str(&text); + + while let Some(pos) = buffer.find("\n\n") { + let event = buffer[..pos].to_string(); + buffer = buffer[pos + 2..].to_string(); + + for line in event.lines() { + if let Some(data) = line.strip_prefix("data: ") { + let (text_delta, is_done, usage) = parser.parse_data(data); + + if usage.is_some() { + final_usage = usage; + } + + if let Some(text) = text_delta { + let _ = tx.send(StreamEvent::TextDelta { text }).await; + } + + if is_done { + let full_content = parser.get_full_content(); + let tool_calls = if parser.has_tool_calls() { + Some(parser.finalize_tool_calls()) + } else { + None + }; + + // 更新会话历史 + self.add_assistant_message_to_session( + session_id, + MessageContent::Text(full_content.clone()), + tool_calls.clone(), + ); + + // 不发送 Done 事件,因为工具循环可能还会继续 + + return Ok(StreamResult { + content: full_content, + tool_calls, + usage: final_usage, + }); + } + } + } + } + } + Err(e) => { + error!("[NativeAgent] 流读取错误: {}", e); + let _ = tx + .send(StreamEvent::Error { + message: format!("流读取错误: {}", e), + }) + .await; + return Err(format!("流读取错误: {}", e)); + } + } + } + + // 流正常结束 + let full_content = parser.get_full_content(); + let tool_calls = if parser.has_tool_calls() { + Some(parser.finalize_tool_calls()) + } else { + None + }; + + self.add_assistant_message_to_session( + session_id, + MessageContent::Text(full_content.clone()), + tool_calls.clone(), + ); + + Ok(StreamResult { + content: full_content, + tool_calls, + usage: final_usage, + }) + } + + /// 从会话构建消息列表(不添加新的用户消息) + fn build_messages_from_session(&self, session: &AgentSession) -> Vec { + let mut messages = Vec::new(); + + // 添加系统提示词 + let system_prompt = session + .system_prompt + .as_ref() + .or(self.config.system_prompt.as_ref()); + if let Some(prompt) = system_prompt { + messages.push(ChatMessage { + role: "system".to_string(), + content: Some(OpenAIMessageContent::Text(prompt.clone())), + tool_calls: None, + tool_call_id: None, + }); + } + + // 添加所有历史消息 + for msg in &session.messages { + messages.push(self.convert_to_chat_message(msg)); + } + + messages } pub fn create_session(&self, model: Option, system_prompt: Option) -> String { @@ -634,7 +1162,7 @@ impl NativeAgentState { &self, request: NativeChatRequest, tx: mpsc::Sender, - ) -> Result<(), String> { + ) -> Result { let (base_url, api_key, config, sessions) = { let guard = self.agent.read(); let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?; @@ -662,6 +1190,44 @@ impl NativeAgentState { temp_agent.chat_stream(request, tx).await } + /// 流式聊天(支持工具调用循环) + /// + /// Requirements: 7.1, 7.2, 7.3, 7.4, 7.5, 7.6 + pub async fn chat_stream_with_tools( + &self, + request: NativeChatRequest, + tx: mpsc::Sender, + tool_loop_engine: &ToolLoopEngine, + ) -> Result { + let (base_url, api_key, config, sessions) = { + let guard = self.agent.read(); + let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?; + ( + agent.base_url.clone(), + agent.api_key.clone(), + agent.config.clone(), + agent.sessions.clone(), + ) + }; + + let temp_agent = NativeAgent { + client: Client::builder() + .timeout(Duration::from_secs(300)) + .connect_timeout(Duration::from_secs(30)) + .no_proxy() + .build() + .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?, + base_url, + api_key, + sessions, + config, + }; + + temp_agent + .chat_stream_with_tools(request, tx, tool_loop_engine) + .await + } + pub fn create_session( &self, model: Option, @@ -712,3 +1278,308 @@ impl NativeAgentState { .and_then(|a| a.get_session_messages(session_id)) } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_sse_parser_text_delta() { + let mut parser = SSEParser::new(); + + // 模拟 SSE 数据 + let data1 = r#"{"choices":[{"delta":{"content":"Hello"}}]}"#; + let data2 = r#"{"choices":[{"delta":{"content":" World"}}]}"#; + let data3 = r#"{"choices":[{"delta":{},"finish_reason":"stop"}]}"#; + + let (text1, done1, _) = parser.parse_data(data1); + assert_eq!(text1, Some("Hello".to_string())); + assert!(!done1); + + let (text2, done2, _) = parser.parse_data(data2); + assert_eq!(text2, Some(" World".to_string())); + assert!(!done2); + + let (text3, done3, _) = parser.parse_data(data3); + assert!(text3.is_none()); + assert!(done3); + + assert_eq!(parser.get_full_content(), "Hello World"); + } + + #[test] + fn test_sse_parser_tool_calls() { + let mut parser = SSEParser::new(); + + // 模拟工具调用的 SSE 数据 + let data1 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_123","type":"function","function":{"name":"bash"}}]}}]}"#; + let data2 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"command\":"}}]}}]}"#; + let data3 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"ls -la\"}"}}]}}]}"#; + let data4 = r#"{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}"#; + + parser.parse_data(data1); + parser.parse_data(data2); + parser.parse_data(data3); + let (_, done, _) = parser.parse_data(data4); + + assert!(done); + assert!(parser.has_tool_calls()); + + let tool_calls = parser.finalize_tool_calls(); + assert_eq!(tool_calls.len(), 1); + assert_eq!(tool_calls[0].id, "call_123"); + assert_eq!(tool_calls[0].function.name, "bash"); + assert_eq!(tool_calls[0].function.arguments, r#"{"command":"ls -la"}"#); + } + + #[test] + fn test_sse_parser_usage() { + let mut parser = SSEParser::new(); + + let data = r#"{"choices":[{"delta":{"content":"Hi"}}],"usage":{"prompt_tokens":10,"completion_tokens":5}}"#; + let (text, _, usage) = parser.parse_data(data); + + assert_eq!(text, Some("Hi".to_string())); + assert!(usage.is_some()); + let usage = usage.unwrap(); + assert_eq!(usage.input_tokens, 10); + assert_eq!(usage.output_tokens, 5); + } + + #[test] + fn test_sse_parser_done_signal() { + let mut parser = SSEParser::new(); + + let (_, done, _) = parser.parse_data("[DONE]"); + assert!(done); + } + + #[test] + fn test_sse_parser_invalid_json() { + let mut parser = SSEParser::new(); + + let (text, done, usage) = parser.parse_data("invalid json"); + assert!(text.is_none()); + assert!(!done); + assert!(usage.is_none()); + } + + #[test] + fn test_sse_parser_multiple_tool_calls() { + let mut parser = SSEParser::new(); + + // 两个工具调用 + let data1 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"bash","arguments":"{}"}}]}}]}"#; + let data2 = r#"{"choices":[{"delta":{"tool_calls":[{"index":1,"id":"call_2","type":"function","function":{"name":"read_file","arguments":"{}"}}]}}]}"#; + + parser.parse_data(data1); + parser.parse_data(data2); + + let tool_calls = parser.finalize_tool_calls(); + assert_eq!(tool_calls.len(), 2); + assert_eq!(tool_calls[0].function.name, "bash"); + assert_eq!(tool_calls[1].function.name, "read_file"); + } +} + +#[cfg(test)] +mod proptests { + use super::*; + use proptest::prelude::*; + + /// 生成有效的文本内容(不包含特殊字符) + fn arb_text_content() -> impl Strategy { + "[a-zA-Z0-9 ,.!?]{1,100}".prop_map(|s| s) + } + + /// 生成文本片段列表 + fn arb_text_chunks() -> impl Strategy> { + prop::collection::vec(arb_text_content(), 1..10) + } + + /// 生成有效的工具名称 + fn arb_tool_name() -> impl Strategy { + prop_oneof![ + Just("bash".to_string()), + Just("read_file".to_string()), + Just("write_file".to_string()), + Just("edit_file".to_string()), + ] + } + + /// 生成有效的工具调用 ID + fn arb_tool_id() -> impl Strategy { + "call_[a-zA-Z0-9]{8}".prop_map(|s| s) + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: agent-tool-calling, Property 1: 流式事件完整性** + /// **Validates: Requirements 1.1, 1.3** + /// + /// *For any* Agent 响应流,流式处理器发送的所有 text_delta 事件的文本拼接后, + /// 应该等于最终的完整响应内容。 + #[test] + fn prop_streaming_text_completeness(chunks in arb_text_chunks()) { + let mut parser = SSEParser::new(); + let mut collected_deltas = String::new(); + + // 模拟流式处理 + for chunk in &chunks { + // 转义 JSON 特殊字符 + let escaped = chunk.replace('\\', "\\\\").replace('"', "\\\""); + let data = format!(r#"{{"choices":[{{"delta":{{"content":"{}"}}}}]}}"#, escaped); + let (text_delta, _, _) = parser.parse_data(&data); + + if let Some(text) = text_delta { + collected_deltas.push_str(&text); + } + } + + // 验证:收集的 text_delta 拼接后等于 parser 的完整内容 + prop_assert_eq!( + collected_deltas, + parser.get_full_content(), + "收集的 text_delta 应该等于完整内容" + ); + + // 验证:完整内容等于原始 chunks 拼接 + let expected = chunks.join(""); + prop_assert_eq!( + parser.get_full_content(), + expected, + "完整内容应该等于原始 chunks 拼接" + ); + } + + /// **Feature: agent-tool-calling, Property 1: 流式事件完整性 - 工具调用** + /// **Validates: Requirements 1.1, 1.3** + /// + /// *For any* 包含工具调用的响应流,工具调用信息应该被正确解析和累积。 + #[test] + fn prop_streaming_tool_calls_completeness( + tool_name in arb_tool_name(), + tool_id in arb_tool_id(), + arg_key in "[a-z]{3,10}", + arg_value in "[a-zA-Z0-9]{1,20}" + ) { + let mut parser = SSEParser::new(); + + // 使用 serde_json 构建正确的 JSON,避免手动转义问题 + let args_json = serde_json::json!({arg_key.clone(): arg_value.clone()}).to_string(); + + // 第一个 chunk: 工具 ID 和名称 + let data1 = serde_json::json!({ + "choices": [{ + "delta": { + "tool_calls": [{ + "index": 0, + "id": tool_id.clone(), + "type": "function", + "function": { + "name": tool_name.clone() + } + }] + } + }] + }).to_string(); + + // 第二个 chunk: 参数的前半部分 + let args_first_half = &args_json[..args_json.len()/2]; + let data2 = serde_json::json!({ + "choices": [{ + "delta": { + "tool_calls": [{ + "index": 0, + "function": { + "arguments": args_first_half + } + }] + } + }] + }).to_string(); + + // 第三个 chunk: 参数的后半部分 + let args_second_half = &args_json[args_json.len()/2..]; + let data3 = serde_json::json!({ + "choices": [{ + "delta": { + "tool_calls": [{ + "index": 0, + "function": { + "arguments": args_second_half + } + }] + } + }] + }).to_string(); + + parser.parse_data(&data1); + parser.parse_data(&data2); + parser.parse_data(&data3); + + prop_assert!(parser.has_tool_calls(), "应该检测到工具调用"); + + let tool_calls = parser.finalize_tool_calls(); + prop_assert_eq!(tool_calls.len(), 1, "应该有一个工具调用"); + prop_assert_eq!(&tool_calls[0].id, &tool_id, "工具调用 ID 应该匹配"); + prop_assert_eq!(&tool_calls[0].function.name, &tool_name, "工具名称应该匹配"); + + // 验证参数被正确累积 + prop_assert_eq!( + &tool_calls[0].function.arguments, + &args_json, + "工具参数应该被正确累积" + ); + } + + /// **Feature: agent-tool-calling, Property 1: 流式事件完整性 - Done 事件** + /// **Validates: Requirements 1.1, 1.3** + /// + /// *For any* 完成的响应流,应该正确识别 finish_reason。 + #[test] + fn prop_streaming_done_detection( + finish_reason in prop_oneof![Just("stop"), Just("tool_calls"), Just("length")] + ) { + let mut parser = SSEParser::new(); + + let data = format!( + r#"{{"choices":[{{"delta":{{}},"finish_reason":"{}"}}]}}"#, + finish_reason + ); + let (_, is_done, _) = parser.parse_data(&data); + + if finish_reason == "stop" || finish_reason == "tool_calls" { + prop_assert!(is_done, "finish_reason={} 应该标记为完成", finish_reason); + } else { + prop_assert!(!is_done, "finish_reason={} 不应该标记为完成", finish_reason); + } + } + + /// **Feature: agent-tool-calling, Property 1: 流式事件完整性 - Usage 统计** + /// **Validates: Requirements 1.3** + /// + /// *For any* 包含 usage 的响应,应该正确解析 token 使用量。 + #[test] + fn prop_streaming_usage_parsing( + input_tokens in 0u32..10000, + output_tokens in 0u32..10000 + ) { + let mut parser = SSEParser::new(); + + let data = format!( + r#"{{"choices":[{{"delta":{{"content":"test"}}}}],"usage":{{"prompt_tokens":{},"completion_tokens":{}}}}}"#, + input_tokens, output_tokens + ); + let (_, _, usage) = parser.parse_data(&data); + + if input_tokens > 0 || output_tokens > 0 { + prop_assert!(usage.is_some(), "应该解析出 usage"); + let usage = usage.unwrap(); + prop_assert_eq!(usage.input_tokens, input_tokens, "input_tokens 应该匹配"); + prop_assert_eq!(usage.output_tokens, output_tokens, "output_tokens 应该匹配"); + } + } + } +} diff --git a/src-tauri/src/agent/tool_loop.rs b/src-tauri/src/agent/tool_loop.rs new file mode 100644 index 000000000..597f70e03 --- /dev/null +++ b/src-tauri/src/agent/tool_loop.rs @@ -0,0 +1,1124 @@ +//! 工具调用循环引擎 +//! +//! 实现 Agent 工具调用循环,自动执行工具并继续对话 +//! 符合 Requirements 7.1, 7.2, 7.3, 7.4, 7.5 +//! +//! ## 功能 +//! - 检测 Agent 响应中的工具调用 +//! - 执行工具并收集结果 +//! - 将工具结果发送回 Agent 继续对话 +//! - 最大迭代限制防止无限循环 + +use crate::agent::tools::{ToolError, ToolRegistry, ToolResult as ToolsResult}; +use crate::agent::types::{ + AgentMessage, MessageContent, StreamEvent, StreamResult, ToolCall, ToolExecutionResult, +}; +use std::sync::Arc; +use thiserror::Error; +use tokio::sync::mpsc; +use tracing::{debug, warn}; + +/// 工具循环错误类型 +#[derive(Debug, Error)] +pub enum ToolLoopError { + /// 超过最大迭代次数 + /// Requirements: 7.5 - THE Tool_Loop SHALL enforce a maximum iteration limit + #[error("超过最大迭代次数限制: {0}")] + MaxIterationsExceeded(usize), + + /// 工具执行错误 + /// Requirements: 7.4 - IF a tool execution fails, THEN THE Tool_Loop SHALL include the error + #[error("工具执行错误: {0}")] + ToolExecution(String), + + /// 工具未找到 + #[error("工具未找到: {0}")] + ToolNotFound(String), + + /// JSON 解析错误 + #[error("JSON 解析错误: {0}")] + JsonParse(String), + + /// 通道发送错误 + #[error("事件发送失败")] + ChannelSend, +} + +/// 工具执行结果(内部使用) +#[derive(Debug, Clone)] +pub struct ToolCallResult { + /// 工具调用 ID + pub tool_call_id: String, + /// 工具名称 + pub tool_name: String, + /// 执行结果 + pub result: ToolsResult, +} + +impl ToolCallResult { + /// 创建新的工具调用结果 + pub fn new(tool_call_id: String, tool_name: String, result: ToolsResult) -> Self { + Self { + tool_call_id, + tool_name, + result, + } + } + + /// 转换为 AgentMessage(tool 角色) + /// + /// Requirements: 7.2 - THE Tool_Loop SHALL send tool results back to the Agent as tool role messages + pub fn to_agent_message(&self) -> AgentMessage { + let content = if self.result.success { + self.result.output.clone() + } else { + format!( + "Error: {}", + self.result.error.as_deref().unwrap_or("Unknown error") + ) + }; + + AgentMessage { + role: "tool".to_string(), + content: MessageContent::Text(content), + timestamp: chrono::Utc::now().to_rfc3339(), + tool_calls: None, + tool_call_id: Some(self.tool_call_id.clone()), + } + } + + /// 转换为 ToolExecutionResult(用于前端显示) + pub fn to_execution_result(&self) -> ToolExecutionResult { + ToolExecutionResult { + success: self.result.success, + output: self.result.output.clone(), + error: self.result.error.clone(), + } + } +} + +/// 工具循环引擎配置 +#[derive(Debug, Clone)] +pub struct ToolLoopConfig { + /// 最大迭代次数 + /// Requirements: 7.5 - THE Tool_Loop SHALL enforce a maximum iteration limit + pub max_iterations: usize, +} + +impl Default for ToolLoopConfig { + fn default() -> Self { + Self { + max_iterations: 25, // 默认最大 25 次迭代 + } + } +} + +impl ToolLoopConfig { + /// 创建新的配置 + pub fn new(max_iterations: usize) -> Self { + Self { max_iterations } + } +} + +/// 工具循环引擎 +/// +/// 负责执行工具调用循环,直到 Agent 产生最终响应 +/// Requirements: 7.1, 7.2, 7.3, 7.4, 7.5 +pub struct ToolLoopEngine { + /// 工具注册表 + registry: Arc, + /// 配置 + config: ToolLoopConfig, +} + +impl ToolLoopEngine { + /// 创建新的工具循环引擎 + pub fn new(registry: Arc) -> Self { + Self { + registry, + config: ToolLoopConfig::default(), + } + } + + /// 使用自定义配置创建 + pub fn with_config(registry: Arc, config: ToolLoopConfig) -> Self { + Self { registry, config } + } + + /// 获取最大迭代次数 + pub fn max_iterations(&self) -> usize { + self.config.max_iterations + } + + /// 检查响应是否包含工具调用 + /// + /// Requirements: 7.1 - WHEN the Agent response contains tool_calls + pub fn has_tool_calls(result: &StreamResult) -> bool { + result.has_tool_calls() + } + + /// 执行单个工具调用 + /// + /// Requirements: 7.1 - THE Tool_Loop SHALL execute each tool and collect results + /// Requirements: 7.4 - IF a tool execution fails, THEN THE Tool_Loop SHALL include the error + pub async fn execute_tool_call(&self, tool_call: &ToolCall) -> ToolCallResult { + let tool_name = &tool_call.function.name; + let tool_id = &tool_call.id; + + debug!("[ToolLoopEngine] 执行工具: {} (id={})", tool_name, tool_id); + + // 解析参数 + let args = match serde_json::from_str::(&tool_call.function.arguments) { + Ok(args) => args, + Err(e) => { + warn!("[ToolLoopEngine] 工具参数解析失败: {} - {}", tool_name, e); + return ToolCallResult::new( + tool_id.clone(), + tool_name.clone(), + ToolsResult::failure(format!("参数解析失败: {}", e)), + ); + } + }; + + // 执行工具 + match self.registry.execute(tool_name, args).await { + Ok(result) => { + debug!( + "[ToolLoopEngine] 工具执行成功: {} success={}", + tool_name, result.success + ); + ToolCallResult::new(tool_id.clone(), tool_name.clone(), result) + } + Err(e) => { + warn!("[ToolLoopEngine] 工具执行失败: {} - {}", tool_name, e); + let error_msg = match &e { + ToolError::NotFound(name) => format!("工具不存在: {}", name), + ToolError::InvalidArguments(msg) => format!("参数无效: {}", msg), + ToolError::ExecutionFailed(msg) => format!("执行失败: {}", msg), + ToolError::Security(msg) => format!("安全错误: {}", msg), + ToolError::Timeout => "执行超时".to_string(), + ToolError::Io(e) => format!("IO 错误: {}", e), + ToolError::Json(e) => format!("JSON 错误: {}", e), + }; + ToolCallResult::new( + tool_id.clone(), + tool_name.clone(), + ToolsResult::failure(error_msg), + ) + } + } + } + + /// 执行所有工具调用 + /// + /// Requirements: 7.1 - THE Tool_Loop SHALL execute each tool and collect results + /// Requirements: 7.6 - WHILE the Tool_Loop is executing, THE Frontend SHALL display the current tool + pub async fn execute_all_tool_calls( + &self, + tool_calls: &[ToolCall], + event_tx: Option<&mpsc::Sender>, + ) -> Vec { + let mut results = Vec::with_capacity(tool_calls.len()); + + for tool_call in tool_calls { + // 发送工具开始事件 + if let Some(tx) = event_tx { + let _ = tx + .send(StreamEvent::ToolStart { + tool_name: tool_call.function.name.clone(), + tool_id: tool_call.id.clone(), + }) + .await; + } + + // 执行工具 + let result = self.execute_tool_call(tool_call).await; + + // 发送工具结束事件 + if let Some(tx) = event_tx { + let _ = tx + .send(StreamEvent::ToolEnd { + tool_id: tool_call.id.clone(), + result: result.to_execution_result(), + }) + .await; + } + + results.push(result); + } + + results + } + + /// 将工具结果转换为 Agent 消息列表 + /// + /// Requirements: 7.2 - THE Tool_Loop SHALL send tool results back to the Agent as tool role messages + pub fn results_to_messages(results: &[ToolCallResult]) -> Vec { + results.iter().map(|r| r.to_agent_message()).collect() + } + + /// 检查是否应该继续循环 + /// + /// Requirements: 7.3 - THE Tool_Loop SHALL continue until the Agent produces a final response without tool_calls + /// Requirements: 7.5 - THE Tool_Loop SHALL enforce a maximum iteration limit + pub fn should_continue(&self, result: &StreamResult, iteration: usize) -> bool { + // 检查最大迭代次数 + if iteration >= self.config.max_iterations { + warn!( + "[ToolLoopEngine] 达到最大迭代次数: {}", + self.config.max_iterations + ); + return false; + } + + // 检查是否有工具调用 + Self::has_tool_calls(result) + } + + /// 创建 assistant 消息(包含工具调用) + pub fn create_assistant_message( + content: &str, + tool_calls: Option>, + ) -> AgentMessage { + AgentMessage { + role: "assistant".to_string(), + content: MessageContent::Text(content.to_string()), + timestamp: chrono::Utc::now().to_rfc3339(), + tool_calls: tool_calls.map(|calls| { + calls + .into_iter() + .map(|tc| crate::agent::types::ToolCall { + id: tc.id, + call_type: tc.call_type, + function: tc.function, + }) + .collect() + }), + tool_call_id: None, + } + } +} + +/// 工具循环状态 +/// +/// 用于跟踪工具循环的执行状态 +#[derive(Debug, Clone)] +pub struct ToolLoopState { + /// 当前迭代次数 + pub iteration: usize, + /// 累计执行的工具调用数 + pub total_tool_calls: usize, + /// 是否已完成 + pub completed: bool, + /// 最终内容 + pub final_content: Option, +} + +impl Default for ToolLoopState { + fn default() -> Self { + Self { + iteration: 0, + total_tool_calls: 0, + completed: false, + final_content: None, + } + } +} + +impl ToolLoopState { + /// 创建新的状态 + pub fn new() -> Self { + Self::default() + } + + /// 增加迭代次数 + pub fn increment_iteration(&mut self) { + self.iteration += 1; + } + + /// 增加工具调用计数 + pub fn add_tool_calls(&mut self, count: usize) { + self.total_tool_calls += count; + } + + /// 标记为完成 + pub fn mark_completed(&mut self, content: String) { + self.completed = true; + self.final_content = Some(content); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::agent::tools::types::{JsonSchema, PropertySchema, ToolDefinition}; + use crate::agent::tools::{Tool, ToolRegistry}; + use crate::agent::types::FunctionCall; + use async_trait::async_trait; + + /// 测试用的 Echo 工具 + struct EchoTool; + + #[async_trait] + impl Tool for EchoTool { + fn definition(&self) -> ToolDefinition { + ToolDefinition::new("echo", "Echo the input message").with_parameters( + JsonSchema::new().add_property( + "message", + PropertySchema::string("The message to echo"), + true, + ), + ) + } + + async fn execute(&self, args: serde_json::Value) -> Result { + let message = args + .get("message") + .and_then(|v| v.as_str()) + .ok_or_else(|| ToolError::InvalidArguments("缺少 message 参数".to_string()))?; + + Ok(ToolsResult::success(message)) + } + } + + /// 测试用的失败工具 + struct FailingTool; + + #[async_trait] + impl Tool for FailingTool { + fn definition(&self) -> ToolDefinition { + ToolDefinition::new("failing", "A tool that always fails") + } + + async fn execute(&self, _args: serde_json::Value) -> Result { + Err(ToolError::ExecutionFailed("故意失败".to_string())) + } + } + + fn create_test_registry() -> Arc { + let registry = ToolRegistry::new(); + registry.register(EchoTool).unwrap(); + registry.register(FailingTool).unwrap(); + Arc::new(registry) + } + + fn create_tool_call(id: &str, name: &str, args: &str) -> ToolCall { + ToolCall { + id: id.to_string(), + call_type: "function".to_string(), + function: FunctionCall { + name: name.to_string(), + arguments: args.to_string(), + }, + } + } + + #[test] + fn test_tool_loop_engine_creation() { + let registry = create_test_registry(); + let engine = ToolLoopEngine::new(registry); + + assert_eq!(engine.max_iterations(), 25); + } + + #[test] + fn test_tool_loop_engine_with_config() { + let registry = create_test_registry(); + let config = ToolLoopConfig::new(10); + let engine = ToolLoopEngine::with_config(registry, config); + + assert_eq!(engine.max_iterations(), 10); + } + + #[test] + fn test_has_tool_calls() { + // 无工具调用 + let result_no_tools = StreamResult::new("Hello".to_string()); + assert!(!ToolLoopEngine::has_tool_calls(&result_no_tools)); + + // 有工具调用 + let result_with_tools = StreamResult::new("".to_string()).with_tool_calls(vec![ToolCall { + id: "call_1".to_string(), + call_type: "function".to_string(), + function: FunctionCall { + name: "echo".to_string(), + arguments: "{}".to_string(), + }, + }]); + assert!(ToolLoopEngine::has_tool_calls(&result_with_tools)); + + // 空工具调用列表 + let result_empty_tools = StreamResult { + content: "".to_string(), + tool_calls: Some(vec![]), + usage: None, + }; + assert!(!ToolLoopEngine::has_tool_calls(&result_empty_tools)); + } + + #[tokio::test] + async fn test_execute_tool_call_success() { + let registry = create_test_registry(); + let engine = ToolLoopEngine::new(registry); + + let tool_call = create_tool_call("call_1", "echo", r#"{"message": "Hello, World!"}"#); + let result = engine.execute_tool_call(&tool_call).await; + + assert_eq!(result.tool_call_id, "call_1"); + assert_eq!(result.tool_name, "echo"); + assert!(result.result.success); + assert_eq!(result.result.output, "Hello, World!"); + } + + #[tokio::test] + async fn test_execute_tool_call_failure() { + let registry = create_test_registry(); + let engine = ToolLoopEngine::new(registry); + + let tool_call = create_tool_call("call_2", "failing", "{}"); + let result = engine.execute_tool_call(&tool_call).await; + + assert_eq!(result.tool_call_id, "call_2"); + assert_eq!(result.tool_name, "failing"); + assert!(!result.result.success); + assert!(result.result.error.is_some()); + } + + #[tokio::test] + async fn test_execute_tool_call_not_found() { + let registry = create_test_registry(); + let engine = ToolLoopEngine::new(registry); + + let tool_call = create_tool_call("call_3", "nonexistent", "{}"); + let result = engine.execute_tool_call(&tool_call).await; + + assert!(!result.result.success); + assert!(result.result.error.as_ref().unwrap().contains("工具不存在")); + } + + #[tokio::test] + async fn test_execute_tool_call_invalid_args() { + let registry = create_test_registry(); + let engine = ToolLoopEngine::new(registry); + + let tool_call = create_tool_call("call_4", "echo", "invalid json"); + let result = engine.execute_tool_call(&tool_call).await; + + assert!(!result.result.success); + assert!(result + .result + .error + .as_ref() + .unwrap() + .contains("参数解析失败")); + } + + #[tokio::test] + async fn test_execute_all_tool_calls() { + let registry = create_test_registry(); + let engine = ToolLoopEngine::new(registry); + + let tool_calls = vec![ + create_tool_call("call_1", "echo", r#"{"message": "First"}"#), + create_tool_call("call_2", "echo", r#"{"message": "Second"}"#), + ]; + + let results = engine.execute_all_tool_calls(&tool_calls, None).await; + + assert_eq!(results.len(), 2); + assert!(results[0].result.success); + assert_eq!(results[0].result.output, "First"); + assert!(results[1].result.success); + assert_eq!(results[1].result.output, "Second"); + } + + #[test] + fn test_results_to_messages() { + let results = vec![ + ToolCallResult::new( + "call_1".to_string(), + "echo".to_string(), + ToolsResult::success("Hello"), + ), + ToolCallResult::new( + "call_2".to_string(), + "failing".to_string(), + ToolsResult::failure("Error occurred"), + ), + ]; + + let messages = ToolLoopEngine::results_to_messages(&results); + + assert_eq!(messages.len(), 2); + + // 第一个消息(成功) + assert_eq!(messages[0].role, "tool"); + assert_eq!(messages[0].content.as_text(), "Hello"); + assert_eq!(messages[0].tool_call_id, Some("call_1".to_string())); + + // 第二个消息(失败) + assert_eq!(messages[1].role, "tool"); + assert!(messages[1].content.as_text().contains("Error")); + assert_eq!(messages[1].tool_call_id, Some("call_2".to_string())); + } + + #[test] + fn test_should_continue() { + let registry = create_test_registry(); + let engine = ToolLoopEngine::with_config(registry, ToolLoopConfig::new(5)); + + // 有工具调用,未达到限制 + let result_with_tools = StreamResult::new("".to_string()).with_tool_calls(vec![ToolCall { + id: "call_1".to_string(), + call_type: "function".to_string(), + function: FunctionCall { + name: "echo".to_string(), + arguments: "{}".to_string(), + }, + }]); + assert!(engine.should_continue(&result_with_tools, 0)); + assert!(engine.should_continue(&result_with_tools, 4)); + + // 达到最大迭代次数 + assert!(!engine.should_continue(&result_with_tools, 5)); + + // 无工具调用 + let result_no_tools = StreamResult::new("Final response".to_string()); + assert!(!engine.should_continue(&result_no_tools, 0)); + } + + #[test] + fn test_tool_loop_state() { + let mut state = ToolLoopState::new(); + + assert_eq!(state.iteration, 0); + assert_eq!(state.total_tool_calls, 0); + assert!(!state.completed); + assert!(state.final_content.is_none()); + + state.increment_iteration(); + assert_eq!(state.iteration, 1); + + state.add_tool_calls(3); + assert_eq!(state.total_tool_calls, 3); + + state.mark_completed("Final content".to_string()); + assert!(state.completed); + assert_eq!(state.final_content, Some("Final content".to_string())); + } + + #[test] + fn test_tool_call_result_to_agent_message() { + // 成功结果 + let success_result = ToolCallResult::new( + "call_1".to_string(), + "echo".to_string(), + ToolsResult::success("Success output"), + ); + let success_msg = success_result.to_agent_message(); + assert_eq!(success_msg.role, "tool"); + assert_eq!(success_msg.content.as_text(), "Success output"); + assert_eq!(success_msg.tool_call_id, Some("call_1".to_string())); + + // 失败结果 + let failure_result = ToolCallResult::new( + "call_2".to_string(), + "failing".to_string(), + ToolsResult::failure("Something went wrong"), + ); + let failure_msg = failure_result.to_agent_message(); + assert_eq!(failure_msg.role, "tool"); + assert!(failure_msg.content.as_text().contains("Error")); + assert!(failure_msg + .content + .as_text() + .contains("Something went wrong")); + } + + #[tokio::test] + async fn test_execute_all_tool_calls_with_events() { + let registry = create_test_registry(); + let engine = ToolLoopEngine::new(registry); + + let (tx, mut rx) = mpsc::channel::(10); + + let tool_calls = vec![create_tool_call("call_1", "echo", r#"{"message": "Test"}"#)]; + + let results = engine.execute_all_tool_calls(&tool_calls, Some(&tx)).await; + + assert_eq!(results.len(), 1); + assert!(results[0].result.success); + + // 检查事件 + let event1 = rx.recv().await.unwrap(); + assert!(matches!(event1, StreamEvent::ToolStart { .. })); + + let event2 = rx.recv().await.unwrap(); + assert!(matches!(event2, StreamEvent::ToolEnd { .. })); + } +} + +#[cfg(test)] +mod proptests { + use super::*; + use crate::agent::tools::types::{JsonSchema, PropertySchema, ToolDefinition}; + use crate::agent::tools::{Tool, ToolError, ToolRegistry, ToolResult as ToolsResult}; + use crate::agent::types::FunctionCall; + use async_trait::async_trait; + use proptest::prelude::*; + + /// 测试用的 Echo 工具(用于属性测试) + struct PropTestEchoTool; + + #[async_trait] + impl Tool for PropTestEchoTool { + fn definition(&self) -> ToolDefinition { + ToolDefinition::new("echo", "Echo the input message").with_parameters( + JsonSchema::new().add_property( + "message", + PropertySchema::string("The message to echo"), + true, + ), + ) + } + + async fn execute(&self, args: serde_json::Value) -> Result { + let message = args + .get("message") + .and_then(|v| v.as_str()) + .ok_or_else(|| ToolError::InvalidArguments("缺少 message 参数".to_string()))?; + + Ok(ToolsResult::success(message)) + } + } + + /// 测试用的计数工具 + struct PropTestCountTool; + + #[async_trait] + impl Tool for PropTestCountTool { + fn definition(&self) -> ToolDefinition { + ToolDefinition::new("count", "Count characters in a string").with_parameters( + JsonSchema::new().add_property( + "text", + PropertySchema::string("The text to count"), + true, + ), + ) + } + + async fn execute(&self, args: serde_json::Value) -> Result { + let text = args + .get("text") + .and_then(|v| v.as_str()) + .ok_or_else(|| ToolError::InvalidArguments("缺少 text 参数".to_string()))?; + + Ok(ToolsResult::success(format!("{}", text.len()))) + } + } + + fn create_proptest_registry() -> Arc { + let registry = ToolRegistry::new(); + registry.register(PropTestEchoTool).unwrap(); + registry.register(PropTestCountTool).unwrap(); + Arc::new(registry) + } + + /// 生成有效的工具调用 ID + fn arb_tool_id() -> impl Strategy { + "call_[a-zA-Z0-9]{8}".prop_map(|s| s) + } + + /// 生成有效的消息内容(用于 echo 工具) + fn arb_message_content() -> impl Strategy { + "[a-zA-Z0-9 ]{1,50}".prop_map(|s| s) + } + + /// 生成工具调用数量 + fn arb_tool_call_count() -> impl Strategy { + 1..=5usize + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: agent-tool-calling, Property 12: 工具循环执行完整性** + /// **Validates: Requirements 7.1, 7.2** + /// + /// *For any* 包含 tool_calls 的 Agent 响应,Tool Loop 应该执行所有工具并将结果 + /// 作为 tool 角色消息发送回 Agent。 + #[test] + fn prop_tool_loop_executes_all_tools( + tool_ids in prop::collection::vec(arb_tool_id(), 1..=5), + messages in prop::collection::vec(arb_message_content(), 1..=5) + ) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let registry = create_proptest_registry(); + let engine = ToolLoopEngine::new(registry); + + // 创建工具调用列表 + let tool_calls: Vec = tool_ids + .iter() + .zip(messages.iter()) + .map(|(id, msg)| ToolCall { + id: id.clone(), + call_type: "function".to_string(), + function: FunctionCall { + name: "echo".to_string(), + arguments: serde_json::json!({"message": msg}).to_string(), + }, + }) + .collect(); + + let num_calls = tool_calls.len(); + + // 执行所有工具调用 + let results = engine.execute_all_tool_calls(&tool_calls, None).await; + + // 验证:结果数量等于工具调用数量 + prop_assert_eq!( + results.len(), + num_calls, + "结果数量应该等于工具调用数量" + ); + + // 验证:每个结果都有正确的 tool_call_id + for (i, result) in results.iter().enumerate() { + prop_assert_eq!( + &result.tool_call_id, + &tool_ids[i], + "工具调用 ID 应该匹配" + ); + } + + // 验证:所有工具都成功执行 + for result in &results { + prop_assert!( + result.result.success, + "工具执行应该成功: {:?}", + result.result.error + ); + } + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 12: 工具循环执行完整性 - 结果转换为消息** + /// **Validates: Requirements 7.1, 7.2** + /// + /// *For any* 工具执行结果,转换为 AgentMessage 后应该具有正确的 role 和 tool_call_id。 + #[test] + fn prop_tool_results_convert_to_messages( + tool_id in arb_tool_id(), + message in arb_message_content() + ) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let registry = create_proptest_registry(); + let engine = ToolLoopEngine::new(registry); + + let tool_call = ToolCall { + id: tool_id.clone(), + call_type: "function".to_string(), + function: FunctionCall { + name: "echo".to_string(), + arguments: serde_json::json!({"message": message}).to_string(), + }, + }; + + // 执行工具 + let result = engine.execute_tool_call(&tool_call).await; + + // 转换为 AgentMessage + let agent_msg = result.to_agent_message(); + + // 验证:role 为 "tool" + prop_assert_eq!( + agent_msg.role, + "tool", + "消息角色应该为 'tool'" + ); + + // 验证:tool_call_id 正确 + prop_assert_eq!( + agent_msg.tool_call_id, + Some(tool_id.clone()), + "tool_call_id 应该匹配" + ); + + // 验证:成功结果的内容包含原始消息 + if result.result.success { + prop_assert!( + agent_msg.content.as_text().contains(&message), + "成功结果的内容应该包含原始消息" + ); + } + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 12: 工具循环执行完整性 - 事件发送** + /// **Validates: Requirements 7.1, 7.6** + /// + /// *For any* 工具执行,应该发送 ToolStart 和 ToolEnd 事件。 + #[test] + fn prop_tool_execution_sends_events( + tool_id in arb_tool_id(), + message in arb_message_content() + ) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let registry = create_proptest_registry(); + let engine = ToolLoopEngine::new(registry); + + let (tx, mut rx) = mpsc::channel::(10); + + let tool_calls = vec![ToolCall { + id: tool_id.clone(), + call_type: "function".to_string(), + function: FunctionCall { + name: "echo".to_string(), + arguments: serde_json::json!({"message": message}).to_string(), + }, + }]; + + // 执行工具调用 + let _ = engine.execute_all_tool_calls(&tool_calls, Some(&tx)).await; + + // 验证:收到 ToolStart 事件 + let event1 = rx.recv().await; + prop_assert!(event1.is_some(), "应该收到 ToolStart 事件"); + if let Some(StreamEvent::ToolStart { tool_name, tool_id: event_tool_id }) = event1 { + prop_assert_eq!(tool_name, "echo", "工具名称应该为 'echo'"); + prop_assert_eq!(event_tool_id, tool_id.clone(), "工具 ID 应该匹配"); + } else { + prop_assert!(false, "第一个事件应该是 ToolStart"); + } + + // 验证:收到 ToolEnd 事件 + let event2 = rx.recv().await; + prop_assert!(event2.is_some(), "应该收到 ToolEnd 事件"); + if let Some(StreamEvent::ToolEnd { tool_id: event_tool_id, result }) = event2 { + prop_assert_eq!(event_tool_id, tool_id.clone(), "工具 ID 应该匹配"); + prop_assert!(result.success, "工具执行应该成功"); + } else { + prop_assert!(false, "第二个事件应该是 ToolEnd"); + } + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 12: 工具循环执行完整性 - 多工具执行顺序** + /// **Validates: Requirements 7.1, 7.2** + /// + /// *For any* 多个工具调用,执行顺序应该与调用顺序一致。 + #[test] + fn prop_tool_execution_order_preserved( + count in 2..=5usize + ) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let registry = create_proptest_registry(); + let engine = ToolLoopEngine::new(registry); + + // 创建多个工具调用,每个使用不同的消息 + let tool_calls: Vec = (0..count) + .map(|i| ToolCall { + id: format!("call_{}", i), + call_type: "function".to_string(), + function: FunctionCall { + name: "echo".to_string(), + arguments: serde_json::json!({"message": format!("msg_{}", i)}).to_string(), + }, + }) + .collect(); + + // 执行所有工具调用 + let results = engine.execute_all_tool_calls(&tool_calls, None).await; + + // 验证:结果顺序与调用顺序一致 + for (i, result) in results.iter().enumerate() { + prop_assert_eq!( + &result.tool_call_id, + &format!("call_{}", i), + "结果顺序应该与调用顺序一致" + ); + prop_assert!( + result.result.output.contains(&format!("msg_{}", i)), + "结果内容应该对应正确的调用" + ); + } + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 13: 工具循环终止** + /// **Validates: Requirements 7.3, 7.5** + /// + /// *For any* 工具循环执行,当 Agent 响应不包含 tool_calls 时,循环应该终止。 + #[test] + fn prop_tool_loop_terminates_without_tool_calls( + content in "[a-zA-Z0-9 ]{1,100}" + ) { + let registry = create_proptest_registry(); + let engine = ToolLoopEngine::new(registry); + + // 创建不包含工具调用的响应 + let result = StreamResult::new(content.clone()); + + // 验证:should_continue 返回 false + prop_assert!( + !engine.should_continue(&result, 0), + "不包含工具调用的响应应该终止循环" + ); + + // 验证:has_tool_calls 返回 false + prop_assert!( + !ToolLoopEngine::has_tool_calls(&result), + "不包含工具调用的响应 has_tool_calls 应该返回 false" + ); + } + + /// **Feature: agent-tool-calling, Property 13: 工具循环终止 - 最大迭代次数** + /// **Validates: Requirements 7.3, 7.5** + /// + /// *For any* 工具循环执行,当达到最大迭代次数时,循环应该终止。 + #[test] + fn prop_tool_loop_terminates_at_max_iterations( + max_iterations in 1..=20usize, + current_iteration in 0..=25usize + ) { + let registry = create_proptest_registry(); + let config = ToolLoopConfig::new(max_iterations); + let engine = ToolLoopEngine::with_config(registry, config); + + // 创建包含工具调用的响应 + let result = StreamResult::new("".to_string()).with_tool_calls(vec![ToolCall { + id: "call_1".to_string(), + call_type: "function".to_string(), + function: FunctionCall { + name: "echo".to_string(), + arguments: r#"{"message": "test"}"#.to_string(), + }, + }]); + + let should_continue = engine.should_continue(&result, current_iteration); + + if current_iteration >= max_iterations { + // 达到或超过最大迭代次数,应该终止 + prop_assert!( + !should_continue, + "达到最大迭代次数 {} 时应该终止循环(当前迭代: {})", + max_iterations, + current_iteration + ); + } else { + // 未达到最大迭代次数,应该继续 + prop_assert!( + should_continue, + "未达到最大迭代次数 {} 时应该继续循环(当前迭代: {})", + max_iterations, + current_iteration + ); + } + } + + /// **Feature: agent-tool-calling, Property 13: 工具循环终止 - 空工具调用列表** + /// **Validates: Requirements 7.3, 7.5** + /// + /// *For any* 包含空工具调用列表的响应,循环应该终止。 + #[test] + fn prop_tool_loop_terminates_with_empty_tool_calls( + content in "[a-zA-Z0-9 ]{1,100}" + ) { + let registry = create_proptest_registry(); + let engine = ToolLoopEngine::new(registry); + + // 创建包含空工具调用列表的响应 + let result = StreamResult { + content: content.clone(), + tool_calls: Some(vec![]), + usage: None, + }; + + // 验证:should_continue 返回 false + prop_assert!( + !engine.should_continue(&result, 0), + "空工具调用列表应该终止循环" + ); + + // 验证:has_tool_calls 返回 false + prop_assert!( + !ToolLoopEngine::has_tool_calls(&result), + "空工具调用列表 has_tool_calls 应该返回 false" + ); + } + + /// **Feature: agent-tool-calling, Property 13: 工具循环终止 - 配置一致性** + /// **Validates: Requirements 7.5** + /// + /// *For any* 配置的最大迭代次数,engine.max_iterations() 应该返回相同的值。 + #[test] + fn prop_tool_loop_config_consistency( + max_iterations in 1..=100usize + ) { + let registry = create_proptest_registry(); + let config = ToolLoopConfig::new(max_iterations); + let engine = ToolLoopEngine::with_config(registry, config); + + prop_assert_eq!( + engine.max_iterations(), + max_iterations, + "max_iterations() 应该返回配置的值" + ); + } + + /// **Feature: agent-tool-calling, Property 13: 工具循环终止 - 边界条件** + /// **Validates: Requirements 7.3, 7.5** + /// + /// *For any* 最大迭代次数,在边界处的行为应该正确。 + #[test] + fn prop_tool_loop_boundary_conditions( + max_iterations in 1..=20usize + ) { + let registry = create_proptest_registry(); + let config = ToolLoopConfig::new(max_iterations); + let engine = ToolLoopEngine::with_config(registry, config); + + // 创建包含工具调用的响应 + let result_with_tools = StreamResult::new("".to_string()).with_tool_calls(vec![ToolCall { + id: "call_1".to_string(), + call_type: "function".to_string(), + function: FunctionCall { + name: "echo".to_string(), + arguments: r#"{"message": "test"}"#.to_string(), + }, + }]); + + // 在 max_iterations - 1 处应该继续 + if max_iterations > 0 { + prop_assert!( + engine.should_continue(&result_with_tools, max_iterations - 1), + "在 max_iterations - 1 处应该继续" + ); + } + + // 在 max_iterations 处应该终止 + prop_assert!( + !engine.should_continue(&result_with_tools, max_iterations), + "在 max_iterations 处应该终止" + ); + + // 在 max_iterations + 1 处应该终止 + prop_assert!( + !engine.should_continue(&result_with_tools, max_iterations + 1), + "在 max_iterations + 1 处应该终止" + ); + } + } +} diff --git a/src-tauri/src/agent/tools/README.md b/src-tauri/src/agent/tools/README.md new file mode 100644 index 000000000..bb311732c --- /dev/null +++ b/src-tauri/src/agent/tools/README.md @@ -0,0 +1,284 @@ +# 工具系统模块 + + + +## 架构说明 + +Agent 工具系统模块,提供工具定义、注册、执行的核心框架。 + +### 设计决策 + +- **可扩展架构**:通过 `Tool` trait 定义工具接口,便于添加新工具 +- **类型安全**:使用 JSON Schema 定义参数,支持必需和可选参数验证 +- **动态注册**:工具可在运行时注册/注销,无需重启 +- **安全优先**:所有工具执行前进行参数验证,SecurityManager 提供路径安全检查 + +## 文件索引 + +| 文件 | 说明 | +|------|------| +| `mod.rs` | 模块入口,导出公共类型 | +| `types.rs` | 工具类型定义(ToolDefinition, ToolCall, ToolResult, ToolError) | +| `registry.rs` | Tool trait 和 ToolRegistry 实现 | +| `security.rs` | 安全管理器(路径验证、符号链接检查、目录遍历防护) | +| `bash.rs` | Bash 命令执行工具(shell 检测、命令执行、超时控制、环境变量设置) | +| `read_file.rs` | 文件读取工具(带行号读取、行范围读取、大文件检测、目录列表、语言检测) | +| `write_file.rs` | 文件写入工具(文件创建/覆盖、父目录自动创建、换行符规范化、尾部换行符保证) | +| `edit_file.rs` | 文件编辑工具(精确字符串替换、多次出现检测、unified diff、历史栈、撤销功能) | +| `prompt.rs` | 工具 Prompt 生成器(System Prompt 工具注入、XML/JSON 格式转换) | + +## 核心类型 + +### 工具定义 +- `ToolDefinition`: 工具定义结构(名称、描述、参数 Schema) +- `JsonSchema`: JSON Schema 参数定义 +- `PropertySchema`: 属性 Schema(类型、描述、默认值、枚举值) + +### 工具调用 +- `ToolCall`: 工具调用请求(ID、名称、参数) +- `ToolResult`: 工具执行结果(成功/失败、输出、错误信息) + +### 错误类型 +- `ToolError`: 工具执行错误(NotFound, InvalidArguments, ExecutionFailed, Security, Timeout) +- `ToolValidationError`: 工具定义验证错误(EmptyName, EmptyDescription, RequiredPropertyNotDefined, DuplicateName) +- `SecurityError`: 安全错误(PathTraversal, OutsideBaseDir, SymlinkNotAllowed, InvalidPath) + +### 工具接口 +- `Tool` trait: 工具接口,包含 `definition()` 和 `execute()` 方法 +- `ToolRegistry`: 工具注册表,管理所有已注册的工具 + +### 安全管理 +- `SecurityManager`: 安全管理器,验证文件操作的安全性 + - `validate_path()`: 完整路径验证(".." 检查、基础目录检查、符号链接检查) + - `quick_check()`: 快速检查(仅检查 ".." 组件) + - `validate_path_no_symlink_check()`: 不检查符号链接的路径验证 + +### Bash 工具 +- `BashTool`: Bash 命令执行工具 + - `execute_command()`: 执行 shell 命令,捕获 stdout/stderr + - `detect_shell()`: 检测用户默认 shell(bash/zsh/powershell) + - `get_non_interactive_env()`: 获取防止交互的环境变量 +- `ShellType`: Shell 类型枚举(Bash, Zsh, PowerShell, Cmd, Sh) +- `BashExecutionResult`: 命令执行结果(stdout, stderr, exit_code, timed_out) + +### 文件读取工具 +- `ReadFileTool`: 文件读取工具 + - `read_file()`: 读取文件内容,支持行范围 + - 自动检测编程语言 + - 大文件推荐使用行范围 + - 目录自动列出内容 +- `ReadFileResult`: 文件读取结果(content, total_lines, start_line, end_line, language, is_directory, recommend_range, truncated) + +### 文件写入工具 +- `WriteFileTool`: 文件写入工具 + - `write_file()`: 创建或覆盖文件 + - 自动创建父目录 + - 换行符规范化(Unix: LF, Windows: CRLF) + - 确保文件以换行符结尾 +- `WriteFileResult`: 文件写入结果(path, bytes_written, line_count, created, overwritten) + +### 文件编辑工具 +- `EditFileTool`: 文件编辑工具 + - `edit_file()`: 精确字符串替换(old_str → new_str) + - `apply_diff()`: 应用 unified diff 格式的变更 + - `undo_edit()`: 撤销上一次编辑 + - `history_count()`: 获取编辑历史数量 + - `clear_history()`: 清除编辑历史 + - 多次出现检测(返回错误要求更多上下文) + - 不存在检测(返回错误和指导) + - 返回变更上下文片段 +- `EditFileResult`: 文件编辑结果(path, old_str_len, new_str_len, context_snippet, diff) +- `UndoResult`: 撤销结果(path, restored_content_len, previous_content_len) + +### Prompt 生成器 +- `ToolPromptGenerator`: 工具 Prompt 生成器 + - `generate_system_prompt()`: 生成包含工具定义的 System Prompt + - `tool_to_xml()`: 将工具定义转换为 XML 格式 + - `tool_to_json()`: 将工具定义转换为 JSON 格式 +- `PromptFormat`: Prompt 输出格式枚举(Xml, Json) +- `generate_tools_prompt()`: 便捷函数,生成工具 Prompt + +## 使用示例 + +### 定义工具 + +```rust +use crate::agent::tools::{Tool, ToolDefinition, ToolResult, ToolError, JsonSchema, PropertySchema}; +use async_trait::async_trait; + +struct EchoTool; + +#[async_trait] +impl Tool for EchoTool { + fn definition(&self) -> ToolDefinition { + ToolDefinition::new("echo", "Echo the input message") + .with_parameters( + JsonSchema::new() + .add_property("message", PropertySchema::string("The message to echo"), true) + ) + } + + async fn execute(&self, args: serde_json::Value) -> Result { + let message = args.get("message") + .and_then(|v| v.as_str()) + .ok_or_else(|| ToolError::InvalidArguments("缺少 message 参数".to_string()))?; + + Ok(ToolResult::success(message)) + } +} +``` + +### 注册和执行工具 + +```rust +use crate::agent::tools::ToolRegistry; + +let registry = ToolRegistry::new(); + +// 注册工具 +registry.register(EchoTool)?; + +// 执行工具 +let result = registry.execute("echo", serde_json::json!({"message": "Hello!"})).await?; +assert!(result.success); +assert_eq!(result.output, "Hello!"); +``` + +### 使用文件读取工具 + +```rust +use crate::agent::tools::{ReadFileTool, SecurityManager}; +use std::sync::Arc; + +let security = Arc::new(SecurityManager::new("/path/to/project")); +let tool = ReadFileTool::new(security); + +// 读取整个文件 +let result = tool.read_file(Path::new("src/main.rs"), None, None)?; +println!("语言: {:?}", result.language); +println!("总行数: {}", result.total_lines); + +// 读取指定行范围 +let result = tool.read_file(Path::new("src/main.rs"), Some(10), Some(20))?; +println!("内容:\n{}", result.content); +``` + +### 使用文件写入工具 + +```rust +use crate::agent::tools::{WriteFileTool, SecurityManager}; +use std::sync::Arc; + +let security = Arc::new(SecurityManager::new("/path/to/project")); +let tool = WriteFileTool::new(security); + +// 写入新文件 +let result = tool.write_file(Path::new("output.txt"), "Hello, World!")?; +println!("创建: {}, 字节数: {}", result.created, result.bytes_written); + +// 覆盖已有文件 +let result = tool.write_file(Path::new("output.txt"), "New content")?; +println!("覆盖: {}", result.overwritten); + +// 自动创建父目录 +let result = tool.write_file(Path::new("a/b/c/nested.txt"), "Nested content")?; +println!("路径: {:?}", result.path); +``` + +### 使用文件编辑工具 + +```rust +use crate::agent::tools::{EditFileTool, SecurityManager}; +use std::sync::Arc; + +let security = Arc::new(SecurityManager::new("/path/to/project")); +let tool = EditFileTool::new(security); + +// 精确字符串替换 +let result = tool.edit_file(Path::new("src/main.rs"), "old_code", "new_code")?; +println!("替换: {} 字节 -> {} 字节", result.old_str_len, result.new_str_len); +println!("变更上下文:\n{}", result.context_snippet); +println!("Diff:\n{}", result.diff); + +// 撤销编辑 +let undo_result = tool.undo_edit(Path::new("src/main.rs"))?; +println!("已恢复: {} 字节", undo_result.restored_content_len); + +// 查看历史记录数量 +let count = tool.history_count(Path::new("src/main.rs")); +println!("历史记录: {} 条", count); +``` + +### 使用 Prompt 生成器 + +```rust +use crate::agent::tools::{ToolPromptGenerator, PromptFormat, ToolDefinition, JsonSchema, PropertySchema}; + +// 创建工具定义 +let tools = vec![ + ToolDefinition::new("bash", "Execute a bash command") + .with_parameters( + JsonSchema::new() + .add_property("command", PropertySchema::string("The command to execute"), true) + ), + ToolDefinition::new("read_file", "Read file contents") + .with_parameters( + JsonSchema::new() + .add_property("path", PropertySchema::string("The file path"), true) + ), +]; + +// 生成 XML 格式的 System Prompt +let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml); +let system_prompt = generator.generate_system_prompt(&tools); +println!("System Prompt:\n{}", system_prompt); + +// 生成 JSON 格式的 System Prompt +let json_generator = ToolPromptGenerator::new().with_format(PromptFormat::Json); +let json_prompt = json_generator.generate_system_prompt(&tools); +println!("JSON Prompt:\n{}", json_prompt); + +// 使用便捷函数 +use crate::agent::tools::generate_tools_prompt; +let prompt = generate_tools_prompt(&tools, PromptFormat::Xml); +``` + +## 需求追溯 + +- Requirements 2.1: 工具定义包含 name, description, JSON Schema parameters +- Requirements 2.2: 注册时验证工具定义 +- Requirements 2.4: 运行时添加工具无需重启 +- Requirements 2.5: 支持必需和可选参数类型验证 +- Requirements 3.1: Bash 工具在用户默认 shell 中执行命令 +- Requirements 3.2: Bash 工具捕获 stdout 和 stderr +- Requirements 3.3: Bash 工具支持超时控制 +- Requirements 3.4: Bash 工具设置防止交互的环境变量 +- Requirements 3.5: Bash 工具返回退出码和错误输出 +- Requirements 3.6: Bash 工具支持可配置的工作目录 +- Requirements 4.1: 文件读取工具返回带行号的内容 +- Requirements 4.2: 文件读取工具支持行范围读取 +- Requirements 4.3: 文件不存在时返回清晰错误信息 +- Requirements 4.4: 大文件推荐使用行范围 +- Requirements 4.5: 检测并报告文件的编程语言 +- Requirements 4.6: 路径为目录时列出目录内容 +- Requirements 5.1: 文件写入工具创建或覆盖文件 +- Requirements 5.2: 文件写入工具自动创建父目录 +- Requirements 5.3: 文件写入工具规范化换行符(Unix: LF, Windows: CRLF) +- Requirements 5.4: 文件写入工具确保文件以换行符结尾 +- Requirements 5.5: 写入失败时返回描述性错误信息 +- Requirements 6.1: 文件编辑工具精确替换匹配的字符串 +- Requirements 6.2: 多次出现时返回错误要求更多上下文 +- Requirements 6.3: 字符串不存在时返回错误和指导 +- Requirements 6.4: 支持 unified diff 格式 +- Requirements 6.5: 维护历史栈支持撤销操作 +- Requirements 6.6: 编辑后返回变更上下文片段 +- Requirements 8.1: 验证所有文件路径防止目录遍历攻击 +- Requirements 8.2: 拒绝包含 ".." 组件的路径 +- Requirements 8.3: 拒绝符号链接操作 +- Requirements 8.4: Bash 工具设置环境变量禁用交互式编辑器和提示 +- Requirements 8.5: 强制执行可配置的基础目录 +- Requirements 2.3: System Prompt 包含所有可用工具定义 + +## 更新提醒 + +任何文件变更后,请更新此文档和相关的上级文档。 diff --git a/src-tauri/src/agent/tools/bash.rs b/src-tauri/src/agent/tools/bash.rs new file mode 100644 index 000000000..8a9e92e10 --- /dev/null +++ b/src-tauri/src/agent/tools/bash.rs @@ -0,0 +1,1062 @@ +//! Bash 工具模块 +//! +//! 提供 Shell 命令执行功能,支持 bash/zsh/powershell +//! 符合 Requirements 3.1, 3.2, 3.3, 3.4, 3.5, 3.6 +//! +//! ## 功能 +//! - Shell 配置检测(bash/zsh/powershell) +//! - 命令执行(捕获 stdout/stderr) +//! - 超时控制 +//! - 防止交互的环境变量设置 + +use super::registry::Tool; +use super::security::SecurityManager; +use super::types::{JsonSchema, PropertySchema, ToolDefinition, ToolError, ToolResult}; +use async_trait::async_trait; +use std::collections::HashMap; +use std::env; +use std::path::PathBuf; +use std::process::Stdio; +use std::sync::Arc; +use std::time::Duration; +use tokio::io::AsyncReadExt; +use tokio::process::Command; +use tokio::time::timeout; +use tracing::{debug, info, warn}; + +/// 默认超时时间(秒) +const DEFAULT_TIMEOUT_SECS: u64 = 120; + +/// 最大输出大小(字节) +const MAX_OUTPUT_SIZE: usize = 1024 * 1024; // 1MB + +/// Shell 类型 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ShellType { + /// Bash shell + Bash, + /// Zsh shell + Zsh, + /// PowerShell (Windows) + PowerShell, + /// Cmd (Windows fallback) + Cmd, + /// Sh (Unix fallback) + Sh, +} + +impl ShellType { + /// 获取 shell 可执行文件路径 + pub fn executable(&self) -> &'static str { + match self { + ShellType::Bash => "bash", + ShellType::Zsh => "zsh", + ShellType::PowerShell => "powershell", + ShellType::Cmd => "cmd", + ShellType::Sh => "sh", + } + } + + /// 获取执行命令的参数 + pub fn command_args(&self) -> Vec<&'static str> { + match self { + ShellType::Bash | ShellType::Zsh | ShellType::Sh => vec!["-c"], + ShellType::PowerShell => vec!["-Command"], + ShellType::Cmd => vec!["/C"], + } + } +} + +/// Bash 执行结果 +#[derive(Debug, Clone)] +pub struct BashExecutionResult { + /// 标准输出 + pub stdout: String, + /// 标准错误 + pub stderr: String, + /// 退出码 + pub exit_code: Option, + /// 是否超时 + pub timed_out: bool, +} + +impl BashExecutionResult { + /// 创建成功结果 + pub fn success(stdout: String, stderr: String, exit_code: i32) -> Self { + Self { + stdout, + stderr, + exit_code: Some(exit_code), + timed_out: false, + } + } + + /// 创建超时结果 + pub fn timeout(stdout: String, stderr: String) -> Self { + Self { + stdout, + stderr, + exit_code: None, + timed_out: true, + } + } + + /// 检查命令是否成功执行 + pub fn is_success(&self) -> bool { + self.exit_code == Some(0) && !self.timed_out + } + + /// 获取合并的输出 + pub fn combined_output(&self) -> String { + let mut output = String::new(); + if !self.stdout.is_empty() { + output.push_str(&self.stdout); + } + if !self.stderr.is_empty() { + if !output.is_empty() { + output.push('\n'); + } + output.push_str(&self.stderr); + } + output + } +} + +/// Bash 工具 +/// +/// 执行 shell 命令并返回结果 +/// Requirements: 3.1, 3.2, 3.3, 3.4, 3.5, 3.6 +pub struct BashTool { + /// 安全管理器 + security: Arc, + /// 默认工作目录 + working_dir: PathBuf, + /// 超时时间(秒) + timeout_secs: u64, + /// Shell 类型 + shell_type: ShellType, +} + +impl BashTool { + /// 创建新的 Bash 工具 + pub fn new(security: Arc) -> Self { + let working_dir = security.base_dir().to_path_buf(); + let shell_type = Self::detect_shell(); + + Self { + security, + working_dir, + timeout_secs: DEFAULT_TIMEOUT_SECS, + shell_type, + } + } + + /// 设置超时时间 + pub fn with_timeout(mut self, timeout_secs: u64) -> Self { + self.timeout_secs = timeout_secs; + self + } + + /// 设置工作目录 + pub fn with_working_dir(mut self, working_dir: PathBuf) -> Self { + self.working_dir = working_dir; + self + } + + /// 设置 shell 类型 + pub fn with_shell_type(mut self, shell_type: ShellType) -> Self { + self.shell_type = shell_type; + self + } + + /// 检测用户默认 shell + /// + /// Requirements: 3.1 - THE Bash_Executor SHALL execute it in the user's default shell + pub fn detect_shell() -> ShellType { + #[cfg(windows)] + { + // Windows 优先使用 PowerShell + if Self::is_shell_available("powershell") { + return ShellType::PowerShell; + } + return ShellType::Cmd; + } + + #[cfg(not(windows))] + { + // Unix 系统检查 SHELL 环境变量 + if let Ok(shell) = env::var("SHELL") { + if shell.contains("zsh") { + return ShellType::Zsh; + } + if shell.contains("bash") { + return ShellType::Bash; + } + } + + // 检查可用的 shell + if Self::is_shell_available("zsh") { + return ShellType::Zsh; + } + if Self::is_shell_available("bash") { + return ShellType::Bash; + } + ShellType::Sh + } + } + + /// 检查 shell 是否可用 + fn is_shell_available(shell: &str) -> bool { + std::process::Command::new("which") + .arg(shell) + .output() + .map(|o| o.status.success()) + .unwrap_or(false) + } + + /// 获取防止交互的环境变量 + /// + /// Requirements: 3.4 - THE Bash_Executor SHALL set appropriate environment variables to prevent interactive prompts + /// Requirements: 8.4 - THE Bash_Executor SHALL set environment variables to disable interactive editors and prompts + pub fn get_non_interactive_env() -> HashMap { + let mut env = HashMap::new(); + + // 标记为非交互终端 + env.insert("GOOSE_TERMINAL".to_string(), "1".to_string()); + + // 禁用 Git 交互提示 + env.insert("GIT_TERMINAL_PROMPT".to_string(), "0".to_string()); + + // 设置非交互编辑器 + env.insert("EDITOR".to_string(), "cat".to_string()); + env.insert("VISUAL".to_string(), "cat".to_string()); + env.insert("GIT_EDITOR".to_string(), "cat".to_string()); + + // 禁用 SSH 交互 + env.insert( + "GIT_SSH_COMMAND".to_string(), + "ssh -o BatchMode=yes".to_string(), + ); + + // 禁用 GPG 交互 + env.insert("GPG_TTY".to_string(), "".to_string()); + + // 设置 CI 环境变量(许多工具会检查这个) + env.insert("CI".to_string(), "true".to_string()); + + // 禁用颜色输出(避免 ANSI 转义序列) + env.insert("NO_COLOR".to_string(), "1".to_string()); + env.insert("TERM".to_string(), "dumb".to_string()); + + // npm/yarn 非交互模式 + env.insert("npm_config_yes".to_string(), "true".to_string()); + + // 禁用 pager + env.insert("PAGER".to_string(), "cat".to_string()); + env.insert("GIT_PAGER".to_string(), "cat".to_string()); + + env + } + + /// 执行命令 + /// + /// Requirements: 3.1, 3.2, 3.3, 3.5 + pub async fn execute_command( + &self, + command: &str, + working_dir: Option<&PathBuf>, + timeout_secs: Option, + ) -> Result { + let work_dir = working_dir.unwrap_or(&self.working_dir); + let timeout_duration = Duration::from_secs(timeout_secs.unwrap_or(self.timeout_secs)); + + info!( + "[BashTool] 执行命令: {} (工作目录: {:?}, 超时: {:?})", + command, work_dir, timeout_duration + ); + + // 构建命令 + let mut cmd = Command::new(self.shell_type.executable()); + + // 添加命令参数 + for arg in self.shell_type.command_args() { + cmd.arg(arg); + } + cmd.arg(command); + + // 设置工作目录 + // Requirements: 3.6 - THE Bash_Executor SHALL support a configurable working directory + cmd.current_dir(work_dir); + + // 设置环境变量 + // Requirements: 3.4 - prevent interactive prompts + for (key, value) in Self::get_non_interactive_env() { + cmd.env(key, value); + } + + // 配置标准输入输出 + cmd.stdin(Stdio::null()); + cmd.stdout(Stdio::piped()); + cmd.stderr(Stdio::piped()); + + // 启动进程 + let mut child = cmd + .spawn() + .map_err(|e| ToolError::ExecutionFailed(format!("无法启动 shell 进程: {}", e)))?; + + // 获取输出流 + let mut stdout = child + .stdout + .take() + .ok_or_else(|| ToolError::ExecutionFailed("无法获取 stdout".to_string()))?; + let mut stderr = child + .stderr + .take() + .ok_or_else(|| ToolError::ExecutionFailed("无法获取 stderr".to_string()))?; + + // 异步读取输出 + let stdout_handle = tokio::spawn(async move { + let mut buffer = Vec::new(); + let _ = stdout.read_to_end(&mut buffer).await; + buffer + }); + + let stderr_handle = tokio::spawn(async move { + let mut buffer = Vec::new(); + let _ = stderr.read_to_end(&mut buffer).await; + buffer + }); + + // 等待进程完成(带超时) + // Requirements: 3.3 - WHEN a command exceeds the timeout limit, THE Bash_Executor SHALL terminate it + let result = timeout(timeout_duration, child.wait()).await; + + match result { + Ok(Ok(status)) => { + // 进程正常完成 + let stdout_bytes = stdout_handle.await.unwrap_or_default(); + let stderr_bytes = stderr_handle.await.unwrap_or_default(); + + let stdout_str = + Self::truncate_output(String::from_utf8_lossy(&stdout_bytes).to_string()); + let stderr_str = + Self::truncate_output(String::from_utf8_lossy(&stderr_bytes).to_string()); + + let exit_code = status.code().unwrap_or(-1); + + debug!( + "[BashTool] 命令完成: exit_code={}, stdout_len={}, stderr_len={}", + exit_code, + stdout_str.len(), + stderr_str.len() + ); + + Ok(BashExecutionResult::success( + stdout_str, stderr_str, exit_code, + )) + } + Ok(Err(e)) => { + // 进程等待失败 + Err(ToolError::ExecutionFailed(format!("等待进程失败: {}", e))) + } + Err(_) => { + // 超时 + warn!("[BashTool] 命令超时,正在终止进程"); + + // 尝试终止进程 + let _ = child.kill().await; + + // 获取已有的输出 + let stdout_bytes = stdout_handle.await.unwrap_or_default(); + let stderr_bytes = stderr_handle.await.unwrap_or_default(); + + let stdout_str = + Self::truncate_output(String::from_utf8_lossy(&stdout_bytes).to_string()); + let stderr_str = + Self::truncate_output(String::from_utf8_lossy(&stderr_bytes).to_string()); + + Ok(BashExecutionResult::timeout(stdout_str, stderr_str)) + } + } + } + + /// 截断过长的输出 + fn truncate_output(output: String) -> String { + if output.len() > MAX_OUTPUT_SIZE { + let truncated = &output[..MAX_OUTPUT_SIZE]; + format!( + "{}\n\n[输出已截断,原始大小: {} 字节]", + truncated, + output.len() + ) + } else { + output + } + } + + /// 获取当前 shell 类型 + pub fn shell_type(&self) -> ShellType { + self.shell_type + } + + /// 获取当前工作目录 + pub fn working_dir(&self) -> &PathBuf { + &self.working_dir + } + + /// 获取超时时间 + pub fn timeout_secs(&self) -> u64 { + self.timeout_secs + } +} + +#[async_trait] +impl Tool for BashTool { + fn definition(&self) -> ToolDefinition { + ToolDefinition::new( + "bash", + "Execute a bash command in the shell. Use this for running system commands, \ + scripts, or any command-line operations. The command will be executed in the \ + configured working directory with a timeout limit.", + ) + .with_parameters( + JsonSchema::new() + .add_property( + "command", + PropertySchema::string( + "The bash command to execute. Can be any valid shell command.", + ), + true, + ) + .add_property( + "timeout", + PropertySchema::integer( + "Optional timeout in seconds. Defaults to 120 seconds.", + ) + .with_default(serde_json::json!(120)), + false, + ), + ) + } + + async fn execute(&self, args: serde_json::Value) -> Result { + // 解析参数 + let command = args + .get("command") + .and_then(|v| v.as_str()) + .ok_or_else(|| ToolError::InvalidArguments("缺少 command 参数".to_string()))?; + + let timeout_secs = args + .get("timeout") + .and_then(|v| v.as_u64()) + .unwrap_or(self.timeout_secs); + + // 执行命令 + let result = self + .execute_command(command, None, Some(timeout_secs)) + .await?; + + // 构建输出 + // Requirements: 3.2 - THE Bash_Executor SHALL capture both stdout and stderr + // Requirements: 3.5 - IF a command fails, THEN THE Bash_Executor SHALL return the exit code and error output + if result.timed_out { + let output = format!( + "命令执行超时({}秒)\n\n已捕获的输出:\n{}", + timeout_secs, + result.combined_output() + ); + return Err(ToolError::Timeout); + } + + let output = if result.is_success() { + result.combined_output() + } else { + format!( + "命令执行失败 (退出码: {})\n\n{}", + result.exit_code.unwrap_or(-1), + result.combined_output() + ) + }; + + if result.is_success() { + Ok(ToolResult::success(output)) + } else { + Ok(ToolResult::failure_with_output( + output.clone(), + format!("命令退出码: {}", result.exit_code.unwrap_or(-1)), + )) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::TempDir; + + fn setup_test_tool() -> (BashTool, TempDir) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + (tool, temp_dir) + } + + #[test] + fn test_shell_detection() { + let shell_type = BashTool::detect_shell(); + // 应该检测到某种 shell + assert!(matches!( + shell_type, + ShellType::Bash + | ShellType::Zsh + | ShellType::Sh + | ShellType::PowerShell + | ShellType::Cmd + )); + } + + #[test] + fn test_non_interactive_env() { + let env = BashTool::get_non_interactive_env(); + + // 检查关键环境变量 + assert_eq!(env.get("GOOSE_TERMINAL"), Some(&"1".to_string())); + assert_eq!(env.get("GIT_TERMINAL_PROMPT"), Some(&"0".to_string())); + assert_eq!(env.get("CI"), Some(&"true".to_string())); + } + + #[test] + fn test_tool_definition() { + let (tool, _temp_dir) = setup_test_tool(); + let def = tool.definition(); + + assert_eq!(def.name, "bash"); + assert!(!def.description.is_empty()); + assert!(def.parameters.required.contains(&"command".to_string())); + } + + #[tokio::test] + async fn test_execute_simple_command() { + let (tool, _temp_dir) = setup_test_tool(); + + let result = tool + .execute_command("echo 'Hello, World!'", None, None) + .await; + + assert!(result.is_ok()); + let result = result.unwrap(); + assert!(result.is_success()); + assert!(result.stdout.contains("Hello, World!")); + } + + #[tokio::test] + async fn test_execute_command_with_stderr() { + let (tool, _temp_dir) = setup_test_tool(); + + // 使用一个会产生 stderr 的命令 + let result = tool + .execute_command("echo 'stdout' && echo 'stderr' >&2", None, None) + .await; + + assert!(result.is_ok()); + let result = result.unwrap(); + assert!(result.is_success()); + assert!(result.stdout.contains("stdout")); + assert!(result.stderr.contains("stderr")); + } + + #[tokio::test] + async fn test_execute_failing_command() { + let (tool, _temp_dir) = setup_test_tool(); + + let result = tool.execute_command("exit 1", None, None).await; + + assert!(result.is_ok()); + let result = result.unwrap(); + assert!(!result.is_success()); + assert_eq!(result.exit_code, Some(1)); + } + + #[tokio::test] + async fn test_execute_with_timeout() { + let (tool, _temp_dir) = setup_test_tool(); + + // 使用一个会超时的命令(1秒超时) + let result = tool.execute_command("sleep 10", None, Some(1)).await; + + assert!(result.is_ok()); + let result = result.unwrap(); + assert!(result.timed_out); + assert!(result.exit_code.is_none()); + } + + #[tokio::test] + async fn test_tool_execute() { + let (tool, _temp_dir) = setup_test_tool(); + + let result = tool + .execute(serde_json::json!({ + "command": "echo 'test'" + })) + .await; + + assert!(result.is_ok()); + let result = result.unwrap(); + assert!(result.success); + assert!(result.output.contains("test")); + } + + #[tokio::test] + async fn test_tool_execute_missing_command() { + let (tool, _temp_dir) = setup_test_tool(); + + let result = tool.execute(serde_json::json!({})).await; + + assert!(result.is_err()); + assert!(matches!(result, Err(ToolError::InvalidArguments(_)))); + } + + #[test] + fn test_shell_type_executable() { + assert_eq!(ShellType::Bash.executable(), "bash"); + assert_eq!(ShellType::Zsh.executable(), "zsh"); + assert_eq!(ShellType::PowerShell.executable(), "powershell"); + } + + #[test] + fn test_shell_type_command_args() { + assert_eq!(ShellType::Bash.command_args(), vec!["-c"]); + assert_eq!(ShellType::Zsh.command_args(), vec!["-c"]); + assert_eq!(ShellType::PowerShell.command_args(), vec!["-Command"]); + assert_eq!(ShellType::Cmd.command_args(), vec!["/C"]); + } + + #[test] + fn test_truncate_output() { + // 短输出不截断 + let short = "Hello".to_string(); + assert_eq!(BashTool::truncate_output(short.clone()), short); + + // 长输出截断 + let long = "x".repeat(MAX_OUTPUT_SIZE + 100); + let truncated = BashTool::truncate_output(long.clone()); + assert!(truncated.len() < long.len()); + assert!(truncated.contains("[输出已截断")); + } + + #[test] + fn test_bash_execution_result() { + let success = BashExecutionResult::success("out".to_string(), "err".to_string(), 0); + assert!(success.is_success()); + assert_eq!(success.exit_code, Some(0)); + + let failure = BashExecutionResult::success("out".to_string(), "err".to_string(), 1); + assert!(!failure.is_success()); + + let timeout = BashExecutionResult::timeout("out".to_string(), "err".to_string()); + assert!(!timeout.is_success()); + assert!(timeout.timed_out); + } + + #[test] + fn test_combined_output() { + let result = BashExecutionResult::success( + "stdout content".to_string(), + "stderr content".to_string(), + 0, + ); + let combined = result.combined_output(); + assert!(combined.contains("stdout content")); + assert!(combined.contains("stderr content")); + } +} + +#[cfg(test)] +mod proptests { + use super::*; + use proptest::prelude::*; + use tempfile::TempDir; + + /// 生成有效的简单命令(echo 命令) + /// 避免以 '-' 开头(会被 echo 解释为选项) + fn arb_echo_content() -> impl Strategy { + // 生成安全的字符串内容(以字母开头,避免特殊 shell 字符) + "[a-zA-Z][a-zA-Z0-9 _]{0,49}".prop_map(|s| s) + } + + /// 生成有效的退出码 + fn arb_exit_code() -> impl Strategy { + 0..=255i32 + } + + /// 生成有效的超时时间(秒) + fn arb_timeout_secs() -> impl Strategy { + 1..=10u64 + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: agent-tool-calling, Property 4: Bash 命令执行捕获** + /// **Validates: Requirements 3.1, 3.2, 3.5** + /// + /// *For any* Bash 命令执行,返回的结果应该包含命令的 stdout 和 stderr 输出, + /// 以及正确的退出码。 + #[test] + fn prop_bash_captures_stdout(content in arb_echo_content()) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + + // 执行 echo 命令 + let command = format!("echo '{}'", content); + let result = tool.execute_command(&command, None, None).await; + + prop_assert!(result.is_ok(), "命令执行应该成功"); + let result = result.unwrap(); + + // 验证 stdout 包含输出内容 + prop_assert!( + result.stdout.contains(&content), + "stdout 应该包含 echo 的内容: expected '{}', got '{}'", + content, + result.stdout + ); + + // 验证退出码为 0 + prop_assert_eq!( + result.exit_code, + Some(0), + "成功命令的退出码应该为 0" + ); + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 4: Bash 命令执行捕获 - stderr** + /// **Validates: Requirements 3.1, 3.2, 3.5** + /// + /// *For any* 输出到 stderr 的命令,结果应该捕获 stderr 内容。 + #[test] + fn prop_bash_captures_stderr(content in arb_echo_content()) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + + // 执行输出到 stderr 的命令 + let command = format!("echo '{}' >&2", content); + let result = tool.execute_command(&command, None, None).await; + + prop_assert!(result.is_ok(), "命令执行应该成功"); + let result = result.unwrap(); + + // 验证 stderr 包含输出内容 + prop_assert!( + result.stderr.contains(&content), + "stderr 应该包含 echo 的内容: expected '{}', got '{}'", + content, + result.stderr + ); + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 4: Bash 命令执行捕获 - 退出码** + /// **Validates: Requirements 3.1, 3.2, 3.5** + /// + /// *For any* 指定退出码的命令,结果应该返回正确的退出码。 + #[test] + fn prop_bash_captures_exit_code(exit_code in arb_exit_code()) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + + // 执行带指定退出码的命令 + let command = format!("exit {}", exit_code); + let result = tool.execute_command(&command, None, None).await; + + prop_assert!(result.is_ok(), "命令执行应该完成(即使退出码非零)"); + let result = result.unwrap(); + + // 验证退出码正确 + prop_assert_eq!( + result.exit_code, + Some(exit_code), + "退出码应该匹配: expected {}, got {:?}", + exit_code, + result.exit_code + ); + + // 验证 is_success 方法正确 + prop_assert_eq!( + result.is_success(), + exit_code == 0, + "is_success() 应该在退出码为 0 时返回 true" + ); + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 4: Bash 命令执行捕获 - stdout 和 stderr 同时** + /// **Validates: Requirements 3.1, 3.2, 3.5** + /// + /// *For any* 同时输出到 stdout 和 stderr 的命令,结果应该分别捕获两者。 + #[test] + fn prop_bash_captures_both_streams( + stdout_content in arb_echo_content(), + stderr_content in arb_echo_content() + ) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + + // 执行同时输出到 stdout 和 stderr 的命令 + let command = format!("echo '{}' && echo '{}' >&2", stdout_content, stderr_content); + let result = tool.execute_command(&command, None, None).await; + + prop_assert!(result.is_ok(), "命令执行应该成功"); + let result = result.unwrap(); + + // 验证 stdout 包含内容 + prop_assert!( + result.stdout.contains(&stdout_content), + "stdout 应该包含内容: expected '{}', got '{}'", + stdout_content, + result.stdout + ); + + // 验证 stderr 包含内容 + prop_assert!( + result.stderr.contains(&stderr_content), + "stderr 应该包含内容: expected '{}', got '{}'", + stderr_content, + result.stderr + ); + + // 验证 combined_output 包含两者 + let combined = result.combined_output(); + prop_assert!( + combined.contains(&stdout_content) && combined.contains(&stderr_content), + "combined_output 应该包含 stdout 和 stderr" + ); + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 4: Bash 命令执行捕获 - 超时** + /// **Validates: Requirements 3.3** + /// + /// *For any* 超时的命令,结果应该标记为超时。 + #[test] + fn prop_bash_timeout_detection(timeout_secs in 1u64..=2u64) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + + // 执行一个会超时的命令(sleep 时间大于超时时间) + let sleep_time = timeout_secs + 5; + let command = format!("sleep {}", sleep_time); + let result = tool.execute_command(&command, None, Some(timeout_secs)).await; + + prop_assert!(result.is_ok(), "超时命令应该返回结果而不是错误"); + let result = result.unwrap(); + + // 验证超时标记 + prop_assert!( + result.timed_out, + "命令应该被标记为超时" + ); + + // 验证退出码为 None(因为进程被终止) + prop_assert!( + result.exit_code.is_none(), + "超时命令的退出码应该为 None" + ); + + // 验证 is_success 返回 false + prop_assert!( + !result.is_success(), + "超时命令的 is_success() 应该返回 false" + ); + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 5: Bash 环境变量设置** + /// **Validates: Requirements 3.4, 8.4** + /// + /// *For any* Bash 命令执行,执行环境应该包含 GOOSE_TERMINAL、GIT_TERMINAL_PROMPT=0 等 + /// 防止交互的环境变量。 + #[test] + fn prop_bash_env_goose_terminal_set(_dummy in 0..100u32) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + + // 检查 GOOSE_TERMINAL 环境变量 + let result = tool.execute_command("echo $GOOSE_TERMINAL", None, None).await; + + prop_assert!(result.is_ok(), "命令执行应该成功"); + let result = result.unwrap(); + + prop_assert!( + result.stdout.trim() == "1", + "GOOSE_TERMINAL 应该设置为 1,实际值: '{}'", + result.stdout.trim() + ); + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 5: Bash 环境变量设置 - GIT_TERMINAL_PROMPT** + /// **Validates: Requirements 3.4, 8.4** + /// + /// *For any* Bash 命令执行,GIT_TERMINAL_PROMPT 应该设置为 0 以禁用 Git 交互提示。 + #[test] + fn prop_bash_env_git_terminal_prompt_disabled(_dummy in 0..100u32) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + + // 检查 GIT_TERMINAL_PROMPT 环境变量 + let result = tool.execute_command("echo $GIT_TERMINAL_PROMPT", None, None).await; + + prop_assert!(result.is_ok(), "命令执行应该成功"); + let result = result.unwrap(); + + prop_assert!( + result.stdout.trim() == "0", + "GIT_TERMINAL_PROMPT 应该设置为 0,实际值: '{}'", + result.stdout.trim() + ); + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 5: Bash 环境变量设置 - CI** + /// **Validates: Requirements 3.4, 8.4** + /// + /// *For any* Bash 命令执行,CI 环境变量应该设置为 true。 + #[test] + fn prop_bash_env_ci_set(_dummy in 0..100u32) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + + // 检查 CI 环境变量 + let result = tool.execute_command("echo $CI", None, None).await; + + prop_assert!(result.is_ok(), "命令执行应该成功"); + let result = result.unwrap(); + + prop_assert!( + result.stdout.trim() == "true", + "CI 应该设置为 true,实际值: '{}'", + result.stdout.trim() + ); + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 5: Bash 环境变量设置 - EDITOR** + /// **Validates: Requirements 3.4, 8.4** + /// + /// *For any* Bash 命令执行,EDITOR 应该设置为非交互式编辑器(cat)。 + #[test] + fn prop_bash_env_editor_non_interactive(_dummy in 0..100u32) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + + // 检查 EDITOR 环境变量 + let result = tool.execute_command("echo $EDITOR", None, None).await; + + prop_assert!(result.is_ok(), "命令执行应该成功"); + let result = result.unwrap(); + + prop_assert!( + result.stdout.trim() == "cat", + "EDITOR 应该设置为 cat,实际值: '{}'", + result.stdout.trim() + ); + + Ok(()) + })?; + } + + /// **Feature: agent-tool-calling, Property 5: Bash 环境变量设置 - 所有关键变量** + /// **Validates: Requirements 3.4, 8.4** + /// + /// *For any* Bash 命令执行,所有防止交互的关键环境变量都应该正确设置。 + #[test] + fn prop_bash_env_all_non_interactive_vars(_dummy in 0..100u32) { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = BashTool::new(security); + + // 检查多个环境变量 + let result = tool.execute_command( + "echo \"GOOSE=$GOOSE_TERMINAL,GIT=$GIT_TERMINAL_PROMPT,CI=$CI,PAGER=$PAGER\"", + None, + None + ).await; + + prop_assert!(result.is_ok(), "命令执行应该成功"); + let result = result.unwrap(); + let output = result.stdout.trim(); + + // 验证所有关键变量 + prop_assert!( + output.contains("GOOSE=1"), + "GOOSE_TERMINAL 应该为 1,输出: '{}'", + output + ); + prop_assert!( + output.contains("GIT=0"), + "GIT_TERMINAL_PROMPT 应该为 0,输出: '{}'", + output + ); + prop_assert!( + output.contains("CI=true"), + "CI 应该为 true,输出: '{}'", + output + ); + prop_assert!( + output.contains("PAGER=cat"), + "PAGER 应该为 cat,输出: '{}'", + output + ); + + Ok(()) + })?; + } + } +} diff --git a/src-tauri/src/agent/tools/edit_file.rs b/src-tauri/src/agent/tools/edit_file.rs new file mode 100644 index 000000000..e7c626e1c --- /dev/null +++ b/src-tauri/src/agent/tools/edit_file.rs @@ -0,0 +1,1628 @@ +//! 文件编辑工具模块 +//! +//! 提供文件精确编辑功能,支持字符串替换、多次出现检测、unified diff、历史栈和撤销 +//! 符合 Requirements 6.1, 6.2, 6.3, 6.4, 6.5, 6.6 +//! +//! ## 功能 +//! - 精确字符串替换(old_str → new_str) +//! - 多次出现检测(返回错误要求更多上下文) +//! - 不存在检测(返回错误和指导) +//! - Unified diff 支持 +//! - 历史栈和撤销功能 +//! - 返回变更上下文片段 + +use super::registry::Tool; +use super::security::SecurityManager; +use super::types::{JsonSchema, PropertySchema, ToolDefinition, ToolError, ToolResult}; +use async_trait::async_trait; +use parking_lot::RwLock; +use std::collections::HashMap; +use std::fs; +use std::path::{Path, PathBuf}; +use std::sync::Arc; +use tracing::{debug, info, warn}; + +/// 上下文行数(显示变更前后的行数) +const CONTEXT_LINES: usize = 3; + +/// 最大历史记录数 +const MAX_HISTORY_SIZE: usize = 100; + +/// 文件编辑工具 +/// +/// 提供精确的文件编辑功能,支持撤销操作 +/// Requirements: 6.1, 6.2, 6.3, 6.4, 6.5, 6.6 +pub struct EditFileTool { + /// 安全管理器 + security: Arc, + /// 编辑历史栈(文件路径 -> 历史记录列表) + history: Arc>>>, +} + +impl EditFileTool { + /// 创建新的文件编辑工具 + pub fn new(security: Arc) -> Self { + Self { + security, + history: Arc::new(RwLock::new(HashMap::new())), + } + } + + /// 编辑文件(精确字符串替换) + /// + /// Requirements: 6.1 - WHEN the Agent calls the edit_file tool with old_str and new_str, + /// THE File_Editor SHALL replace the exact match + /// Requirements: 6.2 - IF old_str appears multiple times, THEN THE File_Editor SHALL + /// return an error requiring more context + /// Requirements: 6.3 - IF old_str does not appear in the file, THEN THE File_Editor SHALL + /// return an error with guidance + pub fn edit_file( + &self, + path: &Path, + old_str: &str, + new_str: &str, + ) -> Result { + // 验证路径安全性 + let validated_path = self + .security + .validate_path(path) + .map_err(|e| ToolError::Security(e.to_string()))?; + + // 检查文件是否存在 + if !validated_path.exists() { + return Err(ToolError::ExecutionFailed(format!( + "文件不存在: {}", + path.display() + ))); + } + + // 读取文件内容 + let original_content = fs::read_to_string(&validated_path).map_err(|e| { + ToolError::ExecutionFailed(format!("无法读取文件 {}: {}", path.display(), e)) + })?; + + // 检查 old_str 出现次数 + // Requirements: 6.2 - IF old_str appears multiple times, THEN return error + let occurrences = count_occurrences(&original_content, old_str); + + if occurrences == 0 { + // Requirements: 6.3 - IF old_str does not appear, THEN return error with guidance + return Err(ToolError::ExecutionFailed(format!( + "在文件 {} 中未找到要替换的字符串。\n\ + 请确保 old_str 与文件内容完全匹配(包括空格和换行符)。\n\ + 提示:可以先使用 read_file 工具查看文件内容。", + path.display() + ))); + } + + if occurrences > 1 { + // Requirements: 6.2 - Multiple occurrences, require more context + let positions = find_occurrence_positions(&original_content, old_str); + let context_snippets = positions + .iter() + .take(3) + .map(|&pos| get_context_snippet(&original_content, pos, old_str.len())) + .collect::>() + .join("\n---\n"); + + return Err(ToolError::ExecutionFailed(format!( + "在文件 {} 中找到 {} 处匹配,需要更多上下文来确定要替换哪一处。\n\ + 请在 old_str 中包含更多周围的内容以唯一标识要替换的位置。\n\n\ + 匹配位置示例:\n{}", + path.display(), + occurrences, + context_snippets + ))); + } + + // 执行替换 + // Requirements: 6.1 - Replace the exact match + let new_content = original_content.replacen(old_str, new_str, 1); + + // 保存历史记录 + // Requirements: 6.5 - THE File_Editor SHALL maintain a history stack for undo operations + self.save_history(&validated_path, &original_content); + + // 写入文件 + fs::write(&validated_path, &new_content).map_err(|e| { + ToolError::ExecutionFailed(format!("无法写入文件 {}: {}", path.display(), e)) + })?; + + // 生成变更上下文片段 + // Requirements: 6.6 - WHEN an edit is applied, THE File_Editor SHALL return a snippet + // showing the changed context + let change_position = original_content.find(old_str).unwrap_or(0); + let context_snippet = generate_change_context(&new_content, change_position, new_str.len()); + + // 生成 unified diff + // Requirements: 6.4 - THE File_Editor SHALL support unified diff format + let diff = generate_unified_diff(path, &original_content, &new_content); + + info!( + "[EditFileTool] 编辑文件: {} (替换 {} 字节 -> {} 字节)", + path.display(), + old_str.len(), + new_str.len() + ); + + Ok(EditFileResult { + path: validated_path, + old_str_len: old_str.len(), + new_str_len: new_str.len(), + context_snippet, + diff, + }) + } + + /// 应用 unified diff + /// + /// Requirements: 6.4 - THE File_Editor SHALL support unified diff format for multi-line changes + pub fn apply_diff(&self, path: &Path, diff: &str) -> Result { + // 验证路径安全性 + let validated_path = self + .security + .validate_path(path) + .map_err(|e| ToolError::Security(e.to_string()))?; + + // 检查文件是否存在 + if !validated_path.exists() { + return Err(ToolError::ExecutionFailed(format!( + "文件不存在: {}", + path.display() + ))); + } + + // 读取文件内容 + let original_content = fs::read_to_string(&validated_path).map_err(|e| { + ToolError::ExecutionFailed(format!("无法读取文件 {}: {}", path.display(), e)) + })?; + + // 解析并应用 diff + let new_content = apply_unified_diff(&original_content, diff)?; + + // 保存历史记录 + self.save_history(&validated_path, &original_content); + + // 写入文件 + fs::write(&validated_path, &new_content).map_err(|e| { + ToolError::ExecutionFailed(format!("无法写入文件 {}: {}", path.display(), e)) + })?; + + info!("[EditFileTool] 应用 diff: {}", path.display()); + + Ok(EditFileResult { + path: validated_path, + old_str_len: original_content.len(), + new_str_len: new_content.len(), + context_snippet: "Diff applied successfully".to_string(), + diff: diff.to_string(), + }) + } + + /// 撤销上一次编辑 + /// + /// Requirements: 6.5 - THE File_Editor SHALL maintain a history stack for undo operations + pub fn undo_edit(&self, path: &Path) -> Result { + // 验证路径安全性 + let validated_path = self + .security + .validate_path(path) + .map_err(|e| ToolError::Security(e.to_string()))?; + + // 获取历史记录 + let previous_content = { + let mut history = self.history.write(); + let file_history = history.get_mut(&validated_path).ok_or_else(|| { + ToolError::ExecutionFailed(format!("没有可撤销的编辑历史: {}", path.display())) + })?; + + file_history.pop().ok_or_else(|| { + ToolError::ExecutionFailed(format!("没有可撤销的编辑历史: {}", path.display())) + })? + }; + + // 读取当前内容 + let current_content = fs::read_to_string(&validated_path).map_err(|e| { + ToolError::ExecutionFailed(format!("无法读取文件 {}: {}", path.display(), e)) + })?; + + // 恢复之前的内容 + fs::write(&validated_path, &previous_content.content).map_err(|e| { + ToolError::ExecutionFailed(format!("无法写入文件 {}: {}", path.display(), e)) + })?; + + info!("[EditFileTool] 撤销编辑: {}", path.display()); + + Ok(UndoResult { + path: validated_path, + restored_content_len: previous_content.content.len(), + previous_content_len: current_content.len(), + }) + } + + /// 获取文件的编辑历史数量 + pub fn history_count(&self, path: &Path) -> usize { + let validated_path = match self.security.validate_path(path) { + Ok(p) => p, + Err(_) => return 0, + }; + + self.history + .read() + .get(&validated_path) + .map(|h| h.len()) + .unwrap_or(0) + } + + /// 清除文件的编辑历史 + pub fn clear_history(&self, path: &Path) { + if let Ok(validated_path) = self.security.validate_path(path) { + let mut history = self.history.write(); + history.remove(&validated_path); + debug!("[EditFileTool] 清除历史: {:?}", validated_path); + } + } + + /// 清除所有编辑历史 + pub fn clear_all_history(&self) { + let mut history = self.history.write(); + history.clear(); + debug!("[EditFileTool] 清除所有历史"); + } + + /// 保存历史记录 + fn save_history(&self, path: &PathBuf, content: &str) { + let mut history = self.history.write(); + let file_history = history.entry(path.clone()).or_insert_with(Vec::new); + + // 限制历史记录数量 + if file_history.len() >= MAX_HISTORY_SIZE { + file_history.remove(0); + } + + file_history.push(EditHistory { + content: content.to_string(), + timestamp: std::time::SystemTime::now(), + }); + } +} + +/// 文件编辑结果 +#[derive(Debug, Clone)] +pub struct EditFileResult { + /// 编辑的文件路径 + pub path: PathBuf, + /// 原字符串长度 + pub old_str_len: usize, + /// 新字符串长度 + pub new_str_len: usize, + /// 变更上下文片段 + pub context_snippet: String, + /// Unified diff + pub diff: String, +} + +/// 撤销结果 +#[derive(Debug, Clone)] +pub struct UndoResult { + /// 文件路径 + pub path: PathBuf, + /// 恢复的内容长度 + pub restored_content_len: usize, + /// 之前的内容长度 + pub previous_content_len: usize, +} + +/// 编辑历史记录 +#[derive(Debug, Clone)] +struct EditHistory { + /// 编辑前的内容 + content: String, + /// 时间戳 + timestamp: std::time::SystemTime, +} + +/// 统计字符串出现次数 +fn count_occurrences(content: &str, pattern: &str) -> usize { + if pattern.is_empty() { + return 0; + } + content.matches(pattern).count() +} + +/// 查找所有出现位置 +fn find_occurrence_positions(content: &str, pattern: &str) -> Vec { + if pattern.is_empty() { + return Vec::new(); + } + content.match_indices(pattern).map(|(pos, _)| pos).collect() +} + +/// 获取指定位置的上下文片段 +fn get_context_snippet(content: &str, position: usize, match_len: usize) -> String { + let lines: Vec<&str> = content.lines().collect(); + + // 找到位置所在的行 + let mut current_pos = 0; + let mut target_line = 0; + + for (i, line) in lines.iter().enumerate() { + let line_end = current_pos + line.len() + 1; // +1 for newline + if position < line_end { + target_line = i; + break; + } + current_pos = line_end; + } + + // 获取上下文行 + let start_line = target_line.saturating_sub(CONTEXT_LINES); + let end_line = (target_line + CONTEXT_LINES + 1).min(lines.len()); + + let mut result = String::new(); + for i in start_line..end_line { + let line_num = i + 1; + let marker = if i == target_line { ">>>" } else { " " }; + result.push_str(&format!("{} {:4} | {}\n", marker, line_num, lines[i])); + } + + result +} + +/// 生成变更上下文 +fn generate_change_context(content: &str, position: usize, new_len: usize) -> String { + get_context_snippet(content, position, new_len) +} + +/// 生成 unified diff +/// +/// Requirements: 6.4 - THE File_Editor SHALL support unified diff format +fn generate_unified_diff(path: &Path, old_content: &str, new_content: &str) -> String { + let old_lines: Vec<&str> = old_content.lines().collect(); + let new_lines: Vec<&str> = new_content.lines().collect(); + + let mut diff = String::new(); + diff.push_str(&format!("--- a/{}\n", path.display())); + diff.push_str(&format!("+++ b/{}\n", path.display())); + + // 简单的行级 diff(找出不同的行) + let mut i = 0; + let mut j = 0; + + while i < old_lines.len() || j < new_lines.len() { + if i < old_lines.len() && j < new_lines.len() && old_lines[i] == new_lines[j] { + i += 1; + j += 1; + continue; + } + + // 找到差异块的起始位置 + let context_start = i.saturating_sub(CONTEXT_LINES); + let old_start = i; + let new_start = j; + + // 找到差异块的结束位置 + let mut old_end = i; + let mut new_end = j; + + // 跳过不同的行 + while old_end < old_lines.len() && new_end < new_lines.len() { + if old_lines.get(old_end) == new_lines.get(new_end) { + // 检查是否有足够的相同行来结束差异块 + let mut same_count = 0; + while old_end + same_count < old_lines.len() + && new_end + same_count < new_lines.len() + && old_lines[old_end + same_count] == new_lines[new_end + same_count] + { + same_count += 1; + if same_count > CONTEXT_LINES * 2 { + break; + } + } + if same_count > CONTEXT_LINES * 2 { + break; + } + } + if old_end < old_lines.len() { + old_end += 1; + } + if new_end < new_lines.len() { + new_end += 1; + } + } + + // 添加上下文后的结束位置 + let context_end_old = (old_end + CONTEXT_LINES).min(old_lines.len()); + let context_end_new = (new_end + CONTEXT_LINES).min(new_lines.len()); + + // 输出 hunk header + diff.push_str(&format!( + "@@ -{},{} +{},{} @@\n", + context_start + 1, + context_end_old - context_start, + context_start + 1, + context_end_new - context_start + )); + + // 输出上下文和差异 + let mut oi = context_start; + let mut ni = context_start; + + while oi < context_end_old || ni < context_end_new { + if oi < old_start && ni < new_start { + // 前置上下文 + if oi < old_lines.len() { + diff.push_str(&format!(" {}\n", old_lines[oi])); + } + oi += 1; + ni += 1; + } else if oi < old_end { + // 删除的行 + if oi < old_lines.len() { + diff.push_str(&format!("-{}\n", old_lines[oi])); + } + oi += 1; + } else if ni < new_end { + // 添加的行 + if ni < new_lines.len() { + diff.push_str(&format!("+{}\n", new_lines[ni])); + } + ni += 1; + } else { + // 后置上下文 + if oi < old_lines.len() && ni < new_lines.len() { + diff.push_str(&format!(" {}\n", old_lines[oi])); + } + oi += 1; + ni += 1; + } + } + + i = old_end; + j = new_end; + } + + diff +} + +/// 应用 unified diff +/// +/// Requirements: 6.4 - THE File_Editor SHALL support unified diff format +fn apply_unified_diff(content: &str, diff: &str) -> Result { + let mut lines: Vec = content.lines().map(|s| s.to_string()).collect(); + let diff_lines: Vec<&str> = diff.lines().collect(); + + let mut i = 0; + while i < diff_lines.len() { + let line = diff_lines[i]; + + // 跳过文件头 + if line.starts_with("---") || line.starts_with("+++") { + i += 1; + continue; + } + + // 解析 hunk header + if line.starts_with("@@") { + // 解析 @@ -start,count +start,count @@ + let parts: Vec<&str> = line.split_whitespace().collect(); + if parts.len() < 3 { + i += 1; + continue; + } + + let old_range = parts[1].trim_start_matches('-'); + let old_start: usize = old_range + .split(',') + .next() + .and_then(|s| s.parse().ok()) + .unwrap_or(1); + + let mut current_line = old_start.saturating_sub(1); + i += 1; + + // 应用 hunk 中的变更 + while i < diff_lines.len() && !diff_lines[i].starts_with("@@") { + let diff_line = diff_lines[i]; + + if diff_line.starts_with('-') { + // 删除行 + if current_line < lines.len() { + lines.remove(current_line); + } + } else if diff_line.starts_with('+') { + // 添加行 + let new_line = diff_line.strip_prefix('+').unwrap_or(""); + lines.insert(current_line, new_line.to_string()); + current_line += 1; + } else if diff_line.starts_with(' ') || diff_line.is_empty() { + // 上下文行 + current_line += 1; + } + + i += 1; + } + } else { + i += 1; + } + } + + Ok(lines.join("\n")) +} + +#[async_trait] +impl Tool for EditFileTool { + fn definition(&self) -> ToolDefinition { + ToolDefinition::new( + "edit_file", + "Make precise edits to an existing file by replacing exact string matches. \ + The old_str must match exactly one location in the file. \ + If old_str appears multiple times, include more surrounding context to uniquely identify the location. \ + Supports undo operations to revert changes.", + ) + .with_parameters( + JsonSchema::new() + .add_property( + "path", + PropertySchema::string( + "The path to the file to edit. Can be relative or absolute.", + ), + true, + ) + .add_property( + "old_str", + PropertySchema::string( + "The exact string to find and replace. Must match exactly one location in the file.", + ), + true, + ) + .add_property( + "new_str", + PropertySchema::string( + "The string to replace old_str with.", + ), + true, + ), + ) + } + + async fn execute(&self, args: serde_json::Value) -> Result { + // 解析参数 + let path_str = args + .get("path") + .and_then(|v| v.as_str()) + .ok_or_else(|| ToolError::InvalidArguments("缺少 path 参数".to_string()))?; + + let old_str = args + .get("old_str") + .and_then(|v| v.as_str()) + .ok_or_else(|| ToolError::InvalidArguments("缺少 old_str 参数".to_string()))?; + + let new_str = args + .get("new_str") + .and_then(|v| v.as_str()) + .ok_or_else(|| ToolError::InvalidArguments("缺少 new_str 参数".to_string()))?; + + let path = PathBuf::from(path_str); + + info!( + "[EditFileTool] 编辑文件: {} (old_str: {} 字节, new_str: {} 字节)", + path_str, + old_str.len(), + new_str.len() + ); + + // 执行编辑 + let result = self.edit_file(&path, old_str, new_str)?; + + // 构建输出 + let output = format!( + "成功编辑文件: {}\n\ + 替换: {} 字节 -> {} 字节\n\n\ + 变更上下文:\n{}\n\n\ + Diff:\n{}", + path_str, result.old_str_len, result.new_str_len, result.context_snippet, result.diff + ); + + debug!("[EditFileTool] 编辑完成: {}", path_str); + + Ok(ToolResult::success(output)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::fs; + use tempfile::TempDir; + + fn setup_test_tool() -> (EditFileTool, TempDir) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + (tool, temp_dir) + } + + #[test] + fn test_tool_definition() { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + let def = tool.definition(); + + assert_eq!(def.name, "edit_file"); + assert!(!def.description.is_empty()); + assert!(def.parameters.required.contains(&"path".to_string())); + assert!(def.parameters.required.contains(&"old_str".to_string())); + assert!(def.parameters.required.contains(&"new_str".to_string())); + } + + #[test] + fn test_edit_file_simple_replacement() { + let (tool, temp_dir) = setup_test_tool(); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "Hello, World!").unwrap(); + + // 执行编辑 + let result = tool.edit_file(Path::new("test.txt"), "World", "Rust"); + assert!(result.is_ok()); + + let result = result.unwrap(); + assert_eq!(result.old_str_len, 5); + assert_eq!(result.new_str_len, 4); + + // 验证文件内容 + let content = fs::read_to_string(&file_path).unwrap(); + assert_eq!(content, "Hello, Rust!"); + } + + #[test] + fn test_edit_file_multiline() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "Line 1\nLine 2\nLine 3").unwrap(); + + // 替换多行内容 + let result = tool.edit_file(Path::new("test.txt"), "Line 2", "Modified Line"); + assert!(result.is_ok()); + + let content = fs::read_to_string(&file_path).unwrap(); + assert_eq!(content, "Line 1\nModified Line\nLine 3"); + } + + #[test] + fn test_edit_file_not_found() { + let (tool, _temp_dir) = setup_test_tool(); + + let result = tool.edit_file(Path::new("nonexistent.txt"), "old", "new"); + assert!(result.is_err()); + assert!(matches!(result, Err(ToolError::ExecutionFailed(_)))); + } + + #[test] + fn test_edit_file_old_str_not_found() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "Hello, World!").unwrap(); + + // 尝试替换不存在的字符串 + let result = tool.edit_file(Path::new("test.txt"), "NotFound", "New"); + assert!(result.is_err()); + + let err = result.unwrap_err(); + if let ToolError::ExecutionFailed(msg) = err { + assert!(msg.contains("未找到")); + } else { + panic!("Expected ExecutionFailed error"); + } + } + + #[test] + fn test_edit_file_multiple_occurrences() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "foo bar foo baz foo").unwrap(); + + // 尝试替换出现多次的字符串 + let result = tool.edit_file(Path::new("test.txt"), "foo", "qux"); + assert!(result.is_err()); + + let err = result.unwrap_err(); + if let ToolError::ExecutionFailed(msg) = err { + assert!(msg.contains("3 处匹配")); + assert!(msg.contains("更多上下文")); + } else { + panic!("Expected ExecutionFailed error"); + } + + // 验证文件未被修改 + let content = fs::read_to_string(&file_path).unwrap(); + assert_eq!(content, "foo bar foo baz foo"); + } + + #[test] + fn test_edit_file_with_context() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "foo bar foo baz foo").unwrap(); + + // 使用更多上下文来唯一标识 + let result = tool.edit_file(Path::new("test.txt"), "bar foo", "bar qux"); + assert!(result.is_ok()); + + let content = fs::read_to_string(&file_path).unwrap(); + assert_eq!(content, "foo bar qux baz foo"); + } + + #[test] + fn test_undo_edit() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + let original_content = "Hello, World!"; + fs::write(&file_path, original_content).unwrap(); + + // 执行编辑 + let result = tool.edit_file(Path::new("test.txt"), "World", "Rust"); + assert!(result.is_ok()); + + // 验证编辑后的内容 + let content = fs::read_to_string(&file_path).unwrap(); + assert_eq!(content, "Hello, Rust!"); + + // 撤销编辑 + let undo_result = tool.undo_edit(Path::new("test.txt")); + assert!(undo_result.is_ok()); + + // 验证内容已恢复 + let content = fs::read_to_string(&file_path).unwrap(); + assert_eq!(content, original_content); + } + + #[test] + fn test_undo_no_history() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "Hello, World!").unwrap(); + + // 尝试撤销(没有历史记录) + let result = tool.undo_edit(Path::new("test.txt")); + assert!(result.is_err()); + } + + #[test] + fn test_multiple_edits_and_undos() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "A B C").unwrap(); + + // 第一次编辑 + tool.edit_file(Path::new("test.txt"), "A", "X").unwrap(); + assert_eq!(fs::read_to_string(&file_path).unwrap(), "X B C"); + + // 第二次编辑 + tool.edit_file(Path::new("test.txt"), "B", "Y").unwrap(); + assert_eq!(fs::read_to_string(&file_path).unwrap(), "X Y C"); + + // 第三次编辑 + tool.edit_file(Path::new("test.txt"), "C", "Z").unwrap(); + assert_eq!(fs::read_to_string(&file_path).unwrap(), "X Y Z"); + + // 撤销第三次 + tool.undo_edit(Path::new("test.txt")).unwrap(); + assert_eq!(fs::read_to_string(&file_path).unwrap(), "X Y C"); + + // 撤销第二次 + tool.undo_edit(Path::new("test.txt")).unwrap(); + assert_eq!(fs::read_to_string(&file_path).unwrap(), "X B C"); + + // 撤销第一次 + tool.undo_edit(Path::new("test.txt")).unwrap(); + assert_eq!(fs::read_to_string(&file_path).unwrap(), "A B C"); + } + + #[test] + fn test_history_count() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "A B C D E").unwrap(); + + assert_eq!(tool.history_count(Path::new("test.txt")), 0); + + tool.edit_file(Path::new("test.txt"), "A", "X").unwrap(); + assert_eq!(tool.history_count(Path::new("test.txt")), 1); + + tool.edit_file(Path::new("test.txt"), "B", "Y").unwrap(); + assert_eq!(tool.history_count(Path::new("test.txt")), 2); + } + + #[test] + fn test_clear_history() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "A B C").unwrap(); + + tool.edit_file(Path::new("test.txt"), "A", "X").unwrap(); + tool.edit_file(Path::new("test.txt"), "B", "Y").unwrap(); + assert_eq!(tool.history_count(Path::new("test.txt")), 2); + + tool.clear_history(Path::new("test.txt")); + assert_eq!(tool.history_count(Path::new("test.txt")), 0); + } + + #[test] + fn test_security_path_traversal() { + let (tool, _temp_dir) = setup_test_tool(); + + let result = tool.edit_file(Path::new("../../../etc/passwd"), "old", "new"); + assert!(result.is_err()); + assert!(matches!(result, Err(ToolError::Security(_)))); + } + + #[test] + fn test_count_occurrences() { + assert_eq!(count_occurrences("foo bar foo baz foo", "foo"), 3); + assert_eq!(count_occurrences("hello world", "foo"), 0); + assert_eq!(count_occurrences("aaa", "a"), 3); + assert_eq!(count_occurrences("aaa", "aa"), 1); // non-overlapping + assert_eq!(count_occurrences("", "foo"), 0); + assert_eq!(count_occurrences("foo", ""), 0); + } + + #[test] + fn test_find_occurrence_positions() { + let positions = find_occurrence_positions("foo bar foo baz foo", "foo"); + assert_eq!(positions, vec![0, 8, 16]); + + let positions = find_occurrence_positions("hello world", "foo"); + assert!(positions.is_empty()); + } + + #[test] + fn test_generate_unified_diff() { + let old_content = "Line 1\nLine 2\nLine 3"; + let new_content = "Line 1\nModified\nLine 3"; + + let diff = generate_unified_diff(Path::new("test.txt"), old_content, new_content); + + assert!(diff.contains("--- a/test.txt")); + assert!(diff.contains("+++ b/test.txt")); + assert!(diff.contains("-Line 2")); + assert!(diff.contains("+Modified")); + } + + #[tokio::test] + async fn test_tool_execute() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "Hello, World!").unwrap(); + + let result = tool + .execute(serde_json::json!({ + "path": "test.txt", + "old_str": "World", + "new_str": "Rust" + })) + .await; + + assert!(result.is_ok()); + let result = result.unwrap(); + assert!(result.success); + assert!(result.output.contains("成功编辑")); + + // 验证文件已修改 + let content = fs::read_to_string(&file_path).unwrap(); + assert_eq!(content, "Hello, Rust!"); + } + + #[tokio::test] + async fn test_tool_execute_missing_args() { + let (tool, _temp_dir) = setup_test_tool(); + + // 缺少 path + let result = tool + .execute(serde_json::json!({ + "old_str": "old", + "new_str": "new" + })) + .await; + assert!(result.is_err()); + + // 缺少 old_str + let result = tool + .execute(serde_json::json!({ + "path": "test.txt", + "new_str": "new" + })) + .await; + assert!(result.is_err()); + + // 缺少 new_str + let result = tool + .execute(serde_json::json!({ + "path": "test.txt", + "old_str": "old" + })) + .await; + assert!(result.is_err()); + } + + #[test] + fn test_edit_preserves_whitespace() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, " indented\n\ttabbed\n").unwrap(); + + let result = tool.edit_file(Path::new("test.txt"), " indented", " more indented"); + assert!(result.is_ok()); + + let content = fs::read_to_string(&file_path).unwrap(); + assert_eq!(content, " more indented\n\ttabbed\n"); + } + + #[test] + fn test_edit_empty_new_str() { + let (tool, temp_dir) = setup_test_tool(); + + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, "Hello, World!").unwrap(); + + // 删除字符串(替换为空) + let result = tool.edit_file(Path::new("test.txt"), ", World", ""); + assert!(result.is_ok()); + + let content = fs::read_to_string(&file_path).unwrap(); + assert_eq!(content, "Hello!"); + } +} + +#[cfg(test)] +mod proptests { + use super::*; + use proptest::prelude::*; + use std::fs; + use tempfile::TempDir; + + /// 生成有效的文件内容(多行文本,每行唯一) + fn arb_file_content_with_unique_lines() -> impl Strategy> { + prop::collection::vec("[a-zA-Z0-9 ,.!?]{5,50}", 3..20).prop_map(|lines| { + lines + .iter() + .enumerate() + .map(|(i, content)| format!("LINE{}_{}", i + 1, content)) + .collect() + }) + } + + /// 生成有效的替换字符串 + fn arb_replacement_str() -> impl Strategy { + "[a-zA-Z0-9 ,.!?]{1,100}" + } + + /// 生成包含重复内容的文件 + fn arb_file_with_duplicates() -> impl Strategy, String)> { + ( + prop::collection::vec("[a-zA-Z0-9]{5,20}", 2..10), + "[a-zA-Z0-9]{3,10}", + ) + .prop_map(|(unique_parts, duplicate)| { + let mut lines = Vec::new(); + for (i, part) in unique_parts.iter().enumerate() { + if i % 2 == 0 { + lines.push(format!("{} {} {}", part, duplicate, part)); + } else { + lines.push(part.clone()); + } + } + (lines, duplicate) + }) + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: agent-tool-calling, Property 9: 文件编辑精确替换** + /// **Validates: Requirements 6.1** + /// + /// *For any* 文件内容和唯一出现的 old_str,edit_file 执行后 + /// 文件中应该不再包含 old_str,而包含 new_str。 + #[test] + fn prop_edit_file_exact_replacement( + lines in arb_file_content_with_unique_lines(), + new_str in arb_replacement_str() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + let content = lines.join("\n"); + fs::write(&file_path, &content).unwrap(); + + // 选择一个唯一的行作为 old_str + if lines.is_empty() { + return Ok(()); + } + let target_line_idx = 0; // 使用第一行 + let old_str = &lines[target_line_idx]; + + // 执行编辑 + let result = tool.edit_file(Path::new("test.txt"), old_str, &new_str); + + prop_assert!( + result.is_ok(), + "编辑应该成功: {:?}", + result.err() + ); + + // 读取编辑后的文件 + let edited_content = fs::read_to_string(&file_path).unwrap(); + + // 验证 old_str 不再存在 + prop_assert!( + !edited_content.contains(old_str), + "编辑后文件不应该包含 old_str: '{}'", + old_str + ); + + // 验证 new_str 存在 + prop_assert!( + edited_content.contains(&new_str), + "编辑后文件应该包含 new_str: '{}'", + new_str + ); + } + + /// **Feature: agent-tool-calling, Property 9: 文件编辑精确替换 - 其他内容不变** + /// **Validates: Requirements 6.1** + /// + /// *For any* 文件内容和唯一出现的 old_str,edit_file 执行后 + /// 除了被替换的部分,其他内容应该保持不变。 + #[test] + fn prop_edit_file_preserves_other_content( + lines in arb_file_content_with_unique_lines(), + new_str in arb_replacement_str() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + let content = lines.join("\n"); + fs::write(&file_path, &content).unwrap(); + + if lines.len() < 2 { + return Ok(()); + } + + // 选择第一行作为 old_str + let old_str = &lines[0]; + + // 执行编辑 + let result = tool.edit_file(Path::new("test.txt"), old_str, &new_str); + prop_assert!(result.is_ok()); + + // 读取编辑后的文件 + let edited_content = fs::read_to_string(&file_path).unwrap(); + + // 验证其他行仍然存在 + for (i, line) in lines.iter().enumerate() { + if i == 0 { + continue; // 跳过被替换的行 + } + prop_assert!( + edited_content.contains(line), + "编辑后文件应该保留第 {} 行: '{}'", + i + 1, + line + ); + } + } + + /// **Feature: agent-tool-calling, Property 9: 文件编辑精确替换 - 返回正确的长度** + /// **Validates: Requirements 6.1** + /// + /// *For any* 编辑操作,返回的 old_str_len 和 new_str_len 应该正确。 + #[test] + fn prop_edit_file_returns_correct_lengths( + lines in arb_file_content_with_unique_lines(), + new_str in arb_replacement_str() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + let content = lines.join("\n"); + fs::write(&file_path, &content).unwrap(); + + if lines.is_empty() { + return Ok(()); + } + + let old_str = &lines[0]; + + // 执行编辑 + let result = tool.edit_file(Path::new("test.txt"), old_str, &new_str); + prop_assert!(result.is_ok()); + + let result = result.unwrap(); + + prop_assert_eq!( + result.old_str_len, + old_str.len(), + "old_str_len 应该等于 old_str 的长度" + ); + + prop_assert_eq!( + result.new_str_len, + new_str.len(), + "new_str_len 应该等于 new_str 的长度" + ); + } + + /// **Feature: agent-tool-calling, Property 9: 文件编辑精确替换 - 生成 diff** + /// **Validates: Requirements 6.1, 6.4** + /// + /// *For any* 编辑操作,应该生成包含变更信息的 diff。 + #[test] + fn prop_edit_file_generates_diff( + lines in arb_file_content_with_unique_lines(), + new_str in arb_replacement_str() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + let content = lines.join("\n"); + fs::write(&file_path, &content).unwrap(); + + if lines.is_empty() { + return Ok(()); + } + + let old_str = &lines[0]; + + // 执行编辑 + let result = tool.edit_file(Path::new("test.txt"), old_str, &new_str); + prop_assert!(result.is_ok()); + + let result = result.unwrap(); + + // 验证 diff 包含文件头 + prop_assert!( + result.diff.contains("--- a/test.txt"), + "diff 应该包含旧文件头" + ); + prop_assert!( + result.diff.contains("+++ b/test.txt"), + "diff 应该包含新文件头" + ); + } + } +} + +#[cfg(test)] +mod proptests_multiple_occurrences { + use super::*; + use proptest::prelude::*; + use std::fs; + use tempfile::TempDir; + + /// 生成包含重复字符串的文件内容 + fn arb_content_with_duplicates() -> impl Strategy { + ( + "[a-zA-Z0-9]{3,15}", // 重复的字符串 + "[a-zA-Z0-9 ]{5,30}", // 唯一的前缀/后缀 + 2..6usize, // 重复次数 + ) + .prop_map(|(duplicate, unique, count)| { + let mut content = String::new(); + for i in 0..count { + content.push_str(&format!( + "{}_{} {} {}_{}\n", + unique, i, duplicate, unique, i + )); + } + (content, duplicate, count) + }) + } + + /// 生成替换字符串 + fn arb_replacement() -> impl Strategy { + "[a-zA-Z0-9]{5,20}" + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: agent-tool-calling, Property 10: 文件编辑多次出现错误** + /// **Validates: Requirements 6.2** + /// + /// *For any* 文件内容中出现多次的字符串 old_str, + /// edit_file 应该返回错误而不修改文件。 + #[test] + fn prop_edit_file_multiple_occurrences_error( + (content, duplicate, count) in arb_content_with_duplicates(), + new_str in arb_replacement() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, &content).unwrap(); + + // 尝试替换出现多次的字符串 + let result = tool.edit_file(Path::new("test.txt"), &duplicate, &new_str); + + // 应该返回错误 + prop_assert!( + result.is_err(), + "出现 {} 次的字符串 '{}' 应该返回错误,但结果是 {:?}", + count, + duplicate, + result + ); + + // 验证错误消息包含出现次数 + if let Err(ToolError::ExecutionFailed(msg)) = result { + prop_assert!( + msg.contains(&format!("{} 处匹配", count)), + "错误消息应该包含出现次数 '{}': {}", + count, + msg + ); + prop_assert!( + msg.contains("更多上下文"), + "错误消息应该提示需要更多上下文: {}", + msg + ); + } else { + prop_assert!(false, "应该返回 ExecutionFailed 错误"); + } + } + + /// **Feature: agent-tool-calling, Property 10: 文件编辑多次出现错误 - 文件未修改** + /// **Validates: Requirements 6.2** + /// + /// *For any* 文件内容中出现多次的字符串 old_str, + /// edit_file 返回错误后文件内容应该保持不变。 + #[test] + fn prop_edit_file_multiple_occurrences_no_modification( + (content, duplicate, _count) in arb_content_with_duplicates(), + new_str in arb_replacement() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, &content).unwrap(); + + // 尝试替换出现多次的字符串 + let _ = tool.edit_file(Path::new("test.txt"), &duplicate, &new_str); + + // 验证文件内容未被修改 + let after_content = fs::read_to_string(&file_path).unwrap(); + prop_assert_eq!( + after_content, + content, + "文件内容应该保持不变" + ); + } + + /// **Feature: agent-tool-calling, Property 10: 文件编辑多次出现错误 - 无历史记录** + /// **Validates: Requirements 6.2** + /// + /// *For any* 失败的编辑操作,不应该添加历史记录。 + #[test] + fn prop_edit_file_multiple_occurrences_no_history( + (content, duplicate, _count) in arb_content_with_duplicates(), + new_str in arb_replacement() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, &content).unwrap(); + + // 验证初始历史记录为空 + prop_assert_eq!( + tool.history_count(Path::new("test.txt")), + 0, + "初始历史记录应该为空" + ); + + // 尝试替换出现多次的字符串(应该失败) + let _ = tool.edit_file(Path::new("test.txt"), &duplicate, &new_str); + + // 验证历史记录仍然为空 + prop_assert_eq!( + tool.history_count(Path::new("test.txt")), + 0, + "失败的编辑不应该添加历史记录" + ); + } + + /// **Feature: agent-tool-calling, Property 10: 文件编辑多次出现错误 - 提供上下文示例** + /// **Validates: Requirements 6.2** + /// + /// *For any* 出现多次的字符串,错误消息应该包含匹配位置的上下文示例。 + #[test] + fn prop_edit_file_multiple_occurrences_shows_context( + (content, duplicate, _count) in arb_content_with_duplicates(), + new_str in arb_replacement() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + fs::write(&file_path, &content).unwrap(); + + // 尝试替换出现多次的字符串 + let result = tool.edit_file(Path::new("test.txt"), &duplicate, &new_str); + + // 验证错误消息包含上下文示例 + if let Err(ToolError::ExecutionFailed(msg)) = result { + prop_assert!( + msg.contains("匹配位置示例"), + "错误消息应该包含匹配位置示例: {}", + msg + ); + } + } + } +} + +#[cfg(test)] +mod proptests_undo { + use super::*; + use proptest::prelude::*; + use std::fs; + use tempfile::TempDir; + + /// 生成有效的文件内容(多行文本,每行唯一) + fn arb_file_content_unique() -> impl Strategy> { + prop::collection::vec("[a-zA-Z0-9 ,.!?]{5,50}", 3..15).prop_map(|lines| { + lines + .iter() + .enumerate() + .map(|(i, content)| format!("UNIQUE{}_{}", i + 1, content)) + .collect() + }) + } + + /// 生成有效的替换字符串 + fn arb_replacement() -> impl Strategy { + "[a-zA-Z0-9 ,.!?]{1,50}" + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: agent-tool-calling, Property 11: 文件编辑撤销 Round-Trip** + /// **Validates: Requirements 6.5** + /// + /// *For any* 成功的文件编辑操作,执行 undo_edit 后 + /// 文件内容应该恢复到编辑前的状态。 + #[test] + fn prop_edit_undo_roundtrip( + lines in arb_file_content_unique(), + new_str in arb_replacement() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + let original_content = lines.join("\n"); + fs::write(&file_path, &original_content).unwrap(); + + if lines.is_empty() { + return Ok(()); + } + + // 选择第一行作为 old_str + let old_str = &lines[0]; + + // 执行编辑 + let edit_result = tool.edit_file(Path::new("test.txt"), old_str, &new_str); + prop_assert!(edit_result.is_ok(), "编辑应该成功"); + + // 验证文件已被修改 + let edited_content = fs::read_to_string(&file_path).unwrap(); + prop_assert_ne!( + edited_content, + original_content.clone(), + "编辑后文件内容应该改变" + ); + + // 执行撤销 + let undo_result = tool.undo_edit(Path::new("test.txt")); + prop_assert!(undo_result.is_ok(), "撤销应该成功"); + + // 验证文件内容已恢复 + let restored_content = fs::read_to_string(&file_path).unwrap(); + prop_assert_eq!( + restored_content, + original_content, + "撤销后文件内容应该恢复到原始状态" + ); + } + + /// **Feature: agent-tool-calling, Property 11: 文件编辑撤销 Round-Trip - 多次编辑** + /// **Validates: Requirements 6.5** + /// + /// *For any* 多次成功的文件编辑操作,每次 undo_edit 应该 + /// 恢复到上一次编辑前的状态。 + #[test] + fn prop_edit_multiple_undo_roundtrip( + lines in arb_file_content_unique(), + new_str1 in arb_replacement(), + new_str2 in arb_replacement() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + let original_content = lines.join("\n"); + fs::write(&file_path, &original_content).unwrap(); + + if lines.len() < 2 { + return Ok(()); + } + + // 第一次编辑 + let old_str1 = &lines[0]; + let edit1_result = tool.edit_file(Path::new("test.txt"), old_str1, &new_str1); + prop_assert!(edit1_result.is_ok(), "第一次编辑应该成功"); + + let after_edit1 = fs::read_to_string(&file_path).unwrap(); + + // 第二次编辑 + let old_str2 = &lines[1]; + let edit2_result = tool.edit_file(Path::new("test.txt"), old_str2, &new_str2); + prop_assert!(edit2_result.is_ok(), "第二次编辑应该成功"); + + // 撤销第二次编辑 + let undo2_result = tool.undo_edit(Path::new("test.txt")); + prop_assert!(undo2_result.is_ok(), "撤销第二次编辑应该成功"); + + let after_undo2 = fs::read_to_string(&file_path).unwrap(); + prop_assert_eq!( + after_undo2, + after_edit1, + "撤销第二次编辑后应该恢复到第一次编辑后的状态" + ); + + // 撤销第一次编辑 + let undo1_result = tool.undo_edit(Path::new("test.txt")); + prop_assert!(undo1_result.is_ok(), "撤销第一次编辑应该成功"); + + let after_undo1 = fs::read_to_string(&file_path).unwrap(); + prop_assert_eq!( + after_undo1, + original_content.clone(), + "撤销第一次编辑后应该恢复到原始状态" + ); + } + + /// **Feature: agent-tool-calling, Property 11: 文件编辑撤销 Round-Trip - 历史记录正确** + /// **Validates: Requirements 6.5** + /// + /// *For any* 编辑操作,历史记录数量应该正确增减。 + #[test] + fn prop_edit_undo_history_count( + lines in arb_file_content_unique(), + new_str in arb_replacement() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + let content = lines.join("\n"); + fs::write(&file_path, &content).unwrap(); + + if lines.is_empty() { + return Ok(()); + } + + // 初始历史记录为 0 + prop_assert_eq!( + tool.history_count(Path::new("test.txt")), + 0, + "初始历史记录应该为 0" + ); + + // 编辑后历史记录为 1 + let old_str = &lines[0]; + tool.edit_file(Path::new("test.txt"), old_str, &new_str).unwrap(); + prop_assert_eq!( + tool.history_count(Path::new("test.txt")), + 1, + "编辑后历史记录应该为 1" + ); + + // 撤销后历史记录为 0 + tool.undo_edit(Path::new("test.txt")).unwrap(); + prop_assert_eq!( + tool.history_count(Path::new("test.txt")), + 0, + "撤销后历史记录应该为 0" + ); + } + + /// **Feature: agent-tool-calling, Property 11: 文件编辑撤销 Round-Trip - 撤销结果正确** + /// **Validates: Requirements 6.5** + /// + /// *For any* 撤销操作,返回的结果应该包含正确的长度信息。 + #[test] + fn prop_edit_undo_result_correct( + lines in arb_file_content_unique(), + new_str in arb_replacement() + ) { + let temp_dir = TempDir::new().unwrap(); + let security = Arc::new(SecurityManager::new(temp_dir.path())); + let tool = EditFileTool::new(security); + + // 创建测试文件 + let file_path = temp_dir.path().join("test.txt"); + let original_content = lines.join("\n"); + fs::write(&file_path, &original_content).unwrap(); + + if lines.is_empty() { + return Ok(()); + } + + // 执行编辑 + let old_str = &lines[0]; + tool.edit_file(Path::new("test.txt"), old_str, &new_str).unwrap(); + + let edited_content = fs::read_to_string(&file_path).unwrap(); + + // 执行撤销 + let undo_result = tool.undo_edit(Path::new("test.txt")).unwrap(); + + // 验证撤销结果 + prop_assert_eq!( + undo_result.restored_content_len, + original_content.len(), + "restored_content_len 应该等于原始内容长度" + ); + prop_assert_eq!( + undo_result.previous_content_len, + edited_content.len(), + "previous_content_len 应该等于编辑后内容长度" + ); + } + } +} diff --git a/src-tauri/src/agent/tools/mod.rs b/src-tauri/src/agent/tools/mod.rs new file mode 100644 index 000000000..d3f82fa2b --- /dev/null +++ b/src-tauri/src/agent/tools/mod.rs @@ -0,0 +1,32 @@ +//! Agent 工具系统模块 +//! +//! 提供工具定义、注册、执行的核心框架 +//! 参考 Claude Code 和 goose 项目的工具系统设计 +//! +//! ## 模块结构 +//! - `types`: 工具类型定义(ToolDefinition, ToolCall, ToolResult 等) +//! - `registry`: 工具注册表和 Tool trait +//! - `security`: 安全管理器(路径验证、符号链接检查等) +//! - `bash`: Bash 命令执行工具 +//! - `read_file`: 文件读取工具 +//! - `write_file`: 文件写入工具 +//! - `edit_file`: 文件编辑工具 +//! - `prompt`: 工具 Prompt 生成器(System Prompt 工具注入) + +pub mod bash; +pub mod edit_file; +pub mod prompt; +pub mod read_file; +pub mod registry; +pub mod security; +pub mod types; +pub mod write_file; + +pub use bash::{BashExecutionResult, BashTool, ShellType}; +pub use edit_file::{EditFileResult, EditFileTool, UndoResult}; +pub use prompt::{generate_tools_prompt, PromptFormat, ToolPromptGenerator}; +pub use read_file::{ReadFileResult, ReadFileTool}; +pub use registry::{Tool, ToolRegistry}; +pub use security::{SecurityError, SecurityManager}; +pub use types::*; +pub use write_file::{WriteFileResult, WriteFileTool}; diff --git a/src-tauri/src/agent/tools/prompt.rs b/src-tauri/src/agent/tools/prompt.rs new file mode 100644 index 000000000..91fa21b97 --- /dev/null +++ b/src-tauri/src/agent/tools/prompt.rs @@ -0,0 +1,613 @@ +//! 工具 Prompt 生成器模块 +//! +//! 提供工具定义到 System Prompt 的转换功能 +//! 符合 Requirements 2.3 - THE System_Prompt SHALL include all available tool definitions +//! +//! ## 功能 +//! - 工具定义到 XML 格式转换 +//! - 工具定义到 JSON 格式转换 +//! - System Prompt 模板生成 + +use super::types::{JsonSchema, ToolDefinition}; + +/// Prompt 输出格式 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub enum PromptFormat { + /// XML 格式(Claude 风格) + #[default] + Xml, + /// JSON 格式(OpenAI 风格) + Json, +} + +/// 工具 Prompt 生成器 +/// +/// 将工具定义转换为 LLM 可理解的 System Prompt 格式 +/// Requirements: 2.3 - THE System_Prompt SHALL include all available tool definitions +pub struct ToolPromptGenerator { + /// 输出格式 + format: PromptFormat, +} + +impl Default for ToolPromptGenerator { + fn default() -> Self { + Self::new() + } +} + +impl ToolPromptGenerator { + /// 创建新的 Prompt 生成器 + pub fn new() -> Self { + Self { + format: PromptFormat::Xml, + } + } + + /// 设置输出格式 + pub fn with_format(mut self, format: PromptFormat) -> Self { + self.format = format; + self + } + + /// 生成包含工具定义的 System Prompt + /// + /// Requirements: 2.3 - THE System_Prompt SHALL include all available tool definitions + /// in a format the LLM can understand + pub fn generate_system_prompt(&self, tools: &[ToolDefinition]) -> String { + match self.format { + PromptFormat::Xml => self.generate_xml_prompt(tools), + PromptFormat::Json => self.generate_json_prompt(tools), + } + } + + /// 生成 XML 格式的 System Prompt(Claude 风格) + fn generate_xml_prompt(&self, tools: &[ToolDefinition]) -> String { + let mut prompt = String::new(); + + // 添加工具使用说明 + prompt.push_str(TOOL_USAGE_INSTRUCTIONS); + prompt.push_str("\n\n"); + + // 添加工具定义 + prompt.push_str("\n"); + for tool in tools { + prompt.push_str(&self.tool_to_xml(tool)); + prompt.push('\n'); + } + prompt.push_str("\n"); + + prompt + } + + /// 生成 JSON 格式的 System Prompt(OpenAI 风格) + fn generate_json_prompt(&self, tools: &[ToolDefinition]) -> String { + let mut prompt = String::new(); + + // 添加工具使用说明 + prompt.push_str(TOOL_USAGE_INSTRUCTIONS); + prompt.push_str("\n\n"); + + // 添加工具定义 + prompt.push_str("Available tools:\n```json\n"); + let tools_json = serde_json::to_string_pretty(tools).unwrap_or_else(|_| "[]".to_string()); + prompt.push_str(&tools_json); + prompt.push_str("\n```\n"); + + prompt + } + + /// 将单个工具定义转换为 XML 格式 + pub fn tool_to_xml(&self, tool: &ToolDefinition) -> String { + let mut xml = String::new(); + + xml.push_str(&format!("\n", escape_xml(&tool.name))); + xml.push_str(&format!( + " {}\n", + escape_xml(&tool.description) + )); + xml.push_str(" \n"); + xml.push_str(&self.json_schema_to_xml(&tool.parameters, 4)); + xml.push_str(" \n"); + xml.push_str(""); + + xml + } + + /// 将 JsonSchema 转换为 XML 格式 + fn json_schema_to_xml(&self, schema: &JsonSchema, indent: usize) -> String { + let mut xml = String::new(); + let indent_str = " ".repeat(indent); + + for (name, prop) in &schema.properties { + let required = if schema.required.contains(name) { + " required=\"true\"" + } else { + "" + }; + + xml.push_str(&format!( + "{}\n", + indent_str, + escape_xml(name), + escape_xml(&prop.prop_type), + required + )); + xml.push_str(&format!( + "{} {}\n", + indent_str, + escape_xml(&prop.description) + )); + + // 添加默认值(如果有) + if let Some(default) = &prop.default { + xml.push_str(&format!( + "{} {}\n", + indent_str, + escape_xml(&default.to_string()) + )); + } + + // 添加枚举值(如果有) + if let Some(enum_values) = &prop.enum_values { + xml.push_str(&format!("{} \n", indent_str)); + for value in enum_values { + xml.push_str(&format!( + "{} {}\n", + indent_str, + escape_xml(&value.to_string()) + )); + } + xml.push_str(&format!("{} \n", indent_str)); + } + + xml.push_str(&format!("{}\n", indent_str)); + } + + xml + } + + /// 将单个工具定义转换为 JSON 格式 + pub fn tool_to_json(&self, tool: &ToolDefinition) -> String { + serde_json::to_string_pretty(tool).unwrap_or_else(|_| "{}".to_string()) + } + + /// 获取当前格式 + pub fn format(&self) -> PromptFormat { + self.format + } +} + +/// 工具使用说明模板 +const TOOL_USAGE_INSTRUCTIONS: &str = r#"You have access to a set of tools that you can use to help accomplish tasks. When you need to use a tool, respond with a tool call in the following format: + + +tool_name + +{ + "param1": "value1", + "param2": "value2" +} + + + +Important guidelines for tool usage: +1. Only use tools when necessary to accomplish the task +2. Provide all required parameters for each tool call +3. Wait for tool results before making additional tool calls that depend on them +4. If a tool call fails, analyze the error and try an alternative approach +5. Always explain your reasoning before and after using tools"#; + +/// XML 特殊字符转义 +fn escape_xml(s: &str) -> String { + s.replace('&', "&") + .replace('<', "<") + .replace('>', ">") + .replace('"', """) + .replace('\'', "'") +} + +/// 从 ToolRegistry 生成 System Prompt 的便捷函数 +pub fn generate_tools_prompt(tools: &[ToolDefinition], format: PromptFormat) -> String { + ToolPromptGenerator::new() + .with_format(format) + .generate_system_prompt(tools) +} + +#[cfg(test)] +mod tests { + use super::super::types::PropertySchema; + use super::*; + + fn create_test_tool() -> ToolDefinition { + ToolDefinition::new("bash", "Execute a bash command in the shell").with_parameters( + JsonSchema::new() + .add_property( + "command", + PropertySchema::string("The bash command to execute"), + true, + ) + .add_property( + "timeout", + PropertySchema::integer("Optional timeout in seconds") + .with_default(serde_json::json!(120)), + false, + ), + ) + } + + fn create_test_tools() -> Vec { + vec![ + create_test_tool(), + ToolDefinition::new("read_file", "Read the contents of a file").with_parameters( + JsonSchema::new() + .add_property( + "path", + PropertySchema::string("The file path to read"), + true, + ) + .add_property( + "start_line", + PropertySchema::integer("Starting line number (1-based)"), + false, + ) + .add_property( + "end_line", + PropertySchema::integer("Ending line number (inclusive)"), + false, + ), + ), + ] + } + + #[test] + fn test_generator_default_format() { + let generator = ToolPromptGenerator::new(); + assert_eq!(generator.format(), PromptFormat::Xml); + } + + #[test] + fn test_generator_with_format() { + let generator = ToolPromptGenerator::new().with_format(PromptFormat::Json); + assert_eq!(generator.format(), PromptFormat::Json); + } + + #[test] + fn test_tool_to_xml() { + let generator = ToolPromptGenerator::new(); + let tool = create_test_tool(); + let xml = generator.tool_to_xml(&tool); + + // 验证 XML 包含工具名称 + assert!(xml.contains("name=\"bash\"")); + // 验证 XML 包含描述 + assert!(xml.contains("Execute a bash command")); + // 验证 XML 包含必需参数 + assert!(xml.contains("name=\"command\"")); + assert!(xml.contains("required=\"true\"")); + // 验证 XML 包含可选参数 + assert!(xml.contains("name=\"timeout\"")); + // 验证 XML 包含默认值 + assert!(xml.contains("120")); + } + + #[test] + fn test_tool_to_json() { + let generator = ToolPromptGenerator::new(); + let tool = create_test_tool(); + let json = generator.tool_to_json(&tool); + + // 验证 JSON 包含工具名称 + assert!(json.contains("\"name\": \"bash\"")); + // 验证 JSON 包含描述 + assert!(json.contains("Execute a bash command")); + // 验证 JSON 包含参数 + assert!(json.contains("\"command\"")); + } + + #[test] + fn test_generate_xml_prompt() { + let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml); + let tools = create_test_tools(); + let prompt = generator.generate_system_prompt(&tools); + + // 验证包含工具使用说明 + assert!(prompt.contains("You have access to a set of tools")); + // 验证包含 tools 标签 + assert!(prompt.contains("")); + assert!(prompt.contains("")); + // 验证包含所有工具 + assert!(prompt.contains("name=\"bash\"")); + assert!(prompt.contains("name=\"read_file\"")); + } + + #[test] + fn test_generate_json_prompt() { + let generator = ToolPromptGenerator::new().with_format(PromptFormat::Json); + let tools = create_test_tools(); + let prompt = generator.generate_system_prompt(&tools); + + // 验证包含工具使用说明 + assert!(prompt.contains("You have access to a set of tools")); + // 验证包含 JSON 代码块 + assert!(prompt.contains("```json")); + // 验证包含所有工具 + assert!(prompt.contains("\"bash\"")); + assert!(prompt.contains("\"read_file\"")); + } + + #[test] + fn test_generate_tools_prompt_convenience_function() { + let tools = create_test_tools(); + + let xml_prompt = generate_tools_prompt(&tools, PromptFormat::Xml); + assert!(xml_prompt.contains("")); + + let json_prompt = generate_tools_prompt(&tools, PromptFormat::Json); + assert!(json_prompt.contains("```json")); + } + + #[test] + fn test_escape_xml() { + assert_eq!(escape_xml("hello"), "hello"); + assert_eq!(escape_xml("