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
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"name": "proxycast",
|
||||
"private": true,
|
||||
"version": "0.24.0",
|
||||
"version": "0.25.0",
|
||||
"type": "module",
|
||||
"repository": {
|
||||
"type": "git",
|
||||
|
||||
@@ -3674,7 +3674,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast"
|
||||
version = "0.24.0"
|
||||
version = "0.25.0"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"arboard",
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
@@ -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 = "-"
|
||||
@@ -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"]
|
||||
@@ -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?;
|
||||
```
|
||||
|
||||
## 更新提醒
|
||||
|
||||
@@ -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::*;
|
||||
|
||||
@@ -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<ToolResult, ToolError> {
|
||||
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 包含所有可用工具定义
|
||||
|
||||
## 更新提醒
|
||||
|
||||
任何文件变更后,请更新此文档和相关的上级文档。
|
||||
@@ -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};
|
||||
@@ -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("<tools>\n");
|
||||
for tool in tools {
|
||||
prompt.push_str(&self.tool_to_xml(tool));
|
||||
prompt.push('\n');
|
||||
}
|
||||
prompt.push_str("</tools>\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!("<tool name=\"{}\">\n", escape_xml(&tool.name)));
|
||||
xml.push_str(&format!(
|
||||
" <description>{}</description>\n",
|
||||
escape_xml(&tool.description)
|
||||
));
|
||||
xml.push_str(" <parameters>\n");
|
||||
xml.push_str(&self.json_schema_to_xml(&tool.parameters, 4));
|
||||
xml.push_str(" </parameters>\n");
|
||||
xml.push_str("</tool>");
|
||||
|
||||
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!(
|
||||
"{}<parameter name=\"{}\" type=\"{}\"{}>\n",
|
||||
indent_str,
|
||||
escape_xml(name),
|
||||
escape_xml(&prop.prop_type),
|
||||
required
|
||||
));
|
||||
xml.push_str(&format!(
|
||||
"{} <description>{}</description>\n",
|
||||
indent_str,
|
||||
escape_xml(&prop.description)
|
||||
));
|
||||
|
||||
// 添加默认值(如果有)
|
||||
if let Some(default) = &prop.default {
|
||||
xml.push_str(&format!(
|
||||
"{} <default>{}</default>\n",
|
||||
indent_str,
|
||||
escape_xml(&default.to_string())
|
||||
));
|
||||
}
|
||||
|
||||
// 添加枚举值(如果有)
|
||||
if let Some(enum_values) = &prop.enum_values {
|
||||
xml.push_str(&format!("{} <enum>\n", indent_str));
|
||||
for value in enum_values {
|
||||
xml.push_str(&format!(
|
||||
"{} <value>{}</value>\n",
|
||||
indent_str,
|
||||
escape_xml(&value.to_string())
|
||||
));
|
||||
}
|
||||
xml.push_str(&format!("{} </enum>\n", indent_str));
|
||||
}
|
||||
|
||||
xml.push_str(&format!("{}</parameter>\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_call>
|
||||
<name>tool_name</name>
|
||||
<arguments>
|
||||
{
|
||||
"param1": "value1",
|
||||
"param2": "value2"
|
||||
}
|
||||
</arguments>
|
||||
</tool_call>
|
||||
|
||||
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<ToolDefinition> {
|
||||
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("<default>120</default>"));
|
||||
}
|
||||
|
||||
#[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("<tools>"));
|
||||
assert!(prompt.contains("</tools>"));
|
||||
// 验证包含所有工具
|
||||
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("<tools>"));
|
||||
|
||||
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("<script>"), "<script>");
|
||||
assert_eq!(escape_xml("a & b"), "a & b");
|
||||
assert_eq!(escape_xml("\"quoted\""), ""quoted"");
|
||||
assert_eq!(escape_xml("it's"), "it's");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_tools() {
|
||||
let generator = ToolPromptGenerator::new();
|
||||
let prompt = generator.generate_system_prompt(&[]);
|
||||
|
||||
// 即使没有工具,也应该包含使用说明
|
||||
assert!(prompt.contains("You have access to a set of tools"));
|
||||
assert!(prompt.contains("<tools>"));
|
||||
assert!(prompt.contains("</tools>"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_with_enum_values() {
|
||||
let tool = ToolDefinition::new("select", "Select an option").with_parameters(
|
||||
JsonSchema::new().add_property(
|
||||
"choice",
|
||||
PropertySchema::string("The choice to make").with_enum(vec![
|
||||
serde_json::json!("option_a"),
|
||||
serde_json::json!("option_b"),
|
||||
]),
|
||||
true,
|
||||
),
|
||||
);
|
||||
|
||||
let generator = ToolPromptGenerator::new();
|
||||
let xml = generator.tool_to_xml(&tool);
|
||||
|
||||
// 验证包含枚举值
|
||||
assert!(xml.contains("<enum>"));
|
||||
// JSON 序列化会包含引号,所以检查转义后的值
|
||||
assert!(
|
||||
xml.contains("option_a"),
|
||||
"XML should contain option_a: {}",
|
||||
xml
|
||||
);
|
||||
assert!(
|
||||
xml.contains("option_b"),
|
||||
"XML should contain option_b: {}",
|
||||
xml
|
||||
);
|
||||
assert!(xml.contains("</enum>"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_prompt_contains_all_tool_names_and_descriptions() {
|
||||
let tools = vec![
|
||||
ToolDefinition::new("tool_a", "Description for tool A"),
|
||||
ToolDefinition::new("tool_b", "Description for tool B"),
|
||||
ToolDefinition::new("tool_c", "Description for tool C"),
|
||||
];
|
||||
|
||||
let generator = ToolPromptGenerator::new();
|
||||
let prompt = generator.generate_system_prompt(&tools);
|
||||
|
||||
// 验证所有工具名称都在 prompt 中
|
||||
for tool in &tools {
|
||||
assert!(
|
||||
prompt.contains(&tool.name),
|
||||
"Prompt should contain tool name: {}",
|
||||
tool.name
|
||||
);
|
||||
assert!(
|
||||
prompt.contains(&tool.description),
|
||||
"Prompt should contain tool description: {}",
|
||||
tool.description
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod proptests {
|
||||
use super::super::types::PropertySchema;
|
||||
use super::*;
|
||||
use proptest::prelude::*;
|
||||
|
||||
/// 生成有效的工具名称
|
||||
fn arb_valid_name() -> impl Strategy<Value = String> {
|
||||
"[a-z][a-z0-9_]{0,30}".prop_map(|s| s)
|
||||
}
|
||||
|
||||
/// 生成有效的工具描述
|
||||
fn arb_valid_description() -> impl Strategy<Value = String> {
|
||||
// 生成不包含 XML 特殊字符的描述,避免转义问题
|
||||
"[a-zA-Z0-9 ,.!?]{1,100}".prop_map(|s| s)
|
||||
}
|
||||
|
||||
/// 生成有效的属性名称
|
||||
fn arb_property_name() -> impl Strategy<Value = String> {
|
||||
"[a-z][a-z0-9_]{0,20}".prop_map(|s| s)
|
||||
}
|
||||
|
||||
/// 生成有效的 PropertySchema
|
||||
fn arb_property_schema() -> impl Strategy<Value = PropertySchema> {
|
||||
prop_oneof![
|
||||
"[a-zA-Z0-9 ]{1,50}".prop_map(|desc| PropertySchema::string(desc)),
|
||||
"[a-zA-Z0-9 ]{1,50}".prop_map(|desc| PropertySchema::number(desc)),
|
||||
"[a-zA-Z0-9 ]{1,50}".prop_map(|desc| PropertySchema::integer(desc)),
|
||||
"[a-zA-Z0-9 ]{1,50}".prop_map(|desc| PropertySchema::boolean(desc)),
|
||||
]
|
||||
}
|
||||
|
||||
/// 生成有效的 JsonSchema
|
||||
fn arb_valid_json_schema() -> impl Strategy<Value = JsonSchema> {
|
||||
prop::collection::vec(
|
||||
(arb_property_name(), arb_property_schema(), any::<bool>()),
|
||||
0..5,
|
||||
)
|
||||
.prop_map(|props| {
|
||||
let mut schema = JsonSchema::new();
|
||||
for (name, prop, required) in props {
|
||||
schema = schema.add_property(name, prop, required);
|
||||
}
|
||||
schema
|
||||
})
|
||||
}
|
||||
|
||||
/// 生成有效的 ToolDefinition
|
||||
fn arb_valid_tool_definition() -> impl Strategy<Value = ToolDefinition> {
|
||||
(
|
||||
arb_valid_name(),
|
||||
arb_valid_description(),
|
||||
arb_valid_json_schema(),
|
||||
)
|
||||
.prop_map(|(name, description, parameters)| ToolDefinition {
|
||||
name,
|
||||
description,
|
||||
parameters,
|
||||
})
|
||||
}
|
||||
|
||||
/// 生成有效的工具定义列表
|
||||
fn arb_tool_definitions() -> impl Strategy<Value = Vec<ToolDefinition>> {
|
||||
prop::collection::vec(arb_valid_tool_definition(), 0..10)
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 3: System Prompt 工具包含**
|
||||
/// **Validates: Requirements 2.3**
|
||||
///
|
||||
/// *For any* 已注册的工具集合,生成的 System Prompt 应该包含所有工具的 name 和 description。
|
||||
#[test]
|
||||
fn prop_system_prompt_contains_all_tool_names(tools in arb_tool_definitions()) {
|
||||
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml);
|
||||
let prompt = generator.generate_system_prompt(&tools);
|
||||
|
||||
// 验证所有工具名称都在 prompt 中
|
||||
for tool in &tools {
|
||||
prop_assert!(
|
||||
prompt.contains(&tool.name),
|
||||
"System Prompt 应该包含工具名称 '{}'\nPrompt:\n{}",
|
||||
tool.name,
|
||||
prompt
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 3: System Prompt 工具包含 - 描述**
|
||||
/// **Validates: Requirements 2.3**
|
||||
///
|
||||
/// *For any* 已注册的工具集合,生成的 System Prompt 应该包含所有工具的 description。
|
||||
#[test]
|
||||
fn prop_system_prompt_contains_all_tool_descriptions(tools in arb_tool_definitions()) {
|
||||
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml);
|
||||
let prompt = generator.generate_system_prompt(&tools);
|
||||
|
||||
// 验证所有工具描述都在 prompt 中
|
||||
for tool in &tools {
|
||||
prop_assert!(
|
||||
prompt.contains(&tool.description),
|
||||
"System Prompt 应该包含工具描述 '{}'\nPrompt:\n{}",
|
||||
tool.description,
|
||||
prompt
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 3: System Prompt 工具包含 - JSON 格式**
|
||||
/// **Validates: Requirements 2.3**
|
||||
///
|
||||
/// *For any* 已注册的工具集合,JSON 格式的 System Prompt 也应该包含所有工具的 name 和 description。
|
||||
#[test]
|
||||
fn prop_system_prompt_json_contains_all_tools(tools in arb_tool_definitions()) {
|
||||
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Json);
|
||||
let prompt = generator.generate_system_prompt(&tools);
|
||||
|
||||
// 验证所有工具名称和描述都在 prompt 中
|
||||
for tool in &tools {
|
||||
prop_assert!(
|
||||
prompt.contains(&tool.name),
|
||||
"JSON System Prompt 应该包含工具名称 '{}'\nPrompt:\n{}",
|
||||
tool.name,
|
||||
prompt
|
||||
);
|
||||
prop_assert!(
|
||||
prompt.contains(&tool.description),
|
||||
"JSON System Prompt 应该包含工具描述 '{}'\nPrompt:\n{}",
|
||||
tool.description,
|
||||
prompt
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 3: System Prompt 工具包含 - 格式一致性**
|
||||
/// **Validates: Requirements 2.3**
|
||||
///
|
||||
/// *For any* 已注册的工具集合,无论使用 XML 还是 JSON 格式,
|
||||
/// 生成的 System Prompt 都应该包含相同的工具信息。
|
||||
#[test]
|
||||
fn prop_system_prompt_format_consistency(tools in arb_tool_definitions()) {
|
||||
let xml_generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml);
|
||||
let json_generator = ToolPromptGenerator::new().with_format(PromptFormat::Json);
|
||||
|
||||
let xml_prompt = xml_generator.generate_system_prompt(&tools);
|
||||
let json_prompt = json_generator.generate_system_prompt(&tools);
|
||||
|
||||
// 两种格式都应该包含所有工具名称和描述
|
||||
for tool in &tools {
|
||||
prop_assert!(
|
||||
xml_prompt.contains(&tool.name) && json_prompt.contains(&tool.name),
|
||||
"两种格式都应该包含工具名称 '{}'",
|
||||
tool.name
|
||||
);
|
||||
prop_assert!(
|
||||
xml_prompt.contains(&tool.description) && json_prompt.contains(&tool.description),
|
||||
"两种格式都应该包含工具描述 '{}'",
|
||||
tool.description
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 3: System Prompt 工具包含 - 工具数量**
|
||||
/// **Validates: Requirements 2.3**
|
||||
///
|
||||
/// *For any* 已注册的工具集合,生成的 System Prompt 中工具名称出现的次数
|
||||
/// 应该至少等于工具数量(每个工具至少出现一次)。
|
||||
#[test]
|
||||
fn prop_system_prompt_tool_count(tools in arb_tool_definitions()) {
|
||||
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml);
|
||||
let prompt = generator.generate_system_prompt(&tools);
|
||||
|
||||
// 统计每个工具名称在 prompt 中出现的次数
|
||||
for tool in &tools {
|
||||
let count = prompt.matches(&tool.name).count();
|
||||
prop_assert!(
|
||||
count >= 1,
|
||||
"工具 '{}' 应该在 System Prompt 中至少出现一次,实际出现 {} 次",
|
||||
tool.name,
|
||||
count
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,390 @@
|
||||
//! 工具注册表和 Tool trait
|
||||
//!
|
||||
//! 提供工具的注册、查找和验证功能
|
||||
//! 符合 Requirements 2.2, 2.4
|
||||
|
||||
use super::types::{ToolDefinition, ToolError, ToolResult, ToolValidationError};
|
||||
use async_trait::async_trait;
|
||||
use parking_lot::RwLock;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
/// 工具 trait
|
||||
///
|
||||
/// 所有工具必须实现此 trait
|
||||
/// Requirements: 2.2 - WHEN tools are registered, THE Tool_Registry SHALL validate the tool definitions
|
||||
#[async_trait]
|
||||
pub trait Tool: Send + Sync {
|
||||
/// 获取工具定义
|
||||
fn definition(&self) -> ToolDefinition;
|
||||
|
||||
/// 执行工具
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `args` - 工具调用参数(JSON 对象)
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(ToolResult)` - 执行结果
|
||||
/// * `Err(ToolError)` - 执行错误
|
||||
async fn execute(&self, args: serde_json::Value) -> Result<ToolResult, ToolError>;
|
||||
|
||||
/// 获取工具名称(便捷方法)
|
||||
fn name(&self) -> String {
|
||||
self.definition().name
|
||||
}
|
||||
|
||||
/// 验证参数
|
||||
///
|
||||
/// 默认实现检查必需参数是否存在
|
||||
fn validate_args(&self, args: &serde_json::Value) -> Result<(), ToolError> {
|
||||
let def = self.definition();
|
||||
let obj = args
|
||||
.as_object()
|
||||
.ok_or_else(|| ToolError::InvalidArguments("参数必须是 JSON 对象".to_string()))?;
|
||||
|
||||
// 检查必需参数
|
||||
for required in &def.parameters.required {
|
||||
if !obj.contains_key(required) {
|
||||
return Err(ToolError::InvalidArguments(format!(
|
||||
"缺少必需参数: {}",
|
||||
required
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// 工具注册表
|
||||
///
|
||||
/// 管理所有已注册的工具
|
||||
/// Requirements: 2.4 - WHEN a new tool is added, THE Tool_Registry SHALL make it available to the Agent without restart
|
||||
pub struct ToolRegistry {
|
||||
tools: RwLock<HashMap<String, Arc<dyn Tool>>>,
|
||||
}
|
||||
|
||||
impl Default for ToolRegistry {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl ToolRegistry {
|
||||
/// 创建新的工具注册表
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
tools: RwLock::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 注册工具
|
||||
///
|
||||
/// Requirements: 2.2 - WHEN tools are registered, THE Tool_Registry SHALL validate the tool definitions
|
||||
pub fn register<T: Tool + 'static>(&self, tool: T) -> Result<(), ToolValidationError> {
|
||||
let definition = tool.definition();
|
||||
|
||||
// 验证工具定义
|
||||
definition.validate()?;
|
||||
|
||||
let name = definition.name.clone();
|
||||
|
||||
// 检查是否已存在
|
||||
{
|
||||
let tools = self.tools.read();
|
||||
if tools.contains_key(&name) {
|
||||
return Err(ToolValidationError::DuplicateName(name));
|
||||
}
|
||||
}
|
||||
|
||||
// 注册工具
|
||||
{
|
||||
let mut tools = self.tools.write();
|
||||
tools.insert(name.clone(), Arc::new(tool));
|
||||
}
|
||||
|
||||
info!("[ToolRegistry] 注册工具: {}", name);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 注册工具(Arc 版本)
|
||||
pub fn register_arc(&self, tool: Arc<dyn Tool>) -> Result<(), ToolValidationError> {
|
||||
let definition = tool.definition();
|
||||
|
||||
// 验证工具定义
|
||||
definition.validate()?;
|
||||
|
||||
let name = definition.name.clone();
|
||||
|
||||
// 检查是否已存在
|
||||
{
|
||||
let tools = self.tools.read();
|
||||
if tools.contains_key(&name) {
|
||||
return Err(ToolValidationError::DuplicateName(name));
|
||||
}
|
||||
}
|
||||
|
||||
// 注册工具
|
||||
{
|
||||
let mut tools = self.tools.write();
|
||||
tools.insert(name.clone(), tool);
|
||||
}
|
||||
|
||||
info!("[ToolRegistry] 注册工具: {}", name);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 注销工具
|
||||
pub fn unregister(&self, name: &str) -> bool {
|
||||
let mut tools = self.tools.write();
|
||||
let removed = tools.remove(name).is_some();
|
||||
if removed {
|
||||
info!("[ToolRegistry] 注销工具: {}", name);
|
||||
}
|
||||
removed
|
||||
}
|
||||
|
||||
/// 获取工具
|
||||
pub fn get(&self, name: &str) -> Option<Arc<dyn Tool>> {
|
||||
self.tools.read().get(name).cloned()
|
||||
}
|
||||
|
||||
/// 检查工具是否存在
|
||||
pub fn contains(&self, name: &str) -> bool {
|
||||
self.tools.read().contains_key(name)
|
||||
}
|
||||
|
||||
/// 获取所有工具定义
|
||||
pub fn list_definitions(&self) -> Vec<ToolDefinition> {
|
||||
self.tools.read().values().map(|t| t.definition()).collect()
|
||||
}
|
||||
|
||||
/// 获取所有工具名称
|
||||
pub fn list_names(&self) -> Vec<String> {
|
||||
self.tools.read().keys().cloned().collect()
|
||||
}
|
||||
|
||||
/// 获取工具数量
|
||||
pub fn len(&self) -> usize {
|
||||
self.tools.read().len()
|
||||
}
|
||||
|
||||
/// 检查是否为空
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.tools.read().is_empty()
|
||||
}
|
||||
|
||||
/// 执行工具
|
||||
///
|
||||
/// 查找并执行指定的工具
|
||||
pub async fn execute(
|
||||
&self,
|
||||
name: &str,
|
||||
args: serde_json::Value,
|
||||
) -> Result<ToolResult, ToolError> {
|
||||
let tool = self
|
||||
.get(name)
|
||||
.ok_or_else(|| ToolError::NotFound(name.to_string()))?;
|
||||
|
||||
debug!("[ToolRegistry] 执行工具: {} args={:?}", name, args);
|
||||
|
||||
// 验证参数
|
||||
tool.validate_args(&args)?;
|
||||
|
||||
// 执行工具
|
||||
let result = tool.execute(args).await?;
|
||||
|
||||
debug!(
|
||||
"[ToolRegistry] 工具执行完成: {} success={}",
|
||||
name, result.success
|
||||
);
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// 验证工具定义
|
||||
///
|
||||
/// 用于在注册前验证工具定义
|
||||
pub fn validate_definition(
|
||||
&self,
|
||||
definition: &ToolDefinition,
|
||||
) -> Result<(), ToolValidationError> {
|
||||
definition.validate()?;
|
||||
|
||||
// 检查名称是否已存在
|
||||
if self.contains(&definition.name) {
|
||||
return Err(ToolValidationError::DuplicateName(definition.name.clone()));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 清空所有工具
|
||||
pub fn clear(&self) {
|
||||
let mut tools = self.tools.write();
|
||||
let count = tools.len();
|
||||
tools.clear();
|
||||
warn!("[ToolRegistry] 清空所有工具: {} 个", count);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
/// 测试用的简单工具
|
||||
struct EchoTool;
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for EchoTool {
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
ToolDefinition::new("echo", "Echo the input message").with_parameters(
|
||||
super::super::types::JsonSchema::new().add_property(
|
||||
"message",
|
||||
super::super::types::PropertySchema::string("The message to echo"),
|
||||
true,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
async fn execute(&self, args: serde_json::Value) -> Result<ToolResult, ToolError> {
|
||||
let message = args
|
||||
.get("message")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| ToolError::InvalidArguments("缺少 message 参数".to_string()))?;
|
||||
|
||||
Ok(ToolResult::success(message))
|
||||
}
|
||||
}
|
||||
|
||||
/// 无效工具(空名称)
|
||||
struct InvalidTool;
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for InvalidTool {
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
ToolDefinition::new("", "Invalid tool with empty name")
|
||||
}
|
||||
|
||||
async fn execute(&self, _args: serde_json::Value) -> Result<ToolResult, ToolError> {
|
||||
Ok(ToolResult::success(""))
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_registry_register_and_get() {
|
||||
let registry = ToolRegistry::new();
|
||||
|
||||
// 注册工具
|
||||
assert!(registry.register(EchoTool).is_ok());
|
||||
assert_eq!(registry.len(), 1);
|
||||
|
||||
// 获取工具
|
||||
let tool = registry.get("echo");
|
||||
assert!(tool.is_some());
|
||||
assert_eq!(tool.unwrap().name(), "echo");
|
||||
|
||||
// 检查存在性
|
||||
assert!(registry.contains("echo"));
|
||||
assert!(!registry.contains("nonexistent"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_registry_reject_invalid_tool() {
|
||||
let registry = ToolRegistry::new();
|
||||
|
||||
// 注册无效工具应该失败
|
||||
let result = registry.register(InvalidTool);
|
||||
assert!(result.is_err());
|
||||
assert!(matches!(result, Err(ToolValidationError::EmptyName)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_registry_reject_duplicate() {
|
||||
let registry = ToolRegistry::new();
|
||||
|
||||
// 第一次注册成功
|
||||
assert!(registry.register(EchoTool).is_ok());
|
||||
|
||||
// 第二次注册应该失败
|
||||
let result = registry.register(EchoTool);
|
||||
assert!(result.is_err());
|
||||
assert!(matches!(result, Err(ToolValidationError::DuplicateName(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_registry_unregister() {
|
||||
let registry = ToolRegistry::new();
|
||||
|
||||
registry.register(EchoTool).unwrap();
|
||||
assert_eq!(registry.len(), 1);
|
||||
|
||||
// 注销工具
|
||||
assert!(registry.unregister("echo"));
|
||||
assert_eq!(registry.len(), 0);
|
||||
assert!(!registry.contains("echo"));
|
||||
|
||||
// 再次注销应该返回 false
|
||||
assert!(!registry.unregister("echo"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_registry_list_definitions() {
|
||||
let registry = ToolRegistry::new();
|
||||
registry.register(EchoTool).unwrap();
|
||||
|
||||
let definitions = registry.list_definitions();
|
||||
assert_eq!(definitions.len(), 1);
|
||||
assert_eq!(definitions[0].name, "echo");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_registry_execute() {
|
||||
let registry = ToolRegistry::new();
|
||||
registry.register(EchoTool).unwrap();
|
||||
|
||||
// 执行工具
|
||||
let result = registry
|
||||
.execute("echo", serde_json::json!({"message": "Hello, World!"}))
|
||||
.await;
|
||||
|
||||
assert!(result.is_ok());
|
||||
let result = result.unwrap();
|
||||
assert!(result.success);
|
||||
assert_eq!(result.output, "Hello, World!");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_registry_execute_not_found() {
|
||||
let registry = ToolRegistry::new();
|
||||
|
||||
let result = registry.execute("nonexistent", serde_json::json!({})).await;
|
||||
|
||||
assert!(result.is_err());
|
||||
assert!(matches!(result, Err(ToolError::NotFound(_))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_registry_execute_missing_required_arg() {
|
||||
let registry = ToolRegistry::new();
|
||||
registry.register(EchoTool).unwrap();
|
||||
|
||||
// 缺少必需参数
|
||||
let result = registry.execute("echo", serde_json::json!({})).await;
|
||||
|
||||
assert!(result.is_err());
|
||||
assert!(matches!(result, Err(ToolError::InvalidArguments(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_registry_clear() {
|
||||
let registry = ToolRegistry::new();
|
||||
registry.register(EchoTool).unwrap();
|
||||
assert_eq!(registry.len(), 1);
|
||||
|
||||
registry.clear();
|
||||
assert_eq!(registry.len(), 0);
|
||||
assert!(registry.is_empty());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,594 @@
|
||||
//! 安全管理器模块
|
||||
//!
|
||||
//! 提供路径验证、目录遍历防护、符号链接检查等安全功能
|
||||
//! 符合 Requirements 8.1, 8.2, 8.3, 8.5
|
||||
|
||||
use std::path::{Component, Path, PathBuf};
|
||||
use thiserror::Error;
|
||||
use tracing::{debug, warn};
|
||||
|
||||
/// 安全错误类型
|
||||
///
|
||||
/// Requirements: 8.6 - IF a security violation is detected, THEN THE Security_Manager SHALL reject the operation with a clear error
|
||||
#[derive(Debug, Error)]
|
||||
pub enum SecurityError {
|
||||
/// 路径遍历攻击
|
||||
/// Requirements: 8.1, 8.2 - THE Security_Manager SHALL validate all file paths to prevent directory traversal attacks
|
||||
#[error("路径遍历攻击: 路径 '{0}' 包含 '..' 组件")]
|
||||
PathTraversal(PathBuf),
|
||||
|
||||
/// 路径超出基础目录
|
||||
/// Requirements: 8.5 - THE Security_Manager SHALL enforce a configurable base directory for all file operations
|
||||
#[error("路径超出基础目录: '{0}' 不在允许的目录范围内")]
|
||||
OutsideBaseDir(PathBuf),
|
||||
|
||||
/// 不允许操作符号链接
|
||||
/// Requirements: 8.3 - THE Security_Manager SHALL reject operations on symlinks to prevent escape attacks
|
||||
#[error("不允许操作符号链接: '{0}'")]
|
||||
SymlinkNotAllowed(PathBuf),
|
||||
|
||||
/// 无效路径
|
||||
#[error("无效路径: {0}")]
|
||||
InvalidPath(String),
|
||||
|
||||
/// IO 错误
|
||||
#[error("IO 错误: {0}")]
|
||||
Io(#[from] std::io::Error),
|
||||
}
|
||||
|
||||
/// 安全管理器
|
||||
///
|
||||
/// 负责验证所有文件操作的安全性
|
||||
/// Requirements: 8.1, 8.2, 8.3, 8.5
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SecurityManager {
|
||||
/// 基础目录(所有文件操作必须在此目录内)
|
||||
base_dir: PathBuf,
|
||||
}
|
||||
|
||||
impl SecurityManager {
|
||||
/// 创建新的安全管理器
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `base_dir` - 基础目录,所有文件操作必须在此目录内
|
||||
pub fn new(base_dir: impl Into<PathBuf>) -> Self {
|
||||
Self {
|
||||
base_dir: base_dir.into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取基础目录
|
||||
pub fn base_dir(&self) -> &Path {
|
||||
&self.base_dir
|
||||
}
|
||||
|
||||
/// 设置基础目录
|
||||
pub fn set_base_dir(&mut self, base_dir: impl Into<PathBuf>) {
|
||||
self.base_dir = base_dir.into();
|
||||
debug!("[SecurityManager] 设置基础目录: {:?}", self.base_dir);
|
||||
}
|
||||
|
||||
/// 验证路径安全性
|
||||
///
|
||||
/// 执行以下检查:
|
||||
/// 1. 检查路径是否包含 ".." 组件(Requirements 8.2)
|
||||
/// 2. 检查路径是否为符号链接(Requirements 8.3)- 在规范化之前检查
|
||||
/// 3. 检查路径是否在基础目录内(Requirements 8.5)
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `path` - 要验证的路径(可以是相对路径或绝对路径)
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(PathBuf)` - 规范化后的安全路径
|
||||
/// * `Err(SecurityError)` - 安全错误
|
||||
pub fn validate_path(&self, path: &Path) -> Result<PathBuf, SecurityError> {
|
||||
// 1. 检查 ".." 组件
|
||||
// Requirements: 8.2 - THE Security_Manager SHALL reject paths containing ".." components
|
||||
if self.contains_parent_dir(path) {
|
||||
warn!("[SecurityManager] 检测到路径遍历攻击: {:?}", path);
|
||||
return Err(SecurityError::PathTraversal(path.to_path_buf()));
|
||||
}
|
||||
|
||||
// 2. 构建完整路径
|
||||
let full_path = if path.is_absolute() {
|
||||
path.to_path_buf()
|
||||
} else {
|
||||
self.base_dir.join(path)
|
||||
};
|
||||
|
||||
// 3. 检查符号链接(在规范化之前检查,因为规范化会解析符号链接)
|
||||
// Requirements: 8.3 - THE Security_Manager SHALL reject operations on symlinks
|
||||
self.check_symlink(&full_path)?;
|
||||
|
||||
// 4. 检查是否在基础目录内
|
||||
// Requirements: 8.5 - THE Security_Manager SHALL enforce a configurable base directory
|
||||
let validated_path = self.check_within_base_dir(&full_path)?;
|
||||
|
||||
debug!(
|
||||
"[SecurityManager] 路径验证通过: {:?} -> {:?}",
|
||||
path, validated_path
|
||||
);
|
||||
Ok(validated_path)
|
||||
}
|
||||
|
||||
/// 检查路径是否包含 ".." 组件
|
||||
///
|
||||
/// Requirements: 8.2 - THE Security_Manager SHALL reject paths containing ".." components
|
||||
fn contains_parent_dir(&self, path: &Path) -> bool {
|
||||
path.components().any(|c| matches!(c, Component::ParentDir))
|
||||
}
|
||||
|
||||
/// 检查路径是否在基础目录内
|
||||
///
|
||||
/// Requirements: 8.5 - THE Security_Manager SHALL enforce a configurable base directory
|
||||
fn check_within_base_dir(&self, path: &Path) -> Result<PathBuf, SecurityError> {
|
||||
// 尝试规范化基础目录
|
||||
let canonical_base = self.base_dir.canonicalize().map_err(|e| {
|
||||
SecurityError::InvalidPath(format!("无法规范化基础目录 {:?}: {}", self.base_dir, e))
|
||||
})?;
|
||||
|
||||
// 尝试规范化目标路径
|
||||
if path.exists() {
|
||||
// 文件存在,直接规范化
|
||||
let canonical_path = path.canonicalize()?;
|
||||
if !canonical_path.starts_with(&canonical_base) {
|
||||
return Err(SecurityError::OutsideBaseDir(path.to_path_buf()));
|
||||
}
|
||||
Ok(canonical_path)
|
||||
} else {
|
||||
// 文件不存在,检查父目录
|
||||
if let Some(parent) = path.parent() {
|
||||
if parent.as_os_str().is_empty() {
|
||||
// 父目录为空,说明是相对路径的单个文件名
|
||||
// 此时完整路径应该在基础目录内
|
||||
return Ok(path.to_path_buf());
|
||||
}
|
||||
|
||||
if parent.exists() {
|
||||
let canonical_parent = parent.canonicalize()?;
|
||||
if !canonical_parent.starts_with(&canonical_base) {
|
||||
return Err(SecurityError::OutsideBaseDir(path.to_path_buf()));
|
||||
}
|
||||
// 返回规范化的父目录 + 文件名
|
||||
if let Some(file_name) = path.file_name() {
|
||||
return Ok(canonical_parent.join(file_name));
|
||||
}
|
||||
}
|
||||
}
|
||||
// 父目录也不存在,返回原路径(后续创建时会再次验证)
|
||||
Ok(path.to_path_buf())
|
||||
}
|
||||
}
|
||||
|
||||
/// 检查路径是否为符号链接
|
||||
///
|
||||
/// Requirements: 8.3 - THE Security_Manager SHALL reject operations on symlinks
|
||||
fn check_symlink(&self, path: &Path) -> Result<(), SecurityError> {
|
||||
if path.exists() {
|
||||
let metadata = path.symlink_metadata()?;
|
||||
if metadata.is_symlink() {
|
||||
warn!("[SecurityManager] 检测到符号链接: {:?}", path);
|
||||
return Err(SecurityError::SymlinkNotAllowed(path.to_path_buf()));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 验证路径安全性(不检查符号链接)
|
||||
///
|
||||
/// 用于某些只需要检查路径遍历和基础目录的场景
|
||||
pub fn validate_path_no_symlink_check(&self, path: &Path) -> Result<PathBuf, SecurityError> {
|
||||
// 1. 检查 ".." 组件
|
||||
if self.contains_parent_dir(path) {
|
||||
return Err(SecurityError::PathTraversal(path.to_path_buf()));
|
||||
}
|
||||
|
||||
// 2. 构建完整路径
|
||||
let full_path = if path.is_absolute() {
|
||||
path.to_path_buf()
|
||||
} else {
|
||||
self.base_dir.join(path)
|
||||
};
|
||||
|
||||
// 3. 检查是否在基础目录内
|
||||
self.check_within_base_dir(&full_path)
|
||||
}
|
||||
|
||||
/// 检查路径是否安全(快速检查,不规范化)
|
||||
///
|
||||
/// 仅检查是否包含 ".." 组件,用于快速过滤明显的攻击
|
||||
pub fn quick_check(&self, path: &Path) -> bool {
|
||||
!self.contains_parent_dir(path)
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for SecurityManager {
|
||||
fn default() -> Self {
|
||||
// 默认使用当前目录作为基础目录
|
||||
Self::new(std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::fs;
|
||||
use tempfile::TempDir;
|
||||
|
||||
/// 创建测试用的临时目录结构
|
||||
fn setup_test_dir() -> TempDir {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
|
||||
// 创建一些测试文件和目录
|
||||
let test_file = temp_dir.path().join("test.txt");
|
||||
fs::write(&test_file, "test content").unwrap();
|
||||
|
||||
let sub_dir = temp_dir.path().join("subdir");
|
||||
fs::create_dir(&sub_dir).unwrap();
|
||||
|
||||
let sub_file = sub_dir.join("nested.txt");
|
||||
fs::write(&sub_file, "nested content").unwrap();
|
||||
|
||||
temp_dir
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_security_manager_creation() {
|
||||
let temp_dir = setup_test_dir();
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
|
||||
assert_eq!(security.base_dir(), temp_dir.path());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_path_within_base_dir() {
|
||||
let temp_dir = setup_test_dir();
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
|
||||
// 相对路径应该通过
|
||||
let result = security.validate_path(Path::new("test.txt"));
|
||||
assert!(result.is_ok());
|
||||
|
||||
// 子目录中的文件也应该通过
|
||||
let result = security.validate_path(Path::new("subdir/nested.txt"));
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_reject_path_traversal() {
|
||||
let temp_dir = setup_test_dir();
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
|
||||
// 包含 ".." 的路径应该被拒绝
|
||||
let result = security.validate_path(Path::new("../etc/passwd"));
|
||||
assert!(matches!(result, Err(SecurityError::PathTraversal(_))));
|
||||
|
||||
// 嵌套的 ".." 也应该被拒绝
|
||||
let result = security.validate_path(Path::new("subdir/../../etc/passwd"));
|
||||
assert!(matches!(result, Err(SecurityError::PathTraversal(_))));
|
||||
|
||||
// 中间包含 ".." 的路径也应该被拒绝
|
||||
let result = security.validate_path(Path::new("subdir/../../../etc/passwd"));
|
||||
assert!(matches!(result, Err(SecurityError::PathTraversal(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_reject_outside_base_dir() {
|
||||
let temp_dir = setup_test_dir();
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
|
||||
// 绝对路径指向基础目录外应该被拒绝
|
||||
let result = security.validate_path(Path::new("/etc/passwd"));
|
||||
assert!(matches!(result, Err(SecurityError::OutsideBaseDir(_))));
|
||||
|
||||
// 另一个临时目录也应该被拒绝
|
||||
let other_temp = TempDir::new().unwrap();
|
||||
let other_file = other_temp.path().join("other.txt");
|
||||
fs::write(&other_file, "other content").unwrap();
|
||||
|
||||
let result = security.validate_path(&other_file);
|
||||
assert!(matches!(result, Err(SecurityError::OutsideBaseDir(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[cfg(unix)]
|
||||
fn test_reject_symlink() {
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
let temp_dir = setup_test_dir();
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
|
||||
// 创建符号链接
|
||||
let link_path = temp_dir.path().join("link.txt");
|
||||
let target_path = temp_dir.path().join("test.txt");
|
||||
symlink(&target_path, &link_path).unwrap();
|
||||
|
||||
// 符号链接应该被拒绝
|
||||
let result = security.validate_path(Path::new("link.txt"));
|
||||
assert!(matches!(result, Err(SecurityError::SymlinkNotAllowed(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_quick_check() {
|
||||
let security = SecurityManager::default();
|
||||
|
||||
// 正常路径应该通过
|
||||
assert!(security.quick_check(Path::new("test.txt")));
|
||||
assert!(security.quick_check(Path::new("subdir/file.txt")));
|
||||
|
||||
// 包含 ".." 的路径应该失败
|
||||
assert!(!security.quick_check(Path::new("../test.txt")));
|
||||
assert!(!security.quick_check(Path::new("subdir/../test.txt")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_set_base_dir() {
|
||||
let temp_dir = setup_test_dir();
|
||||
let mut security = SecurityManager::default();
|
||||
|
||||
security.set_base_dir(temp_dir.path());
|
||||
assert_eq!(security.base_dir(), temp_dir.path());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_new_file_path() {
|
||||
let temp_dir = setup_test_dir();
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
|
||||
// 新文件(不存在)在基础目录内应该通过
|
||||
let result = security.validate_path(Path::new("new_file.txt"));
|
||||
assert!(result.is_ok());
|
||||
|
||||
// 新文件在子目录内也应该通过
|
||||
let result = security.validate_path(Path::new("subdir/new_file.txt"));
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_path_no_symlink_check() {
|
||||
let temp_dir = setup_test_dir();
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
|
||||
// 正常路径应该通过
|
||||
let result = security.validate_path_no_symlink_check(Path::new("test.txt"));
|
||||
assert!(result.is_ok());
|
||||
|
||||
// 包含 ".." 的路径应该被拒绝
|
||||
let result = security.validate_path_no_symlink_check(Path::new("../test.txt"));
|
||||
assert!(matches!(result, Err(SecurityError::PathTraversal(_))));
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod proptests {
|
||||
use super::*;
|
||||
use proptest::prelude::*;
|
||||
use std::fs;
|
||||
use tempfile::TempDir;
|
||||
|
||||
/// 生成有效的文件名(不包含特殊字符)
|
||||
fn arb_valid_filename() -> impl Strategy<Value = String> {
|
||||
"[a-zA-Z][a-zA-Z0-9_-]{0,20}\\.[a-z]{1,4}"
|
||||
}
|
||||
|
||||
/// 生成有效的目录名
|
||||
fn arb_valid_dirname() -> impl Strategy<Value = String> {
|
||||
"[a-zA-Z][a-zA-Z0-9_-]{0,15}"
|
||||
}
|
||||
|
||||
/// 生成包含 ".." 的路径
|
||||
fn arb_path_with_parent_dir() -> impl Strategy<Value = PathBuf> {
|
||||
prop_oneof![
|
||||
// 开头的 ..
|
||||
arb_valid_filename().prop_map(|f| PathBuf::from(format!("../{}", f))),
|
||||
// 中间的 ..
|
||||
(arb_valid_dirname(), arb_valid_filename())
|
||||
.prop_map(|(d, f)| PathBuf::from(format!("{}/../{}", d, f))),
|
||||
// 多个 ..
|
||||
arb_valid_filename().prop_map(|f| PathBuf::from(format!("../../{}", f))),
|
||||
// 嵌套的 ..
|
||||
(arb_valid_dirname(), arb_valid_filename())
|
||||
.prop_map(|(d, f)| PathBuf::from(format!("{}/../../{}", d, f))),
|
||||
]
|
||||
}
|
||||
|
||||
/// 生成不包含 ".." 的相对路径
|
||||
fn arb_safe_relative_path() -> impl Strategy<Value = PathBuf> {
|
||||
prop_oneof![
|
||||
// 单个文件名
|
||||
arb_valid_filename().prop_map(PathBuf::from),
|
||||
// 一级子目录
|
||||
(arb_valid_dirname(), arb_valid_filename())
|
||||
.prop_map(|(d, f)| PathBuf::from(format!("{}/{}", d, f))),
|
||||
// 两级子目录
|
||||
(
|
||||
arb_valid_dirname(),
|
||||
arb_valid_dirname(),
|
||||
arb_valid_filename()
|
||||
)
|
||||
.prop_map(|(d1, d2, f)| PathBuf::from(format!("{}/{}/{}", d1, d2, f))),
|
||||
]
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 14: 路径安全验证**
|
||||
/// **Validates: Requirements 8.1, 8.2, 8.5**
|
||||
///
|
||||
/// *For any* 包含 ".." 组件或指向基础目录外的路径,
|
||||
/// Security Manager 应该拒绝操作并返回安全错误。
|
||||
#[test]
|
||||
fn prop_path_traversal_rejected(path in arb_path_with_parent_dir()) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
|
||||
let result = security.validate_path(&path);
|
||||
|
||||
prop_assert!(
|
||||
matches!(result, Err(SecurityError::PathTraversal(_))),
|
||||
"包含 '..' 的路径 {:?} 应该被拒绝,但结果是 {:?}",
|
||||
path,
|
||||
result
|
||||
);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 14: 路径安全验证 - 安全路径通过**
|
||||
/// **Validates: Requirements 8.1, 8.2, 8.5**
|
||||
///
|
||||
/// *For any* 不包含 ".." 且在基础目录内的路径,
|
||||
/// Security Manager 应该允许操作。
|
||||
#[test]
|
||||
fn prop_safe_path_accepted(path in arb_safe_relative_path()) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
|
||||
// 创建必要的目录结构
|
||||
let full_path = temp_dir.path().join(&path);
|
||||
if let Some(parent) = full_path.parent() {
|
||||
let _ = fs::create_dir_all(parent);
|
||||
}
|
||||
// 创建文件
|
||||
let _ = fs::write(&full_path, "test content");
|
||||
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
let result = security.validate_path(&path);
|
||||
|
||||
prop_assert!(
|
||||
result.is_ok(),
|
||||
"安全路径 {:?} 应该通过验证,但结果是 {:?}",
|
||||
path,
|
||||
result
|
||||
);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 14: 路径安全验证 - 快速检查一致性**
|
||||
/// **Validates: Requirements 8.1, 8.2**
|
||||
///
|
||||
/// *For any* 路径,quick_check 返回 false 当且仅当路径包含 ".." 组件。
|
||||
#[test]
|
||||
fn prop_quick_check_consistency(path in arb_path_with_parent_dir()) {
|
||||
let security = SecurityManager::default();
|
||||
|
||||
prop_assert!(
|
||||
!security.quick_check(&path),
|
||||
"包含 '..' 的路径 {:?} 的 quick_check 应该返回 false",
|
||||
path
|
||||
);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 14: 路径安全验证 - 安全路径快速检查**
|
||||
/// **Validates: Requirements 8.1, 8.2**
|
||||
#[test]
|
||||
fn prop_safe_path_quick_check(path in arb_safe_relative_path()) {
|
||||
let security = SecurityManager::default();
|
||||
|
||||
prop_assert!(
|
||||
security.quick_check(&path),
|
||||
"不包含 '..' 的路径 {:?} 的 quick_check 应该返回 true",
|
||||
path
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// Property 15 符号链接拒绝测试(仅 Unix 平台)
|
||||
#[cfg(unix)]
|
||||
mod symlink_proptests {
|
||||
use super::*;
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
/// 生成符号链接测试场景
|
||||
fn arb_symlink_scenario() -> impl Strategy<Value = (String, String)> {
|
||||
// 生成目标文件名和链接文件名
|
||||
(
|
||||
"[a-zA-Z][a-zA-Z0-9_-]{0,10}\\.[a-z]{1,3}",
|
||||
"[a-zA-Z][a-zA-Z0-9_-]{0,10}_link\\.[a-z]{1,3}",
|
||||
)
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 15: 符号链接拒绝**
|
||||
/// **Validates: Requirements 8.3**
|
||||
///
|
||||
/// *For any* 指向符号链接的路径,Security Manager 应该拒绝操作并返回安全错误。
|
||||
#[test]
|
||||
fn prop_symlink_rejected((target_name, link_name) in arb_symlink_scenario()) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
|
||||
// 创建目标文件
|
||||
let target_path = temp_dir.path().join(&target_name);
|
||||
fs::write(&target_path, "target content").unwrap();
|
||||
|
||||
// 创建符号链接
|
||||
let link_path = temp_dir.path().join(&link_name);
|
||||
symlink(&target_path, &link_path).unwrap();
|
||||
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
let result = security.validate_path(Path::new(&link_name));
|
||||
|
||||
prop_assert!(
|
||||
matches!(result, Err(SecurityError::SymlinkNotAllowed(_))),
|
||||
"符号链接 {:?} 应该被拒绝,但结果是 {:?}",
|
||||
link_name,
|
||||
result
|
||||
);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 15: 符号链接拒绝 - 普通文件通过**
|
||||
/// **Validates: Requirements 8.3**
|
||||
///
|
||||
/// *For any* 普通文件(非符号链接),Security Manager 应该允许操作。
|
||||
#[test]
|
||||
fn prop_regular_file_accepted(filename in "[a-zA-Z][a-zA-Z0-9_-]{0,15}\\.[a-z]{1,4}") {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
|
||||
// 创建普通文件
|
||||
let file_path = temp_dir.path().join(&filename);
|
||||
fs::write(&file_path, "regular file content").unwrap();
|
||||
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
let result = security.validate_path(Path::new(&filename));
|
||||
|
||||
prop_assert!(
|
||||
result.is_ok(),
|
||||
"普通文件 {:?} 应该通过验证,但结果是 {:?}",
|
||||
filename,
|
||||
result
|
||||
);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 15: 符号链接拒绝 - 目录符号链接**
|
||||
/// **Validates: Requirements 8.3**
|
||||
///
|
||||
/// *For any* 指向目录的符号链接,Security Manager 应该拒绝操作。
|
||||
#[test]
|
||||
fn prop_dir_symlink_rejected(
|
||||
(dir_name, link_name) in (
|
||||
"[a-zA-Z][a-zA-Z0-9_-]{0,10}",
|
||||
"[a-zA-Z][a-zA-Z0-9_-]{0,10}_dirlink"
|
||||
)
|
||||
) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
|
||||
// 创建目标目录
|
||||
let target_dir = temp_dir.path().join(&dir_name);
|
||||
fs::create_dir(&target_dir).unwrap();
|
||||
|
||||
// 创建指向目录的符号链接
|
||||
let link_path = temp_dir.path().join(&link_name);
|
||||
symlink(&target_dir, &link_path).unwrap();
|
||||
|
||||
let security = SecurityManager::new(temp_dir.path());
|
||||
let result = security.validate_path(Path::new(&link_name));
|
||||
|
||||
prop_assert!(
|
||||
matches!(result, Err(SecurityError::SymlinkNotAllowed(_))),
|
||||
"目录符号链接 {:?} 应该被拒绝,但结果是 {:?}",
|
||||
link_name,
|
||||
result
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,554 @@
|
||||
//! 工具类型定义
|
||||
//!
|
||||
//! 定义工具系统的核心类型,包括工具定义、调用、结果和错误
|
||||
//! 符合 Requirements 2.1, 2.5
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use thiserror::Error;
|
||||
|
||||
/// 工具定义结构
|
||||
///
|
||||
/// 包含工具的名称、描述和参数 JSON Schema
|
||||
/// Requirements: 2.1 - THE Tool_Definition SHALL include name, description, and JSON Schema parameters
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolDefinition {
|
||||
/// 工具名称(唯一标识)
|
||||
pub name: String,
|
||||
/// 工具描述(供 LLM 理解)
|
||||
pub description: String,
|
||||
/// 参数 JSON Schema
|
||||
pub parameters: JsonSchema,
|
||||
}
|
||||
|
||||
impl ToolDefinition {
|
||||
/// 创建新的工具定义
|
||||
pub fn new(name: impl Into<String>, description: impl Into<String>) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
description: description.into(),
|
||||
parameters: JsonSchema::default(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 设置参数 schema
|
||||
pub fn with_parameters(mut self, parameters: JsonSchema) -> Self {
|
||||
self.parameters = parameters;
|
||||
self
|
||||
}
|
||||
|
||||
/// 验证工具定义是否有效
|
||||
pub fn validate(&self) -> Result<(), ToolValidationError> {
|
||||
if self.name.is_empty() {
|
||||
return Err(ToolValidationError::EmptyName);
|
||||
}
|
||||
if self.description.is_empty() {
|
||||
return Err(ToolValidationError::EmptyDescription);
|
||||
}
|
||||
self.parameters.validate()?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// JSON Schema 参数定义
|
||||
///
|
||||
/// 定义工具参数的类型和结构
|
||||
/// Requirements: 2.5 - THE Tool_Definition SHALL support required and optional parameters with type validation
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct JsonSchema {
|
||||
/// Schema 类型(通常为 "object")
|
||||
#[serde(rename = "type")]
|
||||
pub schema_type: String,
|
||||
/// 属性定义
|
||||
#[serde(default)]
|
||||
pub properties: HashMap<String, PropertySchema>,
|
||||
/// 必需参数列表
|
||||
#[serde(default)]
|
||||
pub required: Vec<String>,
|
||||
}
|
||||
|
||||
impl Default for JsonSchema {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
schema_type: "object".to_string(),
|
||||
properties: HashMap::new(),
|
||||
required: Vec::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl JsonSchema {
|
||||
/// 创建新的 JSON Schema
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
/// 添加属性
|
||||
pub fn add_property(
|
||||
mut self,
|
||||
name: impl Into<String>,
|
||||
prop: PropertySchema,
|
||||
required: bool,
|
||||
) -> Self {
|
||||
let name = name.into();
|
||||
if required {
|
||||
self.required.push(name.clone());
|
||||
}
|
||||
self.properties.insert(name, prop);
|
||||
self
|
||||
}
|
||||
|
||||
/// 验证 schema 是否有效
|
||||
pub fn validate(&self) -> Result<(), ToolValidationError> {
|
||||
// 检查 required 中的字段是否都在 properties 中定义
|
||||
for req in &self.required {
|
||||
if !self.properties.contains_key(req) {
|
||||
return Err(ToolValidationError::RequiredPropertyNotDefined(req.clone()));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// 属性 Schema
|
||||
///
|
||||
/// 定义单个参数的类型、描述和默认值
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct PropertySchema {
|
||||
/// 属性类型(string, number, boolean, array, object)
|
||||
#[serde(rename = "type")]
|
||||
pub prop_type: String,
|
||||
/// 属性描述
|
||||
pub description: String,
|
||||
/// 默认值(可选)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub default: Option<serde_json::Value>,
|
||||
/// 枚举值(可选,用于限制取值范围)
|
||||
#[serde(skip_serializing_if = "Option::is_none", rename = "enum")]
|
||||
pub enum_values: Option<Vec<serde_json::Value>>,
|
||||
}
|
||||
|
||||
impl PropertySchema {
|
||||
/// 创建字符串类型属性
|
||||
pub fn string(description: impl Into<String>) -> Self {
|
||||
Self {
|
||||
prop_type: "string".to_string(),
|
||||
description: description.into(),
|
||||
default: None,
|
||||
enum_values: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建数字类型属性
|
||||
pub fn number(description: impl Into<String>) -> Self {
|
||||
Self {
|
||||
prop_type: "number".to_string(),
|
||||
description: description.into(),
|
||||
default: None,
|
||||
enum_values: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建整数类型属性
|
||||
pub fn integer(description: impl Into<String>) -> Self {
|
||||
Self {
|
||||
prop_type: "integer".to_string(),
|
||||
description: description.into(),
|
||||
default: None,
|
||||
enum_values: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建布尔类型属性
|
||||
pub fn boolean(description: impl Into<String>) -> Self {
|
||||
Self {
|
||||
prop_type: "boolean".to_string(),
|
||||
description: description.into(),
|
||||
default: None,
|
||||
enum_values: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建数组类型属性
|
||||
pub fn array(description: impl Into<String>) -> Self {
|
||||
Self {
|
||||
prop_type: "array".to_string(),
|
||||
description: description.into(),
|
||||
default: None,
|
||||
enum_values: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 设置默认值
|
||||
pub fn with_default(mut self, default: serde_json::Value) -> Self {
|
||||
self.default = Some(default);
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置枚举值
|
||||
pub fn with_enum(mut self, values: Vec<serde_json::Value>) -> Self {
|
||||
self.enum_values = Some(values);
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// 工具调用请求
|
||||
///
|
||||
/// 表示 LLM 发起的工具调用
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolCall {
|
||||
/// 工具调用 ID(用于关联结果)
|
||||
pub id: String,
|
||||
/// 工具名称
|
||||
pub name: String,
|
||||
/// 调用参数(JSON 对象)
|
||||
pub arguments: serde_json::Value,
|
||||
}
|
||||
|
||||
impl ToolCall {
|
||||
/// 创建新的工具调用
|
||||
pub fn new(
|
||||
id: impl Into<String>,
|
||||
name: impl Into<String>,
|
||||
arguments: serde_json::Value,
|
||||
) -> Self {
|
||||
Self {
|
||||
id: id.into(),
|
||||
name: name.into(),
|
||||
arguments,
|
||||
}
|
||||
}
|
||||
|
||||
/// 从 JSON 字符串参数创建
|
||||
pub fn from_json_str(
|
||||
id: impl Into<String>,
|
||||
name: impl Into<String>,
|
||||
arguments_str: &str,
|
||||
) -> Result<Self, serde_json::Error> {
|
||||
let arguments: serde_json::Value = serde_json::from_str(arguments_str)?;
|
||||
Ok(Self {
|
||||
id: id.into(),
|
||||
name: name.into(),
|
||||
arguments,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// 工具执行结果
|
||||
///
|
||||
/// 表示工具执行后的返回值
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ToolResult {
|
||||
/// 是否成功
|
||||
pub success: bool,
|
||||
/// 输出内容
|
||||
pub output: String,
|
||||
/// 错误信息(如果失败)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
impl ToolResult {
|
||||
/// 创建成功结果
|
||||
pub fn success(output: impl Into<String>) -> Self {
|
||||
Self {
|
||||
success: true,
|
||||
output: output.into(),
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建失败结果
|
||||
pub fn failure(error: impl Into<String>) -> Self {
|
||||
let error_msg = error.into();
|
||||
Self {
|
||||
success: false,
|
||||
output: String::new(),
|
||||
error: Some(error_msg),
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建带输出的失败结果
|
||||
pub fn failure_with_output(output: impl Into<String>, error: impl Into<String>) -> Self {
|
||||
Self {
|
||||
success: false,
|
||||
output: output.into(),
|
||||
error: Some(error.into()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 工具错误类型
|
||||
#[derive(Debug, Error)]
|
||||
pub enum ToolError {
|
||||
/// 工具不存在
|
||||
#[error("工具不存在: {0}")]
|
||||
NotFound(String),
|
||||
|
||||
/// 参数验证失败
|
||||
#[error("参数验证失败: {0}")]
|
||||
InvalidArguments(String),
|
||||
|
||||
/// 执行失败
|
||||
#[error("执行失败: {0}")]
|
||||
ExecutionFailed(String),
|
||||
|
||||
/// 安全错误
|
||||
#[error("安全错误: {0}")]
|
||||
Security(String),
|
||||
|
||||
/// 超时
|
||||
#[error("执行超时")]
|
||||
Timeout,
|
||||
|
||||
/// IO 错误
|
||||
#[error("IO 错误: {0}")]
|
||||
Io(#[from] std::io::Error),
|
||||
|
||||
/// JSON 解析错误
|
||||
#[error("JSON 解析错误: {0}")]
|
||||
Json(#[from] serde_json::Error),
|
||||
}
|
||||
|
||||
/// 工具定义验证错误
|
||||
#[derive(Debug, Error)]
|
||||
pub enum ToolValidationError {
|
||||
/// 工具名称为空
|
||||
#[error("工具名称不能为空")]
|
||||
EmptyName,
|
||||
|
||||
/// 工具描述为空
|
||||
#[error("工具描述不能为空")]
|
||||
EmptyDescription,
|
||||
|
||||
/// 必需属性未定义
|
||||
#[error("必需属性 '{0}' 未在 properties 中定义")]
|
||||
RequiredPropertyNotDefined(String),
|
||||
|
||||
/// 重复的工具名称
|
||||
#[error("工具名称 '{0}' 已存在")]
|
||||
DuplicateName(String),
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_tool_definition_validation() {
|
||||
// 有效的工具定义
|
||||
let valid = ToolDefinition::new("bash", "Execute bash commands");
|
||||
assert!(valid.validate().is_ok());
|
||||
|
||||
// 空名称
|
||||
let empty_name = ToolDefinition::new("", "Some description");
|
||||
assert!(matches!(
|
||||
empty_name.validate(),
|
||||
Err(ToolValidationError::EmptyName)
|
||||
));
|
||||
|
||||
// 空描述
|
||||
let empty_desc = ToolDefinition::new("bash", "");
|
||||
assert!(matches!(
|
||||
empty_desc.validate(),
|
||||
Err(ToolValidationError::EmptyDescription)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_json_schema_validation() {
|
||||
// 有效的 schema
|
||||
let valid = JsonSchema::new().add_property(
|
||||
"command",
|
||||
PropertySchema::string("The command to run"),
|
||||
true,
|
||||
);
|
||||
assert!(valid.validate().is_ok());
|
||||
|
||||
// required 中有未定义的属性
|
||||
let invalid = JsonSchema {
|
||||
schema_type: "object".to_string(),
|
||||
properties: HashMap::new(),
|
||||
required: vec!["undefined_prop".to_string()],
|
||||
};
|
||||
assert!(matches!(
|
||||
invalid.validate(),
|
||||
Err(ToolValidationError::RequiredPropertyNotDefined(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_result_creation() {
|
||||
let success = ToolResult::success("Hello, World!");
|
||||
assert!(success.success);
|
||||
assert_eq!(success.output, "Hello, World!");
|
||||
assert!(success.error.is_none());
|
||||
|
||||
let failure = ToolResult::failure("Something went wrong");
|
||||
assert!(!failure.success);
|
||||
assert!(failure.output.is_empty());
|
||||
assert_eq!(failure.error, Some("Something went wrong".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_call_from_json() {
|
||||
let call = ToolCall::from_json_str("call_1", "bash", r#"{"command": "ls -la"}"#).unwrap();
|
||||
assert_eq!(call.id, "call_1");
|
||||
assert_eq!(call.name, "bash");
|
||||
assert_eq!(call.arguments["command"], "ls -la");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_property_schema_builders() {
|
||||
let string_prop = PropertySchema::string("A string value");
|
||||
assert_eq!(string_prop.prop_type, "string");
|
||||
|
||||
let number_prop =
|
||||
PropertySchema::number("A number value").with_default(serde_json::json!(0));
|
||||
assert_eq!(number_prop.prop_type, "number");
|
||||
assert_eq!(number_prop.default, Some(serde_json::json!(0)));
|
||||
|
||||
let enum_prop = PropertySchema::string("A choice")
|
||||
.with_enum(vec![serde_json::json!("a"), serde_json::json!("b")]);
|
||||
assert!(enum_prop.enum_values.is_some());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod proptests {
|
||||
use super::*;
|
||||
use proptest::prelude::*;
|
||||
|
||||
/// 生成有效的工具名称
|
||||
fn arb_valid_name() -> impl Strategy<Value = String> {
|
||||
"[a-z][a-z0-9_]{0,30}".prop_map(|s| s)
|
||||
}
|
||||
|
||||
/// 生成有效的工具描述
|
||||
fn arb_valid_description() -> impl Strategy<Value = String> {
|
||||
".{1,200}".prop_map(|s| s)
|
||||
}
|
||||
|
||||
/// 生成有效的属性名称
|
||||
fn arb_property_name() -> impl Strategy<Value = String> {
|
||||
"[a-z][a-z0-9_]{0,20}".prop_map(|s| s)
|
||||
}
|
||||
|
||||
/// 生成有效的 PropertySchema
|
||||
fn arb_property_schema() -> impl Strategy<Value = PropertySchema> {
|
||||
prop_oneof![
|
||||
".{1,50}".prop_map(|desc| PropertySchema::string(desc)),
|
||||
".{1,50}".prop_map(|desc| PropertySchema::number(desc)),
|
||||
".{1,50}".prop_map(|desc| PropertySchema::integer(desc)),
|
||||
".{1,50}".prop_map(|desc| PropertySchema::boolean(desc)),
|
||||
]
|
||||
}
|
||||
|
||||
/// 生成有效的 JsonSchema
|
||||
fn arb_valid_json_schema() -> impl Strategy<Value = JsonSchema> {
|
||||
prop::collection::vec(
|
||||
(arb_property_name(), arb_property_schema(), any::<bool>()),
|
||||
0..5,
|
||||
)
|
||||
.prop_map(|props| {
|
||||
let mut schema = JsonSchema::new();
|
||||
for (name, prop, required) in props {
|
||||
schema = schema.add_property(name, prop, required);
|
||||
}
|
||||
schema
|
||||
})
|
||||
}
|
||||
|
||||
/// 生成有效的 ToolDefinition
|
||||
fn arb_valid_tool_definition() -> impl Strategy<Value = ToolDefinition> {
|
||||
(
|
||||
arb_valid_name(),
|
||||
arb_valid_description(),
|
||||
arb_valid_json_schema(),
|
||||
)
|
||||
.prop_map(|(name, description, parameters)| ToolDefinition {
|
||||
name,
|
||||
description,
|
||||
parameters,
|
||||
})
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 2: 工具定义验证**
|
||||
/// **Validates: Requirements 2.1, 2.2**
|
||||
///
|
||||
/// *For any* 工具定义,如果缺少 name、description 或 parameters 中的任何一个字段,
|
||||
/// 注册时应该返回验证错误。
|
||||
#[test]
|
||||
fn prop_valid_tool_definition_passes_validation(def in arb_valid_tool_definition()) {
|
||||
// 有效的工具定义应该通过验证
|
||||
prop_assert!(def.validate().is_ok(), "有效的工具定义应该通过验证: {:?}", def);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 2: 工具定义验证 - 空名称**
|
||||
/// **Validates: Requirements 2.1, 2.2**
|
||||
#[test]
|
||||
fn prop_empty_name_fails_validation(description in arb_valid_description()) {
|
||||
let def = ToolDefinition::new("", description);
|
||||
prop_assert!(
|
||||
matches!(def.validate(), Err(ToolValidationError::EmptyName)),
|
||||
"空名称应该返回 EmptyName 错误"
|
||||
);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 2: 工具定义验证 - 空描述**
|
||||
/// **Validates: Requirements 2.1, 2.2**
|
||||
#[test]
|
||||
fn prop_empty_description_fails_validation(name in arb_valid_name()) {
|
||||
let def = ToolDefinition::new(name, "");
|
||||
prop_assert!(
|
||||
matches!(def.validate(), Err(ToolValidationError::EmptyDescription)),
|
||||
"空描述应该返回 EmptyDescription 错误"
|
||||
);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 2: 工具定义验证 - required 属性未定义**
|
||||
/// **Validates: Requirements 2.1, 2.2**
|
||||
#[test]
|
||||
fn prop_undefined_required_property_fails_validation(
|
||||
name in arb_valid_name(),
|
||||
description in arb_valid_description(),
|
||||
undefined_prop in arb_property_name()
|
||||
) {
|
||||
let schema = JsonSchema {
|
||||
schema_type: "object".to_string(),
|
||||
properties: HashMap::new(),
|
||||
required: vec![undefined_prop.clone()],
|
||||
};
|
||||
let def = ToolDefinition {
|
||||
name,
|
||||
description,
|
||||
parameters: schema,
|
||||
};
|
||||
prop_assert!(
|
||||
matches!(def.validate(), Err(ToolValidationError::RequiredPropertyNotDefined(_))),
|
||||
"required 中未定义的属性应该返回 RequiredPropertyNotDefined 错误"
|
||||
);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 2: 工具定义验证 - required 属性已定义**
|
||||
/// **Validates: Requirements 2.1, 2.2, 2.5**
|
||||
#[test]
|
||||
fn prop_defined_required_property_passes_validation(
|
||||
name in arb_valid_name(),
|
||||
description in arb_valid_description(),
|
||||
prop_name in arb_property_name(),
|
||||
prop_schema in arb_property_schema()
|
||||
) {
|
||||
let schema = JsonSchema::new().add_property(prop_name, prop_schema, true);
|
||||
let def = ToolDefinition {
|
||||
name,
|
||||
description,
|
||||
parameters: schema,
|
||||
};
|
||||
prop_assert!(def.validate().is_ok(), "已定义的 required 属性应该通过验证");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,802 @@
|
||||
//! 文件写入工具模块
|
||||
//!
|
||||
//! 提供文件创建和写入功能,支持父目录自动创建、换行符规范化、尾部换行符保证
|
||||
//! 符合 Requirements 5.1, 5.2, 5.3, 5.4, 5.5
|
||||
//!
|
||||
//! ## 功能
|
||||
//! - 文件创建/覆盖
|
||||
//! - 父目录自动创建
|
||||
//! - 换行符规范化(Unix: LF, Windows: CRLF)
|
||||
//! - 尾部换行符保证
|
||||
|
||||
use super::registry::Tool;
|
||||
use super::security::SecurityManager;
|
||||
use super::types::{JsonSchema, PropertySchema, ToolDefinition, ToolError, ToolResult};
|
||||
use async_trait::async_trait;
|
||||
use std::fs;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::Arc;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
/// 文件写入工具
|
||||
///
|
||||
/// 创建或覆盖文件,支持自动创建父目录和换行符规范化
|
||||
/// Requirements: 5.1, 5.2, 5.3, 5.4, 5.5
|
||||
pub struct WriteFileTool {
|
||||
/// 安全管理器
|
||||
security: Arc<SecurityManager>,
|
||||
}
|
||||
|
||||
impl WriteFileTool {
|
||||
/// 创建新的文件写入工具
|
||||
pub fn new(security: Arc<SecurityManager>) -> Self {
|
||||
Self { security }
|
||||
}
|
||||
|
||||
/// 写入文件内容
|
||||
///
|
||||
/// Requirements: 5.1 - THE File_Writer SHALL create or overwrite the file with the provided content
|
||||
/// Requirements: 5.2 - THE File_Writer SHALL create parent directories if they do not exist
|
||||
/// Requirements: 5.3 - THE File_Writer SHALL normalize line endings based on the platform
|
||||
/// Requirements: 5.4 - THE File_Writer SHALL ensure files end with a trailing newline
|
||||
pub fn write_file(&self, path: &Path, content: &str) -> Result<WriteFileResult, ToolError> {
|
||||
// 验证路径安全性(不检查符号链接,因为文件可能不存在)
|
||||
let validated_path = self
|
||||
.security
|
||||
.validate_path_no_symlink_check(path)
|
||||
.map_err(|e| ToolError::Security(e.to_string()))?;
|
||||
|
||||
// 规范化换行符
|
||||
// Requirements: 5.3 - THE File_Writer SHALL normalize line endings based on the platform
|
||||
let normalized_content = normalize_line_endings(content);
|
||||
|
||||
// 确保尾部换行符
|
||||
// Requirements: 5.4 - THE File_Writer SHALL ensure files end with a trailing newline
|
||||
let final_content = ensure_trailing_newline(&normalized_content);
|
||||
|
||||
// 创建父目录
|
||||
// Requirements: 5.2 - THE File_Writer SHALL create parent directories if they do not exist
|
||||
if let Some(parent) = validated_path.parent() {
|
||||
if !parent.exists() {
|
||||
fs::create_dir_all(parent).map_err(|e| {
|
||||
ToolError::ExecutionFailed(format!(
|
||||
"无法创建父目录 {}: {}",
|
||||
parent.display(),
|
||||
e
|
||||
))
|
||||
})?;
|
||||
debug!("[WriteFileTool] 创建父目录: {:?}", parent);
|
||||
}
|
||||
}
|
||||
|
||||
// 检查文件是否已存在(用于返回结果)
|
||||
let file_existed = validated_path.exists();
|
||||
|
||||
// 写入文件
|
||||
// Requirements: 5.1 - THE File_Writer SHALL create or overwrite the file
|
||||
// Requirements: 5.5 - IF the write operation fails, THEN THE File_Writer SHALL return a descriptive error message
|
||||
fs::write(&validated_path, &final_content).map_err(|e| {
|
||||
ToolError::ExecutionFailed(format!("无法写入文件 {}: {}", path.display(), e))
|
||||
})?;
|
||||
|
||||
let bytes_written = final_content.len();
|
||||
let line_count = final_content.lines().count();
|
||||
|
||||
info!(
|
||||
"[WriteFileTool] 写入文件: {} ({} 字节, {} 行, {})",
|
||||
path.display(),
|
||||
bytes_written,
|
||||
line_count,
|
||||
if file_existed { "覆盖" } else { "新建" }
|
||||
);
|
||||
|
||||
Ok(WriteFileResult {
|
||||
path: validated_path,
|
||||
bytes_written,
|
||||
line_count,
|
||||
created: !file_existed,
|
||||
overwritten: file_existed,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// 文件写入结果
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct WriteFileResult {
|
||||
/// 写入的文件路径
|
||||
pub path: PathBuf,
|
||||
/// 写入的字节数
|
||||
pub bytes_written: usize,
|
||||
/// 写入的行数
|
||||
pub line_count: usize,
|
||||
/// 是否为新创建的文件
|
||||
pub created: bool,
|
||||
/// 是否覆盖了已有文件
|
||||
pub overwritten: bool,
|
||||
}
|
||||
|
||||
/// 规范化换行符
|
||||
///
|
||||
/// Requirements: 5.3 - THE File_Writer SHALL normalize line endings based on the platform
|
||||
/// - Unix/macOS: LF (\n)
|
||||
/// - Windows: CRLF (\r\n)
|
||||
fn normalize_line_endings(content: &str) -> String {
|
||||
// 首先将所有换行符统一为 LF
|
||||
let unified = content
|
||||
.replace("\r\n", "\n") // CRLF -> LF
|
||||
.replace("\r", "\n"); // CR -> LF
|
||||
|
||||
// 根据平台转换
|
||||
#[cfg(windows)]
|
||||
{
|
||||
// Windows: LF -> CRLF
|
||||
unified.replace("\n", "\r\n")
|
||||
}
|
||||
|
||||
#[cfg(not(windows))]
|
||||
{
|
||||
// Unix/macOS: 保持 LF
|
||||
unified
|
||||
}
|
||||
}
|
||||
|
||||
/// 确保内容以换行符结尾
|
||||
///
|
||||
/// Requirements: 5.4 - THE File_Writer SHALL ensure files end with a trailing newline
|
||||
fn ensure_trailing_newline(content: &str) -> String {
|
||||
if content.is_empty() {
|
||||
return String::new();
|
||||
}
|
||||
|
||||
#[cfg(windows)]
|
||||
{
|
||||
if content.ends_with("\r\n") {
|
||||
content.to_string()
|
||||
} else if content.ends_with('\n') {
|
||||
// 已有 LF,转换为 CRLF
|
||||
let mut result = content[..content.len() - 1].to_string();
|
||||
result.push_str("\r\n");
|
||||
result
|
||||
} else {
|
||||
let mut result = content.to_string();
|
||||
result.push_str("\r\n");
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(windows))]
|
||||
{
|
||||
if content.ends_with('\n') {
|
||||
content.to_string()
|
||||
} else {
|
||||
let mut result = content.to_string();
|
||||
result.push('\n');
|
||||
result
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取平台的换行符
|
||||
#[allow(dead_code)]
|
||||
fn platform_line_ending() -> &'static str {
|
||||
#[cfg(windows)]
|
||||
{
|
||||
"\r\n"
|
||||
}
|
||||
|
||||
#[cfg(not(windows))]
|
||||
{
|
||||
"\n"
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for WriteFileTool {
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
ToolDefinition::new(
|
||||
"write_file",
|
||||
"Create a new file or overwrite an existing file with the provided content. \
|
||||
Automatically creates parent directories if they don't exist. \
|
||||
Line endings are normalized based on the platform (LF for Unix, CRLF for Windows). \
|
||||
Files are guaranteed to end with a trailing newline.",
|
||||
)
|
||||
.with_parameters(
|
||||
JsonSchema::new()
|
||||
.add_property(
|
||||
"path",
|
||||
PropertySchema::string(
|
||||
"The path to the file to write. Can be relative or absolute.",
|
||||
),
|
||||
true,
|
||||
)
|
||||
.add_property(
|
||||
"content",
|
||||
PropertySchema::string("The content to write to the file."),
|
||||
true,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
async fn execute(&self, args: serde_json::Value) -> Result<ToolResult, ToolError> {
|
||||
// 解析参数
|
||||
let path_str = args
|
||||
.get("path")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| ToolError::InvalidArguments("缺少 path 参数".to_string()))?;
|
||||
|
||||
let content = args
|
||||
.get("content")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| ToolError::InvalidArguments("缺少 content 参数".to_string()))?;
|
||||
|
||||
let path = PathBuf::from(path_str);
|
||||
|
||||
info!("[WriteFileTool] 写入文件: {}", path_str);
|
||||
|
||||
// 写入文件
|
||||
let result = self.write_file(&path, content)?;
|
||||
|
||||
// 构建输出
|
||||
let action = if result.created { "创建" } else { "覆盖" };
|
||||
let output = format!(
|
||||
"成功{}文件: {}\n写入 {} 字节, {} 行",
|
||||
action, path_str, result.bytes_written, result.line_count
|
||||
);
|
||||
|
||||
debug!("[WriteFileTool] 写入完成: {}", output);
|
||||
|
||||
Ok(ToolResult::success(output))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::fs;
|
||||
use tempfile::TempDir;
|
||||
|
||||
fn setup_test_tool() -> (WriteFileTool, TempDir) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let security = Arc::new(SecurityManager::new(temp_dir.path()));
|
||||
let tool = WriteFileTool::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 = WriteFileTool::new(security);
|
||||
let def = tool.definition();
|
||||
|
||||
assert_eq!(def.name, "write_file");
|
||||
assert!(!def.description.is_empty());
|
||||
assert!(def.parameters.required.contains(&"path".to_string()));
|
||||
assert!(def.parameters.required.contains(&"content".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_write_new_file() {
|
||||
let (tool, temp_dir) = setup_test_tool();
|
||||
|
||||
let result = tool.write_file(Path::new("test.txt"), "Hello, World!");
|
||||
assert!(result.is_ok());
|
||||
|
||||
let result = result.unwrap();
|
||||
assert!(result.created);
|
||||
assert!(!result.overwritten);
|
||||
|
||||
// 验证文件内容
|
||||
let file_path = temp_dir.path().join("test.txt");
|
||||
let content = fs::read_to_string(&file_path).unwrap();
|
||||
assert!(content.contains("Hello, World!"));
|
||||
assert!(content.ends_with('\n') || content.ends_with("\r\n"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_overwrite_existing_file() {
|
||||
let (tool, temp_dir) = setup_test_tool();
|
||||
|
||||
// 创建初始文件
|
||||
let file_path = temp_dir.path().join("test.txt");
|
||||
fs::write(&file_path, "Original content").unwrap();
|
||||
|
||||
// 覆盖文件
|
||||
let result = tool.write_file(Path::new("test.txt"), "New content");
|
||||
assert!(result.is_ok());
|
||||
|
||||
let result = result.unwrap();
|
||||
assert!(!result.created);
|
||||
assert!(result.overwritten);
|
||||
|
||||
// 验证文件内容已更新
|
||||
let content = fs::read_to_string(&file_path).unwrap();
|
||||
assert!(content.contains("New content"));
|
||||
assert!(!content.contains("Original"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_parent_directories() {
|
||||
let (tool, temp_dir) = setup_test_tool();
|
||||
|
||||
// 写入嵌套目录中的文件
|
||||
let result = tool.write_file(Path::new("a/b/c/test.txt"), "Nested content");
|
||||
assert!(result.is_ok());
|
||||
|
||||
// 验证目录和文件都已创建
|
||||
let file_path = temp_dir.path().join("a/b/c/test.txt");
|
||||
assert!(file_path.exists());
|
||||
let content = fs::read_to_string(&file_path).unwrap();
|
||||
assert!(content.contains("Nested content"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trailing_newline() {
|
||||
let (tool, temp_dir) = setup_test_tool();
|
||||
|
||||
// 写入不带换行符的内容
|
||||
let result = tool.write_file(Path::new("test.txt"), "No newline");
|
||||
assert!(result.is_ok());
|
||||
|
||||
// 验证文件以换行符结尾
|
||||
let file_path = temp_dir.path().join("test.txt");
|
||||
let content = fs::read_to_string(&file_path).unwrap();
|
||||
|
||||
#[cfg(windows)]
|
||||
assert!(content.ends_with("\r\n"));
|
||||
|
||||
#[cfg(not(windows))]
|
||||
assert!(content.ends_with('\n'));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_normalize_line_endings_crlf_to_lf() {
|
||||
let input = "Line 1\r\nLine 2\r\nLine 3";
|
||||
let result = normalize_line_endings(input);
|
||||
|
||||
#[cfg(windows)]
|
||||
{
|
||||
assert!(result.contains("\r\n"));
|
||||
assert!(!result.contains("\r\n\r\n")); // 不应该有双换行
|
||||
}
|
||||
|
||||
#[cfg(not(windows))]
|
||||
{
|
||||
assert!(!result.contains("\r\n"));
|
||||
assert!(result.contains('\n'));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_normalize_line_endings_mixed() {
|
||||
let input = "Line 1\r\nLine 2\nLine 3\rLine 4";
|
||||
let result = normalize_line_endings(input);
|
||||
|
||||
// 所有换行符应该被统一
|
||||
let line_count = result.lines().count();
|
||||
assert_eq!(line_count, 4);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_content() {
|
||||
let (tool, temp_dir) = setup_test_tool();
|
||||
|
||||
let result = tool.write_file(Path::new("empty.txt"), "");
|
||||
assert!(result.is_ok());
|
||||
|
||||
let file_path = temp_dir.path().join("empty.txt");
|
||||
let content = fs::read_to_string(&file_path).unwrap();
|
||||
assert!(content.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_security_path_traversal() {
|
||||
let (tool, _temp_dir) = setup_test_tool();
|
||||
|
||||
// 尝试路径遍历攻击
|
||||
let result = tool.write_file(Path::new("../../../etc/passwd"), "malicious");
|
||||
assert!(result.is_err());
|
||||
assert!(matches!(result, Err(ToolError::Security(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_bytes_and_line_count() {
|
||||
let (tool, _temp_dir) = setup_test_tool();
|
||||
|
||||
let content = "Line 1\nLine 2\nLine 3";
|
||||
let result = tool.write_file(Path::new("test.txt"), content).unwrap();
|
||||
|
||||
assert_eq!(result.line_count, 3);
|
||||
assert!(result.bytes_written > 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tool_execute() {
|
||||
let (tool, temp_dir) = setup_test_tool();
|
||||
|
||||
let result = tool
|
||||
.execute(serde_json::json!({
|
||||
"path": "test.txt",
|
||||
"content": "Hello from execute!"
|
||||
}))
|
||||
.await;
|
||||
|
||||
assert!(result.is_ok());
|
||||
let result = result.unwrap();
|
||||
assert!(result.success);
|
||||
assert!(result.output.contains("成功"));
|
||||
|
||||
// 验证文件已创建
|
||||
let file_path = temp_dir.path().join("test.txt");
|
||||
assert!(file_path.exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tool_execute_missing_path() {
|
||||
let (tool, _temp_dir) = setup_test_tool();
|
||||
|
||||
let result = tool
|
||||
.execute(serde_json::json!({
|
||||
"content": "Some content"
|
||||
}))
|
||||
.await;
|
||||
|
||||
assert!(result.is_err());
|
||||
assert!(matches!(result, Err(ToolError::InvalidArguments(_))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tool_execute_missing_content() {
|
||||
let (tool, _temp_dir) = setup_test_tool();
|
||||
|
||||
let result = tool
|
||||
.execute(serde_json::json!({
|
||||
"path": "test.txt"
|
||||
}))
|
||||
.await;
|
||||
|
||||
assert!(result.is_err());
|
||||
assert!(matches!(result, Err(ToolError::InvalidArguments(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ensure_trailing_newline() {
|
||||
// 已有换行符
|
||||
let with_newline = "content\n";
|
||||
let result = ensure_trailing_newline(with_newline);
|
||||
|
||||
#[cfg(windows)]
|
||||
assert!(result.ends_with("\r\n"));
|
||||
|
||||
#[cfg(not(windows))]
|
||||
assert!(result.ends_with('\n'));
|
||||
|
||||
// 无换行符
|
||||
let without_newline = "content";
|
||||
let result = ensure_trailing_newline(without_newline);
|
||||
|
||||
#[cfg(windows)]
|
||||
assert!(result.ends_with("\r\n"));
|
||||
|
||||
#[cfg(not(windows))]
|
||||
assert!(result.ends_with('\n'));
|
||||
|
||||
// 空内容
|
||||
let empty = "";
|
||||
let result = ensure_trailing_newline(empty);
|
||||
assert!(result.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod proptests {
|
||||
use super::*;
|
||||
use crate::agent::tools::read_file::ReadFileTool;
|
||||
use proptest::prelude::*;
|
||||
use std::fs;
|
||||
use tempfile::TempDir;
|
||||
|
||||
/// 生成有效的文件内容(多行文本)
|
||||
fn arb_file_content() -> impl Strategy<Value = String> {
|
||||
prop::collection::vec("[a-zA-Z0-9 ,.!?]{1,100}", 1..50).prop_map(|lines| lines.join("\n"))
|
||||
}
|
||||
|
||||
/// 生成有效的文件名
|
||||
fn arb_valid_filename() -> impl Strategy<Value = String> {
|
||||
"[a-zA-Z][a-zA-Z0-9_-]{0,20}\\.[a-z]{1,4}"
|
||||
}
|
||||
|
||||
/// 生成包含各种换行符的内容
|
||||
fn arb_content_with_mixed_line_endings() -> impl Strategy<Value = String> {
|
||||
prop::collection::vec("[a-zA-Z0-9 ,.!?]{1,50}", 1..20).prop_map(|lines| {
|
||||
// 随机使用不同的换行符
|
||||
let mut result = String::new();
|
||||
for (i, line) in lines.iter().enumerate() {
|
||||
result.push_str(line);
|
||||
match i % 3 {
|
||||
0 => result.push('\n'), // LF
|
||||
1 => result.push_str("\r\n"), // CRLF
|
||||
_ => result.push('\r'), // CR
|
||||
}
|
||||
}
|
||||
result
|
||||
})
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 6: 文件读写 Round-Trip**
|
||||
/// **Validates: Requirements 4.1, 5.1**
|
||||
///
|
||||
/// *For any* 有效的文件内容,使用 write_file 写入后再使用 read_file 读取,
|
||||
/// 应该得到等价的内容(考虑换行符规范化)。
|
||||
#[test]
|
||||
fn prop_file_write_read_roundtrip(content in arb_file_content()) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let security = Arc::new(SecurityManager::new(temp_dir.path()));
|
||||
let write_tool = WriteFileTool::new(security.clone());
|
||||
let read_tool = ReadFileTool::new(security);
|
||||
|
||||
let filename = "roundtrip_test.txt";
|
||||
|
||||
// 写入文件
|
||||
let write_result = write_tool.write_file(Path::new(filename), &content);
|
||||
prop_assert!(
|
||||
write_result.is_ok(),
|
||||
"写入文件应该成功: {:?}",
|
||||
write_result.err()
|
||||
);
|
||||
|
||||
// 读取文件
|
||||
let read_result = read_tool.read_file(Path::new(filename), None, None);
|
||||
prop_assert!(
|
||||
read_result.is_ok(),
|
||||
"读取文件应该成功: {:?}",
|
||||
read_result.err()
|
||||
);
|
||||
|
||||
let read_result = read_result.unwrap();
|
||||
|
||||
// 提取实际内容(去除行号格式)
|
||||
let read_lines: Vec<&str> = read_result.content
|
||||
.lines()
|
||||
.map(|line| {
|
||||
// 格式: " N | content",提取 | 后面的内容
|
||||
if let Some(pos) = line.find(" | ") {
|
||||
&line[pos + 3..]
|
||||
} else {
|
||||
line
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
// 规范化原始内容进行比较
|
||||
let normalized_original = normalize_line_endings(&content);
|
||||
let original_lines: Vec<&str> = normalized_original.lines().collect();
|
||||
|
||||
// 比较行数
|
||||
prop_assert_eq!(
|
||||
read_lines.len(),
|
||||
original_lines.len(),
|
||||
"读取的行数应该与写入的行数相同"
|
||||
);
|
||||
|
||||
// 比较每一行的内容
|
||||
for (i, (read_line, orig_line)) in read_lines.iter().zip(original_lines.iter()).enumerate() {
|
||||
prop_assert_eq!(
|
||||
*read_line,
|
||||
*orig_line,
|
||||
"第 {} 行内容应该匹配",
|
||||
i + 1
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 6: 文件读写 Round-Trip - 字节级验证**
|
||||
/// **Validates: Requirements 4.1, 5.1**
|
||||
///
|
||||
/// *For any* 有效的文件内容,写入后直接读取文件字节,
|
||||
/// 应该得到规范化后的内容加上尾部换行符。
|
||||
#[test]
|
||||
fn prop_file_write_read_bytes_roundtrip(content in arb_file_content()) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let security = Arc::new(SecurityManager::new(temp_dir.path()));
|
||||
let write_tool = WriteFileTool::new(security);
|
||||
|
||||
let file_path = temp_dir.path().join("bytes_test.txt");
|
||||
|
||||
// 写入文件
|
||||
let write_result = write_tool.write_file(Path::new("bytes_test.txt"), &content);
|
||||
prop_assert!(write_result.is_ok());
|
||||
|
||||
// 直接读取文件字节
|
||||
let read_bytes = fs::read_to_string(&file_path).unwrap();
|
||||
|
||||
// 计算预期内容
|
||||
let normalized = normalize_line_endings(&content);
|
||||
let expected = ensure_trailing_newline(&normalized);
|
||||
|
||||
prop_assert_eq!(
|
||||
read_bytes,
|
||||
expected,
|
||||
"文件内容应该是规范化后的内容加尾部换行符"
|
||||
);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 6: 文件读写 Round-Trip - 多次写入**
|
||||
/// **Validates: Requirements 5.1**
|
||||
///
|
||||
/// *For any* 两个不同的内容,第二次写入应该完全覆盖第一次的内容。
|
||||
#[test]
|
||||
fn prop_file_overwrite_roundtrip(
|
||||
content1 in arb_file_content(),
|
||||
content2 in arb_file_content()
|
||||
) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let security = Arc::new(SecurityManager::new(temp_dir.path()));
|
||||
let write_tool = WriteFileTool::new(security);
|
||||
|
||||
let file_path = temp_dir.path().join("overwrite_test.txt");
|
||||
|
||||
// 第一次写入
|
||||
let result1 = write_tool.write_file(Path::new("overwrite_test.txt"), &content1);
|
||||
prop_assert!(result1.is_ok());
|
||||
prop_assert!(result1.unwrap().created);
|
||||
|
||||
// 第二次写入(覆盖)
|
||||
let result2 = write_tool.write_file(Path::new("overwrite_test.txt"), &content2);
|
||||
prop_assert!(result2.is_ok());
|
||||
prop_assert!(result2.unwrap().overwritten);
|
||||
|
||||
// 验证文件内容是第二次写入的内容
|
||||
let read_bytes = fs::read_to_string(&file_path).unwrap();
|
||||
let expected = ensure_trailing_newline(&normalize_line_endings(&content2));
|
||||
|
||||
prop_assert_eq!(
|
||||
read_bytes,
|
||||
expected,
|
||||
"文件内容应该是第二次写入的内容"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 8: 文件换行符规范化**
|
||||
/// **Validates: Requirements 5.3, 5.4**
|
||||
///
|
||||
/// *For any* 写入的文件内容,最终文件应该以换行符结尾,
|
||||
/// 且换行符符合平台规范(Unix: LF, Windows: CRLF)。
|
||||
#[test]
|
||||
fn prop_file_trailing_newline(content in arb_file_content()) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let security = Arc::new(SecurityManager::new(temp_dir.path()));
|
||||
let write_tool = WriteFileTool::new(security);
|
||||
|
||||
let file_path = temp_dir.path().join("newline_test.txt");
|
||||
|
||||
// 写入文件
|
||||
let result = write_tool.write_file(Path::new("newline_test.txt"), &content);
|
||||
prop_assert!(result.is_ok());
|
||||
|
||||
// 读取文件内容
|
||||
let read_bytes = fs::read_to_string(&file_path).unwrap();
|
||||
|
||||
// 验证尾部换行符
|
||||
if !content.is_empty() {
|
||||
#[cfg(windows)]
|
||||
prop_assert!(
|
||||
read_bytes.ends_with("\r\n"),
|
||||
"Windows 平台文件应该以 CRLF 结尾"
|
||||
);
|
||||
|
||||
#[cfg(not(windows))]
|
||||
prop_assert!(
|
||||
read_bytes.ends_with('\n'),
|
||||
"Unix 平台文件应该以 LF 结尾"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 8: 文件换行符规范化 - 混合换行符**
|
||||
/// **Validates: Requirements 5.3**
|
||||
///
|
||||
/// *For any* 包含混合换行符(LF, CRLF, CR)的内容,
|
||||
/// 写入后所有换行符应该被统一为平台规范的换行符。
|
||||
#[test]
|
||||
fn prop_file_normalize_mixed_line_endings(content in arb_content_with_mixed_line_endings()) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let security = Arc::new(SecurityManager::new(temp_dir.path()));
|
||||
let write_tool = WriteFileTool::new(security);
|
||||
|
||||
let file_path = temp_dir.path().join("mixed_newline_test.txt");
|
||||
|
||||
// 写入文件
|
||||
let result = write_tool.write_file(Path::new("mixed_newline_test.txt"), &content);
|
||||
prop_assert!(result.is_ok());
|
||||
|
||||
// 读取文件内容
|
||||
let read_bytes = fs::read_to_string(&file_path).unwrap();
|
||||
|
||||
// 验证换行符已被规范化
|
||||
#[cfg(windows)]
|
||||
{
|
||||
// Windows: 不应该有单独的 LF 或 CR
|
||||
let without_crlf = read_bytes.replace("\r\n", "");
|
||||
prop_assert!(
|
||||
!without_crlf.contains('\n') && !without_crlf.contains('\r'),
|
||||
"Windows 平台所有换行符应该是 CRLF"
|
||||
);
|
||||
}
|
||||
|
||||
#[cfg(not(windows))]
|
||||
{
|
||||
// Unix: 不应该有 CR
|
||||
prop_assert!(
|
||||
!read_bytes.contains('\r'),
|
||||
"Unix 平台不应该有 CR 字符"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 8: 文件换行符规范化 - 行数保持**
|
||||
/// **Validates: Requirements 5.3**
|
||||
///
|
||||
/// *For any* 内容,规范化换行符后行数应该保持不变。
|
||||
#[test]
|
||||
fn prop_file_line_count_preserved(content in arb_content_with_mixed_line_endings()) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let security = Arc::new(SecurityManager::new(temp_dir.path()));
|
||||
let write_tool = WriteFileTool::new(security);
|
||||
|
||||
let file_path = temp_dir.path().join("line_count_test.txt");
|
||||
|
||||
// 计算原始行数(统一换行符后)
|
||||
let unified = content
|
||||
.replace("\r\n", "\n")
|
||||
.replace('\r', "\n");
|
||||
let original_line_count = unified.lines().count();
|
||||
|
||||
// 写入文件
|
||||
let result = write_tool.write_file(Path::new("line_count_test.txt"), &content);
|
||||
prop_assert!(result.is_ok());
|
||||
|
||||
// 读取文件并计算行数
|
||||
let read_bytes = fs::read_to_string(&file_path).unwrap();
|
||||
let read_line_count = read_bytes.lines().count();
|
||||
|
||||
prop_assert_eq!(
|
||||
read_line_count,
|
||||
original_line_count,
|
||||
"规范化后行数应该保持不变"
|
||||
);
|
||||
}
|
||||
|
||||
/// **Feature: agent-tool-calling, Property 8: 文件换行符规范化 - 空内容**
|
||||
/// **Validates: Requirements 5.4**
|
||||
///
|
||||
/// *For any* 空内容,写入后文件应该为空(不添加换行符)。
|
||||
#[test]
|
||||
fn prop_empty_content_no_newline(_dummy in Just(())) {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let security = Arc::new(SecurityManager::new(temp_dir.path()));
|
||||
let write_tool = WriteFileTool::new(security);
|
||||
|
||||
let file_path = temp_dir.path().join("empty_test.txt");
|
||||
|
||||
// 写入空内容
|
||||
let result = write_tool.write_file(Path::new("empty_test.txt"), "");
|
||||
prop_assert!(result.is_ok());
|
||||
|
||||
// 读取文件内容
|
||||
let read_bytes = fs::read_to_string(&file_path).unwrap();
|
||||
|
||||
prop_assert!(
|
||||
read_bytes.is_empty(),
|
||||
"空内容写入后文件应该为空"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -193,7 +193,10 @@ pub struct NativeChatResponse {
|
||||
}
|
||||
|
||||
/// Token 使用量
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
///
|
||||
/// 记录 API 调用的 token 消耗
|
||||
/// Requirements: 1.3 - THE Streaming_Handler SHALL emit a done event with token usage statistics
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct TokenUsage {
|
||||
/// 输入 token 数
|
||||
pub input_tokens: u32,
|
||||
@@ -201,17 +204,151 @@ pub struct TokenUsage {
|
||||
pub output_tokens: u32,
|
||||
}
|
||||
|
||||
impl TokenUsage {
|
||||
/// 创建新的 TokenUsage
|
||||
pub fn new(input_tokens: u32, output_tokens: u32) -> Self {
|
||||
Self {
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
}
|
||||
}
|
||||
|
||||
/// 计算总 token 数
|
||||
pub fn total(&self) -> u32 {
|
||||
self.input_tokens + self.output_tokens
|
||||
}
|
||||
}
|
||||
|
||||
/// 流式响应事件
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
///
|
||||
/// 定义流式输出过程中的各种事件类型
|
||||
/// Requirements: 1.1, 1.3, 1.4
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum StreamEvent {
|
||||
/// 文本增量
|
||||
/// Requirements: 1.1 - THE Streaming_Handler SHALL emit text deltas to the frontend in real-time
|
||||
#[serde(rename = "text_delta")]
|
||||
TextDelta { text: String },
|
||||
|
||||
/// 工具调用开始
|
||||
/// Requirements: 7.6 - WHILE the Tool_Loop is executing, THE Frontend SHALL display the current tool being executed
|
||||
#[serde(rename = "tool_start")]
|
||||
ToolStart {
|
||||
/// 工具名称
|
||||
tool_name: String,
|
||||
/// 工具调用 ID
|
||||
tool_id: String,
|
||||
},
|
||||
|
||||
/// 工具调用结束
|
||||
/// Requirements: 7.6 - 工具执行完成后通知前端
|
||||
#[serde(rename = "tool_end")]
|
||||
ToolEnd {
|
||||
/// 工具调用 ID
|
||||
tool_id: String,
|
||||
/// 工具执行结果
|
||||
result: ToolExecutionResult,
|
||||
},
|
||||
|
||||
/// 完成
|
||||
/// Requirements: 1.3 - THE Streaming_Handler SHALL emit a done event with token usage statistics
|
||||
#[serde(rename = "done")]
|
||||
Done { usage: Option<TokenUsage> },
|
||||
|
||||
/// 错误
|
||||
/// Requirements: 1.4 - IF a streaming error occurs, THEN THE Streaming_Handler SHALL emit an error event
|
||||
#[serde(rename = "error")]
|
||||
Error { message: String },
|
||||
}
|
||||
|
||||
/// 工具执行结果(用于 StreamEvent)
|
||||
///
|
||||
/// 简化版的工具结果,用于前端显示
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ToolExecutionResult {
|
||||
/// 是否成功
|
||||
pub success: bool,
|
||||
/// 输出内容
|
||||
pub output: String,
|
||||
/// 错误信息(如果失败)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
impl ToolExecutionResult {
|
||||
/// 创建成功结果
|
||||
pub fn success(output: impl Into<String>) -> Self {
|
||||
Self {
|
||||
success: true,
|
||||
output: output.into(),
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建失败结果
|
||||
pub fn failure(error: impl Into<String>) -> Self {
|
||||
let error_msg = error.into();
|
||||
Self {
|
||||
success: false,
|
||||
output: String::new(),
|
||||
error: Some(error_msg),
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建带输出的失败结果
|
||||
pub fn failure_with_output(output: impl Into<String>, error: impl Into<String>) -> Self {
|
||||
Self {
|
||||
success: false,
|
||||
output: output.into(),
|
||||
error: Some(error.into()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 流式响应结果
|
||||
///
|
||||
/// 流式处理完成后的最终结果
|
||||
/// Requirements: 1.1, 1.3
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct StreamResult {
|
||||
/// 完整的响应内容
|
||||
pub content: String,
|
||||
/// 工具调用列表(如果有)
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_calls: Option<Vec<ToolCall>>,
|
||||
/// Token 使用量
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub usage: Option<TokenUsage>,
|
||||
}
|
||||
|
||||
impl StreamResult {
|
||||
/// 创建新的流式结果
|
||||
pub fn new(content: String) -> Self {
|
||||
Self {
|
||||
content,
|
||||
tool_calls: None,
|
||||
usage: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 设置工具调用
|
||||
pub fn with_tool_calls(mut self, tool_calls: Vec<ToolCall>) -> Self {
|
||||
self.tool_calls = Some(tool_calls);
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置 token 使用量
|
||||
pub fn with_usage(mut self, usage: TokenUsage) -> Self {
|
||||
self.usage = Some(usage);
|
||||
self
|
||||
}
|
||||
|
||||
/// 是否有工具调用
|
||||
pub fn has_tool_calls(&self) -> bool {
|
||||
self.tool_calls
|
||||
.as_ref()
|
||||
.map(|tc| !tc.is_empty())
|
||||
.unwrap_or(false)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"$schema": "https://schema.tauri.app/config/2",
|
||||
"productName": "ProxyCast",
|
||||
"version": "0.24.0",
|
||||
"version": "0.25.0",
|
||||
"identifier": "com.proxycast.app",
|
||||
"build": {
|
||||
"beforeDevCommand": "npm run dev",
|
||||
|
||||
|
After Width: | Height: | Size: 930 B |
|
After Width: | Height: | Size: 1.2 KiB |
|
After Width: | Height: | Size: 8.6 KiB |
|
After Width: | Height: | Size: 1.2 KiB |
|
After Width: | Height: | Size: 807 B |
|
After Width: | Height: | Size: 1.7 KiB |
|
After Width: | Height: | Size: 1.2 KiB |
|
After Width: | Height: | Size: 890 B |
|
After Width: | Height: | Size: 4.3 KiB |
|
After Width: | Height: | Size: 32 KiB |
|
After Width: | Height: | Size: 5.8 KiB |
|
After Width: | Height: | Size: 7.8 KiB |
|
After Width: | Height: | Size: 4.5 KiB |
|
After Width: | Height: | Size: 4.0 KiB |
@@ -1,7 +1,6 @@
|
||||
import React, { useState } from "react";
|
||||
import { Bot, ChevronDown, Check, Box, Settings2 } from "lucide-react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import {
|
||||
Popover,
|
||||
PopoverContent,
|
||||
@@ -28,7 +27,7 @@ export const ChatNavbar: React.FC<ChatNavbarProps> = ({
|
||||
setProviderType,
|
||||
model,
|
||||
setModel,
|
||||
isRunning,
|
||||
isRunning: _isRunning,
|
||||
onToggleHistory,
|
||||
onToggleFullscreen: _onToggleFullscreen,
|
||||
onToggleSettings,
|
||||
@@ -150,13 +149,6 @@ export const ChatNavbar: React.FC<ChatNavbarProps> = ({
|
||||
|
||||
{/* Right: Status & Settings */}
|
||||
<div className="flex items-center gap-2">
|
||||
<Badge
|
||||
variant={isRunning ? "default" : "secondary"}
|
||||
className="h-5 text-[10px] px-1.5"
|
||||
>
|
||||
{isRunning ? "Ready" : "Offline"}
|
||||
</Badge>
|
||||
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
|
||||
@@ -1,5 +1,53 @@
|
||||
import React from "react";
|
||||
import styled from "styled-components";
|
||||
import React, { useState } from "react";
|
||||
import styled, { keyframes, css } from "styled-components";
|
||||
import {
|
||||
Sparkles,
|
||||
ArrowRight,
|
||||
ImageIcon,
|
||||
Video,
|
||||
FileText,
|
||||
PenTool,
|
||||
BrainCircuit,
|
||||
CalendarRange,
|
||||
ChevronDown,
|
||||
Search,
|
||||
Globe,
|
||||
} from "lucide-react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
import {
|
||||
Select,
|
||||
SelectContent,
|
||||
SelectItem,
|
||||
SelectTrigger,
|
||||
SelectValue,
|
||||
} from "@/components/ui/select";
|
||||
import {
|
||||
Popover,
|
||||
PopoverContent,
|
||||
PopoverTrigger,
|
||||
} from "@/components/ui/popover";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
|
||||
// Import Assets
|
||||
import iconXhs from "@/assets/platforms/xhs.png";
|
||||
import iconGzh from "@/assets/platforms/gzh.png";
|
||||
import iconZhihu from "@/assets/platforms/zhihu.png";
|
||||
import iconToutiao from "@/assets/platforms/toutiao.png";
|
||||
import iconJuejin from "@/assets/platforms/juejin.png";
|
||||
import iconCsdn from "@/assets/platforms/csdn.png";
|
||||
|
||||
import modelGemini from "@/assets/models/gemini.png";
|
||||
import modelClaude from "@/assets/models/claude.png";
|
||||
import modelDeepseek from "@/assets/models/deepseek.png";
|
||||
|
||||
// --- Animations ---
|
||||
const fadeIn = keyframes`
|
||||
from { opacity: 0; transform: translateY(10px); }
|
||||
to { opacity: 1; transform: translateY(0); }
|
||||
`;
|
||||
|
||||
// --- Styled Components ---
|
||||
|
||||
const Container = styled.div`
|
||||
display: flex;
|
||||
@@ -7,69 +55,692 @@ const Container = styled.div`
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
flex: 1;
|
||||
padding: 40px;
|
||||
color: hsl(var(--muted-foreground));
|
||||
`;
|
||||
padding: 40px 20px;
|
||||
background-color: hsl(var(--background));
|
||||
overflow-y: auto;
|
||||
position: relative;
|
||||
|
||||
const Logo = styled.img`
|
||||
width: 80px;
|
||||
height: 80px;
|
||||
margin-bottom: 24px;
|
||||
`;
|
||||
|
||||
const Title = styled.h2`
|
||||
font-size: 20px;
|
||||
font-weight: 600;
|
||||
color: hsl(var(--foreground));
|
||||
margin: 0 0 8px 0;
|
||||
`;
|
||||
|
||||
const Description = styled.p`
|
||||
font-size: 14px;
|
||||
color: hsl(var(--muted-foreground));
|
||||
margin: 0;
|
||||
text-align: center;
|
||||
max-width: 300px;
|
||||
`;
|
||||
|
||||
const Tips = styled.div`
|
||||
margin-top: 32px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 8px;
|
||||
`;
|
||||
|
||||
const Tip = styled.div`
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
font-size: 13px;
|
||||
color: hsl(var(--muted-foreground));
|
||||
|
||||
kbd {
|
||||
padding: 2px 6px;
|
||||
border-radius: 4px;
|
||||
background-color: hsl(var(--muted));
|
||||
border: 1px solid hsl(var(--border));
|
||||
font-size: 11px;
|
||||
font-family: monospace;
|
||||
// Subtle mesh background effect
|
||||
&::before {
|
||||
content: "";
|
||||
position: absolute;
|
||||
top: -10%;
|
||||
left: 20%;
|
||||
width: 600px;
|
||||
height: 600px;
|
||||
background: radial-gradient(
|
||||
circle,
|
||||
hsl(var(--primary) / 0.05) 0%,
|
||||
transparent 70%
|
||||
);
|
||||
border-radius: 50%;
|
||||
pointer-events: none;
|
||||
z-index: 0;
|
||||
}
|
||||
`;
|
||||
|
||||
export const EmptyState: React.FC = () => {
|
||||
const ContentWrapper = styled.div`
|
||||
max-width: 900px;
|
||||
width: 100%;
|
||||
position: relative;
|
||||
z-index: 1;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
gap: 36px;
|
||||
animation: ${fadeIn} 0.5s ease-out;
|
||||
`;
|
||||
|
||||
const Header = styled.div`
|
||||
text-align: center;
|
||||
margin-bottom: 8px;
|
||||
`;
|
||||
|
||||
const shimmer = keyframes`
|
||||
0% { background-position: 0% 50%; filter: brightness(100%); }
|
||||
50% { background-position: 100% 50%; filter: brightness(120%); }
|
||||
100% { background-position: 0% 50%; filter: brightness(100%); }
|
||||
`;
|
||||
|
||||
const MainTitle = styled.h1`
|
||||
font-size: 42px;
|
||||
font-weight: 800;
|
||||
color: hsl(var(--foreground));
|
||||
margin-bottom: 16px;
|
||||
letter-spacing: -1px;
|
||||
line-height: 1.15;
|
||||
|
||||
// Advanced Light & Shadow Gradient
|
||||
background: linear-gradient(
|
||||
135deg,
|
||||
hsl(var(--foreground)) 0%,
|
||||
#8b5cf6 25%,
|
||||
#ec4899 50%,
|
||||
#8b5cf6 75%,
|
||||
hsl(var(--foreground)) 100%
|
||||
);
|
||||
background-size: 300% auto;
|
||||
-webkit-background-clip: text;
|
||||
-webkit-text-fill-color: transparent;
|
||||
|
||||
// Animation
|
||||
animation: ${shimmer} 5s ease-in-out infinite;
|
||||
|
||||
// Optical Glow
|
||||
filter: drop-shadow(0 0 20px rgba(139, 92, 246, 0.3));
|
||||
|
||||
span {
|
||||
display: block; // Force new line for the second part naturally if needed, or keep inline
|
||||
background: linear-gradient(to right, #6366f1, #a855f7, #ec4899);
|
||||
-webkit-background-clip: text;
|
||||
-webkit-text-fill-color: transparent;
|
||||
}
|
||||
`;
|
||||
|
||||
// --- Custom Tabs ---
|
||||
const TabsContainer = styled.div`
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
padding: 6px;
|
||||
background-color: hsl(var(--muted) / 0.4);
|
||||
backdrop-filter: blur(10px);
|
||||
border-radius: 16px;
|
||||
border: 1px solid hsl(var(--border) / 0.5);
|
||||
box-shadow:
|
||||
0 4px 6px -1px rgba(0, 0, 0, 0.01),
|
||||
0 2px 4px -1px rgba(0, 0, 0, 0.01);
|
||||
overflow-x: auto;
|
||||
max-width: 100%;
|
||||
scrollbar-width: none; // hide scrollbar
|
||||
`;
|
||||
|
||||
const TabItem = styled.button<{ $active?: boolean }>`
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
padding: 8px 16px;
|
||||
border-radius: 10px;
|
||||
font-size: 13px;
|
||||
font-weight: 500;
|
||||
transition: all 0.25s cubic-bezier(0.25, 1, 0.5, 1);
|
||||
white-space: nowrap;
|
||||
|
||||
${(props) =>
|
||||
props.$active
|
||||
? css`
|
||||
background-color: hsl(var(--background));
|
||||
color: hsl(var(--foreground));
|
||||
box-shadow: 0 4px 12px rgba(0, 0, 0, 0.08);
|
||||
transform: scale(1.02);
|
||||
`
|
||||
: css`
|
||||
color: hsl(var(--muted-foreground));
|
||||
&:hover {
|
||||
background-color: hsl(var(--muted) / 0.5);
|
||||
color: hsl(var(--foreground));
|
||||
}
|
||||
`}
|
||||
`;
|
||||
|
||||
// --- Input Card ---
|
||||
const InputCard = styled.div`
|
||||
width: 100%;
|
||||
position: relative;
|
||||
background-color: hsl(var(--card));
|
||||
border: 1px solid hsl(var(--border) / 0.6);
|
||||
border-radius: 20px;
|
||||
box-shadow:
|
||||
0 20px 40px -5px rgba(0, 0, 0, 0.03),
|
||||
0 8px 16px -4px rgba(0, 0, 0, 0.03);
|
||||
overflow: visible; // Allow dropdowns to overflow
|
||||
transition: all 0.3s cubic-bezier(0.4, 0, 0.2, 1);
|
||||
|
||||
&:hover {
|
||||
box-shadow:
|
||||
0 25px 50px -12px rgba(0, 0, 0, 0.06),
|
||||
0 12px 24px -6px rgba(0, 0, 0, 0.04);
|
||||
border-color: hsl(var(--primary) / 0.3);
|
||||
}
|
||||
|
||||
&:focus-within {
|
||||
border-color: hsl(var(--primary));
|
||||
box-shadow:
|
||||
0 0 0 4px hsl(var(--primary) / 0.1),
|
||||
0 25px 50px -12px rgba(0, 0, 0, 0.08);
|
||||
}
|
||||
`;
|
||||
|
||||
const StyledTextarea = styled(Textarea)`
|
||||
min-height: 150px;
|
||||
padding: 24px 28px;
|
||||
border: none;
|
||||
font-size: 16px;
|
||||
line-height: 1.6;
|
||||
resize: none;
|
||||
background: transparent;
|
||||
color: hsl(var(--foreground));
|
||||
|
||||
&::placeholder {
|
||||
color: hsl(var(--muted-foreground) / 0.7);
|
||||
font-weight: 300;
|
||||
}
|
||||
|
||||
&:focus-visible {
|
||||
ring: 0;
|
||||
outline: none;
|
||||
box-shadow: none;
|
||||
}
|
||||
`;
|
||||
|
||||
const Toolbar = styled.div`
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 12px 20px 16px 20px;
|
||||
background: linear-gradient(to bottom, transparent, hsl(var(--muted) / 0.2));
|
||||
border-bottom-left-radius: 20px;
|
||||
border-bottom-right-radius: 20px;
|
||||
`;
|
||||
|
||||
const ToolLoginLeft = styled.div`
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 10px;
|
||||
flex-wrap: wrap;
|
||||
`;
|
||||
|
||||
// --- Styles for Selectors ---
|
||||
const ColorDot = styled.div<{ $color: string }>`
|
||||
width: 16px;
|
||||
height: 16px;
|
||||
border-radius: 50%;
|
||||
background-color: ${(props) => props.$color};
|
||||
box-shadow: 0 0 0 1px rgba(0, 0, 0, 0.1) inset;
|
||||
`;
|
||||
|
||||
const GridSelect = styled.div`
|
||||
display: grid;
|
||||
grid-template-columns: repeat(3, 1fr);
|
||||
gap: 8px;
|
||||
padding: 8px;
|
||||
`;
|
||||
|
||||
const GridItem = styled.div<{ $active?: boolean }>`
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
padding: 10px;
|
||||
border-radius: 8px;
|
||||
border: 1px solid
|
||||
${(props) => (props.$active ? "hsl(var(--primary))" : "transparent")};
|
||||
background-color: ${(props) =>
|
||||
props.$active ? "hsl(var(--primary)/0.08)" : "hsl(var(--muted)/0.3)"};
|
||||
cursor: pointer;
|
||||
transition: all 0.2s;
|
||||
|
||||
&:hover {
|
||||
background-color: hsl(var(--primary) / 0.05);
|
||||
}
|
||||
`;
|
||||
|
||||
interface EmptyStateProps {
|
||||
input: string;
|
||||
setInput: (value: string) => void;
|
||||
onSend: (value: string) => void;
|
||||
}
|
||||
|
||||
// Scenarios Configuration
|
||||
const CATEGORIES = [
|
||||
{
|
||||
id: "knowledge",
|
||||
label: "知识探索",
|
||||
icon: <BrainCircuit className="w-4 h-4" />,
|
||||
},
|
||||
{
|
||||
id: "planning",
|
||||
label: "计划规划",
|
||||
icon: <CalendarRange className="w-4 h-4" />,
|
||||
},
|
||||
{ id: "social", label: "社媒内容", icon: <PenTool className="w-4 h-4" /> },
|
||||
{ id: "image", label: "图文海报", icon: <ImageIcon className="w-4 h-4" /> },
|
||||
{ id: "office", label: "办公文档", icon: <FileText className="w-4 h-4" /> },
|
||||
{ id: "video", label: "短视频", icon: <Video className="w-4 h-4" /> },
|
||||
];
|
||||
|
||||
export const EmptyState: React.FC<EmptyStateProps> = ({
|
||||
input,
|
||||
setInput,
|
||||
onSend,
|
||||
}) => {
|
||||
const [activeTab, setActiveTab] = useState("knowledge");
|
||||
|
||||
// Local state for parameters (Mocking visual state)
|
||||
const [platform, setPlatform] = useState("xiaohongshu");
|
||||
const [model, setModel] = useState("gemini");
|
||||
const [ratio, setRatio] = useState("3:4");
|
||||
const [style, setStyle] = useState("minimal");
|
||||
const [depth, setDepth] = useState("deep");
|
||||
|
||||
const handleSend = () => {
|
||||
if (!input.trim()) return;
|
||||
let prefix = "";
|
||||
if (activeTab === "social")
|
||||
prefix = `[社媒创作: ${platform}, Model: ${model}] `;
|
||||
if (activeTab === "image") prefix = `[图文生成: ${ratio}, ${style}] `;
|
||||
if (activeTab === "video") prefix = `[视频脚本] `;
|
||||
if (activeTab === "office") prefix = `[办公文档] `;
|
||||
if (activeTab === "knowledge")
|
||||
prefix = `[知识探索: ${depth === "deep" ? "深度" : "快速"}, Model: ${model}] `;
|
||||
if (activeTab === "planning") prefix = `[计划规划] `;
|
||||
|
||||
onSend(prefix + input);
|
||||
};
|
||||
|
||||
const handleKeyDown = (e: React.KeyboardEvent) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault();
|
||||
handleSend();
|
||||
}
|
||||
};
|
||||
|
||||
// Dynamic Placeholder
|
||||
const getPlaceholder = () => {
|
||||
switch (activeTab) {
|
||||
case "knowledge":
|
||||
return "想了解什么?我可以帮你深度搜索、解析概念或总结长文...";
|
||||
case "planning":
|
||||
return "告诉我你的目标,无论是旅行计划、职业规划还是活动筹备...";
|
||||
case "social":
|
||||
return "输入主题,帮你创作小红书爆款文案、公众号文章...";
|
||||
case "image":
|
||||
return "描述画面主体、风格、构图,生成精美海报或插画...";
|
||||
case "video":
|
||||
return "输入视频主题,生成分镜脚本和口播文案...";
|
||||
case "office":
|
||||
return "输入需求,生成周报、汇报PPT大纲或商务邮件...";
|
||||
default:
|
||||
return "输入你的想法...";
|
||||
}
|
||||
};
|
||||
|
||||
// Helper to get platform icon
|
||||
const getPlatformIcon = (val: string) => {
|
||||
if (val === "xiaohongshu") return iconXhs;
|
||||
if (val === "wechat") return iconGzh;
|
||||
if (val === "zhihu") return iconZhihu;
|
||||
if (val === "toutiao") return iconToutiao;
|
||||
if (val === "juejin") return iconJuejin;
|
||||
if (val === "csdn") return iconCsdn;
|
||||
return undefined;
|
||||
};
|
||||
|
||||
// Helper to get platform label
|
||||
const getPlatformLabel = (val: string) => {
|
||||
if (val === "xiaohongshu") return "小红书";
|
||||
if (val === "wechat") return "公众号";
|
||||
if (val === "zhihu") return "知乎";
|
||||
if (val === "toutiao") return "头条";
|
||||
if (val === "juejin") return "掘金";
|
||||
if (val === "csdn") return "CSDN";
|
||||
return val;
|
||||
};
|
||||
|
||||
// Helper to get model icon
|
||||
const getModelIcon = (val: string) => {
|
||||
if (val === "gemini") return modelGemini;
|
||||
if (val === "claude") return modelClaude;
|
||||
if (val === "deepseek") return modelDeepseek;
|
||||
return undefined;
|
||||
};
|
||||
|
||||
// Helper to get model label
|
||||
const getModelLabel = (val: string) => {
|
||||
if (val === "gemini") return "Gemini 3.0 Pro";
|
||||
if (val === "claude") return "Claude 3.5 Sonnet";
|
||||
if (val === "deepseek") return "DeepSeek V3";
|
||||
return val;
|
||||
};
|
||||
|
||||
return (
|
||||
<Container>
|
||||
<Logo src="/logo.png" alt="ProxyCast" />
|
||||
<Title>ProxyCast Agent</Title>
|
||||
<Description>开始一段新的对话,或从左侧选择一个话题继续</Description>
|
||||
<Tips>
|
||||
<Tip>
|
||||
<kbd>Enter</kbd> 发送消息
|
||||
</Tip>
|
||||
<Tip>
|
||||
<kbd>Shift + Enter</kbd> 换行
|
||||
</Tip>
|
||||
</Tips>
|
||||
<ContentWrapper>
|
||||
<Header>
|
||||
<MainTitle>
|
||||
你想在这个平台 <br />
|
||||
<span>完成什么?</span>
|
||||
</MainTitle>
|
||||
</Header>
|
||||
|
||||
<TabsContainer>
|
||||
{CATEGORIES.map((cat) => (
|
||||
<TabItem
|
||||
key={cat.id}
|
||||
$active={activeTab === cat.id}
|
||||
onClick={() => setActiveTab(cat.id)}
|
||||
>
|
||||
<span
|
||||
className={activeTab === cat.id ? "text-primary" : "opacity-70"}
|
||||
>
|
||||
{cat.icon}
|
||||
</span>
|
||||
{cat.label}
|
||||
</TabItem>
|
||||
))}
|
||||
</TabsContainer>
|
||||
|
||||
<InputCard>
|
||||
<StyledTextarea
|
||||
value={input}
|
||||
onChange={(e) => setInput(e.target.value)}
|
||||
onKeyDown={handleKeyDown}
|
||||
placeholder={getPlaceholder()}
|
||||
/>
|
||||
|
||||
<Toolbar>
|
||||
<ToolLoginLeft>
|
||||
{activeTab === "social" && (
|
||||
<>
|
||||
<Select
|
||||
value={platform}
|
||||
onValueChange={setPlatform}
|
||||
closeOnMouseLeave
|
||||
>
|
||||
<SelectTrigger className="h-8 text-xs bg-background border shadow-sm min-w-[120px]">
|
||||
<div className="flex items-center gap-2">
|
||||
{getPlatformIcon(platform) && (
|
||||
<img
|
||||
src={getPlatformIcon(platform)}
|
||||
className="w-4 h-4 rounded-full"
|
||||
/>
|
||||
)}
|
||||
<span>{getPlatformLabel(platform)}</span>
|
||||
</div>
|
||||
</SelectTrigger>
|
||||
<SelectContent className="p-1">
|
||||
<div className="px-2 py-1.5 text-xs text-muted-foreground font-medium">
|
||||
选择要创作的内容平台
|
||||
</div>
|
||||
<SelectItem value="xiaohongshu">
|
||||
<div className="flex items-center gap-2">
|
||||
<img src={iconXhs} className="w-4 h-4 rounded-full" />{" "}
|
||||
小红书
|
||||
</div>
|
||||
</SelectItem>
|
||||
<SelectItem value="wechat">
|
||||
<div className="flex items-center gap-2">
|
||||
<img src={iconGzh} className="w-4 h-4 rounded-full" />{" "}
|
||||
公众号
|
||||
</div>
|
||||
</SelectItem>
|
||||
<SelectItem value="toutiao">
|
||||
<div className="flex items-center gap-2">
|
||||
<img
|
||||
src={iconToutiao}
|
||||
className="w-4 h-4 rounded-full"
|
||||
/>{" "}
|
||||
今日头条
|
||||
</div>
|
||||
</SelectItem>
|
||||
<SelectItem value="zhihu">
|
||||
<div className="flex items-center gap-2">
|
||||
<img
|
||||
src={iconZhihu}
|
||||
className="w-4 h-4 rounded-full"
|
||||
/>{" "}
|
||||
知乎
|
||||
</div>
|
||||
</SelectItem>
|
||||
<SelectItem value="juejin">
|
||||
<div className="flex items-center gap-2">
|
||||
<img
|
||||
src={iconJuejin}
|
||||
className="w-4 h-4 rounded-full"
|
||||
/>{" "}
|
||||
掘金
|
||||
</div>
|
||||
</SelectItem>
|
||||
<SelectItem value="csdn">
|
||||
<div className="flex items-center gap-2">
|
||||
<img
|
||||
src={iconCsdn}
|
||||
className="w-4 h-4 rounded-full"
|
||||
/>{" "}
|
||||
CSDN
|
||||
</div>
|
||||
</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</>
|
||||
)}
|
||||
|
||||
{activeTab === "knowledge" && (
|
||||
<>
|
||||
<Badge
|
||||
variant="secondary"
|
||||
className="cursor-pointer hover:bg-muted font-normal h-8 px-3 gap-1"
|
||||
>
|
||||
<Search className="w-3.5 h-3.5 mr-1" />
|
||||
联网搜索
|
||||
</Badge>
|
||||
<Select value={depth} onValueChange={setDepth}>
|
||||
<SelectTrigger className="h-8 text-xs bg-background border-input shadow-sm w-[110px]">
|
||||
<BrainCircuit className="w-3.5 h-3.5 mr-2 text-muted-foreground" />
|
||||
<SelectValue placeholder="深度" />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
<SelectItem value="deep">深度解析</SelectItem>
|
||||
<SelectItem value="quick">快速概览</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</>
|
||||
)}
|
||||
|
||||
{activeTab === "planning" && (
|
||||
<Badge
|
||||
variant="outline"
|
||||
className="h-8 font-normal text-muted-foreground gap-1"
|
||||
>
|
||||
<Globe className="w-3.5 h-3.5 mr-1" />
|
||||
旅行/职业/活动
|
||||
</Badge>
|
||||
)}
|
||||
|
||||
{activeTab === "image" && (
|
||||
<>
|
||||
<Popover>
|
||||
<PopoverTrigger asChild>
|
||||
<Button
|
||||
variant="outline"
|
||||
size="sm"
|
||||
className="h-8 text-xs font-normal"
|
||||
>
|
||||
<div className="w-3.5 h-3.5 border border-current rounded-[2px] mr-2 flex items-center justify-center text-[6px]">
|
||||
3:4
|
||||
</div>
|
||||
{ratio}
|
||||
<ChevronDown className="w-3 h-3 ml-1 opacity-50" />
|
||||
</Button>
|
||||
</PopoverTrigger>
|
||||
<PopoverContent className="w-64 p-2" align="start">
|
||||
<div className="text-xs font-medium mb-2 px-2 text-muted-foreground">
|
||||
宽高比
|
||||
</div>
|
||||
<GridSelect>
|
||||
{["1:1", "3:4", "4:3", "9:16", "16:9", "21:9"].map(
|
||||
(r) => (
|
||||
<GridItem
|
||||
key={r}
|
||||
$active={ratio === r}
|
||||
onClick={() => setRatio(r)}
|
||||
>
|
||||
<div className="w-5 h-5 border-2 border-current rounded-sm mb-1 opacity-50"></div>
|
||||
<span className="text-xs">{r}</span>
|
||||
</GridItem>
|
||||
),
|
||||
)}
|
||||
</GridSelect>
|
||||
</PopoverContent>
|
||||
</Popover>
|
||||
|
||||
<Popover>
|
||||
<PopoverTrigger asChild>
|
||||
<Button
|
||||
variant="outline"
|
||||
size="sm"
|
||||
className="h-8 text-xs font-normal"
|
||||
>
|
||||
<ColorDot $color="#3b82f6" className="mr-2" />
|
||||
{style === "minimal"
|
||||
? "极简风格"
|
||||
: style === "tech"
|
||||
? "科技质感"
|
||||
: "温暖治愈"}
|
||||
<ChevronDown className="w-3 h-3 ml-1 opacity-50" />
|
||||
</Button>
|
||||
</PopoverTrigger>
|
||||
<PopoverContent className="w-48 p-1" align="start">
|
||||
<div className="p-1">
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
className="w-full justify-start h-8"
|
||||
onClick={() => setStyle("minimal")}
|
||||
>
|
||||
<ColorDot $color="#e2e8f0" className="mr-2" />{" "}
|
||||
极简风格
|
||||
</Button>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
className="w-full justify-start h-8"
|
||||
onClick={() => setStyle("tech")}
|
||||
>
|
||||
<ColorDot $color="#3b82f6" className="mr-2" />{" "}
|
||||
科技质感
|
||||
</Button>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
className="w-full justify-start h-8"
|
||||
onClick={() => setStyle("warm")}
|
||||
>
|
||||
<ColorDot $color="#f59e0b" className="mr-2" />{" "}
|
||||
温暖治愈
|
||||
</Button>
|
||||
</div>
|
||||
</PopoverContent>
|
||||
</Popover>
|
||||
</>
|
||||
)}
|
||||
|
||||
{/* Model Selector using Popover for better control or just a Select */}
|
||||
<Select value={model} onValueChange={setModel} closeOnMouseLeave>
|
||||
<SelectTrigger className="h-8 text-xs bg-background border shadow-sm min-w-[200px] px-2">
|
||||
<div className="flex items-center gap-1.5 text-muted-foreground">
|
||||
{getModelIcon(model) ? (
|
||||
<img src={getModelIcon(model)} className="w-3.5 h-3.5" />
|
||||
) : (
|
||||
<Sparkles className="w-3.5 h-3.5" />
|
||||
)}
|
||||
<span>{getModelLabel(model)}</span>
|
||||
</div>
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
<SelectItem value="gemini">
|
||||
<div className="flex items-center gap-2">
|
||||
<img src={modelGemini} className="w-4 h-4" /> Gemini 3.0
|
||||
Pro
|
||||
</div>
|
||||
</SelectItem>
|
||||
<SelectItem value="claude">
|
||||
<div className="flex items-center gap-2">
|
||||
<img src={modelClaude} className="w-4 h-4" /> Claude 3.5
|
||||
Sonnet
|
||||
</div>
|
||||
</SelectItem>
|
||||
<SelectItem value="deepseek">
|
||||
<div className="flex items-center gap-2">
|
||||
<img src={modelDeepseek} className="w-4 h-4" /> DeepSeek
|
||||
V3
|
||||
</div>
|
||||
</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
|
||||
<Button
|
||||
variant="outline"
|
||||
size="icon"
|
||||
className="h-8 w-8 rounded-full ml-1 bg-background shadow-sm hover:bg-muted"
|
||||
>
|
||||
<Globe className="w-4 h-4 opacity-70" />
|
||||
</Button>
|
||||
</ToolLoginLeft>
|
||||
|
||||
<Button
|
||||
size="sm"
|
||||
onClick={handleSend}
|
||||
disabled={!input.trim()}
|
||||
className="bg-primary hover:bg-primary/90 text-primary-foreground h-9 px-5 rounded-xl shadow-lg shadow-primary/20 transition-all hover:scale-105 active:scale-95"
|
||||
>
|
||||
开始生成
|
||||
<ArrowRight className="h-4 w-4 ml-2" />
|
||||
</Button>
|
||||
</Toolbar>
|
||||
</InputCard>
|
||||
|
||||
{/* Dynamic Inspiration/Tips based on Tab - Styled nicely */}
|
||||
<div className="w-full max-w-[800px] flex flex-wrap gap-3 justify-center">
|
||||
{activeTab === "social" &&
|
||||
["爆款标题生成", "小红书文案", "公众号排版", "评论区回复"].map(
|
||||
(item) => (
|
||||
<Badge
|
||||
key={item}
|
||||
variant="secondary"
|
||||
className="px-4 py-2 text-xs font-normal cursor-pointer hover:bg-muted-foreground/10 transition-colors"
|
||||
>
|
||||
✨ {item}
|
||||
</Badge>
|
||||
),
|
||||
)}
|
||||
{activeTab === "image" &&
|
||||
["海报设计", "插画生成", "UI 界面", "Logo 设计", "摄影修图"].map(
|
||||
(item) => (
|
||||
<Badge
|
||||
key={item}
|
||||
variant="secondary"
|
||||
className="px-4 py-2 text-xs font-normal cursor-pointer hover:bg-muted-foreground/10 transition-colors"
|
||||
>
|
||||
🎨 {item}
|
||||
</Badge>
|
||||
),
|
||||
)}
|
||||
{activeTab === "knowledge" &&
|
||||
["解释量子计算", "总结这篇论文", "如何制定OKR", "分析行业趋势"].map(
|
||||
(item) => (
|
||||
<Badge
|
||||
key={item}
|
||||
variant="secondary"
|
||||
className="px-4 py-2 text-xs font-normal cursor-pointer hover:bg-muted-foreground/10 transition-colors"
|
||||
>
|
||||
🔍 {item}
|
||||
</Badge>
|
||||
),
|
||||
)}
|
||||
{activeTab === "planning" &&
|
||||
["日本旅行计划", "年度职业规划", "婚礼流程表", "健身计划"].map(
|
||||
(item) => (
|
||||
<Badge
|
||||
key={item}
|
||||
variant="secondary"
|
||||
className="px-4 py-2 text-xs font-normal cursor-pointer hover:bg-muted-foreground/10 transition-colors"
|
||||
>
|
||||
📅 {item}
|
||||
</Badge>
|
||||
),
|
||||
)}
|
||||
</div>
|
||||
</ContentWrapper>
|
||||
</Container>
|
||||
);
|
||||
};
|
||||
|
||||
@@ -28,6 +28,8 @@ import {
|
||||
ThinkingContent,
|
||||
} from "../styles";
|
||||
import { MarkdownRenderer } from "./MarkdownRenderer";
|
||||
import { StreamingRenderer } from "./StreamingRenderer";
|
||||
import { TokenUsageDisplay } from "./TokenUsageDisplay";
|
||||
import { Message } from "../types";
|
||||
|
||||
interface MessageListProps {
|
||||
@@ -166,6 +168,14 @@ export const MessageList: React.FC<MessageListProps> = ({
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
) : msg.role === "assistant" ? (
|
||||
/* 使用 StreamingRenderer 渲染 assistant 消息 - Requirements: 9.3, 9.4 */
|
||||
<StreamingRenderer
|
||||
content={msg.content}
|
||||
isStreaming={msg.isThinking}
|
||||
toolCalls={msg.toolCalls}
|
||||
showCursor={msg.isThinking && !msg.content}
|
||||
/>
|
||||
) : (
|
||||
<MarkdownRenderer content={msg.content} />
|
||||
)}
|
||||
@@ -183,6 +193,11 @@ export const MessageList: React.FC<MessageListProps> = ({
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Token 使用量显示 - Requirements: 9.5 */}
|
||||
{msg.role === "assistant" && !msg.isThinking && msg.usage && (
|
||||
<TokenUsageDisplay usage={msg.usage} />
|
||||
)}
|
||||
|
||||
{editingId !== msg.id && (
|
||||
<MessageActions className="message-actions">
|
||||
<Button
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
/**
|
||||
* 流式消息渲染组件
|
||||
*
|
||||
* 实现实时 Markdown 渲染,区分文本响应和工具调用响应
|
||||
* Requirements: 9.3, 9.4
|
||||
*/
|
||||
|
||||
import React, { memo, useMemo } from "react";
|
||||
import styled, { keyframes } from "styled-components";
|
||||
import { MarkdownRenderer } from "./MarkdownRenderer";
|
||||
import { ToolCallList } from "./ToolCallDisplay";
|
||||
import type { ToolCallState } from "@/lib/api/agent";
|
||||
|
||||
// 光标闪烁动画
|
||||
const blink = keyframes`
|
||||
0%, 50% { opacity: 1; }
|
||||
51%, 100% { opacity: 0; }
|
||||
`;
|
||||
|
||||
const StreamingContainer = styled.div`
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 8px;
|
||||
`;
|
||||
|
||||
const TextSection = styled.div`
|
||||
position: relative;
|
||||
`;
|
||||
|
||||
const StreamingCursor = styled.span`
|
||||
display: inline-block;
|
||||
width: 2px;
|
||||
height: 1em;
|
||||
background-color: hsl(var(--primary));
|
||||
margin-left: 2px;
|
||||
vertical-align: text-bottom;
|
||||
animation: ${blink} 1s step-end infinite;
|
||||
`;
|
||||
|
||||
const ToolSection = styled.div`
|
||||
margin-top: 8px;
|
||||
`;
|
||||
|
||||
interface StreamingRendererProps {
|
||||
/** 文本内容 */
|
||||
content: string;
|
||||
/** 是否正在流式输出 */
|
||||
isStreaming?: boolean;
|
||||
/** 工具调用列表 */
|
||||
toolCalls?: ToolCallState[];
|
||||
/** 是否显示光标 */
|
||||
showCursor?: boolean;
|
||||
}
|
||||
|
||||
/**
|
||||
* 流式消息渲染组件
|
||||
*
|
||||
* 支持实时 Markdown 渲染和工具调用显示
|
||||
* Requirements: 9.3 - THE Frontend SHALL distinguish between text responses and tool call responses visually
|
||||
* Requirements: 9.4 - WHEN streaming text, THE Frontend SHALL render markdown formatting in real-time
|
||||
*/
|
||||
export const StreamingRenderer: React.FC<StreamingRendererProps> = memo(
|
||||
({ content, isStreaming = false, toolCalls, showCursor = true }) => {
|
||||
// 判断是否有正在执行的工具
|
||||
const hasRunningTools = useMemo(
|
||||
() => toolCalls?.some((tc) => tc.status === "running") ?? false,
|
||||
[toolCalls],
|
||||
);
|
||||
|
||||
// 判断是否显示光标
|
||||
const shouldShowCursor = isStreaming && showCursor && !hasRunningTools;
|
||||
|
||||
// 判断是否有工具调用
|
||||
const hasToolCalls = toolCalls && toolCalls.length > 0;
|
||||
|
||||
return (
|
||||
<StreamingContainer>
|
||||
{/* 工具调用区域 - 显示在文本之前 */}
|
||||
{hasToolCalls && (
|
||||
<ToolSection>
|
||||
<ToolCallList toolCalls={toolCalls} />
|
||||
</ToolSection>
|
||||
)}
|
||||
|
||||
{/* 文本内容区域 */}
|
||||
{content && (
|
||||
<TextSection>
|
||||
<MarkdownRenderer content={content} />
|
||||
{shouldShowCursor && <StreamingCursor />}
|
||||
</TextSection>
|
||||
)}
|
||||
|
||||
{/* 如果没有内容但正在流式输出,显示光标 */}
|
||||
{!content && isStreaming && showCursor && !hasRunningTools && (
|
||||
<TextSection>
|
||||
<StreamingCursor />
|
||||
</TextSection>
|
||||
)}
|
||||
</StreamingContainer>
|
||||
);
|
||||
},
|
||||
);
|
||||
|
||||
StreamingRenderer.displayName = "StreamingRenderer";
|
||||
|
||||
export default StreamingRenderer;
|
||||
@@ -0,0 +1,69 @@
|
||||
/**
|
||||
* Token 使用量显示组件
|
||||
*
|
||||
* 在响应完成后显示 token 使用量
|
||||
* Requirements: 9.5 - THE Frontend SHALL display token usage statistics after each Agent response
|
||||
*/
|
||||
|
||||
import React from "react";
|
||||
import styled from "styled-components";
|
||||
import { Coins } from "lucide-react";
|
||||
import type { TokenUsage } from "@/lib/api/agent";
|
||||
|
||||
const UsageContainer = styled.div`
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
padding: 4px 10px;
|
||||
border-radius: 6px;
|
||||
background-color: hsl(var(--muted) / 0.5);
|
||||
font-size: 11px;
|
||||
color: hsl(var(--muted-foreground));
|
||||
margin-top: 8px;
|
||||
`;
|
||||
|
||||
const UsageIcon = styled(Coins)`
|
||||
width: 12px;
|
||||
height: 12px;
|
||||
opacity: 0.7;
|
||||
`;
|
||||
|
||||
const UsageText = styled.span`
|
||||
font-family:
|
||||
ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, "Liberation Mono",
|
||||
"Courier New", monospace;
|
||||
`;
|
||||
|
||||
const Separator = styled.span`
|
||||
opacity: 0.5;
|
||||
`;
|
||||
|
||||
interface TokenUsageDisplayProps {
|
||||
usage: TokenUsage;
|
||||
className?: string;
|
||||
}
|
||||
|
||||
/**
|
||||
* Token 使用量显示组件
|
||||
*
|
||||
* 显示输入/输出 token 数量
|
||||
*/
|
||||
export const TokenUsageDisplay: React.FC<TokenUsageDisplayProps> = ({
|
||||
usage,
|
||||
className,
|
||||
}) => {
|
||||
const total = usage.input_tokens + usage.output_tokens;
|
||||
|
||||
return (
|
||||
<UsageContainer className={className}>
|
||||
<UsageIcon />
|
||||
<UsageText>{usage.input_tokens.toLocaleString()} in</UsageText>
|
||||
<Separator>/</Separator>
|
||||
<UsageText>{usage.output_tokens.toLocaleString()} out</UsageText>
|
||||
<Separator>·</Separator>
|
||||
<UsageText>{total.toLocaleString()} total</UsageText>
|
||||
</UsageContainer>
|
||||
);
|
||||
};
|
||||
|
||||
export default TokenUsageDisplay;
|
||||
@@ -0,0 +1,288 @@
|
||||
/**
|
||||
* 工具调用显示组件
|
||||
*
|
||||
* 显示工具执行状态和结果
|
||||
* Requirements: 9.1, 9.2 - 工具执行指示器和结果折叠面板
|
||||
*/
|
||||
|
||||
import React, { useState } from "react";
|
||||
import styled, { keyframes } from "styled-components";
|
||||
import {
|
||||
Terminal,
|
||||
FileText,
|
||||
Edit3,
|
||||
FolderOpen,
|
||||
ChevronDown,
|
||||
ChevronRight,
|
||||
Check,
|
||||
X,
|
||||
Loader2,
|
||||
} from "lucide-react";
|
||||
import type { ToolCallState } from "@/lib/api/agent";
|
||||
|
||||
// 动画
|
||||
const spin = keyframes`
|
||||
from { transform: rotate(0deg); }
|
||||
to { transform: rotate(360deg); }
|
||||
`;
|
||||
|
||||
const pulse = keyframes`
|
||||
0%, 100% { opacity: 1; }
|
||||
50% { opacity: 0.5; }
|
||||
`;
|
||||
|
||||
// 样式组件
|
||||
const ToolCallContainer = styled.div`
|
||||
margin: 8px 0;
|
||||
border: 1px solid hsl(var(--border));
|
||||
border-radius: 8px;
|
||||
overflow: hidden;
|
||||
background-color: hsl(var(--muted) / 0.3);
|
||||
`;
|
||||
|
||||
const ToolCallHeader = styled.div<{ $status: string }>`
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
padding: 10px 12px;
|
||||
cursor: pointer;
|
||||
transition: background-color 0.2s;
|
||||
background-color: ${(props) =>
|
||||
props.$status === "running"
|
||||
? "hsl(var(--primary) / 0.1)"
|
||||
: props.$status === "failed"
|
||||
? "hsl(var(--destructive) / 0.1)"
|
||||
: "transparent"};
|
||||
|
||||
&:hover {
|
||||
background-color: hsl(var(--muted) / 0.5);
|
||||
}
|
||||
`;
|
||||
|
||||
const ToolIcon = styled.div<{ $status: string }>`
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
width: 24px;
|
||||
height: 24px;
|
||||
border-radius: 4px;
|
||||
background-color: ${(props) =>
|
||||
props.$status === "running"
|
||||
? "hsl(var(--primary))"
|
||||
: props.$status === "failed"
|
||||
? "hsl(var(--destructive))"
|
||||
: "hsl(var(--primary) / 0.8)"};
|
||||
color: white;
|
||||
`;
|
||||
|
||||
const SpinningLoader = styled(Loader2)`
|
||||
animation: ${spin} 1s linear infinite;
|
||||
`;
|
||||
|
||||
const ToolName = styled.span`
|
||||
font-size: 13px;
|
||||
font-weight: 500;
|
||||
color: hsl(var(--foreground));
|
||||
flex: 1;
|
||||
`;
|
||||
|
||||
const ToolStatus = styled.span<{ $status: string }>`
|
||||
font-size: 12px;
|
||||
padding: 2px 8px;
|
||||
border-radius: 4px;
|
||||
background-color: ${(props) =>
|
||||
props.$status === "running"
|
||||
? "hsl(var(--primary) / 0.2)"
|
||||
: props.$status === "failed"
|
||||
? "hsl(var(--destructive) / 0.2)"
|
||||
: "hsl(var(--primary) / 0.1)"};
|
||||
color: ${(props) =>
|
||||
props.$status === "running"
|
||||
? "hsl(var(--primary))"
|
||||
: props.$status === "failed"
|
||||
? "hsl(var(--destructive))"
|
||||
: "hsl(var(--primary))"};
|
||||
animation: ${(props) => (props.$status === "running" ? pulse : "none")} 1.5s
|
||||
ease-in-out infinite;
|
||||
`;
|
||||
|
||||
const ExpandIcon = styled.div`
|
||||
color: hsl(var(--muted-foreground));
|
||||
transition: transform 0.2s;
|
||||
`;
|
||||
|
||||
const ToolResultPanel = styled.div<{ $expanded: boolean }>`
|
||||
display: ${(props) => (props.$expanded ? "block" : "none")};
|
||||
border-top: 1px solid hsl(var(--border));
|
||||
background-color: hsl(var(--background));
|
||||
`;
|
||||
|
||||
const ResultContent = styled.pre`
|
||||
margin: 0;
|
||||
padding: 12px;
|
||||
font-size: 12px;
|
||||
font-family:
|
||||
ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, "Liberation Mono",
|
||||
"Courier New", monospace;
|
||||
white-space: pre-wrap;
|
||||
word-break: break-word;
|
||||
max-height: 300px;
|
||||
overflow-y: auto;
|
||||
color: hsl(var(--foreground));
|
||||
line-height: 1.5;
|
||||
`;
|
||||
|
||||
const ErrorContent = styled(ResultContent)`
|
||||
color: hsl(var(--destructive));
|
||||
background-color: hsl(var(--destructive) / 0.05);
|
||||
`;
|
||||
|
||||
const ExecutionTime = styled.span`
|
||||
font-size: 11px;
|
||||
color: hsl(var(--muted-foreground));
|
||||
margin-left: auto;
|
||||
margin-right: 8px;
|
||||
`;
|
||||
|
||||
// 工具图标映射
|
||||
const getToolIcon = (toolName: string) => {
|
||||
const name = toolName.toLowerCase();
|
||||
if (
|
||||
name.includes("bash") ||
|
||||
name.includes("shell") ||
|
||||
name.includes("exec")
|
||||
) {
|
||||
return Terminal;
|
||||
}
|
||||
if (name.includes("read") || name.includes("file")) {
|
||||
return FileText;
|
||||
}
|
||||
if (name.includes("edit") || name.includes("write")) {
|
||||
return Edit3;
|
||||
}
|
||||
if (name.includes("list") || name.includes("dir")) {
|
||||
return FolderOpen;
|
||||
}
|
||||
return Terminal;
|
||||
};
|
||||
|
||||
// 工具名称显示映射
|
||||
const getToolDisplayName = (toolName: string): string => {
|
||||
const nameMap: Record<string, string> = {
|
||||
bash: "执行命令",
|
||||
read_file: "读取文件",
|
||||
write_file: "写入文件",
|
||||
edit_file: "编辑文件",
|
||||
list_directory: "列出目录",
|
||||
};
|
||||
return nameMap[toolName] || toolName;
|
||||
};
|
||||
|
||||
// 状态显示文本
|
||||
const getStatusText = (status: string): string => {
|
||||
switch (status) {
|
||||
case "running":
|
||||
return "执行中...";
|
||||
case "completed":
|
||||
return "完成";
|
||||
case "failed":
|
||||
return "失败";
|
||||
default:
|
||||
return status;
|
||||
}
|
||||
};
|
||||
|
||||
interface ToolCallDisplayProps {
|
||||
toolCall: ToolCallState;
|
||||
defaultExpanded?: boolean;
|
||||
}
|
||||
|
||||
/**
|
||||
* 单个工具调用显示组件
|
||||
*/
|
||||
export const ToolCallDisplay: React.FC<ToolCallDisplayProps> = ({
|
||||
toolCall,
|
||||
defaultExpanded = false,
|
||||
}) => {
|
||||
const [expanded, setExpanded] = useState(defaultExpanded);
|
||||
const IconComponent = getToolIcon(toolCall.name);
|
||||
|
||||
// 计算执行时间
|
||||
const executionTime =
|
||||
toolCall.endTime && toolCall.startTime
|
||||
? Math.round(
|
||||
(toolCall.endTime.getTime() - toolCall.startTime.getTime()) / 1000,
|
||||
)
|
||||
: null;
|
||||
|
||||
const hasResult = toolCall.status !== "running" && toolCall.result;
|
||||
|
||||
return (
|
||||
<ToolCallContainer>
|
||||
<ToolCallHeader
|
||||
$status={toolCall.status}
|
||||
onClick={() => hasResult && setExpanded(!expanded)}
|
||||
>
|
||||
<ToolIcon $status={toolCall.status}>
|
||||
{toolCall.status === "running" ? (
|
||||
<SpinningLoader size={14} />
|
||||
) : toolCall.status === "failed" ? (
|
||||
<X size={14} />
|
||||
) : (
|
||||
<IconComponent size={14} />
|
||||
)}
|
||||
</ToolIcon>
|
||||
|
||||
<ToolName>{getToolDisplayName(toolCall.name)}</ToolName>
|
||||
|
||||
{executionTime !== null && (
|
||||
<ExecutionTime>{executionTime}s</ExecutionTime>
|
||||
)}
|
||||
|
||||
<ToolStatus $status={toolCall.status}>
|
||||
{toolCall.status === "completed" && <Check size={12} />}
|
||||
{getStatusText(toolCall.status)}
|
||||
</ToolStatus>
|
||||
|
||||
{hasResult && (
|
||||
<ExpandIcon>
|
||||
{expanded ? <ChevronDown size={16} /> : <ChevronRight size={16} />}
|
||||
</ExpandIcon>
|
||||
)}
|
||||
</ToolCallHeader>
|
||||
|
||||
{hasResult && (
|
||||
<ToolResultPanel $expanded={expanded}>
|
||||
{toolCall.result?.error ? (
|
||||
<ErrorContent>{toolCall.result.error}</ErrorContent>
|
||||
) : (
|
||||
<ResultContent>
|
||||
{toolCall.result?.output || "(无输出)"}
|
||||
</ResultContent>
|
||||
)}
|
||||
</ToolResultPanel>
|
||||
)}
|
||||
</ToolCallContainer>
|
||||
);
|
||||
};
|
||||
|
||||
interface ToolCallListProps {
|
||||
toolCalls: ToolCallState[];
|
||||
}
|
||||
|
||||
/**
|
||||
* 工具调用列表组件
|
||||
*/
|
||||
export const ToolCallList: React.FC<ToolCallListProps> = ({ toolCalls }) => {
|
||||
if (!toolCalls || toolCalls.length === 0) return null;
|
||||
|
||||
return (
|
||||
<div>
|
||||
{toolCalls.map((tc) => (
|
||||
<ToolCallDisplay key={tc.id} toolCall={tc} />
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
export default ToolCallDisplay;
|
||||
@@ -130,17 +130,27 @@ export function AgentChatPage({
|
||||
/>
|
||||
</ChatContent>
|
||||
) : (
|
||||
<EmptyState />
|
||||
<EmptyState
|
||||
input={input}
|
||||
setInput={setInput}
|
||||
onSend={(text) => {
|
||||
setInput(text);
|
||||
// 使用 setTimeout 确保 state 更新后再发送
|
||||
setTimeout(() => handleSend([], false, false), 0);
|
||||
}}
|
||||
/>
|
||||
)}
|
||||
|
||||
<Inputbar
|
||||
input={input}
|
||||
setInput={setInput}
|
||||
onSend={handleSend}
|
||||
isLoading={isSending}
|
||||
disabled={!processStatus.running && false}
|
||||
onClearMessages={handleClearMessages}
|
||||
/>
|
||||
{hasMessages && (
|
||||
<Inputbar
|
||||
input={input}
|
||||
setInput={setInput}
|
||||
onSend={handleSend}
|
||||
isLoading={isSending}
|
||||
disabled={!processStatus.running && false}
|
||||
onClearMessages={handleClearMessages}
|
||||
/>
|
||||
)}
|
||||
</ChatContainer>
|
||||
</MainArea>
|
||||
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import type { ToolCallState, TokenUsage } from "@/lib/api/agent";
|
||||
|
||||
export interface MessageImage {
|
||||
data: string;
|
||||
mediaType: string;
|
||||
@@ -12,6 +14,10 @@ export interface Message {
|
||||
isThinking?: boolean;
|
||||
thinkingContent?: string;
|
||||
search_results?: any[]; // For potential future use
|
||||
/** 工具调用列表(assistant 消息可能包含) */
|
||||
toolCalls?: ToolCallState[];
|
||||
/** Token 使用量(响应完成后) */
|
||||
usage?: TokenUsage;
|
||||
}
|
||||
|
||||
export interface ChatSession {
|
||||
|
||||
@@ -18,6 +18,7 @@ interface SelectProps {
|
||||
onValueChange?: (value: string) => void;
|
||||
disabled?: boolean;
|
||||
children: React.ReactNode;
|
||||
closeOnMouseLeave?: boolean;
|
||||
}
|
||||
|
||||
const Select: React.FC<SelectProps> = ({
|
||||
@@ -26,6 +27,7 @@ const Select: React.FC<SelectProps> = ({
|
||||
onValueChange,
|
||||
disabled = false,
|
||||
children,
|
||||
closeOnMouseLeave = false,
|
||||
}) => {
|
||||
const [internalValue, setInternalValue] = useState(defaultValue || "");
|
||||
const [open, setOpen] = useState(false);
|
||||
@@ -43,7 +45,12 @@ const Select: React.FC<SelectProps> = ({
|
||||
disabled,
|
||||
}}
|
||||
>
|
||||
<div className="relative">{children}</div>
|
||||
<div
|
||||
className="relative"
|
||||
onMouseLeave={closeOnMouseLeave ? () => setOpen(false) : undefined}
|
||||
>
|
||||
{children}
|
||||
</div>
|
||||
</SelectContext.Provider>
|
||||
);
|
||||
};
|
||||
@@ -133,22 +140,41 @@ const SelectItem: React.FC<SelectItemProps> = ({
|
||||
const context = useContext(SelectContext);
|
||||
if (!context) throw new Error("SelectItem must be used within Select");
|
||||
|
||||
const { onValueChange, setOpen } = context;
|
||||
const { onValueChange, setOpen, value: selectedValue } = context;
|
||||
|
||||
const handleSelect = () => {
|
||||
onValueChange(value);
|
||||
setOpen(false);
|
||||
};
|
||||
|
||||
const isSelected = selectedValue === value;
|
||||
|
||||
return (
|
||||
<div
|
||||
className={cn(
|
||||
"relative flex cursor-default select-none items-center rounded-sm px-2 py-1.5 text-sm outline-none hover:bg-gray-100",
|
||||
"relative flex cursor-default select-none items-center justify-between rounded-sm px-2 py-2 text-sm outline-none transition-colors hover:bg-accent hover:text-accent-foreground data-[state=checked]:bg-accent/50",
|
||||
isSelected && "bg-accent/50 text-accent-foreground",
|
||||
className,
|
||||
)}
|
||||
onClick={handleSelect}
|
||||
>
|
||||
{children}
|
||||
{isSelected && (
|
||||
<span className="flex h-3.5 w-3.5 items-center justify-center ml-2">
|
||||
<svg
|
||||
xmlns="http://www.w3.org/2000/svg"
|
||||
viewBox="0 0 24 24"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
strokeWidth="2"
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
className="h-4 w-4 opacity-100" // Always visible if selected
|
||||
>
|
||||
<polyline points="20 6 9 17 4 12" />
|
||||
</svg>
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
@@ -2,10 +2,164 @@
|
||||
* Agent API
|
||||
*
|
||||
* 原生 Rust Agent 的前端 API 封装
|
||||
* 支持流式输出和工具调用
|
||||
*/
|
||||
|
||||
import { invoke } from "@tauri-apps/api/core";
|
||||
|
||||
// ============================================================
|
||||
// 流式事件类型 (Requirements: 9.1, 9.2, 9.3)
|
||||
// ============================================================
|
||||
|
||||
/**
|
||||
* Token 使用量统计
|
||||
* Requirements: 9.5 - THE Frontend SHALL display token usage statistics after each Agent response
|
||||
*/
|
||||
export interface TokenUsage {
|
||||
/** 输入 token 数 */
|
||||
input_tokens: number;
|
||||
/** 输出 token 数 */
|
||||
output_tokens: number;
|
||||
}
|
||||
|
||||
/**
|
||||
* 工具执行结果
|
||||
* Requirements: 9.2 - THE Frontend SHALL display a collapsible section showing the tool result
|
||||
*/
|
||||
export interface ToolExecutionResult {
|
||||
/** 是否成功 */
|
||||
success: boolean;
|
||||
/** 输出内容 */
|
||||
output: string;
|
||||
/** 错误信息(如果失败) */
|
||||
error?: string;
|
||||
}
|
||||
|
||||
/**
|
||||
* 流式事件类型
|
||||
* Requirements: 9.1, 9.2, 9.3
|
||||
*/
|
||||
export type StreamEvent =
|
||||
| StreamEventTextDelta
|
||||
| StreamEventToolStart
|
||||
| StreamEventToolEnd
|
||||
| StreamEventDone
|
||||
| StreamEventError;
|
||||
|
||||
/**
|
||||
* 文本增量事件
|
||||
* Requirements: 9.3 - THE Frontend SHALL distinguish between text responses and tool call responses visually
|
||||
*/
|
||||
export interface StreamEventTextDelta {
|
||||
type: "text_delta";
|
||||
text: string;
|
||||
}
|
||||
|
||||
/**
|
||||
* 工具调用开始事件
|
||||
* Requirements: 9.1 - WHEN a tool is being executed, THE Frontend SHALL display a tool execution indicator with the tool name
|
||||
*/
|
||||
export interface StreamEventToolStart {
|
||||
type: "tool_start";
|
||||
/** 工具名称 */
|
||||
tool_name: string;
|
||||
/** 工具调用 ID */
|
||||
tool_id: string;
|
||||
}
|
||||
|
||||
/**
|
||||
* 工具调用结束事件
|
||||
* Requirements: 9.2 - WHEN a tool completes, THE Frontend SHALL display a collapsible section showing the tool result
|
||||
*/
|
||||
export interface StreamEventToolEnd {
|
||||
type: "tool_end";
|
||||
/** 工具调用 ID */
|
||||
tool_id: string;
|
||||
/** 工具执行结果 */
|
||||
result: ToolExecutionResult;
|
||||
}
|
||||
|
||||
/**
|
||||
* 完成事件
|
||||
* Requirements: 9.5 - THE Frontend SHALL display token usage statistics after each Agent response
|
||||
*/
|
||||
export interface StreamEventDone {
|
||||
type: "done";
|
||||
/** Token 使用量(可选) */
|
||||
usage?: TokenUsage;
|
||||
}
|
||||
|
||||
/**
|
||||
* 错误事件
|
||||
*/
|
||||
export interface StreamEventError {
|
||||
type: "error";
|
||||
/** 错误信息 */
|
||||
message: string;
|
||||
}
|
||||
|
||||
/**
|
||||
* 工具调用状态(用于 UI 显示)
|
||||
*/
|
||||
export interface ToolCallState {
|
||||
/** 工具调用 ID */
|
||||
id: string;
|
||||
/** 工具名称 */
|
||||
name: string;
|
||||
/** 执行状态 */
|
||||
status: "running" | "completed" | "failed";
|
||||
/** 执行结果(完成后) */
|
||||
result?: ToolExecutionResult;
|
||||
/** 开始时间 */
|
||||
startTime: Date;
|
||||
/** 结束时间(完成后) */
|
||||
endTime?: Date;
|
||||
}
|
||||
|
||||
/**
|
||||
* 解析流式事件
|
||||
* @param data - 原始事件数据
|
||||
* @returns 解析后的流式事件
|
||||
*/
|
||||
export function parseStreamEvent(data: unknown): StreamEvent | null {
|
||||
if (!data || typeof data !== "object") return null;
|
||||
|
||||
const event = data as Record<string, unknown>;
|
||||
const type = event.type as string;
|
||||
|
||||
switch (type) {
|
||||
case "text_delta":
|
||||
return {
|
||||
type: "text_delta",
|
||||
text: (event.text as string) || "",
|
||||
};
|
||||
case "tool_start":
|
||||
return {
|
||||
type: "tool_start",
|
||||
tool_name: (event.tool_name as string) || "",
|
||||
tool_id: (event.tool_id as string) || "",
|
||||
};
|
||||
case "tool_end":
|
||||
return {
|
||||
type: "tool_end",
|
||||
tool_id: (event.tool_id as string) || "",
|
||||
result: event.result as ToolExecutionResult,
|
||||
};
|
||||
case "done":
|
||||
return {
|
||||
type: "done",
|
||||
usage: event.usage as TokenUsage | undefined,
|
||||
};
|
||||
case "error":
|
||||
return {
|
||||
type: "error",
|
||||
message: (event.message as string) || "Unknown error",
|
||||
};
|
||||
default:
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Agent 状态
|
||||
*/
|
||||
|
||||