From 2a7ed06d21f6c4d3912044df9e640462ed682fa6 Mon Sep 17 00:00:00 2001 From: coso Date: Fri, 30 Jan 2026 01:31:09 +0800 Subject: [PATCH] chore: bump version to 0.48.4 --- .gitignore | 1 + AGENTS.md | 24 + IMPLEMENTATION_PLAN.md | 285 ----- docs/README.md | 29 +- docs/aiprompts/README.md | 55 + docs/aiprompts/aster-integration.md | 107 ++ docs/aiprompts/commands.md | 111 ++ docs/aiprompts/components.md | 114 ++ docs/aiprompts/converter.md | 251 ++++ docs/aiprompts/credential-pool.md | 222 ++++ docs/aiprompts/database.md | 106 ++ docs/aiprompts/flow-monitor.md | 261 ++++ docs/aiprompts/hooks.md | 116 ++ docs/aiprompts/lib.md | 89 ++ docs/aiprompts/mcp.md | 94 ++ docs/aiprompts/overview.md | 138 ++ docs/aiprompts/plugins.md | 88 ++ docs/aiprompts/providers.md | 189 +++ docs/aiprompts/server.md | 252 ++++ docs/aiprompts/services.md | 107 ++ docs/aiprompts/terminal.md | 101 ++ docs/test/README.md | 121 ++ docs/test/agent-evaluation.md | 272 ++++ docs/test/e2e-tests.md | 264 ++++ docs/test/integration-tests.md | 229 ++++ docs/test/test-cases/agent-tests.md | 350 +++++ docs/test/test-cases/converter-tests.md | 260 ++++ docs/test/test-cases/provider-tests.md | 283 +++++ docs/test/unit-tests.md | 218 ++++ package.json | 2 +- src-tauri/Cargo.lock | 8 +- src-tauri/Cargo.toml | 2 +- .../crates/core/src/models/provider_type.rs | 18 +- .../agent/native_agent.txt | 7 - .../proptest-regressions/agent/tools/bash.txt | 7 - .../agent/tools/read_file.txt | 8 - .../proptest-regressions/config/tests.txt | 5 +- .../flow_monitor/file_store.txt | 8 - .../flow_monitor/memory_store.txt | 7 - .../flow_monitor/monitor.txt | 7 - .../flow_monitor/stream_rebuilder.txt | 8 - .../proptest-regressions/middleware/tests.txt | 7 - .../proptest-regressions/providers/tests.txt | 7 - .../proptest-regressions/proxy/tests.txt | 7 - .../proptest-regressions/router/tests.txt | 7 - .../proptest-regressions/websocket/tests.txt | 8 - src-tauri/src/agent/README.md | 127 +- src-tauri/src/agent/aster_agent.rs | 114 +- src-tauri/src/agent/aster_state.rs | 106 ++ src-tauri/src/agent/credential_bridge.rs | 428 +++++++ src-tauri/src/agent/mod.rs | 28 +- src-tauri/src/agent/native_agent.rs | 1047 --------------- src-tauri/src/agent/parsers/anthropic_sse.rs | 209 --- src-tauri/src/agent/parsers/mod.rs | 9 - src-tauri/src/agent/parsers/openai_sse.rs | 318 ----- src-tauri/src/agent/protocols/anthropic.rs | 668 ---------- src-tauri/src/agent/protocols/mod.rs | 76 -- src-tauri/src/agent/protocols/openai.rs | 564 -------- src-tauri/src/agent/tool_loop.rs | 1130 ----------------- src-tauri/src/agent/tools/README.md | 284 ----- src-tauri/src/agent/tools/browser.rs | 468 ------- src-tauri/src/agent/tools/mod.rs | 76 -- src-tauri/src/agent/tools/prompt.rs | 643 ---------- src-tauri/src/agent/tools/security.rs | 594 --------- src-tauri/src/agent/tools/term_scrollback.rs | 346 ----- src-tauri/src/agent/tools/terminal.rs | 499 -------- src-tauri/src/agent/tools/test_integration.rs | 59 - src-tauri/src/app/bootstrap.rs | 4 - src-tauri/src/app/runner.rs | 15 +- src-tauri/src/app/setup.rs | 8 +- src-tauri/src/commands/agent_cmd.rs | 292 +---- .../src/commands/api_key_provider_cmd.rs | 16 +- src-tauri/src/commands/aster_agent_cmd.rs | 98 +- src-tauri/src/commands/mod.rs | 1 - src-tauri/src/commands/native_agent_cmd.rs | 510 -------- .../src/database/dao/api_key_provider.rs | 4 + src-tauri/src/providers/openai_custom.rs | 197 ++- src-tauri/src/services/README.md | 9 + .../src/services/api_key_provider_service.rs | 378 +++++- .../agent/chat/components/ChatSettings.tsx | 33 - src/components/agent/chat/config.ts | 21 +- .../api-key/AddCustomProviderModal.tsx | 1 + .../api-key/ApiKeyProviderSection.tsx | 66 +- .../api-key/ProviderConfigForm.test.ts | 1 + .../api-key/ProviderConfigForm.tsx | 2 + .../provider-pool/api-key/ProviderSetting.tsx | 139 +- .../api-key/providerTypeMapping.ts | 6 + .../screenshot-chat/useScreenshotChat.ts | 32 +- src/hooks/useConfigEvents.ts | 10 - src/lib/api/agent.ts | 26 +- src/lib/api/apiKeyProvider.ts | 24 + src/lib/configEventManager.ts | 11 +- src/lib/tauri-mock/core.ts | 29 +- src/lib/types/provider.ts | 1 + 94 files changed, 6126 insertions(+), 8461 deletions(-) delete mode 100644 IMPLEMENTATION_PLAN.md create mode 100644 docs/aiprompts/README.md create mode 100644 docs/aiprompts/aster-integration.md create mode 100644 docs/aiprompts/commands.md create mode 100644 docs/aiprompts/components.md create mode 100644 docs/aiprompts/converter.md create mode 100644 docs/aiprompts/credential-pool.md create mode 100644 docs/aiprompts/database.md create mode 100644 docs/aiprompts/flow-monitor.md create mode 100644 docs/aiprompts/hooks.md create mode 100644 docs/aiprompts/lib.md create mode 100644 docs/aiprompts/mcp.md create mode 100644 docs/aiprompts/overview.md create mode 100644 docs/aiprompts/plugins.md create mode 100644 docs/aiprompts/providers.md create mode 100644 docs/aiprompts/server.md create mode 100644 docs/aiprompts/services.md create mode 100644 docs/aiprompts/terminal.md create mode 100644 docs/test/README.md create mode 100644 docs/test/agent-evaluation.md create mode 100644 docs/test/e2e-tests.md create mode 100644 docs/test/integration-tests.md create mode 100644 docs/test/test-cases/agent-tests.md create mode 100644 docs/test/test-cases/converter-tests.md create mode 100644 docs/test/test-cases/provider-tests.md create mode 100644 docs/test/unit-tests.md delete mode 100644 src-tauri/proptest-regressions/agent/native_agent.txt delete mode 100644 src-tauri/proptest-regressions/agent/tools/bash.txt delete mode 100644 src-tauri/proptest-regressions/agent/tools/read_file.txt delete mode 100644 src-tauri/proptest-regressions/flow_monitor/file_store.txt delete mode 100644 src-tauri/proptest-regressions/flow_monitor/memory_store.txt delete mode 100644 src-tauri/proptest-regressions/flow_monitor/monitor.txt delete mode 100644 src-tauri/proptest-regressions/flow_monitor/stream_rebuilder.txt delete mode 100644 src-tauri/proptest-regressions/middleware/tests.txt delete mode 100644 src-tauri/proptest-regressions/providers/tests.txt delete mode 100644 src-tauri/proptest-regressions/proxy/tests.txt delete mode 100644 src-tauri/proptest-regressions/router/tests.txt delete mode 100644 src-tauri/proptest-regressions/websocket/tests.txt create mode 100644 src-tauri/src/agent/credential_bridge.rs delete mode 100644 src-tauri/src/agent/native_agent.rs delete mode 100644 src-tauri/src/agent/parsers/anthropic_sse.rs delete mode 100644 src-tauri/src/agent/parsers/mod.rs delete mode 100644 src-tauri/src/agent/parsers/openai_sse.rs delete mode 100644 src-tauri/src/agent/protocols/anthropic.rs delete mode 100644 src-tauri/src/agent/protocols/mod.rs delete mode 100644 src-tauri/src/agent/protocols/openai.rs delete mode 100644 src-tauri/src/agent/tool_loop.rs delete mode 100644 src-tauri/src/agent/tools/README.md delete mode 100644 src-tauri/src/agent/tools/browser.rs delete mode 100644 src-tauri/src/agent/tools/mod.rs delete mode 100644 src-tauri/src/agent/tools/prompt.rs delete mode 100644 src-tauri/src/agent/tools/security.rs delete mode 100644 src-tauri/src/agent/tools/term_scrollback.rs delete mode 100644 src-tauri/src/agent/tools/terminal.rs delete mode 100644 src-tauri/src/agent/tools/test_integration.rs delete mode 100644 src-tauri/src/commands/native_agent_cmd.rs diff --git a/.gitignore b/.gitignore index 30494fb1f..71739a661 100644 --- a/.gitignore +++ b/.gitignore @@ -53,3 +53,4 @@ src-tauri/gen .task Taskfile.yml nul +.proptest-regressions \ No newline at end of file diff --git a/AGENTS.md b/AGENTS.md index c88670dac..9579fc45c 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -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) | 工具库 | ## 构建命令 diff --git a/IMPLEMENTATION_PLAN.md b/IMPLEMENTATION_PLAN.md deleted file mode 100644 index c4e7bd14d..000000000 --- a/IMPLEMENTATION_PLAN.md +++ /dev/null @@ -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 流程实现 diff --git a/docs/README.md b/docs/README.md index a2e4a6148..769c15a70 100644 --- a/docs/README.md +++ b/docs/README.md @@ -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` - 内置终端 + ## 更新提醒 任何文件变更后,请更新此文档和相关的上级文档。 diff --git a/docs/aiprompts/README.md b/docs/aiprompts/README.md new file mode 100644 index 000000000..8f0eed56d --- /dev/null +++ b/docs/aiprompts/README.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 +``` + +## 更新提醒 + +任何文件变更后,请更新此文档和相关的上级文档。 diff --git a/docs/aiprompts/aster-integration.md b/docs/aiprompts/aster-integration.md new file mode 100644 index 000000000..2ba715dc0 --- /dev/null +++ b/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) - 凭证池管理 diff --git a/docs/aiprompts/commands.md b/docs/aiprompts/commands.md new file mode 100644 index 000000000..dfc8583d1 --- /dev/null +++ b/docs/aiprompts/commands.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; + +#[tauri::command] +async fn remove_credential(id: String) -> Result<(), String>; + +#[tauri::command] +async fn list_credentials() -> Result, String>; + +#[tauri::command] +async fn refresh_credential(id: String) -> Result<(), String>; + +#[tauri::command] +async fn get_credential_status(id: String) -> Result; +``` + +### 服务器控制 + +```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; + +#[tauri::command] +async fn update_server_config(config: ServerConfig) -> Result<(), String>; +``` + +### 流量监控 + +```rust +#[tauri::command] +async fn get_flow_records(query: FlowQuery) -> Result, String>; + +#[tauri::command] +async fn get_flow_stats(time_range: TimeRange) -> Result; + +#[tauri::command] +async fn clear_flow_records(before: Option) -> Result; +``` + +## 前端调用 + +```typescript +import { invoke } from '@tauri-apps/api/core'; + +// 添加凭证 +const credential = await invoke('add_credential', { + provider: 'kiro', + filePath: '/path/to/credential.json', +}); + +// 获取服务器状态 +const status = await invoke('get_server_status'); + +// 查询流量记录 +const records = await invoke>('get_flow_records', { + query: { page: 1, pageSize: 20 }, +}); +``` + +## 错误处理 + +```rust +// 命令返回 Result +// 错误信息会传递到前端 + +#[tauri::command] +async fn example_command() -> Result { + do_something() + .await + .map_err(|e| e.to_string()) +} +``` + +## 相关文档 + +- [services.md](services.md) - 业务服务 +- [hooks.md](hooks.md) - 前端 Hooks diff --git a/docs/aiprompts/components.md b/docs/aiprompts/components.md new file mode 100644 index 000000000..c0fc1beb0 --- /dev/null +++ b/docs/aiprompts/components.md @@ -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 ( + + ); +} +``` + +### ProviderPool + +凭证池管理组件。 + +```tsx +// src/components/provider-pool/ProviderPoolPanel.tsx +export function ProviderPoolPanel() { + const { credentials, addCredential, removeCredential } = useProviderPool(); + + return ( +
+ + +
+ ); +} +``` + +### FlowMonitor + +流量监控组件。 + +```tsx +// src/components/flow-monitor/FlowMonitorPanel.tsx +export function FlowMonitorPanel() { + const { records, stats, query } = useFlowMonitor(); + + return ( +
+ + + +
+ ); +} +``` + +## 组件规范 + +### 文件命名 + +- 组件文件: `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 ( +
+ {/* JSX */} +
+ ); +} +``` + +## 相关文档 + +- [hooks.md](hooks.md) - React Hooks +- [lib.md](lib.md) - 工具库 diff --git a/docs/aiprompts/converter.md b/docs/aiprompts/converter.md new file mode 100644 index 000000000..8c541aef6 --- /dev/null +++ b/docs/aiprompts/converter.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 { + // 解析 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) - 流式处理 diff --git a/docs/aiprompts/credential-pool.md b/docs/aiprompts/credential-pool.md new file mode 100644 index 000000000..d62c92c92 --- /dev/null +++ b/docs/aiprompts/credential-pool.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, + 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) { + 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, +} + +struct CachedToken { + access_token: String, + expires_at: i64, + refresh_token: String, +} + +impl TokenCacheService { + pub async fn get_or_refresh(&self, credential_id: &str) -> Result { + 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>; + +#[tauri::command] +async fn refresh_credential(id: String) -> Result<()>; + +#[tauri::command] +async fn get_pool_status() -> Result; +``` + +## 相关文档 + +- [providers.md](providers.md) - Provider 系统 +- [services.md](services.md) - 业务服务 +- [database.md](database.md) - 数据库层 diff --git a/docs/aiprompts/database.md b/docs/aiprompts/database.md new file mode 100644 index 000000000..1adc44be0 --- /dev/null +++ b/docs/aiprompts/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>, +} + +impl CredentialDao { + pub fn insert(&self, credential: &Credential) -> Result<()>; + pub fn find_by_id(&self, id: &str) -> Result>; + pub fn find_all(&self) -> Result>; + 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) - 凭证池管理 diff --git a/docs/aiprompts/flow-monitor.md b/docs/aiprompts/flow-monitor.md new file mode 100644 index 000000000..08a136cc9 --- /dev/null +++ b/docs/aiprompts/flow-monitor.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, // 响应数据 + pub status: FlowStatus, // 状态 + pub latency_ms: Option, // 延迟 (毫秒) +} + +pub struct FlowRequest { + pub messages: Vec, // 消息列表 + pub tools: Option>, // 工具定义 + pub stream: bool, // 是否流式 +} + +pub struct FlowResponse { + pub content: String, // 响应内容 + pub tool_calls: Option>, // 工具调用 + 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, + event_sender: mpsc::Sender, +} + +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, + pub model: Option, + pub start_time: Option, + pub end_time: Option, + pub status: Option, + pub page: u32, + pub page_size: u32, +} + +pub async fn query_flows(query: FlowQuery) -> Result> { + // 构建 SQL 查询 + // 执行分页查询 + // 返回结果 +} +``` + +### 统计查询 + +```rust +pub struct FlowStats { + pub total_requests: u64, + pub total_tokens: u64, + pub avg_latency_ms: f64, + pub by_provider: HashMap, + pub by_model: HashMap, +} + +pub async fn get_stats(time_range: TimeRange) -> Result { + // 聚合统计 +} +``` + +## 前端事件 + +### 事件类型 + +```typescript +interface FlowEvent { + type: 'request_started' | 'response_received' | 'error'; + data: FlowRecord; +} +``` + +### 事件监听 + +```typescript +// 前端监听 +import { listen } from '@tauri-apps/api/event'; + +listen('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>; + +#[tauri::command] +async fn get_flow_stats(time_range: TimeRange) -> Result; + +#[tauri::command] +async fn get_flow_detail(id: String) -> Result; + +#[tauri::command] +async fn clear_flow_records(before: Option) -> Result; + +#[tauri::command] +async fn export_flow_records(format: ExportFormat) -> Result; +``` + +## 相关文档 + +- [server.md](server.md) - HTTP 服务器 +- [database.md](database.md) - 数据库层 +- [components.md](components.md) - 前端组件 diff --git a/docs/aiprompts/hooks.md b/docs/aiprompts/hooks.md new file mode 100644 index 000000000..38d22843b --- /dev/null +++ b/docs/aiprompts/hooks.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([]); + const [loading, setLoading] = useState(false); + + const refresh = async () => { + setLoading(true); + const list = await invoke('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([]); + + useEffect(() => { + const unlisten = listen('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('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 命令 diff --git a/docs/aiprompts/lib.md b/docs/aiprompts/lib.md new file mode 100644 index 000000000..f336a7a30 --- /dev/null +++ b/docs/aiprompts/lib.md @@ -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('add_credential', { provider, filePath: path }); +} + +export async function listCredentials() { + return invoke('list_credentials'); +} +``` + +### 流量事件管理 + +```typescript +// src/lib/flowEventManager.ts +class FlowEventManager { + private listeners: Map> = 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) - 组件系统 diff --git a/docs/aiprompts/mcp.md b/docs/aiprompts/mcp.md new file mode 100644 index 000000000..4036e8b01 --- /dev/null +++ b/docs/aiprompts/mcp.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, + config_path: PathBuf, +} + +pub struct McpServer { + name: String, + command: String, + args: Vec, + env: HashMap, + status: ServerStatus, + tools: Vec, +} + +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>; + + /// 调用工具 + pub async fn call_tool( + &self, + server: &str, + tool: &str, + args: Value, + ) -> Result; +} +``` + +## 配置格式 + +```json +{ + "mcpServers": { + "filesystem": { + "command": "npx", + "args": ["-y", "@modelcontextprotocol/server-filesystem", "/path"], + "env": {}, + "disabled": false + } + } +} +``` + +## Tauri 命令 + +```rust +#[tauri::command] +async fn mcp_list_servers() -> Result, 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, String>; + +#[tauri::command] +async fn mcp_call_tool(server: String, tool: String, args: Value) -> Result; +``` + +## 相关文档 + +- [services.md](services.md) - 业务服务 +- [commands.md](commands.md) - Tauri 命令 diff --git a/docs/aiprompts/overview.md b/docs/aiprompts/overview.md new file mode 100644 index 000000000..b6a355a01 --- /dev/null +++ b/docs/aiprompts/overview.md @@ -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) - 数据库层 diff --git a/docs/aiprompts/plugins.md b/docs/aiprompts/plugins.md new file mode 100644 index 000000000..80b78f537 --- /dev/null +++ b/docs/aiprompts/plugins.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; + set(key: string, value: any): Promise; + }; + + // 网络 + http: { + fetch(url: string, options?: RequestInit): Promise; + }; + + // UI + ui: { + showNotification(message: string): void; + showDialog(options: DialogOptions): Promise; + }; +} +``` + +## 相关文档 + +- [components.md](components.md) - 组件系统 +- [services.md](services.md) - 业务服务 diff --git a/docs/aiprompts/providers.md b/docs/aiprompts/providers.md new file mode 100644 index 000000000..5912b2fa9 --- /dev/null +++ b/docs/aiprompts/providers.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; + + /// 刷新 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; +} +``` + +## OAuth Provider 实现 + +### Kiro Provider + +```rust +// 凭证文件结构 +struct KiroCredential { + access_token: String, + refresh_token: String, + expires_at: i64, + client_id: Option, // 从 clientIdHash 合并 + client_secret: Option, // 从 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, // 自定义端点 +} + +// 请求头 +Authorization: Bearer {api_key} +``` + +### Claude Custom + +```rust +// 凭证结构 +struct ClaudeCredential { + api_key: String, + base_url: Option, +} + +// 请求头 +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 服务器 diff --git a/docs/aiprompts/server.md b/docs/aiprompts/server.md new file mode 100644 index 000000000..116af646e --- /dev/null +++ b/docs/aiprompts/server.md @@ -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>, +) -> 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, +} + +// 示例 +{ + "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 系统 diff --git a/docs/aiprompts/services.md b/docs/aiprompts/services.md new file mode 100644 index 000000000..203e8a590 --- /dev/null +++ b/docs/aiprompts/services.md @@ -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, + health_checker: HealthChecker, +} + +impl ProviderPoolService { + /// 获取下一个可用凭证 + pub async fn next_credential(&self, provider: ProviderType) -> Option; + + /// 添加凭证到池 + 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, + db: Arc, +} + +impl TokenCacheService { + /// 获取或刷新 Token + pub async fn get_or_refresh(&self, credential_id: &str) -> Result; + + /// 使 Token 失效 + pub async fn invalidate(&self, credential_id: &str); +} +``` + +### McpService + +```rust +pub struct McpService { + servers: HashMap, +} + +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>; +} +``` + +## 服务注入 + +```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>, + // ... +) -> Result<(), String> { + pool.add_credential(credential).await +} +``` + +## 相关文档 + +- [commands.md](commands.md) - Tauri 命令 +- [credential-pool.md](credential-pool.md) - 凭证池管理 diff --git a/docs/aiprompts/terminal.md b/docs/aiprompts/terminal.md new file mode 100644 index 000000000..6bd789902 --- /dev/null +++ b/docs/aiprompts/terminal.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, +} + +pub struct PtySession { + id: String, + master: PtyMaster, + child: Child, +} + +impl PtyManager { + /// 创建新会话 + pub fn create_session(&mut self, shell: &str) -> Result; + + /// 写入数据 + pub fn write(&self, session_id: &str, data: &[u8]) -> Result<()>; + + /// 读取输出 + pub fn read(&self, session_id: &str) -> Result>; + + /// 调整大小 + 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(null); + const xtermRef = useRef(); + + 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
; +} +``` + +## Tauri 命令 + +```rust +#[tauri::command] +async fn terminal_create(shell: Option) -> Result; + +#[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) - 组件系统 diff --git a/docs/test/README.md b/docs/test/README.md new file mode 100644 index 000000000..006bc0cbc --- /dev/null +++ b/docs/test/README.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**:适用于"每次都必须成功"的场景 diff --git a/docs/test/agent-evaluation.md b/docs/test/agent-evaluation.md new file mode 100644 index 000000000..d33d97ec1 --- /dev/null +++ b/docs/test/agent-evaluation.md @@ -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) diff --git a/docs/test/e2e-tests.md b/docs/test/e2e-tests.md new file mode 100644 index 000000000..4293095a3 --- /dev/null +++ b/docs/test/e2e-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) diff --git a/docs/test/integration-tests.md b/docs/test/integration-tests.md new file mode 100644 index 000000000..ea658e47f --- /dev/null +++ b/docs/test/integration-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) diff --git a/docs/test/test-cases/agent-tests.md b/docs/test/test-cases/agent-tests.md new file mode 100644 index 000000000..40b3b01d2 --- /dev/null +++ b/docs/test/test-cases/agent-tests.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:: +``` diff --git a/docs/test/test-cases/converter-tests.md b/docs/test/test-cases/converter-tests.md new file mode 100644 index 000000000..55c3640c6 --- /dev/null +++ b/docs/test/test-cases/converter-tests.md @@ -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:: +``` diff --git a/docs/test/test-cases/provider-tests.md b/docs/test/test-cases/provider-tests.md new file mode 100644 index 000000000..4b9303f3f --- /dev/null +++ b/docs/test/test-cases/provider-tests.md @@ -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:: +``` diff --git a/docs/test/unit-tests.md b/docs/test/unit-tests.md new file mode 100644 index 000000000..f1eb22373 --- /dev/null +++ b/docs/test/unit-tests.md @@ -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 = 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) diff --git a/package.json b/package.json index 31f689765..cb9e35c4c 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.48.3", + "version": "0.48.4", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index e52466280..4f16a0c33 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -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", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 304000ae7..38f679359 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -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" diff --git a/src-tauri/crates/core/src/models/provider_type.rs b/src-tauri/crates/core/src/models/provider_type.rs index 55254c482..cf8f8702f 100644 --- a/src-tauri/crates/core/src/models/provider_type.rs +++ b/src-tauri/crates/core/src/models/provider_type.rs @@ -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}")), } } } diff --git a/src-tauri/proptest-regressions/agent/native_agent.txt b/src-tauri/proptest-regressions/agent/native_agent.txt deleted file mode 100644 index f1875cabc..000000000 --- a/src-tauri/proptest-regressions/agent/native_agent.txt +++ /dev/null @@ -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" diff --git a/src-tauri/proptest-regressions/agent/tools/bash.txt b/src-tauri/proptest-regressions/agent/tools/bash.txt deleted file mode 100644 index e89b77e13..000000000 --- a/src-tauri/proptest-regressions/agent/tools/bash.txt +++ /dev/null @@ -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 = "-" diff --git a/src-tauri/proptest-regressions/agent/tools/read_file.txt b/src-tauri/proptest-regressions/agent/tools/read_file.txt deleted file mode 100644 index e082bf1cc..000000000 --- a/src-tauri/proptest-regressions/agent/tools/read_file.txt +++ /dev/null @@ -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"] diff --git a/src-tauri/proptest-regressions/config/tests.txt b/src-tauri/proptest-regressions/config/tests.txt index b4f0871a1..e11cd61d5 100644 --- a/src-tauri/proptest-regressions/config/tests.txt +++ b/src-tauri/proptest-regressions/config/tests.txt @@ -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" diff --git a/src-tauri/proptest-regressions/flow_monitor/file_store.txt b/src-tauri/proptest-regressions/flow_monitor/file_store.txt deleted file mode 100644 index a197268a2..000000000 --- a/src-tauri/proptest-regressions/flow_monitor/file_store.txt +++ /dev/null @@ -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 diff --git a/src-tauri/proptest-regressions/flow_monitor/memory_store.txt b/src-tauri/proptest-regressions/flow_monitor/memory_store.txt deleted file mode 100644 index 9664b511e..000000000 --- a/src-tauri/proptest-regressions/flow_monitor/memory_store.txt +++ /dev/null @@ -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" diff --git a/src-tauri/proptest-regressions/flow_monitor/monitor.txt b/src-tauri/proptest-regressions/flow_monitor/monitor.txt deleted file mode 100644 index 78f28e234..000000000 --- a/src-tauri/proptest-regressions/flow_monitor/monitor.txt +++ /dev/null @@ -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 diff --git a/src-tauri/proptest-regressions/flow_monitor/stream_rebuilder.txt b/src-tauri/proptest-regressions/flow_monitor/stream_rebuilder.txt deleted file mode 100644 index 355da8c9d..000000000 --- a/src-tauri/proptest-regressions/flow_monitor/stream_rebuilder.txt +++ /dev/null @@ -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\"}") diff --git a/src-tauri/proptest-regressions/middleware/tests.txt b/src-tauri/proptest-regressions/middleware/tests.txt deleted file mode 100644 index 99b6aa624..000000000 --- a/src-tauri/proptest-regressions/middleware/tests.txt +++ /dev/null @@ -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" diff --git a/src-tauri/proptest-regressions/providers/tests.txt b/src-tauri/proptest-regressions/providers/tests.txt deleted file mode 100644 index 21ac60ee8..000000000 --- a/src-tauri/proptest-regressions/providers/tests.txt +++ /dev/null @@ -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 diff --git a/src-tauri/proptest-regressions/proxy/tests.txt b/src-tauri/proptest-regressions/proxy/tests.txt deleted file mode 100644 index 75b28b39e..000000000 --- a/src-tauri/proptest-regressions/proxy/tests.txt +++ /dev/null @@ -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" diff --git a/src-tauri/proptest-regressions/router/tests.txt b/src-tauri/proptest-regressions/router/tests.txt deleted file mode 100644 index 34d795756..000000000 --- a/src-tauri/proptest-regressions/router/tests.txt +++ /dev/null @@ -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" diff --git a/src-tauri/proptest-regressions/websocket/tests.txt b/src-tauri/proptest-regressions/websocket/tests.txt deleted file mode 100644 index 212af077e..000000000 --- a/src-tauri/proptest-regressions/websocket/tests.txt +++ /dev/null @@ -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 diff --git a/src-tauri/src/agent/README.md b/src-tauri/src/agent/README.md index f290bde5c..1468a0a00 100644 --- a/src-tauri/src/agent/README.md +++ b/src-tauri/src/agent/README.md @@ -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) diff --git a/src-tauri/src/agent/aster_agent.rs b/src-tauri/src/agent/aster_agent.rs index 9f01938ee..6c66ecc11 100644 --- a/src-tauri/src/agent/aster_agent.rs +++ b/src-tauri/src/agent/aster_agent.rs @@ -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(()) } diff --git a/src-tauri/src/agent/aster_state.rs b/src-tauri/src/agent/aster_state.rs index e876acd50..5f538e201 100644 --- a/src-tauri/src/agent/aster_state.rs +++ b/src-tauri/src/agent/aster_state.rs @@ -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, /// Base URL (可选,用于自定义端点) pub base_url: Option, + /// 凭证 UUID(来自凭证池,用于记录使用和健康状态) + pub credential_uuid: Option, } /// Aster Agent 全局状态 @@ -32,6 +40,8 @@ pub struct AsterAgentState { cancel_tokens: Arc>>, /// 当前 Provider 配置 current_provider_config: Arc>>, + /// 凭证桥接器 + 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 { + // 确保 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 类型设置对应的环境变量 diff --git a/src-tauri/src/agent/credential_bridge.rs b/src-tauri/src/agent/credential_bridge.rs new file mode 100644 index 000000000..ef25d27b5 --- /dev/null +++ b/src-tauri/src/agent/credential_bridge.rs @@ -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, + /// Base URL + pub base_url: Option, + /// 凭证 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 { + // 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 { + 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 { + 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 { + 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 { + 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, 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("没有可用凭证")); + } +} diff --git a/src-tauri/src/agent/mod.rs b/src-tauri/src/agent/mod.rs index 0df850afe..ca48cfe0d 100644 --- a/src-tauri/src/agent/mod.rs +++ b/src-tauri/src/agent/mod.rs @@ -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::*; diff --git a/src-tauri/src/agent/native_agent.rs b/src-tauri/src/agent/native_agent.rs deleted file mode 100644 index a3496966c..000000000 --- a/src-tauri/src/agent/native_agent.rs +++ /dev/null @@ -1,1047 +0,0 @@ -//! 原生 Rust Agent 实现 -//! -//! 支持连续对话(Conversation History)和工具调用(Tools) -//! 使用策略模式支持多种 API 协议(OpenAI、Anthropic、Kiro、Gemini) -//! -//! ## 架构设计 -//! - protocols/ - 协议策略实现 -//! - parsers/ - SSE 流解析器 -//! - NativeAgent - 核心 Agent 逻辑 -//! - NativeAgentState - Tauri 状态管理 -//! -//! ## 流式处理 -//! - Requirements: 1.1, 1.3, 1.4 -//! -//! ## 工具调用循环 -//! - Requirements: 7.1, 7.2, 7.3, 7.4, 7.5, 7.6 - -#![allow(dead_code)] - -use crate::agent::protocols::{create_protocol, Protocol}; -use crate::agent::tool_loop::{ToolCallResult, ToolLoopEngine, ToolLoopState}; -use crate::agent::tools::{create_default_registry, ToolRegistry}; -use crate::agent::types::*; -use crate::models::openai::{ - ChatCompletionRequest, ChatCompletionResponse, ChatMessage, ContentPart as OpenAIContentPart, - MessageContent as OpenAIMessageContent, -}; -use parking_lot::RwLock; -use reqwest::Client; -use std::collections::HashMap; -use std::sync::Arc; -use std::time::Duration; -use tokio::sync::mpsc; -use tracing::{debug, error, info, warn}; - -/// 原生 Agent 实现 -pub struct NativeAgent { - client: Client, - base_url: String, - api_key: String, - sessions: Arc>>, - config: AgentConfig, - /// Provider 类型,决定使用哪种协议 - provider_type: ProviderType, - /// 协议处理器 - protocol: Box, - /// Provider ID,用于自定义 Provider 路由(如 "moonshot") - provider_id: Option, -} - -impl NativeAgent { - pub fn new( - base_url: String, - api_key: String, - provider_type: ProviderType, - provider_id: Option, - ) -> Result { - let client = Client::builder() - .timeout(Duration::from_secs(300)) - .connect_timeout(Duration::from_secs(30)) - .no_proxy() - .build() - .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?; - - let protocol = create_protocol(provider_type); - - // 保存原始 base_url,provider_id 将在构建请求 URL 时使用 - let effective_base_url = base_url.clone(); - - info!( - "[NativeAgent] 创建 Agent: base_url={}, effective_base_url={}, provider={:?}, provider_id={:?}, protocol_endpoint={}", - base_url, - effective_base_url, - provider_type, - provider_id, - protocol.endpoint() - ); - - Ok(Self { - client, - base_url: effective_base_url, - api_key, - sessions: Arc::new(RwLock::new(HashMap::new())), - config: AgentConfig::default(), - provider_type, - protocol, - provider_id, - }) - } - - pub fn with_model(mut self, model: String) -> Self { - self.config.model = model; - self - } - - pub fn with_system_prompt(mut self, prompt: String) -> Self { - self.config.system_prompt = Some(prompt); - self - } - - /// 获取 API 请求的有效 base_url - /// - /// 所有 Provider 都使用标准路由,不再使用 Amp CLI 路由前缀 - fn get_effective_base_url(&self) -> String { - self.base_url.clone() - } - - /// 检查是否是自定义 Provider - fn is_custom_provider(&self) -> bool { - if let Some(ref pid) = self.provider_id { - !matches!( - pid.to_lowercase().as_str(), - "openai" - | "claude" - | "anthropic" - | "gemini" - | "kiro" - | "qwen" - | "codex" - | "antigravity" - | "iflow" - ) - } else { - false - } - } - - /// 发送聊天请求(非流式,用于简单场景) - pub async fn chat(&self, request: NativeChatRequest) -> Result { - let model = request.model.unwrap_or_else(|| self.config.model.clone()); - let session_id = request.session_id.clone(); - let has_images = request.images.as_ref().map(|i| i.len()).unwrap_or(0); - - info!( - "[NativeAgent] 发送聊天请求: model={}, session={:?}, images={}", - model, session_id, has_images - ); - - // 获取会话 - let session = if let Some(sid) = &session_id { - self.sessions.read().get(sid).cloned() - } else { - None - }; - - // 构建消息 - let messages = self.build_openai_messages( - session.as_ref(), - &request.message, - request.images.as_deref(), - ); - - let chat_request = ChatCompletionRequest { - model: model.clone(), - messages, - stream: false, - temperature: self.config.temperature, - max_tokens: self.config.max_tokens, - top_p: None, - tools: None, - tool_choice: None, - reasoning_effort: None, - }; - - // 对于自定义 Provider,使用 provider 特定路由 - let url = if self.is_custom_provider() { - format!("{}/chat/completions", self.get_effective_base_url()) - } else { - format!("{}/v1/chat/completions", self.base_url) - }; - - let response = self - .client - .post(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .json(&chat_request) - .send() - .await - .map_err(|e| format!("请求失败: {}", e))?; - - let status = response.status(); - if !status.is_success() { - let body = response.text().await.unwrap_or_default(); - error!("[NativeAgent] 请求失败: {} - {}", status, body); - return Ok(NativeChatResponse { - content: String::new(), - model, - usage: None, - success: false, - error: Some(format!("API 错误 ({}): {}", status, body)), - }); - } - - let body: ChatCompletionResponse = response - .json() - .await - .map_err(|e| format!("解析响应失败: {}", e))?; - - let content = body - .choices - .first() - .and_then(|c| c.message.content.clone()) - .unwrap_or_default(); - - let usage = Some(TokenUsage { - input_tokens: body.usage.prompt_tokens, - output_tokens: body.usage.completion_tokens, - }); - - // 更新会话历史 - if let Some(sid) = session_id { - self.add_message_to_session( - &sid, - "user", - MessageContent::Text(request.message.clone()), - request.images.as_deref(), - ); - self.add_message_to_session( - &sid, - "assistant", - MessageContent::Text(content.clone()), - None, - ); - } - - info!("[NativeAgent] 聊天完成: content_len={}", content.len()); - - Ok(NativeChatResponse { - content, - model: body.model, - usage, - success: true, - error: None, - }) - } - - /// 流式聊天(使用协议策略模式) - /// - /// Requirements: 1.1, 1.3, 1.4 - pub async fn chat_stream( - &self, - request: NativeChatRequest, - tools: Option<&[crate::models::openai::Tool]>, - tx: mpsc::Sender, - ) -> Result { - let model = request - .model - .clone() - .unwrap_or_else(|| self.config.model.clone()); - let session_id = request.session_id.clone(); - - info!( - "[NativeAgent] 发送流式聊天请求: model={}, session={:?}, provider={:?}, tools_count={}", - model, - session_id, - self.provider_type, - tools.map(|t| t.len()).unwrap_or(0) - ); - - // 获取会话 - let session = if let Some(sid) = &session_id { - self.sessions.read().get(sid).cloned() - } else { - None - }; - - // 获取会话历史和配置 - let history: Vec = session - .as_ref() - .map(|s| s.messages.clone()) - .unwrap_or_default(); - - let config = if let Some(ref sess) = session { - let mut cfg = self.config.clone(); - if sess.system_prompt.is_some() { - cfg.system_prompt = sess.system_prompt.clone(); - } - cfg - } else { - self.config.clone() - }; - - // 使用协议策略发送请求 - // 对于自定义 Provider,使用 provider 特定路由 - let effective_base_url = self.get_effective_base_url(); - let result = self - .protocol - .chat_stream( - &self.client, - &effective_base_url, - &self.api_key, - &history, - &request.message, - request.images.as_deref(), - &model, - &config, - tools, - tx, - self.provider_id.as_deref(), - ) - .await?; - - // 更新会话历史 - if let Some(sid) = &session_id { - self.add_message_to_session( - sid, - "user", - MessageContent::Text(request.message.clone()), - request.images.as_deref(), - ); - self.add_assistant_message_to_session( - sid, - MessageContent::Text(result.content.clone()), - result.tool_calls.clone(), - result.reasoning_content.clone(), - ); - } - - Ok(result) - } - - /// 流式聊天(支持工具调用循环) - /// - /// Requirements: 7.1, 7.2, 7.3, 7.4, 7.5, 7.6 - pub async fn chat_stream_with_tools( - &self, - request: NativeChatRequest, - tx: mpsc::Sender, - tool_loop_engine: &ToolLoopEngine, - ) -> Result { - let session_id = request.session_id.clone(); - let mut state = ToolLoopState::new(); - - // 获取工具定义并转换为 OpenAI 格式 - let aster_definitions = tool_loop_engine.registry().get_definitions(); - let tools: Vec = aster_definitions - .into_iter() - .map(|def| crate::models::openai::Tool::Function { - function: crate::models::openai::FunctionDef { - name: def.name, - description: Some(def.description), - parameters: Some(def.input_schema), - }, - }) - .collect(); - let tools_ref = if tools.is_empty() { - None - } else { - Some(tools.as_slice()) - }; - - // 首次请求 - let mut current_result = self - .chat_stream(request.clone(), tools_ref, tx.clone()) - .await?; - - // 工具调用循环 - // Requirements: 7.3 - THE Tool_Loop SHALL continue until the Agent produces a final response without tool_calls - while tool_loop_engine.should_continue(¤t_result, state.iteration) { - state.increment_iteration(); - - let tool_calls = current_result.tool_calls.as_ref().unwrap(); - state.add_tool_calls(tool_calls.len()); - - info!( - "[NativeAgent] 工具循环迭代 {}: 执行 {} 个工具调用", - state.iteration, - tool_calls.len() - ); - - // 执行所有工具调用 - let tool_results = tool_loop_engine - .execute_all_tool_calls(tool_calls, Some(&tx)) - .await; - - // 将工具结果添加到会话 - if let Some(sid) = &session_id { - for result in &tool_results { - self.add_tool_result_to_session(sid, result); - } - } - - // 继续对话 - let continue_request = NativeChatRequest { - session_id: session_id.clone(), - message: String::new(), - model: request.model.clone(), - images: None, - stream: true, - }; - - current_result = self - .chat_stream_continue(continue_request, tools_ref, tx.clone()) - .await?; - } - - // 检查是否因为达到最大迭代次数而停止 - if state.iteration >= tool_loop_engine.max_iterations() && current_result.has_tool_calls() { - warn!( - "[NativeAgent] 达到最大迭代次数 {},强制停止工具循环", - tool_loop_engine.max_iterations() - ); - let _ = tx - .send(StreamEvent::Error { - message: format!( - "达到最大工具调用迭代次数限制 ({})", - tool_loop_engine.max_iterations() - ), - }) - .await; - } - - state.mark_completed(current_result.content.clone()); - - info!( - "[NativeAgent] 工具循环完成: {} 次迭代, {} 个工具调用", - state.iteration, state.total_tool_calls - ); - - // 发送 FinalDone 事件,通知前端整个对话(包括工具循环)已完成 - let _ = tx - .send(StreamEvent::FinalDone { - usage: current_result.usage.clone(), - }) - .await; - - Ok(current_result) - } - - /// 继续流式对话(使用会话历史) - async fn chat_stream_continue( - &self, - request: NativeChatRequest, - tools: Option<&[crate::models::openai::Tool]>, - tx: mpsc::Sender, - ) -> Result { - let model = request.model.unwrap_or_else(|| self.config.model.clone()); - let session_id = request.session_id.as_ref().ok_or("需要 session_id")?; - - debug!( - "[NativeAgent] 继续流式对话: model={}, session={}, tools_count={}", - model, - session_id, - tools.map(|t| t.len()).unwrap_or(0) - ); - - // 获取会话 - let session = self - .sessions - .read() - .get(session_id) - .cloned() - .ok_or_else(|| format!("会话不存在: {}", session_id))?; - - // 获取配置 - let config = { - let mut cfg = self.config.clone(); - if session.system_prompt.is_some() { - cfg.system_prompt = session.system_prompt.clone(); - } - cfg - }; - - // 使用协议策略继续对话 - // 对于自定义 Provider,使用 provider 特定路由 - let effective_base_url = self.get_effective_base_url(); - let result = self - .protocol - .chat_stream_continue( - &self.client, - &effective_base_url, - &self.api_key, - &session.messages, - &model, - &config, - tools, - tx, - self.provider_id.as_deref(), - ) - .await?; - - // 更新会话历史 - self.add_assistant_message_to_session( - session_id, - MessageContent::Text(result.content.clone()), - result.tool_calls.clone(), - result.reasoning_content.clone(), - ); - - Ok(result) - } - - // ==================== 会话管理方法 ==================== - - /// 构建 OpenAI 格式消息(用于非流式请求) - fn build_openai_messages( - &self, - session: Option<&AgentSession>, - user_message: &str, - images: Option<&[ImageData]>, - ) -> Vec { - let mut messages = Vec::new(); - - // 系统提示词 - let system_prompt = session - .and_then(|s| s.system_prompt.as_ref()) - .or(self.config.system_prompt.as_ref()); - if let Some(prompt) = system_prompt { - messages.push(ChatMessage { - role: "system".to_string(), - content: Some(OpenAIMessageContent::Text(prompt.clone())), - tool_calls: None, - tool_call_id: None, - reasoning_content: None, - }); - } - - // 历史消息 - if let Some(sess) = session { - for msg in &sess.messages { - messages.push(self.convert_to_chat_message(msg)); - } - } - - // 用户消息 - 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 - } - - /// 将 AgentMessage 转换为 OpenAI ChatMessage - fn convert_to_chat_message(&self, msg: &AgentMessage) -> ChatMessage { - let content = match &msg.content { - MessageContent::Text(text) => Some(OpenAIMessageContent::Text(text.clone())), - MessageContent::Parts(parts) => { - let openai_parts: Vec = 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)) - } - }; - - 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: msg.reasoning_content.clone(), - } - } - - /// 添加消息到会话 - fn add_message_to_session( - &self, - session_id: &str, - role: &str, - content: MessageContent, - images: Option<&[ImageData]>, - ) { - let mut sessions = self.sessions.write(); - if let Some(session) = sessions.get_mut(session_id) { - let final_content = if let Some(imgs) = images { - let mut parts = vec![ContentPart::Text { - text: content.as_text(), - }]; - for img in imgs { - parts.push(ContentPart::ImageUrl { - image_url: ImageUrl { - url: format!("data:{};base64,{}", img.media_type, img.data), - detail: None, - }, - }); - } - MessageContent::Parts(parts) - } else { - content - }; - - session.messages.push(AgentMessage { - role: role.to_string(), - content: final_content, - timestamp: chrono::Utc::now().to_rfc3339(), - tool_calls: None, - tool_call_id: None, - reasoning_content: None, - }); - session.updated_at = chrono::Utc::now().to_rfc3339(); - } - } - - /// 添加 assistant 消息到会话(支持工具调用和推理内容) - fn add_assistant_message_to_session( - &self, - session_id: &str, - content: MessageContent, - tool_calls: Option>, - reasoning_content: Option, - ) { - let mut sessions = self.sessions.write(); - if let Some(session) = sessions.get_mut(session_id) { - session.messages.push(AgentMessage { - role: "assistant".to_string(), - content, - timestamp: chrono::Utc::now().to_rfc3339(), - tool_calls, - tool_call_id: None, - reasoning_content, - }); - session.updated_at = chrono::Utc::now().to_rfc3339(); - } - } - - /// 添加工具结果消息到会话 - fn add_tool_result_to_session(&self, session_id: &str, tool_result: &ToolCallResult) { - let mut sessions = self.sessions.write(); - if let Some(session) = sessions.get_mut(session_id) { - session.messages.push(tool_result.to_agent_message()); - session.updated_at = chrono::Utc::now().to_rfc3339(); - } - } - - // ==================== 公开会话管理 API ==================== - - pub fn create_session(&self, model: Option, system_prompt: Option) -> String { - let session_id = uuid::Uuid::new_v4().to_string(); - let now = chrono::Utc::now().to_rfc3339(); - let session = AgentSession { - id: session_id.clone(), - model: model.unwrap_or_else(|| self.config.model.clone()), - messages: Vec::new(), - system_prompt, - created_at: now.clone(), - updated_at: now, - }; - - self.sessions.write().insert(session_id.clone(), session); - info!("[NativeAgent] 创建会话: {}", session_id); - - session_id - } - - pub fn get_session(&self, session_id: &str) -> Option { - self.sessions.read().get(session_id).cloned() - } - - pub fn delete_session(&self, session_id: &str) -> bool { - self.sessions.write().remove(session_id).is_some() - } - - pub fn list_sessions(&self) -> Vec { - self.sessions.read().values().cloned().collect() - } - - pub fn clear_session_messages(&self, session_id: &str) -> bool { - let mut sessions = self.sessions.write(); - if let Some(session) = sessions.get_mut(session_id) { - session.messages.clear(); - session.updated_at = chrono::Utc::now().to_rfc3339(); - true - } else { - false - } - } - - pub fn get_session_messages(&self, session_id: &str) -> Option> { - self.sessions - .read() - .get(session_id) - .map(|s| s.messages.clone()) - } -} - -// ==================== Tauri 状态管理 ==================== - -/// Tauri 状态:原生 Agent 管理器 -#[derive(Clone, Default)] -pub struct NativeAgentState { - agent: Arc>>, -} - -impl NativeAgentState { - pub fn new() -> Self { - Self { - agent: Arc::new(RwLock::new(None)), - } - } - - pub fn init( - &self, - base_url: String, - api_key: String, - provider_type: ProviderType, - provider_id: Option, - ) -> Result<(), String> { - let agent = NativeAgent::new(base_url, api_key, provider_type, provider_id)?; - *self.agent.write() = Some(agent); - Ok(()) - } - - /// 使用配置初始化 Agent - /// - /// 从 NativeAgentConfig 加载系统提示词等配置 - pub fn init_with_config( - &self, - base_url: String, - api_key: String, - provider_type: ProviderType, - provider_id: Option, - agent_config: &crate::config::NativeAgentConfig, - ) -> Result<(), String> { - let mut agent = NativeAgent::new(base_url, api_key, provider_type, provider_id)?; - - // 从配置加载系统提示词 - let system_prompt = agent_config.get_effective_system_prompt().or_else(|| { - // 如果配置启用了默认提示词,使用内置默认值 - if agent_config.use_default_system_prompt { - Some(super::types::DEFAULT_SYSTEM_PROMPT.to_string()) - } else { - None - } - }); - - if let Some(prompt) = system_prompt { - agent.config.system_prompt = Some(prompt); - } - - // 从配置加载其他参数 - agent.config.model = agent_config.default_model.clone(); - agent.config.temperature = Some(agent_config.temperature); - agent.config.max_tokens = Some(agent_config.max_tokens); - - *self.agent.write() = Some(agent); - Ok(()) - } - - pub fn is_initialized(&self) -> bool { - self.agent.read().is_some() - } - - /// 获取当前 Agent 的 provider 类型 - pub fn get_provider_type(&self) -> Option { - self.agent.read().as_ref().map(|a| a.provider_type) - } - - /// 获取当前 Agent 的 provider ID - pub fn get_provider_id(&self) -> Option { - self.agent - .read() - .as_ref() - .and_then(|a| a.provider_id.clone()) - } - - pub fn reset(&self) { - *self.agent.write() = None; - } - - /// 获取工具注册表 - pub fn get_tool_registry(&self) -> Result, String> { - self.get_tool_registry_with_mode(false) - } - - /// 获取工具注册表(支持 Terminal 模式) - /// - /// # Arguments - /// * `terminal_mode` - 是否使用 Terminal 模式(使用 TerminalTool 替代 BashTool) - pub fn get_tool_registry_with_mode( - &self, - terminal_mode: bool, - ) -> Result, String> { - let base_dir = dirs::home_dir().ok_or_else(|| "无法获取用户 home 目录".to_string())?; - let registry = if terminal_mode { - // Terminal 模式暂时使用默认注册表 - create_default_registry(base_dir) - } else { - create_default_registry(base_dir) - }; - Ok(Arc::new(registry)) - } - - /// 创建临时 Agent 用于异步操作 - fn create_temp_agent(&self) -> Result { - let guard = self.agent.read(); - let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?; - - let client = Client::builder() - .timeout(Duration::from_secs(300)) - .connect_timeout(Duration::from_secs(30)) - .no_proxy() - .build() - .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?; - - let protocol = create_protocol(agent.provider_type); - - Ok(NativeAgent { - client, - base_url: agent.base_url.clone(), - api_key: agent.api_key.clone(), - sessions: agent.sessions.clone(), - config: agent.config.clone(), - provider_type: agent.provider_type, - protocol, - provider_id: agent.provider_id.clone(), - }) - } - - /// 创建临时 Agent 用于异步操作(支持根据模型名称动态选择协议) - fn create_temp_agent_with_model(&self, model: &str) -> Result { - let guard = self.agent.read(); - let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?; - - let client = Client::builder() - .timeout(Duration::from_secs(300)) - .connect_timeout(Duration::from_secs(30)) - .no_proxy() - .build() - .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?; - - // 如果有自定义 provider_id,尝试从模型名称推断协议类型 - let provider_type = if let Some(provider_id) = &agent.provider_id { - ProviderType::from_provider_and_model(provider_id, model) - } else { - agent.provider_type - }; - - let protocol = create_protocol(provider_type); - - info!( - "[NativeAgent] 创建临时 Agent: model={}, provider_type={:?}, provider_id={:?}, protocol_endpoint={}", - model, - provider_type, - agent.provider_id, - protocol.endpoint() - ); - - Ok(NativeAgent { - client, - base_url: agent.base_url.clone(), - api_key: agent.api_key.clone(), - sessions: agent.sessions.clone(), - config: agent.config.clone(), - provider_type, - protocol, - provider_id: agent.provider_id.clone(), - }) - } - - pub async fn chat(&self, request: NativeChatRequest) -> Result { - let model = request.model.clone().unwrap_or_else(|| { - self.agent - .read() - .as_ref() - .map(|a| a.config.model.clone()) - .unwrap_or_default() - }); - let temp_agent = self.create_temp_agent_with_model(&model)?; - temp_agent.chat(request).await - } - - pub async fn chat_stream( - &self, - request: NativeChatRequest, - tx: mpsc::Sender, - ) -> Result { - let model = request.model.clone().unwrap_or_else(|| { - self.agent - .read() - .as_ref() - .map(|a| a.config.model.clone()) - .unwrap_or_default() - }); - let temp_agent = self.create_temp_agent_with_model(&model)?; - temp_agent.chat_stream(request, None, tx).await - } - - pub async fn chat_stream_with_tools( - &self, - request: NativeChatRequest, - tx: mpsc::Sender, - tool_loop_engine: &ToolLoopEngine, - ) -> Result { - let model = request.model.clone().unwrap_or_else(|| { - self.agent - .read() - .as_ref() - .map(|a| a.config.model.clone()) - .unwrap_or_default() - }); - let temp_agent = self.create_temp_agent_with_model(&model)?; - temp_agent - .chat_stream_with_tools(request, tx, tool_loop_engine) - .await - } - - pub fn create_session( - &self, - model: Option, - system_prompt: Option, - ) -> Result { - let guard = self.agent.read(); - let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?; - Ok(agent.create_session(model, system_prompt)) - } - - pub fn get_session(&self, session_id: &str) -> Result, String> { - let guard = self.agent.read(); - let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?; - Ok(agent.get_session(session_id)) - } - - pub fn delete_session(&self, session_id: &str) -> bool { - let guard = self.agent.read(); - if let Some(agent) = guard.as_ref() { - agent.delete_session(session_id) - } else { - false - } - } - - pub fn list_sessions(&self) -> Vec { - let guard = self.agent.read(); - if let Some(agent) = guard.as_ref() { - agent.list_sessions() - } else { - Vec::new() - } - } - - pub fn clear_session_messages(&self, session_id: &str) -> bool { - let guard = self.agent.read(); - if let Some(agent) = guard.as_ref() { - agent.clear_session_messages(session_id) - } else { - false - } - } - - pub fn get_session_messages(&self, session_id: &str) -> Option> { - let guard = self.agent.read(); - guard - .as_ref() - .and_then(|a| a.get_session_messages(session_id)) - } -} - -#[cfg(test)] -mod tests { - use crate::agent::parsers::OpenAISSEParser; - - #[test] - fn test_sse_parser_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, _, is_done1, _) = parser.parse_data(data1); - assert_eq!(text1, Some("Hello".to_string())); - assert!(!is_done1); - - let (text2, _, is_done2, _) = parser.parse_data(data2); - assert_eq!(text2, Some(" World".to_string())); - assert!(!is_done2); - - let (text3, _, is_done3, _) = parser.parse_data(data3); - assert!(text3.is_none()); - assert!(is_done3); - - assert_eq!(parser.get_full_content(), "Hello World"); - } - - #[test] - fn test_sse_parser_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 (_, _, is_done, _) = parser.parse_data(data4); - - assert!(is_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"}"#); - } -} diff --git a/src-tauri/src/agent/parsers/anthropic_sse.rs b/src-tauri/src/agent/parsers/anthropic_sse.rs deleted file mode 100644 index ef838c3db..000000000 --- a/src-tauri/src/agent/parsers/anthropic_sse.rs +++ /dev/null @@ -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, - /// 当前正在构建的工具调用 - current_tool: Option, - /// Usage 信息 - usage: Option, -} - -/// Anthropic SSE 解析结果 -#[derive(Debug, Clone)] -pub struct AnthropicParseResult { - /// 文本增量 - pub text_delta: Option, - /// 是否完成 - 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 中,用 标签包裹 - let thinking_text = format!("{}", 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 { - 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 { - self.usage.clone() - } -} diff --git a/src-tauri/src/agent/parsers/mod.rs b/src-tauri/src/agent/parsers/mod.rs deleted file mode 100644 index cd898a324..000000000 --- a/src-tauri/src/agent/parsers/mod.rs +++ /dev/null @@ -1,9 +0,0 @@ -//! SSE 流解析器模块 -//! -//! 提供不同协议的 SSE 流解析器 - -mod anthropic_sse; -mod openai_sse; - -pub use anthropic_sse::{AnthropicParseResult, AnthropicSSEParser}; -pub use openai_sse::OpenAISSEParser; diff --git a/src-tauri/src/agent/parsers/openai_sse.rs b/src-tauri/src/agent/parsers/openai_sse.rs deleted file mode 100644 index a46d7dc13..000000000 --- a/src-tauri/src/agent/parsers/openai_sse.rs +++ /dev/null @@ -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, -} - -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, Option, bool, Option) { - 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 - 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 { - 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); - } -} diff --git a/src-tauri/src/agent/protocols/anthropic.rs b/src-tauri/src/agent/protocols/anthropic.rs deleted file mode 100644 index 2be21dd23..000000000 --- a/src-tauri/src/agent/protocols/anthropic.rs +++ /dev/null @@ -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, - max_tokens: u32, - stream: bool, - #[serde(skip_serializing_if = "Option::is_none")] - system: Option, - #[serde(skip_serializing_if = "Option::is_none")] - temperature: Option, - #[serde(skip_serializing_if = "Option::is_none")] - tools: Option>, -} - -/// 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> { - 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 = 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, Option) { - 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, Option) { - 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 { - let mut validated_messages = Vec::new(); - let mut pending_tool_calls: std::collections::HashMap = - 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, - send_done: bool, - ) -> Result { - 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, - provider_id: Option<&str>, - ) -> Result { - 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, - provider_id: Option<&str>, - ) -> Result { - 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" - } -} diff --git a/src-tauri/src/agent/protocols/mod.rs b/src-tauri/src/agent/protocols/mod.rs deleted file mode 100644 index fbdf98e4f..000000000 --- a/src-tauri/src/agent/protocols/mod.rs +++ /dev/null @@ -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, - provider_id: Option<&str>, - ) -> Result; - - /// 继续流式对话(工具调用后) - /// - /// 使用会话历史继续对话,不添加新的用户消息 - 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, - provider_id: Option<&str>, - ) -> Result; - - /// 获取 API 端点 - fn endpoint(&self) -> &'static str; -} - -/// 根据 ProviderType 创建协议处理器 -pub fn create_protocol(provider_type: ProviderType) -> Box { - 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), - } -} diff --git a/src-tauri/src/agent/protocols/openai.rs b/src-tauri/src/agent/protocols/openai.rs deleted file mode 100644 index a8e731e66..000000000 --- a/src-tauri/src/agent/protocols/openai.rs +++ /dev/null @@ -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 = 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 { - 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 { - 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, - send_done: bool, - ) -> Result { - 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::(&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, - provider_id: Option<&str>, - ) -> Result { - 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, - provider_id: Option<&str>, - ) -> Result { - 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" - } -} diff --git a/src-tauri/src/agent/tool_loop.rs b/src-tauri/src/agent/tool_loop.rs deleted file mode 100644 index fa61b4ecf..000000000 --- a/src-tauri/src/agent/tool_loop.rs +++ /dev/null @@ -1,1130 +0,0 @@ -//! 工具调用循环引擎 -//! -//! 实现 Agent 工具调用循环,自动执行工具并继续对话 -//! 符合 Requirements 7.1, 7.2, 7.3, 7.4, 7.5 -//! -//! ## 功能 -//! - 检测 Agent 响应中的工具调用 -//! - 执行工具并收集结果 -//! - 将工具结果发送回 Agent 继续对话 -//! - 最大迭代限制防止无限循环 - -use crate::agent::tools::{ToolContext, ToolRegistry, ToolResult as ToolsResult}; -use crate::agent::types::{ - AgentMessage, MessageContent, StreamEvent, StreamResult, ToolCall, ToolExecutionResult, -}; -use std::sync::Arc; -use thiserror::Error; -use tokio::sync::mpsc; -use tracing::{debug, warn}; - -/// 工具循环错误类型 -#[derive(Debug, Error)] -pub enum ToolLoopError { - /// 超过最大迭代次数 - /// Requirements: 7.5 - THE Tool_Loop SHALL enforce a maximum iteration limit - #[error("超过最大迭代次数限制: {0}")] - MaxIterationsExceeded(usize), - - /// 工具执行错误 - /// Requirements: 7.4 - IF a tool execution fails, THEN THE Tool_Loop SHALL include the error - #[error("工具执行错误: {0}")] - ToolExecution(String), - - /// 工具未找到 - #[error("工具未找到: {0}")] - ToolNotFound(String), - - /// JSON 解析错误 - #[error("JSON 解析错误: {0}")] - JsonParse(String), - - /// 通道发送错误 - #[error("事件发送失败")] - ChannelSend, -} - -/// 工具执行结果(内部使用) -#[derive(Debug, Clone)] -pub struct ToolCallResult { - /// 工具调用 ID - pub tool_call_id: String, - /// 工具名称 - pub tool_name: String, - /// 执行结果 - pub result: ToolsResult, -} - -impl ToolCallResult { - /// 创建新的工具调用结果 - pub fn new(tool_call_id: String, tool_name: String, result: ToolsResult) -> Self { - Self { - tool_call_id, - tool_name, - result, - } - } - - /// 转换为 AgentMessage(tool 角色) - /// - /// Requirements: 7.2 - THE Tool_Loop SHALL send tool results back to the Agent as tool role messages - pub fn to_agent_message(&self) -> AgentMessage { - let content = if self.result.is_success() { - self.result.output.clone().unwrap_or_default() - } else { - format!( - "Error: {}", - self.result.error.as_deref().unwrap_or("Unknown error") - ) - }; - - AgentMessage { - role: "tool".to_string(), - content: MessageContent::Text(content), - timestamp: chrono::Utc::now().to_rfc3339(), - tool_calls: None, - tool_call_id: Some(self.tool_call_id.clone()), - reasoning_content: None, - } - } - - /// 转换为 ToolExecutionResult(用于前端显示) - pub fn to_execution_result(&self) -> ToolExecutionResult { - ToolExecutionResult { - success: self.result.is_success(), - output: self.result.output.clone().unwrap_or_default(), - error: self.result.error.clone(), - } - } -} - -/// 工具循环引擎配置 -#[derive(Debug, Clone)] -pub struct ToolLoopConfig { - /// 最大迭代次数 - /// Requirements: 7.5 - THE Tool_Loop SHALL enforce a maximum iteration limit - pub max_iterations: usize, -} - -impl Default for ToolLoopConfig { - fn default() -> Self { - Self { - max_iterations: 50, // 默认最大 25 次迭代 - } - } -} - -impl ToolLoopConfig { - /// 创建新的配置 - pub fn new(max_iterations: usize) -> Self { - Self { max_iterations } - } -} - -/// 工具循环引擎 -/// -/// 负责执行工具调用循环,直到 Agent 产生最终响应 -/// Requirements: 7.1, 7.2, 7.3, 7.4, 7.5 -pub struct ToolLoopEngine { - /// 工具注册表 - registry: Arc, - /// 配置 - config: ToolLoopConfig, -} - -impl ToolLoopEngine { - /// 创建新的工具循环引擎 - pub fn new(registry: Arc) -> Self { - Self { - registry, - config: ToolLoopConfig::default(), - } - } - - /// 使用自定义配置创建 - pub fn with_config(registry: Arc, config: ToolLoopConfig) -> Self { - Self { registry, config } - } - - /// 获取最大迭代次数 - pub fn max_iterations(&self) -> usize { - self.config.max_iterations - } - - /// 获取工具注册表引用 - pub fn registry(&self) -> &ToolRegistry { - &self.registry - } - - /// 检查响应是否包含工具调用 - /// - /// Requirements: 7.1 - WHEN the Agent response contains tool_calls - pub fn has_tool_calls(result: &StreamResult) -> bool { - result.has_tool_calls() - } - - /// 执行单个工具调用 - /// - /// Requirements: 7.1 - THE Tool_Loop SHALL execute each tool and collect results - /// Requirements: 7.4 - IF a tool execution fails, THEN THE Tool_Loop SHALL include the error - pub async fn execute_tool_call(&self, tool_call: &ToolCall) -> ToolCallResult { - let tool_name = &tool_call.function.name; - let tool_id = &tool_call.id; - - debug!("[ToolLoopEngine] 执行工具: {} (id={})", tool_name, tool_id); - - // 解析参数 - let args = match serde_json::from_str::(&tool_call.function.arguments) { - Ok(args) => args, - Err(e) => { - warn!("[ToolLoopEngine] 工具参数解析失败: {} - {}", tool_name, e); - return ToolCallResult::new( - tool_id.clone(), - tool_name.clone(), - ToolsResult::error(format!("参数解析失败: {}", e)), - ); - } - }; - - // 执行工具 - match self - .registry - .execute(tool_name, args, &ToolContext::default(), None) - .await - { - Ok(result) => { - debug!( - "[ToolLoopEngine] 工具执行成功: {} success={}", - tool_name, - result.is_success() - ); - ToolCallResult::new(tool_id.clone(), tool_name.clone(), result) - } - Err(e) => { - warn!("[ToolLoopEngine] 工具执行失败: {} - {}", tool_name, e); - let error_msg = format!("工具执行失败: {}", e); - ToolCallResult::new( - tool_id.clone(), - tool_name.clone(), - ToolsResult::error(error_msg), - ) - } - } - } - - /// 执行所有工具调用 - /// - /// Requirements: 7.1 - THE Tool_Loop SHALL execute each tool and collect results - /// Requirements: 7.6 - WHILE the Tool_Loop is executing, THE Frontend SHALL display the current tool - pub async fn execute_all_tool_calls( - &self, - tool_calls: &[ToolCall], - event_tx: Option<&mpsc::Sender>, - ) -> Vec { - let mut results = Vec::with_capacity(tool_calls.len()); - - for tool_call in tool_calls { - // 发送工具开始事件 - if let Some(tx) = event_tx { - let _ = tx - .send(StreamEvent::ToolStart { - tool_name: tool_call.function.name.clone(), - tool_id: tool_call.id.clone(), - arguments: Some(tool_call.function.arguments.clone()), - }) - .await; - } - - // 执行工具 - let result = self.execute_tool_call(tool_call).await; - - // 发送工具结束事件 - if let Some(tx) = event_tx { - let _ = tx - .send(StreamEvent::ToolEnd { - tool_id: tool_call.id.clone(), - result: result.to_execution_result(), - }) - .await; - } - - results.push(result); - } - - results - } - - /// 将工具结果转换为 Agent 消息列表 - /// - /// Requirements: 7.2 - THE Tool_Loop SHALL send tool results back to the Agent as tool role messages - pub fn results_to_messages(results: &[ToolCallResult]) -> Vec { - results.iter().map(|r| r.to_agent_message()).collect() - } - - /// 检查是否应该继续循环 - /// - /// Requirements: 7.3 - THE Tool_Loop SHALL continue until the Agent produces a final response without tool_calls - /// Requirements: 7.5 - THE Tool_Loop SHALL enforce a maximum iteration limit - pub fn should_continue(&self, result: &StreamResult, iteration: usize) -> bool { - // 检查最大迭代次数 - if iteration >= self.config.max_iterations { - warn!( - "[ToolLoopEngine] 达到最大迭代次数: {}", - self.config.max_iterations - ); - return false; - } - - // 检查是否有工具调用 - Self::has_tool_calls(result) - } - - /// 创建 assistant 消息(包含工具调用) - pub fn create_assistant_message( - content: &str, - tool_calls: Option>, - ) -> AgentMessage { - AgentMessage { - role: "assistant".to_string(), - content: MessageContent::Text(content.to_string()), - timestamp: chrono::Utc::now().to_rfc3339(), - tool_calls: tool_calls.map(|calls| { - calls - .into_iter() - .map(|tc| crate::agent::types::ToolCall { - id: tc.id, - call_type: tc.call_type, - function: tc.function, - }) - .collect() - }), - tool_call_id: None, - reasoning_content: None, - } - } -} - -/// 工具循环状态 -/// -/// 用于跟踪工具循环的执行状态 -#[derive(Debug, Clone)] -pub struct ToolLoopState { - /// 当前迭代次数 - pub iteration: usize, - /// 累计执行的工具调用数 - pub total_tool_calls: usize, - /// 是否已完成 - pub completed: bool, - /// 最终内容 - pub final_content: Option, -} - -impl Default for ToolLoopState { - fn default() -> Self { - Self { - iteration: 0, - total_tool_calls: 0, - completed: false, - final_content: None, - } - } -} - -impl ToolLoopState { - /// 创建新的状态 - pub fn new() -> Self { - Self::default() - } - - /// 增加迭代次数 - pub fn increment_iteration(&mut self) { - self.iteration += 1; - } - - /// 增加工具调用计数 - pub fn add_tool_calls(&mut self, count: usize) { - self.total_tool_calls += count; - } - - /// 标记为完成 - pub fn mark_completed(&mut self, content: String) { - self.completed = true; - self.final_content = Some(content); - } -} - -// TODO: 重新实现测试,适配 aster-rust 的 Tool trait -// 当前暂时禁用测试,等待完整的工具系统集成 -/* -#[cfg(test)] -mod tests { - use super::*; - use crate::agent::tools::{JsonSchema, PropertySchema, ToolDefinition}; - use crate::agent::tools::{Tool, ToolRegistry}; - use crate::agent::types::FunctionCall; - use async_trait::async_trait; - - /// 测试用的 Echo 工具 - struct EchoTool; - - #[async_trait] - impl Tool for EchoTool { - fn definition(&self) -> ToolDefinition { - ToolDefinition::new("echo", "Echo the input message").with_parameters( - JsonSchema::new().add_property( - "message", - PropertySchema::string("The message to echo"), - true, - ), - ) - } - - async fn execute(&self, args: serde_json::Value) -> Result { - let message = args - .get("message") - .and_then(|v| v.as_str()) - .ok_or_else(|| ToolError::InvalidArguments("缺少 message 参数".to_string()))?; - - Ok(ToolsResult::success(message)) - } - } - - /// 测试用的失败工具 - struct FailingTool; - - #[async_trait] - impl Tool for FailingTool { - fn definition(&self) -> ToolDefinition { - ToolDefinition::new("failing", "A tool that always fails") - } - - async fn execute(&self, _args: serde_json::Value) -> Result { - Err(ToolError::ExecutionFailed("故意失败".to_string())) - } - } - - fn create_test_registry() -> Arc { - let registry = ToolRegistry::new(); - registry.register(EchoTool).unwrap(); - registry.register(FailingTool).unwrap(); - Arc::new(registry) - } - - fn create_tool_call(id: &str, name: &str, args: &str) -> ToolCall { - ToolCall { - id: id.to_string(), - call_type: "function".to_string(), - function: FunctionCall { - name: name.to_string(), - arguments: args.to_string(), - }, - } - } - - #[test] - fn test_tool_loop_engine_creation() { - let registry = create_test_registry(); - let engine = ToolLoopEngine::new(registry); - - assert_eq!(engine.max_iterations(), 25); - } - - #[test] - fn test_tool_loop_engine_with_config() { - let registry = create_test_registry(); - let config = ToolLoopConfig::new(10); - let engine = ToolLoopEngine::with_config(registry, config); - - assert_eq!(engine.max_iterations(), 10); - } - - #[test] - fn test_has_tool_calls() { - // 无工具调用 - let result_no_tools = StreamResult::new("Hello".to_string()); - assert!(!ToolLoopEngine::has_tool_calls(&result_no_tools)); - - // 有工具调用 - let result_with_tools = StreamResult::new("".to_string()).with_tool_calls(vec![ToolCall { - id: "call_1".to_string(), - call_type: "function".to_string(), - function: FunctionCall { - name: "echo".to_string(), - arguments: "{}".to_string(), - }, - }]); - assert!(ToolLoopEngine::has_tool_calls(&result_with_tools)); - - // 空工具调用列表 - let result_empty_tools = StreamResult { - content: "".to_string(), - tool_calls: Some(vec![]), - usage: None, - reasoning_content: None, - }; - assert!(!ToolLoopEngine::has_tool_calls(&result_empty_tools)); - } - - #[tokio::test] - async fn test_execute_tool_call_success() { - let registry = create_test_registry(); - let engine = ToolLoopEngine::new(registry); - - let tool_call = create_tool_call("call_1", "echo", r#"{"message": "Hello, World!"}"#); - let result = engine.execute_tool_call(&tool_call).await; - - assert_eq!(result.tool_call_id, "call_1"); - assert_eq!(result.tool_name, "echo"); - assert!(result.result.success); - assert_eq!(result.result.output, "Hello, World!"); - } - - #[tokio::test] - async fn test_execute_tool_call_failure() { - let registry = create_test_registry(); - let engine = ToolLoopEngine::new(registry); - - let tool_call = create_tool_call("call_2", "failing", "{}"); - let result = engine.execute_tool_call(&tool_call).await; - - assert_eq!(result.tool_call_id, "call_2"); - assert_eq!(result.tool_name, "failing"); - assert!(!result.result.success); - assert!(result.result.error.is_some()); - } - - #[tokio::test] - async fn test_execute_tool_call_not_found() { - let registry = create_test_registry(); - let engine = ToolLoopEngine::new(registry); - - let tool_call = create_tool_call("call_3", "nonexistent", "{}"); - let result = engine.execute_tool_call(&tool_call).await; - - assert!(!result.result.success); - assert!(result.result.error.as_ref().unwrap().contains("工具不存在")); - } - - #[tokio::test] - async fn test_execute_tool_call_invalid_args() { - let registry = create_test_registry(); - let engine = ToolLoopEngine::new(registry); - - let tool_call = create_tool_call("call_4", "echo", "invalid json"); - let result = engine.execute_tool_call(&tool_call).await; - - assert!(!result.result.success); - assert!(result - .result - .error - .as_ref() - .unwrap() - .contains("参数解析失败")); - } - - #[tokio::test] - async fn test_execute_all_tool_calls() { - let registry = create_test_registry(); - let engine = ToolLoopEngine::new(registry); - - let tool_calls = vec![ - create_tool_call("call_1", "echo", r#"{"message": "First"}"#), - create_tool_call("call_2", "echo", r#"{"message": "Second"}"#), - ]; - - let results = engine.execute_all_tool_calls(&tool_calls, None).await; - - assert_eq!(results.len(), 2); - assert!(results[0].result.success); - assert_eq!(results[0].result.output, "First"); - assert!(results[1].result.success); - assert_eq!(results[1].result.output, "Second"); - } - - #[test] - fn test_results_to_messages() { - let results = vec![ - ToolCallResult::new( - "call_1".to_string(), - "echo".to_string(), - ToolsResult::success("Hello"), - ), - ToolCallResult::new( - "call_2".to_string(), - "failing".to_string(), - ToolsResult::failure("Error occurred"), - ), - ]; - - let messages = ToolLoopEngine::results_to_messages(&results); - - assert_eq!(messages.len(), 2); - - // 第一个消息(成功) - assert_eq!(messages[0].role, "tool"); - assert_eq!(messages[0].content.as_text(), "Hello"); - assert_eq!(messages[0].tool_call_id, Some("call_1".to_string())); - - // 第二个消息(失败) - assert_eq!(messages[1].role, "tool"); - assert!(messages[1].content.as_text().contains("Error")); - assert_eq!(messages[1].tool_call_id, Some("call_2".to_string())); - } - - #[test] - fn test_should_continue() { - let registry = create_test_registry(); - let engine = ToolLoopEngine::with_config(registry, ToolLoopConfig::new(5)); - - // 有工具调用,未达到限制 - let result_with_tools = StreamResult::new("".to_string()).with_tool_calls(vec![ToolCall { - id: "call_1".to_string(), - call_type: "function".to_string(), - function: FunctionCall { - name: "echo".to_string(), - arguments: "{}".to_string(), - }, - }]); - assert!(engine.should_continue(&result_with_tools, 0)); - assert!(engine.should_continue(&result_with_tools, 4)); - - // 达到最大迭代次数 - assert!(!engine.should_continue(&result_with_tools, 5)); - - // 无工具调用 - let result_no_tools = StreamResult::new("Final response".to_string()); - assert!(!engine.should_continue(&result_no_tools, 0)); - } - - #[test] - fn test_tool_loop_state() { - let mut state = ToolLoopState::new(); - - assert_eq!(state.iteration, 0); - assert_eq!(state.total_tool_calls, 0); - assert!(!state.completed); - assert!(state.final_content.is_none()); - - state.increment_iteration(); - assert_eq!(state.iteration, 1); - - state.add_tool_calls(3); - assert_eq!(state.total_tool_calls, 3); - - state.mark_completed("Final content".to_string()); - assert!(state.completed); - assert_eq!(state.final_content, Some("Final content".to_string())); - } - - #[test] - fn test_tool_call_result_to_agent_message() { - // 成功结果 - let success_result = ToolCallResult::new( - "call_1".to_string(), - "echo".to_string(), - ToolsResult::success("Success output"), - ); - let success_msg = success_result.to_agent_message(); - assert_eq!(success_msg.role, "tool"); - assert_eq!(success_msg.content.as_text(), "Success output"); - assert_eq!(success_msg.tool_call_id, Some("call_1".to_string())); - - // 失败结果 - let failure_result = ToolCallResult::new( - "call_2".to_string(), - "failing".to_string(), - ToolsResult::failure("Something went wrong"), - ); - let failure_msg = failure_result.to_agent_message(); - assert_eq!(failure_msg.role, "tool"); - assert!(failure_msg.content.as_text().contains("Error")); - assert!(failure_msg - .content - .as_text() - .contains("Something went wrong")); - } - - #[tokio::test] - async fn test_execute_all_tool_calls_with_events() { - let registry = create_test_registry(); - let engine = ToolLoopEngine::new(registry); - - let (tx, mut rx) = mpsc::channel::(10); - - let tool_calls = vec![create_tool_call("call_1", "echo", r#"{"message": "Test"}"#)]; - - let results = engine.execute_all_tool_calls(&tool_calls, Some(&tx)).await; - - assert_eq!(results.len(), 1); - assert!(results[0].result.success); - - // 检查事件 - let event1 = rx.recv().await.unwrap(); - assert!(matches!(event1, StreamEvent::ToolStart { .. })); - - let event2 = rx.recv().await.unwrap(); - assert!(matches!(event2, StreamEvent::ToolEnd { .. })); - } -} - -#[cfg(test)] -mod proptests { - use super::*; - use crate::agent::tools::{JsonSchema, PropertySchema, ToolDefinition}; - use crate::agent::tools::{Tool, ToolRegistry, ToolResult as ToolsResult}; - use crate::agent::types::FunctionCall; - use async_trait::async_trait; - use proptest::prelude::*; - - /// 测试用的 Echo 工具(用于属性测试) - struct PropTestEchoTool; - - #[async_trait] - impl Tool for PropTestEchoTool { - fn definition(&self) -> ToolDefinition { - ToolDefinition::new("echo", "Echo the input message").with_parameters( - JsonSchema::new().add_property( - "message", - PropertySchema::string("The message to echo"), - true, - ), - ) - } - - async fn execute(&self, args: serde_json::Value) -> Result { - let message = args - .get("message") - .and_then(|v| v.as_str()) - .ok_or_else(|| ToolError::InvalidArguments("缺少 message 参数".to_string()))?; - - Ok(ToolsResult::success(message)) - } - } - - /// 测试用的计数工具 - struct PropTestCountTool; - - #[async_trait] - impl Tool for PropTestCountTool { - fn definition(&self) -> ToolDefinition { - ToolDefinition::new("count", "Count characters in a string").with_parameters( - JsonSchema::new().add_property( - "text", - PropertySchema::string("The text to count"), - true, - ), - ) - } - - async fn execute(&self, args: serde_json::Value) -> Result { - let text = args - .get("text") - .and_then(|v| v.as_str()) - .ok_or_else(|| ToolError::InvalidArguments("缺少 text 参数".to_string()))?; - - Ok(ToolsResult::success(format!("{}", text.len()))) - } - } - - fn create_proptest_registry() -> Arc { - let registry = ToolRegistry::new(); - registry.register(PropTestEchoTool).unwrap(); - registry.register(PropTestCountTool).unwrap(); - Arc::new(registry) - } - - /// 生成有效的工具调用 ID - fn arb_tool_id() -> impl Strategy { - "call_[a-zA-Z0-9]{8}".prop_map(|s| s) - } - - /// 生成有效的消息内容(用于 echo 工具) - fn arb_message_content() -> impl Strategy { - "[a-zA-Z0-9 ]{1,50}".prop_map(|s| s) - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: agent-tool-calling, Property 12: 工具循环执行完整性** - /// **Validates: Requirements 7.1, 7.2** - /// - /// *For any* 包含 tool_calls 的 Agent 响应,Tool Loop 应该执行所有工具并将结果 - /// 作为 tool 角色消息发送回 Agent。 - #[test] - fn prop_tool_loop_executes_all_tools( - tool_ids in prop::collection::vec(arb_tool_id(), 1..=5), - messages in prop::collection::vec(arb_message_content(), 1..=5) - ) { - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - let registry = create_proptest_registry(); - let engine = ToolLoopEngine::new(registry); - - // 创建工具调用列表 - let tool_calls: Vec = tool_ids - .iter() - .zip(messages.iter()) - .map(|(id, msg)| ToolCall { - id: id.clone(), - call_type: "function".to_string(), - function: FunctionCall { - name: "echo".to_string(), - arguments: serde_json::json!({"message": msg}).to_string(), - }, - }) - .collect(); - - let num_calls = tool_calls.len(); - - // 执行所有工具调用 - let results = engine.execute_all_tool_calls(&tool_calls, None).await; - - // 验证:结果数量等于工具调用数量 - prop_assert_eq!( - results.len(), - num_calls, - "结果数量应该等于工具调用数量" - ); - - // 验证:每个结果都有正确的 tool_call_id - for (i, result) in results.iter().enumerate() { - prop_assert_eq!( - &result.tool_call_id, - &tool_ids[i], - "工具调用 ID 应该匹配" - ); - } - - // 验证:所有工具都成功执行 - for result in &results { - prop_assert!( - result.result.success, - "工具执行应该成功: {:?}", - result.result.error - ); - } - - Ok(()) - })?; - } - - /// **Feature: agent-tool-calling, Property 12: 工具循环执行完整性 - 结果转换为消息** - /// **Validates: Requirements 7.1, 7.2** - /// - /// *For any* 工具执行结果,转换为 AgentMessage 后应该具有正确的 role 和 tool_call_id。 - #[test] - fn prop_tool_results_convert_to_messages( - tool_id in arb_tool_id(), - message in arb_message_content() - ) { - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - let registry = create_proptest_registry(); - let engine = ToolLoopEngine::new(registry); - - let tool_call = ToolCall { - id: tool_id.clone(), - call_type: "function".to_string(), - function: FunctionCall { - name: "echo".to_string(), - arguments: serde_json::json!({"message": message}).to_string(), - }, - }; - - // 执行工具 - let result = engine.execute_tool_call(&tool_call).await; - - // 转换为 AgentMessage - let agent_msg = result.to_agent_message(); - - // 验证:role 为 "tool" - prop_assert_eq!( - agent_msg.role, - "tool", - "消息角色应该为 'tool'" - ); - - // 验证:tool_call_id 正确 - prop_assert_eq!( - agent_msg.tool_call_id, - Some(tool_id.clone()), - "tool_call_id 应该匹配" - ); - - // 验证:成功结果的内容包含原始消息 - if result.result.success { - prop_assert!( - agent_msg.content.as_text().contains(&message), - "成功结果的内容应该包含原始消息" - ); - } - - Ok(()) - })?; - } - - /// **Feature: agent-tool-calling, Property 12: 工具循环执行完整性 - 事件发送** - /// **Validates: Requirements 7.1, 7.6** - /// - /// *For any* 工具执行,应该发送 ToolStart 和 ToolEnd 事件。 - #[test] - fn prop_tool_execution_sends_events( - tool_id in arb_tool_id(), - message in arb_message_content() - ) { - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - let registry = create_proptest_registry(); - let engine = ToolLoopEngine::new(registry); - - let (tx, mut rx) = mpsc::channel::(10); - - let tool_calls = vec![ToolCall { - id: tool_id.clone(), - call_type: "function".to_string(), - function: FunctionCall { - name: "echo".to_string(), - arguments: serde_json::json!({"message": message}).to_string(), - }, - }]; - - // 执行工具调用 - let _ = engine.execute_all_tool_calls(&tool_calls, Some(&tx)).await; - - // 验证:收到 ToolStart 事件 - let event1 = rx.recv().await; - prop_assert!(event1.is_some(), "应该收到 ToolStart 事件"); - if let Some(StreamEvent::ToolStart { tool_name, tool_id: event_tool_id, .. }) = event1 { - prop_assert_eq!(tool_name, "echo", "工具名称应该为 'echo'"); - prop_assert_eq!(event_tool_id, tool_id.clone(), "工具 ID 应该匹配"); - } else { - prop_assert!(false, "第一个事件应该是 ToolStart"); - } - - // 验证:收到 ToolEnd 事件 - let event2 = rx.recv().await; - prop_assert!(event2.is_some(), "应该收到 ToolEnd 事件"); - if let Some(StreamEvent::ToolEnd { tool_id: event_tool_id, result }) = event2 { - prop_assert_eq!(event_tool_id, tool_id.clone(), "工具 ID 应该匹配"); - prop_assert!(result.success, "工具执行应该成功"); - } else { - prop_assert!(false, "第二个事件应该是 ToolEnd"); - } - - Ok(()) - })?; - } - - /// **Feature: agent-tool-calling, Property 12: 工具循环执行完整性 - 多工具执行顺序** - /// **Validates: Requirements 7.1, 7.2** - /// - /// *For any* 多个工具调用,执行顺序应该与调用顺序一致。 - #[test] - fn prop_tool_execution_order_preserved( - count in 2..=5usize - ) { - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - let registry = create_proptest_registry(); - let engine = ToolLoopEngine::new(registry); - - // 创建多个工具调用,每个使用不同的消息 - let tool_calls: Vec = (0..count) - .map(|i| ToolCall { - id: format!("call_{}", i), - call_type: "function".to_string(), - function: FunctionCall { - name: "echo".to_string(), - arguments: serde_json::json!({"message": format!("msg_{}", i)}).to_string(), - }, - }) - .collect(); - - // 执行所有工具调用 - let results = engine.execute_all_tool_calls(&tool_calls, None).await; - - // 验证:结果顺序与调用顺序一致 - for (i, result) in results.iter().enumerate() { - prop_assert_eq!( - &result.tool_call_id, - &format!("call_{}", i), - "结果顺序应该与调用顺序一致" - ); - prop_assert!( - result.result.output.contains(&format!("msg_{}", i)), - "结果内容应该对应正确的调用" - ); - } - - Ok(()) - })?; - } - - /// **Feature: agent-tool-calling, Property 13: 工具循环终止** - /// **Validates: Requirements 7.3, 7.5** - /// - /// *For any* 工具循环执行,当 Agent 响应不包含 tool_calls 时,循环应该终止。 - #[test] - fn prop_tool_loop_terminates_without_tool_calls( - content in "[a-zA-Z0-9 ]{1,100}" - ) { - let registry = create_proptest_registry(); - let engine = ToolLoopEngine::new(registry); - - // 创建不包含工具调用的响应 - let result = StreamResult::new(content.clone()); - - // 验证:should_continue 返回 false - prop_assert!( - !engine.should_continue(&result, 0), - "不包含工具调用的响应应该终止循环" - ); - - // 验证:has_tool_calls 返回 false - prop_assert!( - !ToolLoopEngine::has_tool_calls(&result), - "不包含工具调用的响应 has_tool_calls 应该返回 false" - ); - } - - /// **Feature: agent-tool-calling, Property 13: 工具循环终止 - 最大迭代次数** - /// **Validates: Requirements 7.3, 7.5** - /// - /// *For any* 工具循环执行,当达到最大迭代次数时,循环应该终止。 - #[test] - fn prop_tool_loop_terminates_at_max_iterations( - max_iterations in 1..=20usize, - current_iteration in 0..=25usize - ) { - let registry = create_proptest_registry(); - let config = ToolLoopConfig::new(max_iterations); - let engine = ToolLoopEngine::with_config(registry, config); - - // 创建包含工具调用的响应 - let result = StreamResult::new("".to_string()).with_tool_calls(vec![ToolCall { - id: "call_1".to_string(), - call_type: "function".to_string(), - function: FunctionCall { - name: "echo".to_string(), - arguments: r#"{"message": "test"}"#.to_string(), - }, - }]); - - let should_continue = engine.should_continue(&result, current_iteration); - - if current_iteration >= max_iterations { - // 达到或超过最大迭代次数,应该终止 - prop_assert!( - !should_continue, - "达到最大迭代次数 {} 时应该终止循环(当前迭代: {})", - max_iterations, - current_iteration - ); - } else { - // 未达到最大迭代次数,应该继续 - prop_assert!( - should_continue, - "未达到最大迭代次数 {} 时应该继续循环(当前迭代: {})", - max_iterations, - current_iteration - ); - } - } - - /// **Feature: agent-tool-calling, Property 13: 工具循环终止 - 空工具调用列表** - /// **Validates: Requirements 7.3, 7.5** - /// - /// *For any* 包含空工具调用列表的响应,循环应该终止。 - #[test] - fn prop_tool_loop_terminates_with_empty_tool_calls( - content in "[a-zA-Z0-9 ]{1,100}" - ) { - let registry = create_proptest_registry(); - let engine = ToolLoopEngine::new(registry); - - // 创建包含空工具调用列表的响应 - let result = StreamResult { - content: content.clone(), - tool_calls: Some(vec![]), - usage: None, - reasoning_content: None, - }; - - // 验证:should_continue 返回 false - prop_assert!( - !engine.should_continue(&result, 0), - "空工具调用列表应该终止循环" - ); - - // 验证:has_tool_calls 返回 false - prop_assert!( - !ToolLoopEngine::has_tool_calls(&result), - "空工具调用列表 has_tool_calls 应该返回 false" - ); - } - - /// **Feature: agent-tool-calling, Property 13: 工具循环终止 - 配置一致性** - /// **Validates: Requirements 7.5** - /// - /// *For any* 配置的最大迭代次数,engine.max_iterations() 应该返回相同的值。 - #[test] - fn prop_tool_loop_config_consistency( - max_iterations in 1..=100usize - ) { - let registry = create_proptest_registry(); - let config = ToolLoopConfig::new(max_iterations); - let engine = ToolLoopEngine::with_config(registry, config); - - prop_assert_eq!( - engine.max_iterations(), - max_iterations, - "max_iterations() 应该返回配置的值" - ); - } - - /// **Feature: agent-tool-calling, Property 13: 工具循环终止 - 边界条件** - /// **Validates: Requirements 7.3, 7.5** - /// - /// *For any* 最大迭代次数,在边界处的行为应该正确。 - #[test] - fn prop_tool_loop_boundary_conditions( - max_iterations in 1..=20usize - ) { - let registry = create_proptest_registry(); - let config = ToolLoopConfig::new(max_iterations); - let engine = ToolLoopEngine::with_config(registry, config); - - // 创建包含工具调用的响应 - let result_with_tools = StreamResult::new("".to_string()).with_tool_calls(vec![ToolCall { - id: "call_1".to_string(), - call_type: "function".to_string(), - function: FunctionCall { - name: "echo".to_string(), - arguments: r#"{"message": "test"}"#.to_string(), - }, - }]); - - // 在 max_iterations - 1 处应该继续 - if max_iterations > 0 { - prop_assert!( - engine.should_continue(&result_with_tools, max_iterations - 1), - "在 max_iterations - 1 处应该继续" - ); - } - - // 在 max_iterations 处应该终止 - prop_assert!( - !engine.should_continue(&result_with_tools, max_iterations), - "在 max_iterations 处应该终止" - ); - - // 在 max_iterations + 1 处应该终止 - prop_assert!( - !engine.should_continue(&result_with_tools, max_iterations + 1), - "在 max_iterations + 1 处应该终止" - ); - } - } -} -*/ diff --git a/src-tauri/src/agent/tools/README.md b/src-tauri/src/agent/tools/README.md deleted file mode 100644 index bb311732c..000000000 --- a/src-tauri/src/agent/tools/README.md +++ /dev/null @@ -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 { - let message = args.get("message") - .and_then(|v| v.as_str()) - .ok_or_else(|| ToolError::InvalidArguments("缺少 message 参数".to_string()))?; - - Ok(ToolResult::success(message)) - } -} -``` - -### 注册和执行工具 - -```rust -use crate::agent::tools::ToolRegistry; - -let registry = ToolRegistry::new(); - -// 注册工具 -registry.register(EchoTool)?; - -// 执行工具 -let result = registry.execute("echo", serde_json::json!({"message": "Hello!"})).await?; -assert!(result.success); -assert_eq!(result.output, "Hello!"); -``` - -### 使用文件读取工具 - -```rust -use crate::agent::tools::{ReadFileTool, SecurityManager}; -use std::sync::Arc; - -let security = Arc::new(SecurityManager::new("/path/to/project")); -let tool = ReadFileTool::new(security); - -// 读取整个文件 -let result = tool.read_file(Path::new("src/main.rs"), None, None)?; -println!("语言: {:?}", result.language); -println!("总行数: {}", result.total_lines); - -// 读取指定行范围 -let result = tool.read_file(Path::new("src/main.rs"), Some(10), Some(20))?; -println!("内容:\n{}", result.content); -``` - -### 使用文件写入工具 - -```rust -use crate::agent::tools::{WriteFileTool, SecurityManager}; -use std::sync::Arc; - -let security = Arc::new(SecurityManager::new("/path/to/project")); -let tool = WriteFileTool::new(security); - -// 写入新文件 -let result = tool.write_file(Path::new("output.txt"), "Hello, World!")?; -println!("创建: {}, 字节数: {}", result.created, result.bytes_written); - -// 覆盖已有文件 -let result = tool.write_file(Path::new("output.txt"), "New content")?; -println!("覆盖: {}", result.overwritten); - -// 自动创建父目录 -let result = tool.write_file(Path::new("a/b/c/nested.txt"), "Nested content")?; -println!("路径: {:?}", result.path); -``` - -### 使用文件编辑工具 - -```rust -use crate::agent::tools::{EditFileTool, SecurityManager}; -use std::sync::Arc; - -let security = Arc::new(SecurityManager::new("/path/to/project")); -let tool = EditFileTool::new(security); - -// 精确字符串替换 -let result = tool.edit_file(Path::new("src/main.rs"), "old_code", "new_code")?; -println!("替换: {} 字节 -> {} 字节", result.old_str_len, result.new_str_len); -println!("变更上下文:\n{}", result.context_snippet); -println!("Diff:\n{}", result.diff); - -// 撤销编辑 -let undo_result = tool.undo_edit(Path::new("src/main.rs"))?; -println!("已恢复: {} 字节", undo_result.restored_content_len); - -// 查看历史记录数量 -let count = tool.history_count(Path::new("src/main.rs")); -println!("历史记录: {} 条", count); -``` - -### 使用 Prompt 生成器 - -```rust -use crate::agent::tools::{ToolPromptGenerator, PromptFormat, ToolDefinition, JsonSchema, PropertySchema}; - -// 创建工具定义 -let tools = vec![ - ToolDefinition::new("bash", "Execute a bash command") - .with_parameters( - JsonSchema::new() - .add_property("command", PropertySchema::string("The command to execute"), true) - ), - ToolDefinition::new("read_file", "Read file contents") - .with_parameters( - JsonSchema::new() - .add_property("path", PropertySchema::string("The file path"), true) - ), -]; - -// 生成 XML 格式的 System Prompt -let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml); -let system_prompt = generator.generate_system_prompt(&tools); -println!("System Prompt:\n{}", system_prompt); - -// 生成 JSON 格式的 System Prompt -let json_generator = ToolPromptGenerator::new().with_format(PromptFormat::Json); -let json_prompt = json_generator.generate_system_prompt(&tools); -println!("JSON Prompt:\n{}", json_prompt); - -// 使用便捷函数 -use crate::agent::tools::generate_tools_prompt; -let prompt = generate_tools_prompt(&tools, PromptFormat::Xml); -``` - -## 需求追溯 - -- Requirements 2.1: 工具定义包含 name, description, JSON Schema parameters -- Requirements 2.2: 注册时验证工具定义 -- Requirements 2.4: 运行时添加工具无需重启 -- Requirements 2.5: 支持必需和可选参数类型验证 -- Requirements 3.1: Bash 工具在用户默认 shell 中执行命令 -- Requirements 3.2: Bash 工具捕获 stdout 和 stderr -- Requirements 3.3: Bash 工具支持超时控制 -- Requirements 3.4: Bash 工具设置防止交互的环境变量 -- Requirements 3.5: Bash 工具返回退出码和错误输出 -- Requirements 3.6: Bash 工具支持可配置的工作目录 -- Requirements 4.1: 文件读取工具返回带行号的内容 -- Requirements 4.2: 文件读取工具支持行范围读取 -- Requirements 4.3: 文件不存在时返回清晰错误信息 -- Requirements 4.4: 大文件推荐使用行范围 -- Requirements 4.5: 检测并报告文件的编程语言 -- Requirements 4.6: 路径为目录时列出目录内容 -- Requirements 5.1: 文件写入工具创建或覆盖文件 -- Requirements 5.2: 文件写入工具自动创建父目录 -- Requirements 5.3: 文件写入工具规范化换行符(Unix: LF, Windows: CRLF) -- Requirements 5.4: 文件写入工具确保文件以换行符结尾 -- Requirements 5.5: 写入失败时返回描述性错误信息 -- Requirements 6.1: 文件编辑工具精确替换匹配的字符串 -- Requirements 6.2: 多次出现时返回错误要求更多上下文 -- Requirements 6.3: 字符串不存在时返回错误和指导 -- Requirements 6.4: 支持 unified diff 格式 -- Requirements 6.5: 维护历史栈支持撤销操作 -- Requirements 6.6: 编辑后返回变更上下文片段 -- Requirements 8.1: 验证所有文件路径防止目录遍历攻击 -- Requirements 8.2: 拒绝包含 ".." 组件的路径 -- Requirements 8.3: 拒绝符号链接操作 -- Requirements 8.4: Bash 工具设置环境变量禁用交互式编辑器和提示 -- Requirements 8.5: 强制执行可配置的基础目录 -- Requirements 2.3: System Prompt 包含所有可用工具定义 - -## 更新提醒 - -任何文件变更后,请更新此文档和相关的上级文档。 diff --git a/src-tauri/src/agent/tools/browser.rs b/src-tauri/src/agent/tools/browser.rs deleted file mode 100644 index b61d0a48a..000000000 --- a/src-tauri/src/agent/tools/browser.rs +++ /dev/null @@ -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, - }, - /// 获取页面文本内容 - GetText { selector: Option }, - /// 执行 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, - /// 页面标题 - pub title: Option, - /// 截图 base64(如果有) - pub screenshot: Option, - /// 错误信息 - pub error: Option, -} - -/// 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 { - 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 { - 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")); - } -} diff --git a/src-tauri/src/agent/tools/mod.rs b/src-tauri/src/agent/tools/mod.rs deleted file mode 100644 index 970ef5cd1..000000000 --- a/src-tauri/src/agent/tools/mod.rs +++ /dev/null @@ -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) -> 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) -> ToolRegistry { - let mut registry = ToolRegistry::new(); - - // 使用 aster-rust 的默认工具注册 - let _shared_history = aster::tools::register_default_tools(&mut registry); - - info!( - "[Tools] 已创建最小工具注册表,共 {} 个工具", - registry.tool_count() - ); - - registry -} diff --git a/src-tauri/src/agent/tools/prompt.rs b/src-tauri/src/agent/tools/prompt.rs deleted file mode 100644 index ad2c6e90f..000000000 --- a/src-tauri/src/agent/tools/prompt.rs +++ /dev/null @@ -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("\n"); - for tool in tools { - prompt.push_str(&self.tool_to_xml(tool)); - prompt.push('\n'); - } - prompt.push_str("\n"); - - prompt - } - - /// 生成 JSON 格式的 System Prompt(OpenAI 风格) - fn generate_json_prompt(&self, tools: &[ToolDefinition]) -> String { - let mut prompt = String::new(); - - // 添加工具使用说明 - prompt.push_str(TOOL_USAGE_INSTRUCTIONS); - prompt.push_str("\n\n"); - - // 添加工具定义 - prompt.push_str("Available tools:\n```json\n"); - let tools_json = serde_json::to_string_pretty(tools).unwrap_or_else(|_| "[]".to_string()); - prompt.push_str(&tools_json); - prompt.push_str("\n```\n"); - - prompt - } - - /// 将单个工具定义转换为 XML 格式 - pub fn tool_to_xml(&self, tool: &ToolDefinition) -> String { - let mut xml = String::new(); - - xml.push_str(&format!("\n", escape_xml(&tool.name))); - xml.push_str(&format!( - " {}\n", - escape_xml(&tool.description) - )); - xml.push_str(" \n"); - xml.push_str(&self.json_schema_to_xml(&tool.parameters, 4)); - xml.push_str(" \n"); - xml.push_str(""); - - xml - } - - /// 将 JsonSchema 转换为 XML 格式 - fn json_schema_to_xml(&self, schema: &JsonSchema, indent: usize) -> String { - let mut xml = String::new(); - let indent_str = " ".repeat(indent); - - for (name, prop) in &schema.properties { - let required = if schema.required.contains(name) { - " required=\"true\"" - } else { - "" - }; - - xml.push_str(&format!( - "{}\n", - indent_str, - escape_xml(name), - escape_xml(&prop.prop_type), - required - )); - xml.push_str(&format!( - "{} {}\n", - indent_str, - escape_xml(&prop.description) - )); - - // 添加默认值(如果有) - if let Some(default) = &prop.default { - xml.push_str(&format!( - "{} {}\n", - indent_str, - escape_xml(&default.to_string()) - )); - } - - // 添加枚举值(如果有) - if let Some(enum_values) = &prop.enum_values { - xml.push_str(&format!("{} \n", indent_str)); - for value in enum_values { - xml.push_str(&format!( - "{} {}\n", - indent_str, - escape_xml(&value.to_string()) - )); - } - xml.push_str(&format!("{} \n", indent_str)); - } - - xml.push_str(&format!("{}\n", indent_str)); - } - - xml - } - - /// 将单个工具定义转换为 JSON 格式 - pub fn tool_to_json(&self, tool: &ToolDefinition) -> String { - serde_json::to_string_pretty(tool).unwrap_or_else(|_| "{}".to_string()) - } - - /// 获取当前格式 - pub fn format(&self) -> PromptFormat { - self.format - } -} - -/// 工具使用说明模板(适合桌面软件) -const TOOL_USAGE_INSTRUCTIONS: &str = r#"你是一个友好的 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 { - vec![ - create_test_tool(), - ToolDefinition::new("read_file", "Read the contents of a file").with_parameters( - JsonSchema::new() - .add_property( - "path", - PropertySchema::string("The file path to read"), - true, - ) - .add_property( - "start_line", - PropertySchema::integer("Starting line number (1-based)"), - false, - ) - .add_property( - "end_line", - PropertySchema::integer("Ending line number (inclusive)"), - false, - ), - ), - ] - } - - #[test] - fn test_generator_default_format() { - let generator = ToolPromptGenerator::new(); - assert_eq!(generator.format(), PromptFormat::Xml); - } - - #[test] - fn test_generator_with_format() { - let generator = ToolPromptGenerator::new().with_format(PromptFormat::Json); - assert_eq!(generator.format(), PromptFormat::Json); - } - - #[test] - fn test_tool_to_xml() { - let generator = ToolPromptGenerator::new(); - let tool = create_test_tool(); - let xml = generator.tool_to_xml(&tool); - - // 验证 XML 包含工具名称 - assert!(xml.contains("name=\"bash\"")); - // 验证 XML 包含描述 - assert!(xml.contains("Execute a bash command")); - // 验证 XML 包含必需参数 - assert!(xml.contains("name=\"command\"")); - assert!(xml.contains("required=\"true\"")); - // 验证 XML 包含可选参数 - assert!(xml.contains("name=\"timeout\"")); - // 验证 XML 包含默认值 - assert!(xml.contains("120")); - } - - #[test] - fn test_tool_to_json() { - let generator = ToolPromptGenerator::new(); - let tool = create_test_tool(); - let json = generator.tool_to_json(&tool); - - // 验证 JSON 包含工具名称 - assert!(json.contains("\"name\": \"bash\"")); - // 验证 JSON 包含描述 - assert!(json.contains("Execute a bash command")); - // 验证 JSON 包含参数 - assert!(json.contains("\"command\"")); - } - - #[test] - fn test_generate_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("")); - assert!(prompt.contains("")); - // 验证包含所有工具 - 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("