chore: bump version to 0.48.4

This commit is contained in:
coso
2026-01-30 01:31:09 +08:00
parent e7385a2467
commit 2a7ed06d21
94 changed files with 6126 additions and 8461 deletions
+1
View File
@@ -53,3 +53,4 @@ src-tauri/gen
.task
Taskfile.yml
nul
.proptest-regressions
+24
View File
@@ -5,6 +5,30 @@
## 基本规则
1. **始终使用中文输出** - 所有回复、注释、文档都使用中文
2. **文件超过 20 行,分批输出** - 避免一次性输出过长内容
3. **先读后写** - 修改文件前必须先读取现有内容
## 详细文档
模块级详细文档位于 `docs/aiprompts/`:
| 文档 | 说明 |
|------|------|
| [overview.md](docs/aiprompts/overview.md) | 项目架构概览 |
| [providers.md](docs/aiprompts/providers.md) | Provider 系统 |
| [credential-pool.md](docs/aiprompts/credential-pool.md) | 凭证池管理 |
| [converter.md](docs/aiprompts/converter.md) | 协议转换 |
| [server.md](docs/aiprompts/server.md) | HTTP 服务器 |
| [flow-monitor.md](docs/aiprompts/flow-monitor.md) | 流量监控 |
| [components.md](docs/aiprompts/components.md) | 组件系统 |
| [hooks.md](docs/aiprompts/hooks.md) | React Hooks |
| [services.md](docs/aiprompts/services.md) | 业务服务 |
| [commands.md](docs/aiprompts/commands.md) | Tauri 命令 |
| [mcp.md](docs/aiprompts/mcp.md) | MCP 服务器 |
| [database.md](docs/aiprompts/database.md) | 数据库层 |
| [terminal.md](docs/aiprompts/terminal.md) | 内置终端 |
| [plugins.md](docs/aiprompts/plugins.md) | 插件系统 |
| [lib.md](docs/aiprompts/lib.md) | 工具库 |
## 构建命令
-285
View File
@@ -1,285 +0,0 @@
# OAuth 插件系统删除实施计划
## 目标说明
**删除内容**:OAuth 插件管理系统("OAuth 插件" 标签页及其相关功能)
**保留内容**:5 个内置 OAuth 提供者的凭证管理功能(Kiro, Gemini, Antigravity, Codex, Claude)
## 架构确认
### ✅ 保留 - 内置 OAuth 提供者
这些提供者有完整的内置实现,不依赖插件系统:
1. **Kiro** - `src-tauri/src/providers/kiro.rs`
2. **Gemini** - `src-tauri/src/providers/gemini.rs`
3. **Antigravity** - `src-tauri/src/providers/antigravity.rs`
4. **Codex** - `src-tauri/src/providers/codex.rs`
5. **Claude** - `src-tauri/src/providers/claude_oauth.rs`
### ❌ 删除 - OAuth 插件系统
这些是可扩展的插件管理系统,允许安装/卸载第三方 OAuth 提供者插件:
- 插件加载器:`oauth_plugin_loader.rs`
- 插件注册表:`credential/registry.rs` 中的插件部分
- 插件管理命令:`oauth_plugin_cmd.rs`
- 插件 UI 组件:`OAuthPluginTab.tsx`, `OAuthPluginContainer.tsx`
- 插件 API:`oauthPlugin.ts`, `useOAuthPlugins.ts`
---
## Stage 1: 删除前端插件 UI 组件
**Goal**: 删除 "OAuth 插件" 标签页及相关 UI 组件
**Success Criteria**: 前端编译无错误,UI 中不再显示 "OAuth 插件" 标签
**Tests**: 应用启动正常,OAuth 凭证管理功能正常
**Status**: ✅ Complete
### 需要删除的文件:
- `src/components/provider-pool/OAuthPluginTab.tsx` - OAuth 插件标签页
- `src/components/plugins/OAuthPluginContainer.tsx` - 插件容器组件
- `src/components/provider-pool/credential-forms/ClaudeOAuthForm.tsx` - 如果仅用于插件
- `src/components/provider-pool/credential-forms/OAuthUrlDisplay.tsx` - 如果仅用于插件
### 需要修改的文件:
- `src/components/provider-pool/ProviderPoolPage.tsx` 或类似的父组件
- 移除 "OAuth 插件" 标签页的引用
- 保留 "OAuth 凭证" 标签页
---
## Stage 2: 删除前端插件 API 和 Hooks
**Goal**: 删除插件管理相关的前端 API 调用和 Hooks
**Success Criteria**: 前端编译无错误,无未使用的导入
**Tests**: 其他 API 调用正常工作
**Status**: Not Started
### 需要删除的文件:
- `src/lib/api/oauthPlugin.ts` - 插件管理 API
- `src/hooks/useOAuthPlugins.ts` - 插件管理 Hook
- `src/hooks/useOAuthCredentials.ts` - 如果仅用于插件凭证
### 需要保留的文件:
- `src/lib/api/credentials.ts` - 保留内置 OAuth 凭证 API
- 其他与 Kiro/Gemini/Antigravity/Codex/Claude 凭证管理相关的 API
---
## Stage 3: 删除后端插件命令
**Goal**: 删除插件管理相关的 Tauri 命令
**Success Criteria**: 后端编译无错误
**Tests**: 内置 OAuth 命令正常工作
**Status**: Not Started
### 需要删除的文件:
- `src-tauri/src/commands/oauth_plugin_cmd.rs` - 完整删除
### 需要保留的文件:
- `src-tauri/src/commands/oauth_cmd.rs` - 保留(内置 OAuth 凭证命令)
### 需要修改的文件:
- `src-tauri/src/commands/mod.rs`
- 移除 `oauth_plugin_cmd` 模块引用
- 移除所有插件相关命令的注册:
- `init_oauth_plugin_system`
- `list_oauth_plugins`
- `get_oauth_plugin`
- `enable_oauth_plugin`
- `disable_oauth_plugin`
- `install_oauth_plugin`
- `uninstall_oauth_plugin`
- `check_oauth_plugin_updates`
- `update_oauth_plugin`
- `reload_oauth_plugins`
- `get_oauth_plugin_config`
- `update_oauth_plugin_config`
- `scan_oauth_plugin_directory`
- `plugin_credential_*` 系列命令
- `plugin_database_*` 系列命令
- `plugin_http_request`
- `plugin_crypto_*` 系列命令
- `plugin_notification`
- `plugin_storage_*` 系列命令
- `plugin_config_*` 系列命令
- `read_plugin_ui_file`
- `src-tauri/src/main.rs` 或 `src-tauri/src/lib.rs`
- 移除 `OAuthPluginManagerState` 的初始化和注册
---
## Stage 4: 删除插件加载器和注册表
**Goal**: 删除插件系统的核心组件
**Success Criteria**: 后端编译无错误
**Tests**: 内置 OAuth 提供者正常工作
**Status**: Not Started
### 需要删除的文件:
- `src-tauri/src/credential/oauth_plugin_loader.rs` - 插件加载器
- `src-tauri/src/credential/plugin.rs` - 插件接口定义
- `src-tauri/src/credential/unified.rs` - 统一凭证接口(如果仅用于插件)
- `src-tauri/src/credential/sdk.rs` - 插件 SDK
### 需要修改的文件:
- `src-tauri/src/credential/mod.rs`
- 移除插件相关模块的导出
- 移除 `get_global_registry`, `init_global_registry` 等插件注册表函数
- `src-tauri/src/credential/registry.rs`
- 移除插件注册表相关代码
- 保留基础凭证注册功能(如果有)
---
## Stage 5: 清理数据库层
**Goal**: 删除插件相关的数据库表和 DAO
**Success Criteria**: 数据库迁移成功,后端编译无错误
**Tests**: 内置 OAuth 凭证的数据库操作正常
**Status**: Not Started
### 需要删除的表:
从 `src-tauri/src/database/schema.rs` 中删除:
- `credential_provider_plugins` - 插件元数据表
- `plugin_credentials` - 插件凭证表
- `plugin_storage` - 插件存储表
- `plugin_event_logs` - 插件事件日志表
### 需要保留的表:
- `provider_pool_credentials` - 内置 OAuth 凭证表(保留)
- 其他与内置提供者相关的表
### 需要删除的文件:
- `src-tauri/src/database/dao/plugin_credential.rs` - 插件凭证 DAO
### 需要修改的文件:
- `src-tauri/src/database/dao/mod.rs`
- 移除 `plugin_credential` 模块引用
- `src-tauri/src/database/mod.rs`
- 移除插件相关的数据库初始化代码
### 数据库迁移:
创建迁移脚本删除插件相关表:
```sql
DROP TABLE IF EXISTS plugin_event_logs;
DROP TABLE IF EXISTS plugin_storage;
DROP TABLE IF EXISTS plugin_credentials;
DROP TABLE IF EXISTS credential_provider_plugins;
```
---
## Stage 6: 清理类型定义
**Goal**: 删除插件相关的类型定义
**Success Criteria**: 代码编译无错误
**Tests**: 应用正常运行
**Status**: Not Started
### 需要修改的文件:
- `src-tauri/crates/core/src/models/provider_type.rs`
- 检查是否有插件特定的 provider 类型,如有则移除
- 保留 Kiro, Gemini, Antigravity, Codex, Claude 的类型定义
- `src-tauri/crates/core/src/models/provider_pool_model.rs`
- 移除插件凭证相关的类型定义
- 保留内置 OAuth 凭证类型
- `src/lib/plugin-sdk/types.ts`
- 如果整个目录仅用于插件 SDK,则删除整个目录
- 否则移除插件相关的类型定义
---
## Stage 7: 清理依赖项
**Goal**: 移除插件系统相关的依赖包
**Success Criteria**: 依赖安装成功,无冗余依赖
**Tests**: 应用正常启动和运行
**Status**: Not Started
### 前端 (package.json):
- 检查是否有插件系统专用的依赖,如有则移除
- 保留 OAuth 凭证管理所需的依赖
### 后端 (Cargo.toml):
- 检查 `src-tauri/Cargo.toml` 中是否有插件系统专用的 crate
- 可能需要移除的依赖:
- 动态加载相关的 crate(如 `libloading`, `dlopen` 等)
- 插件沙箱相关的 crate
---
## Stage 8: 清理文档和脚本
**Goal**: 删除插件相关的文档和脚本
**Success Criteria**: 文档目录清理完成
**Tests**: 无
**Status**: Not Started
### 需要检查和修改的文档:
- `docs/plugins/` 目录 - 如果整个目录仅用于插件文档,则删除
- `docs/content/03.providers/1.overview.md` - 移除插件系统相关的说明
- `docs/content/02.user-guide/4.configuration-example.md` - 移除插件配置示例
- `README.md` - 移除插件系统相关的说明
### 需要保留的文档:
- 关于 Kiro, Gemini, Antigravity, Codex, Claude 的 OAuth 配置文档
### 需要删除的脚本:
- 检查 `scripts/` 目录中是否有插件相关的脚本
---
## Stage 9: 最终验证和清理
**Goal**: 确保所有插件系统代码已删除,内置 OAuth 功能正常
**Success Criteria**: 所有验证项通过
**Tests**: 完整的功能测试
**Status**: Not Started
### 验证清单:
- [ ] 前端应用正常启动
- [ ] 后端应用正常启动
- [ ] 无编译错误或警告
- [ ] UI 中不再显示 "OAuth 插件" 标签页
- [ ] "OAuth 凭证" 标签页正常显示
- [ ] Kiro OAuth 凭证管理正常
- [ ] Gemini OAuth 凭证管理正常
- [ ] Antigravity OAuth 凭证管理正常
- [ ] Codex OAuth 凭证管理正常
- [ ] Claude OAuth 凭证管理正常
- [ ] 数据库迁移成功
- [ ] 代码库中不再有插件系统相关的引用
### 代码搜索验证:
使用以下关键词搜索,确保没有遗漏:
- `oauth_plugin`
- `OAuthPlugin`
- `plugin_credential`
- `PluginCredential`
- `credential_provider_plugins`
- `plugin_storage`
- `plugin_event_logs`
- `OAuthPluginLoader`
- `PluginRegistry`
---
## 注意事项
1. **备份**: 在开始之前,建议创建数据库备份和代码分支
2. **依赖检查**: 仔细检查是否有非插件代码依赖插件模块
3. **测试**: 每个阶段完成后进行编译测试
4. **提交**: 每个阶段完成后创建一个 commit
5. **渐进式**: 按阶段顺序执行,不要跳跃
## 关键区别
### 插件系统(删除)
- 可扩展架构,支持第三方插件
- 插件加载器和注册表
- 插件安装/卸载功能
- 插件 SDK 和权限系统
- 动态加载外部代码
### 内置 OAuth(保留)
- 硬编码的 5 个提供者
- 直接在代码中实现
- 不支持动态加载
- 凭证管理功能
- OAuth 流程实现
+28 -1
View File
@@ -4,20 +4,47 @@
## 架构说明
项目文档目录,包含技术规格、操作指南和文档站点配置。
项目文档目录,包含技术规格、操作指南、AI Agent 文档和文档站点配置。
使用 Nuxt Content 构建文档站点。
## 文件索引
- `aiprompts/` - AI Agent 模块文档(参考 aster-rust 模式)
- `content/` - 文档内容(Markdown)
- `develop/` - 开发文档
- `images/` - 文档图片资源
- `plugins/` - 插件文档
- `prd/` - 产品需求文档
- `tests/` - 测试文档
- `TECH_SPEC.md` - 技术规格文档
- `LLM_FLOW_MONITOR_SPEC.md` - LLM 流量监控规格
- `ops.md` - 运维操作指南
- `plugin-ui-design.md` - 插件 UI 设计文档
- `three-stage-workflow-guide.md` - 三阶段工作流指南
- `app.config.ts` - Nuxt 应用配置
- `nuxt.config.ts` - Nuxt 框架配置
- `package.json` - 文档站点依赖
## aiprompts 文档索引
AI Agent 专用文档,提供模块级别的详细说明:
- `overview.md` - 项目架构概览
- `providers.md` - Provider 系统
- `credential-pool.md` - 凭证池管理
- `converter.md` - 协议转换
- `server.md` - HTTP 服务器
- `flow-monitor.md` - 流量监控
- `components.md` - 组件系统
- `hooks.md` - React Hooks
- `services.md` - 业务服务
- `commands.md` - Tauri 命令
- `mcp.md` - MCP 服务器
- `lib.md` - 工具库
- `plugins.md` - 插件系统
- `database.md` - 数据库层
- `terminal.md` - 内置终端
## 更新提醒
任何文件变更后,请更新此文档和相关的上级文档。
+55
View File
@@ -0,0 +1,55 @@
# aiprompts
<!-- 一旦我所属的文件夹有所变化,请更新我 -->
## 架构说明
AI Agent 专用文档目录,提供模块级别的详细说明。
参考 aster-rust 的 aiprompts 模式设计。
## 文件索引
### 核心系统
- `overview.md` - 项目架构概览
- `providers.md` - Provider 系统(OAuth/API Key 认证)
- `credential-pool.md` - 凭证池管理(负载均衡、健康检查)
- `converter.md` - 协议转换(OpenAI ↔ CW/Claude)
- `server.md` - HTTP 服务器(API 端点)
### 前端模块
- `components.md` - React 组件系统
- `hooks.md` - 自定义 React Hooks
- `lib.md` - 工具库和 API 封装
### 后端模块
- `services.md` - 业务服务层
- `commands.md` - Tauri 命令
- `database.md` - 数据库层(SQLite)
### 功能模块
- `flow-monitor.md` - LLM 流量监控
- `terminal.md` - 内置终端
- `mcp.md` - MCP 服务器管理
- `plugins.md` - 插件系统
### Aster 集成
- `aster-integration.md` - **Aster 框架集成方案**
## 使用方式
AI Agent 在处理特定模块时,应先阅读对应的 aiprompts 文档:
```
# 处理 Provider 相关任务
→ 先读 docs/aiprompts/providers.md
# 处理凭证池相关任务
→ 先读 docs/aiprompts/credential-pool.md
# 处理 Aster Agent 集成
→ 先读 docs/aiprompts/aster-integration.md
```
## 更新提醒
任何文件变更后,请更新此文档和相关的上级文档。
+107
View File
@@ -0,0 +1,107 @@
# Aster 框架集成
## 集成状态 ✅
ProxyCast 已完整集成 aster-rust 框架,包括凭证池桥接。
**后端模块** (`src-tauri/src/agent/`):
- `aster_state.rs` - Agent 状态管理
- `aster_agent.rs` - Agent 包装器
- `event_converter.rs` - 事件转换器
- `credential_bridge.rs` - 凭证池桥接
**Tauri 命令** (`src-tauri/src/commands/aster_agent_cmd.rs`):
- `aster_agent_init` - 初始化 Agent
- `aster_agent_configure_provider` - 手动配置 Provider
- `aster_agent_configure_from_pool` - 从凭证池配置 Provider(推荐)
- `aster_agent_status` - 获取状态
- `aster_agent_chat_stream` - 流式对话
- `aster_agent_stop` - 停止会话
- `aster_session_create/list/get` - 会话管理
## 架构
```
┌─────────────────────────────────────────────────────────────────┐
│ 前端 (React) │
│ ┌─────────────────────────────────────────────────────────────┐│
│ │ sendAsterMessageStream / configureAsterProvider ││
│ └─────────────────────────────────────────────────────────────┘│
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Tauri Commands │
│ ┌─────────────────────────────────────────────────────────────┐│
│ │ aster_agent_cmd.rs ││
│ └─────────────────────────────────────────────────────────────┘│
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Agent 模块 │
│ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │
│ │ aster_state │ │ credential │ │ event_converter │ │
│ │ (状态管理) │ │ _bridge │ │ (事件转换) │ │
│ └──────┬──────┘ └──────┬──────┘ └─────────────────────────┘ │
│ │ │ │
│ ▼ ▼ │
│ ┌─────────────────────────────────────┐ │
│ │ ProxyCast 凭证池 │ │
│ │ - ProviderPoolService │ │
│ │ - ApiKeyProviderService │ │
│ └─────────────────────────────────────┘ │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Aster 框架 │
│ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │
│ │ Agent │ │ Provider │ │ Session │ │
│ │ (核心) │ │ (多种) │ │ (会话) │ │
│ └─────────────┘ └─────────────┘ └─────────────────────────┘ │
└─────────────────────────────────────────────────────────────────┘
```
## 凭证池桥接
### 支持的凭证类型映射
| ProxyCast 凭证类型 | Aster Provider |
|-------------------|----------------|
| OpenAIKey | openai |
| ClaudeKey / AnthropicKey | anthropic |
| KiroOAuth | bedrock |
| GeminiOAuth / GeminiApiKey | google |
| VertexKey | gcpvertexai |
| CodexOAuth | codex |
| ClaudeOAuth | anthropic |
| AntigravityOAuth | google |
### 使用方式
```typescript
// 从凭证池配置(推荐)
const status = await invoke('aster_agent_configure_from_pool', {
request: {
provider_type: 'openai',
model_name: 'gpt-4',
},
session_id: 'my-session',
});
// 流式对话
await invoke('aster_agent_chat_stream', {
request: {
message: 'Hello',
session_id: 'my-session',
event_name: 'agent_stream',
},
});
```
## 相关文档
- [overview.md](overview.md) - 项目架构
- [providers.md](providers.md) - Provider 系统
- [credential-pool.md](credential-pool.md) - 凭证池管理
+111
View File
@@ -0,0 +1,111 @@
# Tauri 命令
## 概述
Tauri 命令是前端与 Rust 后端通信的桥梁,通过 `invoke` 调用。
## 目录结构
```
src-tauri/src/commands/
├── mod.rs # 模块入口
├── credential.rs # 凭证管理命令
├── provider.rs # Provider 命令
├── server.rs # 服务器控制命令
├── flow.rs # 流量监控命令
├── config.rs # 配置命令
├── mcp.rs # MCP 服务器命令
└── terminal.rs # 终端命令
```
## 命令分类
### 凭证管理
```rust
#[tauri::command]
async fn add_credential(
provider: String,
file_path: String,
) -> Result<CredentialInfo, String>;
#[tauri::command]
async fn remove_credential(id: String) -> Result<(), String>;
#[tauri::command]
async fn list_credentials() -> Result<Vec<CredentialInfo>, String>;
#[tauri::command]
async fn refresh_credential(id: String) -> Result<(), String>;
#[tauri::command]
async fn get_credential_status(id: String) -> Result<CredentialStatus, String>;
```
### 服务器控制
```rust
#[tauri::command]
async fn start_server(config: ServerConfig) -> Result<(), String>;
#[tauri::command]
async fn stop_server() -> Result<(), String>;
#[tauri::command]
async fn get_server_status() -> Result<ServerStatus, String>;
#[tauri::command]
async fn update_server_config(config: ServerConfig) -> Result<(), String>;
```
### 流量监控
```rust
#[tauri::command]
async fn get_flow_records(query: FlowQuery) -> Result<PagedResult<FlowRecord>, String>;
#[tauri::command]
async fn get_flow_stats(time_range: TimeRange) -> Result<FlowStats, String>;
#[tauri::command]
async fn clear_flow_records(before: Option<i64>) -> Result<u64, String>;
```
## 前端调用
```typescript
import { invoke } from '@tauri-apps/api/core';
// 添加凭证
const credential = await invoke<CredentialInfo>('add_credential', {
provider: 'kiro',
filePath: '/path/to/credential.json',
});
// 获取服务器状态
const status = await invoke<ServerStatus>('get_server_status');
// 查询流量记录
const records = await invoke<PagedResult<FlowRecord>>('get_flow_records', {
query: { page: 1, pageSize: 20 },
});
```
## 错误处理
```rust
// 命令返回 Result<T, String>
// 错误信息会传递到前端
#[tauri::command]
async fn example_command() -> Result<Data, String> {
do_something()
.await
.map_err(|e| e.to_string())
}
```
## 相关文档
- [services.md](services.md) - 业务服务
- [hooks.md](hooks.md) - 前端 Hooks
+114
View File
@@ -0,0 +1,114 @@
# 组件系统
## 概述
React 组件层,使用 TailwindCSS 和 shadcn/ui。
## 目录结构
```
src/components/
├── ui/ # 基础 UI 组件 (shadcn/ui)
├── provider-pool/ # 凭证池管理
├── flow-monitor/ # 流量监控
├── general-chat/ # 通用对话
├── terminal/ # 内置终端
├── mcp/ # MCP 服务器
├── settings/ # 设置页面
└── AppSidebar.tsx # 全局侧边栏
```
## 核心组件
### AppSidebar
全局图标侧边栏,类似 cherry-studio 风格。
```tsx
// src/components/AppSidebar.tsx
export function AppSidebar() {
return (
<aside className="w-14 bg-sidebar">
<nav className="flex flex-col items-center gap-2">
<SidebarItem icon={Home} to="/" />
<SidebarItem icon={MessageSquare} to="/chat" />
<SidebarItem icon={Settings} to="/settings" />
</nav>
</aside>
);
}
```
### ProviderPool
凭证池管理组件。
```tsx
// src/components/provider-pool/ProviderPoolPanel.tsx
export function ProviderPoolPanel() {
const { credentials, addCredential, removeCredential } = useProviderPool();
return (
<div className="space-y-4">
<CredentialList credentials={credentials} onRemove={removeCredential} />
<AddCredentialDialog onAdd={addCredential} />
</div>
);
}
```
### FlowMonitor
流量监控组件。
```tsx
// src/components/flow-monitor/FlowMonitorPanel.tsx
export function FlowMonitorPanel() {
const { records, stats, query } = useFlowMonitor();
return (
<div className="flex flex-col h-full">
<FlowStats stats={stats} />
<FlowTable records={records} />
<FlowPagination query={query} />
</div>
);
}
```
## 组件规范
### 文件命名
- 组件文件: `PascalCase.tsx`
- Hook 文件: `useCamelCase.ts`
- 工具文件: `camelCase.ts`
### 组件结构
```tsx
// 标准组件结构
interface Props {
// props 定义
}
export function ComponentName({ prop1, prop2 }: Props) {
// hooks
const [state, setState] = useState();
// handlers
const handleClick = () => {};
// render
return (
<div>
{/* JSX */}
</div>
);
}
```
## 相关文档
- [hooks.md](hooks.md) - React Hooks
- [lib.md](lib.md) - 工具库
+251
View File
@@ -0,0 +1,251 @@
# 协议转换
## 概述
协议转换模块实现不同 LLM API 格式之间的双向转换,使客户端可以使用统一的 OpenAI 格式访问各种 Provider。
## 目录结构
```
src-tauri/src/converter/
├── mod.rs # 模块入口
├── protocol_selector.rs # 协议选择器
├── openai_to_cw.rs # OpenAI → CodeWhisperer
├── cw_to_openai.rs # CodeWhisperer → OpenAI
├── anthropic_to_openai.rs # Anthropic → OpenAI
└── openai_to_antigravity.rs # OpenAI → Antigravity
```
## 转换流程
```
┌─────────────────────────────────────────────────────────────────┐
│ 客户端请求 │
│ (OpenAI 格式) │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Protocol Selector │
│ 根据目标 Provider 选择转换器 │
└─────────────────────────────────────────────────────────────────┘
│
┌─────────────────────┼─────────────────────┐
▼ ▼ ▼
┌─────────────────┐ ┌─────────────────┐ ┌─────────────────┐
│ OpenAI → CW │ │ OpenAI → Claude │ │ OpenAI → AG │
└─────────────────┘ └─────────────────┘ └─────────────────┘
│ │ │
▼ ▼ ▼
┌─────────────────────────────────────────────────────────────────┐
│ Provider API │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ 响应转换 (反向) │
│ CW/Claude/AG → OpenAI │
└─────────────────────────────────────────────────────────────────┘
```
## OpenAI → CodeWhisperer
### 请求转换
```rust
// OpenAI 格式
{
"model": "gpt-4",
"messages": [
{"role": "system", "content": "..."},
{"role": "user", "content": "..."}
],
"tools": [...],
"stream": true
}
// CodeWhisperer 格式
{
"conversationState": {
"currentMessage": {
"userInputMessage": {
"content": "...",
"userInputMessageContext": {...}
}
},
"chatTriggerType": "MANUAL",
"customizationArn": "..."
}
}
```
### 工具转换
```rust
// OpenAI function tool
{
"type": "function",
"function": {
"name": "get_weather",
"parameters": {...}
}
}
// CW tool format
{
"name": "get_weather",
"inputSchema": {...}
}
```
### 特殊工具支持
| 工具类型 | OpenAI 格式 | CW 格式 |
|----------|-------------|---------|
| web_search | `{"type": "web_search"}` | 内置支持 |
| web_search_20250305 | Claude Code 格式 | 转换为 CW 格式 |
## OpenAI → Antigravity
### 请求结构
```rust
// Antigravity 请求格式 (参考 CLIProxyAPI)
{
"project": "proxycast",
"request": {
"contents": [...],
"systemInstruction": {...},
"generationConfig": {...},
"tools": [...],
"safetySettings": [...]
},
"model": "gemini-2.0-flash"
}
```
### 工具定义转换
```rust
// OpenAI 格式
{
"type": "function",
"function": {
"name": "tool_name",
"parameters": {...}
}
}
// Antigravity 格式
{
"functionDeclarations": [{
"name": "tool_name",
"parametersJsonSchema": {...} // 注意字段名变化
}]
}
```
### 安全设置
```rust
// 默认安全设置
const DEFAULT_SAFETY_SETTINGS: &[SafetySetting] = &[
SafetySetting {
category: "HARM_CATEGORY_HATE_SPEECH",
threshold: "OFF",
},
SafetySetting {
category: "HARM_CATEGORY_DANGEROUS_CONTENT",
threshold: "OFF",
},
// ...
];
```
## Anthropic → OpenAI
### 响应转换
```rust
// Anthropic 响应
{
"content": [
{"type": "text", "text": "..."},
{"type": "tool_use", "id": "...", "name": "...", "input": {...}}
],
"stop_reason": "end_turn"
}
// OpenAI 响应
{
"choices": [{
"message": {
"role": "assistant",
"content": "...",
"tool_calls": [...]
},
"finish_reason": "stop"
}]
}
```
## 流式响应处理
### SSE 格式转换
```rust
// OpenAI SSE
data: {"choices":[{"delta":{"content":"Hello"}}]}
// CW SSE
data: {"messageMetadata":{"..."},"assistantResponseEvent":{"content":"Hello"}}
```
### 转换器实现
```rust
pub struct StreamConverter {
buffer: String,
state: StreamState,
}
impl StreamConverter {
pub fn process_chunk(&mut self, chunk: &str) -> Vec<String> {
// 解析 SSE 事件
// 转换格式
// 返回 OpenAI 格式的 SSE 事件
}
}
```
## 错误处理
### 错误映射
| Provider 错误 | OpenAI 错误码 |
|---------------|---------------|
| CW ThrottlingException | 429 |
| CW ValidationException | 400 |
| Claude rate_limit_error | 429 |
| Claude invalid_request_error | 400 |
### 错误转换
```rust
pub fn convert_error(provider_error: ProviderError) -> OpenAIError {
match provider_error {
ProviderError::RateLimit => OpenAIError {
code: 429,
message: "Rate limit exceeded",
type_: "rate_limit_error",
},
// ...
}
}
```
## 相关文档
- [providers.md](providers.md) - Provider 系统
- [server.md](server.md) - HTTP 服务器
- [streaming.md](streaming.md) - 流式处理
+222
View File
@@ -0,0 +1,222 @@
# 凭证池管理
## 概述
凭证池管理系统实现多凭证轮询负载均衡、健康检查和自动 Token 刷新。
## 核心组件
```
src-tauri/src/
├── credential/ # 凭证池核心
│ ├── mod.rs
│ ├── pool.rs # 凭证池实现
│ └── health.rs # 健康检查
└── services/
├── provider_pool_service.rs # 池服务
└── token_cache_service.rs # Token 缓存
```
## 凭证池架构
```
┌─────────────────────────────────────────────────────────────────┐
│ ProviderPoolService │
│ ┌─────────────────────────────────────────────────────────────┐│
│ │ Credential Pool ││
│ │ ┌─────────┐ ┌─────────┐ ┌─────────┐ ┌─────────┐ ││
│ │ │ Cred 1 │ │ Cred 2 │ │ Cred 3 │ │ Cred N │ ││
│ │ │ Healthy │ │ Healthy │ │ Expired │ │ Healthy │ ││
│ │ └────┬────┘ └────┬────┘ └────┬────┘ └────┬────┘ ││
│ │ │ │ │ │ ││
│ │ └────────────┴────────────┴────────────┘ ││
│ │ │ ││
│ │ Round Robin ││
│ └─────────────────────────┼───────────────────────────────────┘│
│ │ │
│ ┌─────────────────────────┼───────────────────────────────────┐│
│ │ Health Checker (定时任务) ││
│ │ - Token 过期检查 ││
│ │ - 自动刷新 ││
│ │ - 不健康凭证剔除 ││
│ └─────────────────────────────────────────────────────────────┘│
└─────────────────────────────────────────────────────────────────┘
```
## 负载均衡策略
### Round Robin (轮询)
```rust
pub struct RoundRobinPool {
credentials: Vec<CredentialEntry>,
current_index: AtomicUsize,
}
impl RoundRobinPool {
pub fn next(&self) -> Option<&CredentialEntry> {
let healthy: Vec<_> = self.credentials
.iter()
.filter(|c| c.is_healthy())
.collect();
if healthy.is_empty() {
return None;
}
let index = self.current_index
.fetch_add(1, Ordering::Relaxed) % healthy.len();
Some(healthy[index])
}
}
```
### 权重轮询 (可选)
```rust
pub struct WeightedPool {
credentials: Vec<(CredentialEntry, u32)>, // (凭证, 权重)
}
```
## 健康检查
### 检查项目
| 检查项 | 说明 | 频率 |
|--------|------|------|
| Token 过期 | 检查 expires_at | 每次请求前 |
| Token 刷新 | 尝试刷新过期 Token | Token 过期时 |
| API 可用性 | 发送测试请求 | 定时 (5分钟) |
### 健康状态
```rust
pub enum HealthStatus {
Healthy, // 健康
TokenExpired, // Token 过期
TokenRefreshing, // 正在刷新
RefreshFailed(String), // 刷新失败
Unhealthy(String), // 不健康
Disabled, // 已禁用
}
```
### 自动恢复
```rust
// 健康检查任务
async fn health_check_task(pool: Arc<ProviderPoolService>) {
loop {
for credential in pool.credentials() {
match credential.health_status() {
HealthStatus::TokenExpired => {
// 尝试刷新
if let Err(e) = pool.refresh_token(&credential).await {
credential.set_status(HealthStatus::RefreshFailed(e));
}
}
HealthStatus::RefreshFailed(_) => {
// 重试刷新 (最多 3 次)
if credential.retry_count() < 3 {
pool.retry_refresh(&credential).await;
}
}
_ => {}
}
}
tokio::time::sleep(Duration::from_secs(300)).await;
}
}
```
## Token 缓存
### 缓存策略
```rust
pub struct TokenCacheService {
cache: DashMap<String, CachedToken>,
}
struct CachedToken {
access_token: String,
expires_at: i64,
refresh_token: String,
}
impl TokenCacheService {
pub async fn get_or_refresh(&self, credential_id: &str) -> Result<String> {
if let Some(cached) = self.cache.get(credential_id) {
if !cached.is_expired() {
return Ok(cached.access_token.clone());
}
}
// 刷新并缓存
let new_token = self.refresh(credential_id).await?;
self.cache.insert(credential_id.to_string(), new_token.clone());
Ok(new_token.access_token)
}
}
```
### 数据库持久化
```sql
CREATE TABLE token_cache (
credential_id TEXT PRIMARY KEY,
access_token TEXT NOT NULL,
refresh_token TEXT NOT NULL,
expires_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
```
## 凭证生命周期
```
┌─────────┐ ┌─────────┐ ┌─────────┐ ┌─────────┐
│ 上传 │ ──▶ │ 验证 │ ──▶ │ 激活 │ ──▶ │ 使用中 │
└─────────┘ └─────────┘ └─────────┘ └────┬────┘
│
┌────────────────────────────────┘
│
▼
┌─────────┐ ┌─────────┐ ┌─────────┐
│ 过期 │ ──▶ │ 刷新 │ ──▶ │ 恢复 │
└─────────┘ └────┬────┘ └─────────┘
│
▼ (失败)
┌─────────┐
│ 禁用 │
└─────────┘
```
## API 接口
### Tauri Commands
```rust
#[tauri::command]
async fn add_credential(provider: String, path: String) -> Result<()>;
#[tauri::command]
async fn remove_credential(id: String) -> Result<()>;
#[tauri::command]
async fn list_credentials() -> Result<Vec<CredentialInfo>>;
#[tauri::command]
async fn refresh_credential(id: String) -> Result<()>;
#[tauri::command]
async fn get_pool_status() -> Result<PoolStatus>;
```
## 相关文档
- [providers.md](providers.md) - Provider 系统
- [services.md](services.md) - 业务服务
- [database.md](database.md) - 数据库层
+106
View File
@@ -0,0 +1,106 @@
# 数据库层
## 概述
使用 SQLite (rusqlite) 存储凭证元数据、流量记录等。
## 目录结构
```
src-tauri/src/database/
├── mod.rs # 模块入口
├── schema.rs # 表结构定义
├── migrations.rs # 数据库迁移
└── dao/ # 数据访问对象
├── credential_dao.rs
├── flow_dao.rs
└── config_dao.rs
```
## 表结构
### credentials
```sql
CREATE TABLE credentials (
id TEXT PRIMARY KEY,
provider TEXT NOT NULL,
name TEXT NOT NULL,
file_path TEXT NOT NULL,
status TEXT NOT NULL DEFAULT 'active',
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
```
### token_cache
```sql
CREATE TABLE token_cache (
credential_id TEXT PRIMARY KEY,
access_token TEXT NOT NULL,
refresh_token TEXT,
expires_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
FOREIGN KEY (credential_id) REFERENCES credentials(id)
);
```
### flow_records
```sql
CREATE TABLE flow_records (
id TEXT PRIMARY KEY,
timestamp INTEGER NOT NULL,
provider TEXT NOT NULL,
model TEXT NOT NULL,
request_json TEXT NOT NULL,
response_json TEXT,
status TEXT NOT NULL,
latency_ms INTEGER,
prompt_tokens INTEGER,
completion_tokens INTEGER,
created_at INTEGER NOT NULL
);
CREATE INDEX idx_flow_timestamp ON flow_records(timestamp);
```
## DAO 模式
```rust
pub struct CredentialDao {
conn: Arc<Mutex<Connection>>,
}
impl CredentialDao {
pub fn insert(&self, credential: &Credential) -> Result<()>;
pub fn find_by_id(&self, id: &str) -> Result<Option<Credential>>;
pub fn find_all(&self) -> Result<Vec<Credential>>;
pub fn update(&self, credential: &Credential) -> Result<()>;
pub fn delete(&self, id: &str) -> Result<()>;
}
```
## 数据库迁移
```rust
pub fn run_migrations(conn: &Connection) -> Result<()> {
let version = get_schema_version(conn)?;
if version < 1 {
conn.execute_batch(include_str!("migrations/001_initial.sql"))?;
}
if version < 2 {
conn.execute_batch(include_str!("migrations/002_add_flow.sql"))?;
}
set_schema_version(conn, CURRENT_VERSION)?;
Ok(())
}
```
## 相关文档
- [services.md](services.md) - 业务服务
- [credential-pool.md](credential-pool.md) - 凭证池管理
+261
View File
@@ -0,0 +1,261 @@
# 流量监控
## 概述
流量监控模块拦截和记录所有 LLM API 请求,提供 Token 统计、历史查询和分析功能。
## 目录结构
```
src-tauri/src/flow_monitor/
├── mod.rs # 模块入口
├── interceptor.rs # 请求拦截器
├── storage.rs # 存储层
├── query.rs # 查询接口
└── stats.rs # 统计分析
```
## 架构
```
┌─────────────────────────────────────────────────────────────────┐
│ HTTP 请求 │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Flow Interceptor │
│ ┌─────────────────────────────────────────────────────────────┐│
│ │ 请求拦截 ││
│ │ - 请求 ID 生成 ││
│ │ - 请求体捕获 ││
│ │ - 时间戳记录 ││
│ └─────────────────────────────────────────────────────────────┘│
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Provider 处理 │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Flow Interceptor │
│ ┌─────────────────────────────────────────────────────────────┐│
│ │ 响应拦截 ││
│ │ - 响应体捕获 ││
│ │ - Token 计数 ││
│ │ - 延迟计算 ││
│ └─────────────────────────────────────────────────────────────┘│
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Flow Storage │
│ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │
│ │ SQLite │ │ 内存缓存 │ │ 事件发送 │ │
│ │ 持久化 │ │ (最近 N 条) │ │ (前端通知) │ │
│ └─────────────┘ └─────────────┘ └─────────────────────────┘ │
└─────────────────────────────────────────────────────────────────┘
```
## 数据模型
### FlowRecord
```rust
pub struct FlowRecord {
pub id: String, // 请求 ID
pub timestamp: i64, // 时间戳
pub provider: String, // Provider 类型
pub model: String, // 模型名称
pub request: FlowRequest, // 请求数据
pub response: Option<FlowResponse>, // 响应数据
pub status: FlowStatus, // 状态
pub latency_ms: Option<u64>, // 延迟 (毫秒)
}
pub struct FlowRequest {
pub messages: Vec<Message>, // 消息列表
pub tools: Option<Vec<Tool>>, // 工具定义
pub stream: bool, // 是否流式
}
pub struct FlowResponse {
pub content: String, // 响应内容
pub tool_calls: Option<Vec<ToolCall>>, // 工具调用
pub usage: TokenUsage, // Token 使用
}
pub struct TokenUsage {
pub prompt_tokens: u32, // 输入 Token
pub completion_tokens: u32, // 输出 Token
pub total_tokens: u32, // 总 Token
}
```
### 数据库表
```sql
CREATE TABLE flow_records (
id TEXT PRIMARY KEY,
timestamp INTEGER NOT NULL,
provider TEXT NOT NULL,
model TEXT NOT NULL,
request_json TEXT NOT NULL,
response_json TEXT,
status TEXT NOT NULL,
latency_ms INTEGER,
prompt_tokens INTEGER,
completion_tokens INTEGER,
total_tokens INTEGER,
created_at INTEGER NOT NULL
);
CREATE INDEX idx_flow_timestamp ON flow_records(timestamp);
CREATE INDEX idx_flow_provider ON flow_records(provider);
CREATE INDEX idx_flow_model ON flow_records(model);
```
## 拦截器实现
```rust
pub struct FlowInterceptor {
storage: Arc<FlowStorage>,
event_sender: mpsc::Sender<FlowEvent>,
}
impl FlowInterceptor {
pub async fn intercept_request(&self, req: &Request) -> String {
let request_id = generate_uuid();
let record = FlowRecord {
id: request_id.clone(),
timestamp: current_timestamp(),
provider: extract_provider(req),
model: extract_model(req),
request: parse_request(req),
response: None,
status: FlowStatus::Pending,
latency_ms: None,
};
self.storage.insert(&record).await;
self.event_sender.send(FlowEvent::RequestStarted(record)).await;
request_id
}
pub async fn intercept_response(
&self,
request_id: &str,
response: &Response,
latency: Duration,
) {
let flow_response = parse_response(response);
self.storage.update_response(
request_id,
&flow_response,
latency.as_millis() as u64,
).await;
self.event_sender.send(FlowEvent::ResponseReceived {
request_id: request_id.to_string(),
response: flow_response,
}).await;
}
}
```
## 查询接口
### 分页查询
```rust
pub struct FlowQuery {
pub provider: Option<String>,
pub model: Option<String>,
pub start_time: Option<i64>,
pub end_time: Option<i64>,
pub status: Option<FlowStatus>,
pub page: u32,
pub page_size: u32,
}
pub async fn query_flows(query: FlowQuery) -> Result<PagedResult<FlowRecord>> {
// 构建 SQL 查询
// 执行分页查询
// 返回结果
}
```
### 统计查询
```rust
pub struct FlowStats {
pub total_requests: u64,
pub total_tokens: u64,
pub avg_latency_ms: f64,
pub by_provider: HashMap<String, ProviderStats>,
pub by_model: HashMap<String, ModelStats>,
}
pub async fn get_stats(time_range: TimeRange) -> Result<FlowStats> {
// 聚合统计
}
```
## 前端事件
### 事件类型
```typescript
interface FlowEvent {
type: 'request_started' | 'response_received' | 'error';
data: FlowRecord;
}
```
### 事件监听
```typescript
// 前端监听
import { listen } from '@tauri-apps/api/event';
listen<FlowEvent>('flow-event', (event) => {
switch (event.payload.type) {
case 'request_started':
addPendingRequest(event.payload.data);
break;
case 'response_received':
updateRequest(event.payload.data);
break;
}
});
```
## Tauri Commands
```rust
#[tauri::command]
async fn get_flow_records(query: FlowQuery) -> Result<PagedResult<FlowRecord>>;
#[tauri::command]
async fn get_flow_stats(time_range: TimeRange) -> Result<FlowStats>;
#[tauri::command]
async fn get_flow_detail(id: String) -> Result<FlowRecord>;
#[tauri::command]
async fn clear_flow_records(before: Option<i64>) -> Result<u64>;
#[tauri::command]
async fn export_flow_records(format: ExportFormat) -> Result<String>;
```
## 相关文档
- [server.md](server.md) - HTTP 服务器
- [database.md](database.md) - 数据库层
- [components.md](components.md) - 前端组件
+116
View File
@@ -0,0 +1,116 @@
# React Hooks
## 概述
自定义 Hooks 封装业务逻辑,通过 Tauri invoke 与后端通信。
## 目录结构
```
src/hooks/
├── index.ts # 导出入口
├── useProviderPool.ts # 凭证池管理
├── useOAuthCredentials.ts # OAuth 凭证
├── useFlowEvents.ts # 流量事件
├── useMcpServers.ts # MCP 服务器
├── useDeepLink.ts # Deep Link 处理
├── useSound.ts # 音效管理
└── useTauri.ts # Tauri 通用
```
## 核心 Hooks
### useProviderPool
```typescript
export function useProviderPool() {
const [credentials, setCredentials] = useState<Credential[]>([]);
const [loading, setLoading] = useState(false);
const refresh = async () => {
setLoading(true);
const list = await invoke<Credential[]>('list_credentials');
setCredentials(list);
setLoading(false);
};
const addCredential = async (provider: string, path: string) => {
await invoke('add_credential', { provider, filePath: path });
await refresh();
};
const removeCredential = async (id: string) => {
await invoke('remove_credential', { id });
await refresh();
};
useEffect(() => { refresh(); }, []);
return { credentials, loading, addCredential, removeCredential, refresh };
}
```
### useFlowEvents
```typescript
export function useFlowEvents() {
const [records, setRecords] = useState<FlowRecord[]>([]);
useEffect(() => {
const unlisten = listen<FlowEvent>('flow-event', (event) => {
setRecords(prev => [event.payload.data, ...prev].slice(0, 100));
});
return () => { unlisten.then(fn => fn()); };
}, []);
return { records };
}
```
### useDeepLink
```typescript
export function useDeepLink() {
useEffect(() => {
const unlisten = listen<string>('deep-link', async (event) => {
const url = new URL(event.payload);
if (url.pathname === '/oauth/callback') {
await handleOAuthCallback(url.searchParams);
}
});
return () => { unlisten.then(fn => fn()); };
}, []);
}
```
## Hook 规范
### 命名约定
- 以 `use` 开头
- 描述功能: `useProviderPool`, `useFlowEvents`
### 返回值
```typescript
// 返回对象,包含状态和操作
return {
// 状态
data,
loading,
error,
// 操作
refresh,
add,
remove,
};
```
## 相关文档
- [components.md](components.md) - 组件系统
- [commands.md](commands.md) - Tauri 命令
+89
View File
@@ -0,0 +1,89 @@
# 工具库
## 概述
前端工具库和 API 封装层。
## 目录结构
```
src/lib/
├── api/ # API 封装
│ ├── apiKeyProvider.ts
│ └── pluginUI.ts
├── config/ # 配置
│ └── providers.ts
├── types/ # 类型定义
│ └── provider.ts
├── errors/ # 错误处理
│ └── playwrightErrors.ts
├── plugin-ui/ # 插件 UI 系统
├── tauri/ # Tauri 命令封装
├── utils/ # 工具函数
├── flowEventManager.ts # 流量事件管理
├── terminal-api.ts # 终端 API
└── utils.ts # 通用工具
```
## 核心模块
### Tauri 命令封装
```typescript
// src/lib/tauri/credentials.ts
export async function addCredential(provider: string, path: string) {
return invoke<CredentialInfo>('add_credential', { provider, filePath: path });
}
export async function listCredentials() {
return invoke<CredentialInfo[]>('list_credentials');
}
```
### 流量事件管理
```typescript
// src/lib/flowEventManager.ts
class FlowEventManager {
private listeners: Map<string, Set<FlowEventListener>> = new Map();
subscribe(event: string, listener: FlowEventListener) {
if (!this.listeners.has(event)) {
this.listeners.set(event, new Set());
}
this.listeners.get(event)!.add(listener);
return () => this.listeners.get(event)?.delete(listener);
}
emit(event: string, data: any) {
this.listeners.get(event)?.forEach(listener => listener(data));
}
}
export const flowEventManager = new FlowEventManager();
```
### 工具函数
```typescript
// src/lib/utils.ts
export function cn(...classes: (string | undefined)[]) {
return classes.filter(Boolean).join(' ');
}
export function formatBytes(bytes: number) {
if (bytes < 1024) return `${bytes} B`;
if (bytes < 1024 * 1024) return `${(bytes / 1024).toFixed(1)} KB`;
return `${(bytes / 1024 / 1024).toFixed(1)} MB`;
}
export function formatDuration(ms: number) {
if (ms < 1000) return `${ms}ms`;
return `${(ms / 1000).toFixed(2)}s`;
}
```
## 相关文档
- [hooks.md](hooks.md) - React Hooks
- [components.md](components.md) - 组件系统
+94
View File
@@ -0,0 +1,94 @@
# MCP 服务器
## 概述
MCP (Model Context Protocol) 服务器管理模块。
## 目录结构
```
src-tauri/src/services/
├── mcp_service.rs # MCP 服务管理
└── mcp_sync.rs # 配置同步
src/components/mcp/
├── McpPanel.tsx # MCP 管理面板
├── McpServerList.tsx # 服务器列表
└── McpToolList.tsx # 工具列表
```
## MCP 服务
```rust
pub struct McpService {
servers: HashMap<String, McpServer>,
config_path: PathBuf,
}
pub struct McpServer {
name: String,
command: String,
args: Vec<String>,
env: HashMap<String, String>,
status: ServerStatus,
tools: Vec<Tool>,
}
impl McpService {
/// 启动服务器
pub async fn start(&mut self, name: &str) -> Result<()>;
/// 停止服务器
pub async fn stop(&mut self, name: &str) -> Result<()>;
/// 列出工具
pub async fn list_tools(&self, name: &str) -> Result<Vec<Tool>>;
/// 调用工具
pub async fn call_tool(
&self,
server: &str,
tool: &str,
args: Value,
) -> Result<Value>;
}
```
## 配置格式
```json
{
"mcpServers": {
"filesystem": {
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-filesystem", "/path"],
"env": {},
"disabled": false
}
}
}
```
## Tauri 命令
```rust
#[tauri::command]
async fn mcp_list_servers() -> Result<Vec<McpServerInfo>, String>;
#[tauri::command]
async fn mcp_start_server(name: String) -> Result<(), String>;
#[tauri::command]
async fn mcp_stop_server(name: String) -> Result<(), String>;
#[tauri::command]
async fn mcp_list_tools(server: String) -> Result<Vec<Tool>, String>;
#[tauri::command]
async fn mcp_call_tool(server: String, tool: String, args: Value) -> Result<Value, String>;
```
## 相关文档
- [services.md](services.md) - 业务服务
- [commands.md](commands.md) - Tauri 命令
+138
View File
@@ -0,0 +1,138 @@
# ProxyCast 项目架构概览
## 概述
ProxyCast 是一个 Tauri 桌面应用,作为 LLM API 代理网关,支持多 Provider 凭证池管理、协议转换、流量监控等功能。
## 项目结构
```
proxycast/
├── src/ # React 前端
│ ├── components/ # UI 组件
│ ├── pages/ # 页面组件
│ ├── hooks/ # React Hooks
│ ├── lib/ # 工具库
│ └── stores/ # 状态管理
├── src-tauri/ # Rust 后端
│ └── src/
│ ├── commands/ # Tauri 命令
│ ├── providers/ # Provider 实现
│ ├── services/ # 业务服务
│ ├── converter/ # 协议转换
│ ├── server/ # HTTP 服务器
│ └── ...
├── plugins/ # 插件目录
└── docs/ # 文档
```
## 核心模块
### 后端 (src-tauri/src/)
| 模块 | 说明 |
|------|------|
| `providers/` | LLM Provider 认证和 API 实现 |
| `services/` | 业务服务层 |
| `converter/` | 协议转换 (OpenAI ↔ CW/Claude) |
| `server/` | HTTP API 服务器 |
| `credential/` | 凭证池管理 |
| `flow_monitor/` | 流量监控 |
| `terminal/` | 内置终端 |
### 前端 (src/)
| 模块 | 说明 |
|------|------|
| `components/` | React 组件 |
| `hooks/` | 业务逻辑 Hooks |
| `lib/` | 工具函数和 API 封装 |
| `pages/` | 页面组件 |
## 数据流
```
┌─────────────────────────────────────────────────────────────────┐
│ 客户端请求 (Cursor/Continue) │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ HTTP Server │
│ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │
│ │ OpenAI API │ │ Claude API │ │ Flow Monitor │ │
│ │ 兼容端点 │ │ 兼容端点 │ │ (请求拦截) │ │
│ └──────┬──────┘ └──────┬──────┘ └───────────┬─────────────┘ │
└─────────┼────────────────┼─────────────────────┼────────────────┘
│ │ │
▼ ▼ ▼
┌─────────────────────────────────────────────────────────────────┐
│ Router / Processor │
│ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │
│ │ 模型路由 │ │ 协议转换 │ │ 弹性策略 │ │
│ │ (规则匹配) │ │ (Converter) │ │ (重试/超时) │ │
│ └──────┬──────┘ └──────┬──────┘ └───────────┬─────────────┘ │
└─────────┼────────────────┼─────────────────────┼────────────────┘
│ │ │
▼ ▼ ▼
┌─────────────────────────────────────────────────────────────────┐
│ Provider Pool Service │
│ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │
│ │ 凭证轮询 │ │ 健康检查 │ │ Token 刷新 │ │
│ │ (负载均衡) │ │ (自动剔除) │ │ (OAuth) │ │
│ └──────┬──────┘ └──────┬──────┘ └───────────┬─────────────┘ │
└─────────┼────────────────┼─────────────────────┼────────────────┘
│ │ │
▼ ▼ ▼
┌─────────────────────────────────────────────────────────────────┐
│ Providers │
│ ┌─────────┐ ┌─────────┐ ┌─────────┐ ┌─────────┐ │
│ │ Kiro │ │ Gemini │ │ Claude │ │ OpenAI │ ... │
│ └─────────┘ └─────────┘ └─────────┘ └─────────┘ │
└─────────────────────────────────────────────────────────────────┘
```
## 关键特性
### 1. 多 Provider 支持
- OAuth: Kiro, Gemini, Qwen, Claude, Antigravity
- API Key: OpenAI, Claude, 自定义
### 2. 凭证池管理
- 多凭证轮询负载均衡
- 自动健康检查和剔除
- Token 自动刷新
### 3. 协议转换
- OpenAI ↔ CodeWhisperer
- OpenAI ↔ Claude
- OpenAI ↔ Antigravity
### 4. 流量监控
- 请求/响应拦截
- Token 统计
- 历史查询
## 文档索引
### 核心系统
- [providers.md](providers.md) - Provider 系统
- [credential-pool.md](credential-pool.md) - 凭证池管理
- [converter.md](converter.md) - 协议转换
- [server.md](server.md) - HTTP 服务器
### 前端模块
- [components.md](components.md) - 组件系统
- [hooks.md](hooks.md) - React Hooks
- [lib.md](lib.md) - 工具库
### 功能模块
- [flow-monitor.md](flow-monitor.md) - 流量监控
- [terminal.md](terminal.md) - 内置终端
- [mcp.md](mcp.md) - MCP 服务器
- [plugins.md](plugins.md) - 插件系统
### 配置与服务
- [commands.md](commands.md) - Tauri 命令
- [services.md](services.md) - 业务服务
- [database.md](database.md) - 数据库层
+88
View File
@@ -0,0 +1,88 @@
# 插件系统
## 概述
插件系统支持扩展 ProxyCast 功能,包含声明式 UI 系统。
## 目录结构
```
src-tauri/src/plugin/
├── mod.rs # 模块入口
├── loader.rs # 插件加载器
├── runtime.rs # 插件运行时
└── ui/ # 声明式 UI
├── types.rs
└── renderer.rs
plugins/ # 插件目录
└── example/
├── manifest.json
└── main.js
```
## 插件清单
```json
{
"name": "example-plugin",
"version": "1.0.0",
"description": "示例插件",
"main": "main.js",
"permissions": ["network", "storage"],
"ui": {
"settings": "settings.json"
}
}
```
## 声明式 UI
```json
{
"type": "form",
"fields": [
{
"name": "apiKey",
"type": "password",
"label": "API Key",
"required": true
},
{
"name": "enabled",
"type": "switch",
"label": "启用",
"default": true
}
]
}
```
## 插件 API
```typescript
// 插件可用的 API
interface PluginAPI {
// 存储
storage: {
get(key: string): Promise<any>;
set(key: string, value: any): Promise<void>;
};
// 网络
http: {
fetch(url: string, options?: RequestInit): Promise<Response>;
};
// UI
ui: {
showNotification(message: string): void;
showDialog(options: DialogOptions): Promise<any>;
};
}
```
## 相关文档
- [components.md](components.md) - 组件系统
- [services.md](services.md) - 业务服务
+189
View File
@@ -0,0 +1,189 @@
# Provider 系统
## 概述
Provider 系统负责与各 LLM 服务商的认证和 API 交互。支持 OAuth 和 API Key 两种认证方式。
## 目录结构
```
src-tauri/src/providers/
├── mod.rs # 模块入口和 Provider 枚举
├── traits.rs # Provider trait 定义
├── error.rs # 错误类型
├── kiro.rs # Kiro/CodeWhisperer OAuth
├── gemini.rs # Gemini OAuth
├── qwen.rs # Qwen OAuth
├── antigravity.rs # Antigravity OAuth
├── claude_oauth.rs # Claude OAuth
├── claude_custom.rs # Claude API Key
├── openai_custom.rs # OpenAI API Key
├── codex.rs # Codex Provider
├── iflow.rs # iFlow Provider
├── vertex.rs # Vertex AI Provider
└── tests.rs # 单元测试
```
## Provider 枚举
```rust
pub enum ProviderType {
Kiro, // Kiro/CodeWhisperer OAuth
Gemini, // Google Gemini OAuth
Qwen, // 通义千问 OAuth
Antigravity, // Antigravity (Gemini CLI) OAuth
ClaudeOAuth, // Claude OAuth
ClaudeCustom, // Claude API Key
OpenAICustom, // OpenAI API Key
Codex, // Codex
IFlow, // iFlow
Vertex, // Vertex AI
}
```
## Provider Trait
```rust
pub trait Provider: Send + Sync {
/// 获取 Provider 类型
fn provider_type(&self) -> ProviderType;
/// 加载凭证
async fn load_credential(&self, path: &Path) -> Result<CredentialData>;
/// 刷新 Token
async fn refresh_token(&self, credential: &mut CredentialData) -> Result<()>;
/// 检查 Token 是否过期
fn is_token_expired(&self, credential: &CredentialData) -> bool;
/// 发送 API 请求
async fn send_request(&self, credential: &CredentialData, request: &Request) -> Result<Response>;
}
```
## OAuth Provider 实现
### Kiro Provider
```rust
// 凭证文件结构
struct KiroCredential {
access_token: String,
refresh_token: String,
expires_at: i64,
client_id: Option<String>, // 从 clientIdHash 合并
client_secret: Option<String>, // 从 clientIdHash 合并
}
// Token 刷新流程
1. 检查 expires_at 是否过期
2. 使用 refresh_token 请求新 token
3. 更新凭证文件
```
### Gemini Provider
```rust
// OAuth 端点
const AUTH_URL: &str = "https://accounts.google.com/o/oauth2/v2/auth";
const TOKEN_URL: &str = "https://oauth2.googleapis.com/token";
// 凭证文件结构
struct GeminiCredential {
access_token: String,
refresh_token: String,
expires_at: i64,
}
```
## API Key Provider 实现
### OpenAI Custom
```rust
// 凭证结构
struct OpenAICredential {
api_key: String,
base_url: Option<String>, // 自定义端点
}
// 请求头
Authorization: Bearer {api_key}
```
### Claude Custom
```rust
// 凭证结构
struct ClaudeCredential {
api_key: String,
base_url: Option<String>,
}
// 请求头
x-api-key: {api_key}
anthropic-version: 2023-06-01
```
## 凭证管理策略
### 方案 B: 独立副本策略
```
原始凭证文件 (用户上传)
│
▼
┌─────────────────────────────────────┐
│ 合并 clientIdHash 中的 │
│ client_id / client_secret │
└─────────────────────────────────────┘
│
▼
副本凭证文件 (credentials/ 目录)
│
▼
独立刷新和管理
```
优点:
- 每个副本完全独立
- 支持多账号场景
- 不影响原始文件
## 健康检查
```rust
// 健康检查逻辑
async fn health_check(&self, credential: &CredentialData) -> HealthStatus {
// 1. 检查 Token 是否过期
if self.is_token_expired(credential) {
return HealthStatus::TokenExpired;
}
// 2. 尝试刷新 Token
if let Err(e) = self.refresh_token(credential).await {
return HealthStatus::RefreshFailed(e);
}
// 3. 发送测试请求
match self.send_test_request(credential).await {
Ok(_) => HealthStatus::Healthy,
Err(e) => HealthStatus::Unhealthy(e),
}
}
```
## 添加新 Provider
1. 在 `providers/` 创建新模块文件
2. 实现 `Provider` trait
3. 在 `ProviderType` 枚举添加新类型
4. 在 `ProviderPoolService` 注册健康检查
5. 更新前端 Provider 选择器
## 相关文档
- [credential-pool.md](credential-pool.md) - 凭证池管理
- [converter.md](converter.md) - 协议转换
- [server.md](server.md) - HTTP 服务器
+252
View File
@@ -0,0 +1,252 @@
# HTTP 服务器
## 概述
HTTP 服务器提供 OpenAI 和 Claude 兼容的 API 端点,支持流式响应。
## 目录结构
```
src-tauri/src/
├── server/
│ ├── mod.rs # 服务器入口
│ ├── routes.rs # 路由定义
│ ├── handlers.rs # 请求处理器
│ └── middleware.rs # 中间件
├── server_utils.rs # 工具函数
└── streaming/ # 流式响应
├── mod.rs
└── sse.rs
```
## API 端点
### OpenAI 兼容端点
| 端点 | 方法 | 说明 |
|------|------|------|
| `/v1/chat/completions` | POST | 聊天补全 |
| `/v1/models` | GET | 模型列表 |
| `/v1/embeddings` | POST | 文本嵌入 |
### Claude 兼容端点
| 端点 | 方法 | 说明 |
|------|------|------|
| `/v1/messages` | POST | 消息 API |
### 管理端点
| 端点 | 方法 | 说明 |
|------|------|------|
| `/health` | GET | 健康检查 |
| `/metrics` | GET | 指标统计 |
## 请求处理流程
```
┌─────────────────────────────────────────────────────────────────┐
│ HTTP 请求 │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Middleware │
│ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │
│ │ 认证 │ │ 日志 │ │ 流量监控 │ │
│ └─────────────┘ └─────────────┘ └─────────────────────────┘ │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Router │
│ 根据路径和模型选择处理器 │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Handler │
│ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │
│ │ 请求验证 │ │ 协议转换 │ │ Provider 调用 │ │
│ └─────────────┘ └─────────────┘ └─────────────────────────┘ │
└─────────────────────────────────────────────────────────────────┘
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ 响应 │
│ ┌─────────────┐ ┌─────────────┐ │
│ │ JSON 响应 │ │ SSE 流式 │ │
│ └─────────────┘ └─────────────┘ │
└─────────────────────────────────────────────────────────────────┘
```
## 服务器配置
```rust
pub struct ServerConfig {
pub host: String, // 监听地址
pub port: u16, // 监听端口
pub cors_enabled: bool, // CORS 支持
pub max_body_size: usize, // 最大请求体
pub timeout: Duration, // 请求超时
}
// 默认配置
impl Default for ServerConfig {
fn default() -> Self {
Self {
host: "127.0.0.1".to_string(),
port: 8080,
cors_enabled: true,
max_body_size: 10 * 1024 * 1024, // 10MB
timeout: Duration::from_secs(300),
}
}
}
```
## 中间件
### 认证中间件
```rust
pub async fn auth_middleware(req: Request, next: Next) -> Response {
// 检查 Authorization header
let auth_header = req.headers().get("Authorization");
match auth_header {
Some(value) => {
// 验证 Bearer token
if validate_token(value) {
next.run(req).await
} else {
Response::unauthorized()
}
}
None => Response::unauthorized(),
}
}
```
### 流量监控中间件
```rust
pub async fn flow_monitor_middleware(req: Request, next: Next) -> Response {
let start = Instant::now();
let request_id = generate_request_id();
// 记录请求
flow_monitor.record_request(&request_id, &req).await;
let response = next.run(req).await;
// 记录响应
flow_monitor.record_response(&request_id, &response, start.elapsed()).await;
response
}
```
## 流式响应
### SSE 实现
```rust
pub async fn stream_response(
provider_stream: impl Stream<Item = Result<Bytes>>,
) -> impl IntoResponse {
let stream = provider_stream.map(|chunk| {
match chunk {
Ok(data) => {
// 转换为 OpenAI SSE 格式
let converted = convert_to_openai_sse(&data);
Ok::<_, Error>(Event::default().data(converted))
}
Err(e) => Err(e),
}
});
Sse::new(stream)
.keep_alive(KeepAlive::default())
}
```
### 流式转换
```rust
// Provider 响应 → OpenAI SSE
pub fn convert_stream_chunk(chunk: &ProviderChunk) -> String {
let delta = ChatCompletionChunk {
id: chunk.id.clone(),
choices: vec![Choice {
delta: Delta {
content: chunk.content.clone(),
tool_calls: chunk.tool_calls.clone(),
},
finish_reason: chunk.finish_reason.clone(),
}],
};
format!("data: {}\n\n", serde_json::to_string(&delta).unwrap())
}
```
## 错误处理
### 错误响应格式
```rust
#[derive(Serialize)]
pub struct ErrorResponse {
pub error: ErrorDetail,
}
#[derive(Serialize)]
pub struct ErrorDetail {
pub message: String,
pub r#type: String,
pub code: Option<String>,
}
// 示例
{
"error": {
"message": "Rate limit exceeded",
"type": "rate_limit_error",
"code": "429"
}
}
```
### 错误处理器
```rust
pub async fn error_handler(err: Error) -> Response {
let (status, error_response) = match err {
Error::Validation(msg) => (
StatusCode::BAD_REQUEST,
ErrorResponse::new("invalid_request_error", msg),
),
Error::RateLimit => (
StatusCode::TOO_MANY_REQUESTS,
ErrorResponse::new("rate_limit_error", "Rate limit exceeded"),
),
Error::Provider(e) => (
StatusCode::BAD_GATEWAY,
ErrorResponse::new("provider_error", e.to_string()),
),
_ => (
StatusCode::INTERNAL_SERVER_ERROR,
ErrorResponse::new("internal_error", "Internal server error"),
),
};
(status, Json(error_response)).into_response()
}
```
## 相关文档
- [converter.md](converter.md) - 协议转换
- [flow-monitor.md](flow-monitor.md) - 流量监控
- [providers.md](providers.md) - Provider 系统
+107
View File
@@ -0,0 +1,107 @@
# 业务服务
## 概述
业务服务层封装核心业务逻辑,被 Tauri 命令调用。
## 目录结构
```
src-tauri/src/services/
├── mod.rs # 模块入口
├── provider_pool_service.rs # 凭证池服务
├── token_cache_service.rs # Token 缓存
├── mcp_service.rs # MCP 服务器管理
├── prompt_service.rs # Prompt 管理
├── skill_service.rs # 技能管理
├── usage_service.rs # 使用量统计
├── backup_service.rs # 备份服务
├── update_check_service.rs # 自动更新检查
└── general_chat/ # 通用对话服务
```
## 核心服务
### ProviderPoolService
```rust
pub struct ProviderPoolService {
pools: HashMap<ProviderType, CredentialPool>,
health_checker: HealthChecker,
}
impl ProviderPoolService {
/// 获取下一个可用凭证
pub async fn next_credential(&self, provider: ProviderType) -> Option<Credential>;
/// 添加凭证到池
pub async fn add_credential(&self, credential: Credential) -> Result<()>;
/// 移除凭证
pub async fn remove_credential(&self, id: &str) -> Result<()>;
/// 启动健康检查
pub fn start_health_check(&self);
}
```
### TokenCacheService
```rust
pub struct TokenCacheService {
cache: DashMap<String, CachedToken>,
db: Arc<Database>,
}
impl TokenCacheService {
/// 获取或刷新 Token
pub async fn get_or_refresh(&self, credential_id: &str) -> Result<String>;
/// 使 Token 失效
pub async fn invalidate(&self, credential_id: &str);
}
```
### McpService
```rust
pub struct McpService {
servers: HashMap<String, McpServer>,
}
impl McpService {
/// 启动 MCP 服务器
pub async fn start_server(&self, config: McpConfig) -> Result<()>;
/// 停止 MCP 服务器
pub async fn stop_server(&self, name: &str) -> Result<()>;
/// 列出工具
pub async fn list_tools(&self, server: &str) -> Result<Vec<Tool>>;
}
```
## 服务注入
```rust
// 在 main.rs 中初始化
let pool_service = Arc::new(ProviderPoolService::new());
let token_cache = Arc::new(TokenCacheService::new(db.clone()));
app.manage(pool_service);
app.manage(token_cache);
// 在命令中使用
#[tauri::command]
async fn add_credential(
pool: State<'_, Arc<ProviderPoolService>>,
// ...
) -> Result<(), String> {
pool.add_credential(credential).await
}
```
## 相关文档
- [commands.md](commands.md) - Tauri 命令
- [credential-pool.md](credential-pool.md) - 凭证池管理
+101
View File
@@ -0,0 +1,101 @@
# 内置终端
## 概述
内置终端模块提供 PTY 管理和会话管理功能。
## 目录结构
```
src-tauri/src/terminal/
├── mod.rs # 模块入口
├── pty.rs # PTY 管理
├── session.rs # 会话管理
└── commands.rs # 终端命令
src/components/terminal/
├── Terminal.tsx # 终端组件
└── TerminalTabs.tsx # 多标签管理
```
## PTY 管理
```rust
pub struct PtyManager {
sessions: HashMap<String, PtySession>,
}
pub struct PtySession {
id: String,
master: PtyMaster,
child: Child,
}
impl PtyManager {
/// 创建新会话
pub fn create_session(&mut self, shell: &str) -> Result<String>;
/// 写入数据
pub fn write(&self, session_id: &str, data: &[u8]) -> Result<()>;
/// 读取输出
pub fn read(&self, session_id: &str) -> Result<Vec<u8>>;
/// 调整大小
pub fn resize(&self, session_id: &str, cols: u16, rows: u16) -> Result<()>;
/// 关闭会话
pub fn close_session(&mut self, session_id: &str) -> Result<()>;
}
```
## 前端组件
```tsx
// src/components/terminal/Terminal.tsx
export function Terminal({ sessionId }: { sessionId: string }) {
const termRef = useRef<HTMLDivElement>(null);
const xtermRef = useRef<XTerm>();
useEffect(() => {
const xterm = new XTerm();
xterm.open(termRef.current!);
xtermRef.current = xterm;
// 监听输出
listen(`terminal-output-${sessionId}`, (event) => {
xterm.write(event.payload);
});
// 发送输入
xterm.onData((data) => {
invoke('terminal_write', { sessionId, data });
});
return () => xterm.dispose();
}, [sessionId]);
return <div ref={termRef} className="h-full" />;
}
```
## Tauri 命令
```rust
#[tauri::command]
async fn terminal_create(shell: Option<String>) -> Result<String, String>;
#[tauri::command]
async fn terminal_write(session_id: String, data: String) -> Result<(), String>;
#[tauri::command]
async fn terminal_resize(session_id: String, cols: u16, rows: u16) -> Result<(), String>;
#[tauri::command]
async fn terminal_close(session_id: String) -> Result<(), String>;
```
## 相关文档
- [commands.md](commands.md) - Tauri 命令
- [components.md](components.md) - 组件系统
+121
View File
@@ -0,0 +1,121 @@
# ProxyCast 测试体系
> 基于 Anthropic AI Agent 评估指南与 Orchids Bridge 项目实践
## 概述
ProxyCast 作为 AI API 代理和 Agent 集成平台,需要一套完整的测试体系来确保:
- API 代理的正确性和稳定性
- 凭证池管理的可靠性
- Aster Agent 集成的功能完整性
- 协议转换的准确性
## 测试分层
```
┌─────────────────────────────────────────────────────────────────┐
│ ProxyCast 测试金字塔 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ ┌─────────┐ │
│ │ E2E │ 端到端测试 │
│ │ 测试 │ (Tauri + 前端) │
│ ─┴─────────┴─ │
│ ┌─────────────┐ │
│ │ 集成测试 │ API 服务器、凭证池 │
│ ─┴─────────────┴─ │
│ ┌─────────────────┐ │
│ │ 单元测试 │ 转换器、Provider、工具 │
│ ─┴─────────────────┴─ │
│ │
└─────────────────────────────────────────────────────────────────┘
```
## 目录结构
```
docs/test/
├── README.md # 本文件 - 测试体系概览
├── unit-tests.md # 单元测试指南
├── integration-tests.md # 集成测试指南
├── e2e-tests.md # 端到端测试指南
├── agent-evaluation.md # Agent 评估指南(核心文档)
└── test-cases/ # 测试用例模板
├── converter-tests.md # 协议转换器测试用例
├── provider-tests.md # Provider 测试用例
└── agent-tests.md # Agent 测试用例
```
## 文档索引
| 文档 | 说明 | 适用场景 |
|------|------|----------|
| [unit-tests.md](unit-tests.md) | 单元测试指南 | 独立模块测试 |
| [integration-tests.md](integration-tests.md) | 集成测试指南 | 模块间协作测试 |
| [e2e-tests.md](e2e-tests.md) | E2E 测试指南 | 完整用户流程测试 |
| [agent-evaluation.md](agent-evaluation.md) | Agent 评估指南 | AI Agent 行为评估 |
| [test-cases/converter-tests.md](test-cases/converter-tests.md) | 转换器测试用例 | OpenAI ↔ Claude 转换 |
| [test-cases/provider-tests.md](test-cases/provider-tests.md) | Provider 测试用例 | OAuth 和 API 调用 |
| [test-cases/agent-tests.md](test-cases/agent-tests.md) | Agent 测试用例 | Aster Agent 集成 |
## 快速开始
### 运行 Rust 测试
```bash
cd src-tauri && cargo test
```
### 运行前端测试
```bash
npm test
```
### 运行代码检查
```bash
# Rust
cd src-tauri && cargo clippy
# 前端
npm run lint
```
## 核心测试模块
| 模块 | 测试重点 | 文档 |
|------|----------|------|
| 协议转换 | OpenAI ↔ Claude 转换正确性 | [converter-tests.md](test-cases/converter-tests.md) |
| Provider 系统 | OAuth 刷新、API 调用 | [provider-tests.md](test-cases/provider-tests.md) |
| 凭证池 | 轮询、健康检查、负载均衡 | [integration-tests.md](integration-tests.md) |
| Aster Agent | 流式响应、工具调用 | [agent-tests.md](test-cases/agent-tests.md) |
## 测试原则
基于 [Anthropic AI Agent 评估指南](https://www.anthropic.com/engineering/demystifying-evals-for-ai-agents) 和 Orchids Bridge 项目实践:
1. **评估结果,而非路径** - Agent 可能找到更好的方法,不要过度约束执行路径
2. **平衡问题集** - 测试"应该做"和"不应该做"两种情况
3. **隔离测试环境** - 每个测试独立状态,避免测试间污染
4. **从 Bug 到测试** - 每个修复的 Bug 都应该有对应测试用例
5. **处理非确定性** - 使用 pass@k 和 pass^k 指标评估 Agent 行为
6. **多层防护** - 结合自动评估、监控、人工审查
## 评分器类型
| 类型 | 适用场景 | 优点 | 缺点 |
|------|----------|------|------|
| **代码评分器** | 确定性验证 | 快速、可复现 | 对有效变体脆弱 |
| **模型评分器** | 语义评估 | 灵活、可扩展 | 非确定性、需校准 |
| **人工评分器** | 复杂判断 | 金标准质量 | 昂贵、慢 |
## 评估指标
```
pass@k = P(至少 1 次成功 | k 次尝试) = 1 - (1 - p)^k
pass^k = P(全部成功 | k 次尝试) = p^k
```
- **pass@k**:适用于"找到一个解决方案就行"的场景
- **pass^k**:适用于"每次都必须成功"的场景
+272
View File
@@ -0,0 +1,272 @@
# ProxyCast Agent 评估指南
> 基于 Anthropic AI Agent 评估指南的实践
## 概述
ProxyCast 集成了 Aster Agent,需要专门的评估体系来确保 Agent 行为的正确性和稳定性。本指南基于 Anthropic 官方评估指南和 Orchids Bridge 项目的实践经验。
## 核心概念
### 评估术语
| 术语 | 定义 | ProxyCast 示例 |
|------|------|----------------|
| **Task** | 单个测试任务 | "使用 Agent 读取文件并总结" |
| **Trial** | 对任务的一次尝试 | 同一任务运行 5 次 |
| **Grader** | 评分器 | 代码检查、LLM 判断 |
| **Transcript** | 完整记录 | Agent 的所有消息和工具调用 |
| **Outcome** | 最终结果 | 任务是否完成 |
### 评分器类型
```
┌─────────────────────────────────────────────────────────────────┐
│ 评分器类型 │
├─────────────────┬─────────────────┬─────────────────────────────┤
│ 代码评分器 │ 模型评分器 │ 人工评分器 │
├─────────────────┼─────────────────┼─────────────────────────────┤
│ • 工具调用验证 │ • 回答质量评估 │ • 复杂任务评审 │
│ • 输出格式检查 │ • 语义相似度 │ • 边界情况判断 │
│ • 状态断言 │ • 多轮对话评估 │ • 用户体验评估 │
└─────────────────┴─────────────────┴─────────────────────────────┘
```
## 评估场景
### 1. 工具调用评估
验证 Agent 正确调用工具:
```rust
#[cfg(test)]
mod agent_tool_tests {
use super::*;
#[tokio::test]
async fn test_file_read_tool_call() {
let agent = create_test_agent().await;
let response = agent.chat("请读取 /test/file.txt 的内容").await;
// 验证工具调用
assert!(response.tool_calls.iter().any(|tc| {
tc.name == "read_file" &&
tc.args.get("path") == Some(&"/test/file.txt".into())
}));
}
#[tokio::test]
async fn test_no_unnecessary_tool_calls() {
let agent = create_test_agent().await;
// 简单问题不应该调用工具
let response = agent.chat("1 + 1 等于多少?").await;
assert!(response.tool_calls.is_empty());
}
}
```
### 2. 流式响应评估
验证流式输出的正确性:
```rust
#[tokio::test]
async fn test_streaming_response_format() {
let agent = create_test_agent().await;
let mut stream = agent.chat_stream("你好").await;
let mut events = Vec::new();
while let Some(event) = stream.next().await {
events.push(event);
}
// 验证事件序列
assert!(events.iter().any(|e| matches!(e, StreamEvent::Start)));
assert!(events.iter().any(|e| matches!(e, StreamEvent::Delta(_))));
assert!(events.iter().any(|e| matches!(e, StreamEvent::Stop)));
}
#[tokio::test]
async fn test_streaming_content_accumulation() {
let agent = create_test_agent().await;
let mut stream = agent.chat_stream("写一首短诗").await;
let mut content = String::new();
while let Some(event) = stream.next().await {
if let StreamEvent::Delta(delta) = event {
content.push_str(&delta);
}
}
// 验证内容非空且有意义
assert!(!content.is_empty());
assert!(content.len() > 20);
}
```
### 3. 错误处理评估
验证 Agent 正确处理错误:
```rust
#[tokio::test]
async fn test_invalid_tool_graceful_handling() {
let agent = create_test_agent().await;
// 请求不存在的文件
let response = agent.chat("读取 /nonexistent/file.txt").await;
// Agent 应该优雅处理错误
assert!(response.content.contains("文件不存在") ||
response.content.contains("无法找到"));
}
#[tokio::test]
async fn test_timeout_handling() {
let agent = create_test_agent_with_timeout(Duration::from_secs(1)).await;
// 长时间任务应该超时
let result = agent.chat("执行一个需要很长时间的任务").await;
assert!(result.is_err() || result.unwrap().content.contains("超时"));
}
```
## 评估指标
### pass@k 与 pass^k
```
pass@k = P(至少 1 次成功 | k 次尝试)
pass^k = P(全部成功 | k 次尝试)
```
**应用场景**:
- **pass@k**:代码生成、创意任务(找到一个解决方案即可)
- **pass^k**:关键操作、用户交互(每次都必须成功)
### 评估脚本
```rust
async fn evaluate_task(task: &Task, trials: usize) -> EvalResult {
let mut successes = 0;
let mut transcripts = Vec::new();
for _ in 0..trials {
let agent = create_fresh_agent().await;
let transcript = agent.run_task(task).await;
let passed = task.grader.evaluate(&transcript);
if passed {
successes += 1;
}
transcripts.push(transcript);
}
EvalResult {
task_id: task.id.clone(),
trials,
successes,
pass_at_k: 1.0 - (1.0 - successes as f64 / trials as f64).powi(trials as i32),
pass_pow_k: (successes as f64 / trials as f64).powi(trials as i32),
transcripts,
}
}
```
## 测试套件组织
### 能力评估 vs 回归评估
| 类型 | 目标 | 初始通过率 | 用途 |
|------|------|-----------|------|
| **能力评估** | Agent 能做什么? | 低 | 推动改进 |
| **回归评估** | Agent 还能做以前能做的吗? | ~100% | 防止退化 |
### 测试套件结构
```
tests/agent/
├── capability/ # 能力评估
│ ├── file_operations.rs # 文件操作能力
│ ├── code_generation.rs # 代码生成能力
│ └── reasoning.rs # 推理能力
├── regression/ # 回归评估
│ ├── basic_chat.rs # 基础对话
│ ├── tool_calls.rs # 工具调用
│ └── streaming.rs # 流式响应
└── edge_cases/ # 边界情况
├── error_handling.rs
└── timeout.rs
```
## 评估原则
### 1. 评估结果,而非路径
```rust
// ❌ 错误:检查具体的工具调用顺序
fn test_bad() {
assert_eq!(transcript[0].tool, "list_files");
assert_eq!(transcript[1].tool, "read_file");
}
// ✅ 正确:检查最终结果
fn test_good() {
assert!(outcome.file_content.contains("expected content"));
}
```
### 2. 平衡问题集
```rust
// 测试"应该做"
#[test]
fn test_should_read_file_when_asked() { ... }
// 测试"不应该做"
#[test]
fn test_should_not_read_file_without_permission() { ... }
```
### 3. 从 Bug 到测试
每个修复的 Bug 都应该有对应的测试用例:
```rust
// Bug: Agent 在文件不存在时无限重试
// 修复后添加测试
#[test]
fn test_no_infinite_retry_on_missing_file() {
let agent = create_test_agent();
let response = agent.chat("读取 /nonexistent.txt").await;
// 验证重试次数有限
assert!(response.tool_calls.len() <= 3);
}
```
## 运行评估
```bash
# 运行所有 Agent 评估
cd src-tauri && cargo test agent::
# 运行能力评估
cargo test agent::capability::
# 运行回归评估
cargo test agent::regression::
# 运行多次试验
cargo test agent:: -- --test-threads=1 --nocapture
```
## 下一步
- [测试用例:Agent](test-cases/agent-tests.md)
- [单元测试指南](unit-tests.md)
+264
View File
@@ -0,0 +1,264 @@
# ProxyCast E2E 测试指南
> 端到端测试验证完整用户流程
## 概述
E2E 测试模拟真实用户操作,验证从前端到后端的完整流程。ProxyCast 使用 Tauri 框架,E2E 测试需要覆盖:
- 桌面应用启动和初始化
- 用户界面交互
- API 代理完整流程
- 凭证管理流程
## 测试框架
### Tauri E2E 测试
使用 `tauri-driver` 进行自动化测试:
```bash
# 安装依赖
cargo install tauri-driver
# 运行 E2E 测试
npm run test:e2e
```
### 测试配置
```javascript
// playwright.config.ts
import { defineConfig } from '@playwright/test';
export default defineConfig({
testDir: './tests/e2e',
timeout: 30000,
use: {
baseURL: 'tauri://localhost',
},
});
```
## 测试场景
### 1. 应用启动流程
```typescript
import { test, expect } from '@playwright/test';
test.describe('应用启动', () => {
test('应用正常启动并显示主界面', async ({ page }) => {
// 等待应用加载
await page.waitForSelector('[data-testid="main-layout"]');
// 验证核心组件存在
await expect(page.locator('[data-testid="sidebar"]')).toBeVisible();
await expect(page.locator('[data-testid="content-area"]')).toBeVisible();
});
test('首次启动显示欢迎引导', async ({ page }) => {
// 清除本地存储模拟首次启动
await page.evaluate(() => localStorage.clear());
await page.reload();
await expect(page.locator('[data-testid="welcome-modal"]')).toBeVisible();
});
});
```
### 2. 凭证管理流程
```typescript
test.describe('凭证管理', () => {
test('添加 Kiro 凭证', async ({ page }) => {
// 打开凭证管理
await page.click('[data-testid="credentials-tab"]');
await page.click('[data-testid="add-credential-btn"]');
// 选择 Provider
await page.click('[data-testid="provider-kiro"]');
// 上传凭证文件
const fileInput = page.locator('input[type="file"]');
await fileInput.setInputFiles('./tests/fixtures/test-credential.json');
// 验证凭证添加成功
await expect(page.locator('[data-testid="credential-item"]')).toBeVisible();
await expect(page.locator('text=test@example.com')).toBeVisible();
});
test('删除凭证', async ({ page }) => {
// 假设已有凭证
await page.click('[data-testid="credentials-tab"]');
// 删除凭证
await page.click('[data-testid="credential-menu"]');
await page.click('[data-testid="delete-credential"]');
await page.click('[data-testid="confirm-delete"]');
// 验证凭证已删除
await expect(page.locator('[data-testid="credential-item"]')).not.toBeVisible();
});
});
```
### 3. API 代理流程
```typescript
test.describe('API 代理', () => {
test('启动代理服务器', async ({ page }) => {
await page.click('[data-testid="server-tab"]');
await page.click('[data-testid="start-server-btn"]');
// 等待服务器启动
await expect(page.locator('text=服务器运行中')).toBeVisible();
await expect(page.locator('[data-testid="server-port"]')).toContainText('8080');
});
test('代理请求成功', async ({ page, request }) => {
// 启动服务器
await page.click('[data-testid="start-server-btn"]');
await page.waitForSelector('text=服务器运行中');
// 发送测试请求
const response = await request.post('http://localhost:8080/v1/chat/completions', {
headers: {
'Content-Type': 'application/json',
'Authorization': 'Bearer test-key',
},
data: {
model: 'gpt-4',
messages: [{ role: 'user', content: 'Hello' }],
},
});
expect(response.ok()).toBeTruthy();
});
});
```
### 4. Agent 对话流程
```typescript
test.describe('Agent 对话', () => {
test('发送消息并接收响应', async ({ page }) => {
await page.click('[data-testid="agent-tab"]');
// 输入消息
await page.fill('[data-testid="message-input"]', '你好,请介绍一下自己');
await page.click('[data-testid="send-btn"]');
// 等待响应
await expect(page.locator('[data-testid="assistant-message"]')).toBeVisible({
timeout: 30000,
});
});
test('流式响应正确显示', async ({ page }) => {
await page.click('[data-testid="agent-tab"]');
await page.fill('[data-testid="message-input"]', '写一首短诗');
await page.click('[data-testid="send-btn"]');
// 验证流式显示(内容逐渐增加)
const messageEl = page.locator('[data-testid="assistant-message"]');
let prevLength = 0;
for (let i = 0; i < 5; i++) {
await page.waitForTimeout(500);
const text = await messageEl.textContent();
expect(text?.length).toBeGreaterThan(prevLength);
prevLength = text?.length || 0;
}
});
});
```
## 测试数据管理
### Fixtures
```
tests/
├── fixtures/
│ ├── test-credential.json # 测试凭证
│ ├── mock-responses/ # Mock API 响应
│ │ ├── chat-completion.json
│ │ └── streaming-response.txt
│ └── test-config.json # 测试配置
└── e2e/
└── *.spec.ts
```
### Mock 服务
```typescript
// tests/mocks/api-server.ts
import { setupServer } from 'msw/node';
import { rest } from 'msw';
export const mockServer = setupServer(
rest.post('*/v1/chat/completions', (req, res, ctx) => {
return res(
ctx.json({
id: 'test-id',
choices: [{
message: { role: 'assistant', content: 'Mock response' },
}],
})
);
})
);
```
## 运行 E2E 测试
```bash
# 构建应用
npm run build
# 运行 E2E 测试
npm run test:e2e
# 运行特定测试
npm run test:e2e -- --grep "凭证管理"
# 生成测试报告
npm run test:e2e -- --reporter=html
```
## CI/CD 集成
```yaml
# .github/workflows/e2e.yml
name: E2E Tests
on: [push, pull_request]
jobs:
e2e:
runs-on: macos-latest
steps:
- uses: actions/checkout@v4
- name: Setup Node
uses: actions/setup-node@v4
with:
node-version: '20'
- name: Setup Rust
uses: dtolnay/rust-toolchain@stable
- name: Install dependencies
run: npm ci
- name: Build app
run: npm run build
- name: Run E2E tests
run: npm run test:e2e
```
## 下一步
- [Agent 评估指南](agent-evaluation.md)
- [测试用例:Agent](test-cases/agent-tests.md)
+229
View File
@@ -0,0 +1,229 @@
# ProxyCast 集成测试指南
> 测试模块间的协作和数据流
## 概述
集成测试验证多个模块协同工作的正确性,主要覆盖:
- API 服务器端点
- 凭证池管理
- Provider 与服务层交互
- 数据库操作
## 测试场景
### 1. API 服务器集成
```rust
#[cfg(test)]
mod api_integration_tests {
use super::*;
use axum::http::StatusCode;
use tower::ServiceExt;
#[tokio::test]
async fn test_chat_completion_endpoint() {
let app = create_test_app().await;
let request = Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("Content-Type", "application/json")
.header("Authorization", "Bearer test-key")
.body(Body::from(r#"{
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}]
}"#))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_streaming_response() {
let app = create_test_app().await;
let request = Request::builder()
.method("POST")
.uri("/v1/chat/completions")
.header("Content-Type", "application/json")
.body(Body::from(r#"{
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"stream": true
}"#))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response.headers().get("content-type").unwrap(),
"text/event-stream"
);
}
}
```
### 2. 凭证池集成
```rust
#[cfg(test)]
mod credential_pool_tests {
use super::*;
#[tokio::test]
async fn test_credential_rotation() {
let pool = CredentialPool::new();
// 添加多个凭证
pool.add_credential(create_test_credential("cred1")).await;
pool.add_credential(create_test_credential("cred2")).await;
pool.add_credential(create_test_credential("cred3")).await;
// 验证轮询
let first = pool.get_next().await.unwrap();
let second = pool.get_next().await.unwrap();
let third = pool.get_next().await.unwrap();
let fourth = pool.get_next().await.unwrap();
// 第四次应该回到第一个
assert_eq!(first.id, fourth.id);
}
#[tokio::test]
async fn test_unhealthy_credential_skipped() {
let pool = CredentialPool::new();
let healthy = create_test_credential("healthy");
let unhealthy = create_test_credential("unhealthy");
pool.add_credential(healthy.clone()).await;
pool.add_credential(unhealthy.clone()).await;
// 标记为不健康
pool.mark_unhealthy(&unhealthy.id).await;
// 应该只返回健康的凭证
for _ in 0..10 {
let cred = pool.get_next().await.unwrap();
assert_eq!(cred.id, healthy.id);
}
}
}
```
### 3. Provider 与数据库集成
```rust
#[cfg(test)]
mod provider_db_tests {
use super::*;
#[tokio::test]
async fn test_token_persistence() {
let db = create_test_db().await;
let provider = KiroProvider::new(db.clone());
// 刷新 Token
let token = provider.refresh_token("test-refresh-token").await.unwrap();
// 验证 Token 被保存到数据库
let saved = db.get_token("kiro", "test-id").await.unwrap();
assert_eq!(saved.access_token, token.access_token);
}
#[tokio::test]
async fn test_credential_state_sync() {
let db = create_test_db().await;
let service = ProviderPoolService::new(db.clone());
// 添加凭证
service.add_credential(create_test_credential()).await.unwrap();
// 验证数据库状态
let credentials = db.list_credentials("kiro").await.unwrap();
assert_eq!(credentials.len(), 1);
assert_eq!(credentials[0].status, "active");
}
}
```
## 测试环境设置
### 测试数据库
```rust
async fn create_test_db() -> Database {
let db = Database::new(":memory:").await.unwrap();
db.run_migrations().await.unwrap();
db
}
```
### Mock HTTP 服务
```rust
use wiremock::{MockServer, Mock, ResponseTemplate};
use wiremock::matchers::{method, path};
async fn setup_mock_oauth_server() -> MockServer {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/oauth/token"))
.respond_with(ResponseTemplate::new(200)
.set_body_json(json!({
"access_token": "test-token",
"expires_in": 3600
})))
.mount(&mock_server)
.await;
mock_server
}
```
## 测试数据管理
### Fixtures
```rust
fn create_test_credential(id: &str) -> Credential {
Credential {
id: id.to_string(),
provider: "kiro".to_string(),
email: "test@example.com".to_string(),
access_token: Some("test-access-token".to_string()),
refresh_token: Some("test-refresh-token".to_string()),
expires_at: Some(Utc::now() + Duration::hours(1)),
status: "active".to_string(),
}
}
fn create_expired_credential(id: &str) -> Credential {
let mut cred = create_test_credential(id);
cred.expires_at = Some(Utc::now() - Duration::hours(1));
cred
}
```
## 运行集成测试
```bash
# 运行所有集成测试
cd src-tauri && cargo test --test integration
# 运行特定测试
cargo test --test integration test_credential_rotation
# 并行运行(注意数据库隔离)
cargo test --test integration -- --test-threads=1
```
## 下一步
- [E2E 测试指南](e2e-tests.md)
- [Agent 评估指南](agent-evaluation.md)
+350
View File
@@ -0,0 +1,350 @@
# Agent 测试用例
> Aster Agent 集成的测试用例
## 概述
Agent 测试验证 Aster Agent 在 ProxyCast 中的集成,包括:
- 基础对话功能
- 流式响应
- 工具调用
- 错误处理
- 状态管理
## 测试用例
### 1. 基础对话
#### TC-AGENT-001: 简单对话
```rust
#[tokio::test]
async fn test_simple_chat() {
let state = create_test_agent_state().await;
let response = state.chat("你好").await.unwrap();
assert!(!response.content.is_empty());
assert_eq!(response.role, "assistant");
}
```
#### TC-AGENT-002: 多轮对话
```rust
#[tokio::test]
async fn test_multi_turn_chat() {
let state = create_test_agent_state().await;
// 第一轮
let r1 = state.chat("我叫小明").await.unwrap();
assert!(!r1.content.is_empty());
// 第二轮 - 应该记住上下文
let r2 = state.chat("我叫什么名字?").await.unwrap();
assert!(r2.content.contains("小明"));
}
```
#### TC-AGENT-003: 系统提示词
```rust
#[tokio::test]
async fn test_system_prompt() {
let state = create_test_agent_state().await;
state.set_system_prompt("你是一个诗人,只用诗歌回答问题").await;
let response = state.chat("今天天气怎么样?").await.unwrap();
// 响应应该有诗歌风格(包含换行或韵律)
assert!(response.content.contains('\n') || response.content.len() > 50);
}
```
### 2. 流式响应
#### TC-AGENT-010: 流式输出
```rust
#[tokio::test]
async fn test_streaming_output() {
let state = create_test_agent_state().await;
let mut stream = state.chat_stream("写一首短诗").await.unwrap();
let mut chunks = Vec::new();
while let Some(chunk) = stream.next().await {
chunks.push(chunk);
}
// 应该有多个 chunk
assert!(chunks.len() > 1);
// 合并后应该是完整内容
let full_content: String = chunks.iter()
.filter_map(|c| c.as_text())
.collect();
assert!(!full_content.is_empty());
}
```
#### TC-AGENT-011: 流式事件顺序
```rust
#[tokio::test]
async fn test_streaming_event_order() {
let state = create_test_agent_state().await;
let mut stream = state.chat_stream("你好").await.unwrap();
let mut events = Vec::new();
while let Some(event) = stream.next().await {
events.push(event);
}
// 验证事件顺序
let has_start = events.iter().any(|e| matches!(e, StreamEvent::Start));
let has_delta = events.iter().any(|e| matches!(e, StreamEvent::Delta(_)));
let has_stop = events.iter().any(|e| matches!(e, StreamEvent::Stop));
assert!(has_start);
assert!(has_delta);
assert!(has_stop);
}
```
#### TC-AGENT-012: 流式取消
```rust
#[tokio::test]
async fn test_streaming_cancellation() {
let state = create_test_agent_state().await;
let mut stream = state.chat_stream("写一篇长文章").await.unwrap();
// 只读取前几个 chunk
let mut count = 0;
while let Some(_) = stream.next().await {
count += 1;
if count >= 3 {
break;
}
}
// 取消流
drop(stream);
// 状态应该正确清理
assert!(state.is_idle().await);
}
```
### 3. 工具调用
#### TC-AGENT-020: 文件读取工具
```rust
#[tokio::test]
async fn test_file_read_tool() {
let state = create_test_agent_state().await;
// 创建测试文件
let test_file = create_temp_file("test content").await;
let response = state.chat(&format!("读取文件 {}", test_file.path())).await.unwrap();
// 应该调用了读取工具并返回内容
assert!(response.content.contains("test content") ||
response.tool_calls.iter().any(|tc| tc.name == "read_file"));
}
```
#### TC-AGENT-021: 文件写入工具
```rust
#[tokio::test]
async fn test_file_write_tool() {
let state = create_test_agent_state().await;
let temp_dir = create_temp_dir().await;
let file_path = temp_dir.join("output.txt");
let response = state.chat(&format!(
"在 {} 创建一个文件,内容是 'Hello World'",
file_path.display()
)).await.unwrap();
// 验证文件被创建
assert!(file_path.exists());
let content = std::fs::read_to_string(&file_path).unwrap();
assert!(content.contains("Hello World"));
}
```
#### TC-AGENT-022: 工具调用失败处理
```rust
#[tokio::test]
async fn test_tool_call_failure() {
let state = create_test_agent_state().await;
// 请求读取不存在的文件
let response = state.chat("读取 /nonexistent/file.txt").await.unwrap();
// Agent 应该优雅处理错误
assert!(response.content.contains("不存在") ||
response.content.contains("找不到") ||
response.content.contains("无法"));
}
```
### 4. 错误处理
#### TC-AGENT-030: 网络错误恢复
```rust
#[tokio::test]
async fn test_network_error_recovery() {
let state = create_test_agent_state_with_flaky_network().await;
// 第一次可能失败
let result1 = state.chat("你好").await;
// 重试应该成功
let result2 = state.chat("你好").await;
assert!(result1.is_ok() || result2.is_ok());
}
```
#### TC-AGENT-031: 超时处理
```rust
#[tokio::test]
async fn test_timeout_handling() {
let state = create_test_agent_state_with_timeout(Duration::from_secs(1)).await;
// 长任务应该超时
let result = state.chat("执行一个需要很长时间的复杂任务").await;
assert!(result.is_err() ||
result.unwrap().content.contains("超时"));
}
```
#### TC-AGENT-032: 无效输入处理
```rust
#[tokio::test]
async fn test_invalid_input() {
let state = create_test_agent_state().await;
// 空消息
let result = state.chat("").await;
assert!(result.is_err() || !result.unwrap().content.is_empty());
// 超长消息
let long_msg = "x".repeat(1_000_000);
let result = state.chat(&long_msg).await;
// 应该处理或拒绝,不应该崩溃
assert!(result.is_ok() || result.is_err());
}
```
### 5. 状态管理
#### TC-AGENT-040: 会话隔离
```rust
#[tokio::test]
async fn test_session_isolation() {
let state1 = create_test_agent_state().await;
let state2 = create_test_agent_state().await;
// 在 state1 中设置上下文
state1.chat("我叫小明").await.unwrap();
// state2 不应该知道这个信息
let response = state2.chat("我叫什么名字?").await.unwrap();
assert!(!response.content.contains("小明"));
}
```
#### TC-AGENT-041: 会话清理
```rust
#[tokio::test]
async fn test_session_cleanup() {
let state = create_test_agent_state().await;
// 建立上下文
state.chat("我叫小明").await.unwrap();
// 清理会话
state.clear_session().await;
// 上下文应该被清除
let response = state.chat("我叫什么名字?").await.unwrap();
assert!(!response.content.contains("小明"));
}
```
#### TC-AGENT-042: 并发请求
```rust
#[tokio::test]
async fn test_concurrent_requests() {
let state = Arc::new(create_test_agent_state().await);
let handles: Vec<_> = (0..5).map(|i| {
let state = state.clone();
tokio::spawn(async move {
state.chat(&format!("问题 {}", i)).await
})
}).collect();
let results: Vec<_> = futures::future::join_all(handles).await;
// 所有请求应该成功或有序失败
for result in results {
assert!(result.is_ok());
}
}
```
## 测试矩阵
| 测试 ID | 场景 | 类型 | 优先级 |
|---------|------|------|--------|
| TC-AGENT-001 | 简单对话 | 功能 | P0 |
| TC-AGENT-002 | 多轮对话 | 功能 | P0 |
| TC-AGENT-010 | 流式输出 | 功能 | P0 |
| TC-AGENT-020 | 文件读取 | 工具 | P1 |
| TC-AGENT-030 | 网络错误 | 错误处理 | P1 |
| TC-AGENT-040 | 会话隔离 | 状态 | P1 |
## 测试辅助函数
```rust
async fn create_test_agent_state() -> AsterAgentState {
let config = AsterConfig {
model: "test-model".into(),
api_key: "test-key".into(),
..Default::default()
};
AsterAgentState::new(config).await.unwrap()
}
async fn create_temp_file(content: &str) -> TempFile {
let file = TempFile::new().await.unwrap();
file.write_all(content.as_bytes()).await.unwrap();
file
}
```
## 运行测试
```bash
cd src-tauri && cargo test agent::
```
+260
View File
@@ -0,0 +1,260 @@
# 协议转换器测试用例
> OpenAI ↔ Claude 协议转换的测试用例
## 概述
协议转换器是 ProxyCast 的核心模块,负责在不同 API 格式之间转换。测试需要覆盖:
- 消息格式转换
- 流式响应转换
- 工具调用转换
- 边界情况处理
## 测试用例
### 1. 消息格式转换
#### TC-CONV-001: 基础消息转换
```rust
#[test]
fn test_openai_to_claude_basic_message() {
let openai_msg = OpenAIMessage {
role: "user".to_string(),
content: "Hello, world!".to_string(),
};
let claude_msg = convert_to_claude(&openai_msg);
assert_eq!(claude_msg.role, "user");
assert_eq!(claude_msg.content, "Hello, world!");
}
```
#### TC-CONV-002: System 消息处理
```rust
#[test]
fn test_system_message_extraction() {
let messages = vec![
OpenAIMessage { role: "system".into(), content: "You are helpful.".into() },
OpenAIMessage { role: "user".into(), content: "Hi".into() },
];
let (system, user_msgs) = extract_system_message(&messages);
assert_eq!(system, Some("You are helpful.".to_string()));
assert_eq!(user_msgs.len(), 1);
}
```
#### TC-CONV-003: 多轮对话转换
```rust
#[test]
fn test_multi_turn_conversation() {
let messages = vec![
OpenAIMessage { role: "user".into(), content: "Hello".into() },
OpenAIMessage { role: "assistant".into(), content: "Hi there!".into() },
OpenAIMessage { role: "user".into(), content: "How are you?".into() },
];
let claude_msgs = convert_messages(&messages);
assert_eq!(claude_msgs.len(), 3);
assert_eq!(claude_msgs[0].role, "user");
assert_eq!(claude_msgs[1].role, "assistant");
assert_eq!(claude_msgs[2].role, "user");
}
```
### 2. 流式响应转换
#### TC-CONV-010: SSE 事件格式
```rust
#[test]
fn test_sse_event_format() {
let delta = TextDelta { text: "Hello".to_string() };
let sse = format_sse_event(&delta);
assert!(sse.starts_with("data: "));
assert!(sse.ends_with("\n\n"));
assert!(sse.contains("\"delta\""));
}
```
#### TC-CONV-011: 流式开始事件
```rust
#[test]
fn test_stream_start_event() {
let event = create_stream_start_event("msg-123");
assert_eq!(event.event_type, "message_start");
assert!(event.data.contains("msg-123"));
}
```
#### TC-CONV-012: 流式结束事件
```rust
#[test]
fn test_stream_stop_event() {
let event = create_stream_stop_event("end_turn");
assert_eq!(event.event_type, "message_stop");
assert!(event.data.contains("end_turn"));
}
```
### 3. 工具调用转换
#### TC-CONV-020: 工具定义转换
```rust
#[test]
fn test_tool_definition_conversion() {
let openai_tool = OpenAITool {
r#type: "function".into(),
function: OpenAIFunction {
name: "get_weather".into(),
description: "Get weather info".into(),
parameters: json!({
"type": "object",
"properties": {
"location": { "type": "string" }
}
}),
},
};
let claude_tool = convert_tool(&openai_tool);
assert_eq!(claude_tool.name, "get_weather");
assert_eq!(claude_tool.description, "Get weather info");
}
```
#### TC-CONV-021: 工具调用响应转换
```rust
#[test]
fn test_tool_call_response_conversion() {
let claude_tool_use = ClaudeToolUse {
id: "tool-123".into(),
name: "get_weather".into(),
input: json!({"location": "Beijing"}),
};
let openai_tool_call = convert_tool_call(&claude_tool_use);
assert_eq!(openai_tool_call.id, "tool-123");
assert_eq!(openai_tool_call.function.name, "get_weather");
}
```
#### TC-CONV-022: 工具结果转换
```rust
#[test]
fn test_tool_result_conversion() {
let openai_result = OpenAIToolResult {
tool_call_id: "tool-123".into(),
content: "Sunny, 25°C".into(),
};
let claude_result = convert_tool_result(&openai_result);
assert_eq!(claude_result.tool_use_id, "tool-123");
assert_eq!(claude_result.content, "Sunny, 25°C");
}
```
### 4. 边界情况
#### TC-CONV-030: 空消息处理
```rust
#[test]
fn test_empty_message_content() {
let msg = OpenAIMessage {
role: "user".into(),
content: "".into(),
};
let result = convert_to_claude(&msg);
// 空内容应该被正确处理
assert!(result.content.is_empty());
}
```
#### TC-CONV-031: 特殊字符处理
```rust
#[test]
fn test_special_characters() {
let msg = OpenAIMessage {
role: "user".into(),
content: "Hello\n\t\"world\"\\test".into(),
};
let result = convert_to_claude(&msg);
// 特殊字符应该被保留
assert!(result.content.contains('\n'));
assert!(result.content.contains('\t'));
assert!(result.content.contains('"'));
}
```
#### TC-CONV-032: Unicode 处理
```rust
#[test]
fn test_unicode_content() {
let msg = OpenAIMessage {
role: "user".into(),
content: "你好世界 🌍 مرحبا".into(),
};
let result = convert_to_claude(&msg);
assert_eq!(result.content, "你好世界 🌍 مرحبا");
}
```
#### TC-CONV-033: 大消息处理
```rust
#[test]
fn test_large_message() {
let large_content = "x".repeat(100_000);
let msg = OpenAIMessage {
role: "user".into(),
content: large_content.clone(),
};
let result = convert_to_claude(&msg);
assert_eq!(result.content.len(), 100_000);
}
```
## 测试矩阵
| 测试 ID | 场景 | 输入 | 期望输出 | 优先级 |
|---------|------|------|----------|--------|
| TC-CONV-001 | 基础消息 | user 消息 | 正确转换 | P0 |
| TC-CONV-002 | System 消息 | system + user | 正确提取 | P0 |
| TC-CONV-010 | SSE 格式 | 文本增量 | 正确格式 | P0 |
| TC-CONV-020 | 工具定义 | OpenAI 工具 | Claude 工具 | P1 |
| TC-CONV-030 | 空消息 | 空内容 | 不崩溃 | P1 |
| TC-CONV-032 | Unicode | 多语言 | 正确保留 | P1 |
## 运行测试
```bash
cd src-tauri && cargo test converter::
```
+283
View File
@@ -0,0 +1,283 @@
# Provider 测试用例
> OAuth 认证和 API 调用的测试用例
## 概述
Provider 模块负责与各个 AI 服务提供商的交互,包括:
- OAuth 认证流程
- Token 刷新
- API 调用
- 错误处理
## 测试用例
### 1. Kiro Provider
#### TC-KIRO-001: 凭证加载
```rust
#[test]
fn test_kiro_credential_loading() {
let credential_json = r#"{
"access_token": "test-access",
"refresh_token": "test-refresh",
"expires_at": "2026-01-30T12:00:00Z"
}"#;
let cred = KiroCredential::from_json(credential_json).unwrap();
assert_eq!(cred.access_token, "test-access");
assert_eq!(cred.refresh_token, "test-refresh");
}
```
#### TC-KIRO-002: Token 刷新
```rust
#[tokio::test]
async fn test_kiro_token_refresh() {
let mock_server = setup_mock_oauth_server().await;
let provider = KiroProvider::new_with_endpoint(&mock_server.uri());
let new_token = provider.refresh_token("old-refresh-token").await.unwrap();
assert!(!new_token.access_token.is_empty());
assert!(new_token.expires_at > Utc::now());
}
```
#### TC-KIRO-003: 过期 Token 检测
```rust
#[test]
fn test_kiro_token_expiry_check() {
let expired_cred = KiroCredential {
access_token: "test".into(),
refresh_token: "test".into(),
expires_at: Utc::now() - Duration::hours(1),
};
assert!(expired_cred.is_expired());
let valid_cred = KiroCredential {
access_token: "test".into(),
refresh_token: "test".into(),
expires_at: Utc::now() + Duration::hours(1),
};
assert!(!valid_cred.is_expired());
}
```
### 2. Gemini Provider
#### TC-GEMINI-001: OAuth 流程
```rust
#[tokio::test]
async fn test_gemini_oauth_flow() {
let mock_server = setup_mock_google_oauth().await;
let provider = GeminiProvider::new_with_endpoint(&mock_server.uri());
let auth_url = provider.get_auth_url();
assert!(auth_url.contains("accounts.google.com"));
assert!(auth_url.contains("scope="));
}
```
#### TC-GEMINI-002: API 调用
```rust
#[tokio::test]
async fn test_gemini_api_call() {
let mock_server = setup_mock_gemini_api().await;
let provider = GeminiProvider::new_with_endpoint(&mock_server.uri());
let response = provider.chat(&[
Message { role: "user".into(), content: "Hello".into() }
]).await.unwrap();
assert!(!response.content.is_empty());
}
```
### 3. OpenAI Provider
#### TC-OPENAI-001: API Key 验证
```rust
#[test]
fn test_openai_api_key_validation() {
// 有效的 API Key
assert!(OpenAIProvider::validate_api_key("sk-1234567890abcdef"));
// 无效的 API Key
assert!(!OpenAIProvider::validate_api_key("invalid"));
assert!(!OpenAIProvider::validate_api_key(""));
}
```
#### TC-OPENAI-002: 流式响应处理
```rust
#[tokio::test]
async fn test_openai_streaming() {
let mock_server = setup_mock_openai_streaming().await;
let provider = OpenAIProvider::new_with_endpoint(&mock_server.uri());
let mut stream = provider.chat_stream(&[
Message { role: "user".into(), content: "Hello".into() }
]).await.unwrap();
let mut chunks = Vec::new();
while let Some(chunk) = stream.next().await {
chunks.push(chunk);
}
assert!(!chunks.is_empty());
}
```
### 4. 错误处理
#### TC-PROV-ERR-001: 网络错误
```rust
#[tokio::test]
async fn test_network_error_handling() {
let provider = KiroProvider::new_with_endpoint("http://invalid-host:9999");
let result = provider.refresh_token("test").await;
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), ProviderError::NetworkError(_)));
}
```
#### TC-PROV-ERR-002: 认证错误
```rust
#[tokio::test]
async fn test_auth_error_handling() {
let mock_server = setup_mock_oauth_error(401).await;
let provider = KiroProvider::new_with_endpoint(&mock_server.uri());
let result = provider.refresh_token("invalid-token").await;
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), ProviderError::AuthError(_)));
}
```
#### TC-PROV-ERR-003: 速率限制
```rust
#[tokio::test]
async fn test_rate_limit_handling() {
let mock_server = setup_mock_rate_limit().await;
let provider = OpenAIProvider::new_with_endpoint(&mock_server.uri());
let result = provider.chat(&[
Message { role: "user".into(), content: "Hello".into() }
]).await;
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), ProviderError::RateLimited(_)));
}
```
### 5. 凭证池集成
#### TC-POOL-001: 凭证轮询
```rust
#[tokio::test]
async fn test_credential_rotation() {
let pool = CredentialPool::new();
pool.add(create_credential("cred1")).await;
pool.add(create_credential("cred2")).await;
let first = pool.get_next().await.unwrap();
let second = pool.get_next().await.unwrap();
let third = pool.get_next().await.unwrap();
assert_ne!(first.id, second.id);
assert_eq!(first.id, third.id); // 回到第一个
}
```
#### TC-POOL-002: 健康检查
```rust
#[tokio::test]
async fn test_health_check() {
let pool = CredentialPool::new();
let cred = create_credential("test");
pool.add(cred.clone()).await;
// 标记为不健康
pool.mark_unhealthy(&cred.id).await;
// 不应该返回不健康的凭证
let result = pool.get_next().await;
assert!(result.is_none());
}
```
## 测试矩阵
| 测试 ID | Provider | 场景 | 优先级 |
|---------|----------|------|--------|
| TC-KIRO-001 | Kiro | 凭证加载 | P0 |
| TC-KIRO-002 | Kiro | Token 刷新 | P0 |
| TC-GEMINI-001 | Gemini | OAuth 流程 | P0 |
| TC-OPENAI-001 | OpenAI | API Key 验证 | P0 |
| TC-PROV-ERR-001 | 通用 | 网络错误 | P1 |
| TC-POOL-001 | 凭证池 | 轮询 | P0 |
## Mock 服务设置
```rust
async fn setup_mock_oauth_server() -> MockServer {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/oauth/token"))
.respond_with(ResponseTemplate::new(200)
.set_body_json(json!({
"access_token": "new-access-token",
"refresh_token": "new-refresh-token",
"expires_in": 3600
})))
.mount(&server)
.await;
server
}
async fn setup_mock_oauth_error(status: u16) -> MockServer {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/oauth/token"))
.respond_with(ResponseTemplate::new(status)
.set_body_json(json!({
"error": "invalid_grant",
"error_description": "Token expired"
})))
.mount(&server)
.await;
server
}
```
## 运行测试
```bash
cd src-tauri && cargo test provider::
```
+218
View File
@@ -0,0 +1,218 @@
# ProxyCast 单元测试指南
> 针对独立模块的确定性测试
## 概述
单元测试是测试金字塔的基础,覆盖最小的可测试单元。ProxyCast 的单元测试主要针对:
- 协议转换器
- Provider 模块
- 工具函数
- 数据结构
## Rust 单元测试
### 运行命令
```bash
# 运行所有测试
cd src-tauri && cargo test
# 运行特定模块测试
cargo test converter::
cargo test provider::
# 显示详细输出
cargo test -- --nocapture
```
### 测试文件位置
```
src-tauri/src/
├── converter/
│ ├── mod.rs
│ └── tests.rs # 转换器测试
├── providers/
│ ├── kiro/
│ │ └── tests.rs # Kiro Provider 测试
│ └── gemini/
│ └── tests.rs # Gemini Provider 测试
└── services/
└── tests.rs # 服务层测试
```
### 测试模板
```rust
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_conversion() {
let input = OpenAIMessage {
role: "user".to_string(),
content: "Hello".to_string(),
};
let result = convert_to_claude(&input);
assert_eq!(result.role, "user");
assert!(result.content.contains("Hello"));
}
#[test]
fn test_edge_case_empty_content() {
let input = OpenAIMessage {
role: "user".to_string(),
content: "".to_string(),
};
let result = convert_to_claude(&input);
// 空内容应该被正确处理
assert!(result.content.is_empty());
}
}
```
## 前端单元测试
### 运行命令
```bash
# 运行所有测试
npm test
# 运行特定文件
npm test -- src/lib/utils.test.ts
# 监听模式
npm test -- --watch
```
### 测试文件位置
```
src/
├── lib/
│ ├── utils.ts
│ └── utils.test.ts # 工具函数测试
├── hooks/
│ ├── useCredentials.ts
│ └── useCredentials.test.ts
└── components/
└── __tests__/ # 组件测试
```
### 测试模板
```typescript
import { describe, it, expect } from 'vitest';
import { formatCredentialName, validateApiKey } from './utils';
describe('formatCredentialName', () => {
it('should format kiro credential name', () => {
const result = formatCredentialName('kiro', 'user@example.com');
expect(result).toBe('Kiro (user@example.com)');
});
it('should handle empty email', () => {
const result = formatCredentialName('kiro', '');
expect(result).toBe('Kiro');
});
});
describe('validateApiKey', () => {
it('should accept valid OpenAI key', () => {
expect(validateApiKey('sk-1234567890abcdef')).toBe(true);
});
it('should reject invalid key', () => {
expect(validateApiKey('invalid')).toBe(false);
});
});
```
## 测试原则
### 1. 单一职责
每个测试只验证一个行为:
```rust
// ✅ 好:单一职责
#[test]
fn test_token_refresh_updates_expiry() {
// 只测试过期时间更新
}
#[test]
fn test_token_refresh_preserves_scope() {
// 只测试 scope 保留
}
// ❌ 差:多个职责
#[test]
fn test_token_refresh() {
// 测试过期时间、scope、错误处理...
}
```
### 2. 独立性
测试之间不应该有依赖:
```rust
// ✅ 好:每个测试独立
#[test]
fn test_a() {
let state = TestState::new();
// ...
}
#[test]
fn test_b() {
let state = TestState::new();
// ...
}
// ❌ 差:共享状态
static mut SHARED_STATE: Option<TestState> = None;
```
### 3. 可读性
测试名称应该描述行为:
```rust
// ✅ 好:描述性名称
#[test]
fn test_expired_token_triggers_refresh()
#[test]
fn test_invalid_credentials_returns_error()
// ❌ 差:模糊名称
#[test]
fn test_token()
#[test]
fn test_error()
```
## 覆盖率目标
| 模块 | 目标覆盖率 | 说明 |
|------|-----------|------|
| converter | 90%+ | 核心转换逻辑 |
| providers | 80%+ | OAuth 流程 |
| services | 70%+ | 业务逻辑 |
| utils | 95%+ | 工具函数 |
## 下一步
- [集成测试指南](integration-tests.md)
- [测试用例:转换器](test-cases/converter-tests.md)
- [测试用例:Provider](test-cases/provider-tests.md)
+1 -1
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.48.3",
"version": "0.48.4",
"type": "module",
"repository": {
"type": "git",
+4 -4
View File
@@ -180,7 +180,7 @@ checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50"
[[package]]
name = "aster"
version = "0.4.3"
version = "0.4.5"
dependencies = [
"ahash",
"anyhow",
@@ -4966,7 +4966,7 @@ version = "0.7.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ff32365de1b6743cb203b710788263c44a03de03802daf96092f2da4fe6ba4d7"
dependencies = [
"proc-macro-crate 3.4.0",
"proc-macro-crate 2.0.2",
"proc-macro2",
"quote",
"syn 2.0.114",
@@ -6147,7 +6147,7 @@ dependencies = [
[[package]]
name = "proxycast-core"
version = "0.48.3"
version = "0.48.4"
dependencies = [
"chrono",
"dirs 5.0.1",
@@ -6163,7 +6163,7 @@ dependencies = [
[[package]]
name = "proxycast-infra"
version = "0.48.3"
version = "0.48.4"
dependencies = [
"chrono",
"dashmap 5.5.3",
+1 -1
View File
@@ -3,7 +3,7 @@ members = ["crates/*"]
resolver = "2"
[workspace.package]
version = "0.48.3"
version = "0.48.4"
edition = "2021"
authors = ["you"]
repository = "https://github.com/aiclientproxy/proxycast"
@@ -74,7 +74,23 @@ impl std::str::FromStr for ProviderType {
"azure_openai" | "azure-openai" => Ok(ProviderType::AzureOpenai),
"aws_bedrock" | "aws-bedrock" => Ok(ProviderType::AwsBedrock),
"ollama" => Ok(ProviderType::Ollama),
_ => Err(format!("Invalid provider: {s}")),
// OpenAI 兼容的第三方 Provider 映射到 OpenAI
"deepseek" | "deep_seek" | "deep-seek" => Ok(ProviderType::OpenAI),
"qwen" | "tongyi" | "dashscope" => Ok(ProviderType::OpenAI),
"zhipu" | "glm" | "chatglm" => Ok(ProviderType::OpenAI),
"moonshot" | "kimi" => Ok(ProviderType::OpenAI),
"baichuan" => Ok(ProviderType::OpenAI),
"minimax" => Ok(ProviderType::OpenAI),
"yi" | "01ai" => Ok(ProviderType::OpenAI),
"stepfun" | "step" => Ok(ProviderType::OpenAI),
"groq" => Ok(ProviderType::OpenAI),
"together" | "togetherai" => Ok(ProviderType::OpenAI),
"fireworks" | "fireworksai" => Ok(ProviderType::OpenAI),
"perplexity" => Ok(ProviderType::OpenAI),
"siliconflow" => Ok(ProviderType::OpenAI),
"oneapi" | "one-api" | "newapi" | "new-api" => Ok(ProviderType::OpenAI),
"custom" | "custom_openai" => Ok(ProviderType::OpenAI),
_ => Err(format!("Unknown provider: {s}")),
}
}
}
@@ -1,7 +0,0 @@
# 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"
@@ -1,7 +0,0 @@
# 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 = "-"
@@ -1,8 +0,0 @@
# Seeds for failure cases proptest has generated in the past. It is
# automatically read and these particular cases re-run before any
# novel cases are generated.
#
# It is recommended to check this file in to source control so that
# everyone who runs the test benefits from these saved cases.
cc 7191f893879cb1f039c1686c2b2538314a7a1382e6b9cf1403e2a522d8c47c6c # shrinks to lines = [" C iG 0c0m", " FMf", "9 2 SD1GJ2X6zWI Mn", "4 PG2u6c M4e 3c U2gNc2DwEM1Aith", "CjGFRg F9L", "0y0aKdrF", " ", "jD9CWhlaRf", "SB urH38uh2qk ubPF 4 h1", "G Io fV633I3EViAtUk ekP2 C RJ", "oBNMH ls 45O0VU", "zgpqmoHl JA C5hq JMfE AVS", "qi4 ZnR f0w kv546Hi UNV8 ui k4bx 8mPr 13m", " 5 kZT n 7y1Eo zTf 8 ", "b N6 4Hu HFwkJ15", "3Tj 9WYz0 45 GYtb3GV LVx0 xpwh RYlu8Wo", " Wft 10nYZ Q40", "FD9dF1u bmlWx mdgJ VakP7v23Q", "595taG8u9x4IB 6Hu hv jy 2 0rfDcRf UL", "ql95v A", "pzbaL t7lV t mz n8 s890x 9OIHP p", "RpMgr n6y5 Zw7pt 34aV7m7cK H uhK Z6 o x ", "6LdW7FkF 8EdNr qcCav qp 0US69SCdJdTr", " Kz3 z vd8yI58Q kH 3nf4gJlnU", "JfR35hUSE SgZ C WH Xc5Ud 7f xlesr", "s1Mfb25X7 XBix V", "7BXzRq09oe brNC oUIzWMtwJkDZ8I7", "O", "6m2ai", "O 3Ft671 rM R6laKqS4ef rN20Qr", "Y82AvuHqrm6l Nfk6j a3 J4 0 2IfheC yz62 0", "R9OxTMZU67y 3 p HZ2EUMyZx468UIVB4gn", " m qg noz", "7TuXaK H yGvdYm4i rE9 ", "TE4CSD8H1KXNX4 24 o Iw448QnR c ", "72T", "642tY3FG Pl89X6 oq6iW9Z3UoaP N2M gs8tM7 8nQ6G6", "y 7WQguyJ E8D2 CZ", "O9 Xfp", "VIVXE SNN25D7 x9mJ 3TKdzIZA5aA0", "4T79 XYGsV0wAxU 3 1UG RZM", " PdN JA0R 4zPQ7 Q CBDKXjp4gnxZ 3", "7f j0ahlBI4tn4SS", "8v2lyDgaafHSGQb2lc4Q 6L0TKP7s yqC1 8P2", " XndE 4AA4eht9bIaoAO838 yginQ2CR3 Zh ", " p q8 GvYx2c507XrKCd2U97 73", "7 1Rnek2 y 02 1", "tFR", "5 bQ5", "01yM9Uo3KrMJ 08Jqd N 1Lm2q 05 7eT", "5DnH 6i70GsUE5Gcidwjd0 05Xg3yMiJnLl4g", "5hqDP0YC w17 AG31 U XV mN36d02YBkB8GEM14 AId1", "9GoS ", "T M oR B4Qb Uv0Mk7VsD2Ei 3 ", "2RI14wxg d3 3MtXJo IJ W y ", "EEs UFknE049Y5n", " wh ZirEFtZ67qquw ", "6t x3Yqz 7 b SjqXz1w k 8xH5ycwq WPdgR 73j", "Y KUL40NPzWud", "xjnx2Ow b XEGoVdePIBRwVmv srfK OI8 4P7Fh1", "Dh5a N5D1 rFXz0hKt4t7 4fr a NDDz", "iHuv4H ezvpK P 1EDX 2MT9EPb7hYR7v4", " 76FC41f1 K9B7 ea E5a 6K6 d1", "42 BkFMSr87 uixTdNnd115sSbr4c ", "ewc1 vIf", "DCBwDi8J d0UID OP 4l", "x DtFa6 O7O6stg3", "rDo Dlzn7 6 y g", "mjzRp7QiKy", "C Y0Z U9 28C QNUUWe4g cFqK8hjTHA9GbEq5Jn5Lx2i", "BFY7KGi42X67", "4 KX eO CJ8z f2x12TC HP6gD", "lj3PfjIZOvf MA93FheQo8HgmTV OPjL51q7IE 8X 4c", "hI TJ utF FVNHdkz w9JGBtTCae n4p2mvpz8H w8", " rc 5gQ19wGT k OYpCfDv S8TRN9JlF uU", "NBS t5D9 K fWPT6e3dBjIGsL81r", "8y 29182j3F2 Ty5 mrhvVZ6uo2iQMZ Plt e 1D03L 55 ", "K c aJ999JO XG 55Q dO8KVtyqu3e3jamJ K c5 ib3 P ", " P9OM8zVt0 Uk rxc2QBADHy k7h6b0 zQ8VXWvAJ ", " YyyIOm0", " zw ynSJ v08d IPY7F3", "P 5m eXGbx", "3gwH7 b18D 5SQU W", "0GND44b Eum5t88epJ8nrpc f8eiwvJ", "W h n5iP XgC", "nsU6u7qsa6dLB6 66l lELD49 F4z0 4SeYCDHgFS 6 4vek7 ", "SQ", "Wuk0IEfOW p4Sqc d jZP6C3210i3D6b u hM 7sKE"]
cc 22df1900e88351d89c99fcbe46488f3238532ecb4de326a180fa8e9adfc33350 # shrinks to lines = [" ", "H1SiOq hfWM1 4Rn3gXwADK", "g", "ubEzO2l", "L3dIUqqayi20Wy61PXnszI V 1Itm 4a 1 Y2m7", "l 3 YiL2lWHy nY984S0eH bkcIyJj Hs OEy83A2Y7NbnIaf", "38eOI nEF qx6AvWI5ZKKu5 j64Vg9eSTpIG", "a 44ICR8 p5fWunU co0 mZO4 xQq", "UD 8iEGV v1GVR4X ", "Smu7Hwde6DxdEu5 iGq n8Q 5 i3zpaBMF1b3LB ipR", " g7q 9vug3 ma ByY 6Y ", "9D rMa21kZmZ mWvY5w560HU N2J5fIjp QG040IJ4", "0r d342PHch hxT8 P20Eeng36mI4xrc4l", " Uv3Xhc i 1gm eWw", "s8 ", "o50vB F5N cf7 3 G dy C L9l Hs7F 3tz INB", " k9 45f Po Z BI njV jt ", " QRTc7 YsqG3mjj2 C8 ZD0Bg r6 gu vL", "6A3lB5lq 3 L ZDSWt go56cLAi X3Mu qm9t", "263", "2Qv06h1OF3 U EP1Jb EXhvyZJgw QtFHI4", "C3", "y sHbK4op QQcwH2k t 4y Aap 1E49 v", " 0LqE4 G3KumS EsvU Y Mykc7AG", " lja c", "aA cu8IGU k PBr3IdaNl75", "yN4n mgNln60 bn7 0", " 53jpSbNjC BScZ 7 a DCVfxJi s4pjL WCb p4s", "9DZ 2cmDI2vn", "mut Cr 0 J9 y5 G otM3b qd 6LId1o BwC61HH", " 0O 7SiAE2q x7QYax7H", "16 3XtzIK6z 16 t3 0CO9jdnPQshTR U", " OZkB9niv 5cs ", " 42E gZ ", "Y alH6x34J917 6 7tkpeo3xYq Y f DwQ aJO ", "y4 7rW yrX23GWC bF4G OCVyV4q"]
@@ -4,7 +4,4 @@
#
# It is recommended to check this file in to source control so that
# everyone who runs the test benefits from these saved cases.
cc c887db8633047b16f94b48762dc9bfe3c65bb683295d9047265207226e4e0f0b # shrinks to path = "~/."
cc 16635f3c212b2769f7fb51aec0142ac6e9e1a66d2f9e3118829cfba4ba86e4f8 # shrinks to (yaml_with_comments, original_comments) = ("# 0\nserver:\n# \n host: 127.0.0.1\n port: 1\n api_key: 0__AAA0_\nproviders:\n kiro:\n# a\n enabled: false\n region: us-east-1\n gemini:\n enabled: false\n credentials_path: 0-Aaa\n# 1\n qwen:\n enabled: false\n openai:\n enabled: false\n claude:\n enabled: false\ndefault_provider: kiro\nrouting:\n default_provider: kiro\n rules: []\n model_aliases: {}\n exclusions: {}\nretry:\n max_retries: 1\n base_delay_ms: 1\n max_delay_ms: 5000\n auto_switch_provider: false\nlogging:\n enabled: false\n level: debug\n retention_days: 1\n include_request_body: false\ninjection:\n enabled: false\n rules: []\nauth_dir: ~/.proxycast/auth\ncredential_pool: {}", ["# 0", "# ", "# a", "# 1"]), new_config = Config { server: ServerConfig { host: "127.0.0.1", port: 1, api_key: "a_a--a-a" }, providers: ProvidersConfig { kiro: ProviderConfig { enabled: false, credentials_path: None, region: None, project_id: None }, gemini: ProviderConfig { enabled: false, credentials_path: None, region: None, project_id: None }, qwen: ProviderConfig { enabled: false, credentials_path: None, region: None, project_id: Some("J3JRS6") }, openai: CustomProviderConfig { enabled: false, api_key: None, base_url: None }, claude: CustomProviderConfig { enabled: true, api_key: None, base_url: None } }, default_provider: "kiro", routing: RoutingConfig { default_provider: "kiro", rules: [RoutingRuleConfig { pattern: "cmbvvlhoabjthfkoczp-*", provider: "gemini", priority: 33 }], model_aliases: {}, exclusions: {} }, retry: RetrySettings { max_retries: 86, base_delay_ms: 4635, max_delay_ms: 9567, auto_switch_provider: false }, logging: LoggingConfig { enabled: false, level: "debug", retention_days: 16, include_request_body: false }, injection: InjectionSettings { enabled: false, rules: [] }, auth_dir: "~/.proxycast/auth", credential_pool: CredentialPoolConfig { kiro: [], gemini: [], qwen: [], openai: [], claude: [] } }
cc d22f0e24d166175ada35e5f5c4874b91f97fc63e0b40e6207ad9c70c07ecddda # shrinks to content = "{\"version\": }"
cc 09e08b21269b3921b7f80569c988d0d1591575a63704940a9e5aca7aa55268a3 # shrinks to subpath = "."
cc 0d8594955233ffc968ac57a1a1d9dcff20ab597104e3a2787bca02618e9fcf06 # shrinks to provider = "qwen"
@@ -1,8 +0,0 @@
# 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 93bc59c922237bf2c4ffc251c66c90f6bdcf7433a6a188b325fe65f38a8a6fd4 # shrinks to flow_count = 1
cc 5f7ef8a17a79bad4599803df567d02ed7c7a29bae560e3e90dfb8b2f920cadf2 # shrinks to flow_count = 5
@@ -1,7 +0,0 @@
# 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 4ef289c82d7068ccd05f549e93999e0d88c53e9d4ad4ebbcac1b33d00993d259 # shrinks to prefix = "ot"
@@ -1,7 +0,0 @@
# 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 a025afdf39438a61f1e2fdf53f6ee71ea4d4add523c76c9d060acfce58b901d6 # shrinks to initial_window = 30, new_window = 10, request_count = 14
@@ -1,8 +0,0 @@
# 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 be956d5aa14123ea1af5d3e310121b3dec52581c16da860aa2622f1689edd520 # shrinks to tool_call = ("call_00aa00aa", "aa_", "{\"value\":\"aAaA_a\"}")
cc f35fbd7673ecfb016ab0ee416718b5561162ae8ac1b04fbffe0fcebd8467a341 # shrinks to tool_call = ("call_a000a0a0", "__a", "{\"value\":\"aaAaaA\"}")
@@ -1,7 +0,0 @@
# 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 eac52c1d42823583ca0588a698ae96e4360b6cd26ff9f36c289a3ca06fc1a276 # shrinks to secret_key = "__O__Ew02R17oSvs94e--h"
@@ -1,7 +0,0 @@
# 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 10c52f014c7e9b4cd049a9802452d417b5774e6f42173c2c9329544d8ac4340c # shrinks to lead_time_mins = 21, time_offset_secs = 1260
@@ -1,7 +0,0 @@
# 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 99af5a6dad66f0b5a2650223417a2e83a7a3bfa8b6eab8ad57a88023367e739c # shrinks to url = "http://08:1024"
@@ -1,7 +0,0 @@
# 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 98a27aecaee1a9145362c2fa8a876f39010daa37133f7dd9f9713f6c2f8118f7 # shrinks to provider = "google", version = "v1"
@@ -1,8 +0,0 @@
# 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 bba7155a7ed5c22d19990ebdd9935e761dbdc41bc1e11076fbbf4218b82c8002 # shrinks to request_id = "0-AAaA-a", (original_data, sse_body) = ([" "], "data: \n\n")
cc 4f0d0f49ef961c396cdf1aa7ccbd2a01b525e96ba5a2048d18860dbe763abe37 # shrinks to request_id = "a00000aa", data = " ", index = 0
+63 -64
View File
@@ -4,85 +4,84 @@
## 架构说明
AI Agent 集成模块,提供原生 Rust Agent 功能,支持**连续对话**和**工具调用循环**。
AI Agent 集成模块,基于 aster-rust 框架实现。
### 设计决策
- **原生 Rust 实现**:直接在 Rust 中处理 Agent 功能,复用现有 provider 和流式处理能力
- **会话管理**:支持多会话,每个会话独立维护消息历史和系统提示词
- **连续对话**:每次请求携带 session_id,自动包含历史消息
- **Aster 框架**:使用 aster-rust 框架获得多 Provider、工具系统、会话管理等能力
- **凭证池桥接**:自动从 ProxyCast 凭证池选择凭证配置 Aster Provider
- **流式响应**:通过 Tauri 事件系统向前端推送流式内容
- **工具系统**:可扩展的工具定义和执行框架,支持 Bash、文件操作等
- **工具调用循环**:自动执行工具调用并继续对话,直到产生最终响应
## 文件索引
| 文件/目录 | 说明 |
| 文件 | 说明 |
|------|------|
| `mod.rs` | 模块入口,导出公共类型 |
| `types.rs` | Agent 相关类型定义(会话、消息、工具、配置) |
| `native_agent.rs` | 原生 Rust Agent 实现(NativeAgent、NativeAgentState) |
| `tool_loop.rs` | 工具调用循环引擎(ToolLoopEngine、ToolLoopConfig) |
| `tools/` | 工具系统子模块(类型定义、注册表、具体工具实现) |
| `types.rs` | Agent 相关类型定义 |
| `aster_state.rs` | Aster Agent 状态管理(Provider 配置、取消令牌) |
| `aster_agent.rs` | Aster Agent 包装器(会话管理) |
| `event_converter.rs` | Aster 事件到 Tauri 事件转换 |
| `credential_bridge.rs` | 凭证池桥接(连接 ProxyCast 凭证池与 Aster Provider) |
## 核心类型
## 使用方式
### 会话管理
- `AgentSession`: 会话状态,包含消息历史和系统提示词
- `AgentMessage`: 消息结构,支持文本、图片、工具调用
### 消息内容
- `MessageContent`: 消息内容(文本或多部分)
- `ContentPart`: 内容部分(文本/图片)
### 工具系统
- `ToolDefinition`: 工具定义(名称、描述、参数 Schema)
- `JsonSchema`: JSON Schema 参数定义
- `PropertySchema`: 属性 Schema(类型、描述、默认值)
- `ToolCall`: 工具调用请求
- `ToolResult`: 工具执行结果
- `ToolError`: 工具错误类型
- `Tool` trait: 工具接口(definition + execute)
- `ToolRegistry`: 工具注册表(注册、查找、验证、执行)
### 工具调用循环
- `ToolLoopEngine`: 工具循环引擎,执行工具调用并继续对话
- `ToolLoopConfig`: 循环配置(最大迭代次数等)
- `ToolLoopState`: 循环状态跟踪
- `ToolCallResult`: 工具调用结果
### Agent 实现
- `NativeAgent`: Agent 核心实现
- `NativeAgentState`: Tauri 状态管理器
## 使用示例
### 从凭证池配置(推荐)
```rust
// 创建会话
let session_id = agent_state.create_session(
Some("claude-sonnet-4-20250514".to_string()),
Some("你是一个有帮助的助手".to_string()),
)?;
// 初始化
state.init_agent().await?;
// 发送消息(自动包含历史)
let request = NativeChatRequest {
session_id: Some(session_id.clone()),
message: "你好".to_string(),
model: None,
images: None,
stream: false,
};
let response = agent_state.chat(request).await?;
// 从凭证池自动选择凭证并配置 Provider
let config = state
.configure_provider_from_pool(&db, "openai", "gpt-4", &session_id)
.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?;
// config.credential_uuid 包含使用的凭证 UUID
```
## 更新提醒
### 手动配置
任何文件变更后,请更新此文档和相关的上级文档。
```rust
// 初始化
state.init_agent().await?;
// 手动配置 Provider
let config = ProviderConfig {
provider_name: "openai".to_string(),
model_name: "gpt-4".to_string(),
api_key: Some("sk-...".to_string()),
base_url: None,
credential_uuid: None,
};
state.configure_provider(config, &session_id).await?;
```
### 发送消息
```rust
let user_message = Message::user().with_text("Hello");
let session_config = SessionConfigBuilder::new(&session_id).build();
let stream = agent.reply(user_message, session_config, Some(cancel_token)).await?;
```
## Tauri 命令
| 命令 | 说明 |
|------|------|
| `aster_agent_init` | 初始化 Agent |
| `aster_agent_configure_provider` | 手动配置 Provider |
| `aster_agent_configure_from_pool` | 从凭证池配置 Provider(推荐) |
| `aster_agent_chat_stream` | 流式对话 |
| `aster_agent_stop` | 停止会话 |
| `aster_session_create/list/get` | 会话管理 |
## 凭证池桥接
`credential_bridge.rs` 模块将 ProxyCast 凭证池与 Aster Provider 系统连接:
- 自动从凭证池选择可用凭证
- 支持 OAuth 和 API Key 两种凭证类型
- 自动刷新过期的 OAuth Token
- 记录凭证使用和健康状态
详见 [aster-integration.md](../../../docs/aiprompts/aster-integration.md)
+54 -60
View File
@@ -4,13 +4,11 @@
//! 处理消息发送、事件流转换和会话管理
use crate::agent::aster_state::{AsterAgentState, SessionConfigBuilder};
use crate::agent::event_converter::TauriAgentEvent;
use aster::agents::SessionConfig;
use aster::conversation::message::Message;
use aster::session::SessionManager;
use futures::StreamExt;
use std::path::PathBuf;
use tauri::{AppHandle, Emitter};
use tokio_util::sync::CancellationToken;
/// Aster Agent 包装器
///
@@ -36,78 +34,74 @@ impl AsterAgentWrapper {
session_id: String,
event_name: String,
) -> Result<(), String> {
// 确保 Agent 已初始化
// 1. 初始化检查
if !state.is_initialized().await {
state.init_agent().await?;
}
// 创建取消令牌
// 2. 创建取消令牌
let cancel_token = state.create_cancel_token(&session_id).await;
// 创建用户消息
// 3. 构建消息和配置
let user_message = Message::user().with_text(&message);
// 创建会话配置
let session_config = SessionConfigBuilder::new(&session_id).build();
// 使用 with_agent 方法获取 Agent 并处理
let app_clone = app.clone();
let event_name_clone = event_name.clone();
let cancel_token_clone = cancel_token.clone();
// 4. 获取 Agent 引用(关键步骤)
let agent_arc = state.get_agent_arc();
let guard = agent_arc.read().await;
let agent = guard.as_ref().ok_or("Agent not initialized")?;
let result = state
.with_agent(|agent| {
// 注意:这里我们需要异步处理,但 with_agent 是同步的
// 我们需要重新设计这个接口
})
// 5. 调用 Agent::reply
let stream_result = agent
.reply(user_message, session_config, Some(cancel_token.clone()))
.await;
// 由于 with_agent 的限制,我们需要使用不同的方法
// 直接在这里处理流
Self::process_reply_internal(
state,
&app_clone,
user_message,
session_config,
cancel_token_clone,
event_name_clone,
)
.await?;
// 6. 处理流式响应
match stream_result {
Ok(mut stream) => {
while let Some(event_result) = stream.next().await {
match event_result {
Ok(agent_event) => {
// 转换并发送事件到前端
let tauri_events =
crate::agent::event_converter::convert_agent_event(agent_event);
for tauri_event in tauri_events {
if let Err(e) = app.emit(&event_name, &tauri_event) {
tracing::error!("[AsterAgentWrapper] 发送事件失败: {}", e);
}
}
}
Err(e) => {
// 发送错误事件
let error_event =
crate::agent::event_converter::TauriAgentEvent::Error {
message: format!("Stream error: {}", e),
};
let _ = app.emit(&event_name, &error_event);
}
}
}
// 清理取消令牌
state.remove_cancel_token(&session_id).await;
Ok(())
}
/// 内部处理回复的方法
async fn process_reply_internal(
state: &AsterAgentState,
app: &AppHandle,
user_message: Message,
session_config: SessionConfig,
cancel_token: CancellationToken,
event_name: String,
) -> Result<(), String> {
// 这里我们需要一个更好的方式来访问 Agent
// 暂时使用一个简化的实现
// 发送开始事件
let start_event = TauriAgentEvent::TextDelta {
text: String::new(),
};
let _ = app.emit(&event_name, &start_event);
// TODO: 实现完整的 Agent 调用
// 由于 Agent.reply() 需要 &self,而我们的 with_agent 方法不支持异步
// 我们需要重新设计 AsterAgentState 的接口
// 发送完成事件
let done_event = TauriAgentEvent::FinalDone { usage: None };
if let Err(e) = app.emit(&event_name, &done_event) {
tracing::error!("Failed to emit final done event: {}", e);
// 发送完成事件
let done_event =
crate::agent::event_converter::TauriAgentEvent::FinalDone { usage: None };
let _ = app.emit(&event_name, &done_event);
}
Err(e) => {
// 发送错误事件并返回错误
let error_event = crate::agent::event_converter::TauriAgentEvent::Error {
message: format!("Agent error: {}", e),
};
let _ = app.emit(&event_name, &error_event);
return Err(format!("Agent error: {}", e));
}
}
// guard 在作用域结束时自动释放
// 7. 清理取消令牌
state.remove_cancel_token(&session_id).await;
Ok(())
}
+106
View File
@@ -2,6 +2,7 @@
//!
//! 管理 Aster Agent 实例和相关状态
//! 提供 Tauri 应用与 Aster 框架的桥接
//! 支持从 ProxyCast 凭证池自动选择凭证
use aster::agents::{Agent, SessionConfig};
use aster::model::ModelConfig;
@@ -9,6 +10,11 @@ use std::sync::Arc;
use tokio::sync::RwLock;
use tokio_util::sync::CancellationToken;
use crate::agent::credential_bridge::{
create_aster_provider, AsterProviderConfig, CredentialBridge, CredentialBridgeError,
};
use crate::database::DbConnection;
/// Provider 配置信息
#[derive(Debug, Clone)]
pub struct ProviderConfig {
@@ -20,6 +26,8 @@ pub struct ProviderConfig {
pub api_key: Option<String>,
/// Base URL (可选,用于自定义端点)
pub base_url: Option<String>,
/// 凭证 UUID(来自凭证池,用于记录使用和健康状态)
pub credential_uuid: Option<String>,
}
/// Aster Agent 全局状态
@@ -32,6 +40,8 @@ pub struct AsterAgentState {
cancel_tokens: Arc<RwLock<std::collections::HashMap<String, CancellationToken>>>,
/// 当前 Provider 配置
current_provider_config: Arc<RwLock<Option<ProviderConfig>>>,
/// 凭证桥接器
credential_bridge: CredentialBridge,
}
impl Default for AsterAgentState {
@@ -47,6 +57,7 @@ impl AsterAgentState {
agent: Arc::new(RwLock::new(None)),
cancel_tokens: Arc::new(RwLock::new(std::collections::HashMap::new())),
current_provider_config: Arc::new(RwLock::new(None)),
credential_bridge: CredentialBridge::new(),
}
}
@@ -108,6 +119,101 @@ impl AsterAgentState {
Ok(())
}
/// 从凭证池配置 Provider
///
/// 自动从 ProxyCast 凭证池选择可用凭证并配置 Aster Provider
///
/// # 参数
/// - `db`: 数据库连接
/// - `provider_type`: Provider 类型 (openai, anthropic, kiro 等)
/// - `model`: 模型名称
/// - `session_id`: 会话 ID
pub async fn configure_provider_from_pool(
&self,
db: &DbConnection,
provider_type: &str,
model: &str,
session_id: &str,
) -> Result<AsterProviderConfig, String> {
// 确保 Agent 已初始化
self.init_agent().await?;
// 从凭证池选择凭证并获取配置
let aster_config = self
.credential_bridge
.select_and_configure(db, provider_type, model)
.await
.map_err(|e| format!("从凭证池选择凭证失败: {}", e))?;
// 创建 Provider
let provider = create_aster_provider(&aster_config)
.await
.map_err(|e| format!("创建 Provider 失败: {}", e))?;
// 更新 Agent 的 Provider
let agent_guard = self.agent.read().await;
if let Some(agent) = agent_guard.as_ref() {
agent
.update_provider(provider, session_id)
.await
.map_err(|e| format!("更新 Provider 失败: {}", e))?;
}
// 保存当前配置
let config = ProviderConfig {
provider_name: aster_config.provider_name.clone(),
model_name: aster_config.model_name.clone(),
api_key: aster_config.api_key.clone(),
base_url: aster_config.base_url.clone(),
credential_uuid: Some(aster_config.credential_uuid.clone()),
};
let mut config_guard = self.current_provider_config.write().await;
*config_guard = Some(config);
// 记录凭证使用
if let Err(e) = self
.credential_bridge
.record_usage(db, &aster_config.credential_uuid)
{
tracing::warn!("[AsterAgent] 记录凭证使用失败: {}", e);
}
tracing::info!(
"[AsterAgent] 从凭证池配置 Provider 成功: {} / {} (凭证: {})",
aster_config.provider_name,
aster_config.model_name,
aster_config.credential_uuid
);
Ok(aster_config)
}
/// 标记当前凭证为健康
pub fn mark_current_healthy(&self, db: &DbConnection, model: Option<&str>) {
if let Ok(config_guard) = self.current_provider_config.try_read() {
if let Some(config) = config_guard.as_ref() {
if let Some(uuid) = &config.credential_uuid {
if let Err(e) = self.credential_bridge.mark_healthy(db, uuid, model) {
tracing::warn!("[AsterAgent] 标记凭证健康失败: {}", e);
}
}
}
}
}
/// 标记当前凭证为不健康
pub fn mark_current_unhealthy(&self, db: &DbConnection, error: Option<&str>) {
if let Ok(config_guard) = self.current_provider_config.try_read() {
if let Some(config) = config_guard.as_ref() {
if let Some(uuid) = &config.credential_uuid {
if let Err(e) = self.credential_bridge.mark_unhealthy(db, uuid, error) {
tracing::warn!("[AsterAgent] 标记凭证不健康失败: {}", e);
}
}
}
}
}
/// 设置 Provider 相关的环境变量
fn set_provider_env_vars(&self, config: &ProviderConfig) {
// 根据 provider 类型设置对应的环境变量
+428
View File
@@ -0,0 +1,428 @@
//! 凭证池桥接模块
//!
//! 将 ProxyCast 凭证池与 Aster Provider 系统连接
//! 支持从凭证池自动选择凭证并配置 Aster Provider
//!
//! ## 功能
//! - 从凭证池选择可用凭证
//! - 将凭证转换为 Aster Provider 配置
//! - 支持 OAuth 和 API Key 两种凭证类型
//! - 自动刷新过期的 OAuth Token
use crate::database::DbConnection;
use crate::models::provider_pool_model::{CredentialData, PoolProviderType, ProviderCredential};
use crate::services::api_key_provider_service::ApiKeyProviderService;
use crate::services::provider_pool_service::ProviderPoolService;
use aster::model::ModelConfig;
use aster::providers::base::Provider;
use std::sync::Arc;
/// 凭证桥接错误
#[derive(Debug, Clone)]
pub enum CredentialBridgeError {
/// 没有可用凭证
NoCredentials(String),
/// 凭证类型不支持
UnsupportedCredentialType(String),
/// Provider 创建失败
ProviderCreationFailed(String),
/// Token 刷新失败
TokenRefreshFailed(String),
/// 数据库错误
DatabaseError(String),
}
impl std::fmt::Display for CredentialBridgeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NoCredentials(msg) => write!(f, "没有可用凭证: {}", msg),
Self::UnsupportedCredentialType(msg) => write!(f, "不支持的凭证类型: {}", msg),
Self::ProviderCreationFailed(msg) => write!(f, "Provider 创建失败: {}", msg),
Self::TokenRefreshFailed(msg) => write!(f, "Token 刷新失败: {}", msg),
Self::DatabaseError(msg) => write!(f, "数据库错误: {}", msg),
}
}
}
impl std::error::Error for CredentialBridgeError {}
/// Aster Provider 配置
#[derive(Debug, Clone)]
pub struct AsterProviderConfig {
/// Provider 名称 (openai, anthropic, google 等)
pub provider_name: String,
/// 模型名称
pub model_name: String,
/// API Key
pub api_key: Option<String>,
/// Base URL
pub base_url: Option<String>,
/// 凭证 UUID(用于记录使用和健康状态)
pub credential_uuid: String,
}
/// 凭证池桥接器
///
/// 负责从 ProxyCast 凭证池选择凭证并转换为 Aster Provider 配置
pub struct CredentialBridge {
pool_service: ProviderPoolService,
api_key_service: ApiKeyProviderService,
}
impl Default for CredentialBridge {
fn default() -> Self {
Self::new()
}
}
impl CredentialBridge {
pub fn new() -> Self {
Self {
pool_service: ProviderPoolService::new(),
api_key_service: ApiKeyProviderService::new(),
}
}
/// 从凭证池选择凭证并创建 Aster Provider 配置
///
/// # 参数
/// - `db`: 数据库连接
/// - `provider_type`: Provider 类型 (openai, anthropic, kiro, deepseek 等)
/// - `model`: 模型名称
///
/// # 返回
/// 成功时返回 AsterProviderConfig,失败时返回错误
pub async fn select_and_configure(
&self,
db: &DbConnection,
provider_type: &str,
model: &str,
) -> Result<AsterProviderConfig, CredentialBridgeError> {
// 1. 从凭证池选择凭证
// 将 provider_type 同时作为 provider_id_hint 传递,支持 60+ API Key Provider
// 例如 "deepseek", "moonshot", "qwen" 等
let credential = self
.pool_service
.select_credential_with_fallback(
db,
&self.api_key_service,
provider_type,
Some(model),
Some(provider_type), // 传递 provider_id_hint 支持智能降级
None,
)
.await
.map_err(|e| CredentialBridgeError::DatabaseError(e))?
.ok_or_else(|| {
CredentialBridgeError::NoCredentials(format!(
"没有找到 {} 类型的可用凭证",
provider_type
))
})?;
// 2. 转换为 Aster Provider 配置
self.credential_to_config(&credential, model, db).await
}
/// 将 ProxyCast 凭证转换为 Aster Provider 配置
async fn credential_to_config(
&self,
credential: &ProviderCredential,
model: &str,
db: &DbConnection,
) -> Result<AsterProviderConfig, CredentialBridgeError> {
let (provider_name, api_key, base_url) = match &credential.credential {
// OpenAI API Key
CredentialData::OpenAIKey { api_key, base_url } => (
"openai".to_string(),
Some(api_key.clone()),
base_url.clone(),
),
// Claude/Anthropic API Key
CredentialData::ClaudeKey { api_key, base_url }
| CredentialData::AnthropicKey { api_key, base_url } => (
"anthropic".to_string(),
Some(api_key.clone()),
base_url.clone(),
),
// Kiro OAuth - 需要获取 access_token
CredentialData::KiroOAuth { creds_file_path } => {
let token = self
.get_kiro_token(creds_file_path, db, &credential.uuid)
.await?;
// Kiro 使用 CodeWhisperer API,映射到 bedrock provider
("bedrock".to_string(), Some(token), None)
}
// Gemini OAuth
CredentialData::GeminiOAuth {
creds_file_path, ..
} => {
let token = self.get_oauth_token(creds_file_path).await?;
("google".to_string(), Some(token), None)
}
// Gemini API Key
CredentialData::GeminiApiKey {
api_key, base_url, ..
} => (
"google".to_string(),
Some(api_key.clone()),
base_url.clone(),
),
// Vertex AI
CredentialData::VertexKey {
api_key, base_url, ..
} => (
"gcpvertexai".to_string(),
Some(api_key.clone()),
base_url.clone(),
),
// Codex OAuth
CredentialData::CodexOAuth {
creds_file_path,
api_base_url,
} => {
let token = self.get_codex_token(creds_file_path).await?;
("codex".to_string(), Some(token), api_base_url.clone())
}
// Claude OAuth
CredentialData::ClaudeOAuth { creds_file_path } => {
let token = self.get_oauth_token(creds_file_path).await?;
("anthropic".to_string(), Some(token), None)
}
// Antigravity OAuth
CredentialData::AntigravityOAuth {
creds_file_path, ..
} => {
let token = self.get_oauth_token(creds_file_path).await?;
("google".to_string(), Some(token), None)
}
};
Ok(AsterProviderConfig {
provider_name,
model_name: model.to_string(),
api_key,
base_url,
credential_uuid: credential.uuid.clone(),
})
}
/// 获取 Kiro OAuth Token
async fn get_kiro_token(
&self,
creds_path: &str,
db: &DbConnection,
uuid: &str,
) -> Result<String, CredentialBridgeError> {
use crate::providers::kiro::KiroProvider;
let mut provider = KiroProvider::new();
provider
.load_credentials_from_path(creds_path)
.await
.map_err(|e| {
CredentialBridgeError::TokenRefreshFailed(format!("加载 Kiro 凭证失败: {}", e))
})?;
// 检查 token 是否过期,如果过期则刷新
if provider.is_token_expired() {
tracing::info!("[CredentialBridge] Kiro token 已过期,尝试刷新");
self.pool_service
.refresh_kiro_token(creds_path)
.await
.map_err(|e| CredentialBridgeError::TokenRefreshFailed(e))?;
// 重新加载凭证
provider
.load_credentials_from_path(creds_path)
.await
.map_err(|e| {
CredentialBridgeError::TokenRefreshFailed(format!("重新加载凭证失败: {}", e))
})?;
}
provider.credentials.access_token.ok_or_else(|| {
CredentialBridgeError::TokenRefreshFailed("缺少 access_token".to_string())
})
}
/// 获取通用 OAuth Token
async fn get_oauth_token(&self, creds_path: &str) -> Result<String, CredentialBridgeError> {
let content = std::fs::read_to_string(creds_path).map_err(|e| {
CredentialBridgeError::TokenRefreshFailed(format!("读取凭证文件失败: {}", e))
})?;
let creds: serde_json::Value = serde_json::from_str(&content).map_err(|e| {
CredentialBridgeError::TokenRefreshFailed(format!("解析凭证失败: {}", e))
})?;
creds["access_token"]
.as_str()
.map(String::from)
.ok_or_else(|| {
CredentialBridgeError::TokenRefreshFailed("凭证中缺少 access_token".to_string())
})
}
/// 获取 Codex OAuth Token
async fn get_codex_token(&self, creds_path: &str) -> Result<String, CredentialBridgeError> {
use crate::providers::codex::CodexProvider;
let mut provider = CodexProvider::new();
provider
.load_credentials_from_path(creds_path)
.await
.map_err(|e| {
CredentialBridgeError::TokenRefreshFailed(format!("加载 Codex 凭证失败: {}", e))
})?;
provider.ensure_valid_token().await.map_err(|e| {
CredentialBridgeError::TokenRefreshFailed(format!("获取 Codex token 失败: {}", e))
})
}
/// 记录凭证使用
pub fn record_usage(&self, db: &DbConnection, uuid: &str) -> Result<(), CredentialBridgeError> {
self.pool_service
.record_usage(db, uuid)
.map_err(|e| CredentialBridgeError::DatabaseError(e))
}
/// 标记凭证为健康
pub fn mark_healthy(
&self,
db: &DbConnection,
uuid: &str,
model: Option<&str>,
) -> Result<(), CredentialBridgeError> {
self.pool_service
.mark_healthy(db, uuid, model)
.map_err(|e| CredentialBridgeError::DatabaseError(e))
}
/// 标记凭证为不健康
pub fn mark_unhealthy(
&self,
db: &DbConnection,
uuid: &str,
error: Option<&str>,
) -> Result<(), CredentialBridgeError> {
self.pool_service
.mark_unhealthy(db, uuid, error)
.map_err(|e| CredentialBridgeError::DatabaseError(e))
}
}
/// 从 AsterProviderConfig 创建 Aster Provider
///
/// 设置环境变量并调用 aster::providers::create
pub async fn create_aster_provider(
config: &AsterProviderConfig,
) -> Result<Arc<dyn Provider>, CredentialBridgeError> {
// 设置环境变量
set_provider_env_vars(config);
// 创建 ModelConfig
let model_config = ModelConfig::new(&config.model_name).map_err(|e| {
CredentialBridgeError::ProviderCreationFailed(format!("创建 ModelConfig 失败: {}", e))
})?;
// 创建 Provider
aster::providers::create(&config.provider_name, model_config)
.await
.map_err(|e| {
CredentialBridgeError::ProviderCreationFailed(format!("创建 Provider 失败: {}", e))
})
}
/// 设置 Provider 环境变量
fn set_provider_env_vars(config: &AsterProviderConfig) {
let env_key = match config.provider_name.as_str() {
"openai" => "OPENAI_API_KEY",
"anthropic" => "ANTHROPIC_API_KEY",
"google" => "GOOGLE_API_KEY",
"bedrock" => "AWS_ACCESS_KEY_ID", // Bedrock 使用 AWS 凭证
"gcpvertexai" => "GOOGLE_API_KEY",
"codex" => "OPENAI_API_KEY", // Codex 兼容 OpenAI
_ => "OPENAI_API_KEY", // 默认使用 OpenAI 格式
};
if let Some(api_key) = &config.api_key {
std::env::set_var(env_key, api_key);
}
// 设置 base_url
// Aster 的 OpenAI Provider 使用 OPENAI_HOST 环境变量
if let Some(base_url) = &config.base_url {
match config.provider_name.as_str() {
"openai" => {
// OpenAI 兼容的 Provider 使用 OPENAI_HOST
std::env::set_var("OPENAI_HOST", base_url);
tracing::info!("[CredentialBridge] 设置 OPENAI_HOST={}", base_url);
}
"anthropic" => {
std::env::set_var("ANTHROPIC_BASE_URL", base_url);
}
_ => {
// 其他 Provider 使用通用格式
let base_url_key = format!(
"{}_BASE_URL",
config.provider_name.to_uppercase().replace('-', "_")
);
std::env::set_var(&base_url_key, base_url);
}
}
}
}
/// Provider 类型映射
///
/// 将 ProxyCast PoolProviderType 映射到 Aster Provider 名称
pub fn map_pool_type_to_aster(pool_type: &PoolProviderType) -> &'static str {
match pool_type {
PoolProviderType::Kiro => "bedrock",
PoolProviderType::Gemini => "google",
PoolProviderType::Antigravity => "google",
PoolProviderType::OpenAI => "openai",
PoolProviderType::Claude => "anthropic",
PoolProviderType::Anthropic => "anthropic",
PoolProviderType::AnthropicCompatible => "anthropic",
PoolProviderType::Vertex => "gcpvertexai",
PoolProviderType::GeminiApiKey => "google",
PoolProviderType::Codex => "codex",
PoolProviderType::ClaudeOAuth => "anthropic",
PoolProviderType::AzureOpenai => "azure",
PoolProviderType::AwsBedrock => "bedrock",
PoolProviderType::Ollama => "ollama",
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_map_pool_type_to_aster() {
assert_eq!(map_pool_type_to_aster(&PoolProviderType::OpenAI), "openai");
assert_eq!(
map_pool_type_to_aster(&PoolProviderType::Claude),
"anthropic"
);
assert_eq!(map_pool_type_to_aster(&PoolProviderType::Gemini), "google");
assert_eq!(map_pool_type_to_aster(&PoolProviderType::Kiro), "bedrock");
}
#[test]
fn test_credential_bridge_error_display() {
let err = CredentialBridgeError::NoCredentials("test".to_string());
assert!(err.to_string().contains("没有可用凭证"));
}
}
+9 -19
View File
@@ -1,33 +1,23 @@
//! AI Agent 集成模块
//!
//! 使用策略模式支持多种 API 协议(OpenAI、Anthropic、Kiro、Gemini)
//! 包含工具系统、流式处理和工具调用循环
//! 基于 aster-rust 框架实现 Agent 功能
//!
//! ## 架构设计
//! - protocols/ - 协议策略实现(策略模式)
//! - parsers/ - SSE 流解析器
//! - native_agent - 核心 Agent 逻辑
//! - tool_loop - 工具调用循环
//! - tools/ - 工具实现
//! - aster_state - Aster Agent 状态管理(新)
//! - aster_agent - Aster Agent 包装器(新)
//! - event_converter - Aster 事件转换器(新)
//! - aster_state - Aster Agent 状态管理
//! - aster_agent - Aster Agent 包装器
//! - event_converter - Aster 事件转换器
//! - credential_bridge - 凭证池桥接(连接 ProxyCast 凭证池与 Aster Provider)
pub mod aster_agent;
pub mod aster_state;
pub mod credential_bridge;
pub mod event_converter;
pub mod native_agent;
pub mod parsers;
pub mod protocols;
pub mod tool_loop;
pub mod tools;
pub mod types;
pub use aster_agent::{AsterAgentWrapper, SessionDetail, SessionInfo};
pub use aster_state::AsterAgentState;
pub use credential_bridge::{
create_aster_provider, AsterProviderConfig, CredentialBridge, CredentialBridgeError,
};
pub use event_converter::{convert_agent_event, TauriAgentEvent};
pub use native_agent::{NativeAgent, NativeAgentState};
pub use parsers::{AnthropicSSEParser, OpenAISSEParser};
pub use protocols::{create_protocol, AnthropicProtocol, OpenAIProtocol, Protocol};
pub use tool_loop::{ToolCallResult, ToolLoopConfig, ToolLoopEngine, ToolLoopError, ToolLoopState};
pub use types::*;
File diff suppressed because it is too large Load Diff
@@ -1,209 +0,0 @@
//! Anthropic SSE 流解析器
//!
//! 解析 Anthropic Messages API 的 Server-Sent Events 流
use crate::agent::types::{FunctionCall, TokenUsage, ToolCall};
use crate::models::anthropic::{AnthropicContentBlock, AnthropicDelta, AnthropicStreamEvent};
use tracing::{debug, warn};
/// Anthropic 工具调用构建器
#[derive(Debug, Clone, Default)]
struct AnthropicToolCallBuilder {
id: String,
name: String,
input_json: String,
}
/// Anthropic SSE 流解析器
///
/// 解析 Anthropic Messages API 的 SSE 流
#[derive(Debug, Default)]
pub struct AnthropicSSEParser {
/// 累积的完整内容
full_content: String,
/// 累积的工具调用
tool_calls: Vec<ToolCall>,
/// 当前正在构建的工具调用
current_tool: Option<AnthropicToolCallBuilder>,
/// Usage 信息
usage: Option<TokenUsage>,
}
/// Anthropic SSE 解析结果
#[derive(Debug, Clone)]
pub struct AnthropicParseResult {
/// 文本增量
pub text_delta: Option<String>,
/// 是否完成
pub is_done: bool,
/// 工具调用开始(id, name)
pub tool_start: Option<(String, String)>,
}
impl AnthropicSSEParser {
pub fn new() -> Self {
Self::default()
}
/// 解析 SSE 数据行
///
/// 返回解析结果
pub fn parse_data(&mut self, data: &str) -> AnthropicParseResult {
if data.trim().is_empty() {
return AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: None,
};
}
let event: AnthropicStreamEvent = match serde_json::from_str(data) {
Ok(e) => e,
Err(e) => {
warn!("[AnthropicSSEParser] 解析事件失败: {} - data: {}", e, data);
return AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: None,
};
}
};
match event {
AnthropicStreamEvent::MessageStart { message } => {
debug!("[AnthropicSSEParser] 消息开始: id={}", message.id);
AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: None,
}
}
AnthropicStreamEvent::ContentBlockStart {
index,
content_block,
} => match content_block {
AnthropicContentBlock::ToolUse { id, name, .. } => {
debug!(
"[AnthropicSSEParser] 工具调用开始: id={}, name={}",
id, name
);
self.current_tool = Some(AnthropicToolCallBuilder {
id: id.clone(),
name: name.clone(),
input_json: String::new(),
});
AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: Some((id, name)),
}
}
AnthropicContentBlock::Text { .. } => {
debug!("[AnthropicSSEParser] 文本块开始: index={}", index);
AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: None,
}
}
_ => AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: None,
},
},
AnthropicStreamEvent::ContentBlockDelta { index: _, delta } => match delta {
AnthropicDelta::TextDelta { text } => {
self.full_content.push_str(&text);
AnthropicParseResult {
text_delta: Some(text),
is_done: false,
tool_start: None,
}
}
AnthropicDelta::InputJsonDelta { partial_json } => {
if let Some(ref mut tool) = self.current_tool {
tool.input_json.push_str(&partial_json);
}
AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: None,
}
}
AnthropicDelta::ThinkingDelta { thinking } => {
// 将思考内容添加到 full_content 中,用 <think> 标签包裹
let thinking_text = format!("<think>{}</think>", thinking);
self.full_content.push_str(&thinking_text);
AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: None,
}
}
AnthropicDelta::SignatureDelta { .. } => {
// 忽略签名 delta
AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: None,
}
}
},
AnthropicStreamEvent::ContentBlockStop { index: _ } => {
// 如果有正在构建的工具调用,完成它
if let Some(tool) = self.current_tool.take() {
self.tool_calls.push(ToolCall {
id: tool.id,
call_type: "function".to_string(),
function: FunctionCall {
name: tool.name,
arguments: tool.input_json,
},
});
}
AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: None,
}
}
AnthropicStreamEvent::MessageDelta { delta: _, usage } => {
self.usage = Some(TokenUsage::new(usage.input_tokens, usage.output_tokens));
AnthropicParseResult {
text_delta: None,
is_done: false,
tool_start: None,
}
}
AnthropicStreamEvent::MessageStop => {
debug!("[AnthropicSSEParser] 消息结束");
AnthropicParseResult {
text_delta: None,
is_done: true,
tool_start: None,
}
}
}
}
/// 完成解析,返回最终的工具调用列表
pub fn finalize_tool_calls(&mut self) -> Vec<ToolCall> {
std::mem::take(&mut self.tool_calls)
}
/// 获取完整内容
pub fn get_full_content(&self) -> String {
self.full_content.clone()
}
/// 是否有工具调用
pub fn has_tool_calls(&self) -> bool {
!self.tool_calls.is_empty() || self.current_tool.is_some()
}
/// 获取 usage
pub fn get_usage(&self) -> Option<TokenUsage> {
self.usage.clone()
}
}
-9
View File
@@ -1,9 +0,0 @@
//! SSE 流解析器模块
//!
//! 提供不同协议的 SSE 流解析器
mod anthropic_sse;
mod openai_sse;
pub use anthropic_sse::{AnthropicParseResult, AnthropicSSEParser};
pub use openai_sse::OpenAISSEParser;
-318
View File
@@ -1,318 +0,0 @@
//! OpenAI SSE 流解析器
//!
//! 解析 OpenAI 兼容 API 的 Server-Sent Events 流
//! Requirements: 1.1, 1.3, 1.4
use crate::agent::types::{FunctionCall, TokenUsage, ToolCall};
use serde_json::Value;
use std::collections::HashMap;
use tracing::warn;
/// 工具调用增量数据
#[derive(Debug, Clone, Default)]
struct ToolCallDelta {
/// 工具调用索引
#[allow(dead_code)]
index: usize,
/// 工具调用 ID
id: String,
/// 工具类型
call_type: String,
/// 函数名
function_name: String,
/// 函数参数(累积的 JSON 字符串)
function_arguments: String,
}
/// OpenAI SSE 流解析器
///
/// 解析 Server-Sent Events 流,提取 text_delta 和 tool_calls
#[derive(Debug, Default)]
pub struct OpenAISSEParser {
/// 累积的完整内容
full_content: String,
/// 累积的推理内容(DeepSeek R1 等模型)
reasoning_content: String,
/// 当前正在构建的工具调用索引
current_tool_indices: HashMap<usize, ToolCallDelta>,
}
impl OpenAISSEParser {
pub fn new() -> Self {
Self::default()
}
/// 解析 SSE 数据行
///
/// 返回 (text_delta, reasoning_delta, is_done, usage)
/// - text_delta: 普通文本内容增量
/// - reasoning_delta: 推理内容增量(DeepSeek reasoner 等模型)
pub fn parse_data(
&mut self,
data: &str,
) -> (Option<String>, Option<String>, bool, Option<TokenUsage>) {
if data.trim() == "[DONE]" {
return (None, None, true, None);
}
let json: Value = match serde_json::from_str(data) {
Ok(v) => v,
Err(e) => {
warn!("[OpenAISSEParser] 解析 JSON 失败: {} - data: {}", e, data);
return (None, None, false, None);
}
};
// 提取 usage 信息(如果存在)
let usage = json.get("usage").and_then(|u| {
let input = u.get("prompt_tokens").and_then(|v| v.as_u64()).unwrap_or(0) as u32;
let output = u
.get("completion_tokens")
.and_then(|v| v.as_u64())
.unwrap_or(0) as u32;
if input > 0 || output > 0 {
Some(TokenUsage::new(input, output))
} else {
None
}
});
// 检查是否有 choices
let choices = match json.get("choices").and_then(|c| c.as_array()) {
Some(c) => c,
None => return (None, None, false, usage),
};
if choices.is_empty() {
return (None, None, false, usage);
}
let choice = &choices[0];
let delta = match choice.get("delta") {
Some(d) => d,
None => return (None, None, false, usage),
};
// 检查 finish_reason
let finish_reason = choice
.get("finish_reason")
.and_then(|f| f.as_str())
.unwrap_or("");
let is_done = finish_reason == "stop" || finish_reason == "tool_calls";
// 提取文本内容
let text_delta = delta
.get("content")
.and_then(|c| c.as_str())
.filter(|s| !s.is_empty())
.map(|s| {
self.full_content.push_str(s);
s.to_string()
});
// 提取推理内容(DeepSeek reasoner 等模型)
let reasoning_delta = delta
.get("reasoning_content")
.and_then(|c| c.as_str())
.filter(|s| !s.is_empty())
.map(|s| {
self.reasoning_content.push_str(s);
s.to_string()
});
// 提取工具调用
if let Some(tool_calls) = delta.get("tool_calls").and_then(|tc| tc.as_array()) {
for tc in tool_calls {
self.parse_tool_call_delta(tc);
}
}
(text_delta, reasoning_delta, is_done, usage)
}
/// 解析工具调用增量
fn parse_tool_call_delta(&mut self, tc: &Value) {
let index = tc.get("index").and_then(|i| i.as_u64()).unwrap_or(0) as usize;
// 获取或创建工具调用
let tool_call = self
.current_tool_indices
.entry(index)
.or_insert_with(|| ToolCallDelta {
index,
..Default::default()
});
// 更新 ID
if let Some(id) = tc.get("id").and_then(|i| i.as_str()) {
tool_call.id = id.to_string();
}
// 更新类型
if let Some(t) = tc.get("type").and_then(|t| t.as_str()) {
tool_call.call_type = t.to_string();
}
// 更新函数信息
if let Some(function) = tc.get("function") {
if let Some(name) = function.get("name").and_then(|n| n.as_str()) {
tool_call.function_name = name.to_string();
}
if let Some(args) = function.get("arguments").and_then(|a| a.as_str()) {
tool_call.function_arguments.push_str(args);
}
}
}
/// 完成解析,返回最终的工具调用列表
pub fn finalize_tool_calls(&mut self) -> Vec<ToolCall> {
// 按索引排序并转换为 ToolCall
let mut indices: Vec<_> = self.current_tool_indices.keys().cloned().collect();
indices.sort();
indices
.into_iter()
.filter_map(|idx| {
let delta = self.current_tool_indices.get(&idx)?;
if delta.id.is_empty() || delta.function_name.is_empty() {
return None;
}
Some(ToolCall {
id: delta.id.clone(),
call_type: if delta.call_type.is_empty() {
"function".to_string()
} else {
delta.call_type.clone()
},
function: FunctionCall {
name: delta.function_name.clone(),
arguments: delta.function_arguments.clone(),
},
})
})
.collect()
}
/// 获取完整内容
pub fn get_full_content(&self) -> String {
self.full_content.clone()
}
/// 获取推理内容(DeepSeek R1 等模型)
pub fn get_reasoning_content(&self) -> Option<String> {
if self.reasoning_content.is_empty() {
None
} else {
Some(self.reasoning_content.clone())
}
}
/// 是否有工具调用
pub fn has_tool_calls(&self) -> bool {
!self.current_tool_indices.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_text_delta() {
let mut parser = OpenAISSEParser::new();
let data1 = r#"{"choices":[{"delta":{"content":"Hello"}}]}"#;
let data2 = r#"{"choices":[{"delta":{"content":" World"}}]}"#;
let data3 = r#"{"choices":[{"delta":{},"finish_reason":"stop"}]}"#;
let (text1, _, done1, _) = parser.parse_data(data1);
assert_eq!(text1, Some("Hello".to_string()));
assert!(!done1);
let (text2, _, done2, _) = parser.parse_data(data2);
assert_eq!(text2, Some(" World".to_string()));
assert!(!done2);
let (text3, _, done3, _) = parser.parse_data(data3);
assert!(text3.is_none());
assert!(done3);
assert_eq!(parser.get_full_content(), "Hello World");
}
#[test]
fn test_reasoning_content() {
let mut parser = OpenAISSEParser::new();
let data1 = r#"{"choices":[{"delta":{"reasoning_content":"Let me think"}}]}"#;
let data2 = r#"{"choices":[{"delta":{"reasoning_content":" about this"}}]}"#;
let data3 = r#"{"choices":[{"delta":{"content":"The answer is 42"}}]}"#;
let (text1, reasoning1, done1, _) = parser.parse_data(data1);
assert!(text1.is_none());
assert_eq!(reasoning1, Some("Let me think".to_string()));
assert!(!done1);
let (text2, reasoning2, done2, _) = parser.parse_data(data2);
assert!(text2.is_none());
assert_eq!(reasoning2, Some(" about this".to_string()));
assert!(!done2);
let (text3, reasoning3, done3, _) = parser.parse_data(data3);
assert_eq!(text3, Some("The answer is 42".to_string()));
assert!(reasoning3.is_none());
assert!(!done3);
assert_eq!(parser.get_full_content(), "The answer is 42");
assert_eq!(
parser.get_reasoning_content(),
Some("Let me think about this".to_string())
);
}
#[test]
fn test_tool_calls() {
let mut parser = OpenAISSEParser::new();
let data1 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_123","type":"function","function":{"name":"bash"}}]}}]}"#;
let data2 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"command\":"}}]}}]}"#;
let data3 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"ls -la\"}"}}]}}]}"#;
let data4 = r#"{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}"#;
parser.parse_data(data1);
parser.parse_data(data2);
parser.parse_data(data3);
let (_, _, done, _) = parser.parse_data(data4);
assert!(done);
assert!(parser.has_tool_calls());
let tool_calls = parser.finalize_tool_calls();
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0].id, "call_123");
assert_eq!(tool_calls[0].function.name, "bash");
assert_eq!(tool_calls[0].function.arguments, r#"{"command":"ls -la"}"#);
}
#[test]
fn test_usage() {
let mut parser = OpenAISSEParser::new();
let data = r#"{"choices":[{"delta":{"content":"Hi"}}],"usage":{"prompt_tokens":10,"completion_tokens":5}}"#;
let (text, _, _, usage) = parser.parse_data(data);
assert_eq!(text, Some("Hi".to_string()));
assert!(usage.is_some());
let usage = usage.unwrap();
assert_eq!(usage.input_tokens, 10);
assert_eq!(usage.output_tokens, 5);
}
#[test]
fn test_done_signal() {
let mut parser = OpenAISSEParser::new();
let (_, _, done, _) = parser.parse_data("[DONE]");
assert!(done);
}
}
-668
View File
@@ -1,668 +0,0 @@
//! Anthropic 协议实现
//!
//! 实现 Anthropic Messages API 协议
//! 适用于 Claude、Claude OAuth 等 Anthropic 服务
use super::Protocol;
use crate::agent::parsers::AnthropicSSEParser;
use crate::agent::types::{
AgentConfig, AgentMessage, ContentPart, ImageData, MessageContent, StreamEvent, StreamResult,
};
use crate::models::anthropic::AnthropicMessage;
use crate::models::openai::Tool;
use async_trait::async_trait;
use futures::StreamExt;
use reqwest::Client;
use serde::Serialize;
use tokio::sync::mpsc;
use tracing::{debug, error, info, warn};
/// Anthropic Messages API 请求
#[derive(Debug, Serialize)]
struct AnthropicMessagesRequest {
model: String,
messages: Vec<AnthropicMessage>,
max_tokens: u32,
stream: bool,
#[serde(skip_serializing_if = "Option::is_none")]
system: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
tools: Option<Vec<AnthropicTool>>,
}
/// Anthropic 工具定义
#[derive(Debug, Serialize)]
struct AnthropicTool {
name: String,
description: String,
input_schema: serde_json::Value,
}
/// Anthropic 协议处理器
pub struct AnthropicProtocol {
/// 是否使用数组格式的 system 字段
/// 标准 Anthropic: system: "prompt"
/// 兼容格式: system: [{"type": "text", "text": "prompt"}]
use_array_system_format: bool,
}
impl AnthropicProtocol {
/// 创建标准 Anthropic 协议处理器
pub fn new() -> Self {
Self {
use_array_system_format: false,
}
}
/// 创建兼容格式的 Anthropic 协议处理器
pub fn with_array_system_format() -> Self {
Self {
use_array_system_format: true,
}
}
/// 构建 system 字段
fn build_system_field(&self, prompt: &str) -> serde_json::Value {
if self.use_array_system_format {
serde_json::json!([{
"type": "text",
"text": prompt
}])
} else {
serde_json::json!(prompt)
}
}
/// 将 OpenAI Tool 转换为 Anthropic Tool
fn convert_tools(tools: Option<&[Tool]>) -> Option<Vec<AnthropicTool>> {
tools.map(|t| {
t.iter()
.filter_map(|tool| match tool {
Tool::Function { function } => Some(AnthropicTool {
name: function.name.clone(),
description: function.description.clone().unwrap_or_default(),
input_schema: function.parameters.clone().unwrap_or(serde_json::json!({
"type": "object",
"properties": {}
})),
}),
// WebSearch 工具不支持转换为 Anthropic 格式,跳过
Tool::WebSearch | Tool::WebSearch20250305 => None,
})
.collect()
})
}
/// 将 AgentMessage 转换为 Anthropic Message
fn convert_to_anthropic_message(msg: &AgentMessage) -> AnthropicMessage {
let content = match &msg.content {
MessageContent::Text(text) => {
// 处理工具结果消息
if msg.role == "tool" {
// Anthropic 使用 tool_result content block
if let Some(tool_call_id) = &msg.tool_call_id {
serde_json::json!([{
"type": "tool_result",
"tool_use_id": tool_call_id,
"content": text
}])
} else {
// 如果没有 tool_call_id,这可能是一个错误的工具结果消息
warn!("[AnthropicProtocol] 工具结果消息缺少 tool_call_id");
serde_json::json!(text)
}
} else if msg.role == "assistant" {
// 处理 assistant 消息
let mut blocks = Vec::new();
if !text.is_empty() {
blocks.push(serde_json::json!({
"type": "text",
"text": text
}));
}
// 添加工具调用
if let Some(tool_calls) = &msg.tool_calls {
for tc in tool_calls {
let input: serde_json::Value =
serde_json::from_str(&tc.function.arguments)
.unwrap_or(serde_json::json!({}));
blocks.push(serde_json::json!({
"type": "tool_use",
"id": tc.id,
"name": tc.function.name,
"input": input
}));
}
}
if blocks.is_empty() {
serde_json::json!("")
} else if blocks.len() == 1 && msg.tool_calls.is_none() {
serde_json::json!(text)
} else {
serde_json::json!(blocks)
}
} else {
serde_json::json!(text)
}
}
MessageContent::Parts(parts) => {
let blocks: Vec<serde_json::Value> = parts
.iter()
.map(|p| match p {
ContentPart::Text { text } => serde_json::json!({
"type": "text",
"text": text
}),
ContentPart::ImageUrl { image_url } => {
// 解析 data URL
if let Some(rest) = image_url.url.strip_prefix("data:") {
if let Some(comma_idx) = rest.find(',') {
let media_type = rest[..comma_idx]
.strip_suffix(";base64")
.unwrap_or(&rest[..comma_idx]);
let data = &rest[comma_idx + 1..];
return serde_json::json!({
"type": "image",
"source": {
"type": "base64",
"media_type": media_type,
"data": data
}
});
}
}
// 普通 URL
serde_json::json!({
"type": "image",
"source": {
"type": "url",
"url": image_url.url
}
})
}
})
.collect();
serde_json::json!(blocks)
}
};
// Anthropic 没有 "tool" 角色,需要转换为 "user"
let role = if msg.role == "tool" {
"user".to_string()
} else {
msg.role.clone()
};
AnthropicMessage { role, content }
}
/// 构建消息列表
fn build_messages(
&self,
history: &[AgentMessage],
user_message: &str,
images: Option<&[ImageData]>,
config: &AgentConfig,
) -> (Vec<AnthropicMessage>, Option<serde_json::Value>) {
let mut messages = Vec::new();
// 系统提示词(Anthropic 使用单独的 system 字段)
let system_prompt = config
.system_prompt
.as_ref()
.map(|s| self.build_system_field(s));
// 验证和修复消息序列中的 tool_use/tool_result 配对
let validated_history = self.validate_tool_message_pairs(history);
// 添加历史消息(跳过 system 消息)
for msg in &validated_history {
if msg.role == "system" {
continue;
}
messages.push(Self::convert_to_anthropic_message(msg));
}
// 添加当前用户消息
let user_content = if let Some(imgs) = images {
let mut parts = vec![serde_json::json!({
"type": "text",
"text": user_message
})];
for img in imgs {
parts.push(serde_json::json!({
"type": "image",
"source": {
"type": "base64",
"media_type": img.media_type,
"data": img.data
}
}));
}
serde_json::json!(parts)
} else {
serde_json::json!(user_message)
};
messages.push(AnthropicMessage {
role: "user".to_string(),
content: user_content,
});
(messages, system_prompt)
}
/// 从历史构建消息(不添加新用户消息)
fn build_messages_from_history(
&self,
history: &[AgentMessage],
config: &AgentConfig,
) -> (Vec<AnthropicMessage>, Option<serde_json::Value>) {
let mut messages = Vec::new();
// 系统提示词
let system_prompt = config
.system_prompt
.as_ref()
.map(|s| self.build_system_field(s));
// 验证和修复消息序列中的 tool_use/tool_result 配对
let validated_history = self.validate_tool_message_pairs(history);
// 添加所有历史消息(跳过 system)
for msg in &validated_history {
if msg.role == "system" {
continue;
}
messages.push(Self::convert_to_anthropic_message(msg));
}
(messages, system_prompt)
}
/// 验证并修复消息序列中的 tool_use/tool_result 配对
///
/// Claude API 要求每个 tool_use 都必须紧跟一个对应的 tool_result
fn validate_tool_message_pairs(&self, history: &[AgentMessage]) -> Vec<AgentMessage> {
let mut validated_messages = Vec::new();
let mut pending_tool_calls: std::collections::HashMap<String, bool> =
std::collections::HashMap::new();
for msg in history {
match msg.role.as_str() {
"assistant" => {
// 检查是否有工具调用
if let Some(tool_calls) = &msg.tool_calls {
for tc in tool_calls {
pending_tool_calls.insert(tc.id.clone(), false);
}
}
validated_messages.push(msg.clone());
}
"tool" => {
// 检查工具结果是否有对应的工具调用
if let Some(tool_call_id) = &msg.tool_call_id {
if pending_tool_calls.contains_key(tool_call_id) {
pending_tool_calls.insert(tool_call_id.clone(), true);
validated_messages.push(msg.clone());
} else {
warn!(
"[AnthropicProtocol] 发现孤立的工具结果消息,tool_call_id: {}",
tool_call_id
);
// 跳过孤立的工具结果消息
}
} else {
warn!("[AnthropicProtocol] 工具结果消息缺少 tool_call_id");
// 跳过无效的工具结果消息
}
}
_ => {
validated_messages.push(msg.clone());
}
}
}
// 检查是否有未配对的工具调用
for (tool_call_id, has_result) in &pending_tool_calls {
if !has_result {
warn!(
"[AnthropicProtocol] 发现未配对的工具调用,tool_call_id: {},添加默认工具结果",
tool_call_id
);
// 为未配对的工具调用添加默认结果
let default_result = AgentMessage {
role: "tool".to_string(),
content: MessageContent::Text("工具执行超时或失败".to_string()),
timestamp: chrono::Utc::now().to_rfc3339(),
tool_calls: None,
tool_call_id: Some(tool_call_id.clone()),
reasoning_content: None,
};
validated_messages.push(default_result);
}
}
validated_messages
}
/// 处理 SSE 流
async fn process_stream(
response: reqwest::Response,
tx: mpsc::Sender<StreamEvent>,
send_done: bool,
) -> Result<StreamResult, String> {
let mut stream = response.bytes_stream();
let mut buffer = String::new();
let mut parser = AnthropicSSEParser::new();
while let Some(chunk) = stream.next().await {
match chunk {
Ok(bytes) => {
let text = String::from_utf8_lossy(&bytes);
buffer.push_str(&text);
// 处理完整的 SSE 事件
while let Some(pos) = buffer.find("\n\n") {
let event_block = buffer[..pos].to_string();
buffer = buffer[pos + 2..].to_string();
// 提取 event 类型和 data
let mut event_type = String::new();
let mut data = String::new();
for line in event_block.lines() {
if let Some(e) = line.strip_prefix("event: ") {
event_type = e.to_string();
} else if let Some(d) = line.strip_prefix("data: ") {
data = d.to_string();
}
}
if data.is_empty() {
continue;
}
debug!(
"[AnthropicProtocol] SSE event={}, data={}",
event_type, data
);
let result = parser.parse_data(&data);
// 发送工具开始事件
if let Some((tool_id, tool_name)) = result.tool_start {
let _ = tx
.send(StreamEvent::ToolStart {
tool_name,
tool_id,
arguments: None,
})
.await;
}
// 发送文本增量
if let Some(text) = result.text_delta {
let _ = tx.send(StreamEvent::TextDelta { text }).await;
}
// 检查是否完成
if result.is_done {
let full_content = parser.get_full_content();
let tool_calls = if parser.has_tool_calls() {
Some(parser.finalize_tool_calls())
} else {
None
};
let usage = parser.get_usage();
if send_done {
let _ = tx
.send(StreamEvent::Done {
usage: usage.clone(),
})
.await;
}
return Ok(StreamResult {
content: full_content,
tool_calls,
usage,
reasoning_content: None,
});
}
}
}
Err(e) => {
error!("[AnthropicProtocol] 流读取错误: {}", e);
let _ = tx
.send(StreamEvent::Error {
message: format!("流读取错误: {}", e),
})
.await;
return Err(format!("流读取错误: {}", e));
}
}
}
// 流正常结束
let full_content = parser.get_full_content();
let tool_calls = if parser.has_tool_calls() {
Some(parser.finalize_tool_calls())
} else {
None
};
let usage = parser.get_usage();
if send_done {
let _ = tx
.send(StreamEvent::Done {
usage: usage.clone(),
})
.await;
}
Ok(StreamResult {
content: full_content,
tool_calls,
usage,
reasoning_content: None,
})
}
}
#[async_trait]
impl Protocol for AnthropicProtocol {
async fn chat_stream(
&self,
client: &Client,
base_url: &str,
api_key: &str,
messages: &[AgentMessage],
user_message: &str,
images: Option<&[ImageData]>,
model: &str,
config: &AgentConfig,
tools: Option<&[Tool]>,
tx: mpsc::Sender<StreamEvent>,
provider_id: Option<&str>,
) -> Result<StreamResult, String> {
info!(
"[AnthropicProtocol] 发送流式请求: model={}, history_len={}, tools_count={}, provider_id={:?}",
model,
messages.len(),
tools.map(|t| t.len()).unwrap_or(0),
provider_id
);
let (anthropic_messages, system) =
self.build_messages(messages, user_message, images, config);
let anthropic_tools = Self::convert_tools(tools);
let request = AnthropicMessagesRequest {
model: model.to_string(),
messages: anthropic_messages,
max_tokens: config.max_tokens.unwrap_or(4096),
stream: true,
system,
temperature: config.temperature,
tools: anthropic_tools,
};
let url = format!("{}{}", base_url, self.endpoint());
info!(
"[AnthropicProtocol] 请求详情: url={}, model={}, messages_count={}, has_tools={}",
url,
model,
request.messages.len(),
request.tools.is_some()
);
let mut req_builder = client
.post(&url)
.header("Authorization", format!("Bearer {}", api_key))
.header("Content-Type", "application/json")
.header("anthropic-version", "2023-06-01");
let response = req_builder
.json(&request)
.send()
.await
.map_err(|e| format!("请求失败: {}", e))?;
let status = response.status();
if !status.is_success() {
let body = response.text().await.unwrap_or_default();
error!("[AnthropicProtocol] 请求失败: {} - {}", status, body);
let _ = tx
.send(StreamEvent::Error {
message: format!("API 错误 ({}): {}", status, body),
})
.await;
return Err(format!("API 错误: {}", status));
}
Self::process_stream(response, tx, true).await
}
async fn chat_stream_continue(
&self,
client: &Client,
base_url: &str,
api_key: &str,
messages: &[AgentMessage],
model: &str,
config: &AgentConfig,
tools: Option<&[Tool]>,
tx: mpsc::Sender<StreamEvent>,
provider_id: Option<&str>,
) -> Result<StreamResult, String> {
debug!(
"[AnthropicProtocol] 继续流式对话: model={}, history_len={}, tools_count={}, provider_id={:?}",
model,
messages.len(),
tools.map(|t| t.len()).unwrap_or(0),
provider_id
);
let (anthropic_messages, system) = self.build_messages_from_history(messages, config);
let anthropic_tools = Self::convert_tools(tools);
let request = AnthropicMessagesRequest {
model: model.to_string(),
messages: anthropic_messages,
max_tokens: config.max_tokens.unwrap_or(4096),
stream: true,
system,
temperature: config.temperature,
tools: anthropic_tools,
};
let url = format!("{}{}", base_url, self.endpoint());
info!(
"[AnthropicProtocol] 继续请求详情: url={}, model={}, messages_count={}, has_tools={}",
url,
model,
request.messages.len(),
request.tools.is_some()
);
// 调试:打印消息序列
for (i, msg) in request.messages.iter().enumerate() {
debug!(
"[AnthropicProtocol] Message {}: role={}, content_type={}",
i,
msg.role,
if msg.content.is_string() {
"string"
} else {
"array"
}
);
if let Some(content_array) = msg.content.as_array() {
for (j, block) in content_array.iter().enumerate() {
if let Some(block_type) = block.get("type").and_then(|t| t.as_str()) {
debug!(
"[AnthropicProtocol] Message {} Block {}: type={}",
i, j, block_type
);
if block_type == "tool_result" {
if let Some(tool_use_id) =
block.get("tool_use_id").and_then(|id| id.as_str())
{
debug!(
"[AnthropicProtocol] tool_result for tool_use_id: {}",
tool_use_id
);
}
} else if block_type == "tool_use" {
if let Some(tool_id) = block.get("id").and_then(|id| id.as_str()) {
debug!("[AnthropicProtocol] tool_use with id: {}", tool_id);
}
}
}
}
}
}
let mut req_builder = client
.post(&url)
.header("Authorization", format!("Bearer {}", api_key))
.header("Content-Type", "application/json")
.header("anthropic-version", "2023-06-01");
let response = req_builder
.json(&request)
.send()
.await
.map_err(|e| format!("请求失败: {}", e))?;
let status = response.status();
if !status.is_success() {
let body = response.text().await.unwrap_or_default();
error!("[AnthropicProtocol] 请求失败: {} - {}", status, body);
let _ = tx
.send(StreamEvent::Error {
message: format!("API 错误 ({}): {}", status, body),
})
.await;
return Err(format!("API 错误: {}", status));
}
// 继续对话时不发送 Done 事件
Self::process_stream(response, tx, false).await
}
fn endpoint(&self) -> &'static str {
"/v1/messages"
}
}
-76
View File
@@ -1,76 +0,0 @@
//! 协议策略模块
//!
//! 使用策略模式处理不同 API 协议(OpenAI、Anthropic、Kiro、Gemini)
mod anthropic;
mod openai;
pub use anthropic::AnthropicProtocol;
pub use openai::OpenAIProtocol;
use crate::agent::types::{
AgentConfig, AgentMessage, ImageData, ProviderType, StreamEvent, StreamResult,
};
use crate::models::openai::Tool;
use async_trait::async_trait;
use reqwest::Client;
use tokio::sync::mpsc;
/// 协议处理器 trait
///
/// 定义了所有协议必须实现的方法
#[async_trait]
pub trait Protocol: Send + Sync {
/// 流式聊天
///
/// 发送消息并通过 channel 返回流式响应
async fn chat_stream(
&self,
client: &Client,
base_url: &str,
api_key: &str,
messages: &[AgentMessage],
user_message: &str,
images: Option<&[ImageData]>,
model: &str,
config: &AgentConfig,
tools: Option<&[Tool]>,
tx: mpsc::Sender<StreamEvent>,
provider_id: Option<&str>,
) -> Result<StreamResult, String>;
/// 继续流式对话(工具调用后)
///
/// 使用会话历史继续对话,不添加新的用户消息
async fn chat_stream_continue(
&self,
client: &Client,
base_url: &str,
api_key: &str,
messages: &[AgentMessage],
model: &str,
config: &AgentConfig,
tools: Option<&[Tool]>,
tx: mpsc::Sender<StreamEvent>,
provider_id: Option<&str>,
) -> Result<StreamResult, String>;
/// 获取 API 端点
fn endpoint(&self) -> &'static str;
}
/// 根据 ProviderType 创建协议处理器
pub fn create_protocol(provider_type: ProviderType) -> Box<dyn Protocol> {
match provider_type {
// Claude 和 Kiro 使用标准 Anthropic SSE 协议
ProviderType::Claude | ProviderType::ClaudeOauth | ProviderType::Kiro => {
Box::new(AnthropicProtocol::new())
}
// Anthropic 兼容格式(system 为数组格式)
ProviderType::AnthropicCompatible => {
Box::new(AnthropicProtocol::with_array_system_format())
}
// 其他使用 OpenAI 兼容协议
_ => Box::new(OpenAIProtocol),
}
}
-564
View File
@@ -1,564 +0,0 @@
//! OpenAI 协议实现
//!
//! 实现 OpenAI Chat Completions API 协议
//! 适用于 OpenAI、Qwen、Codex、Antigravity、IFlow、Kiro 等兼容服务
use super::Protocol;
use crate::agent::parsers::OpenAISSEParser;
use crate::agent::types::{
AgentConfig, AgentMessage, ContentPart, ImageData, MessageContent, StreamEvent, StreamResult,
};
use crate::models::openai::{
ChatCompletionRequest, ChatMessage, ContentPart as OpenAIContentPart,
MessageContent as OpenAIMessageContent, Tool,
};
use async_trait::async_trait;
use futures::StreamExt;
use reqwest::Client;
use tokio::sync::mpsc;
use tracing::{debug, error, info};
/// OpenAI 协议处理器
pub struct OpenAIProtocol;
impl OpenAIProtocol {
/// 将 AgentMessage 转换为 OpenAI ChatMessage
fn convert_to_chat_message(msg: &AgentMessage, is_deepseek_reasoner: bool) -> ChatMessage {
let content = match &msg.content {
MessageContent::Text(text) => Some(OpenAIMessageContent::Text(text.clone())),
MessageContent::Parts(parts) => {
let openai_parts: Vec<OpenAIContentPart> = parts
.iter()
.map(|p| match p {
ContentPart::Text { text } => {
OpenAIContentPart::Text { text: text.clone() }
}
ContentPart::ImageUrl { image_url } => OpenAIContentPart::ImageUrl {
image_url: crate::models::openai::ImageUrl {
url: image_url.url.clone(),
detail: image_url.detail.clone(),
},
},
})
.collect();
Some(OpenAIMessageContent::Parts(openai_parts))
}
};
// DeepSeek reasoner 模型要求 assistant 消息必须包含 reasoning_content 字段
// 参考: https://api-docs.deepseek.com/guides/thinking_mode#tool-calls
let reasoning_content = if is_deepseek_reasoner && msg.role == "assistant" {
// 对于 DeepSeek reasoner,确保 reasoning_content 始终有值(即使为空字符串)
Some(msg.reasoning_content.clone().unwrap_or_default())
} else {
msg.reasoning_content.clone()
};
ChatMessage {
role: msg.role.clone(),
content,
tool_calls: msg.tool_calls.as_ref().map(|calls| {
calls
.iter()
.map(|tc| crate::models::openai::ToolCall {
id: tc.id.clone(),
call_type: tc.call_type.clone(),
function: crate::models::openai::FunctionCall {
name: tc.function.name.clone(),
arguments: tc.function.arguments.clone(),
},
})
.collect()
}),
tool_call_id: msg.tool_call_id.clone(),
reasoning_content,
}
}
/// 检查模型是否是 DeepSeek reasoner 模型
fn is_deepseek_reasoner(model: &str) -> bool {
model.contains("deepseek-reasoner") || model.contains("deepseek-r1")
}
/// 构建消息列表
fn build_messages(
history: &[AgentMessage],
user_message: &str,
images: Option<&[ImageData]>,
config: &AgentConfig,
model: &str,
) -> Vec<ChatMessage> {
let mut messages = Vec::new();
let is_deepseek_reasoner = Self::is_deepseek_reasoner(model);
// 添加系统提示词
if let Some(prompt) = &config.system_prompt {
messages.push(ChatMessage {
role: "system".to_string(),
content: Some(OpenAIMessageContent::Text(prompt.clone())),
tool_calls: None,
tool_call_id: None,
reasoning_content: None,
});
}
// 添加历史消息
for msg in history {
messages.push(Self::convert_to_chat_message(msg, is_deepseek_reasoner));
}
// 添加当前用户消息
let user_msg = if let Some(imgs) = images {
let mut parts = vec![OpenAIContentPart::Text {
text: user_message.to_string(),
}];
for img in imgs {
parts.push(OpenAIContentPart::ImageUrl {
image_url: crate::models::openai::ImageUrl {
url: format!("data:{};base64,{}", img.media_type, img.data),
detail: None,
},
});
}
ChatMessage {
role: "user".to_string(),
content: Some(OpenAIMessageContent::Parts(parts)),
tool_calls: None,
tool_call_id: None,
reasoning_content: None,
}
} else {
ChatMessage {
role: "user".to_string(),
content: Some(OpenAIMessageContent::Text(user_message.to_string())),
tool_calls: None,
tool_call_id: None,
reasoning_content: None,
}
};
messages.push(user_msg);
messages
}
/// 从历史构建消息(不添加新用户消息)
fn build_messages_from_history(
history: &[AgentMessage],
config: &AgentConfig,
model: &str,
) -> Vec<ChatMessage> {
let mut messages = Vec::new();
let is_deepseek_reasoner = Self::is_deepseek_reasoner(model);
// 添加系统提示词
if let Some(prompt) = &config.system_prompt {
messages.push(ChatMessage {
role: "system".to_string(),
content: Some(OpenAIMessageContent::Text(prompt.clone())),
tool_calls: None,
tool_call_id: None,
reasoning_content: None,
});
}
// 添加所有历史消息
for msg in history {
messages.push(Self::convert_to_chat_message(msg, is_deepseek_reasoner));
}
messages
}
/// 处理 SSE 流
async fn process_stream(
response: reqwest::Response,
tx: mpsc::Sender<StreamEvent>,
send_done: bool,
) -> Result<StreamResult, String> {
let mut stream = response.bytes_stream();
let mut buffer = String::new();
let mut parser = OpenAISSEParser::new();
let mut final_usage = None;
eprintln!("[OpenAIProtocol] 开始处理 SSE 流...");
while let Some(chunk) = stream.next().await {
match chunk {
Ok(bytes) => {
let text = String::from_utf8_lossy(&bytes);
// 安全截断:使用 char_indices 找到有效的 UTF-8 字符边界
let truncated = if text.len() > 200 {
let mut end = 200;
for (i, _) in text.char_indices() {
if i <= 200 {
end = i;
} else {
break;
}
}
format!("{}...", &text[..end])
} else {
text.to_string()
};
eprintln!(
"[OpenAIProtocol] 收到 chunk: {} bytes, 内容: {}",
bytes.len(),
truncated
);
buffer.push_str(&text);
// 检查是否是非流式响应(直接返回完整 JSON)
// 非流式响应以 { 开头,不是 SSE 格式
if buffer.trim().starts_with('{') && !buffer.contains("data: ") {
// 尝试解析为完整的 ChatCompletionResponse
if let Ok(response) = serde_json::from_str::<
crate::models::openai::ChatCompletionResponse,
>(&buffer)
{
eprintln!("[OpenAIProtocol] 检测到非流式响应,直接解析");
let content = response
.choices
.first()
.and_then(|c| c.message.content.clone())
.unwrap_or_default();
// 发送完整内容作为 TextDelta
if !content.is_empty() {
let _ = tx
.send(StreamEvent::TextDelta {
text: content.clone(),
})
.await;
}
let usage = Some(crate::agent::types::TokenUsage {
input_tokens: response.usage.prompt_tokens,
output_tokens: response.usage.completion_tokens,
});
if send_done {
let _ = tx
.send(StreamEvent::Done {
usage: usage.clone(),
})
.await;
}
return Ok(StreamResult {
content,
tool_calls: None,
usage,
reasoning_content: None,
});
}
}
// 处理完整的 SSE 事件(以 \n\n 分隔)
while let Some(pos) = buffer.find("\n\n") {
let event = buffer[..pos].to_string();
buffer = buffer[pos + 2..].to_string();
for line in event.lines() {
if let Some(data) = line.strip_prefix("data: ") {
debug!("[OpenAIProtocol] SSE data: {}", data);
let (text_delta, reasoning_delta, is_done, usage) =
parser.parse_data(data);
if usage.is_some() {
final_usage = usage;
}
// 发送推理内容增量(DeepSeek reasoner 等模型)
if let Some(reasoning) = reasoning_delta {
let _ = tx
.send(StreamEvent::ReasoningDelta { text: reasoning })
.await;
}
// 发送普通文本内容增量
if let Some(text) = text_delta {
let _ = tx.send(StreamEvent::TextDelta { text }).await;
}
if is_done {
let full_content = parser.get_full_content();
let reasoning_content = parser.get_reasoning_content();
let tool_calls = if parser.has_tool_calls() {
Some(parser.finalize_tool_calls())
} else {
None
};
// 如果没有普通内容但有推理内容,使用推理内容作为最终内容
let final_content = if full_content.is_empty() {
reasoning_content.clone().unwrap_or_default()
} else {
full_content
};
if send_done {
let _ = tx
.send(StreamEvent::Done {
usage: final_usage.clone(),
})
.await;
}
return Ok(StreamResult {
content: final_content,
tool_calls,
usage: final_usage,
reasoning_content,
});
}
}
}
}
}
Err(e) => {
error!("[OpenAIProtocol] 流读取错误: {}", e);
let _ = tx
.send(StreamEvent::Error {
message: format!("流读取错误: {}", e),
})
.await;
return Err(format!("流读取错误: {}", e));
}
}
}
// 流正常结束但没有收到 [DONE]
// 检查 buffer 中是否还有未处理的非流式响应
if !buffer.trim().is_empty() && buffer.trim().starts_with('{') {
if let Ok(response) =
serde_json::from_str::<crate::models::openai::ChatCompletionResponse>(&buffer)
{
eprintln!("[OpenAIProtocol] 流结束时检测到非流式响应");
let content = response
.choices
.first()
.and_then(|c| c.message.content.clone())
.unwrap_or_default();
if !content.is_empty() {
let _ = tx
.send(StreamEvent::TextDelta {
text: content.clone(),
})
.await;
}
let usage = Some(crate::agent::types::TokenUsage {
input_tokens: response.usage.prompt_tokens,
output_tokens: response.usage.completion_tokens,
});
if send_done {
let _ = tx
.send(StreamEvent::Done {
usage: usage.clone(),
})
.await;
}
return Ok(StreamResult {
content,
tool_calls: None,
usage,
reasoning_content: None,
});
}
}
let full_content = parser.get_full_content();
let reasoning_content = parser.get_reasoning_content();
let tool_calls = if parser.has_tool_calls() {
Some(parser.finalize_tool_calls())
} else {
None
};
// 如果没有普通内容但有推理内容,使用推理内容作为最终内容
let final_content = if full_content.is_empty() {
reasoning_content.clone().unwrap_or_default()
} else {
full_content
};
if send_done {
let _ = tx
.send(StreamEvent::Done {
usage: final_usage.clone(),
})
.await;
}
Ok(StreamResult {
content: final_content,
tool_calls,
usage: final_usage,
reasoning_content,
})
}
}
#[async_trait]
impl Protocol for OpenAIProtocol {
async fn chat_stream(
&self,
client: &Client,
base_url: &str,
api_key: &str,
messages: &[AgentMessage],
user_message: &str,
images: Option<&[ImageData]>,
model: &str,
config: &AgentConfig,
tools: Option<&[Tool]>,
tx: mpsc::Sender<StreamEvent>,
provider_id: Option<&str>,
) -> Result<StreamResult, String> {
info!(
"[OpenAIProtocol] 发送流式请求: model={}, history_len={}, tools_count={}, provider_id={:?}",
model,
messages.len(),
tools.map(|t| t.len()).unwrap_or(0),
provider_id
);
let chat_messages = Self::build_messages(messages, user_message, images, config, model);
let request = ChatCompletionRequest {
model: model.to_string(),
messages: chat_messages,
stream: true,
temperature: config.temperature,
max_tokens: config.max_tokens,
top_p: None,
tools: tools.map(|t| t.to_vec()),
tool_choice: if tools.is_some() {
Some(serde_json::json!("auto"))
} else {
None
},
reasoning_effort: None,
};
let url = format!("{}{}", base_url, self.endpoint());
eprintln!(
"[OpenAIProtocol] 发送请求到: {} model={} stream={} provider_id={:?}",
url, model, request.stream, provider_id
);
let mut req_builder = client
.post(&url)
.header("Authorization", format!("Bearer {}", api_key))
.header("Content-Type", "application/json");
// 添加 X-Provider-Id header 用于精确路由
if let Some(pid) = provider_id {
req_builder = req_builder.header("X-Provider-Id", pid);
}
let response = req_builder.json(&request).send().await.map_err(|e| {
eprintln!("[OpenAIProtocol] 请求发送失败: {}", e);
format!("请求失败: {}", e)
})?;
let status = response.status();
eprintln!("[OpenAIProtocol] 响应状态: {}", status);
if !status.is_success() {
let body = response.text().await.unwrap_or_default();
error!("[OpenAIProtocol] 请求失败: {} - {}", status, body);
let _ = tx
.send(StreamEvent::Error {
message: format!("API 错误 ({}): {}", status, body),
})
.await;
return Err(format!("API 错误: {}", status));
}
Self::process_stream(response, tx, true).await
}
async fn chat_stream_continue(
&self,
client: &Client,
base_url: &str,
api_key: &str,
messages: &[AgentMessage],
model: &str,
config: &AgentConfig,
tools: Option<&[Tool]>,
tx: mpsc::Sender<StreamEvent>,
provider_id: Option<&str>,
) -> Result<StreamResult, String> {
debug!(
"[OpenAIProtocol] 继续流式对话: model={}, history_len={}, tools_count={}, provider_id={:?}",
model,
messages.len(),
tools.map(|t| t.len()).unwrap_or(0),
provider_id
);
let chat_messages = Self::build_messages_from_history(messages, config, model);
let request = ChatCompletionRequest {
model: model.to_string(),
messages: chat_messages,
stream: true,
temperature: config.temperature,
max_tokens: config.max_tokens,
top_p: None,
tools: tools.map(|t| t.to_vec()),
tool_choice: if tools.is_some() {
Some(serde_json::json!("auto"))
} else {
None
},
reasoning_effort: None,
};
let url = format!("{}{}", base_url, self.endpoint());
let mut req_builder = client
.post(&url)
.header("Authorization", format!("Bearer {}", api_key))
.header("Content-Type", "application/json");
// 添加 X-Provider-Id header 用于精确路由
if let Some(pid) = provider_id {
req_builder = req_builder.header("X-Provider-Id", pid);
}
let response = req_builder
.json(&request)
.send()
.await
.map_err(|e| format!("请求失败: {}", e))?;
let status = response.status();
if !status.is_success() {
let body = response.text().await.unwrap_or_default();
error!("[OpenAIProtocol] 请求失败: {} - {}", status, body);
let _ = tx
.send(StreamEvent::Error {
message: format!("API 错误 ({}): {}", status, body),
})
.await;
return Err(format!("API 错误: {}", status));
}
// 继续对话时不发送 Done 事件(工具循环可能还会继续)
Self::process_stream(response, tx, false).await
}
fn endpoint(&self) -> &'static str {
"/v1/chat/completions"
}
}
File diff suppressed because it is too large Load Diff
-284
View File
@@ -1,284 +0,0 @@
# 工具系统模块
<!-- 一旦我所属的文件夹有所变化,请更新我 -->
## 架构说明
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 包含所有可用工具定义
## 更新提醒
任何文件变更后,请更新此文档和相关的上级文档。
-468
View File
@@ -1,468 +0,0 @@
//! Browser 工具模块
//!
//! 提供浏览器自动化功能,基于 Playwright
//! 专为 AI Agent 设计,提供结构化的页面快照
#![allow(dead_code)]
use super::registry::Tool;
use super::types::{JsonSchema, PropertySchema, ToolDefinition, ToolError, ToolResult};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use std::process::Stdio;
use std::time::Duration;
use tokio::process::Command;
use tokio::time::timeout;
use tracing::info;
/// 默认超时时间(秒)
const DEFAULT_TIMEOUT_SECS: u64 = 30;
/// Playwright 脚本目录
const PLAYWRIGHT_SCRIPTS_DIR: &str = "scripts/playwright";
/// 浏览器操作类型
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum BrowserAction {
/// 打开页面
Open { url: String },
/// 获取页面快照(AI 友好的可访问性树)
Snapshot {
#[serde(default)]
interactive_only: bool,
},
/// 点击元素
Click { selector: String },
/// 填充表单
Fill { selector: String, value: String },
/// 输入文本(逐字符)
Type { selector: String, text: String },
/// 按键
Press { key: String },
/// 滚动页面
Scroll {
direction: ScrollDirection,
#[serde(default = "default_scroll_amount")]
amount: i32,
},
/// 等待元素
WaitFor {
selector: String,
#[serde(default = "default_wait_timeout")]
timeout_ms: u64,
},
/// 截图
Screenshot {
#[serde(default)]
full_page: bool,
path: Option<String>,
},
/// 获取页面文本内容
GetText { selector: Option<String> },
/// 执行 JavaScript
Evaluate { script: String },
/// 关闭浏览器
Close,
}
fn default_scroll_amount() -> i32 {
500
}
fn default_wait_timeout() -> u64 {
5000
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ScrollDirection {
Up,
Down,
Left,
Right,
}
/// 浏览器操作结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BrowserResult {
/// 是否成功
pub success: bool,
/// 输出内容
pub output: String,
/// 页面 URL
pub url: Option<String>,
/// 页面标题
pub title: Option<String>,
/// 截图 base64(如果有)
pub screenshot: Option<String>,
/// 错误信息
pub error: Option<String>,
}
/// Browser 工具
///
/// 提供浏览器自动化功能,专为 AI Agent 设计
pub struct BrowserTool {
/// Playwright 脚本路径
script_path: PathBuf,
/// 超时时间(秒)
timeout_secs: u64,
/// 是否使用 headless 模式
headless: bool,
}
impl BrowserTool {
/// 创建新的 Browser 工具
pub fn new() -> Self {
// 获取脚本路径(相对于项目根目录)
let script_path = std::env::current_dir()
.unwrap_or_default()
.join(PLAYWRIGHT_SCRIPTS_DIR)
.join("browser-tool.mjs");
Self {
script_path,
timeout_secs: DEFAULT_TIMEOUT_SECS,
headless: true,
}
}
/// 设置脚本路径
pub fn with_script_path(mut self, path: PathBuf) -> Self {
self.script_path = path;
self
}
/// 设置超时时间
pub fn with_timeout(mut self, timeout_secs: u64) -> Self {
self.timeout_secs = timeout_secs;
self
}
/// 设置是否 headless
pub fn with_headless(mut self, headless: bool) -> Self {
self.headless = headless;
self
}
/// 执行浏览器操作
async fn execute_action(&self, action: &BrowserAction) -> Result<BrowserResult, ToolError> {
let action_json = serde_json::to_string(action)
.map_err(|e| ToolError::ExecutionFailed(format!("序列化操作失败: {}", e)))?;
info!("[BrowserTool] 执行操作: {:?}", action);
// 构建命令
let mut cmd = Command::new("node");
cmd.arg(&self.script_path);
cmd.arg("--action");
cmd.arg(&action_json);
if self.headless {
cmd.arg("--headless");
}
cmd.stdin(Stdio::null());
cmd.stdout(Stdio::piped());
cmd.stderr(Stdio::piped());
// 执行命令
let timeout_duration = Duration::from_secs(self.timeout_secs);
let result = timeout(timeout_duration, cmd.output()).await;
match result {
Ok(Ok(output)) => {
let stdout = String::from_utf8_lossy(&output.stdout);
let stderr = String::from_utf8_lossy(&output.stderr);
if output.status.success() {
// 解析 JSON 输出
serde_json::from_str(&stdout).map_err(|e| {
ToolError::ExecutionFailed(format!(
"解析输出失败: {}\nstdout: {}\nstderr: {}",
e, stdout, stderr
))
})
} else {
Err(ToolError::ExecutionFailed(format!(
"浏览器操作失败: {}",
stderr
)))
}
}
Ok(Err(e)) => Err(ToolError::ExecutionFailed(format!("执行命令失败: {}", e))),
Err(_) => Err(ToolError::Timeout),
}
}
}
impl Default for BrowserTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for BrowserTool {
fn definition(&self) -> ToolDefinition {
ToolDefinition::new(
"browser",
"Control a web browser for automation tasks. Use this to navigate websites, \
interact with elements, fill forms, and extract information. The 'snapshot' \
action returns an accessibility tree with element references (like @e1, @e2) \
that can be used in subsequent actions.",
)
.with_parameters(
JsonSchema::new()
.add_property(
"action",
PropertySchema::string(
"The browser action to perform. One of: open, snapshot, click, fill, \
type, press, scroll, wait_for, screenshot, get_text, evaluate, close",
),
true,
)
.add_property(
"url",
PropertySchema::string("URL to open (for 'open' action)"),
false,
)
.add_property(
"selector",
PropertySchema::string(
"Element selector. Can be CSS selector, XPath, or element reference \
like @e1 from snapshot output",
),
false,
)
.add_property(
"value",
PropertySchema::string("Value to fill (for 'fill' action)"),
false,
)
.add_property(
"text",
PropertySchema::string("Text to type (for 'type' action)"),
false,
)
.add_property(
"key",
PropertySchema::string(
"Key to press (for 'press' action), e.g., 'Enter', 'Tab', 'Escape'",
),
false,
)
.add_property(
"direction",
PropertySchema::string("Scroll direction: up, down, left, right"),
false,
)
.add_property(
"script",
PropertySchema::string("JavaScript code to evaluate (for 'evaluate' action)"),
false,
)
.add_property(
"interactive_only",
PropertySchema::boolean(
"Only include interactive elements in snapshot (default: false)",
),
false,
)
.add_property(
"full_page",
PropertySchema::boolean("Capture full page screenshot (default: false)"),
false,
),
)
}
async fn execute(&self, args: serde_json::Value) -> Result<ToolResult, ToolError> {
let action_str = args
.get("action")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidArguments("缺少 action 参数".to_string()))?;
// 解析操作
let action = match action_str {
"open" => {
let url = args.get("url").and_then(|v| v.as_str()).ok_or_else(|| {
ToolError::InvalidArguments("open 操作需要 url 参数".to_string())
})?;
BrowserAction::Open {
url: url.to_string(),
}
}
"snapshot" => {
let interactive_only = args
.get("interactive_only")
.and_then(|v| v.as_bool())
.unwrap_or(false);
BrowserAction::Snapshot { interactive_only }
}
"click" => {
let selector = args
.get("selector")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidArguments("click 操作需要 selector 参数".to_string())
})?;
BrowserAction::Click {
selector: selector.to_string(),
}
}
"fill" => {
let selector = args
.get("selector")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidArguments("fill 操作需要 selector 参数".to_string())
})?;
let value = args.get("value").and_then(|v| v.as_str()).ok_or_else(|| {
ToolError::InvalidArguments("fill 操作需要 value 参数".to_string())
})?;
BrowserAction::Fill {
selector: selector.to_string(),
value: value.to_string(),
}
}
"type" => {
let selector = args
.get("selector")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidArguments("type 操作需要 selector 参数".to_string())
})?;
let text = args.get("text").and_then(|v| v.as_str()).ok_or_else(|| {
ToolError::InvalidArguments("type 操作需要 text 参数".to_string())
})?;
BrowserAction::Type {
selector: selector.to_string(),
text: text.to_string(),
}
}
"press" => {
let key = args.get("key").and_then(|v| v.as_str()).ok_or_else(|| {
ToolError::InvalidArguments("press 操作需要 key 参数".to_string())
})?;
BrowserAction::Press {
key: key.to_string(),
}
}
"scroll" => {
let direction = args
.get("direction")
.and_then(|v| v.as_str())
.unwrap_or("down");
let direction = match direction {
"up" => ScrollDirection::Up,
"down" => ScrollDirection::Down,
"left" => ScrollDirection::Left,
"right" => ScrollDirection::Right,
_ => ScrollDirection::Down,
};
let amount = args
.get("amount")
.and_then(|v| v.as_i64())
.map(|v| v as i32)
.unwrap_or(500);
BrowserAction::Scroll { direction, amount }
}
"wait_for" => {
let selector = args
.get("selector")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidArguments("wait_for 操作需要 selector 参数".to_string())
})?;
let timeout_ms = args
.get("timeout_ms")
.and_then(|v| v.as_u64())
.unwrap_or(5000);
BrowserAction::WaitFor {
selector: selector.to_string(),
timeout_ms,
}
}
"screenshot" => {
let full_page = args
.get("full_page")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let path = args.get("path").and_then(|v| v.as_str()).map(String::from);
BrowserAction::Screenshot { full_page, path }
}
"get_text" => {
let selector = args
.get("selector")
.and_then(|v| v.as_str())
.map(String::from);
BrowserAction::GetText { selector }
}
"evaluate" => {
let script = args.get("script").and_then(|v| v.as_str()).ok_or_else(|| {
ToolError::InvalidArguments("evaluate 操作需要 script 参数".to_string())
})?;
BrowserAction::Evaluate {
script: script.to_string(),
}
}
"close" => BrowserAction::Close,
_ => {
return Err(ToolError::InvalidArguments(format!(
"未知的操作: {}",
action_str
)));
}
};
// 执行操作
let result = self.execute_action(&action).await?;
// 构建输出
let mut output = result.output;
if let Some(url) = &result.url {
output = format!("URL: {}\n{}", url, output);
}
if let Some(title) = &result.title {
output = format!("Title: {}\n{}", title, output);
}
if result.success {
Ok(ToolResult::success(output))
} else {
Ok(ToolResult::failure_with_output(
output,
result.error.unwrap_or_else(|| "未知错误".to_string()),
))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tool_definition() {
let tool = BrowserTool::new();
let def = tool.definition();
assert_eq!(def.name, "browser");
assert!(!def.description.is_empty());
assert!(def.parameters.required.contains(&"action".to_string()));
}
#[test]
fn test_action_serialization() {
let action = BrowserAction::Open {
url: "https://example.com".to_string(),
};
let json = serde_json::to_string(&action).unwrap();
assert!(json.contains("open"));
assert!(json.contains("https://example.com"));
}
}
-76
View File
@@ -1,76 +0,0 @@
//! Agent 工具系统模块
//!
//! 基于 aster-rust 框架的工具系统集成
//! 直接使用 aster-rust 提供的工具实现和注册表
//!
//! ## 架构说明
//! - 使用 aster-rust 的 ToolRegistry 和 Tool trait
//! - 直接注册 aster-rust 提供的所有工具
//! - 保持与现有 ProxyCast 接口的兼容性
// 重新导出 aster-rust 的工具系统
pub use aster::tools::*;
// 保持兼容性的类型别名和重新导出
pub use aster::tools::Tool;
pub use aster::tools::ToolContext;
pub use aster::tools::ToolDefinition;
pub use aster::tools::ToolError;
pub use aster::tools::ToolRegistry;
pub use aster::tools::ToolResult;
// 为了兼容性,创建一些类型别名
pub type JsonSchema = serde_json::Value;
pub type PropertySchema = serde_json::Value;
// 保留现有的特殊工具(暂时注释掉,需要适配 aster-rust 接口)
// pub mod browser;
// pub mod prompt;
// pub mod security;
// pub mod term_scrollback;
// pub mod terminal;
use std::path::Path;
use tracing::info;
#[cfg(test)]
mod test_integration;
/// 创建包含所有 aster-rust 工具的注册表
///
/// # Arguments
/// * `_base_dir` - 基础目录,所有文件操作必须在此目录内
///
/// # Returns
/// 包含 aster-rust 所有工具的注册表
pub fn create_default_registry(_base_dir: impl AsRef<Path>) -> ToolRegistry {
let mut registry = ToolRegistry::new();
// 注册所有 aster-rust 工具
let config = aster::tools::ToolRegistrationConfig::default();
let _shared_history = aster::tools::register_all_tools(&mut registry, config);
info!(
"[Tools] 已创建 aster-rust 工具注册表,共 {} 个工具",
registry.tool_count()
);
registry
}
/// 创建简化的工具注册表(仅核心工具)
///
/// 只包含最基本的 aster-rust 工具
pub fn create_minimal_registry(_base_dir: impl AsRef<Path>) -> ToolRegistry {
let mut registry = ToolRegistry::new();
// 使用 aster-rust 的默认工具注册
let _shared_history = aster::tools::register_default_tools(&mut registry);
info!(
"[Tools] 已创建最小工具注册表,共 {} 个工具",
registry.tool_count()
);
registry
}
-643
View File
@@ -1,643 +0,0 @@
//! 工具 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
///
/// 注意:工具定义已通过 API 的 tools 字段发送,不需要在 system prompt 中重复
/// 此方法只返回工具使用指导
pub fn generate_system_prompt(&self, _tools: &[ToolDefinition]) -> String {
// 只返回使用指导,工具定义由 API 原生处理
TOOL_USAGE_INSTRUCTIONS.to_string()
}
/// 生成包含工具定义的完整 System Prompt(旧版本,保留兼容性)
#[allow(dead_code)]
pub fn generate_full_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#"你是一个友好的 AI 助手。
# 核心原则
1. **自然交流**:对于问候、闲聊、问答,直接用文字回复,不要调用任何工具
2. **显式授权**:只有当用户**明确提供**文件路径或目录时,才能操作
3. **不要主动探索**:不要自作主张读取目录或文件来"了解环境"
# 可用工具
- **read_file**:读取用户指定的文件或目录
- **write_file**:创建/覆盖用户指定的文件
- **edit_file**:修改用户指定的文件
- **bash**:执行用户要求的命令
# 重要限制
⚠️ **禁止行为**:
- 用户说"你好"时,不要读取任何文件
- 用户没有给路径时,不要自己猜测或使用 "."
- 不要为了"打招呼"或"了解用户"而调用工具
✅ **正确做法**:
- 用户说"你好" → 直接回复问候
- 用户说"看看 /path/to/file" → 调用 read_file
- 用户说"列出目录内容" → 询问用户要查看哪个目录
# 输出格式
- 使用 Markdown 格式
- 简洁明了
- 使用中文回复"#;
/// 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_system_prompt() {
let generator = ToolPromptGenerator::new();
let tools = create_test_tools();
let prompt = generator.generate_system_prompt(&tools);
// 验证包含工具使用说明(新版本只返回指导,不包含工具定义)
assert!(prompt.contains("你是一个友好的 AI 助手"));
assert!(prompt.contains("可用工具"));
assert!(prompt.contains("read_file")); // 在说明中提到
}
#[test]
fn test_generate_full_xml_prompt() {
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml);
let tools = create_test_tools();
let prompt = generator.generate_full_system_prompt(&tools);
// 验证包含工具使用说明
assert!(prompt.contains("可用工具"));
// 验证包含 tools 标签
assert!(prompt.contains("<tools>"));
assert!(prompt.contains("</tools>"));
// 验证包含所有工具
assert!(prompt.contains("name=\"bash\""));
assert!(prompt.contains("name=\"read_file\""));
}
#[test]
fn test_generate_full_json_prompt() {
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Json);
let tools = create_test_tools();
let prompt = generator.generate_full_system_prompt(&tools);
// 验证包含工具使用说明
assert!(prompt.contains("可用工具"));
// 验证包含 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();
// generate_tools_prompt 使用 generate_system_prompt,只返回指导
let prompt = generate_tools_prompt(&tools, PromptFormat::Xml);
assert!(prompt.contains("可用工具"));
}
#[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("你是一个友好的 AI 助手"));
assert!(prompt.contains("可用工具"));
}
#[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_full_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();
// 使用 generate_full_system_prompt 来包含工具定义
let prompt = generator.generate_full_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: Full System Prompt 工具包含**
/// **Validates: Requirements 2.3**
///
/// *For any* 已注册的工具集合,生成的完整 System Prompt 应该包含所有工具的 name 和 description。
#[test]
fn prop_full_system_prompt_contains_all_tool_names(tools in arb_tool_definitions()) {
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml);
let prompt = generator.generate_full_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: Full System Prompt 工具包含 - 描述**
/// **Validates: Requirements 2.3**
///
/// *For any* 已注册的工具集合,生成的完整 System Prompt 应该包含所有工具的 description。
#[test]
fn prop_full_system_prompt_contains_all_tool_descriptions(tools in arb_tool_definitions()) {
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml);
let prompt = generator.generate_full_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: Full System Prompt 工具包含 - JSON 格式**
/// **Validates: Requirements 2.3**
///
/// *For any* 已注册的工具集合,JSON 格式的完整 System Prompt 也应该包含所有工具的 name 和 description。
#[test]
fn prop_full_system_prompt_json_contains_all_tools(tools in arb_tool_definitions()) {
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Json);
let prompt = generator.generate_full_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: Full System Prompt 工具包含 - 格式一致性**
/// **Validates: Requirements 2.3**
///
/// *For any* 已注册的工具集合,无论使用 XML 还是 JSON 格式,
/// 生成的完整 System Prompt 都应该包含相同的工具信息。
#[test]
fn prop_full_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_full_system_prompt(&tools);
let json_prompt = json_generator.generate_full_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: Full System Prompt 工具包含 - 工具数量**
/// **Validates: Requirements 2.3**
///
/// *For any* 已注册的工具集合,生成的完整 System Prompt 中工具名称出现的次数
/// 应该至少等于工具数量(每个工具至少出现一次)。
#[test]
fn prop_full_system_prompt_tool_count(tools in arb_tool_definitions()) {
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml);
let prompt = generator.generate_full_system_prompt(&tools);
// 统计每个工具名称在 prompt 中出现的次数
for tool in &tools {
let count = prompt.matches(&tool.name).count();
prop_assert!(
count >= 1,
"工具 '{}' 应该在 System Prompt 中至少出现一次,实际出现 {} 次",
tool.name,
count
);
}
}
}
}
-594
View File
@@ -1,594 +0,0 @@
//! 安全管理器模块
//!
//! 提供路径验证、目录遍历防护、符号链接检查等安全功能
//! 符合 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
);
}
}
}
}
@@ -1,346 +0,0 @@
//! Terminal Scrollback 工具模块
//!
//! 提供只读访问终端输出历史的功能
//! 参考 Waveterm 的 term_get_scrollback 工具设计
use super::registry::Tool;
use super::types::{JsonSchema, PropertySchema, ToolDefinition, ToolError, ToolResult};
use async_trait::async_trait;
use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::oneshot;
use tokio::time::timeout;
use tracing::{debug, info, warn};
/// 默认超时时间(秒)
const DEFAULT_TIMEOUT_SECS: u64 = 30;
/// 获取滚动缓冲区的请求
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GetScrollbackRequest {
/// 请求 ID
pub request_id: String,
/// 终端会话 ID
pub session_id: String,
/// 起始行(可选,默认从最后 200 行开始)
pub line_start: Option<usize>,
/// 行数(可选,默认 200 行)
pub count: Option<usize>,
}
/// 获取滚动缓冲区的响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GetScrollbackResponse {
/// 请求 ID
pub request_id: String,
/// 是否成功
pub success: bool,
/// 总行数
pub total_lines: usize,
/// 起始行
pub line_start: usize,
/// 结束行
pub line_end: usize,
/// 内容
pub content: String,
/// 是否有更多数据
pub has_more: bool,
/// 错误信息
pub error: Option<String>,
}
/// 待处理的请求(内部使用)
pub(crate) struct PendingRequest {
/// 响应发送器
response_tx: oneshot::Sender<GetScrollbackResponse>,
}
/// Terminal Scrollback 工具
///
/// 允许 AI 读取终端的输出历史,而不是直接执行命令
pub struct TermScrollbackTool {
/// 待处理的请求
pending_requests: Arc<RwLock<HashMap<String, PendingRequest>>>,
/// 默认超时时间(秒)
timeout_secs: u64,
/// Tauri AppHandle(用于发送事件)
app_handle: Arc<RwLock<Option<tauri::AppHandle>>>,
}
impl TermScrollbackTool {
/// 创建新的 TermScrollbackTool
pub fn new() -> Self {
Self {
pending_requests: Arc::new(RwLock::new(HashMap::new())),
timeout_secs: DEFAULT_TIMEOUT_SECS,
app_handle: Arc::new(RwLock::new(None)),
}
}
/// 设置超时时间
pub fn with_timeout(mut self, timeout_secs: u64) -> Self {
self.timeout_secs = timeout_secs;
self
}
/// 设置 Tauri AppHandle
pub fn set_app_handle(&self, handle: tauri::AppHandle) {
let mut app_handle = self.app_handle.write();
*app_handle = Some(handle);
eprintln!("[TermScrollbackTool] AppHandle 已设置");
tracing::info!("[TermScrollbackTool] AppHandle 已设置");
}
/// 检查 AppHandle 是否已设置
pub fn is_app_handle_set(&self) -> bool {
let app_handle = self.app_handle.read();
app_handle.is_some()
}
/// 处理响应(由前端调用)
pub fn handle_response(&self, response: GetScrollbackResponse) {
let request_id = response.request_id.clone();
let pending = {
let mut requests = self.pending_requests.write();
requests.remove(&request_id)
};
if let Some(pending) = pending {
if pending.response_tx.send(response).is_err() {
warn!(
"[TermScrollbackTool] 发送响应失败,接收端已关闭: {}",
request_id
);
}
} else {
warn!("[TermScrollbackTool] 未找到待处理的请求: {}", request_id);
}
}
/// 获取终端滚动缓冲区
async fn get_scrollback(
&self,
session_id: &str,
line_start: Option<usize>,
count: Option<usize>,
) -> Result<GetScrollbackResponse, ToolError> {
let request_id = uuid::Uuid::new_v4().to_string();
info!(
"[TermScrollbackTool] 请求获取滚动缓冲区: session_id={}, request_id={}",
session_id, request_id
);
// 创建响应通道
let (response_tx, response_rx) = oneshot::channel();
// 添加到待处理列表
{
let mut requests = self.pending_requests.write();
requests.insert(request_id.clone(), PendingRequest { response_tx });
}
// 构建请求
let request = GetScrollbackRequest {
request_id: request_id.clone(),
session_id: session_id.to_string(),
line_start,
count,
};
// 发送事件到前端
{
let app_handle = self.app_handle.read();
eprintln!(
"[TermScrollbackTool] 检查 AppHandle: is_some={}",
app_handle.is_some()
);
if let Some(handle) = app_handle.as_ref() {
use tauri::Emitter;
eprintln!("[TermScrollbackTool] 尝试发送事件到前端: {}", request_id);
if let Err(e) = handle.emit("term_get_scrollback_request", &request) {
// 清理待处理请求
let mut requests = self.pending_requests.write();
requests.remove(&request_id);
warn!("[TermScrollbackTool] 发送事件到前端失败: {}", e);
return Err(ToolError::ExecutionFailed(format!(
"无法发送请求到前端:{}",
e
)));
}
debug!("[TermScrollbackTool] 已发送请求到前端: {}", request_id);
} else {
// 清理待处理请求
let mut requests = self.pending_requests.write();
requests.remove(&request_id);
eprintln!("[TermScrollbackTool] AppHandle 为 None,无法发送事件");
warn!("[TermScrollbackTool] AppHandle 未设置");
return Err(ToolError::ExecutionFailed(
"TermScrollbackTool 未正确初始化".to_string(),
));
}
}
// 等待响应(带超时)
let timeout_duration = Duration::from_secs(self.timeout_secs);
match timeout(timeout_duration, response_rx).await {
Ok(Ok(response)) => {
debug!(
"[TermScrollbackTool] 收到响应: request_id={}, success={}",
response.request_id, response.success
);
Ok(response)
}
Ok(Err(_)) => {
// 通道关闭
let mut requests = self.pending_requests.write();
requests.remove(&request_id);
Err(ToolError::ExecutionFailed("响应通道已关闭".to_string()))
}
Err(_) => {
// 超时
warn!("[TermScrollbackTool] 请求超时: {}", request_id);
let mut requests = self.pending_requests.write();
requests.remove(&request_id);
Err(ToolError::Timeout)
}
}
}
}
impl Default for TermScrollbackTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for TermScrollbackTool {
fn definition(&self) -> ToolDefinition {
ToolDefinition::new(
"term_get_scrollback",
"Get the terminal output history (scrollback buffer). This is a READ-ONLY tool \
that allows you to view what has been output in the terminal.\n\n\
IMPORTANT: This tool does NOT execute commands. It only reads the terminal output. \
Use this to:\n\
- Check the results of commands that the user has run\n\
- View error messages and logs\n\
- Understand the current state of the terminal\n\n\
The user must manually execute commands in their terminal. You can suggest commands \
for the user to run, but you cannot execute them directly with this tool.",
)
.with_parameters(
JsonSchema::new()
.add_property(
"session_id",
PropertySchema::string(
"The terminal session ID to read from. This is provided by the system.",
),
true,
)
.add_property(
"line_start",
PropertySchema::integer(
"Optional starting line number. If not specified, returns the last 200 lines.",
),
false,
)
.add_property(
"count",
PropertySchema::integer(
"Optional number of lines to return. Defaults to 200 lines.",
)
.with_default(serde_json::json!(200)),
false,
),
)
}
async fn execute(&self, args: serde_json::Value) -> Result<ToolResult, ToolError> {
// 解析参数
let session_id = args
.get("session_id")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidArguments("缺少 session_id 参数".to_string()))?;
let line_start = args
.get("line_start")
.and_then(|v| v.as_u64())
.map(|v| v as usize);
let count = args
.get("count")
.and_then(|v| v.as_u64())
.map(|v| v as usize);
// 获取滚动缓冲区
let response = self.get_scrollback(session_id, line_start, count).await?;
// 构建输出
if response.success {
let mut output = format!(
"Terminal output (lines {}-{} of {}):\n\n",
response.line_start, response.line_end, response.total_lines
);
output.push_str(&response.content);
if response.has_more {
output.push_str(
"\n\n[More output available. Use line_start parameter to fetch earlier lines.]",
);
}
Ok(ToolResult::success(output))
} else {
let error_msg = response
.error
.unwrap_or_else(|| "Unknown error".to_string());
Err(ToolError::ExecutionFailed(error_msg))
}
}
}
/// 全局 TermScrollbackTool 实例
static TERM_SCROLLBACK_TOOL: once_cell::sync::Lazy<Arc<TermScrollbackTool>> =
once_cell::sync::Lazy::new(|| {
eprintln!("[TermScrollbackTool] 创建全局实例");
Arc::new(TermScrollbackTool::new())
});
/// 获取全局 TermScrollbackTool 实例
pub fn get_term_scrollback_tool() -> Arc<TermScrollbackTool> {
eprintln!("[TermScrollbackTool] 获取全局实例");
Arc::clone(&TERM_SCROLLBACK_TOOL)
}
/// 设置全局 TermScrollbackTool 的 AppHandle
pub fn set_term_scrollback_tool_app_handle(handle: tauri::AppHandle) {
eprintln!("[TermScrollbackTool] 设置全局 AppHandle");
tracing::info!("[TermScrollbackTool] 设置全局 AppHandle");
// 强制初始化全局实例(如果还未初始化)
let _ = &*TERM_SCROLLBACK_TOOL;
TERM_SCROLLBACK_TOOL.set_app_handle(handle);
// 验证设置是否成功
let app_handle = TERM_SCROLLBACK_TOOL.app_handle.read();
if app_handle.is_some() {
eprintln!("[TermScrollbackTool] AppHandle 设置成功,已验证");
tracing::info!("[TermScrollbackTool] AppHandle 设置成功,已验证");
} else {
eprintln!("[TermScrollbackTool] 警告:AppHandle 设置后仍为 None");
tracing::error!("[TermScrollbackTool] 警告:AppHandle 设置后仍为 None");
}
}
/// 处理滚动缓冲区响应(由 Tauri 命令调用)
pub fn handle_term_scrollback_response(response: GetScrollbackResponse) {
TERM_SCROLLBACK_TOOL.handle_response(response);
}
-499
View File
@@ -1,499 +0,0 @@
//! Terminal 工具模块
//!
//! 提供终端命令执行功能,通过 Tauri 事件与前端通信
//! 支持命令审批流程,命令在用户的实际终端中执行
//!
//! ## 工作流程
//! 1. AI 调用 terminal 工具
//! 2. 工具发送事件到前端请求执行命令
//! 3. 前端显示审批 UI
//! 4. 用户批准后,命令发送到实际终端
//! 5. 终端执行结果返回给 AI
use super::registry::Tool;
use super::types::{JsonSchema, PropertySchema, ToolDefinition, ToolError, ToolResult};
use async_trait::async_trait;
use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::oneshot;
use tokio::time::timeout;
use tracing::{debug, info, warn};
/// 默认超时时间(秒)
const DEFAULT_TIMEOUT_SECS: u64 = 120;
/// 命令执行请求
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TerminalCommandRequest {
/// 请求 ID
pub request_id: String,
/// 要执行的命令
pub command: String,
/// 工作目录(可选)
pub working_dir: Option<String>,
/// 超时时间(秒)
pub timeout_secs: u64,
}
/// 命令执行响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TerminalCommandResponse {
/// 请求 ID
pub request_id: String,
/// 是否成功
pub success: bool,
/// 输出内容
pub output: String,
/// 错误信息
pub error: Option<String>,
/// 退出码
pub exit_code: Option<i32>,
/// 是否被用户拒绝
pub rejected: bool,
}
/// 命令执行状态
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CommandStatus {
/// 等待审批
Pending,
/// 已批准,执行中
Executing,
/// 已完成
Completed,
/// 已拒绝
Rejected,
/// 超时
Timeout,
}
/// 待处理的命令(内部使用)
pub(crate) struct PendingCommand {
/// 响应发送器
response_tx: oneshot::Sender<TerminalCommandResponse>,
}
/// 已执行命令的记录
#[derive(Debug, Clone)]
struct ExecutedCommand {
/// 命令内容
command: String,
/// 执行时间
executed_at: std::time::Instant,
/// 是否成功
success: bool,
}
/// 重复命令检测的时间窗口(秒)
const DUPLICATE_DETECTION_WINDOW_SECS: u64 = 30;
/// Terminal 工具
///
/// 通过 Tauri 事件与前端通信,在用户终端中执行命令
pub struct TerminalTool {
/// 待处理的命令
pending_commands: Arc<RwLock<HashMap<String, PendingCommand>>>,
/// 已执行的命令历史(用于检测重复)
executed_commands: Arc<RwLock<Vec<ExecutedCommand>>>,
/// 默认超时时间(秒)
timeout_secs: u64,
/// Tauri AppHandle(用于发送事件)
app_handle: Arc<RwLock<Option<tauri::AppHandle>>>,
}
impl TerminalTool {
/// 创建新的 Terminal 工具
pub fn new() -> Self {
Self {
pending_commands: Arc::new(RwLock::new(HashMap::new())),
executed_commands: Arc::new(RwLock::new(Vec::new())),
timeout_secs: DEFAULT_TIMEOUT_SECS,
app_handle: Arc::new(RwLock::new(None)),
}
}
/// 设置超时时间
pub fn with_timeout(mut self, timeout_secs: u64) -> Self {
self.timeout_secs = timeout_secs;
self
}
/// 设置 Tauri AppHandle
pub fn set_app_handle(&self, handle: tauri::AppHandle) {
let mut app_handle = self.app_handle.write();
*app_handle = Some(handle);
eprintln!("[TerminalTool] AppHandle 已设置");
tracing::info!("[TerminalTool] AppHandle 已设置");
}
/// 检查 AppHandle 是否已设置
pub fn is_app_handle_set(&self) -> bool {
let app_handle = self.app_handle.read();
app_handle.is_some()
}
/// 处理命令响应(由前端调用)
pub fn handle_response(&self, response: TerminalCommandResponse) {
let request_id = response.request_id.clone();
let pending = {
let mut commands = self.pending_commands.write();
commands.remove(&request_id)
};
if let Some(pending) = pending {
if pending.response_tx.send(response).is_err() {
warn!("[TerminalTool] 发送响应失败,接收端已关闭: {}", request_id);
}
} else {
warn!("[TerminalTool] 未找到待处理的命令: {}", request_id);
}
}
/// 检查命令是否在最近执行过(用于防止重复执行)
fn check_duplicate_command(&self, command: &str) -> Option<&'static str> {
let now = std::time::Instant::now();
let window = Duration::from_secs(DUPLICATE_DETECTION_WINDOW_SECS);
// 清理过期的记录
{
let mut history = self.executed_commands.write();
history.retain(|cmd| now.duration_since(cmd.executed_at) < window);
}
// 检查是否有重复
let history = self.executed_commands.read();
for cmd in history.iter() {
if cmd.command == command && cmd.success {
return Some(
"This exact command was already executed successfully within the last 30 seconds. \
Do NOT re-execute it. If you need to verify the result, use the term_get_scrollback \
tool to read the terminal output."
);
}
}
None
}
/// 记录已执行的命令
fn record_executed_command(&self, command: &str, success: bool) {
let mut history = self.executed_commands.write();
history.push(ExecutedCommand {
command: command.to_string(),
executed_at: std::time::Instant::now(),
success,
});
// 限制历史记录大小
if history.len() > 100 {
history.remove(0);
}
}
/// 执行命令
async fn execute_command(
&self,
command: &str,
working_dir: Option<&str>,
timeout_secs: u64,
) -> Result<TerminalCommandResponse, ToolError> {
let request_id = uuid::Uuid::new_v4().to_string();
info!(
"[TerminalTool] 请求执行命令: {} (request_id: {}, timeout: {}s)",
command, request_id, timeout_secs
);
// 创建响应通道
let (response_tx, response_rx) = oneshot::channel();
// 添加到待处理列表
{
let mut commands = self.pending_commands.write();
commands.insert(request_id.clone(), PendingCommand { response_tx });
}
// 构建请求
let request = TerminalCommandRequest {
request_id: request_id.clone(),
command: command.to_string(),
working_dir: working_dir.map(|s| s.to_string()),
timeout_secs,
};
// 发送事件到前端
{
let app_handle = self.app_handle.read();
eprintln!(
"[TerminalTool] 检查 AppHandle: is_some={}",
app_handle.is_some()
);
if let Some(handle) = app_handle.as_ref() {
use tauri::Emitter;
eprintln!("[TerminalTool] 尝试发送事件到前端: {}", request_id);
if let Err(e) = handle.emit("terminal_command_request", &request) {
// 清理待处理命令
let mut commands = self.pending_commands.write();
commands.remove(&request_id);
// 返回失败响应而不是错误,避免 Agent 重试
warn!("[TerminalTool] 发送事件到前端失败: {}", e);
return Ok(TerminalCommandResponse {
request_id: request_id.clone(),
success: false,
output: String::new(),
error: Some(format!("无法发送命令到终端:{}。请检查应用配置。", e)),
exit_code: Some(-1),
rejected: false,
});
}
debug!("[TerminalTool] 已发送命令请求到前端: {}", request_id);
} else {
// 清理待处理命令
let mut commands = self.pending_commands.write();
commands.remove(&request_id);
eprintln!("[TerminalTool] AppHandle 为 None,无法发送事件");
// 返回失败响应而不是错误,避免 Agent 重试
warn!("[TerminalTool] AppHandle 未设置");
return Ok(TerminalCommandResponse {
request_id: request_id.clone(),
success: false,
output: String::new(),
error: Some(
"终端工具未正确初始化。这是一个应用程序配置问题,请联系开发者。\n\
作为替代方案,我可以为您提供命令建议,但无法直接执行。"
.to_string(),
),
exit_code: Some(-1),
rejected: false,
});
}
}
// 等待响应(带超时)
let timeout_duration = Duration::from_secs(timeout_secs);
match timeout(timeout_duration, response_rx).await {
Ok(Ok(response)) => {
debug!(
"[TerminalTool] 收到响应: request_id={}, success={}",
response.request_id, response.success
);
Ok(response)
}
Ok(Err(_)) => {
// 通道关闭
let mut commands = self.pending_commands.write();
commands.remove(&request_id);
Err(ToolError::ExecutionFailed("响应通道已关闭".to_string()))
}
Err(_) => {
// 超时
warn!("[TerminalTool] 命令执行超时: {}", request_id);
let mut commands = self.pending_commands.write();
commands.remove(&request_id);
Err(ToolError::Timeout)
}
}
}
}
impl Default for TerminalTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for TerminalTool {
fn definition(&self) -> ToolDefinition {
ToolDefinition::new(
"terminal",
"Execute a command in the user's terminal. The command will be sent to the \
active terminal session and requires user approval before execution.\n\n\
CRITICAL RULES:\n\
1. NEVER re-execute a command that has already succeeded. When you receive \
'[COMMAND EXECUTED SUCCESSFULLY]' in the response, the command is DONE.\n\
2. Each command runs exactly ONCE. Do not retry successful commands.\n\
3. If you need to verify the result, use the term_get_scrollback tool to \
read the terminal output instead of re-running the command.\n\n\
Use this for running system commands, scripts, or any command-line operations \
that should be visible to the user.",
)
.with_parameters(
JsonSchema::new()
.add_property(
"command",
PropertySchema::string(
"The command to execute in the terminal. Can be any valid shell command.",
),
true,
)
.add_property(
"working_dir",
PropertySchema::string(
"Optional working directory for the command. If not specified, \
uses the terminal's current directory.",
),
false,
)
.add_property(
"timeout",
PropertySchema::integer(
"Optional timeout in seconds. Defaults to 120 seconds.",
)
.with_default(serde_json::json!(120)),
false,
),
)
}
async fn execute(&self, args: serde_json::Value) -> Result<ToolResult, ToolError> {
// 解析参数
let command = args
.get("command")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidArguments("缺少 command 参数".to_string()))?;
let working_dir = args.get("working_dir").and_then(|v| v.as_str());
let timeout_secs = args
.get("timeout")
.and_then(|v| v.as_u64())
.unwrap_or(self.timeout_secs);
// 检查是否是重复命令
if let Some(duplicate_msg) = self.check_duplicate_command(command) {
warn!("[TerminalTool] 检测到重复命令: {}", command);
return Ok(ToolResult::success(format!(
"[DUPLICATE COMMAND BLOCKED]\n{}\n\nOriginal command: {}",
duplicate_msg, command
)));
}
// 执行命令
let response = self
.execute_command(command, working_dir, timeout_secs)
.await?;
// 记录已执行的命令
self.record_executed_command(command, response.success);
// 构建输出
if response.rejected {
return Ok(ToolResult::failure_with_output(
"用户拒绝执行此命令".to_string(),
"命令被用户拒绝".to_string(),
));
}
if response.success {
Ok(ToolResult::success(response.output))
} else {
let error_msg = response.error.unwrap_or_else(|| {
format!(
"命令执行失败 (退出码: {})",
response.exit_code.unwrap_or(-1)
)
});
Ok(ToolResult::failure_with_output(response.output, error_msg))
}
}
}
/// 全局 TerminalTool 实例
static TERMINAL_TOOL: once_cell::sync::Lazy<Arc<TerminalTool>> = once_cell::sync::Lazy::new(|| {
eprintln!("[TerminalTool] 创建全局实例");
Arc::new(TerminalTool::new())
});
/// 获取全局 TerminalTool 实例
pub fn get_terminal_tool() -> Arc<TerminalTool> {
eprintln!("[TerminalTool] 获取全局实例");
Arc::clone(&TERMINAL_TOOL)
}
/// 设置全局 TerminalTool 的 AppHandle
pub fn set_terminal_tool_app_handle(handle: tauri::AppHandle) {
eprintln!("[TerminalTool] 设置全局 AppHandle");
tracing::info!("[TerminalTool] 设置全局 AppHandle");
// 强制初始化全局实例(如果还未初始化)
let _ = &*TERMINAL_TOOL;
TERMINAL_TOOL.set_app_handle(handle);
// 验证设置是否成功
let app_handle = TERMINAL_TOOL.app_handle.read();
if app_handle.is_some() {
eprintln!("[TerminalTool] AppHandle 设置成功,已验证");
tracing::info!("[TerminalTool] AppHandle 设置成功,已验证");
} else {
eprintln!("[TerminalTool] 警告:AppHandle 设置后仍为 None");
tracing::error!("[TerminalTool] 警告:AppHandle 设置后仍为 None");
}
}
/// 处理终端命令响应(由 Tauri 命令调用)
pub fn handle_terminal_command_response(response: TerminalCommandResponse) {
TERMINAL_TOOL.handle_response(response);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tool_definition() {
let tool = TerminalTool::new();
let def = tool.definition();
assert_eq!(def.name, "terminal");
assert!(!def.description.is_empty());
assert!(def.parameters.required.contains(&"command".to_string()));
}
#[test]
fn test_terminal_command_request_serialization() {
let request = TerminalCommandRequest {
request_id: "test-123".to_string(),
command: "echo hello".to_string(),
working_dir: Some("/home/user".to_string()),
timeout_secs: 60,
};
let json = serde_json::to_string(&request).unwrap();
let parsed: TerminalCommandRequest = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.request_id, "test-123");
assert_eq!(parsed.command, "echo hello");
assert_eq!(parsed.working_dir, Some("/home/user".to_string()));
assert_eq!(parsed.timeout_secs, 60);
}
#[test]
fn test_terminal_command_response_serialization() {
let response = TerminalCommandResponse {
request_id: "test-123".to_string(),
success: true,
output: "hello\n".to_string(),
error: None,
exit_code: Some(0),
rejected: false,
};
let json = serde_json::to_string(&response).unwrap();
let parsed: TerminalCommandResponse = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.request_id, "test-123");
assert!(parsed.success);
assert_eq!(parsed.output, "hello\n");
assert_eq!(parsed.exit_code, Some(0));
assert!(!parsed.rejected);
}
}
@@ -1,59 +0,0 @@
//! aster-rust 工具系统集成测试
//!
//! 验证 ProxyCast 与 aster-rust 工具系统的集成是否正常工作
use super::*;
use std::path::PathBuf;
#[tokio::test]
async fn test_create_default_registry() {
let base_dir = PathBuf::from("/tmp");
let registry = create_default_registry(&base_dir);
// 验证注册表已创建
let definitions = registry.get_definitions();
assert!(
!definitions.is_empty(),
"工具注册表应该包含 aster-rust 工具"
);
tracing::info!("✅ 工具注册表创建成功,包含 {} 个工具", definitions.len());
}
#[tokio::test]
async fn test_create_minimal_registry() {
let base_dir = PathBuf::from("/tmp");
let registry = create_minimal_registry(&base_dir);
// 验证注册表已创建
let definitions = registry.get_definitions();
assert!(!definitions.is_empty(), "最小工具注册表应该包含基础工具");
tracing::info!(
"✅ 最小工具注册表创建成功,包含 {} 个工具",
definitions.len()
);
}
#[tokio::test]
async fn test_tool_registry_basic_operations() {
let base_dir = PathBuf::from("/tmp");
let registry = create_default_registry(&base_dir);
// 测试基本操作
let definitions = registry.get_definitions();
let tool_count = definitions.len();
assert!(tool_count > 0, "应该有工具被注册");
// 测试是否为空
assert!(!definitions.is_empty(), "注册表不应该为空");
tracing::info!("✅ 工具注册表基本操作测试通过,共 {} 个工具", tool_count);
// 验证工具定义的基本结构
for def in definitions.iter().take(3) {
assert!(!def.name.is_empty(), "工具名称不应该为空");
assert!(!def.description.is_empty(), "工具描述不应该为空");
tracing::debug!("工具: {} - {}", def.name, def.description);
}
}
-4
View File
@@ -6,7 +6,6 @@ use std::sync::Arc;
use tokio::sync::RwLock;
use crate::agent::AsterAgentState;
use crate::agent::NativeAgentState;
use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState;
use crate::commands::connect_cmd::ConnectStateWrapper;
use crate::commands::context_memory::ContextMemoryServiceState;
@@ -138,7 +137,6 @@ pub struct AppStates {
pub bookmark_manager: BookmarkManagerState,
pub enhanced_stats_service: EnhancedStatsServiceState,
pub batch_operations: BatchOperationsState,
pub native_agent: NativeAgentState,
pub aster_agent: AsterAgentState,
pub orchestrator: OrchestratorState,
pub connect_state: ConnectStateWrapper,
@@ -219,7 +217,6 @@ pub fn init_states(config: &Config) -> Result<AppStates, String> {
) = init_flow_monitor(&provider_pool_service_state, &db, &plugin_installer_state)?;
// 其他状态
let native_agent_state = NativeAgentState::new();
let aster_agent_state = AsterAgentState::new();
let orchestrator_state = OrchestratorState::new();
@@ -292,7 +289,6 @@ pub fn init_states(config: &Config) -> Result<AppStates, String> {
bookmark_manager: bookmark_manager_state,
enhanced_stats_service: enhanced_stats_service_state,
batch_operations: batch_operations_state,
native_agent: native_agent_state,
aster_agent: aster_agent_state,
orchestrator: orchestrator_state,
connect_state,
+2 -13
View File
@@ -70,7 +70,6 @@ pub fn run() {
bookmark_manager: bookmark_manager_state,
enhanced_stats_service: enhanced_stats_service_state,
batch_operations: batch_operations_state,
native_agent: native_agent_state,
aster_agent: aster_agent_state,
orchestrator: orchestrator_state,
connect_state,
@@ -154,7 +153,6 @@ pub fn run() {
.manage(bookmark_manager_state)
.manage(enhanced_stats_service_state)
.manage(batch_operations_state)
.manage(native_agent_state)
.manage(aster_agent_state)
.manage(orchestrator_state)
.manage(connect_state)
@@ -814,6 +812,7 @@ pub fn run() {
commands::api_key_provider_cmd::delete_legacy_api_key_credential,
// API Key Provider connection test command
commands::api_key_provider_cmd::test_api_key_provider_connection,
commands::api_key_provider_cmd::test_api_key_provider_chat,
// Route commands
commands::route_cmd::get_available_routes,
commands::route_cmd::get_route_curl_examples,
@@ -1052,21 +1051,11 @@ pub fn run() {
// TODO: 重新启用这些命令,适配 aster-rust 工具系统
// commands::agent_cmd::agent_terminal_command_response,
// commands::agent_cmd::agent_term_scrollback_response,
// Native Agent commands
commands::native_agent_cmd::native_agent_init,
commands::native_agent_cmd::native_agent_status,
commands::native_agent_cmd::native_agent_reset,
commands::native_agent_cmd::native_agent_chat,
commands::native_agent_cmd::native_agent_chat_stream,
commands::native_agent_cmd::native_agent_create_session,
commands::native_agent_cmd::native_agent_get_session,
commands::native_agent_cmd::native_agent_delete_session,
commands::native_agent_cmd::native_agent_list_sessions,
commands::native_agent_cmd::agent_permission_response,
// Aster Agent commands
commands::aster_agent_cmd::aster_agent_init,
commands::aster_agent_cmd::aster_agent_status,
commands::aster_agent_cmd::aster_agent_configure_provider,
commands::aster_agent_cmd::aster_agent_configure_from_pool,
commands::aster_agent_cmd::aster_agent_chat_stream,
commands::aster_agent_cmd::aster_agent_stop,
commands::aster_agent_cmd::aster_session_create,
+4 -4
View File
@@ -6,7 +6,7 @@ use std::sync::Arc;
use tauri::{App, Manager};
// use crate::agent::tools::{set_term_scrollback_tool_app_handle, set_terminal_tool_app_handle};
use crate::agent::NativeAgentState;
use crate::agent::AsterAgentState;
use crate::database;
use crate::flow_monitor::FlowInterceptor;
use crate::services::provider_pool_service::ProviderPoolService;
@@ -48,9 +48,9 @@ pub fn setup_app(
}
}
// 初始化 NativeAgentState
let native_agent_state = NativeAgentState::new();
app.manage(native_agent_state);
// 初始化 AsterAgentState
let aster_agent_state = AsterAgentState::new();
app.manage(aster_agent_state);
// TODO: 重新实现 TerminalTool 和 TermScrollbackTool 的 AppHandle 设置
// 当前暂时注释掉,等待适配 aster-rust 工具系统
+54 -238
View File
@@ -1,16 +1,9 @@
//! Agent 命令模块
//!
//! 提供原生 Agent 的 Tauri 命令(兼容旧 API)
//! 提供 Agent 的 Tauri 命令(兼容旧 API)
//! 内部使用 Aster Agent 实现
// TODO: 重新实现工具响应处理,适配 aster-rust 工具系统
// use crate::agent::tools::{
// handle_term_scrollback_response, handle_terminal_command_response, GetScrollbackResponse,
// TerminalCommandResponse,
// };
use crate::agent::{
AgentMessage, AgentSession, ImageData, NativeAgentState, NativeChatRequest, ProviderType,
};
use crate::commands::network_cmd::get_local_url;
use crate::agent::{AgentMessage, AgentSession, AsterAgentState};
use crate::database::dao::agent::AgentDao;
use crate::database::DbConnection;
use crate::AppState;
@@ -35,23 +28,21 @@ pub struct CreateSessionResponse {
pub model: Option<String>,
}
/// 启动 Agent(原生实现,无需外部进程)
/// 启动 Agent(使用 Aster 实现)
#[tauri::command]
pub async fn agent_start_process(
agent_state: State<'_, NativeAgentState>,
agent_state: State<'_, AsterAgentState>,
app_state: State<'_, AppState>,
_port: Option<u16>,
) -> Result<AgentProcessStatus, String> {
tracing::info!("[Agent] 初始化原生 Agent");
tracing::info!("[Agent] 初始化 Aster Agent");
let (host, port, api_key, running, default_provider) = {
let (host, port, running) = {
let state = app_state.read().await;
(
state.config.server.host.clone(),
state.config.server.port,
state.running_api_key.clone(),
state.running,
state.config.routing.default_provider.clone(),
)
};
@@ -59,16 +50,9 @@ pub async fn agent_start_process(
return Err("ProxyCast API Server 未运行,请先启动服务器".to_string());
}
let api_key = api_key.ok_or_else(|| "ProxyCast API Server 未配置 API Key".to_string())?;
let base_url = get_local_url(&host, port);
let provider_type = ProviderType::from_str(&default_provider);
agent_state.init_agent().await?;
agent_state.init(
base_url.clone(),
api_key,
provider_type,
Some(default_provider),
)?;
let base_url = format!("http://{}:{}", host, port);
Ok(AgentProcessStatus {
running: true,
@@ -79,23 +63,26 @@ pub async fn agent_start_process(
/// 停止 Agent
#[tauri::command]
pub async fn agent_stop_process(agent_state: State<'_, NativeAgentState>) -> Result<(), String> {
tracing::info!("[Agent] 停止原生 Agent");
agent_state.reset();
pub async fn agent_stop_process(_agent_state: State<'_, AsterAgentState>) -> Result<(), String> {
tracing::info!("[Agent] 停止 Aster Agent(无操作,Agent 保持活跃)");
// Aster Agent 不需要显式停止
Ok(())
}
/// 获取 Agent 状态
#[tauri::command]
pub async fn agent_get_process_status(
agent_state: State<'_, NativeAgentState>,
agent_state: State<'_, AsterAgentState>,
app_state: State<'_, AppState>,
) -> Result<AgentProcessStatus, String> {
let initialized = agent_state.is_initialized();
let initialized = agent_state.is_initialized().await;
if initialized {
let state = app_state.read().await;
let base_url = get_local_url(&state.config.server.host, state.config.server.port);
let base_url = format!(
"http://{}:{}",
state.config.server.host, state.config.server.port
);
Ok(AgentProcessStatus {
running: true,
base_url: Some(base_url),
@@ -121,8 +108,7 @@ pub struct SkillInfo {
/// 创建 Agent 会话
#[tauri::command]
pub async fn agent_create_session(
agent_state: State<'_, NativeAgentState>,
app_state: State<'_, AppState>,
agent_state: State<'_, AsterAgentState>,
db: State<'_, DbConnection>,
provider_type: String,
model: Option<String>,
@@ -136,39 +122,28 @@ pub async fn agent_create_session(
skills.as_ref().map(|s| s.len())
);
// 如果未初始化,自动初始化
if !agent_state.is_initialized() {
let (host, port, api_key, running, default_provider) = {
let state = app_state.read().await;
(
state.config.server.host.clone(),
state.config.server.port,
state.running_api_key.clone(),
state.running,
state.config.routing.default_provider.clone(),
)
};
// 初始化 Agent
agent_state.init_agent().await?;
if !running {
return Err("ProxyCast API Server 未运行".to_string());
}
// 生成会话 ID
let session_id = uuid::Uuid::new_v4().to_string();
let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?;
let base_url = get_local_url(&host, port);
let provider_type = ProviderType::from_str(&default_provider);
agent_state.init(base_url, api_key, provider_type, Some(default_provider))?;
}
// 从凭证池配置 Provider
let model_name = model
.clone()
.unwrap_or_else(|| "claude-sonnet-4-20250514".to_string());
let aster_config = agent_state
.configure_provider_from_pool(&db, &provider_type, &model_name, &session_id)
.await?;
// 构建包含 Skills 的 System Prompt
let final_system_prompt = build_system_prompt_with_skills(system_prompt, skills.as_ref());
let session_id = agent_state.create_session(model.clone(), final_system_prompt.clone())?;
// 保存会话到数据库
let now = chrono::Utc::now().to_rfc3339();
let session = AgentSession {
id: session_id.clone(),
model: model.clone().unwrap_or_else(|| "default".to_string()),
model: model_name.clone(),
messages: Vec::new(),
system_prompt: final_system_prompt,
created_at: now.clone(),
@@ -185,9 +160,9 @@ pub async fn agent_create_session(
Ok(CreateSessionResponse {
session_id,
credential_name: "ProxyCast".to_string(),
credential_uuid: "native-agent".to_string(),
credential_uuid: aster_config.credential_uuid,
provider_type,
model,
model: Some(model_name),
})
}
@@ -234,98 +209,19 @@ pub struct ImageInputParam {
}
/// 发送消息到 Agent
///
/// 注意:此命令已废弃,请使用 aster_agent_chat_stream
#[tauri::command]
pub async fn agent_send_message(
agent_state: State<'_, NativeAgentState>,
app_state: State<'_, AppState>,
session_id: Option<String>,
message: String,
images: Option<Vec<ImageInputParam>>,
model: Option<String>,
web_search: Option<bool>,
thinking: Option<bool>,
_agent_state: State<'_, AsterAgentState>,
_session_id: Option<String>,
_message: String,
_images: Option<Vec<ImageInputParam>>,
_model: Option<String>,
_web_search: Option<bool>,
_thinking: Option<bool>,
) -> Result<String, String> {
let images_count = images.as_ref().map(|v| v.len()).unwrap_or(0);
let images_sizes: Vec<usize> = images
.as_ref()
.map(|imgs| imgs.iter().map(|i| i.data.len()).collect())
.unwrap_or_default();
tracing::info!(
"[Agent] 发送消息: len={}, session={:?}, images_count={}, images_sizes={:?}, web_search={:?}, thinking={:?}",
message.len(),
session_id,
images_count,
images_sizes,
web_search,
thinking
);
// 如果未初始化,自动初始化
if !agent_state.is_initialized() {
let (host, port, api_key, running, default_provider) = {
let state = app_state.read().await;
(
state.config.server.host.clone(),
state.config.server.port,
state.running_api_key.clone(),
state.running,
state.config.routing.default_provider.clone(),
)
};
if !running {
return Err("ProxyCast API Server 未运行".to_string());
}
let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?;
let base_url = get_local_url(&host, port);
let provider_type = ProviderType::from_str(&default_provider);
agent_state.init(base_url, api_key, provider_type, Some(default_provider))?;
}
// 根据启用的模式构建最终消息
let web_search_enabled = web_search.unwrap_or(false);
let thinking_enabled = thinking.unwrap_or(false);
let final_message = match (web_search_enabled, thinking_enabled) {
(true, true) => format!(
"[深度思考 + 联网搜索模式] 请深入分析问题,并搜索网络获取最新信息,然后给出详细的回答:\n\n{}",
message
),
(true, false) => format!(
"[联网搜索模式] 请先搜索网络获取最新信息,然后回答以下问题:\n\n{}",
message
),
(false, true) => format!(
"[深度思考模式] 请深入分析这个问题,考虑多个角度,给出详细的推理过程和结论:\n\n{}",
message
),
(false, false) => message,
};
let request = NativeChatRequest {
session_id,
message: final_message,
model,
images: images.map(|imgs| {
imgs.into_iter()
.map(|img| ImageData {
data: img.data,
media_type: img.media_type,
})
.collect()
}),
stream: false,
};
let response = agent_state.chat(request).await?;
if response.success {
Ok(response.content)
} else {
Err(response.error.unwrap_or_else(|| "未知错误".to_string()))
}
Err("此命令已废弃,请使用 aster_agent_chat_stream 进行流式对话".to_string())
}
/// 会话信息
@@ -347,14 +243,13 @@ pub async fn agent_list_sessions(db: State<'_, DbConnection>) -> Result<Vec<Sess
let sessions =
AgentDao::list_sessions(&conn).map_err(|e| format!("获取会话列表失败: {}", e))?;
// 获取每个会话的消息数量
let result: Vec<SessionInfo> = sessions
.into_iter()
.map(|s| {
let messages_count = AgentDao::get_message_count(&conn, &s.id).unwrap_or(0);
SessionInfo {
session_id: s.id,
provider_type: "native".to_string(),
provider_type: "aster".to_string(),
model: Some(s.model),
created_at: s.created_at.clone(),
last_activity: s.updated_at,
@@ -369,37 +264,35 @@ pub async fn agent_list_sessions(db: State<'_, DbConnection>) -> Result<Vec<Sess
/// 获取会话详情
#[tauri::command]
pub async fn agent_get_session(
agent_state: State<'_, NativeAgentState>,
db: State<'_, DbConnection>,
session_id: String,
) -> Result<SessionInfo, String> {
let session = agent_state
.get_session(&session_id)?
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?;
let session = AgentDao::get_session(&conn, &session_id)
.map_err(|e| format!("获取会话失败: {}", e))?
.ok_or_else(|| "会话不存在".to_string())?;
let messages_count = AgentDao::get_message_count(&conn, &session_id).unwrap_or(0);
Ok(SessionInfo {
session_id: session.id,
provider_type: "native".to_string(),
provider_type: "aster".to_string(),
model: Some(session.model),
created_at: session.created_at.clone(),
last_activity: session.created_at,
messages_count: session.messages.len(),
last_activity: session.updated_at,
messages_count,
})
}
/// 删除会话
#[tauri::command]
pub async fn agent_delete_session(
agent_state: State<'_, NativeAgentState>,
db: State<'_, DbConnection>,
session_id: String,
) -> Result<(), String> {
// 从内存中删除
agent_state.delete_session(&session_id);
// 从数据库中删除
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?;
AgentDao::delete_session(&conn, &session_id).map_err(|e| format!("删除会话失败: {}", e))?;
Ok(())
}
@@ -410,84 +303,7 @@ pub async fn agent_get_session_messages(
session_id: String,
) -> Result<Vec<AgentMessage>, String> {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?;
let messages =
AgentDao::get_messages(&conn, &session_id).map_err(|e| format!("获取消息失败: {}", e))?;
Ok(messages)
}
/// 处理终端命令响应
///
/// 前端在用户批准/拒绝命令后调用此命令,将结果传递给 TerminalTool
#[tauri::command]
pub async fn agent_terminal_command_response(
request_id: String,
success: bool,
output: String,
error: Option<String>,
exit_code: Option<i32>,
rejected: bool,
) -> Result<(), String> {
tracing::info!(
"[Agent] 收到终端命令响应: request_id={}, success={}, rejected={}",
request_id,
success,
rejected
);
// TODO: 重新实现终端命令响应处理,适配 aster-rust 工具系统
// let response = TerminalCommandResponse {
// request_id,
// success,
// output,
// error,
// exit_code,
// rejected,
// };
// handle_terminal_command_response(response);
tracing::warn!("[Agent] 终端命令响应处理暂时禁用,等待适配 aster-rust 工具系统");
Ok(())
}
/// 前端返回终端滚动缓冲区数据
#[tauri::command]
pub async fn agent_term_scrollback_response(
request_id: String,
success: bool,
total_lines: usize,
line_start: usize,
line_end: usize,
content: String,
has_more: bool,
error: Option<String>,
) -> Result<(), String> {
tracing::info!(
"[Agent] 收到终端滚动缓冲区响应: request_id={}, success={}, lines={}-{}",
request_id,
success,
line_start,
line_end
);
// TODO: 重新实现终端滚动缓冲区响应处理,适配 aster-rust 工具系统
// let response = GetScrollbackResponse {
// request_id,
// success,
// total_lines,
// line_start,
// line_end,
// content,
// has_more,
// error,
// };
// handle_term_scrollback_response(response);
tracing::warn!("[Agent] 终端滚动缓冲区响应处理暂时禁用,等待适配 aster-rust 工具系统");
Ok(())
}
+15 -1
View File
@@ -10,7 +10,7 @@ use crate::database::dao::api_key_provider::{
};
use crate::database::DbConnection;
use crate::services::api_key_provider_service::{
ApiKeyProviderService, ConnectionTestResult, ImportResult,
ApiKeyProviderService, ChatTestResult, ConnectionTestResult, ImportResult,
};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
@@ -612,3 +612,17 @@ pub async fn test_api_key_provider_connection(
.test_connection(&db, &provider_id, model_name)
.await
}
#[tauri::command]
pub async fn test_api_key_provider_chat(
db: State<'_, DbConnection>,
service: State<'_, ApiKeyProviderServiceState>,
provider_id: String,
model_name: Option<String>,
prompt: String,
) -> Result<ChatTestResult, String> {
service
.0
.test_chat(&db, &provider_id, model_name, prompt)
.await
}
+92 -6
View File
@@ -2,18 +2,51 @@
//!
//! 提供基于 Aster 框架的 Tauri 命令
//! 这是新的对话系统实现,与 native_agent_cmd.rs 并行存在
//! 支持从 ProxyCast 凭证池自动选择凭证
use crate::agent::aster_state::{ProviderConfig, SessionConfigBuilder};
use crate::agent::event_converter::convert_agent_event;
use crate::agent::{
AsterAgentState, AsterAgentWrapper, SessionDetail, SessionInfo, TauriAgentEvent,
};
use crate::database::DbConnection;
use aster::conversation::message::Message;
use aster::session::SessionManager;
use futures::StreamExt;
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use tauri::{AppHandle, Emitter, State};
/// 确保 session 在 Aster 数据库中存在
/// 如果不存在则创建新的 session
async fn ensure_session_exists(session_id: &str) -> Result<String, String> {
// 尝试获取现有 session
match SessionManager::get_session(session_id, false).await {
Ok(_) => {
tracing::debug!("[AsterAgent] Session 已存在: {}", session_id);
Ok(session_id.to_string())
}
Err(_) => {
// Session 不存在,创建新的
tracing::info!(
"[AsterAgent] Session 不存在,创建新 session: {}",
session_id
);
let working_dir = std::env::current_dir().unwrap_or_else(|_| PathBuf::from("."));
let session = SessionManager::create_session(
working_dir,
"New Chat".to_string(),
aster::session::SessionType::User,
)
.await
.map_err(|e| format!("创建 session 失败: {}", e))?;
tracing::info!("[AsterAgent] 创建新 session: {}", session.id);
Ok(session.id)
}
}
}
/// Aster Agent 状态信息
#[derive(Debug, Serialize)]
pub struct AsterAgentStatus {
@@ -21,6 +54,9 @@ pub struct AsterAgentStatus {
pub provider_configured: bool,
pub provider_name: Option<String>,
pub model_name: Option<String>,
/// 凭证 UUID(来自凭证池)
#[serde(skip_serializing_if = "Option::is_none")]
pub credential_uuid: Option<String>,
}
/// Provider 配置请求
@@ -34,6 +70,15 @@ pub struct ConfigureProviderRequest {
pub base_url: Option<String>,
}
/// 从凭证池配置 Provider 的请求
#[derive(Debug, Deserialize)]
pub struct ConfigureFromPoolRequest {
/// Provider 类型 (openai, anthropic, kiro, gemini 等)
pub provider_type: String,
/// 模型名称
pub model_name: String,
}
/// 初始化 Aster Agent
#[tauri::command]
pub async fn aster_agent_init(
@@ -52,6 +97,7 @@ pub async fn aster_agent_init(
provider_configured: provider_config.is_some(),
provider_name: provider_config.as_ref().map(|c| c.provider_name.clone()),
model_name: provider_config.as_ref().map(|c| c.model_name.clone()),
credential_uuid: provider_config.and_then(|c| c.credential_uuid),
})
}
@@ -73,6 +119,7 @@ pub async fn aster_agent_configure_provider(
model_name: request.model_name,
api_key: request.api_key,
base_url: request.base_url,
credential_uuid: None,
};
state
@@ -84,6 +131,41 @@ pub async fn aster_agent_configure_provider(
provider_configured: true,
provider_name: Some(config.provider_name),
model_name: Some(config.model_name),
credential_uuid: None,
})
}
/// 从凭证池配置 Aster Agent 的 Provider
///
/// 自动从 ProxyCast 凭证池选择可用凭证并配置 Aster Provider
#[tauri::command]
pub async fn aster_agent_configure_from_pool(
state: State<'_, AsterAgentState>,
db: State<'_, DbConnection>,
request: ConfigureFromPoolRequest,
session_id: String,
) -> Result<AsterAgentStatus, String> {
tracing::info!(
"[AsterAgent] 从凭证池配置 Provider: {} / {}",
request.provider_type,
request.model_name
);
let aster_config = state
.configure_provider_from_pool(
&db,
&request.provider_type,
&request.model_name,
&session_id,
)
.await?;
Ok(AsterAgentStatus {
initialized: true,
provider_configured: true,
provider_name: Some(aster_config.provider_name),
model_name: Some(aster_config.model_name),
credential_uuid: Some(aster_config.credential_uuid),
})
}
@@ -98,6 +180,7 @@ pub async fn aster_agent_status(
provider_configured: provider_config.is_some(),
provider_name: provider_config.as_ref().map(|c| c.provider_name.clone()),
model_name: provider_config.as_ref().map(|c| c.model_name.clone()),
credential_uuid: provider_config.and_then(|c| c.credential_uuid),
})
}
@@ -139,6 +222,10 @@ pub async fn aster_agent_chat_stream(
state.init_agent().await?;
}
// 确保 session 在 Aster 数据库中存在
// 如果 session 不存在,自动创建
let session_id = ensure_session_exists(&request.session_id).await?;
// 如果提供了 Provider 配置,则配置 Provider
if let Some(provider_config) = &request.provider_config {
let config = ProviderConfig {
@@ -146,10 +233,9 @@ pub async fn aster_agent_chat_stream(
model_name: provider_config.model_name.clone(),
api_key: provider_config.api_key.clone(),
base_url: provider_config.base_url.clone(),
credential_uuid: None,
};
state
.configure_provider(config, &request.session_id)
.await?;
state.configure_provider(config, &session_id).await?;
}
// 检查 Provider 是否已配置
@@ -158,13 +244,13 @@ pub async fn aster_agent_chat_stream(
}
// 创建取消令牌
let cancel_token = state.create_cancel_token(&request.session_id).await;
let cancel_token = state.create_cancel_token(&session_id).await;
// 创建用户消息
let user_message = Message::user().with_text(&request.message);
// 创建会话配置
let session_config = SessionConfigBuilder::new(&request.session_id).build();
let session_config = SessionConfigBuilder::new(&session_id).build();
// 获取 Agent Arc 并保持 guard 在整个流处理期间存活
let agent_arc = state.get_agent_arc();
@@ -225,7 +311,7 @@ pub async fn aster_agent_chat_stream(
// guard 会在函数结束时自动释放(stream_result 先释放)
// 清理取消令牌
state.remove_cancel_token(&request.session_id).await;
state.remove_cancel_token(&session_id).await;
Ok(())
}
-1
View File
@@ -17,7 +17,6 @@ pub mod model_cmd;
pub mod model_registry_cmd;
pub mod models_cmd;
pub mod music_cmd;
pub mod native_agent_cmd;
pub mod network_cmd;
pub mod oauth_cmd;
pub mod orchestrator_cmd;
-510
View File
@@ -1,510 +0,0 @@
//! 原生 Agent 命令模块
//!
//! 提供原生 Rust Agent 的 Tauri 命令,替代 aster sidecar 方案
use crate::agent::{
AgentMessage, AgentSession, ImageData, MessageContent, NativeAgentState, NativeChatRequest,
NativeChatResponse, ProviderType, StreamEvent, ToolLoopEngine,
};
use crate::commands::network_cmd::get_local_url;
use crate::database::dao::agent::AgentDao;
use crate::database::dao::api_key_provider::ApiKeyProviderDao;
use crate::database::DbConnection;
use crate::AppState;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use tauri::{Emitter, State};
use tokio::sync::mpsc;
#[derive(Debug, Serialize)]
pub struct NativeAgentStatus {
pub initialized: bool,
pub base_url: Option<String>,
}
#[tauri::command]
pub async fn native_agent_init(
agent_state: State<'_, NativeAgentState>,
app_state: State<'_, AppState>,
db: State<'_, crate::database::DbConnection>,
) -> Result<NativeAgentStatus, String> {
tracing::info!("[NativeAgent] 初始化 Agent");
let (host, port, api_key, running, default_provider, agent_config) = {
let state = app_state.read().await;
(
state.config.server.host.clone(),
state.config.server.port,
state.running_api_key.clone(),
state.running,
state.config.routing.default_provider.clone(),
state.config.agent.clone(),
)
};
if !running {
return Err("ProxyCast API Server 未运行,请先启动服务器".to_string());
}
let api_key = api_key.ok_or_else(|| "ProxyCast API Server 未配置 API Key".to_string())?;
let base_url = get_local_url(&host, port);
// 对于自定义 Provider ID(如 custom-xxx),使用 Anthropic 兼容协议
let provider_type = if default_provider.starts_with("custom-") {
tracing::info!(
"[NativeAgent] 自定义 Provider ID '{}',使用 Anthropic 兼容协议",
default_provider
);
ProviderType::AnthropicCompatible
} else {
// 对于标准 Provider,从数据库查询类型
match crate::database::dao::api_key_provider::ApiKeyProviderDao::get_provider_by_id(
&*db.lock().map_err(|e| e.to_string())?,
&default_provider,
) {
Ok(Some(provider)) => {
// 从数据库的 provider_type 转换为 ProviderType
match provider.provider_type {
crate::database::dao::api_key_provider::ApiProviderType::Anthropic |
crate::database::dao::api_key_provider::ApiProviderType::AnthropicCompatible => {
tracing::info!(
"[NativeAgent] 从数据库获取 Provider 类型: {:?} (Anthropic)",
provider.provider_type
);
ProviderType::Claude
}
crate::database::dao::api_key_provider::ApiProviderType::Gemini => {
tracing::info!(
"[NativeAgent] 从数据库获取 Provider 类型: {:?} (Gemini)",
provider.provider_type
);
ProviderType::Gemini
}
_ => {
tracing::info!(
"[NativeAgent] 从数据库获取 Provider 类型: {:?} (OpenAI)",
provider.provider_type
);
ProviderType::OpenAI
}
}
}
Ok(None) => {
tracing::warn!(
"[NativeAgent] 数据库中未找到 Provider '{}',使用字符串解析",
default_provider
);
ProviderType::from_str(&default_provider)
}
Err(e) => {
tracing::warn!(
"[NativeAgent] 从数据库查询 Provider 失败: {},使用字符串解析",
e
);
ProviderType::from_str(&default_provider)
}
}
};
tracing::info!(
"[NativeAgent] 初始化 Agent: base_url={}, provider={:?}, use_default_prompt={}",
base_url,
provider_type,
agent_config.use_default_system_prompt
);
// 使用带配置的初始化方法
agent_state.init_with_config(
base_url.clone(),
api_key,
provider_type,
Some(default_provider.to_string()),
&agent_config,
)?;
tracing::info!("[NativeAgent] Agent 初始化成功: {}", base_url);
Ok(NativeAgentStatus {
initialized: true,
base_url: Some(base_url),
})
}
#[tauri::command]
pub async fn native_agent_status(
agent_state: State<'_, NativeAgentState>,
) -> Result<NativeAgentStatus, String> {
Ok(NativeAgentStatus {
initialized: agent_state.is_initialized(),
base_url: None,
})
}
#[tauri::command]
pub async fn native_agent_reset(agent_state: State<'_, NativeAgentState>) -> Result<(), String> {
agent_state.reset();
tracing::info!("[NativeAgent] Agent 已重置");
Ok(())
}
#[derive(Debug, Deserialize)]
pub struct ImageInputParam {
pub data: String,
pub media_type: String,
}
#[tauri::command]
pub async fn native_agent_chat(
agent_state: State<'_, NativeAgentState>,
app_state: State<'_, AppState>,
message: String,
model: Option<String>,
images: Option<Vec<ImageInputParam>>,
) -> Result<NativeChatResponse, String> {
tracing::info!(
"[NativeAgent] 发送消息: message_len={}, model={:?}",
message.len(),
model
);
// 如果 Agent 未初始化,自动初始化
if !agent_state.is_initialized() {
let (host, port, api_key, running, default_provider) = {
let state = app_state.read().await;
(
state.config.server.host.clone(),
state.config.server.port,
state.running_api_key.clone(),
state.running,
state.config.routing.default_provider.clone(),
)
};
if !running {
return Err("ProxyCast API Server 未运行".to_string());
}
let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?;
let base_url = get_local_url(&host, port);
let provider_type = ProviderType::from_str(&default_provider);
agent_state.init(base_url, api_key, provider_type, Some(default_provider))?;
}
let request = NativeChatRequest {
session_id: None,
message,
model,
images: images.map(|imgs| {
imgs.into_iter()
.map(|img| ImageData {
data: img.data,
media_type: img.media_type,
})
.collect()
}),
stream: false,
};
// 使用 chat_sync 方法避免跨 await 持有锁
agent_state.chat(request).await
}
#[tauri::command]
pub async fn native_agent_chat_stream(
app_handle: tauri::AppHandle,
agent_state: State<'_, NativeAgentState>,
app_state: State<'_, AppState>,
db: State<'_, DbConnection>,
message: String,
event_name: String,
session_id: Option<String>,
model: Option<String>,
images: Option<Vec<ImageInputParam>>,
provider: Option<String>,
terminal_mode: Option<bool>,
) -> Result<(), String> {
let terminal_mode = terminal_mode.unwrap_or(false);
tracing::info!(
"[NativeAgent] 发送流式消息: message_len={}, model={:?}, provider={:?}, event={}, session={:?}, terminal_mode={}",
message.len(),
model,
provider,
event_name,
session_id,
terminal_mode
);
// 获取配置信息
let (host, port, api_key, running, default_provider) = {
let state = app_state.read().await;
(
state.config.server.host.clone(),
state.config.server.port,
state.running_api_key.clone(),
state.running,
state.config.routing.default_provider.clone(),
)
};
if !running {
return Err("ProxyCast API Server 未运行".to_string());
}
let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?;
// 使用前端传递的 provider,如果没有则使用默认值
let provider_str = provider.unwrap_or(default_provider);
// 尝试从数据库查询 Provider 的类型(用于确定协议)
let provider_type = {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?;
if let Ok(Some(api_provider)) = ApiKeyProviderDao::get_provider_by_id(&conn, &provider_str)
{
// 根据 API Key Provider 的 type 确定协议
let api_type = api_provider.provider_type.to_string();
tracing::info!(
"[NativeAgent] 从数据库获取 Provider 类型: {} -> {}",
provider_str,
api_type
);
ProviderType::from_str(&api_type)
} else {
// 数据库中没有找到,使用默认解析
ProviderType::from_str(&provider_str)
}
};
tracing::info!(
"[NativeAgent] 使用 provider: {:?} (原始值: {})",
provider_type,
provider_str
);
// 如果 Agent 未初始化,或者 provider 发生变化,重新初始化
// 使用 provider_str 而不是 provider_type 来判断,因为自定义 Provider 的 type 都是 OpenAI
let need_reinit = if !agent_state.is_initialized() {
tracing::info!("[NativeAgent] Agent 未初始化,需要初始化");
true
} else if let Some(current_provider_id) = agent_state.get_provider_id() {
if current_provider_id != provider_str {
tracing::info!(
"[NativeAgent] Provider 发生变化: {} -> {},需要重新初始化",
current_provider_id,
provider_str
);
true
} else {
false
}
} else {
true
};
if need_reinit {
let base_url = get_local_url(&host, port);
agent_state.init(base_url, api_key, provider_type, Some(provider_str.clone()))?;
}
// 获取工具注册表(用于创建 ToolLoopEngine)
// 如果是 terminal_mode,使用 TerminalTool 替代 BashTool
let tool_registry = agent_state.get_tool_registry_with_mode(terminal_mode)?;
// 保存用户消息到数据库
let session_id_for_db = session_id.clone();
let message_for_db = message.clone();
let db_clone = Arc::clone(&db);
if let Some(ref sid) = session_id_for_db {
let conn = db_clone
.lock()
.map_err(|e| format!("数据库锁定失败: {}", e))?;
let user_message = AgentMessage {
role: "user".to_string(),
content: MessageContent::Text(message_for_db.clone()),
timestamp: chrono::Utc::now().to_rfc3339(),
tool_calls: None,
tool_call_id: None,
reasoning_content: None,
};
if let Err(e) = AgentDao::add_message(&conn, sid, &user_message) {
tracing::warn!("[NativeAgent] 保存用户消息到数据库失败: {}", e);
}
}
let request = NativeChatRequest {
session_id, // 使用前端传递的 session_id 以保持上下文
message,
model,
images: images.map(|imgs| {
imgs.into_iter()
.map(|img| ImageData {
data: img.data,
media_type: img.media_type,
})
.collect()
}),
stream: true,
};
// 克隆 agent_state 用于后台任务(共享 sessions)
let agent_state_clone = agent_state.inner().clone();
let session_id_for_task = session_id_for_db.clone();
// 在后台任务中处理流式响应
let event_name_clone = event_name.clone();
eprintln!(
"[native_agent_chat_stream] 启动后台任务, event_name={}",
event_name_clone
);
tauri::async_runtime::spawn(async move {
eprintln!("[native_agent_chat_stream] 后台任务开始执行");
// 创建工具循环引擎(使用共享的 tool_registry)
let tool_loop_engine = ToolLoopEngine::new(tool_registry);
eprintln!("[native_agent_chat_stream] 工具循环引擎创建成功");
let (tx, mut rx) = mpsc::channel::<StreamEvent>(100);
// 用于收集完整的助手响应
let mut full_content = String::new();
// 使用 agent_state 的方法(共享 sessions)
eprintln!(
"[native_agent_chat_stream] 开始 chat_stream_with_tools, request.session_id={:?}",
request.session_id
);
let stream_task = tokio::spawn(async move {
agent_state_clone
.chat_stream_with_tools(request, tx, &tool_loop_engine)
.await
});
eprintln!("[native_agent_chat_stream] 开始接收流式事件...");
// 注意:不要在收到 Done 事件后立即 break,因为工具循环可能还在执行
// 继续接收直到 channel 关闭(stream_task 完成)
while let Some(event) = rx.recv().await {
eprintln!("[native_agent_chat_stream] 收到事件: {:?}", event);
tracing::debug!(
"[NativeAgent] 收到流式事件: {:?}, 发送到: {}",
event,
event_name_clone
);
// 收集文本增量
if let StreamEvent::TextDelta { ref text } = event {
full_content.push_str(text);
}
if let Err(e) = app_handle.emit(&event_name_clone, &event) {
tracing::error!("[NativeAgent] 发送事件失败: {}", e);
eprintln!("[native_agent_chat_stream] 发送事件失败: {}", e);
break;
}
tracing::debug!("[NativeAgent] 事件发送成功");
// 只在 Error 时 break,Done 不 break 因为工具循环可能还会发送更多事件
if matches!(event, StreamEvent::Error { .. }) {
tracing::info!("[NativeAgent] 流式响应错误,停止接收");
eprintln!("[native_agent_chat_stream] 流式响应错误");
break;
}
}
eprintln!("[native_agent_chat_stream] channel 关闭,事件接收完成");
eprintln!("[native_agent_chat_stream] 等待 stream_task 完成...");
match stream_task.await {
Ok(result) => {
eprintln!("[native_agent_chat_stream] stream_task 完成: {:?}", result);
// 保存助手消息到数据库
if let Some(ref sid) = session_id_for_task {
if !full_content.is_empty() {
if let Ok(conn) = db_clone.lock() {
let assistant_message = AgentMessage {
role: "assistant".to_string(),
content: MessageContent::Text(full_content.clone()),
timestamp: chrono::Utc::now().to_rfc3339(),
tool_calls: None,
tool_call_id: None,
reasoning_content: None,
};
if let Err(e) = AgentDao::add_message(&conn, sid, &assistant_message) {
tracing::warn!("[NativeAgent] 保存助手消息到数据库失败: {}", e);
} else {
tracing::info!("[NativeAgent] 助手消息已保存到数据库");
}
}
}
}
}
Err(e) => {
eprintln!("[native_agent_chat_stream] stream_task 错误: {}", e);
}
}
eprintln!("[native_agent_chat_stream] 后台任务结束");
});
Ok(())
}
#[tauri::command]
pub async fn native_agent_create_session(
agent_state: State<'_, NativeAgentState>,
model: Option<String>,
system_prompt: Option<String>,
) -> Result<String, String> {
agent_state.create_session(model, system_prompt)
}
#[tauri::command]
pub async fn native_agent_get_session(
agent_state: State<'_, NativeAgentState>,
session_id: String,
) -> Result<Option<AgentSession>, String> {
agent_state.get_session(&session_id)
}
#[tauri::command]
pub async fn native_agent_delete_session(
agent_state: State<'_, NativeAgentState>,
session_id: String,
) -> Result<bool, String> {
Ok(agent_state.delete_session(&session_id))
}
#[tauri::command]
pub async fn native_agent_list_sessions(
agent_state: State<'_, NativeAgentState>,
) -> Result<Vec<AgentSession>, String> {
Ok(agent_state.list_sessions())
}
/// 权限确认响应请求
#[derive(Debug, Deserialize)]
pub struct PermissionResponseRequest {
pub request_id: String,
pub confirmed: bool,
pub response: Option<String>,
}
/// 发送权限确认响应
#[tauri::command]
pub async fn agent_permission_response(
_agent_state: State<'_, NativeAgentState>,
request_id: String,
confirmed: bool,
response: Option<String>,
) -> Result<(), String> {
tracing::info!(
"[NativeAgent] 权限确认响应: id={}, confirmed={}, response={:?}",
request_id,
confirmed,
response
);
// TODO: 实现权限确认响应逻辑
// 这需要与工具循环引擎集成,将用户的确认响应传递给等待中的工具
// 目前先记录日志并返回成功
Ok(())
}
@@ -19,6 +19,8 @@ use serde::{Deserialize, Serialize};
pub enum ApiProviderType {
Openai,
OpenaiResponse,
/// Codex CLI 协议(使用 /responses 端点)
Codex,
Anthropic,
/// Anthropic 兼容格式(支持 system 数组格式等变体)
AnthropicCompatible,
@@ -36,6 +38,7 @@ impl std::fmt::Display for ApiProviderType {
match self {
ApiProviderType::Openai => write!(f, "openai"),
ApiProviderType::OpenaiResponse => write!(f, "openai-response"),
ApiProviderType::Codex => write!(f, "codex"),
ApiProviderType::Anthropic => write!(f, "anthropic"),
ApiProviderType::AnthropicCompatible => write!(f, "anthropic-compatible"),
ApiProviderType::Gemini => write!(f, "gemini"),
@@ -56,6 +59,7 @@ impl std::str::FromStr for ApiProviderType {
match s.to_lowercase().as_str() {
"openai" => Ok(ApiProviderType::Openai),
"openai-response" => Ok(ApiProviderType::OpenaiResponse),
"codex" => Ok(ApiProviderType::Codex),
"anthropic" => Ok(ApiProviderType::Anthropic),
"anthropic-compatible" => Ok(ApiProviderType::AnthropicCompatible),
"gemini" => Ok(ApiProviderType::Gemini),
+178 -19
View File
@@ -1,9 +1,11 @@
//! OpenAI Custom Provider (自定义 OpenAI 兼容 API)
use crate::models::openai::ChatCompletionRequest;
use reqwest::Client;
use reqwest::StatusCode;
use serde::{Deserialize, Serialize};
use std::error::Error;
use std::time::Duration;
use url::Url;
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct OpenAICustomConfig {
@@ -100,6 +102,96 @@ impl OpenAICustomProvider {
}
}
fn build_url_fallback_without_v1(&self, endpoint: &str) -> Option<String> {
let url = self.build_url(endpoint);
if url.contains("/v1/") {
Some(url.replacen("/v1/", "/", 1))
} else {
None
}
}
fn build_url_from_base(base_url: &str, endpoint: &str) -> String {
let base = base_url.trim_end_matches('/');
let has_version = base
.rsplit('/')
.next()
.map(|last_segment| {
last_segment.starts_with('v')
&& last_segment.len() >= 2
&& last_segment[1..].chars().all(|c| c.is_ascii_digit())
})
.unwrap_or(false);
if has_version {
format!("{}/{}", base, endpoint)
} else {
format!("{}/v1/{}", base, endpoint)
}
}
fn base_url_parent(&self) -> Option<String> {
let base = self.get_base_url();
let base = base.trim();
let mut url = Url::parse(base)
.or_else(|_| Url::parse(&format!("http://{}", base)))
.ok()?;
let path = url.path().trim_end_matches('/');
if path.is_empty() || path == "/" {
return None;
}
let mut segments: Vec<&str> = path.split('/').filter(|s| !s.is_empty()).collect();
if segments.is_empty() {
return None;
}
segments.pop();
let new_path = if segments.is_empty() {
"/".to_string()
} else {
format!("/{}", segments.join("/"))
};
url.set_path(&new_path);
url.set_query(None);
url.set_fragment(None);
Some(url.to_string().trim_end_matches('/').to_string())
}
fn build_urls_with_fallbacks(&self, endpoint: &str) -> Vec<String> {
let mut urls: Vec<String> = Vec::new();
let primary = self.build_url(endpoint);
urls.push(primary.clone());
if let Some(no_v1) = self.build_url_fallback_without_v1(endpoint) {
if no_v1 != primary {
urls.push(no_v1);
}
}
if let Some(parent_base) = self.base_url_parent() {
let u = Self::build_url_from_base(&parent_base, endpoint);
if !urls.iter().any(|x| x == &u) {
urls.push(u.clone());
}
if u.contains("/v1/") {
let u2 = u.replacen("/v1/", "/", 1);
if !urls.iter().any(|x| x == &u2) {
urls.push(u2);
}
}
}
urls
}
/// 调用 OpenAI API(使用类型化请求)
pub async fn call_api(
&self,
@@ -111,18 +203,32 @@ impl OpenAICustomProvider {
.as_ref()
.ok_or("OpenAI API key not configured")?;
let url = self.build_url("chat/completions");
let urls = self.build_urls_with_fallbacks("chat/completions");
let mut last_resp: Option<reqwest::Response> = None;
let resp = self
.client
.post(&url)
.header("Authorization", format!("Bearer {api_key}"))
.header("Content-Type", "application/json")
.json(request)
.send()
.await?;
eprintln!(
"[OPENAI_CUSTOM] call_api testing with model: {}",
request.model
);
Ok(resp)
for url in &urls {
eprintln!("[OPENAI_CUSTOM] call_api trying URL: {}", url);
let resp = self
.client
.post(url)
.header("Authorization", format!("Bearer {api_key}"))
.header("Content-Type", "application/json")
.json(request)
.send()
.await?;
if resp.status() != StatusCode::NOT_FOUND {
return Ok(resp);
}
last_resp = Some(resp);
}
Ok(last_resp.ok_or("Request failed")?)
}
pub async fn chat_completions(
@@ -152,6 +258,22 @@ impl OpenAICustomProvider {
.send()
.await?;
if resp.status() == StatusCode::NOT_FOUND {
if let Some(fallback_url) = self.build_url_fallback_without_v1("chat/completions") {
if fallback_url != url {
let resp2 = self
.client
.post(&fallback_url)
.header("Authorization", format!("Bearer {api_key}"))
.header("Content-Type", "application/json")
.json(request)
.send()
.await?;
return Ok(resp2);
}
}
}
Ok(resp)
}
@@ -162,22 +284,37 @@ impl OpenAICustomProvider {
.as_ref()
.ok_or("OpenAI API key not configured")?;
let url = self.build_url("models");
let urls = self.build_urls_with_fallbacks("models");
let mut tried_urls: Vec<String> = Vec::new();
let mut resp: Option<reqwest::Response> = None;
eprintln!("[OPENAI_CUSTOM] list_models URL: {}", url);
for url in urls {
eprintln!("[OPENAI_CUSTOM] list_models URL: {}", url);
tried_urls.push(url.clone());
let r = self
.client
.get(&url)
.header("Authorization", format!("Bearer {api_key}"))
.send()
.await?;
if r.status() != StatusCode::NOT_FOUND {
resp = Some(r);
break;
}
resp = Some(r);
}
let resp = self
.client
.get(&url)
.header("Authorization", format!("Bearer {api_key}"))
.send()
.await?;
let resp = resp.ok_or("Request failed")?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
eprintln!("[OPENAI_CUSTOM] list_models 失败: {} - {}", status, body);
return Err(format!("Failed to list models: {status} - {body}").into());
return Err(format!(
"Failed to list models: {status} - {body} (tried: {})",
tried_urls.join(", ")
)
.into());
}
let data: serde_json::Value = resp.json().await?;
@@ -235,6 +372,28 @@ impl StreamingProvider for OpenAICustomProvider {
.await
.map_err(|e| ProviderError::from_reqwest_error(&e))?;
let resp = if resp.status() == StatusCode::NOT_FOUND {
if let Some(fallback_url) = self.build_url_fallback_without_v1("chat/completions") {
if fallback_url != url {
self.client
.post(&fallback_url)
.header("Authorization", format!("Bearer {api_key}"))
.header("Content-Type", "application/json")
.header("Accept", "text/event-stream")
.json(&stream_request)
.send()
.await
.map_err(|e| ProviderError::from_reqwest_error(&e))?
} else {
resp
}
} else {
resp
}
} else {
resp
};
// 检查响应状态
let status = resp.status();
if !status.is_success() {
+9
View File
@@ -31,6 +31,15 @@
- `update_check_service.rs` - 自动更新检查服务(每日检查、系统通知)
- `update_window.rs` - 更新提醒独立窗口管理
## Aster Agent 集成
Aster Agent 集成位于 `src-tauri/src/agent/` 目录:
- `aster_state.rs` - Agent 状态管理
- `aster_agent.rs` - Agent 包装器
- `event_converter.rs` - 事件转换器
详见 `docs/aiprompts/aster-integration.md`
## 更新提醒
任何文件变更后,请更新此文档和相关的上级文档。
@@ -37,6 +37,40 @@ pub struct ConnectionTestResult {
pub models: Option<Vec<String>>,
}
#[cfg(test)]
mod tests {
use super::ApiKeyProviderService;
#[test]
fn test_build_codex_responses_request_input_list() {
let req = ApiKeyProviderService::build_codex_responses_request("gpt-5", "hello");
assert!(req.get("input").is_some());
let input = req["input"].as_array().expect("input should be array");
assert_eq!(input.len(), 1);
assert_eq!(input[0]["role"].as_str(), Some("user"));
assert_eq!(input[0]["content"][0]["type"].as_str(), Some("input_text"));
assert_eq!(input[0]["content"][0]["text"].as_str(), Some("hello"));
}
#[test]
fn test_parse_codex_responses_sse_content_delta() {
let body = "data: {\"type\":\"response.output_text.delta\",\"delta\":\"hi\"}\n\n\
data: {\"type\":\"response.output_text.delta\",\"delta\":\"!\"}\n\n\
data: [DONE]\n";
let content = ApiKeyProviderService::parse_codex_responses_sse_content(body);
assert_eq!(content, "hi!");
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ChatTestResult {
pub success: bool,
pub latency_ms: Option<u64>,
pub error: Option<String>,
pub content: Option<String>,
pub raw: Option<String>,
}
// ============================================================================
// 加密服务
// ============================================================================
@@ -153,6 +187,270 @@ impl ApiKeyProviderService {
}
}
pub async fn test_chat(
&self,
db: &DbConnection,
provider_id: &str,
model_name: Option<String>,
prompt: String,
) -> Result<ChatTestResult, String> {
use std::time::Instant;
let provider_with_keys = self
.get_provider(db, provider_id)?
.ok_or_else(|| format!("Provider not found: {}", provider_id))?;
let provider = &provider_with_keys.provider;
let api_key = self
.get_next_api_key(db, provider_id)?
.ok_or_else(|| "没有可用的 API Key".to_string())?;
let test_model = model_name.or_else(|| provider.custom_models.first().cloned());
let test_model =
test_model.ok_or_else(|| "缺少模型名称:请在自定义模型中填写一个模型名".to_string())?;
let start = Instant::now();
// Codex 协议直接走 /responses 端点
let result = if provider.provider_type == ApiProviderType::Codex {
self.test_codex_responses_endpoint(&api_key, &provider.api_host, &test_model, &prompt)
.await
} else {
self.test_openai_chat_once(&api_key, &provider.api_host, &test_model, &prompt)
.await
};
let latency_ms = start.elapsed().as_millis() as u64;
match result {
Ok((content, raw)) => Ok(ChatTestResult {
success: true,
latency_ms: Some(latency_ms),
error: None,
content: Some(content),
raw: Some(raw),
}),
Err(e) => Ok(ChatTestResult {
success: false,
latency_ms: Some(latency_ms),
error: Some(e),
content: None,
raw: None,
}),
}
}
async fn test_openai_chat_once(
&self,
api_key: &str,
api_host: &str,
model: &str,
prompt: &str,
) -> Result<(String, String), String> {
use crate::models::openai::{ChatCompletionRequest, ChatMessage, MessageContent};
use crate::providers::openai_custom::OpenAICustomProvider;
let provider =
OpenAICustomProvider::with_config(api_key.to_string(), Some(api_host.to_string()));
let request = ChatCompletionRequest {
model: model.to_string(),
messages: vec![ChatMessage {
role: "user".to_string(),
content: Some(MessageContent::Text(prompt.to_string())),
tool_calls: None,
tool_call_id: None,
reasoning_content: None,
}],
temperature: Some(0.2),
max_tokens: Some(64),
top_p: None,
stream: false,
tools: None,
tool_choice: None,
reasoning_effort: None,
};
let resp = provider
.call_api(&request)
.await
.map_err(|e| format!("API 调用失败: {}", e))?;
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
if status.is_success() {
let parsed: serde_json::Value = serde_json::from_str(&body)
.map_err(|e| format!("解析响应失败: {} - {}", e, body))?;
let content = parsed["choices"]
.as_array()
.and_then(|arr| arr.first())
.and_then(|c| c["message"]["content"].as_str())
.unwrap_or("")
.to_string();
return Ok((content, body));
}
// 部分上游(如某些 relay)强制要求 stream=true
if status.as_u16() == 400 && body.contains("Stream must be set to true") {
let mut request2 = request.clone();
request2.stream = true;
let resp2 = provider
.call_api(&request2)
.await
.map_err(|e| format!("API 调用失败: {}", e))?;
let status2 = resp2.status();
let body2 = resp2.text().await.unwrap_or_default();
if !status2.is_success() {
return Err(format!("API 返回错误: {} - {}", status2, body2));
}
let content = Self::parse_chat_completions_sse_content(&body2);
return Ok((content, body2));
}
// 部分上游(如 Codex relay)不支持 messages 参数,需要走 /responses 端点
if status.as_u16() == 400 && body.contains("Unsupported parameter: messages") {
return self
.test_codex_responses_endpoint(api_key, api_host, model, prompt)
.await;
}
Err(format!("API 返回错误: {} - {}", status, body))
}
fn parse_chat_completions_sse_content(body: &str) -> String {
let mut out = String::new();
for line in body.lines() {
let line = line.trim();
if !line.starts_with("data:") {
continue;
}
let data = line.trim_start_matches("data:").trim();
if data.is_empty() || data == "[DONE]" {
continue;
}
if let Ok(v) = serde_json::from_str::<serde_json::Value>(data) {
if let Some(s) = v["choices"][0]["delta"]["content"].as_str() {
out.push_str(s);
} else if let Some(s) = v["choices"][0]["message"]["content"].as_str() {
out.push_str(s);
}
}
}
out
}
fn build_codex_responses_request(model: &str, prompt: &str) -> serde_json::Value {
serde_json::json!({
"model": model,
"input": [
{
"role": "user",
"content": [
{
"type": "input_text",
"text": prompt
}
]
}
],
"stream": true,
"max_output_tokens": 64
})
}
/// 测试 Codex /responses 端点(用于不支持 messages 参数的上游)
async fn test_codex_responses_endpoint(
&self,
api_key: &str,
api_host: &str,
model: &str,
prompt: &str,
) -> Result<(String, String), String> {
// 构建 /responses 端点 URL
let base = api_host.trim_end_matches('/');
let url = if base.ends_with("/v1") {
format!("{}/responses", base)
} else if base.ends_with("/openai") {
format!("{}/v1/responses", base)
} else {
format!("{}/v1/responses", base)
};
// Codex Responses 格式请求体(input 必须是列表)
let request_body = Self::build_codex_responses_request(model, prompt);
let client = reqwest::Client::new();
let resp = client
.post(&url)
.header("Authorization", format!("Bearer {}", api_key))
.header("Content-Type", "application/json")
.json(&request_body)
.send()
.await
.map_err(|e| format!("API 调用失败: {}", e))?;
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
if !status.is_success() {
return Err(format!("API 返回错误: {} - {}", status, body));
}
// 解析 Codex SSE 响应
let content = Self::parse_codex_responses_sse_content(&body);
Ok((content, body))
}
fn parse_codex_responses_sse_content(body: &str) -> String {
let mut out = String::new();
for line in body.lines() {
let line = line.trim();
if !line.starts_with("data:") {
continue;
}
let data = line.trim_start_matches("data:").trim();
if data.is_empty() || data == "[DONE]" {
continue;
}
if let Ok(v) = serde_json::from_str::<serde_json::Value>(data) {
// Codex responses 格式: {"type": "response.output_text.delta", "delta": "..."}
if let Some(s) = v["delta"].as_str() {
out.push_str(s);
}
// 或者完整响应格式
if let Some(arr) = v["output"].as_array() {
for item in arr {
if item["type"].as_str() == Some("message") {
if let Some(content_arr) = item["content"].as_array() {
for c in content_arr {
if c["type"].as_str() == Some("output_text") {
if let Some(text) = c["text"].as_str() {
out.push_str(text);
}
}
}
}
}
}
}
}
}
out
}
// ==================== Provider 操作 ====================
/// 初始化系统 Provider
@@ -1329,17 +1627,50 @@ impl ApiKeyProviderService {
self.test_gemini_connection(&api_key, &provider.api_host)
.await
}
ApiProviderType::Codex => {
// Codex 协议直接走 /responses 端点
let test_model = model_name
.or_else(|| provider.custom_models.first().cloned())
.ok_or_else(|| "缺少模型名称:请在自定义模型中填写一个模型名".to_string())?;
self.test_codex_responses_endpoint(&api_key, &provider.api_host, &test_model, "hi")
.await
.map(|_| vec![test_model])
}
_ => {
// OpenAI 兼容类型,优先使用 /models 端点
eprintln!("[TEST_CONNECTION] model_name param: {:?}", model_name);
eprintln!(
"[TEST_CONNECTION] provider.custom_models: {:?}",
provider.custom_models
);
let models_result = self
.test_openai_models_endpoint(&api_key, &provider.api_host)
.await;
// 如果 /models 端点失败且有模型名称,尝试发送测试请求
if models_result.is_err() && model_name.is_some() {
let test_model = model_name.unwrap();
self.test_openai_chat_completion(&api_key, &provider.api_host, &test_model)
.await
eprintln!("[TEST_CONNECTION] models_result: {:?}", models_result);
// 如果 /models 端点失败:
// 1) 优先用传入的 model_name
// 2) 否则如果 Provider 配置了 custom_models,则用第一个模型降级测试 chat/completions
if models_result.is_err() {
let test_model = model_name.or_else(|| provider.custom_models.first().cloned());
eprintln!("[TEST_CONNECTION] fallback test_model: {:?}", test_model);
if let Some(test_model) = test_model {
let chat_result = self
.test_openai_chat_completion(&api_key, &provider.api_host, &test_model)
.await;
eprintln!(
"[TEST_CONNECTION] chat_completion result: {:?}",
chat_result
);
chat_result
} else {
models_result
}
} else {
models_result
}
@@ -1404,42 +1735,9 @@ impl ApiKeyProviderService {
api_host: &str,
model: &str,
) -> Result<Vec<String>, String> {
use crate::models::openai::{ChatCompletionRequest, ChatMessage, MessageContent};
use crate::providers::openai_custom::OpenAICustomProvider;
let provider =
OpenAICustomProvider::with_config(api_key.to_string(), Some(api_host.to_string()));
let request = ChatCompletionRequest {
model: model.to_string(),
messages: vec![ChatMessage {
role: "user".to_string(),
content: Some(MessageContent::Text("hi".to_string())),
tool_calls: None,
tool_call_id: None,
reasoning_content: None,
}],
temperature: None,
max_tokens: Some(1),
top_p: None,
stream: false,
tools: None,
tool_choice: None,
reasoning_effort: None,
};
let response = provider
.call_api(&request)
self.test_openai_chat_once(api_key, api_host, model, "hi")
.await
.map_err(|e| format!("API 调用失败: {}", e))?;
if response.status().is_success() {
Ok(vec![model.to_string()])
} else {
let status = response.status();
let body = response.text().await.unwrap_or_default();
Err(format!("API 返回错误: {} - {}", status, body))
}
.map(|_| vec![model.to_string()])
}
/// 测试 Claude Key 的客户端兼容性
@@ -24,7 +24,6 @@ import {
CollapsibleContent,
CollapsibleTrigger,
} from "@/components/ui/collapsible";
import { getAgentBackend, setAgentBackend, type AgentBackend } from "../config";
// --- Styled Components ---
@@ -143,14 +142,6 @@ interface ChatSettingsProps {
export const ChatSettings: React.FC<ChatSettingsProps> = ({ onClose }) => {
// Local state for UI toggles (Mocking functional settings)
const [fontSize, setFontSize] = useState([14]);
const [agentBackend, setBackend] = useState<AgentBackend>(getAgentBackend());
const handleBackendChange = (value: AgentBackend) => {
setBackend(value);
setAgentBackend(value);
// 提示用户需要刷新页面
window.location.reload();
};
return (
<SettingsContainer>
@@ -170,30 +161,6 @@ export const ChatSettings: React.FC<ChatSettingsProps> = ({ onClose }) => {
</Header>
<ScrollArea className="flex-1">
{/* Agent Backend Settings */}
<CollapsibleSection title="Agent 后端">
<SettingRow>
<div>
<div className="label">Agent 引擎</div>
<div className="desc">切换后需刷新页面</div>
</div>
<Select
value={agentBackend}
onValueChange={(v) => handleBackendChange(v as AgentBackend)}
>
<SelectTrigger className="w-[100px] h-7 text-xs">
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="native">Native</SelectItem>
<SelectItem value="aster">Aster</SelectItem>
</SelectContent>
</Select>
</SettingRow>
</CollapsibleSection>
<Separator />
{/* Message Settings */}
<CollapsibleSection title="消息设置">
<SettingRow>
+10 -11
View File
@@ -1,35 +1,34 @@
/**
* Agent 后端配置
*
* 用于切换 Native 和 Aster 后端
* 现在只使用 Aster 后端,保留配置接口以便未来扩展
*/
export type AgentBackend = "native" | "aster";
export type AgentBackend = "aster";
// 默认使用 Native 后端,可以通过 localStorage 切换
// 默认使用 Aster 后端
const STORAGE_KEY = "proxycast_agent_backend";
/**
* 获取当前 Agent 后端
* 现在固定返回 aster
*/
export function getAgentBackend(): AgentBackend {
const stored = localStorage.getItem(STORAGE_KEY);
if (stored === "aster" || stored === "native") {
return stored;
}
return "native"; // 默认
return "aster";
}
/**
* 设置 Agent 后端
* 保留接口但不再生效
*/
export function setAgentBackend(backend: AgentBackend): void {
localStorage.setItem(STORAGE_KEY, backend);
export function setAgentBackend(_backend: AgentBackend): void {
localStorage.setItem(STORAGE_KEY, "aster");
}
/**
* 是否使用 Aster 后端
* 现在固定返回 true
*/
export function useAsterBackend(): boolean {
return getAgentBackend() === "aster";
return true;
}
@@ -54,6 +54,7 @@ const PROVIDER_TYPES: { value: ProviderType; label: string }[] = [
const PROVIDER_TYPE_EXTRA_FIELDS: Record<ProviderType, string[]> = {
openai: [],
"openai-response": [],
codex: [],
anthropic: [],
"anthropic-compatible": [], // Anthropic 兼容格式,无需额外字段
gemini: [],
@@ -139,11 +139,27 @@ export const ApiKeyProviderSection = forwardRef<
}
// 如果 Provider 配置了自定义模型,使用第一个模型进行测试
const modelName =
let modelName =
provider.custom_models && provider.custom_models.length > 0
? provider.custom_models[0]
: undefined;
// 兜底:自定义模型可能还在防抖保存中(provider.custom_models 还未更新)
// 直接从输入框读取当前值,确保连接测试可用
if (!modelName) {
const input = document.getElementById(
"custom-models",
) as HTMLInputElement | null;
const raw = input?.value ?? "";
const parsed = raw
.split(",")
.map((m) => m.trim())
.filter((m) => m.length > 0);
if (parsed.length > 0) {
modelName = parsed[0];
}
}
// 调用后端连接测试 API
const result = await apiKeyProviderApi.testConnection(
providerId,
@@ -166,6 +182,53 @@ export const ApiKeyProviderSection = forwardRef<
[selectedProvider],
);
const handleTestChat = useCallback(
async (providerId: string, prompt: string) => {
const provider = selectedProvider;
if (!provider || provider.api_keys.length === 0) {
return {
success: false,
error: "没有可用的 API Key",
};
}
let modelName =
provider.custom_models && provider.custom_models.length > 0
? provider.custom_models[0]
: undefined;
if (!modelName) {
const input = document.getElementById(
"custom-models",
) as HTMLInputElement | null;
const raw = input?.value ?? "";
const parsed = raw
.split(",")
.map((m) => m.trim())
.filter((m) => m.length > 0);
if (parsed.length > 0) {
modelName = parsed[0];
}
}
try {
return await apiKeyProviderApi.testChat(providerId, modelName, prompt);
} catch (e) {
const msg =
e instanceof Error
? e.message
: typeof e === "string"
? e
: JSON.stringify(e);
return {
success: false,
error: msg || "对话测试失败",
};
}
},
[selectedProvider],
);
// ===== 删除 Provider =====
const handleDeleteProviderClick = useCallback(() => {
if (selectedProvider && !selectedProvider.is_system) {
@@ -209,6 +272,7 @@ export const ApiKeyProviderSection = forwardRef<
onDeleteApiKey={deleteApiKey}
onToggleApiKey={toggleApiKey}
onTestConnection={handleTestConnection}
onTestChat={handleTestChat}
onDeleteProvider={handleDeleteProviderClick}
loading={loading}
className="h-full"
@@ -50,6 +50,7 @@ const providerTypeArbitrary: fc.Arbitrary<ProviderType> = fc.constantFrom(
const EXPECTED_EXTRA_FIELDS: Record<ProviderType, string[]> = {
openai: [],
"openai-response": [],
codex: [],
anthropic: [],
"anthropic-compatible": [],
gemini: [],
@@ -35,6 +35,7 @@ const DEBOUNCE_DELAY = 500;
const PROVIDER_TYPES: { value: ProviderType; label: string }[] = [
{ value: "openai", label: "OpenAI 兼容" },
{ value: "openai-response", label: "OpenAI Responses API" },
{ value: "codex", label: "Codex CLI" },
{ value: "anthropic", label: "Anthropic" },
{ value: "anthropic-compatible", label: "Anthropic 兼容" },
{ value: "gemini", label: "Gemini" },
@@ -50,6 +51,7 @@ const PROVIDER_TYPES: { value: ProviderType; label: string }[] = [
const PROVIDER_TYPE_FIELDS: Record<ProviderType, string[]> = {
openai: [],
"openai-response": [],
codex: [],
anthropic: [],
"anthropic-compatible": [], // Anthropic 兼容格式,无需额外字段
gemini: [],
@@ -7,10 +7,19 @@
* **Validates: Requirements 4.1, 6.3, 6.4**
*/
import React from "react";
import React, { useState } from "react";
import { cn } from "@/lib/utils";
import { Switch } from "@/components/ui/switch";
import { Button } from "@/components/ui/button";
import {
Dialog,
DialogContent,
DialogHeader,
DialogTitle,
DialogDescription,
DialogFooter,
} from "@/components/ui/dialog";
import { Textarea } from "@/components/ui/textarea";
import { Trash2 } from "lucide-react";
import { ProviderIcon } from "@/icons/providers";
import { ApiKeyList } from "./ApiKeyList";
@@ -21,6 +30,7 @@ import {
} from "./ConnectionTestButton";
import { ProviderModelList } from "./ProviderModelList";
import type {
ChatTestResult,
ProviderWithKeysDisplay,
UpdateProviderRequest,
} from "@/lib/api/apiKeyProvider";
@@ -46,6 +56,8 @@ export interface ProviderSettingProps {
onToggleApiKey?: (keyId: string, enabled: boolean) => void;
/** 测试连接回调 */
onTestConnection?: (providerId: string) => Promise<ConnectionTestResult>;
/** 对话测试回调 */
onTestChat?: (providerId: string, prompt: string) => Promise<ChatTestResult>;
/** 删除自定义 Provider 回调 */
onDeleteProvider?: (providerId: string) => void;
/** 是否正在加载 */
@@ -86,10 +98,39 @@ export const ProviderSetting: React.FC<ProviderSettingProps> = ({
onDeleteApiKey,
onToggleApiKey,
onTestConnection,
onTestChat,
onDeleteProvider,
loading = false,
className,
}) => {
const [chatDialogOpen, setChatDialogOpen] = useState(false);
const [chatPrompt, setChatPrompt] = useState("hello");
const [chatTesting, setChatTesting] = useState(false);
const [chatResult, setChatResult] = useState<ChatTestResult | null>(null);
const handleChatTest = async () => {
if (!onTestChat || chatTesting || !provider) return;
setChatTesting(true);
setChatResult(null);
try {
const res = await onTestChat(provider.id, chatPrompt);
setChatResult(res);
} catch (e) {
const msg =
e instanceof Error
? e.message
: typeof e === "string"
? e
: JSON.stringify(e);
setChatResult({
success: false,
error: msg || "对话测试失败",
});
} finally {
setChatTesting(false);
}
};
// 空状态
if (!provider) {
return (
@@ -218,20 +259,98 @@ export const ProviderSetting: React.FC<ProviderSettingProps> = ({
{/* 连接测试 */}
<section data-testid="connection-test-section">
<h4 className="text-sm font-medium text-foreground mb-3">连接测试</h4>
<ConnectionTestButton
providerId={provider.id}
onTest={onTestConnection}
disabled={
loading ||
!provider.enabled ||
(provider.api_keys?.length ?? 0) === 0
}
/>
<div className="flex gap-2">
<ConnectionTestButton
providerId={provider.id}
onTest={onTestConnection}
disabled={
loading ||
!provider.enabled ||
(provider.api_keys?.length ?? 0) === 0
}
className="flex-1"
/>
<Button
variant="outline"
size="sm"
disabled={
loading ||
!provider.enabled ||
(provider.api_keys?.length ?? 0) === 0 ||
!onTestChat
}
onClick={() => setChatDialogOpen(true)}
>
对话测试
</Button>
</div>
{(provider.api_keys?.length ?? 0) === 0 && (
<p className="text-xs text-muted-foreground mt-2">
请先添加 API Key 后再进行连接测试
</p>
)}
<Dialog open={chatDialogOpen} onOpenChange={setChatDialogOpen}>
<DialogContent className="sm:max-w-[700px] p-6">
<DialogHeader className="mb-4">
<DialogTitle>对话测试</DialogTitle>
<DialogDescription>
发送一条最小对话请求,直接查看返回内容或原始错误,便于排查模型/权限/路由问题。
</DialogDescription>
</DialogHeader>
<div className="space-y-3">
<Textarea
value={chatPrompt}
onChange={(e) => setChatPrompt(e.target.value)}
className="h-[120px]"
/>
{chatResult?.error && (
<div className="p-2 rounded-md bg-red-50 border border-red-200 text-xs text-red-600">
<p className="font-medium">错误详情:</p>
<p className="mt-1 break-all">{chatResult.error}</p>
</div>
)}
{chatResult?.success && (
<div className="p-2 rounded-md bg-green-50 border border-green-200 text-xs text-green-700">
<p className="font-medium">
返回内容
{chatResult.latency_ms !== undefined
? ` (${chatResult.latency_ms}ms)`
: ""}
:
</p>
<p className="mt-1 whitespace-pre-wrap break-words">
{chatResult.content || ""}
</p>
</div>
)}
{chatResult?.raw && (
<Textarea
value={chatResult.raw}
readOnly
className="h-[180px] font-mono text-xs"
/>
)}
</div>
<DialogFooter className="mt-4">
<Button
variant="outline"
onClick={() => setChatDialogOpen(false)}
disabled={chatTesting}
>
关闭
</Button>
<Button
onClick={handleChatTest}
disabled={chatTesting || !chatPrompt.trim()}
>
{chatTesting ? "发送中..." : "发送"}
</Button>
</DialogFooter>
</DialogContent>
</Dialog>
</section>
{/* 分隔线 */}
@@ -85,6 +85,7 @@ const PROVIDER_TYPE_TO_REGISTRY_ID: Record<string, string> = {
"anthropic-compatible": "anthropic", // Anthropic 兼容格式
openai: "openai",
"openai-response": "openai",
codex: "codex",
gemini: "gemini",
};
@@ -100,6 +101,11 @@ export function mapProviderIdToRegistryId(
providerId: string,
providerType?: string,
): string {
// Codex 协议优先走 codex 模型资源
if (providerType === "codex") {
return "codex";
}
// 优先使用 Provider ID 映射
if (PROVIDER_ID_TO_REGISTRY_ID[providerId]) {
return PROVIDER_ID_TO_REGISTRY_ID[providerId];
@@ -225,20 +225,24 @@ export function useScreenshotChat(): UseScreenshotChatReturn {
? [{ data: imageBase64, mediaType: "image/png" }]
: [];
// 发送流式请求
await safeInvoke("native_agent_chat_stream", {
message,
eventName,
sessionId,
model: "claude-sonnet-4-5",
images:
images.length > 0
? images.map((img) => ({
data: img.data,
media_type: img.mediaType,
}))
: undefined,
provider: "claude",
// 发送流式请求(使用 Aster Agent)
await safeInvoke("aster_agent_chat_stream", {
request: {
message,
session_id: sessionId,
event_name: eventName,
images:
images.length > 0
? images.map((img) => ({
data: img.data,
media_type: img.mediaType,
}))
: undefined,
provider_config: {
provider_name: "anthropic",
model_name: "claude-sonnet-4-5",
},
},
});
} catch (err) {
console.error("发送消息失败:", err);
-10
View File
@@ -10,7 +10,6 @@ import {
RetryChangeEvent,
FullReloadEvent,
CredentialPoolChangeEvent,
NativeAgentChangeEvent,
} from "@/lib/configEventManager";
interface UseConfigEventsOptions {
@@ -25,7 +24,6 @@ interface UseConfigEventsOptions {
onLoggingChanged?: (event: LoggingChangeEvent) => void;
onRetryChanged?: (event: RetryChangeEvent) => void;
onCredentialPoolChanged?: (event: CredentialPoolChangeEvent) => void;
onNativeAgentChanged?: (event: NativeAgentChangeEvent) => void;
/** 通用事件回调(用于处理所有事件) */
onAnyChange?: (event: ConfigChangeEvent) => void;
}
@@ -83,7 +81,6 @@ export function useConfigEvents(
onLoggingChanged,
onRetryChanged,
onCredentialPoolChanged,
onNativeAgentChanged,
onAnyChange,
} = options;
@@ -106,7 +103,6 @@ export function useConfigEvents(
onLoggingChanged,
onRetryChanged,
onCredentialPoolChanged,
onNativeAgentChanged,
onAnyChange,
});
@@ -121,7 +117,6 @@ export function useConfigEvents(
onLoggingChanged,
onRetryChanged,
onCredentialPoolChanged,
onNativeAgentChanged,
onAnyChange,
};
}, [
@@ -133,7 +128,6 @@ export function useConfigEvents(
onLoggingChanged,
onRetryChanged,
onCredentialPoolChanged,
onNativeAgentChanged,
onAnyChange,
]);
@@ -170,9 +164,6 @@ export function useConfigEvents(
case "CredentialPoolChanged":
callbacksRef.current.onCredentialPoolChanged?.(event.data);
break;
case "NativeAgentChanged":
callbacksRef.current.onNativeAgentChanged?.(event.data);
break;
}
}, []);
@@ -251,5 +242,4 @@ export type {
RetryChangeEvent,
FullReloadEvent,
CredentialPoolChangeEvent,
NativeAgentChangeEvent,
};
+17 -9
View File
@@ -378,6 +378,8 @@ export async function sendAgentMessage(
* });
* await sendAgentMessageStream(message, eventName, sessionId, model, undefined, provider);
* ```
*
* @deprecated 请使用 sendAsterMessageStream 代替
*/
export async function sendAgentMessageStream(
message: string,
@@ -386,16 +388,22 @@ export async function sendAgentMessageStream(
model?: string,
images?: ImageInput[],
provider?: string,
terminalMode?: boolean,
_terminalMode?: boolean,
): Promise<void> {
return await safeInvoke("native_agent_chat_stream", {
message,
eventName,
sessionId,
model,
images,
provider,
terminalMode,
// 使用 Aster Agent 实现
return await safeInvoke("aster_agent_chat_stream", {
request: {
message,
session_id: sessionId || "default",
event_name: eventName,
images,
provider_config: provider
? {
provider_name: provider,
model_name: model || "claude-sonnet-4-20250514",
}
: undefined,
},
});
}
+24
View File
@@ -114,6 +114,14 @@ export interface ImportResult {
errors: string[];
}
export interface ChatTestResult {
success: boolean;
latency_ms?: number;
error?: string;
content?: string;
raw?: string;
}
// ============================================================================
// API 函数
// ============================================================================
@@ -300,7 +308,23 @@ export const apiKeyProviderApi = {
): Promise<ConnectionTestResult> {
return safeInvoke("test_api_key_provider_connection", {
providerId,
provider_id: providerId,
modelName,
model_name: modelName,
});
},
async testChat(
providerId: string,
modelName: string | undefined,
prompt: string,
): Promise<ChatTestResult> {
return safeInvoke("test_api_key_provider_chat", {
providerId,
provider_id: providerId,
modelName,
model_name: modelName,
prompt,
});
},
};
+1 -10
View File
@@ -86,14 +86,6 @@ export interface CredentialPoolChangeEvent {
source: ConfigChangeSource;
}
/** Native Agent 配置变更事件 */
export interface NativeAgentChangeEvent {
default_model: string;
temperature: number;
max_tokens: number;
source: ConfigChangeSource;
}
/** 配置变更事件联合类型 */
export type ConfigChangeEvent =
| { type: "FullReload"; data: FullReloadEvent }
@@ -104,8 +96,7 @@ export type ConfigChangeEvent =
| { type: "LoggingChanged"; data: LoggingChangeEvent }
| { type: "RetryChanged"; data: RetryChangeEvent }
| { type: "AmpConfigChanged"; data: AmpConfigChangeEvent }
| { type: "CredentialPoolChanged"; data: CredentialPoolChangeEvent }
| { type: "NativeAgentChanged"; data: NativeAgentChangeEvent };
| { type: "CredentialPoolChanged"; data: CredentialPoolChangeEvent };
/** 事件回调类型 */
type ConfigEventCallback = (event: ConfigChangeEvent) => void;
+20 -9
View File
@@ -138,16 +138,27 @@ const defaultMocks: Record<string, any> = {
agent_chat_stream: () => ({}),
agent_terminal_command_response: () => ({}),
agent_term_scrollback_response: () => ({}),
native_agent_chat_stream: () => ({}),
// aster Agent
aster_agent_init: () => ({ success: true }),
aster_agent_status: () => ({ initialized: false }),
aster_agent_reset: () => ({ success: true }),
aster_agent_create_session: () => ({ session_id: "mock-aster-session" }),
aster_agent_send_message: () => ({ message_id: "mock-message-id" }),
aster_agent_extend_system_prompt: () => ({ success: true }),
aster_agent_list_providers: () => [],
// Aster Agent
aster_agent_init: () => ({ initialized: true, provider_configured: false }),
aster_agent_status: () => ({
initialized: false,
provider_configured: false,
}),
aster_agent_configure_provider: () => ({
initialized: true,
provider_configured: true,
}),
aster_agent_configure_from_pool: () => ({
initialized: true,
provider_configured: true,
}),
aster_agent_chat_stream: () => ({}),
aster_agent_stop: () => true,
aster_session_create: () => "mock-aster-session",
aster_session_list: () => [],
aster_session_get: () => ({ id: "mock", messages: [] }),
aster_agent_confirm: () => ({}),
// 终端相关
create_terminal_session: () => ({ uuid: "mock-terminal-uuid" }),
+1
View File
@@ -18,6 +18,7 @@
export type ProviderType =
| "openai" // 标准 OpenAI Chat Completions API
| "openai-response" // OpenAI Responses API (支持 Reasoning)
| "codex" // Codex CLI 协议 (使用 /responses 端点)
| "anthropic" // Anthropic Messages API
| "anthropic-compatible" // Anthropic 兼容格式 (system 为数组格式)
| "gemini" // Google Gemini API