mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
chore: bump version to 0.48.4
This commit is contained in:
@@ -53,3 +53,4 @@ src-tauri/gen
|
||||
.task
|
||||
Taskfile.yml
|
||||
nul
|
||||
.proptest-regressions
|
||||
@@ -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) | 工具库 |
|
||||
|
||||
## 构建命令
|
||||
|
||||
|
||||
@@ -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
@@ -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` - 内置终端
|
||||
|
||||
## 更新提醒
|
||||
|
||||
任何文件变更后,请更新此文档和相关的上级文档。
|
||||
|
||||
@@ -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
|
||||
```
|
||||
|
||||
## 更新提醒
|
||||
|
||||
任何文件变更后,请更新此文档和相关的上级文档。
|
||||
@@ -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) - 凭证池管理
|
||||
@@ -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
|
||||
@@ -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) - 工具库
|
||||
@@ -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) - 流式处理
|
||||
@@ -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) - 数据库层
|
||||
@@ -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) - 凭证池管理
|
||||
@@ -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) - 前端组件
|
||||
@@ -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 命令
|
||||
@@ -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) - 组件系统
|
||||
@@ -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 命令
|
||||
@@ -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) - 数据库层
|
||||
@@ -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) - 业务服务
|
||||
@@ -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 服务器
|
||||
@@ -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 系统
|
||||
@@ -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) - 凭证池管理
|
||||
@@ -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) - 组件系统
|
||||
@@ -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**:适用于"每次都必须成功"的场景
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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::
|
||||
```
|
||||
@@ -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::
|
||||
```
|
||||
@@ -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::
|
||||
```
|
||||
@@ -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
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"name": "proxycast",
|
||||
"private": true,
|
||||
"version": "0.48.3",
|
||||
"version": "0.48.4",
|
||||
"type": "module",
|
||||
"repository": {
|
||||
"type": "git",
|
||||
|
||||
Generated
+4
-4
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
|
||||
@@ -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 类型设置对应的环境变量
|
||||
|
||||
@@ -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("没有可用凭证"));
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -1,9 +0,0 @@
|
||||
//! SSE 流解析器模块
|
||||
//!
|
||||
//! 提供不同协议的 SSE 流解析器
|
||||
|
||||
mod anthropic_sse;
|
||||
mod openai_sse;
|
||||
|
||||
pub use anthropic_sse::{AnthropicParseResult, AnthropicSSEParser};
|
||||
pub use openai_sse::OpenAISSEParser;
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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 包含所有可用工具定义
|
||||
|
||||
## 更新提醒
|
||||
|
||||
任何文件变更后,请更新此文档和相关的上级文档。
|
||||
@@ -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"));
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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('&', "&")
|
||||
.replace('<', "<")
|
||||
.replace('>', ">")
|
||||
.replace('"', """)
|
||||
.replace('\'', "'")
|
||||
}
|
||||
|
||||
/// 从 ToolRegistry 生成 System Prompt 的便捷函数
|
||||
pub fn generate_tools_prompt(tools: &[ToolDefinition], format: PromptFormat) -> String {
|
||||
ToolPromptGenerator::new()
|
||||
.with_format(format)
|
||||
.generate_system_prompt(tools)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::super::types::PropertySchema;
|
||||
use super::*;
|
||||
|
||||
fn create_test_tool() -> ToolDefinition {
|
||||
ToolDefinition::new("bash", "Execute a bash command in the shell").with_parameters(
|
||||
JsonSchema::new()
|
||||
.add_property(
|
||||
"command",
|
||||
PropertySchema::string("The bash command to execute"),
|
||||
true,
|
||||
)
|
||||
.add_property(
|
||||
"timeout",
|
||||
PropertySchema::integer("Optional timeout in seconds")
|
||||
.with_default(serde_json::json!(120)),
|
||||
false,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
fn create_test_tools() -> Vec<ToolDefinition> {
|
||||
vec![
|
||||
create_test_tool(),
|
||||
ToolDefinition::new("read_file", "Read the contents of a file").with_parameters(
|
||||
JsonSchema::new()
|
||||
.add_property(
|
||||
"path",
|
||||
PropertySchema::string("The file path to read"),
|
||||
true,
|
||||
)
|
||||
.add_property(
|
||||
"start_line",
|
||||
PropertySchema::integer("Starting line number (1-based)"),
|
||||
false,
|
||||
)
|
||||
.add_property(
|
||||
"end_line",
|
||||
PropertySchema::integer("Ending line number (inclusive)"),
|
||||
false,
|
||||
),
|
||||
),
|
||||
]
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_generator_default_format() {
|
||||
let generator = ToolPromptGenerator::new();
|
||||
assert_eq!(generator.format(), PromptFormat::Xml);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_generator_with_format() {
|
||||
let generator = ToolPromptGenerator::new().with_format(PromptFormat::Json);
|
||||
assert_eq!(generator.format(), PromptFormat::Json);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_to_xml() {
|
||||
let generator = ToolPromptGenerator::new();
|
||||
let tool = create_test_tool();
|
||||
let xml = generator.tool_to_xml(&tool);
|
||||
|
||||
// 验证 XML 包含工具名称
|
||||
assert!(xml.contains("name=\"bash\""));
|
||||
// 验证 XML 包含描述
|
||||
assert!(xml.contains("Execute a bash command"));
|
||||
// 验证 XML 包含必需参数
|
||||
assert!(xml.contains("name=\"command\""));
|
||||
assert!(xml.contains("required=\"true\""));
|
||||
// 验证 XML 包含可选参数
|
||||
assert!(xml.contains("name=\"timeout\""));
|
||||
// 验证 XML 包含默认值
|
||||
assert!(xml.contains("<default>120</default>"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_to_json() {
|
||||
let generator = ToolPromptGenerator::new();
|
||||
let tool = create_test_tool();
|
||||
let json = generator.tool_to_json(&tool);
|
||||
|
||||
// 验证 JSON 包含工具名称
|
||||
assert!(json.contains("\"name\": \"bash\""));
|
||||
// 验证 JSON 包含描述
|
||||
assert!(json.contains("Execute a bash command"));
|
||||
// 验证 JSON 包含参数
|
||||
assert!(json.contains("\"command\""));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_generate_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>"), "<script>");
|
||||
assert_eq!(escape_xml("a & b"), "a & b");
|
||||
assert_eq!(escape_xml("\"quoted\""), ""quoted"");
|
||||
assert_eq!(escape_xml("it's"), "it's");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_tools() {
|
||||
let generator = ToolPromptGenerator::new();
|
||||
let prompt = generator.generate_system_prompt(&[]);
|
||||
|
||||
// 即使没有工具,也应该包含使用说明
|
||||
assert!(prompt.contains("你是一个友好的 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
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 工具系统
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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,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
@@ -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,
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
});
|
||||
},
|
||||
};
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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" }),
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user