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
This commit is contained in:
coso
2025-12-31 19:35:02 +08:00
parent c621779c75
commit 602fd07384
46 changed files with 10614 additions and 122 deletions
+1 -1
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.24.0",
"version": "0.25.0",
"type": "module",
"repository": {
"type": "git",
+1 -1
View File
@@ -3674,7 +3674,7 @@ dependencies = [
[[package]]
name = "proxycast"
version = "0.24.0"
version = "0.25.0"
dependencies = [
"anyhow",
"arboard",
+1 -1
View File
@@ -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"]
+28 -5
View File
@@ -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?;
```
## 更新提醒
+4
View File
@@ -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::*;
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+284
View File
@@ -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 包含所有可用工具定义
## 更新提醒
任何文件变更后,请更新此文档和相关的上级文档。
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+32
View File
@@ -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};
+613
View File
@@ -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('&', "&amp;")
.replace('<', "&lt;")
.replace('>', "&gt;")
.replace('"', "&quot;")
.replace('\'', "&apos;")
}
/// 从 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>"), "&lt;script&gt;");
assert_eq!(escape_xml("a & b"), "a &amp; b");
assert_eq!(escape_xml("\"quoted\""), "&quot;quoted&quot;");
assert_eq!(escape_xml("it's"), "it&apos;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
);
}
}
}
}
File diff suppressed because it is too large Load Diff
+390
View File
@@ -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());
}
}
+594
View File
@@ -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
);
}
}
}
}
+554
View File
@@ -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 属性应该通过验证");
}
}
}
+802
View File
@@ -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(),
"空内容写入后文件应该为空"
);
}
}
}
+139 -2
View File
@@ -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 -1
View File
@@ -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",
Binary file not shown.

After

Width:  |  Height:  |  Size: 930 B

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 8.6 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 807 B

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.7 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 890 B

Binary file not shown.

After

Width:  |  Height:  |  Size: 4.3 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 32 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 5.8 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 7.8 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 4.5 KiB

Binary file not shown.

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;
+19 -9
View File
@@ -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>
+6
View File
@@ -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 {
+29 -3
View File
@@ -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>
);
};
+154
View File
@@ -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 状态
*/