From e287fd3d392a1b6c07ec10486ade185cfc4b6ec3 Mon Sep 17 00:00:00 2001 From: coso Date: Sat, 7 Feb 2026 02:01:04 +0800 Subject: [PATCH] release: v0.58.0 --- docs/product-overview.md | 285 ++ docs/testing/skills-e2e-testing.md | 205 ++ package.json | 2 +- src-tauri/Cargo.lock | 89 +- src-tauri/Cargo.toml | 16 +- .../crates/core/src/models/provider_type.rs | 25 + src-tauri/src/agent/README.md | 39 +- src-tauri/src/agent/aster_agent.rs | 23 +- src-tauri/src/agent/aster_state.rs | 192 +- src-tauri/src/agent/credential_bridge.rs | 77 +- src-tauri/src/agent/mcp_bridge.rs | 46 + src-tauri/src/agent/mod.rs | 5 + src-tauri/src/agent/subagent_scheduler.rs | 283 ++ src-tauri/src/app/bootstrap.rs | 7 + src-tauri/src/app/runner.rs | 35 + src-tauri/src/commands/agent_cmd.rs | 32 +- src-tauri/src/commands/aster_agent_cmd.rs | 247 +- src-tauri/src/commands/external_tools_cmd.rs | 192 ++ src-tauri/src/commands/mcp_cmd.rs | 469 ++++ src-tauri/src/commands/mod.rs | 3 + src-tauri/src/commands/skill_cmd.rs | 13 +- src-tauri/src/commands/skill_exec_cmd.rs | 712 +++++ src-tauri/src/commands/subagent_cmd.rs | 100 + src-tauri/src/commands/switch_cmd.rs | 9 +- .../src/database/dao/brand_persona_dao.rs | 3 +- src-tauri/src/database/migration.rs | 49 + src-tauri/src/database/mod.rs | 12 + src-tauri/src/errors/mod.rs | 1 + .../src/flow_monitor/stream_rebuilder.rs | 9 - src-tauri/src/lib.rs | 6 + src-tauri/src/mcp/README.md | 42 + src-tauri/src/mcp/client.rs | 399 +++ src-tauri/src/mcp/manager.rs | 2426 +++++++++++++++++ src-tauri/src/mcp/mod.rs | 31 + src-tauri/src/mcp/tool_converter.rs | 180 ++ src-tauri/src/mcp/types.rs | 244 ++ src-tauri/src/models/mcp_model.rs | 143 + src-tauri/src/models/project_model.rs | 7 + src-tauri/src/services/aster_session_store.rs | 44 +- src-tauri/src/services/live_sync.rs | 69 +- src-tauri/src/services/mcp_service.rs | 59 + src-tauri/src/services/mcp_sync.rs | 10 +- src-tauri/src/services/switch.rs | 173 ++ src-tauri/src/skills/README.md | 96 + src-tauri/src/skills/execution_callback.rs | 320 +++ src-tauri/src/skills/llm_provider.rs | 597 ++++ src-tauri/src/skills/mod.rs | 32 + src-tauri/src/skills/skill_loader.rs | 231 ++ src-tauri/tauri.conf.json | 2 +- ...pi_key_provider_tests.proptest-regressions | 13 + src/App.tsx | 12 + src/components/AppSidebar.tsx | 22 +- src/components/README.md | 2 +- .../agent/chat/components/ChatSidebar.tsx | 6 +- .../agent/chat/hooks/skillCommand.ts | 531 ++++ .../agent/chat/hooks/skillSettings.ts | 98 + .../agent/chat/hooks/useAgentChat.ts | 41 +- .../agent/chat/hooks/useAsterAgentChat.ts | 54 +- src/components/api-server/ApiServerPage.tsx | 63 +- src/components/api-server/RoutesTab.tsx | 302 -- .../canvas/canvasUtils.test.ts | 35 +- src/components/mcp/McpPage.tsx | 16 +- src/components/mcp/McpPanel.tsx | 195 ++ src/components/mcp/McpPromptsBrowser.tsx | 325 +++ src/components/mcp/McpResourcesBrowser.tsx | 261 ++ src/components/mcp/McpServerList.tsx | 207 ++ src/components/mcp/McpToolCaller.tsx | 258 ++ src/components/mcp/McpToolsBrowser.tsx | 234 ++ src/components/mcp/README.md | 29 + src/components/mcp/index.ts | 6 + .../provider-pool/ProviderPoolPage.tsx | 2 +- .../settings/ExternalToolsSettings.tsx | 240 ++ src/components/settings/SettingsPage.tsx | 4 + src/components/settings/index.ts | 1 + src/components/skills/README.md | 37 + src/components/skills/SkillCard.tsx | 48 +- .../skills/SkillExecutionDialog.tsx | 490 ++++ src/components/skills/SkillsPage.tsx | 45 +- src/components/skills/WorkflowProgress.tsx | 330 +++ src/components/skills/index.ts | 10 + src/components/subagent/SubAgentProgress.tsx | 172 ++ src/components/subagent/index.ts | 13 + src/components/ui/sonner.tsx | 1 + src/components/vibe/VibePage.tsx | 556 ++++ src/components/vibe/index.ts | 5 + src/hooks/README.md | 73 + src/hooks/index.ts | 7 + src/hooks/useMcp.ts | 334 +++ src/hooks/useSkillExecution.ts | 346 +++ src/hooks/useSkills.ts | 2 +- src/hooks/useSubAgentScheduler.ts | 262 ++ src/hooks/useSwitch.ts | 25 +- src/icons/providers/index.tsx | 16 +- src/lib/api/agent.ts | 90 +- src/lib/api/externalTools.ts | 76 + src/lib/api/mcp.ts | 170 ++ src/lib/api/routes.ts | 37 - src/lib/api/skill-execution.ts | 271 ++ src/lib/api/skills.ts | 11 +- src/lib/tauri-mock/core.ts | 12 + src/types/page.ts | 2 + vite.config.ts | 19 + 102 files changed, 14162 insertions(+), 546 deletions(-) create mode 100644 docs/product-overview.md create mode 100644 docs/testing/skills-e2e-testing.md create mode 100644 src-tauri/src/agent/mcp_bridge.rs create mode 100644 src-tauri/src/agent/subagent_scheduler.rs create mode 100644 src-tauri/src/commands/external_tools_cmd.rs create mode 100644 src-tauri/src/commands/skill_exec_cmd.rs create mode 100644 src-tauri/src/commands/subagent_cmd.rs create mode 100644 src-tauri/src/mcp/README.md create mode 100644 src-tauri/src/mcp/client.rs create mode 100644 src-tauri/src/mcp/manager.rs create mode 100644 src-tauri/src/mcp/mod.rs create mode 100644 src-tauri/src/mcp/tool_converter.rs create mode 100644 src-tauri/src/mcp/types.rs create mode 100644 src-tauri/src/skills/README.md create mode 100644 src-tauri/src/skills/execution_callback.rs create mode 100644 src-tauri/src/skills/llm_provider.rs create mode 100644 src-tauri/src/skills/mod.rs create mode 100644 src-tauri/src/skills/skill_loader.rs create mode 100644 src-tauri/tests/api_key_provider_tests.proptest-regressions create mode 100644 src/components/agent/chat/hooks/skillCommand.ts create mode 100644 src/components/agent/chat/hooks/skillSettings.ts delete mode 100644 src/components/api-server/RoutesTab.tsx create mode 100644 src/components/mcp/McpPanel.tsx create mode 100644 src/components/mcp/McpPromptsBrowser.tsx create mode 100644 src/components/mcp/McpResourcesBrowser.tsx create mode 100644 src/components/mcp/McpServerList.tsx create mode 100644 src/components/mcp/McpToolCaller.tsx create mode 100644 src/components/mcp/McpToolsBrowser.tsx create mode 100644 src/components/mcp/README.md create mode 100644 src/components/settings/ExternalToolsSettings.tsx create mode 100644 src/components/skills/README.md create mode 100644 src/components/skills/SkillExecutionDialog.tsx create mode 100644 src/components/skills/WorkflowProgress.tsx create mode 100644 src/components/subagent/SubAgentProgress.tsx create mode 100644 src/components/subagent/index.ts create mode 100644 src/components/vibe/VibePage.tsx create mode 100644 src/components/vibe/index.ts create mode 100644 src/hooks/useMcp.ts create mode 100644 src/hooks/useSkillExecution.ts create mode 100644 src/hooks/useSubAgentScheduler.ts create mode 100644 src/lib/api/externalTools.ts delete mode 100644 src/lib/api/routes.ts create mode 100644 src/lib/api/skill-execution.ts diff --git a/docs/product-overview.md b/docs/product-overview.md new file mode 100644 index 000000000..a02bde669 --- /dev/null +++ b/docs/product-overview.md @@ -0,0 +1,285 @@ +# ProxyCast AI 创作工作站 - 产品介绍 + +> 版本: 1.0.0 +> 更新: 2026-02-04 +> 用途: 客户演示、产品介绍 + +--- + +## 一、产品定位 + +**中文创作者的本地 AI 工作站** + +核心理念:**AI 增强人,而非替代人** + +设计原则: +- **对话即创作** - 用自然语言描述需求,AI 理解意图 +- **一个对话,多种画布** - 同一对话可切换不同创作画布 +- **人机协作** - AI 建议透明可审查,用户掌控最终决策 + +--- + +## 二、六大创作画布 + +ProxyCast 提供 **6 种专业画布**,覆盖主流内容创作场景: + +| 画布类型 | 图标 | 适用场景 | 核心能力 | +|---------|------|---------|---------| +| **通用对话** | 💬 | 日常问答、头脑风暴 | 智能对话、知识问答 | +| **社媒内容** | 📱 | 公众号、小红书、知乎 | 6 步工作流、多平台适配 | +| **图文海报** | 🖼️ | 营销海报、社交图片 | 可视化设计、多尺寸导出 | +| **音乐歌词** | 🎵 | 歌词创作、简谱编曲 | 旋律学习、Suno 导出 | +| **短剧脚本** | 🎬 | 短视频、微短剧 | 场景管理、对白编辑 | +| **小说创作** | 📖 | 网文、长篇小说 | 章节管理、大纲规划 | + +--- + +## 三、各画布详细能力 + +### 3.1 📱 社媒内容画布 + +**6 步引导式创作流程**: + +``` +选题研究 → 竞品分析 → 大纲生成 → 初稿写作 → 多轮优化 → 平台发布 + 10% 20% 40% 70% 90% 100% +``` + +**多平台一键适配**: + +| 平台 | 特点 | 自动处理 | +|-----|------|---------| +| 公众号 | 深度长文 | 外链转二维码、排版优化 | +| 小红书 | 种草短文 | emoji 风格、话题标签 | +| 知乎 | 专业问答 | 引用来源、脚注格式 | +| 小说平台 | 章节连载 | 作者说、字数统计 | + +**专业 Agent 协作**: + +| Agent | 功能 | +|-------|------| +| 选题 Agent | 热点分析、选题推荐 | +| 标题 Agent | 爆款标题优化 | +| 开头 Agent | 吸睛开场设计 | +| 互动 Agent | 评论区引导 | +| 金句 Agent | 金句提取与优化 | + + + +### 3.2 🎵 音乐歌词画布 + +**支持歌曲类型**: + +| 类型 | 说明 | +|-----|------| +| 流行 (pop) | 主流流行音乐 | +| 民谣 (folk) | 抒情民谣 | +| 摇滚 (rock) | 摇滚乐 | +| 古风 (guofeng) | 中国风 | +| 说唱 (rap) | Hip-hop | +| R&B | 节奏蓝调 | +| 电子 (electronic) | 电子音乐 | + +**三种创作模式**: + +| 模式 | 说明 | 适合人群 | +|-----|------|---------| +| 教练模式 | AI 逐段引导创作 | 新手创作者 | +| 快速模式 | AI 直接生成完整歌词 | 追求效率 | +| 混合模式 | AI 生成框架,用户填充细节 | 专业创作者 | + +**四种视图模式**: +- 🎤 纯歌词视图 - 专注歌词编辑 +- 🎼 简谱视图 - 数字简谱展示 +- 🎸 吉他谱视图 - 和弦指法图 +- 🎹 钢琴谱视图 - 钢琴键位标注 + +**旋律学习功能**: +- 上传 MIDI/MP3 参考曲目 +- AI 分析旋律特征(调式、节奏、音程) +- 智能借鉴风格创作新曲 +- 一致性评分(结构、风格、旋律适配度) + +**导出格式**: +- PDF 歌词本 +- MIDI 文件 +- MusicXML +- **Suno 提示词** - 直接生成 AI 音乐 +- **Tunee 素材包** - 对话素材导出 + + + +### 3.3 🖼️ 图文海报画布 + +**基于 Fabric.js 的专业设计器**: +- 文字元素 - 多字体、多样式 +- 图片元素 - 裁剪、滤镜 +- 形状元素 - 矩形、圆形、线条 +- 背景元素 - 纯色、渐变、图片 + +**专业功能**: +- 图层管理 - 上移、下移、置顶、置底 +- 对齐工具 - 左对齐、居中、右对齐、分布 +- 多页面支持 - 批量设计 + +**预设尺寸**: + +| 平台 | 比例 | 像素 | +|-----|------|-----| +| 小红书封面 | 3:4 | 1080×1440 | +| 公众号头图 | 2.35:1 | 900×383 | +| 朋友圈 | 1:1 | 1080×1080 | +| 自定义 | 任意 | 自定义 | + +**导出格式**:PNG、JPEG、PDF + + + +### 3.4 🎬 短剧脚本画布 + +**专业剧本格式**: +- 场景管理 - 内景/外景、日/夜/晨/昏 +- 角色对白编辑 +- 表演指示标注(括号内) +- 情绪标记 + +**结构化编辑示例**: + +``` +第1场:咖啡厅(日) +*女主角坐在窗边,若有所思* + +女主:(叹气)为什么事情总是这样... +男主:(走近)你还好吗? +``` + +**场景元素**: +- 场景编号 +- 地点描述 +- 时间设定 +- 场景描述 +- 对白列表 + + + +### 3.5 📖 小说创作画布 + +**长篇创作支持**: +- 章节管理 - 拖拽排序、批量操作 +- 大纲树形结构 - 多级展开 +- 字数统计 - 章节/全书 +- 版本历史 - 随时回溯 + +**章节状态**: +- 草稿 (draft) - 创作中 +- 已完成 (completed) - 定稿 + +**创作辅助**: +- 世界观设定 +- 角色档案 +- 剧情线索追踪 +- AI 续写建议 + +--- + +## 四、通用能力 + +### 4.1 人设系统 + +**人设配置项**: +- 名称与简介 +- 写作风格描述 +- 语气设定 +- 目标读者画像 +- 禁用词列表 +- 偏好词列表 +- 示例文章(供 AI 学习) +- 适用平台 + +**使用方式**: +- 项目级默认人设 +- 话题级人设覆盖 +- 多人设快速切换 + + + +### 4.2 素材库 + +**支持素材类型**: +- 文档 (document) - PDF、Word、Markdown +- 图片 (image) - PNG、JPEG、GIF +- 文本 (text) - 纯文本片段 +- 数据 (data) - Excel、CSV +- 链接 (link) - 网页引用 + +**管理功能**: +- 标签分类 +- 描述备注 +- 预览查看 +- 写作时一键引用 + +### 4.3 项目管理 + +**层级关系**: +``` +项目 (Project) - 内容容器 + └── 话题 (Topic) - 内容载体 + └── 消息 (Message) - 对话记录 + └── 产出物 (Artifact) - 生成内容 +``` + +**项目类型**: +- general - 通用 +- social - 社媒内容 +- novel - 小说创作 +- drama - 短剧脚本 +- document - 办公文档 +- paper - 学术论文 +- music - 歌词曲谱 +- poster - 图文海报 + + + +--- + +## 五、产品亮点 + +| 特性 | 说明 | +|-----|------| +| **本地运行** | 数据安全,无需上传云端 | +| **多画布** | 6 种专业画布,覆盖主流场景 | +| **人机协作** | AI 建议透明可审查,用户掌控最终决策 | +| **多平台适配** | 一份内容,自动适配多个发布平台 | +| **专业导出** | 支持 Suno、Tunee 等 AI 音乐平台 | +| **项目化管理** | 人设/素材/排版 项目级复用 | + +--- + +## 六、目标用户 + +| 用户群体 | 典型场景 | +|---------|---------| +| 自媒体创作者 | 公众号、小红书、知乎日更 | +| 音乐创作者 | 歌词创作、编曲辅助 | +| 短剧编剧 | 微短剧、短视频脚本 | +| 网文作者 | 小说连载、大纲规划 | +| 设计师 | 营销海报、社交图片 | +| 内容运营 | 品牌文案、多平台分发 | + +--- + +## 七、技术架构(简述) + +- **前端**:React + TypeScript + Vite + TailwindCSS +- **后端**:Rust + Tauri +- **数据库**:SQLite(本地存储) +- **AI 框架**:集成 Aster-Rust Agent 框架 + +--- + +## 相关文档 + +- [社媒内容创作 PRD](prd/ai-content-creator.md) +- [SheMedia 工作流设计](prd/shemei/workflow.md) +- [画布系统架构](../src/components/content-creator/canvas/README.md) +- [统一内容系统](prd/unified-content-system.md) diff --git a/docs/testing/skills-e2e-testing.md b/docs/testing/skills-e2e-testing.md new file mode 100644 index 000000000..5afd222e4 --- /dev/null +++ b/docs/testing/skills-e2e-testing.md @@ -0,0 +1,205 @@ +# Skills 集成 E2E 测试指南 + +本文档指导如何手动进行 Skills 集成功能的端到端测试。 + +## 架构说明 + +ProxyCast 的 Skills 集成基于 aster-rust 框架的 `SkillTool`: + +``` +用户消息 → AI Agent → SkillTool → global_registry → 执行 Skill + ↑ + | + load_proxycast_skills() 加载 ~/.proxycast/skills/ +``` + +关键组件: +- `AsterAgentState::load_proxycast_skills()` - 启动时加载 Skills +- `AsterAgentState::reload_proxycast_skills()` - 安装/卸载后刷新 +- `aster::skills::global_registry()` - 全局 Skill 注册表 +- `aster::skills::SkillTool` - AI 调用 Skills 的工具 + +## 前置条件 + +1. ProxyCast 应用已构建并可运行 +2. 至少配置了一个可用的 AI Provider(如 OpenAI API Key) +3. 终端可以访问 `~/.proxycast/skills/` 目录 + +## 测试场景 + +### 场景 1:Skills 自动加载 + +**目的**:验证 Agent 初始化时能正确加载 Skills + +**步骤**: + +1. 创建测试 Skill: +```bash +mkdir -p ~/.proxycast/skills/test-greeting +cat > ~/.proxycast/skills/test-greeting/SKILL.md << 'EOF' +--- +name: test-greeting +description: 一个简单的问候技能,用于测试 Skills 集成 +--- + +# 问候技能 + +当用户请求问候时,使用以下格式回复: + +"你好!我是 ProxyCast 助手,很高兴为你服务!" + +请始终使用中文回复。 +EOF +``` + +2. 启动 ProxyCast 应用: +```bash +cd proxycast && npm run tauri dev +``` + +3. 打开开发者工具(Cmd+Option+I),查看控制台日志 + +4. **预期结果**: + - 日志中应显示 `[AsterAgent] 成功加载 1 个 ProxyCast Skills 到 global_registry` + - 日志中应显示 `[AsterAgent] 已注册 Skill: user:test-greeting` + +### 场景 2:AI 自动调用 Skill + +**目的**:验证 AI 能根据用户意图自动调用 Skill + +**步骤**: + +1. 确保测试 Skill 已创建(见场景 1) + +2. 在 ProxyCast 聊天界面发送消息: + ``` + 请用问候技能跟我打个招呼 + ``` + +3. **预期结果**: + - AI 应该识别到 `test-greeting` Skill + - AI 应该调用 Skill 并返回问候语 + - 响应中应包含 "你好!我是 ProxyCast 助手" + +### 场景 3:通过斜杠命令调用 Skill + +**目的**:验证用户可以通过 `/skill-name` 显式调用 Skill + +**步骤**: + +1. 在聊天界面发送: + ``` + /test-greeting + ``` + +2. **预期结果**: + - AI 应该直接执行 `test-greeting` Skill + - 返回问候语 + +### 场景 4:安装新 Skill 后动态刷新 + +**目的**:验证安装新 Skill 后 AI 能立即发现 + +**步骤**: + +1. 在 ProxyCast 运行时,创建新 Skill: +```bash +mkdir -p ~/.proxycast/skills/test-calculator +cat > ~/.proxycast/skills/test-calculator/SKILL.md << 'EOF' +--- +name: test-calculator +description: 一个简单的计算器技能 +--- + +# 计算器技能 + +当用户请求计算时,执行数学运算并返回结果。 + +支持:加法、减法、乘法、除法 +EOF +``` + +2. 在 ProxyCast Skills 页面点击刷新(或重新进入页面) + +3. 发送消息: + ``` + 请用计算器技能帮我算 123 + 456 + ``` + +4. **预期结果**: + - AI 应该能发现新安装的 `test-calculator` Skill + - AI 应该调用该 Skill 并返回计算结果 + +### 场景 5:卸载 Skill 后不再可用 + +**目的**:验证卸载 Skill 后 AI 不再能调用 + +**步骤**: + +1. 删除测试 Skill: +```bash +rm -rf ~/.proxycast/skills/test-greeting +``` + +2. 在 ProxyCast Skills 页面点击刷新 + +3. 发送消息: + ``` + /test-greeting + ``` + +4. **预期结果**: + - AI 应该提示找不到该 Skill + - 或者 AI 应该说明该 Skill 不可用 + +## 清理测试数据 + +测试完成后,清理测试 Skills: + +```bash +rm -rf ~/.proxycast/skills/test-greeting +rm -rf ~/.proxycast/skills/test-calculator +``` + +## 常见问题排查 + +### Skills 没有被加载 + +1. 检查目录是否存在:`ls -la ~/.proxycast/skills/` +2. 检查 SKILL.md 文件格式是否正确 +3. 查看应用日志中是否有错误信息 + +### AI 没有调用 Skill + +1. 确认 Skill 已被加载(查看启动日志) +2. 尝试使用更明确的指令,如 "使用 xxx 技能" +3. 检查 Skill 的 `description` 是否清晰描述了用途 + +### 动态刷新不生效 + +1. 确认调用了 `reload_proxycast_skills()` +2. 检查日志中是否有刷新相关的输出 +3. 尝试重启应用 + +## 自动化测试(未来计划) + +后续可以使用 Playwright 或 Tauri 的测试框架实现自动化 E2E 测试: + +```typescript +// 示例:Playwright E2E 测试 +test('AI should auto-invoke skill based on intent', async ({ page }) => { + // 1. 创建测试 Skill + await createTestSkill('test-greeting'); + + // 2. 启动应用 + await launchProxyCast(); + + // 3. 发送消息 + await page.fill('[data-testid="chat-input"]', '请用问候技能跟我打招呼'); + await page.click('[data-testid="send-button"]'); + + // 4. 验证响应 + await expect(page.locator('[data-testid="chat-message"]')) + .toContainText('你好!我是 ProxyCast 助手'); +}); +``` diff --git a/package.json b/package.json index 74d190eee..dc8fc3ab0 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.57.0", + "version": "0.58.0", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index d470c8f0e..a93ec697f 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -202,7 +202,7 @@ checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" [[package]] name = "aster" -version = "0.8.0" +version = "0.10.0" dependencies = [ "ahash", "anyhow", @@ -255,7 +255,7 @@ dependencies = [ "rand 0.8.5", "regex", "reqwest 0.12.28", - "rmcp", + "rmcp 0.12.0", "schemars 1.2.0", "scraper", "serde", @@ -2112,7 +2112,7 @@ dependencies = [ "dtoa-short", "itoa", "matches", - "phf 0.10.1", + "phf 0.8.0", "proc-macro2", "quote", "smallvec", @@ -2128,7 +2128,7 @@ dependencies = [ "cssparser-macros", "dtoa-short", "itoa", - "phf 0.11.3", + "phf 0.8.0", "smallvec", ] @@ -3988,7 +3988,7 @@ dependencies = [ "js-sys", "log", "wasm-bindgen", - "windows-core 0.57.0", + "windows-core 0.56.0", ] [[package]] @@ -5307,7 +5307,7 @@ version = "0.7.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ff32365de1b6743cb203b710788263c44a03de03802daf96092f2da4fe6ba4d7" dependencies = [ - "proc-macro-crate 2.0.2", + "proc-macro-crate 1.3.1", "proc-macro2", "quote", "syn 2.0.114", @@ -6024,7 +6024,9 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3dfb61232e34fcb633f43d12c58f83c1df82962dcdfa565a4e866ffc17dafe12" dependencies = [ + "phf_macros 0.8.0", "phf_shared 0.8.0", + "proc-macro-hack", ] [[package]] @@ -6033,9 +6035,7 @@ version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fabbf1ead8a5bcbc20f5f8b939ee3f5b0f6f281b6ad3468b84656b658b455259" dependencies = [ - "phf_macros 0.10.0", "phf_shared 0.10.0", - "proc-macro-hack", ] [[package]] @@ -6139,12 +6139,12 @@ dependencies = [ [[package]] name = "phf_macros" -version = "0.10.0" +version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "58fdf3184dd560f160dd73922bea2d5cd6e8f064bf4b13110abd81b03697b4e0" +checksum = "7f6fde18ff429ffc8fe78e2bf7f8b7a5a5a6e2a8b58bc5a9ac69198bbda9189c" dependencies = [ - "phf_generator 0.10.0", - "phf_shared 0.10.0", + "phf_generator 0.8.0", + "phf_shared 0.8.0", "proc-macro-hack", "proc-macro2", "quote", @@ -6495,6 +6495,20 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "process-wrap" +version = "8.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a3ef4f2f0422f23a82ec9f628ea2acd12871c81a9362b02c43c1aa86acfc3ba1" +dependencies = [ + "futures", + "indexmap 2.13.0", + "nix 0.30.1", + "tokio", + "tracing", + "windows 0.61.3", +] + [[package]] name = "process-wrap" version = "9.0.0" @@ -6545,7 +6559,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8a56d757972c98b346a9b766e3f02746cde6dd1cd1d1d563472929fdd74bec4d" dependencies = [ "anyhow", - "itertools 0.14.0", + "itertools 0.12.1", "proc-macro2", "quote", "syn 2.0.114", @@ -6553,7 +6567,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.57.0" +version = "0.58.0" dependencies = [ "anyhow", "arboard", @@ -6592,6 +6606,7 @@ dependencies = [ "rand 0.8.5", "regex", "reqwest 0.12.28", + "rmcp 0.6.4", "rusqlite", "rustls-pemfile 2.2.0", "scopeguard", @@ -6635,7 +6650,7 @@ dependencies = [ [[package]] name = "proxycast-core" -version = "0.57.0" +version = "0.58.0" dependencies = [ "chrono", "dirs 5.0.1", @@ -6651,7 +6666,7 @@ dependencies = [ [[package]] name = "proxycast-infra" -version = "0.57.0" +version = "0.58.0" dependencies = [ "chrono", "dashmap 5.5.3", @@ -7180,6 +7195,29 @@ dependencies = [ "windows-sys 0.52.0", ] +[[package]] +name = "rmcp" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41ab0892f4938752b34ae47cb53910b1b0921e55e77ddb6e44df666cab17939f" +dependencies = [ + "base64 0.22.1", + "chrono", + "futures", + "paste", + "pin-project-lite", + "process-wrap 8.2.1", + "rmcp-macros 0.6.4", + "schemars 1.2.0", + "serde", + "serde_json", + "thiserror 2.0.17", + "tokio", + "tokio-stream", + "tokio-util", + "tracing", +] + [[package]] name = "rmcp" version = "0.12.0" @@ -7194,9 +7232,9 @@ dependencies = [ "oauth2", "pastey", "pin-project-lite", - "process-wrap", + "process-wrap 9.0.0", "reqwest 0.12.28", - "rmcp-macros", + "rmcp-macros 0.12.0", "schemars 1.2.0", "serde", "serde_json", @@ -7209,6 +7247,19 @@ dependencies = [ "url", ] +[[package]] +name = "rmcp-macros" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1827cd98dab34cade0513243c6fe0351f0f0b2c9d6825460bcf45b42804bdda0" +dependencies = [ + "darling 0.21.3", + "proc-macro2", + "quote", + "serde_json", + "syn 2.0.114", +] + [[package]] name = "rmcp-macros" version = "0.12.0" @@ -7969,7 +8020,7 @@ version = "3.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8b1fdf65dd6331831494dd616b30351c38e96e45921a27745cf98490458b90bb" dependencies = [ - "dirs 6.0.0", + "dirs 4.0.0", ] [[package]] diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 811b2815d..f923f1e54 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -3,7 +3,7 @@ members = ["crates/*"] resolver = "2" [workspace.package] -version = "0.57.0" +version = "0.58.0" edition = "2021" authors = ["you"] repository = "https://github.com/aiclientproxy/proxycast" @@ -103,9 +103,12 @@ enigo = "0.3" # Aster Agent Framework # 开发时使用本地 aster-rust,CI/CD 使用远程 GitHub 仓库 # 本地开发: path = "../../../astercloud/aster-rust/crates/aster" (相对 src-tauri/) -# CI/CD: git = "https://github.com/astercloud/aster-rust", tag = "v0.7.1" -# aster = { version = "0.5.1", path = "../../../astercloud/aster-rust/crates/aster" } -aster = { git = "https://github.com/astercloud/aster-rust", tag = "v0.7.1" } +# CI/CD: git = "https://github.com/astercloud/aster-rust", tag = "v0.10.0" +aster = { version = "0.10.0", path = "../../../astercloud/aster-rust/crates/aster" } +# aster = { git = "https://github.com/astercloud/aster-rust", tag = "v0.10.0" } + +# MCP (Model Context Protocol) +rmcp = { version = "0.6", features = ["client", "transport-io", "transport-child-process"] } # Tauri @@ -164,7 +167,7 @@ version = "2.4" [package] name = "proxycast" -version = "0.57.0" +version = "0.58.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" @@ -278,6 +281,9 @@ cpal.workspace = true # Aster Agent Framework aster.workspace = true +# MCP (Model Context Protocol) +rmcp.workspace = true + # Windows specific dependencies for browser interceptor and machine ID management [target.'cfg(windows)'.dependencies] windows.workspace = true diff --git a/src-tauri/crates/core/src/models/provider_type.rs b/src-tauri/crates/core/src/models/provider_type.rs index cf8f8702f..f0c765483 100644 --- a/src-tauri/crates/core/src/models/provider_type.rs +++ b/src-tauri/crates/core/src/models/provider_type.rs @@ -90,6 +90,9 @@ impl std::str::FromStr for ProviderType { "siliconflow" => Ok(ProviderType::OpenAI), "oneapi" | "one-api" | "newapi" | "new-api" => Ok(ProviderType::OpenAI), "custom" | "custom_openai" => Ok(ProviderType::OpenAI), + // 自定义 Provider(UUID 格式,如 custom-ba4e7574-dd00-4784-945a-0f383dfa1272) + // 这些是用户通过 API Key Provider 添加的自定义服务 + s if s.starts_with("custom-") => Ok(ProviderType::OpenAI), _ => Err(format!("Unknown provider: {s}")), } } @@ -154,6 +157,28 @@ mod tests { assert!("invalid".parse::().is_err()); } + #[test] + fn test_custom_provider_uuid_format() { + // 自定义 Provider UUID 格式应该映射到 OpenAI + assert_eq!( + "custom-ba4e7574-dd00-4784-945a-0f383dfa1272" + .parse::() + .unwrap(), + ProviderType::OpenAI + ); + assert_eq!( + "custom-12345678-1234-1234-1234-123456789abc" + .parse::() + .unwrap(), + ProviderType::OpenAI + ); + // 普通 custom 也应该映射到 OpenAI + assert_eq!( + "custom".parse::().unwrap(), + ProviderType::OpenAI + ); + } + #[test] fn test_provider_type_display() { assert_eq!(ProviderType::Kiro.to_string(), "kiro"); diff --git a/src-tauri/src/agent/README.md b/src-tauri/src/agent/README.md index 1468a0a00..19a53454f 100644 --- a/src-tauri/src/agent/README.md +++ b/src-tauri/src/agent/README.md @@ -11,6 +11,7 @@ AI Agent 集成模块,基于 aster-rust 框架实现。 - **Aster 框架**:使用 aster-rust 框架获得多 Provider、工具系统、会话管理等能力 - **凭证池桥接**:自动从 ProxyCast 凭证池选择凭证配置 Aster Provider - **流式响应**:通过 Tauri 事件系统向前端推送流式内容 +- **Skills 集成**:自动加载 ProxyCast Skills 到 aster-rust,使 AI 能够自动调用 ## 文件索引 @@ -18,32 +19,56 @@ AI Agent 集成模块,基于 aster-rust 框架实现。 |------|------| | `mod.rs` | 模块入口,导出公共类型 | | `types.rs` | Agent 相关类型定义 | -| `aster_state.rs` | Aster Agent 状态管理(Provider 配置、取消令牌) | +| `aster_state.rs` | Aster Agent 状态管理(Provider 配置、取消令牌、Skills 加载) | | `aster_agent.rs` | Aster Agent 包装器(会话管理) | | `event_converter.rs` | Aster 事件到 Tauri 事件转换 | | `credential_bridge.rs` | 凭证池桥接(连接 ProxyCast 凭证池与 Aster Provider) | +## Skills 集成 + +### 自动加载机制 + +Agent 初始化时自动加载 `~/.proxycast/skills/` 目录下的 Skills: + +```rust +// init_agent_with_db() 内部调用 +Self::load_proxycast_skills(); +``` + +### AI 自动调用 + +aster-rust 的 `SkillTool` 会从 `global_registry` 读取可用 Skills,AI 可以: +- 根据用户意图自动选择合适的 Skill +- 通过 `/skill-name` 命令显式调用 Skill + +### 动态刷新 + +安装/卸载 Skills 后自动刷新: + +```rust +// skill_cmd.rs 中调用 +AsterAgentState::reload_proxycast_skills(); +``` + ## 使用方式 ### 从凭证池配置(推荐) ```rust -// 初始化 -state.init_agent().await?; +// 初始化(同时加载 Skills) +state.init_agent_with_db(&db).await?; // 从凭证池自动选择凭证并配置 Provider let config = state .configure_provider_from_pool(&db, "openai", "gpt-4", &session_id) .await?; - -// config.credential_uuid 包含使用的凭证 UUID ``` ### 手动配置 ```rust // 初始化 -state.init_agent().await?; +state.init_agent_with_db(&db).await?; // 手动配置 Provider let config = ProviderConfig { @@ -53,7 +78,7 @@ let config = ProviderConfig { base_url: None, credential_uuid: None, }; -state.configure_provider(config, &session_id).await?; +state.configure_provider(config, &session_id, &db).await?; ``` ### 发送消息 diff --git a/src-tauri/src/agent/aster_agent.rs b/src-tauri/src/agent/aster_agent.rs index 980c5fe02..a9ddc4774 100644 --- a/src-tauri/src/agent/aster_agent.rs +++ b/src-tauri/src/agent/aster_agent.rs @@ -144,15 +144,19 @@ impl AsterAgentWrapper { Ok(sessions .into_iter() - .map(|s| SessionInfo { - id: s.id, - name: s.title.unwrap_or_else(|| "未命名".to_string()), - created_at: chrono::DateTime::parse_from_rfc3339(&s.created_at) - .map(|dt| dt.timestamp()) - .unwrap_or(0), - updated_at: chrono::DateTime::parse_from_rfc3339(&s.updated_at) - .map(|dt| dt.timestamp()) - .unwrap_or(0), + .map(|s| { + let messages_count = AgentDao::get_message_count(&conn, &s.id).unwrap_or(0); + SessionInfo { + id: s.id, + name: s.title.unwrap_or_else(|| "未命名".to_string()), + created_at: chrono::DateTime::parse_from_rfc3339(&s.created_at) + .map(|dt| dt.timestamp()) + .unwrap_or(0), + updated_at: chrono::DateTime::parse_from_rfc3339(&s.updated_at) + .map(|dt| dt.timestamp()) + .unwrap_or(0), + messages_count, + } }) .collect()) } @@ -192,6 +196,7 @@ pub struct SessionInfo { pub name: String, pub created_at: i64, pub updated_at: i64, + pub messages_count: usize, } /// 会话详情(包含消息) diff --git a/src-tauri/src/agent/aster_state.rs b/src-tauri/src/agent/aster_state.rs index 43e81297d..5c2e4c72b 100644 --- a/src-tauri/src/agent/aster_state.rs +++ b/src-tauri/src/agent/aster_state.rs @@ -15,10 +15,16 @@ //! 包括名称、语言偏好、产品描述等。这是架构层面的正确做法, //! 而不是简单地追加提示词。 //! +//! ## Skills 集成 +//! +//! Agent 初始化时会自动加载 `~/.proxycast/skills/` 目录下的 Skills 到 +//! aster-rust 的 global_registry,使 AI 能够自动发现和调用这些 Skills。 +//! //! 参考文档:`docs/prd/chat-architecture-redesign.md` use aster::agents::{Agent, AgentIdentity, SessionConfig}; use aster::model::ModelConfig; +use aster::skills::{global_registry, load_skills_from_directory, SkillSource}; use std::sync::Arc; use tokio::sync::RwLock; use tokio_util::sync::CancellationToken; @@ -80,6 +86,7 @@ impl AsterAgentState { /// /// 创建 Agent 并注入 ProxyCastSessionStore,确保消息存储到 ProxyCast 数据库。 /// 同时设置 ProxyCast 专属的 Agent 身份(名称、语言、描述)。 + /// 自动加载 `~/.proxycast/skills/` 目录下的 Skills 到 aster-rust 的 global_registry。 /// /// **推荐使用此方法**而不是 `init_agent()`。 /// @@ -106,9 +113,12 @@ impl AsterAgentState { let identity = Self::create_proxycast_identity(); agent.set_identity(identity).await; + // 加载 ProxyCast Skills 到 aster-rust 的 global_registry + Self::load_proxycast_skills(); + *agent_guard = Some(agent); tracing::info!( - "[AsterAgent] Agent 初始化成功,已注入 ProxyCastSessionStore 和 ProxyCast 身份" + "[AsterAgent] Agent 初始化成功,已注入 ProxyCastSessionStore、ProxyCast 身份和 Skills" ); } else { tracing::debug!("[AsterAgent] Agent 已初始化,跳过"); @@ -116,6 +126,60 @@ impl AsterAgentState { Ok(()) } + /// 加载 ProxyCast Skills 到 aster-rust 的 global_registry + /// + /// 从 `~/.proxycast/skills/` 目录加载 Skills,使 AI 能够自动发现和调用。 + fn load_proxycast_skills() { + let home = match dirs::home_dir() { + Some(h) => h, + None => { + tracing::warn!("[AsterAgent] 无法获取 home 目录,跳过 Skills 加载"); + return; + } + }; + + let skills_dir = home.join(".proxycast").join("skills"); + if !skills_dir.exists() { + tracing::info!( + "[AsterAgent] ProxyCast Skills 目录不存在: {:?},跳过加载", + skills_dir + ); + return; + } + + // 从 ProxyCast skills 目录加载 Skills + let skills = load_skills_from_directory(&skills_dir, SkillSource::User); + let skill_count = skills.len(); + + if skill_count == 0 { + tracing::info!("[AsterAgent] ProxyCast Skills 目录为空,无 Skills 可加载"); + return; + } + + // 注册到 global_registry + let registry = global_registry(); + if let Ok(mut registry_guard) = registry.write() { + for skill in skills { + let skill_name = skill.skill_name.clone(); + registry_guard.register(skill); + tracing::debug!("[AsterAgent] 已注册 Skill: {}", skill_name); + } + tracing::info!( + "[AsterAgent] 成功加载 {} 个 ProxyCast Skills 到 global_registry", + skill_count + ); + } else { + tracing::error!("[AsterAgent] 无法获取 global_registry 写锁,Skills 加载失败"); + } + } + + /// 重新加载 ProxyCast Skills + /// + /// 当用户安装或卸载 Skills 后调用此方法刷新 registry。 + pub fn reload_proxycast_skills() { + Self::load_proxycast_skills(); + } + /// 创建 ProxyCast 专属的 Agent 身份配置 fn create_proxycast_identity() -> AgentIdentity { AgentIdentity::new("ProxyCast 助手") @@ -296,17 +360,29 @@ impl AsterAgentState { /// 设置 Provider 相关的环境变量 fn set_provider_env_vars(&self, config: &ProviderConfig) { + tracing::info!( + "[AsterAgent] set_provider_env_vars: provider_name={}, model_name={}, has_api_key={}, base_url={:?}", + config.provider_name, + config.model_name, + config.api_key.is_some(), + config.base_url + ); + // 根据 provider 类型设置对应的环境变量 let env_key = match config.provider_name.as_str() { "openai" => "OPENAI_API_KEY", "anthropic" => "ANTHROPIC_API_KEY", "google" => "GOOGLE_API_KEY", - "deepseek" | "custom_deepseek" => "DEEPSEEK_API_KEY", - "groq" => "GROQ_API_KEY", - "mistral" => "MISTRAL_API_KEY", + "deepseek" | "custom_deepseek" => "OPENAI_API_KEY", // DeepSeek 使用 OpenAI 兼容 API + "groq" => "OPENAI_API_KEY", // Groq 使用 OpenAI 兼容 API + "mistral" => "OPENAI_API_KEY", // Mistral 使用 OpenAI 兼容 API "openrouter" => "OPENROUTER_API_KEY", "ollama" => return, // Ollama 不需要 API Key _ => { + tracing::warn!( + "[AsterAgent] 未知的 provider_name: {}, 使用通用 OpenAI 格式", + config.provider_name + ); // 通用 OpenAI 兼容格式 if let Some(api_key) = &config.api_key { std::env::set_var("OPENAI_API_KEY", api_key); @@ -318,6 +394,8 @@ impl AsterAgentState { } }; + tracing::info!("[AsterAgent] 设置环境变量: {}=***", env_key); + if let Some(api_key) = &config.api_key { std::env::set_var(env_key, api_key); } @@ -336,6 +414,15 @@ impl AsterAgentState { self.current_provider_config.read().await.clone() } + /// 清除当前 Provider 配置 + /// + /// 用于切换凭证后重置状态,下次对话时会重新从凭证池选择凭证 + pub async fn clear_provider_config(&self) { + let mut config_guard = self.current_provider_config.write().await; + *config_guard = None; + tracing::info!("[AsterAgent] Provider 配置已清除"); + } + /// 检查 Provider 是否已配置 pub async fn is_provider_configured(&self) -> bool { self.current_provider_config.read().await.is_some() @@ -541,6 +628,8 @@ ProxyCast 是一个 AI 代理服务应用,帮助用户: #[cfg(test)] mod tests { use super::*; + use std::fs; + use tempfile::TempDir; #[tokio::test] async fn test_aster_state_init() { @@ -566,4 +655,99 @@ mod tests { state.remove_cancel_token(session_id).await; assert!(!state.cancel_session(session_id).await); } + + // ========================================================================= + // Skills 集成测试 + // ========================================================================= + + /// 测试辅助函数:创建测试用的 Skill 目录 + fn create_test_skill(skills_dir: &std::path::Path, skill_name: &str, description: &str) { + let skill_path = skills_dir.join(skill_name); + fs::create_dir_all(&skill_path).unwrap(); + let skill_md = format!( + r#"--- +name: {} +description: {} +--- + +# {} + +这是一个测试 Skill。 +"#, + skill_name, description, skill_name + ); + fs::write(skill_path.join("SKILL.md"), skill_md).unwrap(); + } + + /// 测试:load_skills_from_directory 能正确加载 Skills + #[test] + fn test_load_skills_from_directory() { + let temp_dir = TempDir::new().unwrap(); + let skills_dir = temp_dir.path(); + + // 创建测试 Skills + create_test_skill(skills_dir, "test-skill-1", "第一个测试技能"); + create_test_skill(skills_dir, "test-skill-2", "第二个测试技能"); + + // 加载 Skills + let skills = load_skills_from_directory(skills_dir, SkillSource::User); + + // 验证 + assert_eq!(skills.len(), 2); + let names: Vec<_> = skills.iter().map(|s| s.display_name.as_str()).collect(); + assert!(names.contains(&"test-skill-1")); + assert!(names.contains(&"test-skill-2")); + } + + /// 测试:空目录返回空列表 + #[test] + fn test_load_skills_empty_directory() { + let temp_dir = TempDir::new().unwrap(); + let skills = load_skills_from_directory(temp_dir.path(), SkillSource::User); + assert!(skills.is_empty()); + } + + /// 测试:不存在的目录返回空列表 + #[test] + fn test_load_skills_nonexistent_directory() { + let nonexistent = std::path::Path::new("/nonexistent/path/to/skills"); + let skills = load_skills_from_directory(nonexistent, SkillSource::User); + assert!(skills.is_empty()); + } + + /// 测试:global_registry 能正确注册和查找 Skills + #[test] + fn test_global_registry_register_and_find() { + let temp_dir = TempDir::new().unwrap(); + let skills_dir = temp_dir.path(); + + // 创建测试 Skill + create_test_skill(skills_dir, "registry-test-skill", "注册表测试技能"); + + // 加载并注册到 global_registry + let skills = load_skills_from_directory(skills_dir, SkillSource::User); + let registry = global_registry(); + + if let Ok(mut registry_guard) = registry.write() { + for skill in skills { + registry_guard.register(skill); + } + } + + // 验证能找到注册的 Skill + if let Ok(registry_guard) = registry.read() { + let found = registry_guard.find("registry-test-skill"); + assert!(found.is_some()); + assert_eq!(found.unwrap().display_name, "registry-test-skill"); + } + } + + /// 测试:reload_proxycast_skills 不会 panic(即使目录不存在) + #[test] + fn test_reload_proxycast_skills_no_panic() { + // 这个测试确保 reload_proxycast_skills 在各种情况下都不会 panic + // 即使 ~/.proxycast/skills/ 目录不存在 + AsterAgentState::reload_proxycast_skills(); + // 如果没有 panic,测试通过 + } } diff --git a/src-tauri/src/agent/credential_bridge.rs b/src-tauri/src/agent/credential_bridge.rs index f0e6e2ded..fc9f18d94 100644 --- a/src-tauri/src/agent/credential_bridge.rs +++ b/src-tauri/src/agent/credential_bridge.rs @@ -119,8 +119,9 @@ impl CredentialBridge { )) })?; - // 2. 转换为 Aster Provider 配置 - self.credential_to_config(&credential, model, db).await + // 2. 转换为 Aster Provider 配置,传递 provider_type 以便正确识别 Provider + self.credential_to_config(&credential, model, provider_type, db) + .await } /// 将 ProxyCast 凭证转换为 Aster Provider 配置 @@ -128,15 +129,31 @@ impl CredentialBridge { &self, credential: &ProviderCredential, model: &str, + provider_type_hint: &str, db: &DbConnection, ) -> Result { + tracing::info!( + "[CredentialBridge] credential_to_config: provider_type_hint={}, credential_type={:?}", + provider_type_hint, + credential.provider_type + ); + 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(), - ), + // OpenAI API Key - 根据 provider_type_hint 确定实际的 Provider + CredentialData::OpenAIKey { api_key, base_url } => { + // 使用 provider_type_hint 来确定 aster provider 名称 + let provider = map_provider_type_to_aster(provider_type_hint); + tracing::info!( + "[CredentialBridge] OpenAIKey: provider_type_hint={} -> aster_provider={}", + provider_type_hint, + provider + ); + ( + provider.to_string(), + Some(api_key.clone()), + base_url.clone(), + ) + } // Claude/Anthropic API Key CredentialData::ClaudeKey { api_key, base_url } @@ -343,6 +360,13 @@ pub async fn create_aster_provider( /// 设置 Provider 环境变量 fn set_provider_env_vars(config: &AsterProviderConfig) { + tracing::info!( + "[CredentialBridge] set_provider_env_vars: provider_name={}, has_api_key={}, base_url={:?}", + config.provider_name, + config.api_key.is_some(), + config.base_url + ); + let env_key = match config.provider_name.as_str() { "openai" => "OPENAI_API_KEY", "anthropic" => "ANTHROPIC_API_KEY", @@ -350,9 +374,15 @@ fn set_provider_env_vars(config: &AsterProviderConfig) { "bedrock" => "AWS_ACCESS_KEY_ID", // Bedrock 使用 AWS 凭证 "gcpvertexai" => "GOOGLE_API_KEY", "codex" => "OPENAI_API_KEY", // Codex 兼容 OpenAI - _ => "OPENAI_API_KEY", // 默认使用 OpenAI 格式 + "deepseek" | "custom_deepseek" => "OPENAI_API_KEY", // DeepSeek 使用 OpenAI 兼容 API + "groq" => "OPENAI_API_KEY", // Groq 使用 OpenAI 兼容 API + "mistral" => "OPENAI_API_KEY", // Mistral 使用 OpenAI 兼容 API + "openrouter" => "OPENROUTER_API_KEY", + _ => "OPENAI_API_KEY", // 默认使用 OpenAI 格式 }; + tracing::info!("[CredentialBridge] 设置环境变量: {}=***", env_key); + if let Some(api_key) = &config.api_key { std::env::set_var(env_key, api_key); } @@ -403,6 +433,35 @@ pub fn map_pool_type_to_aster(pool_type: &PoolProviderType) -> &'static str { } } +/// 将 provider_type 字符串映射到 Aster Provider 名称 +/// +/// 支持 60+ API Key Provider,包括 deepseek, moonshot, qwen 等 +fn map_provider_type_to_aster(provider_type: &str) -> &'static str { + match provider_type { + // 标准 Provider + "openai" => "openai", + "anthropic" | "claude" => "anthropic", + "google" | "gemini" => "google", + "bedrock" | "kiro" => "bedrock", + "gcpvertexai" | "vertex" => "gcpvertexai", + "codex" => "codex", + "azure" | "azure-openai" => "azure", + "ollama" => "ollama", + + // DeepSeek - 使用 openai 兼容 provider(Aster 会通过 alias 映射) + "deepseek" | "custom_deepseek" => "openai", + + // 其他 OpenAI 兼容 Provider - 使用 openai provider + // 这些 Provider 都使用 OpenAI 兼容 API,通过 base_url 区分 + "groq" => "openai", + "mistral" => "openai", + "openrouter" => "openrouter", + + // 默认使用 openai(OpenAI 兼容格式) + _ => "openai", + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/src-tauri/src/agent/mcp_bridge.rs b/src-tauri/src/agent/mcp_bridge.rs new file mode 100644 index 000000000..9447ff94e --- /dev/null +++ b/src-tauri/src/agent/mcp_bridge.rs @@ -0,0 +1,46 @@ +//! MCP 桥接客户端 +//! +//! 实现 Aster 的 McpClientTrait,将工具调用转发到 +//! ProxyCast 已有的 MCP RunningService,避免重复启动进程。 + +use aster::agents::mcp_client::{Error, McpClientTrait}; +use rmcp::model::{ + CallToolResult, GetPromptResult, InitializeResult, JsonObject, + ListPromptsResult, ListResourcesResult, ListToolsResult, + ReadResourceResult, ServerNotification, +}; +use rmcp::service::RunningService; +use rmcp::RoleClient; +use serde_json::Value; +use std::sync::Arc; +use tokio::sync::{mpsc, Mutex}; +use tokio_util::sync::CancellationToken; + +use crate::mcp::client::ProxyCastMcpClient; + +/// MCP 桥接客户端 +/// +/// 持有 ProxyCast 的 RunningService 引用, +/// 将 Aster 的工具调用转发到已有的 MCP 连接。 +pub struct McpBridgeClient { + /// 服务器名称 + name: String, + /// ProxyCast 的 rmcp RunningService + service: Arc>, + /// 服务器初始化信息 + server_info: Option, +} + +impl McpBridgeClient { + pub fn new( + name: String, + service: Arc>, + server_info: Option, + ) -> Self { + Self { + name, + service, + server_info, + } + } +} diff --git a/src-tauri/src/agent/mod.rs b/src-tauri/src/agent/mod.rs index ca48cfe0d..6bc5e9b1b 100644 --- a/src-tauri/src/agent/mod.rs +++ b/src-tauri/src/agent/mod.rs @@ -7,11 +7,13 @@ //! - aster_agent - Aster Agent 包装器 //! - event_converter - Aster 事件转换器 //! - credential_bridge - 凭证池桥接(连接 ProxyCast 凭证池与 Aster Provider) +//! - subagent_scheduler - SubAgent 调度器集成 pub mod aster_agent; pub mod aster_state; pub mod credential_bridge; pub mod event_converter; +pub mod subagent_scheduler; pub mod types; pub use aster_agent::{AsterAgentWrapper, SessionDetail, SessionInfo}; @@ -20,4 +22,7 @@ pub use credential_bridge::{ create_aster_provider, AsterProviderConfig, CredentialBridge, CredentialBridgeError, }; pub use event_converter::{convert_agent_event, TauriAgentEvent}; +pub use subagent_scheduler::{ + ProxyCastScheduler, ProxyCastSubAgentExecutor, SubAgentProgressEvent, +}; pub use types::*; diff --git a/src-tauri/src/agent/subagent_scheduler.rs b/src-tauri/src/agent/subagent_scheduler.rs new file mode 100644 index 000000000..127c23f5f --- /dev/null +++ b/src-tauri/src/agent/subagent_scheduler.rs @@ -0,0 +1,283 @@ +//! SubAgent 调度器集成 +//! +//! 将 aster-rust 的 SubAgent 调度器与 ProxyCast 凭证池集成 +//! +//! ## 功能 +//! - 自动从凭证池选择健康凭证 +//! - 支持凭证 fallback 策略 +//! - 集成 Tauri 事件系统进行进度通知 + +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use aster::agents::context::AgentContext; +use aster::agents::subagent_scheduler::{ + SchedulerConfig, SchedulerError, SchedulerExecutionResult, SchedulerResult, SubAgentExecutor, + SubAgentResult, SubAgentScheduler, SubAgentTask, TokenUsage as SchedulerTokenUsage, +}; +use aster::conversation::message::Message; +use chrono::Utc; +use tauri::{AppHandle, Emitter}; +use tokio::sync::RwLock; +use tracing::{debug, info, warn}; + +use crate::agent::credential_bridge::{ + create_aster_provider, AsterProviderConfig, CredentialBridge, +}; +use crate::database::DbConnection; + +/// ProxyCast SubAgent 执行器 +/// +/// 实现 aster-rust 的 SubAgentExecutor trait, +/// 集成 ProxyCast 凭证池进行 LLM 调用 +pub struct ProxyCastSubAgentExecutor { + /// 凭证桥接器 + credential_bridge: CredentialBridge, + /// 数据库连接 + db: DbConnection, + /// 默认模型 + default_model: String, + /// 默认 Provider 类型 + default_provider: String, + /// Tauri AppHandle(用于事件通知) + app_handle: Option, +} + +impl ProxyCastSubAgentExecutor { + /// 创建新的执行器 + pub fn new(db: DbConnection) -> Self { + Self { + credential_bridge: CredentialBridge::new(), + db, + default_model: "claude-sonnet-4-20250514".to_string(), + default_provider: "anthropic".to_string(), + app_handle: None, + } + } + + /// 设置 Tauri AppHandle + pub fn with_app_handle(mut self, handle: AppHandle) -> Self { + self.app_handle = Some(handle); + self + } + + /// 设置默认模型 + pub fn with_default_model(mut self, model: impl Into) -> Self { + self.default_model = model.into(); + self + } + + /// 设置默认 Provider + pub fn with_default_provider(mut self, provider: impl Into) -> Self { + self.default_provider = provider.into(); + self + } + + /// 从凭证池选择凭证 + async fn select_credential(&self, task: &SubAgentTask) -> SchedulerResult { + // 根据任务类型和模型选择 provider + let model = task.model.as_deref().unwrap_or(&self.default_model); + let provider_type = &self.default_provider; + + // 使用 CredentialBridge 选择凭证 + let config = self + .credential_bridge + .select_and_configure(&self.db, provider_type, model) + .await + .map_err(|e| SchedulerError::ProviderError(e.to_string()))?; + + Ok(config) + } + + /// 发送 Tauri 事件 + #[allow(dead_code)] + fn emit_event(&self, event_name: &str, payload: impl serde::Serialize + Clone) { + if let Some(handle) = &self.app_handle { + if let Err(e) = handle.emit(event_name, payload) { + warn!("发送 Tauri 事件失败: {}", e); + } + } + } +} + +#[async_trait::async_trait] +impl SubAgentExecutor for ProxyCastSubAgentExecutor { + async fn execute_task( + &self, + task: &SubAgentTask, + context: &AgentContext, + ) -> SchedulerResult { + let start_time = Utc::now(); + info!("执行 SubAgent 任务: {}", task.id); + + // 选择凭证 + let provider_config = self.select_credential(task).await?; + debug!("使用凭证: {}", provider_config.credential_uuid); + + // 创建 provider + let provider = create_aster_provider(&provider_config) + .await + .map_err(|e| SchedulerError::ProviderError(e.to_string()))?; + + // 构建提示 + let system_prompt = context.system_prompt.clone().unwrap_or_default(); + let user_message = Message::user().with_text(&task.prompt); + + // 调用 LLM(使用 complete 方法) + let (response_msg, usage) = provider + .complete(&system_prompt, &[user_message], &[]) + .await + .map_err(|e| SchedulerError::ProviderError(e.to_string()))?; + + let response = response_msg.as_concat_text(); + + let end_time = Utc::now(); + let duration = (end_time - start_time).to_std().unwrap_or(Duration::ZERO); + + // 生成摘要 + let summary = if task.return_summary { + Some(self.generate_summary(&response, task)) + } else { + None + }; + + // 转换 token 使用 + let token_usage = Some(SchedulerTokenUsage { + input_tokens: usage.usage.input_tokens.unwrap_or(0) as usize, + output_tokens: usage.usage.output_tokens.unwrap_or(0) as usize, + total_tokens: usage.usage.total_tokens.unwrap_or(0) as usize, + }); + + Ok(SubAgentResult { + task_id: task.id.clone(), + success: true, + output: Some(response), + summary, + error: None, + duration, + retries: 0, + started_at: start_time, + completed_at: end_time, + token_usage, + metadata: HashMap::new(), + }) + } +} + +impl ProxyCastSubAgentExecutor { + /// 生成摘要 + fn generate_summary(&self, output: &str, task: &SubAgentTask) -> String { + // 简单摘要:取前 500 字符 + let max_len = 500; + if output.chars().count() <= max_len { + format!("任务 {} 完成:\n{}", task.id, output) + } else { + let truncated: String = output.chars().take(max_len - 3).collect(); + format!("任务 {} 完成:\n{}...", task.id, truncated) + } + } +} + +/// ProxyCast SubAgent 调度器包装器 +pub struct ProxyCastScheduler { + /// 内部调度器 + scheduler: Arc>>>, + /// 数据库连接 + db: DbConnection, + /// Tauri AppHandle + app_handle: Option, +} + +impl ProxyCastScheduler { + /// 创建新的调度器 + pub fn new(db: DbConnection) -> Self { + Self { + scheduler: Arc::new(RwLock::new(None)), + db, + app_handle: None, + } + } + + /// 设置 Tauri AppHandle + pub fn with_app_handle(mut self, handle: AppHandle) -> Self { + self.app_handle = Some(handle); + self + } + + /// 初始化调度器 + pub async fn init(&self, config: Option) { + let executor = ProxyCastSubAgentExecutor::new(self.db.clone()); + let executor = if let Some(handle) = &self.app_handle { + executor.with_app_handle(handle.clone()) + } else { + executor + }; + + let config = config.unwrap_or_default(); + + // 创建调度器并设置事件回调 + let app_handle = self.app_handle.clone(); + let scheduler = + SubAgentScheduler::new(config, executor).with_event_callback(move |event| { + if let Some(handle) = &app_handle { + let _ = handle.emit("subagent-scheduler-event", &event); + } + }); + + *self.scheduler.write().await = Some(scheduler); + info!("ProxyCast SubAgent 调度器初始化完成"); + } + + /// 执行任务 + pub async fn execute( + &self, + tasks: Vec, + parent_context: Option<&AgentContext>, + ) -> SchedulerResult { + let scheduler = self.scheduler.read().await; + let scheduler = scheduler + .as_ref() + .ok_or_else(|| SchedulerError::ContextError("调度器未初始化".to_string()))?; + + scheduler.execute(tasks, parent_context).await + } + + /// 取消执行 + pub async fn cancel(&self) { + if let Some(scheduler) = self.scheduler.read().await.as_ref() { + scheduler.cancel().await; + } + } +} + +/// Tauri 事件:SubAgent 进度 +#[derive(Debug, Clone, serde::Serialize)] +#[serde(rename_all = "camelCase")] +pub struct SubAgentProgressEvent { + /// 总任务数 + pub total: usize, + /// 已完成数 + pub completed: usize, + /// 失败数 + pub failed: usize, + /// 运行中数 + pub running: usize, + /// 进度百分比 + pub percentage: f64, + /// 当前任务 + pub current_tasks: Vec, +} + +impl From for SubAgentProgressEvent { + fn from(p: aster::agents::subagent_scheduler::SchedulerProgress) -> Self { + Self { + total: p.total, + completed: p.completed, + failed: p.failed, + running: p.running, + percentage: p.percentage, + current_tasks: p.current_tasks, + } + } +} diff --git a/src-tauri/src/app/bootstrap.rs b/src-tauri/src/app/bootstrap.rs index ac397caaf..331b6ebe8 100644 --- a/src-tauri/src/app/bootstrap.rs +++ b/src-tauri/src/app/bootstrap.rs @@ -34,6 +34,7 @@ use crate::flow_monitor::{ QuickFilterManager, RotationConfig, SessionManager, }; use crate::logger; +use crate::mcp::McpManagerState; use crate::plugin; use crate::server; use crate::services::api_key_provider_service::ApiKeyProviderService; @@ -151,6 +152,7 @@ pub struct AppStates { pub context_memory_service: ContextMemoryServiceState, pub tool_hooks_service: ToolHooksServiceState, pub recording_service: RecordingServiceState, + pub mcp_manager: McpManagerState, // 用于 setup hook 的共享实例 pub shared_stats: Arc>, pub shared_tokens: Arc>, @@ -288,6 +290,10 @@ pub fn init_states(config: &Config) -> Result { // 录音服务(使用独立线程 + channel 通信解决 cpal::Stream 不是 Send 的问题) let recording_service_state = create_recording_service_state(); + // 初始化 MCP 客户端管理器(延迟设置 AppHandle,在 setup hook 中完成) + let mcp_manager = crate::mcp::McpClientManager::new(None); + let mcp_manager_state: McpManagerState = Arc::new(tokio::sync::Mutex::new(mcp_manager)); + Ok(AppStates { state, logs, @@ -324,6 +330,7 @@ pub fn init_states(config: &Config) -> Result { context_memory_service: context_memory_service_state, tool_hooks_service: tool_hooks_service_state, recording_service: recording_service_state, + mcp_manager: mcp_manager_state, shared_stats, shared_tokens, shared_logger, diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index cd204c811..572bbdf0d 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -82,6 +82,7 @@ pub fn run() { context_memory_service, tool_hooks_service, recording_service, + mcp_manager: mcp_manager_state, shared_stats, shared_tokens, shared_logger, @@ -166,6 +167,7 @@ pub fn run() { .manage(context_memory_service) .manage(tool_hooks_service) .manage(recording_service) + .manage(mcp_manager_state) .on_window_event(move |window, event| { // 处理窗口关闭事件 if let tauri::WindowEvent::CloseRequested { api, .. } = event { @@ -227,6 +229,16 @@ pub fn run() { tracing::info!("[启动] GlobalConfigManager AppHandle 已设置"); } + // 设置 MCP Manager 的 AppHandle(用于发送 mcp:* 事件) + if let Some(mcp_manager) = app.try_state::() { + let app_handle = app.handle().clone(); + tauri::async_runtime::block_on(async { + let mut manager = mcp_manager.lock().await; + manager.set_app_handle(app_handle); + }); + tracing::info!("[启动] MCP Manager AppHandle 已设置"); + } + // 初始化截图对话模块 // _Requirements: 7.3_ { @@ -728,6 +740,19 @@ pub fn run() { commands::mcp_cmd::toggle_mcp_server, commands::mcp_cmd::import_mcp_from_app, commands::mcp_cmd::sync_all_mcp_to_live, + // MCP 生命周期管理命令 + commands::mcp_cmd::mcp_list_servers_with_status, + commands::mcp_cmd::mcp_start_server, + commands::mcp_cmd::mcp_stop_server, + // MCP 工具管理命令 + commands::mcp_cmd::mcp_list_tools, + commands::mcp_cmd::mcp_call_tool, + // MCP 提示词管理命令 + commands::mcp_cmd::mcp_list_prompts, + commands::mcp_cmd::mcp_get_prompt, + // MCP 资源管理命令 + commands::mcp_cmd::mcp_list_resources, + commands::mcp_cmd::mcp_read_resource, // Prompt commands commands::prompt_cmd::get_prompts, commands::prompt_cmd::upsert_prompt, @@ -750,6 +775,10 @@ pub fn run() { commands::skill_cmd::add_skill_repo, commands::skill_cmd::remove_skill_repo, commands::skill_cmd::get_installed_proxycast_skills, + // Skill Execution commands + commands::skill_exec_cmd::execute_skill, + commands::skill_exec_cmd::list_executable_skills, + commands::skill_exec_cmd::get_skill_detail, // Provider Pool commands commands::provider_pool_cmd::get_provider_pool_overview, commands::provider_pool_cmd::get_provider_pool_credentials, @@ -1069,6 +1098,7 @@ pub fn run() { // Aster Agent commands commands::aster_agent_cmd::aster_agent_init, commands::aster_agent_cmd::aster_agent_status, + commands::aster_agent_cmd::aster_agent_reset, commands::aster_agent_cmd::aster_agent_configure_provider, commands::aster_agent_cmd::aster_agent_configure_from_pool, commands::aster_agent_cmd::aster_agent_chat_stream, @@ -1349,6 +1379,11 @@ pub fn run() { commands::asr_cmd::delete_asr_credential, commands::asr_cmd::set_default_asr_credential, commands::asr_cmd::test_asr_credential, + // External Tools commands (Codex CLI 等外部工具) + commands::external_tools_cmd::check_codex_cli_status, + commands::external_tools_cmd::open_codex_cli_login, + commands::external_tools_cmd::open_codex_cli_logout, + commands::external_tools_cmd::get_external_tools, // Voice Input commands crate::voice::commands::get_voice_input_config, crate::voice::commands::save_voice_input_config, diff --git a/src-tauri/src/commands/agent_cmd.rs b/src-tauri/src/commands/agent_cmd.rs index f575d6728..9c9b26351 100644 --- a/src-tauri/src/commands/agent_cmd.rs +++ b/src-tauri/src/commands/agent_cmd.rs @@ -10,6 +10,24 @@ use crate::AppState; use serde::{Deserialize, Serialize}; use tauri::State; +/// 安全截断字符串,确保不会在多字节字符中间切割 +/// +/// # 参数 +/// - `s`: 要截断的字符串 +/// - `max_chars`: 最大字符数(按 Unicode 字符计算,非字节) +/// +/// # 返回 +/// 截断后的字符串,如果被截断则添加 "..." 后缀 +fn truncate_string(s: &str, max_chars: usize) -> String { + let char_count = s.chars().count(); + if char_count <= max_chars { + s.to_string() + } else { + let truncated: String = s.chars().take(max_chars).collect(); + format!("{}...", truncated) + } +} + /// Agent 进程状态响应 #[derive(Debug, Serialize)] pub struct AgentProcessStatus { @@ -360,11 +378,8 @@ pub async fn agent_generate_title( "助手" }; let content = msg.content.as_text(); - let truncated_content = if content.len() > 100 { - format!("{}...", &content[..100]) - } else { - content - }; + // 使用字符边界安全截断,避免在多字节字符中间切割 + let truncated_content = truncate_string(&content, 100); conversation.push_str(&format!("{role}:{truncated_content}\n")); } @@ -372,11 +387,8 @@ pub async fn agent_generate_title( // 这里简化处理:使用第一条用户消息的前 15 个字作为默认标题 if let Some(first_user_msg) = chat_messages.iter().find(|msg| msg.role == "user") { let content = first_user_msg.content.as_text(); - let title = if content.len() > 15 { - format!("{}...", &content[..15]) - } else { - content - }; + // 使用字符边界安全截断 + let title = truncate_string(&content, 15); Ok(title) } else { Ok("新话题".to_string()) diff --git a/src-tauri/src/commands/aster_agent_cmd.rs b/src-tauri/src/commands/aster_agent_cmd.rs index fa197683e..bdcf57c95 100644 --- a/src-tauri/src/commands/aster_agent_cmd.rs +++ b/src-tauri/src/commands/aster_agent_cmd.rs @@ -11,6 +11,9 @@ use crate::agent::{ }; use crate::database::dao::agent::AgentDao; use crate::database::DbConnection; +use crate::mcp::{McpManagerState, McpServerConfig}; +use crate::services::mcp_service::McpService; +use aster::agents::extension::{Envs, ExtensionConfig}; use aster::conversation::message::Message; use futures::StreamExt; use serde::{Deserialize, Serialize}; @@ -155,6 +158,28 @@ pub async fn aster_agent_status( }) } +/// 重置 Aster Agent +/// +/// 清除当前 Provider 配置,下次对话时会重新从凭证池选择凭证。 +/// 用于切换凭证后无需重启应用即可生效。 +#[tauri::command] +pub async fn aster_agent_reset( + state: State<'_, AsterAgentState>, +) -> Result { + tracing::info!("[AsterAgent] 重置 Agent Provider 配置"); + + // 清除当前 Provider 配置 + state.clear_provider_config().await; + + Ok(AsterAgentStatus { + initialized: state.is_initialized().await, + provider_configured: false, + provider_name: None, + model_name: None, + credential_uuid: None, + }) +} + /// 发送消息请求参数 #[derive(Debug, Deserialize)] pub struct AsterChatRequest { @@ -186,6 +211,7 @@ pub async fn aster_agent_chat_stream( app: AppHandle, state: State<'_, AsterAgentState>, db: State<'_, DbConnection>, + mcp_manager: State<'_, McpManagerState>, request: AsterChatRequest, ) -> Result<(), String> { tracing::info!( @@ -217,6 +243,23 @@ pub async fn aster_agent_chat_stream( // 同时 get_session 也会自动创建不存在的 session let session_id = &request.session_id; + // 启动并注入 MCP extensions 到 Aster Agent + let (_start_ok, start_fail) = ensure_proxycast_mcp_servers_running(&db, &mcp_manager).await; + if start_fail > 0 { + tracing::warn!( + "[AsterAgent] 部分 MCP server 自动启动失败 ({} 失败),后续可用工具可能不完整", + start_fail + ); + } + + let (_mcp_ok, mcp_fail) = inject_mcp_extensions(&state, &mcp_manager).await; + if mcp_fail > 0 { + tracing::warn!( + "[AsterAgent] 部分 MCP extension 注入失败 ({} 失败),Agent 可能无法使用某些 MCP 工具", + mcp_fail + ); + } + // 构建 system_prompt:优先使用项目上下文,其次使用 session 的 system_prompt let system_prompt = { let db_conn = db.lock().map_err(|e| format!("获取数据库连接失败: {e}"))?; @@ -276,6 +319,13 @@ pub async fn aster_agent_chat_stream( // 如果提供了 Provider 配置,则配置 Provider if let Some(provider_config) = &request.provider_config { + tracing::info!( + "[AsterAgent] 收到 provider_config: provider_name={}, model_name={}, has_api_key={}, base_url={:?}", + provider_config.provider_name, + provider_config.model_name, + provider_config.api_key.is_some(), + provider_config.base_url + ); let config = ProviderConfig { provider_name: provider_config.provider_name.clone(), model_name: provider_config.model_name.clone(), @@ -283,7 +333,20 @@ pub async fn aster_agent_chat_stream( base_url: provider_config.base_url.clone(), credential_uuid: None, }; - state.configure_provider(config, session_id, &db).await?; + // 如果前端提供了 api_key,直接使用;否则从凭证池选择凭证 + if provider_config.api_key.is_some() { + state.configure_provider(config, session_id, &db).await?; + } else { + // 没有 api_key,使用凭证池(provider_name 作为 provider_type) + state + .configure_provider_from_pool( + &db, + &provider_config.provider_name, + &provider_config.model_name, + session_id, + ) + .await?; + } } // 检查 Provider 是否已配置 @@ -297,7 +360,7 @@ pub async fn aster_agent_chat_stream( // 创建用户消息 let user_message = Message::user().with_text(&request.message); - // 创建会话配置,包含 system_prompt + // 创建会话配置 let mut session_config_builder = SessionConfigBuilder::new(session_id); if let Some(prompt) = system_prompt { session_config_builder = session_config_builder.system_prompt(prompt); @@ -451,3 +514,183 @@ mod tests { assert_eq!(request.event_name, "agent_stream"); } } + +/// 将 ProxyCast 已运行的 MCP servers 注入到 Aster Agent 作为 extensions +/// +/// 获取 McpClientManager 中所有已运行的 server 配置, +/// 转换为 Aster 的 ExtensionConfig::Stdio 并注册到 Agent。 +/// +/// 关键:将当前进程的 PATH 等环境变量合并到 MCP server 的 env 中, +/// 确保 Aster 启动的子进程能找到 npx/uvx 等命令。 +/// +/// 返回 (成功数, 失败数) +async fn inject_mcp_extensions( + state: &AsterAgentState, + mcp_manager: &McpManagerState, +) -> (usize, usize) { + let manager = mcp_manager.lock().await; + let running_servers = manager.get_running_servers().await; + + if running_servers.is_empty() { + tracing::debug!("[AsterAgent] 没有运行中的 MCP servers,跳过注入"); + return (0, 0); + } + + let agent_arc = state.get_agent_arc(); + let guard = agent_arc.read().await; + let agent = match guard.as_ref() { + Some(a) => a, + None => { + tracing::warn!("[AsterAgent] Agent 未初始化,无法注入 MCP extensions"); + return (0, running_servers.len()); + } + }; + + let mut success_count = 0usize; + let mut fail_count = 0usize; + + for server_name in &running_servers { + // 检查是否已注册(避免重复注册) + let ext_configs = agent.get_extension_configs().await; + if ext_configs.iter().any(|c| c.name() == *server_name) { + tracing::debug!("[AsterAgent] MCP extension '{}' 已注册,跳过", server_name); + success_count += 1; + continue; + } + + if let Some(config) = manager.get_client_config(server_name).await { + // 合并当前进程的关键环境变量到 MCP server 的 env 中 + // 确保子进程能找到 npx/uvx/node 等命令 + let mut merged_env = config.env.clone(); + for key in &["PATH", "HOME", "USER", "SHELL", "NODE_PATH", "NVM_DIR"] { + if !merged_env.contains_key(*key) { + if let Ok(val) = std::env::var(key) { + merged_env.insert(key.to_string(), val); + } + } + } + + tracing::info!( + "[AsterAgent] 注入 MCP extension '{}': cmd='{}', args={:?}, env_keys={:?}", + server_name, + config.command, + config.args, + merged_env.keys().collect::>() + ); + + // 增加超时时间:npx 首次下载可能需要较长时间 + let timeout = std::cmp::max(config.timeout, 60); + + let extension = ExtensionConfig::Stdio { + name: server_name.clone(), + description: format!("MCP Server: {server_name}"), + cmd: config.command.clone(), + args: config.args.clone(), + envs: Envs::new(merged_env), + env_keys: vec![], + timeout: Some(timeout), + bundled: Some(false), + available_tools: vec![], + }; + + match agent.add_extension(extension).await { + Ok(_) => { + tracing::info!("[AsterAgent] 成功注入 MCP extension: {}", server_name); + success_count += 1; + } + Err(e) => { + tracing::error!( + "[AsterAgent] 注入 MCP extension '{}' 失败: {}。\ + cmd='{}', args={:?}。请检查命令是否在 PATH 中可用。", + server_name, + e, + config.command, + config.args + ); + fail_count += 1; + } + } + } else { + tracing::warn!("[AsterAgent] 无法获取 MCP server '{}' 的配置", server_name); + fail_count += 1; + } + } + + if fail_count > 0 { + tracing::warn!( + "[AsterAgent] MCP 注入结果: {} 成功, {} 失败", + success_count, + fail_count + ); + } else { + tracing::info!( + "[AsterAgent] MCP 注入完成: {} 个 extension 全部成功", + success_count + ); + } + + (success_count, fail_count) +} + +/// 确保 ProxyCast 可用的 MCP servers 已启动 +/// +/// 启动启用了 `enabled_proxycast` 的服务器。 +async fn ensure_proxycast_mcp_servers_running( + db: &DbConnection, + mcp_manager: &McpManagerState, +) -> (usize, usize) { + let servers = match McpService::get_all(db) { + Ok(items) => items, + Err(e) => { + tracing::warn!("[AsterAgent] 读取 MCP 配置失败,跳过自动启动: {}", e); + return (0, 0); + } + }; + + if servers.is_empty() { + return (0, 0); + } + + let candidates: Vec<&crate::models::McpServer> = + servers.iter().filter(|s| s.enabled_proxycast).collect(); + + if candidates.is_empty() { + return (0, 0); + } + + let manager = mcp_manager.lock().await; + let mut success_count = 0usize; + let mut fail_count = 0usize; + + for server in candidates { + if manager.is_server_running(&server.name).await { + continue; + } + + let parsed = server.parse_config(); + let config = McpServerConfig { + command: parsed.command, + args: parsed.args, + env: parsed.env, + cwd: parsed.cwd, + timeout: parsed.timeout, + }; + + match manager.start_server(&server.name, &config).await { + Ok(_) => { + tracing::info!("[AsterAgent] MCP server 已自动启动: {}", server.name); + success_count += 1; + } + Err(e) => { + tracing::error!( + "[AsterAgent] MCP server 自动启动失败: {} => {}", + server.name, + e + ); + fail_count += 1; + } + } + } + + (success_count, fail_count) +} diff --git a/src-tauri/src/commands/external_tools_cmd.rs b/src-tauri/src/commands/external_tools_cmd.rs new file mode 100644 index 000000000..f0e395904 --- /dev/null +++ b/src-tauri/src/commands/external_tools_cmd.rs @@ -0,0 +1,192 @@ +//! 外部 CLI 工具管理命令 +//! +//! 管理 Codex CLI 等外部工具的状态检查和配置 +//! 这些工具有自己的认证系统,不通过 ProxyCast 凭证池管理 + +use serde::{Deserialize, Serialize}; +use std::process::Stdio; +use tokio::process::Command; + +/// Codex CLI 状态 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct CodexCliStatus { + /// CLI 是否已安装 + pub installed: bool, + /// CLI 版本 + pub version: Option, + /// 是否已登录 + pub logged_in: bool, + /// 登录方式(api_key 或 oauth) + pub auth_type: Option, + /// API Key 前缀(如果使用 API Key 登录) + pub api_key_prefix: Option, + /// 错误信息 + pub error: Option, +} + +impl Default for CodexCliStatus { + fn default() -> Self { + Self { + installed: false, + version: None, + logged_in: false, + auth_type: None, + api_key_prefix: None, + error: None, + } + } +} + +/// 检查 Codex CLI 状态 +#[tauri::command] +pub async fn check_codex_cli_status() -> Result { + let mut status = CodexCliStatus::default(); + + // 1. 检查 codex 命令是否存在 + let version_result = Command::new("codex") + .arg("--version") + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .output() + .await; + + match version_result { + Ok(output) => { + if output.status.success() { + status.installed = true; + let version_str = String::from_utf8_lossy(&output.stdout); + // 解析版本号,格式通常是 "codex x.y.z" 或直接 "x.y.z" + status.version = Some(version_str.trim().to_string()); + } else { + status.error = Some("Codex CLI 未正确安装".to_string()); + return Ok(status); + } + } + Err(e) => { + status.error = Some(format!( + "Codex CLI 未安装。请运行: npm i -g @openai/codex\n错误: {}", + e + )); + return Ok(status); + } + } + + // 2. 检查登录状态 + let login_result = Command::new("codex") + .args(["login", "status"]) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .output() + .await; + + match login_result { + Ok(output) => { + let stdout = String::from_utf8_lossy(&output.stdout); + let stderr = String::from_utf8_lossy(&output.stderr); + let combined = format!("{}{}", stdout, stderr); + + tracing::debug!("[CodexCli] login status output: {}", combined); + + // 解析登录状态 + // 示例输出: "Logged in using an API key - cr_4453c***0b3a7" + // 或: "Not logged in" + if combined.contains("Logged in") { + status.logged_in = true; + + if combined.contains("API key") || combined.contains("api key") { + status.auth_type = Some("api_key".to_string()); + // 提取 API Key 前缀 + if let Some(key_part) = combined.split('-').last() { + let key = key_part.trim(); + if !key.is_empty() { + status.api_key_prefix = Some(key.to_string()); + } + } + } else if combined.contains("OAuth") || combined.contains("oauth") { + status.auth_type = Some("oauth".to_string()); + } else { + status.auth_type = Some("unknown".to_string()); + } + } else { + status.logged_in = false; + } + } + Err(e) => { + tracing::warn!("[CodexCli] 检查登录状态失败: {}", e); + // 不设置 error,因为 CLI 已安装,只是无法检查登录状态 + } + } + + Ok(status) +} + +/// 打开 Codex CLI 登录(在终端中执行) +#[tauri::command] +pub async fn open_codex_cli_login() -> Result { + // 返回登录命令,让前端在终端中执行 + Ok("codex login".to_string()) +} + +/// 打开 Codex CLI 登出 +#[tauri::command] +pub async fn open_codex_cli_logout() -> Result { + Ok("codex logout".to_string()) +} + +/// 外部工具列表 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ExternalTool { + /// 工具 ID + pub id: String, + /// 显示名称 + pub name: String, + /// 描述 + pub description: String, + /// 是否已安装 + pub installed: bool, + /// 是否已配置/登录 + pub configured: bool, + /// 安装命令 + pub install_command: String, + /// 配置命令 + pub config_command: String, + /// 文档链接 + pub doc_url: String, +} + +/// 获取外部工具列表 +#[tauri::command] +pub async fn get_external_tools() -> Result, String> { + let mut tools = Vec::new(); + + // Codex CLI + let codex_status = check_codex_cli_status().await.unwrap_or_default(); + tools.push(ExternalTool { + id: "codex-cli".to_string(), + name: "Codex CLI".to_string(), + description: "OpenAI Codex 命令行工具,支持 Agent 模式和工具调用".to_string(), + installed: codex_status.installed, + configured: codex_status.logged_in, + install_command: "npm i -g @openai/codex".to_string(), + config_command: "codex login".to_string(), + doc_url: "https://github.com/openai/codex".to_string(), + }); + + // 可以在这里添加更多外部工具... + + Ok(tools) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn test_codex_cli_status() { + // 这个测试依赖于本地环境 + let status = check_codex_cli_status().await; + assert!(status.is_ok()); + let status = status.unwrap(); + println!("Codex CLI Status: {:?}", status); + } +} diff --git a/src-tauri/src/commands/mcp_cmd.rs b/src-tauri/src/commands/mcp_cmd.rs index f71c69236..e60a82c04 100644 --- a/src-tauri/src/commands/mcp_cmd.rs +++ b/src-tauri/src/commands/mcp_cmd.rs @@ -1,7 +1,48 @@ +//! MCP Tauri 命令 +//! +//! 本模块提供 MCP 相关的 Tauri 命令接口,包括: +//! - 服务器配置 CRUD 操作 +//! - 服务器生命周期管理(启动、停止) +//! - 服务器状态查询 +//! - 工具管理(列表、调用) +//! - 提示词管理(列表、获取内容) +//! - 资源管理(列表、读取内容) +//! +//! # 命令分类 +//! +//! ## 配置管理命令 +//! - `get_mcp_servers`: 获取所有 MCP 服务器配置 +//! - `add_mcp_server`: 添加新的 MCP 服务器配置 +//! - `update_mcp_server`: 更新 MCP 服务器配置 +//! - `delete_mcp_server`: 删除 MCP 服务器配置 +//! - `toggle_mcp_server`: 切换服务器在特定应用中的启用状态 +//! +//! ## 生命周期管理命令 +//! - `mcp_list_servers_with_status`: 获取所有服务器及其运行状态 +//! - `mcp_start_server`: 启动指定的 MCP 服务器 +//! - `mcp_stop_server`: 停止指定的 MCP 服务器 +//! +//! ## 工具管理命令 +//! - `mcp_list_tools`: 获取所有可用工具 +//! - `mcp_call_tool`: 调用指定工具 +//! +//! ## 提示词管理命令 +//! - `mcp_list_prompts`: 获取所有可用提示词 +//! - `mcp_get_prompt`: 获取提示词内容 +//! +//! ## 资源管理命令 +//! - `mcp_list_resources`: 获取所有可用资源 +//! - `mcp_read_resource`: 读取资源内容 + use crate::database::DbConnection; +use crate::mcp::{ + McpManagerState, McpPromptDefinition, McpPromptResult, McpResourceContent, + McpResourceDefinition, McpServerConfig, McpServerInfo, McpToolDefinition, McpToolResult, +}; use crate::models::McpServer; use crate::services::mcp_service::McpService; use tauri::State; +use tracing::{debug, error, info}; #[tauri::command] pub fn get_mcp_servers(db: State<'_, DbConnection>) -> Result, String> { @@ -42,3 +83,431 @@ pub fn import_mcp_from_app(db: State<'_, DbConnection>, app_type: String) -> Res pub fn sync_all_mcp_to_live(db: State<'_, DbConnection>) -> Result<(), String> { McpService::sync_all_to_live(&db) } + +// ============================================================================ +// 服务器生命周期管理命令 +// ============================================================================ + +/// 获取所有 MCP 服务器配置及其运行状态 +/// +/// 从数据库获取所有配置的 MCP 服务器,并查询每个服务器的运行状态。 +/// +/// # Arguments +/// +/// * `db` - 数据库连接状态 +/// * `mcp_manager` - MCP 管理器状态 +/// +/// # Returns +/// +/// 返回包含运行状态的服务器信息列表。 +/// +/// # Requirements +/// +/// - **9.1**: THE mcp_list_servers command SHALL return all configured MCP servers with status +#[tauri::command] +pub async fn mcp_list_servers_with_status( + db: State<'_, DbConnection>, + mcp_manager: State<'_, McpManagerState>, +) -> Result, String> { + info!("获取所有 MCP 服务器及状态"); + + // 1. 从数据库获取所有服务器配置 + let servers = McpService::get_all(&db)?; + + // 2. 获取管理器锁 + let manager = mcp_manager.lock().await; + + // 3. 构建带状态的服务器信息列表 + let mut result: Vec = Vec::new(); + + for server in servers { + // 解析服务器配置 + let config = parse_server_config(&server.server_config); + + // 检查服务器是否正在运行 + let is_running = manager.is_server_running(&server.name).await; + + // 获取服务器能力信息(如果正在运行) + let server_info = if is_running { + manager.get_client_capabilities(&server.name).await + } else { + None + }; + + result.push(McpServerInfo { + id: server.id, + name: server.name, + description: server.description, + config, + is_running, + server_info, + enabled_proxycast: server.enabled_proxycast, + enabled_claude: server.enabled_claude, + enabled_codex: server.enabled_codex, + enabled_gemini: server.enabled_gemini, + }); + } + + debug!(server_count = result.len(), "返回服务器列表"); + Ok(result) +} + +/// 启动 MCP 服务器 +/// +/// 根据服务器名称从数据库获取配置,然后启动服务器进程。 +/// +/// # Arguments +/// +/// * `db` - 数据库连接状态 +/// * `mcp_manager` - MCP 管理器状态 +/// * `name` - 服务器名称 +/// +/// # Returns +/// +/// 成功返回 Ok(()),失败返回错误信息。 +/// +/// # Requirements +/// +/// - **9.2**: THE mcp_start_server command SHALL start a specified MCP server +#[tauri::command] +pub async fn mcp_start_server( + db: State<'_, DbConnection>, + mcp_manager: State<'_, McpManagerState>, + name: String, +) -> Result<(), String> { + info!(server_name = %name, "启动 MCP 服务器命令"); + + // 1. 从数据库获取服务器配置 + let servers = McpService::get_all(&db)?; + let server = servers + .iter() + .find(|s| s.name == name) + .ok_or_else(|| format!("服务器配置不存在: {}", name))?; + + // 2. 解析服务器配置 + let config = parse_server_config(&server.server_config); + + // 3. 获取管理器锁并启动服务器 + let manager = mcp_manager.lock().await; + manager.start_server(&name, &config).await.map_err(|e| { + error!(server_name = %name, error = %e, "启动 MCP 服务器失败"); + e.to_string() + })?; + + info!(server_name = %name, "MCP 服务器启动成功"); + Ok(()) +} + +/// 停止 MCP 服务器 +/// +/// 根据服务器名称停止正在运行的服务器进程。 +/// +/// # Arguments +/// +/// * `mcp_manager` - MCP 管理器状态 +/// * `name` - 服务器名称 +/// +/// # Returns +/// +/// 成功返回 Ok(()),失败返回错误信息。 +/// 如果服务器未运行,也返回 Ok()(幂等操作)。 +/// +/// # Requirements +/// +/// - **9.3**: THE mcp_stop_server command SHALL stop a specified MCP server +#[tauri::command] +pub async fn mcp_stop_server( + mcp_manager: State<'_, McpManagerState>, + name: String, +) -> Result<(), String> { + info!(server_name = %name, "停止 MCP 服务器命令"); + + // 获取管理器锁并停止服务器 + let manager = mcp_manager.lock().await; + manager.stop_server(&name).await.map_err(|e| { + error!(server_name = %name, error = %e, "停止 MCP 服务器失败"); + e.to_string() + })?; + + info!(server_name = %name, "MCP 服务器已停止"); + Ok(()) +} + +// ============================================================================ +// 辅助函数 +// ============================================================================ + +/// 解析服务器配置 JSON 为 McpServerConfig +/// +/// 将数据库中存储的 JSON 配置解析为结构化的 McpServerConfig。 +/// 如果解析失败,返回默认配置。 +/// +/// # Arguments +/// +/// * `config_value` - JSON 格式的服务器配置 +/// +/// # Returns +/// +/// 返回解析后的 McpServerConfig,如果解析失败则返回默认值。 +fn parse_server_config(config_value: &serde_json::Value) -> McpServerConfig { + serde_json::from_value(config_value.clone()).unwrap_or_else(|e| { + debug!(error = %e, "解析服务器配置失败,使用默认值"); + McpServerConfig { + command: config_value + .get("command") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(), + args: config_value + .get("args") + .and_then(|v| v.as_array()) + .map(|arr| { + arr.iter() + .filter_map(|v| v.as_str().map(|s| s.to_string())) + .collect() + }) + .unwrap_or_default(), + env: config_value + .get("env") + .and_then(|v| v.as_object()) + .map(|obj| { + obj.iter() + .filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_string()))) + .collect() + }) + .unwrap_or_default(), + cwd: config_value + .get("cwd") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()), + timeout: config_value + .get("timeout") + .and_then(|v| v.as_u64()) + .unwrap_or(30), + } + }) +} + +// ============================================================================ +// 工具管理命令 +// ============================================================================ + +/// 获取所有可用工具 +/// +/// 从所有运行中的 MCP 服务器获取工具定义列表。 +/// 工具定义包含名称、描述和输入参数 schema。 +/// +/// # Arguments +/// +/// * `mcp_manager` - MCP 管理器状态 +/// +/// # Returns +/// +/// 返回所有可用工具的定义列表。 +/// +/// # Requirements +/// +/// - **9.4**: THE mcp_list_tools command SHALL return all available tools from running servers +#[tauri::command] +pub async fn mcp_list_tools( + mcp_manager: State<'_, McpManagerState>, +) -> Result, String> { + info!("获取所有 MCP 工具列表"); + + let manager = mcp_manager.lock().await; + let tools = manager.list_tools().await.map_err(|e| { + error!(error = %e, "获取工具列表失败"); + e.to_string() + })?; + + debug!(tool_count = tools.len(), "返回工具列表"); + Ok(tools) +} + +/// 调用 MCP 工具 +/// +/// 根据工具名称和参数调用指定的 MCP 工具。 +/// 工具名称可能包含服务器前缀(格式为 "server_toolname")。 +/// +/// # Arguments +/// +/// * `mcp_manager` - MCP 管理器状态 +/// * `tool_name` - 工具名称 +/// * `arguments` - 工具参数(JSON 对象) +/// +/// # Returns +/// +/// 返回工具调用结果,包含内容和错误状态。 +/// +/// # Requirements +/// +/// - **9.5**: THE mcp_call_tool command SHALL call a tool and return the result +#[tauri::command] +pub async fn mcp_call_tool( + mcp_manager: State<'_, McpManagerState>, + tool_name: String, + arguments: serde_json::Value, +) -> Result { + info!(tool_name = %tool_name, "调用 MCP 工具命令"); + + let manager = mcp_manager.lock().await; + let result = manager + .call_tool(&tool_name, arguments) + .await + .map_err(|e| { + error!(tool_name = %tool_name, error = %e, "调用工具失败"); + e.to_string() + })?; + + info!( + tool_name = %tool_name, + is_error = result.is_error, + "工具调用完成" + ); + Ok(result) +} + +// ============================================================================ +// 提示词管理命令 +// ============================================================================ + +/// 获取所有可用提示词 +/// +/// 从所有运行中的 MCP 服务器获取提示词定义列表。 +/// 提示词定义包含名称、描述和参数列表。 +/// +/// # Arguments +/// +/// * `mcp_manager` - MCP 管理器状态 +/// +/// # Returns +/// +/// 返回所有可用提示词的定义列表。 +/// +/// # Requirements +/// +/// - **9.6**: THE mcp_list_prompts command SHALL return all available prompts from running servers +#[tauri::command] +pub async fn mcp_list_prompts( + mcp_manager: State<'_, McpManagerState>, +) -> Result, String> { + info!("获取所有 MCP 提示词列表"); + + let manager = mcp_manager.lock().await; + let prompts = manager.list_prompts().await.map_err(|e| { + error!(error = %e, "获取提示词列表失败"); + e.to_string() + })?; + + debug!(prompt_count = prompts.len(), "返回提示词列表"); + Ok(prompts) +} + +/// 获取提示词内容 +/// +/// 根据提示词名称和参数获取提示词内容。 +/// 提示词名称可能包含服务器前缀(格式为 "server_promptname")。 +/// +/// # Arguments +/// +/// * `mcp_manager` - MCP 管理器状态 +/// * `name` - 提示词名称 +/// * `arguments` - 提示词参数(JSON 对象) +/// +/// # Returns +/// +/// 返回提示词内容,包含描述和消息列表。 +/// +/// # Requirements +/// +/// - **9.7**: THE mcp_get_prompt command SHALL return prompt content with argument substitution +#[tauri::command] +pub async fn mcp_get_prompt( + mcp_manager: State<'_, McpManagerState>, + name: String, + arguments: serde_json::Map, +) -> Result { + info!(prompt_name = %name, "获取 MCP 提示词内容命令"); + + let manager = mcp_manager.lock().await; + let result = manager.get_prompt(&name, arguments).await.map_err(|e| { + error!(prompt_name = %name, error = %e, "获取提示词内容失败"); + e.to_string() + })?; + + info!( + prompt_name = %name, + message_count = result.messages.len(), + "提示词内容获取完成" + ); + Ok(result) +} + +// ============================================================================ +// 资源管理命令 +// ============================================================================ + +/// 获取所有可用资源 +/// +/// 从所有运行中的 MCP 服务器获取资源定义列表。 +/// 资源定义包含 URI、名称、描述和 MIME 类型。 +/// +/// # Arguments +/// +/// * `mcp_manager` - MCP 管理器状态 +/// +/// # Returns +/// +/// 返回所有可用资源的定义列表。 +/// +/// # Requirements +/// +/// - **9.8**: THE mcp_list_resources command SHALL return all available resources from running servers +#[tauri::command] +pub async fn mcp_list_resources( + mcp_manager: State<'_, McpManagerState>, +) -> Result, String> { + info!("获取所有 MCP 资源列表"); + + let manager = mcp_manager.lock().await; + let resources = manager.list_resources().await.map_err(|e| { + error!(error = %e, "获取资源列表失败"); + e.to_string() + })?; + + debug!(resource_count = resources.len(), "返回资源列表"); + Ok(resources) +} + +/// 读取资源内容 +/// +/// 根据资源 URI 读取资源内容。 +/// +/// # Arguments +/// +/// * `mcp_manager` - MCP 管理器状态 +/// * `uri` - 资源 URI +/// +/// # Returns +/// +/// 返回资源内容,包含 URI、MIME 类型和内容(文本或二进制)。 +/// +/// # Requirements +/// +/// - **9.9**: THE mcp_read_resource command SHALL return resource content by URI +#[tauri::command] +pub async fn mcp_read_resource( + mcp_manager: State<'_, McpManagerState>, + uri: String, +) -> Result { + info!(uri = %uri, "读取 MCP 资源内容命令"); + + let manager = mcp_manager.lock().await; + let result = manager.read_resource(&uri).await.map_err(|e| { + error!(uri = %uri, error = %e, "读取资源内容失败"); + e.to_string() + })?; + + info!(uri = %uri, "资源内容读取完成"); + Ok(result) +} diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 603afc889..48c8d0632 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -10,6 +10,7 @@ pub mod connect_cmd; pub mod connection_cmd; pub mod content_cmd; pub mod context_memory; +pub mod external_tools_cmd; pub mod flow_monitor_cmd; pub mod general_chat_cmd; pub mod injection_cmd; @@ -37,6 +38,8 @@ pub mod route_cmd; pub mod screenshot_cmd; pub mod session_files_cmd; pub mod skill_cmd; +pub mod skill_exec_cmd; +pub mod subagent_cmd; pub mod switch_cmd; pub mod telemetry_cmd; pub mod template_cmd; diff --git a/src-tauri/src/commands/skill_cmd.rs b/src-tauri/src/commands/skill_cmd.rs index ee12037b6..ac9b8074a 100644 --- a/src-tauri/src/commands/skill_cmd.rs +++ b/src-tauri/src/commands/skill_cmd.rs @@ -1,3 +1,4 @@ +use crate::agent::aster_state::AsterAgentState; use crate::database::dao::skills::SkillDao; use crate::database::DbConnection; use crate::models::{AppType, Skill, SkillRepo, SkillState}; @@ -66,7 +67,7 @@ pub async fn get_skills( db: State<'_, DbConnection>, skill_service: State<'_, SkillServiceState>, ) -> Result, String> { - get_skills_for_app(db, skill_service, "claude".to_string()).await + get_skills_for_app(db, skill_service, "proxycast".to_string()).await } #[tauri::command] @@ -120,7 +121,7 @@ pub async fn install_skill( skill_service: State<'_, SkillServiceState>, directory: String, ) -> Result { - install_skill_for_app(db, skill_service, "claude".to_string(), directory).await + install_skill_for_app(db, skill_service, "proxycast".to_string(), directory).await } #[tauri::command] @@ -186,12 +187,15 @@ pub async fn install_skill_for_app( SkillDao::update_skill_state(&conn, &key, &state).map_err(|e| e.to_string())?; } + // 刷新 aster-rust 的 global_registry,使 AI 能够发现新安装的 Skill + AsterAgentState::reload_proxycast_skills(); + Ok(true) } #[tauri::command] pub fn uninstall_skill(db: State<'_, DbConnection>, directory: String) -> Result { - uninstall_skill_for_app(db, "claude".to_string(), directory) + uninstall_skill_for_app(db, "proxycast".to_string(), directory) } #[tauri::command] @@ -215,6 +219,9 @@ pub fn uninstall_skill_for_app( let conn = db.lock().map_err(|e| e.to_string())?; SkillDao::update_skill_state(&conn, &key, &state).map_err(|e| e.to_string())?; + // 刷新 aster-rust 的 global_registry,移除已卸载的 Skill + AsterAgentState::reload_proxycast_skills(); + Ok(true) } diff --git a/src-tauri/src/commands/skill_exec_cmd.rs b/src-tauri/src/commands/skill_exec_cmd.rs new file mode 100644 index 000000000..c6aeda244 --- /dev/null +++ b/src-tauri/src/commands/skill_exec_cmd.rs @@ -0,0 +1,712 @@ +//! Skill 执行 Tauri 命令模块 +//! +//! 本模块提供 Skill 执行相关的 Tauri 命令,包括: +//! - `execute_skill`: 执行指定的 Skill +//! - `list_executable_skills`: 列出所有可执行的 Skills +//! - `get_skill_detail`: 获取 Skill 详情 +//! +//! ## 依赖 +//! - `AsterAgentState`: Aster Agent 状态管理,提供完整的工具集支持 +//! - `TauriExecutionCallback`: 执行进度回调 +//! - `ProviderPoolService`: 凭证池服务 +//! +//! ## Requirements +//! - 3.1: execute_skill 命令接受 skill_name 和 user_input 参数 +//! - 4.1: list_executable_skills 返回所有可执行的 skills +//! - 5.1: get_skill_detail 接受 skill_name 参数 + +use futures::StreamExt; +use serde::{Deserialize, Serialize}; +use tauri::{Emitter, State}; +use uuid::Uuid; + +use aster::conversation::message::Message; + +use crate::agent::aster_state::SessionConfigBuilder; +use crate::agent::event_converter::convert_agent_event; +use crate::agent::{AsterAgentState, TauriAgentEvent}; +use crate::database::DbConnection; +use crate::skills::{ + find_skill_by_name, get_proxycast_skills_dir, load_skills_from_directory, ExecutionCallback, + TauriExecutionCallback, +}; +#[cfg(test)] +use crate::skills::{ + load_skill_from_file, parse_allowed_tools, parse_boolean, parse_skill_frontmatter, +}; + +// ============================================================================ +// 公开类型定义 +// ============================================================================ + +/// 可执行 Skill 信息 +/// +/// 用于 list_executable_skills 命令的返回类型 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ExecutableSkillInfo { + /// Skill 名称(唯一标识) + pub name: String, + /// 显示名称 + pub display_name: String, + /// Skill 描述 + pub description: String, + /// 执行模式:prompt, workflow, agent + pub execution_mode: String, + /// 是否有 workflow 定义 + pub has_workflow: bool, + /// 指定的 Provider(可选) + pub provider: Option, + /// 指定的 Model(可选) + pub model: Option, + /// 参数提示(可选) + pub argument_hint: Option, +} + +/// Workflow 步骤信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct WorkflowStepInfo { + /// 步骤 ID + pub id: String, + /// 步骤名称 + pub name: String, + /// 依赖的步骤 ID 列表 + pub dependencies: Vec, +} + +/// Skill 详情信息 +/// +/// 用于 get_skill_detail 命令的返回类型 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SkillDetailInfo { + /// 基本信息 + #[serde(flatten)] + pub basic: ExecutableSkillInfo, + /// Markdown 内容 + pub markdown_content: String, + /// Workflow 步骤(如果有) + pub workflow_steps: Option>, + /// 允许的工具列表(可选) + pub allowed_tools: Option>, + /// 使用场景说明(可选) + pub when_to_use: Option, +} + +/// 步骤执行结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct StepResult { + /// 步骤 ID + pub step_id: String, + /// 步骤名称 + pub step_name: String, + /// 是否成功 + pub success: bool, + /// 输出内容 + pub output: Option, + /// 错误信息 + pub error: Option, +} + +/// Skill 执行结果 +/// +/// 用于 execute_skill 命令的返回类型 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SkillExecutionResult { + /// 是否成功 + pub success: bool, + /// 最终输出 + pub output: Option, + /// 错误信息 + pub error: Option, + /// 已完成的步骤结果 + pub steps_completed: Vec, +} + +/// 执行 Skill +/// +/// 加载并执行指定的 Skill,使用 Aster Agent 系统提供完整的工具集支持。 +/// +/// # Arguments +/// * `app_handle` - Tauri AppHandle,用于发送事件 +/// * `db` - 数据库连接 +/// * `aster_state` - Aster Agent 状态 +/// * `skill_name` - Skill 名称 +/// * `user_input` - 用户输入 +/// * `provider_override` - 可选的 Provider 覆盖 +/// * `session_id` - 可选的会话 ID(用于复用当前聊天上下文) +/// +/// # Returns +/// * `Ok(SkillExecutionResult)` - 执行结果 +/// * `Err(String)` - 错误信息 +/// +/// # Requirements +/// - 3.1: 接受 skill_name 和 user_input 参数 +/// - 3.2: 从 registry 加载 skill +/// - 3.3: 使用 Aster Agent 执行(支持工具调用) +/// - 3.5: 返回 SkillExecutionResult +#[tauri::command] +pub async fn execute_skill( + app_handle: tauri::AppHandle, + db: State<'_, DbConnection>, + aster_state: State<'_, AsterAgentState>, + skill_name: String, + user_input: String, + provider_override: Option, + model_override: Option, + execution_id: Option, + session_id: Option, +) -> Result { + // 生成执行 ID,并优先复用前端会话 ID(提升 /skill 与主会话上下文一致性) + let execution_id = execution_id.unwrap_or_else(|| Uuid::new_v4().to_string()); + let session_id = session_id.unwrap_or_else(|| format!("skill-exec-{}", Uuid::new_v4())); + + tracing::info!( + "[execute_skill] 开始执行 Skill: name={}, execution_id={}, session_id={}, provider_override={:?}, model_override={:?}", + skill_name, + execution_id, + session_id, + provider_override, + model_override + ); + + // 1. 从 registry 加载 skill(Requirements 3.2) + let skill = find_skill_by_name(&skill_name)?; + + // 检查是否禁用了模型调用 + if skill.disable_model_invocation { + return Err(format!("Skill '{}' 已禁用模型调用,无法执行", skill_name)); + } + + // 2. 创建 TauriExecutionCallback + let callback = TauriExecutionCallback::new(app_handle.clone(), execution_id.clone()); + + // 3. 初始化 Agent(如果未初始化) + if !aster_state.is_initialized().await { + tracing::info!("[execute_skill] Agent 未初始化,开始初始化..."); + aster_state.init_agent_with_db(&db).await?; + tracing::info!("[execute_skill] Agent 初始化完成"); + } + + // 4. 配置 Provider(从凭证池选择,支持 fallback) + let preferred_provider = provider_override + .or_else(|| skill.provider.clone()) + .unwrap_or_else(|| "anthropic".to_string()); + + let preferred_model = model_override + .or_else(|| skill.model.clone()) + .unwrap_or_else(|| "claude-sonnet-4-20250514".to_string()); + + // 支持工具调用的 Provider fallback 列表 + // 注意:provider 名称需要与 ProviderType::FromStr 匹配 + let fallback_providers: Vec<(&str, &str)> = vec![ + ("anthropic", "claude-sonnet-4-20250514"), + ("openai", "gpt-4o"), + ("gemini", "gemini-2.0-flash"), + ]; + + let mut configure_result = aster_state + .configure_provider_from_pool(&db, &preferred_provider, &preferred_model, &session_id) + .await; + + if configure_result.is_err() { + tracing::warn!( + "[execute_skill] 首选 Provider {} 配置失败: {:?},尝试 fallback", + preferred_provider, + configure_result.as_ref().err() + ); + + for (fb_provider, fb_model) in &fallback_providers { + if *fb_provider == preferred_provider { + continue; + } + match aster_state + .configure_provider_from_pool(&db, fb_provider, fb_model, &session_id) + .await + { + Ok(config) => { + tracing::info!( + "[execute_skill] Fallback 到 {} / {} 成功", + fb_provider, + fb_model + ); + configure_result = Ok(config); + break; + } + Err(e) => { + tracing::warn!("[execute_skill] Fallback {} 也失败: {}", fb_provider, e); + } + } + } + } + + configure_result.map_err(|e| { + format!("无法配置任何可用的 Provider(需要支持工具调用的 Provider,如 Anthropic、OpenAI 或 Google): {e}") + })?; + + tracing::info!( + "[execute_skill] Provider 配置成功: preferred={}, model={}", + preferred_provider, + preferred_model + ); + + // 5. 发送步骤开始事件 + callback.on_step_start("main", &skill.display_name, 1, 1); + + // 6. 构建 SessionConfig,将 skill 内容作为 system_prompt + let session_config = SessionConfigBuilder::new(&session_id) + .system_prompt(&skill.markdown_content) + .build(); + + // 7. 创建用户消息 + let user_message = Message::user().with_text(&user_input); + + // 8. 获取 Agent 并执行 + let agent_arc = aster_state.get_agent_arc(); + let guard = agent_arc.read().await; + let agent = guard.as_ref().ok_or("Agent not initialized")?; + + // 创建取消令牌 + let cancel_token = aster_state.create_cancel_token(&session_id).await; + + // 获取事件流 + let stream_result = agent + .reply(user_message, session_config, Some(cancel_token.clone())) + .await; + + // 9. 处理流式事件并收集结果 + let mut final_output = String::new(); + let mut has_error = false; + let mut error_message: Option = None; + + // 用于发送流式事件的 event_name + let event_name = format!("skill-exec-{}", execution_id); + + match stream_result { + Ok(mut stream) => { + while let Some(event_result) = stream.next().await { + match event_result { + Ok(agent_event) => { + // 转换 Aster 事件为 Tauri 事件 + let tauri_events = convert_agent_event(agent_event); + + for tauri_event in tauri_events { + // 收集文本输出 + if let TauriAgentEvent::TextDelta { ref text } = tauri_event { + final_output.push_str(text); + } + + // 发送事件到前端 + if let Err(e) = app_handle.emit(&event_name, &tauri_event) { + tracing::error!("[execute_skill] 发送事件失败: {}", e); + } + } + } + Err(e) => { + has_error = true; + error_message = Some(format!("Stream error: {e}")); + tracing::error!("[execute_skill] 流处理错误: {}", e); + } + } + } + + // 发送完成事件 + let done_event = TauriAgentEvent::FinalDone { usage: None }; + if let Err(e) = app_handle.emit(&event_name, &done_event) { + tracing::error!("[execute_skill] 发送完成事件失败: {}", e); + } + } + Err(e) => { + has_error = true; + error_message = Some(format!("Agent error: {e}")); + tracing::error!("[execute_skill] Agent 错误: {}", e); + } + } + + // 清理取消令牌 + aster_state.remove_cancel_token(&session_id).await; + + // 10. 返回执行结果(Requirements 3.5) + if has_error { + let err_msg = error_message.unwrap_or_else(|| "Unknown error".to_string()); + callback.on_step_error("main", &err_msg, false); + callback.on_complete(false, None, Some(&err_msg)); + + tracing::error!( + "[execute_skill] Skill 执行失败: name={}, error={}", + skill_name, + err_msg + ); + + Ok(SkillExecutionResult { + success: false, + output: None, + error: Some(err_msg.clone()), + steps_completed: vec![StepResult { + step_id: "main".to_string(), + step_name: skill.display_name, + success: false, + output: None, + error: Some(err_msg), + }], + }) + } else { + callback.on_step_complete("main", &final_output); + callback.on_complete(true, Some(&final_output), None); + + tracing::info!( + "[execute_skill] Skill 执行成功: name={}, output_len={}", + skill_name, + final_output.len() + ); + + Ok(SkillExecutionResult { + success: true, + output: Some(final_output.clone()), + error: None, + steps_completed: vec![StepResult { + step_id: "main".to_string(), + step_name: skill.display_name, + success: true, + output: Some(final_output), + error: None, + }], + }) + } +} + +/// 列出可执行的 Skills +/// +/// 返回所有可以执行的 Skills 列表,过滤掉 disable_model_invocation=true 的 Skills。 +/// +/// # Returns +/// * `Ok(Vec)` - 可执行的 Skills 列表 +/// * `Err(String)` - 错误信息 +/// +/// # Requirements +/// - 4.1: 返回所有可执行的 skills +/// - 4.2: 包含 name, description, execution_mode +/// - 4.3: 指示是否有 workflow 定义 +/// - 4.4: 过滤 disable_model_invocation=true 的 skills +#[tauri::command] +pub async fn list_executable_skills() -> Result, String> { + let skills_dir = + get_proxycast_skills_dir().ok_or_else(|| "无法获取 Skills 目录".to_string())?; + + // 加载所有 skills + let all_skills = load_skills_from_directory(&skills_dir); + + // 过滤掉 disable_model_invocation=true 的 skills(Requirements 4.4) + let executable_skills: Vec = all_skills + .into_iter() + .filter(|s| !s.disable_model_invocation) + .map(|s| ExecutableSkillInfo { + name: s.skill_name, + display_name: s.display_name, + description: s.description, + execution_mode: s.execution_mode.clone(), + has_workflow: s.execution_mode == "workflow", + provider: s.provider, + model: s.model, + argument_hint: s.argument_hint, + }) + .collect(); + + tracing::info!( + "[list_executable_skills] 返回 {} 个可执行 Skills", + executable_skills.len() + ); + + Ok(executable_skills) +} + +/// 获取 Skill 详情 +/// +/// 根据 skill_name 返回完整的 Skill 详情信息。 +/// +/// # Arguments +/// * `skill_name` - Skill 名称 +/// +/// # Returns +/// * `Ok(SkillDetailInfo)` - Skill 详情 +/// * `Err(String)` - 错误信息(如 skill 不存在) +/// +/// # Requirements +/// - 5.1: 接受 skill_name 参数 +/// - 5.2: 返回完整的 SkillDefinition +/// - 5.3: 包含 workflow steps 信息(如果有) +/// - 5.4: skill 不存在时返回错误 +#[tauri::command] +pub async fn get_skill_detail(skill_name: String) -> Result { + // 查找 skill(Requirements 5.1, 5.4) + let skill = find_skill_by_name(&skill_name)?; + + // 转换为 SkillDetailInfo(Requirements 5.2, 5.3) + let detail = SkillDetailInfo { + basic: ExecutableSkillInfo { + name: skill.skill_name, + display_name: skill.display_name, + description: skill.description, + execution_mode: skill.execution_mode.clone(), + has_workflow: skill.execution_mode == "workflow", + provider: skill.provider, + model: skill.model, + argument_hint: skill.argument_hint, + }, + markdown_content: skill.markdown_content, + workflow_steps: None, // TODO: 解析 workflow 步骤(如果有) + allowed_tools: skill.allowed_tools, + when_to_use: skill.when_to_use, + }; + + tracing::info!("[get_skill_detail] 返回 Skill 详情: name={}", skill_name); + + Ok(detail) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_executable_skill_info_serialization() { + let info = ExecutableSkillInfo { + name: "test-skill".to_string(), + display_name: "Test Skill".to_string(), + description: "A test skill".to_string(), + execution_mode: "prompt".to_string(), + has_workflow: false, + provider: None, + model: None, + argument_hint: Some("Enter your query".to_string()), + }; + + let json = serde_json::to_string(&info).unwrap(); + assert!(json.contains("test-skill")); + assert!(json.contains("Test Skill")); + } + + #[test] + fn test_skill_execution_result_serialization() { + let result = SkillExecutionResult { + success: true, + output: Some("Hello, world!".to_string()), + error: None, + steps_completed: vec![StepResult { + step_id: "step-1".to_string(), + step_name: "Process".to_string(), + success: true, + output: Some("Done".to_string()), + error: None, + }], + }; + + let json = serde_json::to_string(&result).unwrap(); + assert!(json.contains("\"success\":true")); + assert!(json.contains("Hello, world!")); + assert!(json.contains("step-1")); + } + + #[test] + fn test_skill_detail_info_serialization() { + let detail = SkillDetailInfo { + basic: ExecutableSkillInfo { + name: "workflow-skill".to_string(), + display_name: "Workflow Skill".to_string(), + description: "A workflow skill".to_string(), + execution_mode: "workflow".to_string(), + has_workflow: true, + provider: Some("claude".to_string()), + model: Some("claude-sonnet-4-5-20250514".to_string()), + argument_hint: None, + }, + markdown_content: "# Workflow Skill\n\nThis is a workflow skill.".to_string(), + workflow_steps: Some(vec![ + WorkflowStepInfo { + id: "step-1".to_string(), + name: "Initialize".to_string(), + dependencies: vec![], + }, + WorkflowStepInfo { + id: "step-2".to_string(), + name: "Process".to_string(), + dependencies: vec!["step-1".to_string()], + }, + ]), + allowed_tools: Some(vec!["read_file".to_string(), "write_file".to_string()]), + when_to_use: Some("Use this skill for complex workflows".to_string()), + }; + + let json = serde_json::to_string(&detail).unwrap(); + assert!(json.contains("workflow-skill")); + assert!(json.contains("workflow_steps")); + assert!(json.contains("step-1")); + assert!(json.contains("step-2")); + } + + #[test] + fn test_parse_skill_frontmatter_basic() { + let content = r#"--- +name: test-skill +description: A test skill +model: claude-sonnet-4-5-20250514 +provider: claude +--- + +# Test Skill + +This is the body content. +"#; + let (fm, body) = parse_skill_frontmatter(content); + assert_eq!(fm.name, Some("test-skill".to_string())); + assert_eq!(fm.description, Some("A test skill".to_string())); + assert_eq!(fm.model, Some("claude-sonnet-4-5-20250514".to_string())); + assert_eq!(fm.provider, Some("claude".to_string())); + assert!(body.contains("# Test Skill")); + assert!(body.contains("This is the body content.")); + } + + #[test] + fn test_parse_skill_frontmatter_no_frontmatter() { + let content = "# Just content\nNo frontmatter here."; + let (fm, body) = parse_skill_frontmatter(content); + assert!(fm.name.is_none()); + assert_eq!(body, content); + } + + #[test] + fn test_parse_skill_frontmatter_with_quotes() { + let content = r#"--- +name: "quoted-name" +description: 'single quoted' +--- +Body +"#; + let (fm, _) = parse_skill_frontmatter(content); + assert_eq!(fm.name, Some("quoted-name".to_string())); + assert_eq!(fm.description, Some("single quoted".to_string())); + } + + #[test] + fn test_parse_allowed_tools() { + assert_eq!(parse_allowed_tools(None), None); + assert_eq!(parse_allowed_tools(Some("")), None); + assert_eq!( + parse_allowed_tools(Some("tool1")), + Some(vec!["tool1".to_string()]) + ); + assert_eq!( + parse_allowed_tools(Some("tool1, tool2, tool3")), + Some(vec![ + "tool1".to_string(), + "tool2".to_string(), + "tool3".to_string() + ]) + ); + } + + #[test] + fn test_parse_boolean() { + assert!(!parse_boolean(None, false)); + assert!(parse_boolean(None, true)); + assert!(parse_boolean(Some("true"), false)); + assert!(parse_boolean(Some("TRUE"), false)); + assert!(parse_boolean(Some("1"), false)); + assert!(parse_boolean(Some("yes"), false)); + assert!(!parse_boolean(Some("false"), true)); + assert!(!parse_boolean(Some("no"), true)); + } + + #[test] + fn test_load_skill_from_file() { + use tempfile::TempDir; + + let temp_dir = TempDir::new().unwrap(); + let skill_dir = temp_dir.path().join("my-skill"); + std::fs::create_dir(&skill_dir).unwrap(); + + let skill_file = skill_dir.join("SKILL.md"); + std::fs::write( + &skill_file, + r#"--- +name: my-skill +description: Test skill description +allowed-tools: tool1, tool2 +model: gpt-4 +provider: openai +--- + +# My Skill + +Instructions here. +"#, + ) + .unwrap(); + + let skill = load_skill_from_file("my-skill", &skill_file).unwrap(); + + assert_eq!(skill.skill_name, "my-skill"); + assert_eq!(skill.display_name, "my-skill"); + assert_eq!(skill.description, "Test skill description"); + assert_eq!( + skill.allowed_tools, + Some(vec!["tool1".to_string(), "tool2".to_string()]) + ); + assert_eq!(skill.model, Some("gpt-4".to_string())); + assert_eq!(skill.provider, Some("openai".to_string())); + assert!(!skill.disable_model_invocation); + assert_eq!(skill.execution_mode, "prompt"); + } + + #[test] + fn test_load_skills_from_directory() { + use tempfile::TempDir; + + let temp_dir = TempDir::new().unwrap(); + let skills_dir = temp_dir.path(); + + // 创建 skill 1 + let skill1_dir = skills_dir.join("skill-one"); + std::fs::create_dir(&skill1_dir).unwrap(); + std::fs::write( + skill1_dir.join("SKILL.md"), + r#"--- +name: skill-one +description: First skill +--- +Content 1 +"#, + ) + .unwrap(); + + // 创建 skill 2 + let skill2_dir = skills_dir.join("skill-two"); + std::fs::create_dir(&skill2_dir).unwrap(); + std::fs::write( + skill2_dir.join("SKILL.md"), + r#"--- +name: skill-two +description: Second skill +disable-model-invocation: true +--- +Content 2 +"#, + ) + .unwrap(); + + let skills = load_skills_from_directory(skills_dir); + + assert_eq!(skills.len(), 2); + let names: Vec<_> = skills.iter().map(|s| s.skill_name.as_str()).collect(); + assert!(names.contains(&"skill-one")); + assert!(names.contains(&"skill-two")); + + // 验证 disable_model_invocation 被正确解析 + let skill_two = skills.iter().find(|s| s.skill_name == "skill-two").unwrap(); + assert!(skill_two.disable_model_invocation); + } + + #[test] + fn test_load_skills_from_nonexistent_directory() { + let skills = load_skills_from_directory(std::path::Path::new("/nonexistent/path")); + assert!(skills.is_empty()); + } +} diff --git a/src-tauri/src/commands/subagent_cmd.rs b/src-tauri/src/commands/subagent_cmd.rs new file mode 100644 index 000000000..fc7283123 --- /dev/null +++ b/src-tauri/src/commands/subagent_cmd.rs @@ -0,0 +1,100 @@ +//! SubAgent 调度器命令 +//! +//! 提供 SubAgent 调度功能的 Tauri 命令接口 + +use std::sync::Arc; +use tauri::{AppHandle, State}; +use tokio::sync::RwLock; + +use aster::agents::context::AgentContext; +use aster::agents::subagent_scheduler::{SchedulerConfig, SchedulerExecutionResult, SubAgentTask}; + +use crate::agent::subagent_scheduler::ProxyCastScheduler; +use crate::database::DbConnection; + +/// SubAgent 调度器状态 +pub struct SubAgentSchedulerState { + #[allow(dead_code)] + scheduler: Arc>>, +} + +impl SubAgentSchedulerState { + pub fn new() -> Self { + Self { + scheduler: Arc::new(RwLock::new(None)), + } + } +} + +impl Default for SubAgentSchedulerState { + fn default() -> Self { + Self::new() + } +} + +/// 初始化 SubAgent 调度器 +#[allow(dead_code)] +#[tauri::command] +pub async fn init_subagent_scheduler( + app: AppHandle, + db: State<'_, DbConnection>, + state: State<'_, SubAgentSchedulerState>, + config: Option, +) -> Result<(), String> { + let scheduler = ProxyCastScheduler::new(db.inner().clone()).with_app_handle(app); + + scheduler.init(config).await; + + *state.scheduler.write().await = Some(scheduler); + + Ok(()) +} + +/// 执行 SubAgent 任务 +#[allow(dead_code)] +#[tauri::command] +pub async fn execute_subagent_tasks( + app: AppHandle, + db: State<'_, DbConnection>, + state: State<'_, SubAgentSchedulerState>, + tasks: Vec, + config: Option, +) -> Result { + // 确保调度器已初始化 + let scheduler_guard = state.scheduler.read().await; + + if scheduler_guard.is_none() { + drop(scheduler_guard); + // 自动初始化 + let scheduler = ProxyCastScheduler::new(db.inner().clone()).with_app_handle(app); + scheduler.init(config.clone()).await; + *state.scheduler.write().await = Some(scheduler); + } + + let scheduler_guard = state.scheduler.read().await; + let scheduler = scheduler_guard + .as_ref() + .ok_or_else(|| "调度器初始化失败".to_string())?; + + // 创建父上下文 + let parent_context = AgentContext::new(); + + // 执行任务 + scheduler + .execute(tasks, Some(&parent_context)) + .await + .map_err(|e| e.to_string()) +} + +/// 取消 SubAgent 任务 +#[allow(dead_code)] +#[tauri::command] +pub async fn cancel_subagent_tasks(state: State<'_, SubAgentSchedulerState>) -> Result<(), String> { + let scheduler_guard = state.scheduler.read().await; + + if let Some(scheduler) = scheduler_guard.as_ref() { + scheduler.cancel().await; + } + + Ok(()) +} diff --git a/src-tauri/src/commands/switch_cmd.rs b/src-tauri/src/commands/switch_cmd.rs index 41ceec18d..7dffc8a14 100644 --- a/src-tauri/src/commands/switch_cmd.rs +++ b/src-tauri/src/commands/switch_cmd.rs @@ -43,13 +43,14 @@ pub fn delete_switch_provider( SwitchService::delete_provider(&db, &app_type, &id) } +/// 切换 Provider(异步版本,优化 Windows 性能) #[tauri::command] -pub fn switch_provider( +pub async fn switch_provider( db: State<'_, DbConnection>, app_type: String, id: String, ) -> Result<(), String> { - SwitchService::switch_provider(&db, &app_type, &id) + SwitchService::switch_provider_async(&db, &app_type, &id).await } #[tauri::command] @@ -88,7 +89,7 @@ pub fn check_config_sync_status( /// 从外部配置同步到 ProxyCast #[tauri::command] -pub fn sync_from_external_config( +pub async fn sync_from_external_config( db: State<'_, DbConnection>, app_type: String, ) -> Result { @@ -102,7 +103,7 @@ pub fn sync_from_external_config( .map_err(|e| format!("Failed to sync from external: {e}"))?; // 切换到外部检测到的 provider - SwitchService::switch_provider(&db, &app_type, &external_provider)?; + SwitchService::switch_provider_async(&db, &app_type, &external_provider).await?; Ok(format!("已同步到外部配置的 provider: {external_provider}")) } diff --git a/src-tauri/src/database/dao/brand_persona_dao.rs b/src-tauri/src/database/dao/brand_persona_dao.rs index c5ecb0bad..863f344e6 100644 --- a/src-tauri/src/database/dao/brand_persona_dao.rs +++ b/src-tauri/src/database/dao/brand_persona_dao.rs @@ -10,7 +10,7 @@ use uuid::Uuid; use crate::errors::project_error::PersonaError; use crate::models::project_model::{ BrandPersona, BrandPersonaExtension, BrandPersonaTemplate, BrandTone, - CreateBrandExtensionRequest, DesignConfig, Persona, UpdateBrandExtensionRequest, VisualConfig, + CreateBrandExtensionRequest, DesignConfig, UpdateBrandExtensionRequest, VisualConfig, }; use super::persona_dao::PersonaDao; @@ -418,6 +418,7 @@ mod tests { use super::*; use crate::database::schema::create_tables; use crate::models::project_model::CreatePersonaRequest; + use crate::models::Persona; /// 创建测试数据库连接 fn setup_test_db() -> Connection { diff --git a/src-tauri/src/database/migration.rs b/src-tauri/src/database/migration.rs index b11b7c89b..1c955555f 100644 --- a/src-tauri/src/database/migration.rs +++ b/src-tauri/src/database/migration.rs @@ -497,6 +497,55 @@ pub fn cleanup_legacy_api_key_credentials(conn: &Connection) -> Result Result { + // 检查是否已经迁移过 + let migrated: bool = conn + .query_row( + "SELECT value FROM settings WHERE key = 'migrated_mcp_proxycast_enabled'", + [], + |row| row.get::<_, String>(0), + ) + .map(|v| v == "true") + .unwrap_or(false); + + if migrated { + tracing::debug!("[迁移] MCP proxycast 启用状态已迁移过,跳过"); + return Ok(0); + } + + let updated = conn + .execute( + "UPDATE mcp_servers + SET enabled_proxycast = 1 + WHERE enabled_proxycast = 0 + AND (enabled_claude = 1 OR enabled_codex = 1 OR enabled_gemini = 1)", + [], + ) + .map_err(|e| format!("修复 MCP enabled_proxycast 失败: {e}"))?; + + conn.execute( + "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_mcp_proxycast_enabled', 'true')", + [], + ) + .map_err(|e| format!("标记 MCP proxycast 迁移完成失败: {e}"))?; + + tracing::info!( + "[迁移] MCP proxycast 启用状态修复完成,更新 {} 条记录", + updated + ); + + Ok(updated) +} + /// 当前模型注册表版本 /// 每次更新模型数据结构或添加新 Provider 时,增加此版本号 const MODEL_REGISTRY_VERSION: &str = "2026.01.16.1"; diff --git a/src-tauri/src/database/mod.rs b/src-tauri/src/database/mod.rs index 785f30c0b..71bcff7c3 100644 --- a/src-tauri/src/database/mod.rs +++ b/src-tauri/src/database/mod.rs @@ -84,6 +84,18 @@ pub fn init_database() -> Result { } } + // 修复历史 MCP 导入数据(补齐 enabled_proxycast) + match migration::migrate_mcp_proxycast_enabled(&conn) { + Ok(count) => { + if count > 0 { + tracing::info!("[数据库] 已修复 {} 条 MCP ProxyCast 启用状态", count); + } + } + Err(e) => { + tracing::warn!("[数据库] MCP ProxyCast 启用状态修复失败(非致命): {}", e); + } + } + // 执行统一内容系统迁移(创建默认项目,迁移话题) // _Requirements: 2.1, 2.2, 2.3, 2.4_ match migration_v2::migrate_unified_content_system(&conn) { diff --git a/src-tauri/src/errors/mod.rs b/src-tauri/src/errors/mod.rs index 61511c45a..2af54cc79 100644 --- a/src-tauri/src/errors/mod.rs +++ b/src-tauri/src/errors/mod.rs @@ -8,4 +8,5 @@ pub mod project_error; // 重新导出常用错误类型 +#[allow(unused_imports)] pub use project_error::{MaterialError, MigrationError, PersonaError, ProjectError, TemplateError}; diff --git a/src-tauri/src/flow_monitor/stream_rebuilder.rs b/src-tauri/src/flow_monitor/stream_rebuilder.rs index 0ab23e4c6..ab9648840 100644 --- a/src-tauri/src/flow_monitor/stream_rebuilder.rs +++ b/src-tauri/src/flow_monitor/stream_rebuilder.rs @@ -69,15 +69,6 @@ struct ToolCallBuilder { } impl ToolCallBuilder { - fn new() -> Self { - Self { - id: None, - tool_type: "function".to_string(), - function_name: None, - arguments: String::new(), - } - } - fn build(self) -> Option { let id = self.id?; let name = self.function_name?; diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index ebcf8d9cd..20ebb5730 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -50,6 +50,12 @@ pub mod tray; pub mod voice; pub mod workspace; +// Skills 集成模块 +pub mod skills; + +// MCP 集成模块 +pub mod mcp; + // 内部模块 mod commands; mod config; diff --git a/src-tauri/src/mcp/README.md b/src-tauri/src/mcp/README.md new file mode 100644 index 000000000..33611c9e6 --- /dev/null +++ b/src-tauri/src/mcp/README.md @@ -0,0 +1,42 @@ +# MCP 模块 + +MCP(Model Context Protocol)集成模块,提供 MCP 协议的客户端实现。 + +## 模块结构 + +| 文件 | 说明 | +|------|------| +| `mod.rs` | 模块导出和文档 | +| `types.rs` | MCP 数据类型定义(配置、工具、提示词、资源、错误) | +| `client.rs` | MCP 客户端实现(rmcp ClientHandler) | +| `manager.rs` | MCP 客户端管理器(连接池、缓存、生命周期) | +| `tool_converter.rs` | 工具格式转换器(OpenAI/Anthropic/Gemini) | + +## 功能概览 + +### 服务器生命周期管理 +- 启动/停止 MCP 服务器进程 +- stdio 传输连接 +- 状态监控和事件通知 + +### 工具管理 +- 工具发现和缓存 +- 工具调用路由 +- 名称冲突解决(服务器前缀) + +### 格式转换 +- MCP → OpenAI function calling +- MCP → Anthropic tool use +- MCP → Gemini function declaration + +## 依赖 + +- `rmcp`: Rust MCP SDK +- `tokio`: 异步运行时 +- `serde`: 序列化/反序列化 +- `thiserror`: 错误类型定义 + +## 相关文档 + +- 设计文档: `.kiro/specs/mcp-integration/design.md` +- 需求文档: `.kiro/specs/mcp-integration/requirements.md` diff --git a/src-tauri/src/mcp/client.rs b/src-tauri/src/mcp/client.rs new file mode 100644 index 000000000..bf660b032 --- /dev/null +++ b/src-tauri/src/mcp/client.rs @@ -0,0 +1,399 @@ +//! MCP 客户端实现 +//! +//! 本模块实现 rmcp 的 ClientHandler trait,处理: +//! - 客户端信息返回 +//! - 进度通知处理 +//! - 日志消息处理 +//! - 与 Tauri 事件系统的集成 + +#![allow(dead_code)] + +use rmcp::{ + model::{ + ClientCapabilities, ClientInfo, Implementation, LoggingMessageNotification, + LoggingMessageNotificationMethod, LoggingMessageNotificationParam, ProgressNotification, + ProgressNotificationMethod, ProgressNotificationParam, ProtocolVersion, ServerNotification, + }, + service::NotificationContext, + ClientHandler, RoleClient, +}; +use std::sync::Arc; +use tauri::Emitter; +use tokio::sync::{mpsc, Mutex}; +use tracing::{debug, info, warn}; + +/// 进度通知事件 Payload +#[derive(Debug, Clone, serde::Serialize)] +pub struct McpProgressPayload { + pub server_name: String, + pub progress_token: String, + pub progress: f64, + pub total: Option, + pub message: Option, +} + +/// 日志消息事件 Payload +#[derive(Debug, Clone, serde::Serialize)] +pub struct McpLogMessagePayload { + pub server_name: String, + pub level: String, + pub logger: Option, + pub data: serde_json::Value, +} + +/// ProxyCast MCP 客户端处理器 +/// +/// 实现 rmcp::ClientHandler trait,处理 MCP 服务器的通知和回调 +pub struct ProxyCastMcpClient { + /// Tauri AppHandle(用于发送事件) + app_handle: Option, + /// 服务器名称(用于事件标识) + server_name: String, + /// 通知订阅者(用于内部通知分发) + notification_handlers: Arc>>>, +} + +impl ProxyCastMcpClient { + /// 创建新的 MCP 客户端处理器 + /// + /// # Arguments + /// * `server_name` - MCP 服务器名称,用于事件标识 + /// * `app_handle` - Tauri AppHandle,用于发送事件到前端 + pub fn new(server_name: String, app_handle: Option) -> Self { + Self { + app_handle, + server_name, + notification_handlers: Arc::new(Mutex::new(Vec::new())), + } + } + + /// 获取通知处理器的引用(用于订阅通知) + pub fn notification_handlers(&self) -> Arc>>> { + self.notification_handlers.clone() + } + + /// 订阅服务器通知 + /// + /// 返回一个接收器,用于接收来自 MCP 服务器的通知 + pub async fn subscribe(&self) -> mpsc::Receiver { + let (tx, rx) = mpsc::channel(16); + self.notification_handlers.lock().await.push(tx); + rx + } + + /// 发送 Tauri 事件到前端 + fn emit_event(&self, event: &str, payload: T) { + if let Some(ref app_handle) = self.app_handle { + if let Err(e) = app_handle.emit(event, payload) { + warn!( + server_name = %self.server_name, + event = %event, + error = %e, + "发送 Tauri 事件失败" + ); + } + } + } +} + +impl ClientHandler for ProxyCastMcpClient { + /// 返回客户端信息 + /// + /// 提供 ProxyCast 客户端的标识信息,包括: + /// - 协议版本 + /// - 客户端能力(采样支持) + /// - 客户端实现信息 + fn get_info(&self) -> ClientInfo { + ClientInfo { + protocol_version: ProtocolVersion::V_2025_03_26, + capabilities: ClientCapabilities::builder().enable_sampling().build(), + client_info: Implementation { + name: "proxycast".to_string(), + version: env!("CARGO_PKG_VERSION").to_string(), + icons: None, + title: Some("ProxyCast MCP Client".to_string()), + website_url: Some("https://github.com/aiclientproxy/proxycast".to_string()), + }, + } + } + + /// 处理进度通知 + /// + /// 当 MCP 服务器发送进度更新时调用此方法。 + /// 进度信息会: + /// 1. 记录到日志 + /// 2. 发送到 Tauri 事件系统(前端可监听) + /// 3. 分发给内部通知订阅者 + async fn on_progress( + &self, + params: ProgressNotificationParam, + context: NotificationContext, + ) { + // 记录进度日志 + debug!( + server_name = %self.server_name, + progress_token = ?params.progress_token, + progress = params.progress, + total = ?params.total, + "收到 MCP 进度通知" + ); + + // 发送 Tauri 事件到前端 + let payload = McpProgressPayload { + server_name: self.server_name.clone(), + progress_token: format!("{:?}", params.progress_token), + progress: params.progress, + total: params.total, + message: None, + }; + self.emit_event("mcp:progress", payload); + + // 分发给内部通知订阅者 + let notification = ServerNotification::ProgressNotification(ProgressNotification { + params: params.clone(), + method: ProgressNotificationMethod, + extensions: context.extensions.clone(), + }); + + let handlers = self.notification_handlers.lock().await; + for handler in handlers.iter() { + let _ = handler.try_send(notification.clone()); + } + } + + /// 处理日志消息通知 + /// + /// 当 MCP 服务器发送日志消息时调用此方法。 + /// 日志消息会: + /// 1. 根据级别记录到本地日志 + /// 2. 发送到 Tauri 事件系统(前端可监听) + /// 3. 分发给内部通知订阅者 + async fn on_logging_message( + &self, + params: LoggingMessageNotificationParam, + context: NotificationContext, + ) { + // 根据日志级别记录 + let level_str = format!("{:?}", params.level); + match params.level { + rmcp::model::LoggingLevel::Debug => { + debug!( + server_name = %self.server_name, + logger = ?params.logger, + data = ?params.data, + "MCP 服务器日志 [DEBUG]" + ); + } + rmcp::model::LoggingLevel::Info => { + info!( + server_name = %self.server_name, + logger = ?params.logger, + data = ?params.data, + "MCP 服务器日志 [INFO]" + ); + } + rmcp::model::LoggingLevel::Notice => { + info!( + server_name = %self.server_name, + logger = ?params.logger, + data = ?params.data, + "MCP 服务器日志 [NOTICE]" + ); + } + rmcp::model::LoggingLevel::Warning => { + warn!( + server_name = %self.server_name, + logger = ?params.logger, + data = ?params.data, + "MCP 服务器日志 [WARNING]" + ); + } + rmcp::model::LoggingLevel::Error => { + tracing::error!( + server_name = %self.server_name, + logger = ?params.logger, + data = ?params.data, + "MCP 服务器日志 [ERROR]" + ); + } + rmcp::model::LoggingLevel::Critical => { + tracing::error!( + server_name = %self.server_name, + logger = ?params.logger, + data = ?params.data, + "MCP 服务器日志 [CRITICAL]" + ); + } + rmcp::model::LoggingLevel::Alert => { + tracing::error!( + server_name = %self.server_name, + logger = ?params.logger, + data = ?params.data, + "MCP 服务器日志 [ALERT]" + ); + } + rmcp::model::LoggingLevel::Emergency => { + tracing::error!( + server_name = %self.server_name, + logger = ?params.logger, + data = ?params.data, + "MCP 服务器日志 [EMERGENCY]" + ); + } + } + + // 发送 Tauri 事件到前端 + let payload = McpLogMessagePayload { + server_name: self.server_name.clone(), + level: level_str, + logger: params.logger.clone(), + data: params.data.clone(), + }; + self.emit_event("mcp:log_message", payload); + + // 分发给内部通知订阅者 + let notification = + ServerNotification::LoggingMessageNotification(LoggingMessageNotification { + params: params.clone(), + method: LoggingMessageNotificationMethod, + extensions: context.extensions.clone(), + }); + + let handlers = self.notification_handlers.lock().await; + for handler in handlers.iter() { + let _ = handler.try_send(notification.clone()); + } + } +} + +/// MCP 客户端包装器 +/// +/// 封装 rmcp 客户端和相关状态 +pub struct McpClientWrapper { + /// 服务器名称 + pub server_name: String, + /// 服务器配置 + pub config: super::types::McpServerConfig, + /// 子进程句柄 + pub process: Option, + /// 服务器能力信息 + pub server_info: Option, + /// 客户端处理器 + pub client_handler: Arc, + /// rmcp 运行服务(用于发送请求) + pub running_service: + Option>, +} + +impl McpClientWrapper { + /// 创建新的客户端包装器 + pub fn new( + server_name: String, + config: super::types::McpServerConfig, + app_handle: Option, + ) -> Self { + let client_handler = Arc::new(ProxyCastMcpClient::new(server_name.clone(), app_handle)); + + Self { + server_name, + config, + process: None, + server_info: None, + client_handler, + running_service: None, + } + } + + /// 获取客户端处理器的引用 + pub fn handler(&self) -> Arc { + self.client_handler.clone() + } + + /// 设置子进程句柄 + pub fn set_process(&mut self, process: tokio::process::Child) { + self.process = Some(process); + } + + /// 设置服务器能力信息 + pub fn set_server_info(&mut self, info: super::types::McpServerCapabilities) { + self.server_info = Some(info); + } + + /// 设置 rmcp 运行服务 + pub fn set_running_service( + &mut self, + service: rmcp::service::RunningService, + ) { + self.running_service = Some(service); + } + + /// 获取 rmcp 运行服务的引用 + pub fn running_service( + &self, + ) -> Option<&rmcp::service::RunningService> { + self.running_service.as_ref() + } + + /// 终止子进程 + pub async fn kill_process(&mut self) -> Result<(), std::io::Error> { + if let Some(ref mut process) = self.process { + process.kill().await?; + } + self.process = None; + self.running_service = None; + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_client_info() { + let client = ProxyCastMcpClient::new("test-server".to_string(), None); + let info = client.get_info(); + + assert_eq!(info.client_info.name, "proxycast"); + assert_eq!(info.client_info.version, env!("CARGO_PKG_VERSION")); + assert_eq!( + info.client_info.title, + Some("ProxyCast MCP Client".to_string()) + ); + assert_eq!(info.protocol_version, ProtocolVersion::V_2025_03_26); + } + + #[test] + fn test_client_wrapper_creation() { + let config = super::super::types::McpServerConfig { + command: "test-command".to_string(), + args: vec!["--arg1".to_string()], + env: std::collections::HashMap::new(), + cwd: None, + timeout: 30, + }; + + let wrapper = McpClientWrapper::new("test-server".to_string(), config.clone(), None); + + assert_eq!(wrapper.server_name, "test-server"); + assert_eq!(wrapper.config.command, "test-command"); + assert!(wrapper.process.is_none()); + assert!(wrapper.server_info.is_none()); + } + + #[tokio::test] + async fn test_notification_subscription() { + let client = ProxyCastMcpClient::new("test-server".to_string(), None); + + // 订阅通知 + let mut rx = client.subscribe().await; + + // 验证订阅者已添加 + let handlers = client.notification_handlers.lock().await; + assert_eq!(handlers.len(), 1); + drop(handlers); + + // 验证接收器可用(不会阻塞) + assert!(rx.try_recv().is_err()); // 应该是空的 + } +} diff --git a/src-tauri/src/mcp/manager.rs b/src-tauri/src/mcp/manager.rs new file mode 100644 index 000000000..83f0d0e99 --- /dev/null +++ b/src-tauri/src/mcp/manager.rs @@ -0,0 +1,2426 @@ +//! MCP 客户端管理器 +//! +//! 本模块提供 MCP 客户端的集中管理,包括: +//! - 服务器生命周期管理(启动、停止、重启) +//! - 客户端连接池管理 +//! - 工具定义缓存 +//! - Tauri 事件发送 +//! +//! # 架构设计 +//! +//! ```text +//! ┌─────────────────────────────────────────────────────────┐ +//! │ McpClientManager │ +//! │ ┌─────────────────────────────────────────────────┐ │ +//! │ │ clients (连接池) │ │ +//! │ │ ┌─────────┐ ┌─────────┐ ┌─────────┐ │ │ +//! │ │ │ Client1 │ │ Client2 │ │ Client3 │ │ │ +//! │ │ └─────────┘ └─────────┘ └─────────┘ │ │ +//! │ └─────────────────────────────────────────────────┘ │ +//! │ ┌─────────────────────────────────────────────────┐ │ +//! │ │ tool_cache (工具缓存) │ │ +//! │ │ 缓存所有运行中服务器的工具定义 │ │ +//! │ └─────────────────────────────────────────────────┘ │ +//! └─────────────────────────────────────────────────────────┘ +//! ``` + +#![allow(dead_code)] + +use std::collections::HashMap; +use std::process::Stdio; +use std::sync::Arc; +use std::time::Duration; +use tauri::Emitter; +use tokio::io::AsyncReadExt; +use tokio::process::Command; +use tokio::sync::RwLock; +use tracing::{debug, error, info, warn}; + +use rmcp::transport::TokioChildProcess; +use rmcp::ServiceExt; + +use super::client::McpClientWrapper; +use super::types::*; + +/// MCP 客户端管理器 +/// +/// 负责管理所有 MCP 服务器的连接和生命周期。 +/// +/// # 功能 +/// +/// - **连接池管理**: 维护所有运行中的 MCP 客户端连接 +/// - **工具缓存**: 缓存工具定义以避免重复查询 +/// - **事件通知**: 通过 Tauri 事件系统通知前端状态变化 +/// +/// # 线程安全 +/// +/// 所有内部状态都使用 `Arc>` 包装,支持并发访问。 +/// +/// # 示例 +/// +/// ```rust,ignore +/// let manager = McpClientManager::new(Some(app_handle)); +/// +/// // 启动服务器 +/// manager.start_server("my-server", &config).await?; +/// +/// // 获取工具列表 +/// let tools = manager.list_tools().await?; +/// +/// // 调用工具 +/// let result = manager.call_tool("my-tool", args).await?; +/// +/// // 停止服务器 +/// manager.stop_server("my-server").await?; +/// ``` +pub struct McpClientManager { + /// 运行中的客户端 (server_name -> client) + /// + /// 使用 HashMap 存储所有活跃的 MCP 客户端连接。 + /// 键为服务器名称,值为客户端包装器。 + clients: Arc>>, + + /// 工具定义缓存 + /// + /// 缓存所有运行中服务器的工具定义。 + /// 当服务器启动或停止时,缓存会被失效。 + /// 使用 Option 表示缓存状态: + /// - None: 缓存无效,需要重新获取 + /// - Some(tools): 缓存有效 + tool_cache: Arc>>>, + + /// Tauri AppHandle(用于发送事件) + /// + /// 用于向前端发送 MCP 相关事件,如: + /// - mcp:server_started + /// - mcp:server_stopped + /// - mcp:server_error + /// - mcp:tools_updated + app_handle: Option, +} + +impl McpClientManager { + /// 创建新的管理器实例 + /// + /// # Arguments + /// + /// * `app_handle` - Tauri AppHandle,用于发送事件到前端。 + /// 如果为 None,则不会发送事件。 + /// + /// # Returns + /// + /// 返回初始化的 McpClientManager 实例,连接池和缓存均为空。 + pub fn new(app_handle: Option) -> Self { + info!("创建 MCP 客户端管理器"); + Self { + clients: Arc::new(RwLock::new(HashMap::new())), + tool_cache: Arc::new(RwLock::new(None)), + app_handle, + } + } + + /// 设置 Tauri AppHandle(用于发送前端事件) + pub fn set_app_handle(&mut self, app_handle: tauri::AppHandle) { + self.app_handle = Some(app_handle); + } + + // ======================================================================== + // 连接池管理方法 + // ======================================================================== + + /// 获取客户端连接池的只读引用 + /// + /// 用于需要遍历所有客户端的场景。 + pub fn clients(&self) -> Arc>> { + self.clients.clone() + } + + /// 获取指定服务器的客户端(检查是否存在) + /// + /// # Arguments + /// + /// * `name` - 服务器名称 + /// + /// # Returns + /// + /// 如果服务器正在运行,返回 true;否则返回 false。 + /// + /// 注意:由于 McpClientWrapper 包含不可克隆的字段(如 tokio::process::Child), + /// 我们不能直接返回客户端的克隆。如需操作客户端,请使用 clients() 获取连接池引用。 + pub async fn has_client(&self, name: &str) -> bool { + let clients = self.clients.read().await; + clients.contains_key(name) + } + + /// 获取指定服务器的配置(如果存在) + /// + /// # Arguments + /// + /// * `name` - 服务器名称 + /// + /// # Returns + /// + /// 如果服务器正在运行,返回 Some(配置的克隆); + /// 否则返回 None。 + pub async fn get_client_config(&self, name: &str) -> Option { + let clients = self.clients.read().await; + clients.get(name).map(|c| c.config.clone()) + } + + /// 获取指定服务器的能力信息(如果存在) + /// + /// # Arguments + /// + /// * `name` - 服务器名称 + /// + /// # Returns + /// + /// 如果服务器正在运行且有能力信息,返回 Some(能力信息的克隆); + /// 否则返回 None。 + pub async fn get_client_capabilities(&self, name: &str) -> Option { + let clients = self.clients.read().await; + clients.get(name).and_then(|c| c.server_info.clone()) + } + + /// 添加客户端到连接池 + /// + /// # Arguments + /// + /// * `name` - 服务器名称 + /// * `client` - 客户端包装器 + /// + /// # Returns + /// + /// 如果服务器已存在,返回错误;否则添加成功。 + pub async fn add_client(&self, name: String, client: McpClientWrapper) -> Result<(), McpError> { + let mut clients = self.clients.write().await; + if clients.contains_key(&name) { + return Err(McpError::ServerAlreadyRunning(name)); + } + debug!(server_name = %name, "添加客户端到连接池"); + clients.insert(name, client); + Ok(()) + } + + /// 从连接池移除客户端 + /// + /// # Arguments + /// + /// * `name` - 服务器名称 + /// + /// # Returns + /// + /// 如果服务器存在,返回移除的客户端包装器; + /// 否则返回 None。 + pub async fn remove_client(&self, name: &str) -> Option { + let mut clients = self.clients.write().await; + let removed = clients.remove(name); + if removed.is_some() { + debug!(server_name = %name, "从连接池移除客户端"); + } + removed + } + + /// 获取所有运行中的服务器名称 + /// + /// # Returns + /// + /// 返回所有运行中服务器的名称列表。 + pub async fn get_running_servers(&self) -> Vec { + let clients = self.clients.read().await; + clients.keys().cloned().collect() + } + + /// 获取运行中的服务器数量 + pub async fn running_server_count(&self) -> usize { + let clients = self.clients.read().await; + clients.len() + } + + // ======================================================================== + // 缓存管理方法 + // ======================================================================== + + /// 失效工具缓存 + /// + /// 当服务器启动或停止时调用此方法, + /// 确保下次获取工具列表时会重新查询所有服务器。 + pub async fn invalidate_tool_cache(&self) { + let mut cache = self.tool_cache.write().await; + if cache.is_some() { + debug!("失效工具缓存"); + } + *cache = None; + } + + /// 检查工具缓存是否有效 + pub async fn is_tool_cache_valid(&self) -> bool { + let cache = self.tool_cache.read().await; + cache.is_some() + } + + /// 获取缓存的工具列表(如果有效) + /// + /// # Returns + /// + /// 如果缓存有效,返回 Some(工具列表); + /// 否则返回 None。 + pub async fn get_cached_tools(&self) -> Option> { + let cache = self.tool_cache.read().await; + cache.clone() + } + + /// 更新工具缓存 + /// + /// # Arguments + /// + /// * `tools` - 新的工具列表 + pub async fn update_tool_cache(&self, tools: Vec) { + let mut cache = self.tool_cache.write().await; + debug!(tool_count = tools.len(), "更新工具缓存"); + *cache = Some(tools); + } + + // ======================================================================== + // 事件发送方法 + // ======================================================================== + + /// 发送 Tauri 事件到前端 + /// + /// # Arguments + /// + /// * `event` - 事件名称 + /// * `payload` - 事件数据 + pub fn emit_event(&self, event: &str, payload: T) { + if let Some(ref app_handle) = self.app_handle { + if let Err(e) = app_handle.emit(event, payload) { + warn!( + event = %event, + error = %e, + "发送 Tauri 事件失败" + ); + } else { + debug!(event = %event, "发送 Tauri 事件"); + } + } + } + + /// 发送服务器启动事件 + pub fn emit_server_started( + &self, + server_name: &str, + server_info: Option, + ) { + info!(server_name = %server_name, "MCP 服务器已启动"); + self.emit_event( + "mcp:server_started", + McpServerStartedPayload { + server_name: server_name.to_string(), + server_info, + }, + ); + } + + /// 发送服务器停止事件 + pub fn emit_server_stopped(&self, server_name: &str) { + info!(server_name = %server_name, "MCP 服务器已停止"); + self.emit_event( + "mcp:server_stopped", + McpServerStoppedPayload { + server_name: server_name.to_string(), + }, + ); + } + + /// 发送服务器错误事件 + pub fn emit_server_error(&self, server_name: &str, error: &str) { + warn!(server_name = %server_name, error = %error, "MCP 服务器错误"); + self.emit_event( + "mcp:server_error", + McpServerErrorPayload { + server_name: server_name.to_string(), + error: error.to_string(), + }, + ); + } + + /// 发送工具列表更新事件 + pub fn emit_tools_updated(&self, tools: Vec) { + debug!(tool_count = tools.len(), "工具列表已更新"); + self.emit_event("mcp:tools_updated", McpToolsUpdatedPayload { tools }); + } + + // ======================================================================== + // 服务器生命周期管理方法 + // ======================================================================== + + /// 启动 MCP 服务器 + /// + /// # Arguments + /// + /// * `name` - 服务器名称 + /// * `config` - 服务器配置 + /// + /// # Returns + /// + /// 成功返回 Ok(()),失败返回错误。 + /// + /// # 实现步骤(Task 4.2) + /// + /// 1. 检查服务器是否已运行 + /// 2. 启动子进程 + /// 3. 建立 stdio 连接 + /// 4. 初始化 MCP 客户端 + /// 5. 失效工具缓存 + /// 6. 发送 mcp:server_started 事件 + pub async fn start_server(&self, name: &str, config: &McpServerConfig) -> Result<(), McpError> { + info!(server_name = %name, command = %config.command, "启动 MCP 服务器"); + + // 1. 检查服务器是否已运行 + if self.is_server_running(name).await { + return Err(McpError::ServerAlreadyRunning(name.to_string())); + } + + // 2. 构建命令 + let mut command = Command::new(&config.command); + command.args(&config.args); + + // 设置环境变量 + for (key, value) in &config.env { + command.env(key, value); + } + + // macOS GUI 应用的 PATH 通常不完整,需要补充常见的命令路径 + // 确保 npx/node/uvx 等命令可被找到 + if !config.env.contains_key("PATH") { + let current_path = std::env::var("PATH").unwrap_or_default(); + let home = std::env::var("HOME").unwrap_or_else(|_| "/Users/unknown".to_string()); + let extra_paths = [ + format!("{home}/.nvm/versions/node/*/bin"), + format!("{home}/.local/bin"), + format!("{home}/.cargo/bin"), + format!("{home}/Library/pnpm"), + format!("{home}/.bun/bin"), + "/usr/local/bin".to_string(), + "/opt/homebrew/bin".to_string(), + "/opt/homebrew/sbin".to_string(), + ]; + // 用 glob 展开 nvm 路径,取最新版本 + let mut resolved_paths: Vec = Vec::new(); + for p in &extra_paths { + if p.contains('*') { + if let Ok(entries) = glob::glob(p) { + let mut matched: Vec = entries + .filter_map(|e| e.ok()) + .map(|e| e.to_string_lossy().to_string()) + .collect(); + matched.sort(); + if let Some(last) = matched.last() { + resolved_paths.push(last.clone()); + } + } + } else if std::path::Path::new(p).exists() { + resolved_paths.push(p.clone()); + } + } + if !resolved_paths.is_empty() { + let merged = if current_path.is_empty() { + resolved_paths.join(":") + } else { + format!("{}:{}", resolved_paths.join(":"), current_path) + }; + command.env("PATH", &merged); + debug!(server_name = %name, "补充 PATH: {}", merged); + } + } + + // 设置工作目录 + if let Some(ref cwd) = config.cwd { + command.current_dir(cwd); + } + + // Unix 系统设置进程组(使子进程独立于父进程组) + #[cfg(unix)] + command.process_group(0); + + // 3. 启动子进程并建立 stdio 连接 + let spawn_result = TokioChildProcess::builder(command) + .stderr(Stdio::piped()) + .spawn(); + + let (transport, mut stderr_opt) = match spawn_result { + Ok(result) => result, + Err(e) => { + let error_msg = format!("无法启动服务器进程: {}", e); + error!(server_name = %name, error = %e, "启动 MCP 服务器进程失败"); + self.emit_server_error(name, &error_msg); + return Err(McpError::ProcessSpawnFailed(error_msg)); + } + }; + + // 启动 stderr 读取任务(用于错误诊断) + let stderr_task = if let Some(mut stderr) = stderr_opt.take() { + Some(tokio::spawn(async move { + let mut all_stderr = Vec::new(); + let _ = stderr.read_to_end(&mut all_stderr).await; + String::from_utf8_lossy(&all_stderr).into_owned() + })) + } else { + None + }; + + // 4. 初始化 MCP 客户端 + let client_handler = + super::client::ProxyCastMcpClient::new(name.to_string(), self.app_handle.clone()); + + // 连接超时:至少 60 秒,避免 npx 首次下载时超时 + let timeout_secs = std::cmp::max(config.timeout, 60); + let timeout = Duration::from_secs(timeout_secs); + let connect_result = tokio::time::timeout(timeout, client_handler.serve(transport)).await; + + let running_service = match connect_result { + Ok(Ok(service)) => service, + Ok(Err(e)) => { + // 获取 stderr 内容用于诊断 + let stderr_content = if let Some(task) = stderr_task { + task.await.unwrap_or_default() + } else { + String::new() + }; + + let error_msg = if stderr_content.is_empty() { + format!("MCP 连接失败: {}", e) + } else { + format!("MCP 连接失败: {}. Stderr: {}", e, stderr_content) + }; + + error!( + server_name = %name, + error = %e, + stderr = %stderr_content, + "MCP 客户端初始化失败" + ); + self.emit_server_error(name, &error_msg); + return Err(McpError::ConnectionFailed(error_msg)); + } + Err(_) => { + let error_msg = format!("MCP 连接超时({}秒)", timeout_secs); + error!(server_name = %name, timeout = timeout_secs, "MCP 连接超时"); + self.emit_server_error(name, &error_msg); + return Err(McpError::Timeout); + } + }; + + // 获取服务器信息 + let server_info = running_service + .peer_info() + .map(|info| McpServerCapabilities { + name: info.server_info.name.clone(), + version: info.server_info.version.clone(), + supports_tools: info + .capabilities + .tools + .as_ref() + .map(|_| true) + .unwrap_or(false), + supports_prompts: info + .capabilities + .prompts + .as_ref() + .map(|_| true) + .unwrap_or(false), + supports_resources: info + .capabilities + .resources + .as_ref() + .map(|_| true) + .unwrap_or(false), + }); + + // 创建客户端包装器 + let mut wrapper = super::client::McpClientWrapper::new( + name.to_string(), + config.clone(), + self.app_handle.clone(), + ); + if let Some(ref info) = server_info { + wrapper.set_server_info(info.clone()); + } + wrapper.set_running_service(running_service); + + // 添加到连接池 + self.add_client(name.to_string(), wrapper).await?; + + // 5. 失效工具缓存 + self.invalidate_tool_cache().await; + + // 6. 发送 mcp:server_started 事件 + self.emit_server_started(name, server_info); + + info!(server_name = %name, "MCP 服务器启动成功"); + Ok(()) + } + + /// 停止 MCP 服务器 + /// + /// # Arguments + /// + /// * `name` - 服务器名称 + /// + /// # Returns + /// + /// 成功返回 Ok(()),失败返回错误。 + /// 如果服务器未运行,也返回 Ok()(幂等操作)。 + /// + /// # 实现步骤(Task 4.2) + /// + /// 1. 检查服务器是否在运行 + /// 2. 终止子进程 + /// 3. 清理客户端连接 + /// 4. 失效工具缓存 + /// 5. 发送 mcp:server_stopped 事件 + pub async fn stop_server(&self, name: &str) -> Result<(), McpError> { + info!(server_name = %name, "停止 MCP 服务器"); + + // 1. 检查服务器是否在运行 + if !self.is_server_running(name).await { + debug!(server_name = %name, "服务器未运行,跳过停止操作"); + return Ok(()); // 幂等操作 + } + + // 2. 从连接池移除客户端 + let mut wrapper = match self.remove_client(name).await { + Some(w) => w, + None => { + debug!(server_name = %name, "客户端已被移除"); + return Ok(()); + } + }; + + // 3. 取消 rmcp 服务(如果存在) + if let Some(ref service) = wrapper.running_service { + let cancellation_token = service.cancellation_token(); + cancellation_token.cancel(); + debug!(server_name = %name, "已取消 MCP 服务"); + } + + // 4. 终止子进程 + if let Err(e) = wrapper.kill_process().await { + warn!( + server_name = %name, + error = %e, + "终止子进程时出错(可能已退出)" + ); + // 不返回错误,因为进程可能已经退出 + } + + // 5. 失效工具缓存 + self.invalidate_tool_cache().await; + + // 6. 发送 mcp:server_stopped 事件 + self.emit_server_stopped(name); + + info!(server_name = %name, "MCP 服务器已停止"); + Ok(()) + } + + /// 检查服务器是否在运行 + /// + /// # Arguments + /// + /// * `name` - 服务器名称 + /// + /// # Returns + /// + /// 如果服务器正在运行返回 true,否则返回 false。 + pub async fn is_server_running(&self, name: &str) -> bool { + let clients = self.clients.read().await; + clients.contains_key(name) + } + + /// 重启 MCP 服务器 + /// + /// 先停止服务器,然后重新启动。 + /// + /// # Arguments + /// + /// * `name` - 服务器名称 + /// * `config` - 服务器配置 + /// + /// # Returns + /// + /// 成功返回 Ok(()),失败返回错误。 + pub async fn restart_server( + &self, + name: &str, + config: &McpServerConfig, + ) -> Result<(), McpError> { + // 先停止(忽略未运行的错误) + let _ = self.stop_server(name).await; + // 再启动 + self.start_server(name, config).await + } + + // ======================================================================== + // 工具管理方法 + // ======================================================================== + + /// 获取所有工具定义 + /// + /// 从所有运行中的服务器获取工具定义,并使用缓存优化性能。 + /// + /// # Returns + /// + /// 返回所有可用工具的定义列表。 + /// + /// # 实现步骤(Task 4.3) + /// + /// 1. 检查缓存是否有效 + /// 2. 如果缓存有效,直接返回缓存 + /// 3. 从所有运行中的服务器获取工具 + /// 4. 解决名称冲突(添加服务器前缀) + /// 5. 更新缓存 + /// 6. 发送 mcp:tools_updated 事件 + /// 7. 返回工具列表 + pub async fn list_tools(&self) -> Result, McpError> { + // 1. 检查缓存是否有效 + if let Some(cached_tools) = self.get_cached_tools().await { + debug!(tool_count = cached_tools.len(), "返回缓存的工具列表"); + return Ok(cached_tools); + } + + // 2. 从所有运行中的服务器获取工具 + let mut all_tools: Vec = Vec::new(); + let clients = self.clients.read().await; + + for (server_name, wrapper) in clients.iter() { + // 检查服务器是否支持工具 + if let Some(ref info) = wrapper.server_info { + if !info.supports_tools { + debug!(server_name = %server_name, "服务器不支持工具,跳过"); + continue; + } + } + + // 获取 rmcp 服务 + let service = match wrapper.running_service() { + Some(s) => s, + None => { + warn!(server_name = %server_name, "服务器无运行服务,跳过"); + continue; + } + }; + + // 调用 list_tools(使用 list_all_tools 获取所有工具) + match service.list_all_tools().await { + Ok(tools) => { + debug!( + server_name = %server_name, + tool_count = tools.len(), + "获取服务器工具列表成功" + ); + for tool in tools { + all_tools.push(McpToolDefinition { + name: tool.name.to_string(), + description: tool + .description + .clone() + .map(|s| s.to_string()) + .unwrap_or_default(), + input_schema: serde_json::Value::Object((*tool.input_schema).clone()), + server_name: server_name.clone(), + }); + } + } + Err(e) => { + warn!( + server_name = %server_name, + error = %e, + "获取服务器工具列表失败" + ); + // 继续处理其他服务器,不中断 + } + } + } + drop(clients); + + // 3. 解决名称冲突(添加服务器前缀) + let resolved_tools = Self::resolve_tool_name_conflicts(all_tools); + + // 4. 更新缓存 + self.update_tool_cache(resolved_tools.clone()).await; + + // 5. 发送 mcp:tools_updated 事件 + self.emit_tools_updated(resolved_tools.clone()); + + info!(tool_count = resolved_tools.len(), "工具列表已更新"); + Ok(resolved_tools) + } + + /// 解决工具名称冲突 + /// + /// 当多个服务器提供同名工具时,为冲突的工具名称添加服务器前缀。 + /// + /// # Arguments + /// + /// * `tools` - 原始工具列表 + /// + /// # Returns + /// + /// 返回解决冲突后的工具列表。 + fn resolve_tool_name_conflicts(tools: Vec) -> Vec { + use std::collections::HashSet; + + // 统计每个工具名称出现的次数 + let mut name_counts: HashMap = HashMap::new(); + for tool in &tools { + *name_counts.entry(tool.name.clone()).or_insert(0) += 1; + } + + // 找出有冲突的名称 + let conflicting_names: HashSet = name_counts + .into_iter() + .filter(|(_, count)| *count > 1) + .map(|(name, _)| name) + .collect(); + + // 为冲突的工具添加服务器前缀 + tools + .into_iter() + .map(|mut tool| { + if conflicting_names.contains(&tool.name) { + debug!( + original_name = %tool.name, + server_name = %tool.server_name, + "工具名称冲突,添加服务器前缀" + ); + tool.name = format!("{}_{}", tool.server_name, tool.name); + } + tool + }) + .collect() + } + + /// 调用工具 + /// + /// # Arguments + /// + /// * `tool_name` - 工具名称(可能包含服务器前缀) + /// * `arguments` - 工具参数 + /// + /// # Returns + /// + /// 返回工具调用结果。 + /// + /// # 实现步骤(Task 4.3) + /// + /// 1. 解析工具名称,确定目标服务器 + /// 2. 路由到正确的客户端 + /// 3. 执行工具调用 + /// 4. 转换结果为 McpToolResult + /// 5. 返回结果 + pub async fn call_tool( + &self, + tool_name: &str, + arguments: serde_json::Value, + ) -> Result { + info!(tool_name = %tool_name, "调用 MCP 工具"); + + // 1. 解析工具名称,确定目标服务器和实际工具名 + let (server_name, actual_tool_name) = self.resolve_tool_target(tool_name).await?; + + debug!( + tool_name = %tool_name, + server_name = %server_name, + actual_tool_name = %actual_tool_name, + "解析工具目标" + ); + + // 2. 获取目标服务器的客户端 + let clients = self.clients.read().await; + let wrapper = clients + .get(&server_name) + .ok_or_else(|| McpError::ServerNotRunning(server_name.clone()))?; + + let service = wrapper + .running_service() + .ok_or_else(|| McpError::ServerNotRunning(server_name.clone()))?; + + // 3. 构建工具调用参数 + let args = match arguments { + serde_json::Value::Object(map) => Some(map), + serde_json::Value::Null => None, + _ => { + return Err(McpError::ToolCallFailed( + "参数必须是 JSON 对象或 null".to_string(), + )); + } + }; + + let call_param = rmcp::model::CallToolRequestParam { + name: actual_tool_name.clone().into(), + arguments: args, + }; + + // 4. 执行工具调用 + let result = service.call_tool(call_param).await.map_err(|e| { + error!( + tool_name = %actual_tool_name, + server_name = %server_name, + error = %e, + "工具调用失败" + ); + McpError::ToolCallFailed(format!("{}", e)) + })?; + + // 5. 转换结果为 McpToolResult + let mcp_result = Self::convert_call_tool_result(result); + + info!( + tool_name = %actual_tool_name, + server_name = %server_name, + is_error = mcp_result.is_error, + "工具调用完成" + ); + + Ok(mcp_result) + } + + /// 解析工具目标(服务器名称和实际工具名) + /// + /// # Arguments + /// + /// * `tool_name` - 工具名称(可能包含服务器前缀,格式为 "server_toolname") + /// + /// # Returns + /// + /// 返回 (服务器名称, 实际工具名) 元组。 + /// + /// # 解析逻辑 + /// + /// 1. 如果工具名包含下划线,尝试解析为 "server_toolname" 格式 + /// 2. 检查解析出的服务器是否存在 + /// 3. 如果服务器存在,使用解析结果 + /// 4. 如果服务器不存在,在所有服务器中查找该工具 + async fn resolve_tool_target(&self, tool_name: &str) -> Result<(String, String), McpError> { + let clients = self.clients.read().await; + + // 尝试解析带前缀的工具名(格式:server_toolname) + if let Some(underscore_pos) = tool_name.find('_') { + let potential_server = &tool_name[..underscore_pos]; + let potential_tool = &tool_name[underscore_pos + 1..]; + + // 检查是否存在该服务器 + if clients.contains_key(potential_server) && !potential_tool.is_empty() { + return Ok((potential_server.to_string(), potential_tool.to_string())); + } + } + + // 没有前缀或前缀不匹配,在所有服务器中查找该工具 + for (server_name, wrapper) in clients.iter() { + if let Some(service) = wrapper.running_service() { + // 尝试获取工具列表并查找 + if let Ok(tools) = service.list_all_tools().await { + if tools.iter().any(|t| t.name.as_ref() == tool_name) { + return Ok((server_name.clone(), tool_name.to_string())); + } + } + } + } + + // 工具未找到 + Err(McpError::ToolNotFound(tool_name.to_string())) + } + + /// 转换 rmcp CallToolResult 为 McpToolResult + fn convert_call_tool_result(result: rmcp::model::CallToolResult) -> McpToolResult { + let content: Vec = result + .content + .into_iter() + .map(|c| Self::convert_content(c)) + .collect(); + + McpToolResult { + content, + is_error: result.is_error.unwrap_or(false), + } + } + + /// 转换 rmcp Content 为 McpContent + fn convert_content(content: rmcp::model::Content) -> McpContent { + // Content 是 Annotated,需要访问内部的 raw 字段 + match content.raw { + rmcp::model::RawContent::Text(text_content) => McpContent::Text { + text: text_content.text, + }, + rmcp::model::RawContent::Image(image_content) => McpContent::Image { + data: image_content.data, + mime_type: image_content.mime_type, + }, + rmcp::model::RawContent::Resource(resource_content) => { + let (uri, text, blob) = match resource_content.resource { + rmcp::model::ResourceContents::TextResourceContents { uri, text, .. } => { + (uri, Some(text), None) + } + rmcp::model::ResourceContents::BlobResourceContents { uri, blob, .. } => { + (uri, None, Some(blob)) + } + }; + McpContent::Resource { uri, text, blob } + } + rmcp::model::RawContent::Audio(audio_content) => { + // 将音频内容作为 Image 类型处理(因为 McpContent 没有 Audio 变体) + McpContent::Image { + data: audio_content.data, + mime_type: audio_content.mime_type, + } + } + rmcp::model::RawContent::ResourceLink(resource_link) => McpContent::Resource { + uri: resource_link.uri.clone(), + text: Some(resource_link.name.clone()), + blob: None, + }, + } + } + + // ======================================================================== + // 提示词管理方法 + // ======================================================================== + + /// 获取所有提示词 + /// + /// 从所有运行中的服务器获取提示词定义。 + /// + /// # Returns + /// + /// 返回所有可用提示词的定义列表。 + /// + /// # 实现步骤(Task 4.4) + /// + /// 1. 遍历所有运行中的服务器 + /// 2. 检查服务器是否支持提示词 + /// 3. 调用 list_all_prompts 获取提示词列表 + /// 4. 转换为 McpPromptDefinition 格式 + /// 5. 返回合并后的提示词列表 + pub async fn list_prompts(&self) -> Result, McpError> { + info!("获取所有 MCP 提示词"); + + let mut all_prompts: Vec = Vec::new(); + let clients = self.clients.read().await; + + for (server_name, wrapper) in clients.iter() { + // 检查服务器是否支持提示词 + if let Some(ref info) = wrapper.server_info { + if !info.supports_prompts { + debug!(server_name = %server_name, "服务器不支持提示词,跳过"); + continue; + } + } + + // 获取 rmcp 服务 + let service = match wrapper.running_service() { + Some(s) => s, + None => { + warn!(server_name = %server_name, "服务器无运行服务,跳过"); + continue; + } + }; + + // 调用 list_all_prompts 获取所有提示词 + match service.list_all_prompts().await { + Ok(prompts) => { + debug!( + server_name = %server_name, + prompt_count = prompts.len(), + "获取服务器提示词列表成功" + ); + for prompt in prompts { + all_prompts.push(Self::convert_prompt_to_definition( + prompt, + server_name.clone(), + )); + } + } + Err(e) => { + warn!( + server_name = %server_name, + error = %e, + "获取服务器提示词列表失败" + ); + // 继续处理其他服务器,不中断 + } + } + } + + info!(prompt_count = all_prompts.len(), "提示词列表已获取"); + Ok(all_prompts) + } + + /// 将 rmcp Prompt 转换为 McpPromptDefinition + fn convert_prompt_to_definition( + prompt: rmcp::model::Prompt, + server_name: String, + ) -> McpPromptDefinition { + let arguments = prompt + .arguments + .unwrap_or_default() + .into_iter() + .map(|arg| McpPromptArgument { + name: arg.name, + description: arg.description, + required: arg.required.unwrap_or(false), + }) + .collect(); + + McpPromptDefinition { + name: prompt.name.to_string(), + description: prompt.description.map(|s| s.to_string()), + arguments, + server_name, + } + } + + /// 获取提示词内容 + /// + /// # Arguments + /// + /// * `name` - 提示词名称(可能包含服务器前缀,格式为 "server_promptname") + /// * `arguments` - 提示词参数 + /// + /// # Returns + /// + /// 返回提示词内容,包含描述和消息列表。 + /// + /// # 实现步骤(Task 4.4) + /// + /// 1. 解析提示词名称,确定目标服务器 + /// 2. 验证必需参数是否提供 + /// 3. 调用服务器的 get_prompt 方法 + /// 4. 转换结果为 McpPromptResult + /// 5. 返回结果 + pub async fn get_prompt( + &self, + name: &str, + arguments: serde_json::Map, + ) -> Result { + info!(prompt_name = %name, "获取 MCP 提示词内容"); + + // 1. 解析提示词名称,确定目标服务器和实际提示词名 + let (server_name, actual_prompt_name) = self.resolve_prompt_target(name).await?; + + debug!( + prompt_name = %name, + server_name = %server_name, + actual_prompt_name = %actual_prompt_name, + "解析提示词目标" + ); + + // 2. 获取目标服务器的客户端 + let clients = self.clients.read().await; + let wrapper = clients + .get(&server_name) + .ok_or_else(|| McpError::ServerNotRunning(server_name.clone()))?; + + let service = wrapper + .running_service() + .ok_or_else(|| McpError::ServerNotRunning(server_name.clone()))?; + + // 3. 构建 get_prompt 请求参数 + let args: Option> = if arguments.is_empty() { + None + } else { + Some(arguments) + }; + + let get_prompt_param = rmcp::model::GetPromptRequestParam { + name: actual_prompt_name.clone().into(), + arguments: args, + }; + + // 4. 调用 get_prompt + let result = service.get_prompt(get_prompt_param).await.map_err(|e| { + error!( + prompt_name = %actual_prompt_name, + server_name = %server_name, + error = %e, + "获取提示词失败" + ); + McpError::ToolCallFailed(format!("获取提示词失败: {}", e)) + })?; + + // 5. 转换结果为 McpPromptResult + let mcp_result = Self::convert_get_prompt_result(result); + + info!( + prompt_name = %actual_prompt_name, + server_name = %server_name, + message_count = mcp_result.messages.len(), + "提示词获取完成" + ); + + Ok(mcp_result) + } + + /// 解析提示词目标(服务器名称和实际提示词名) + /// + /// # Arguments + /// + /// * `prompt_name` - 提示词名称(可能包含服务器前缀,格式为 "server_promptname") + /// + /// # Returns + /// + /// 返回 (服务器名称, 实际提示词名) 元组。 + async fn resolve_prompt_target(&self, prompt_name: &str) -> Result<(String, String), McpError> { + let clients = self.clients.read().await; + + // 尝试解析带前缀的提示词名(格式:server_promptname) + if let Some(underscore_pos) = prompt_name.find('_') { + let potential_server = &prompt_name[..underscore_pos]; + let potential_prompt = &prompt_name[underscore_pos + 1..]; + + // 检查是否存在该服务器 + if clients.contains_key(potential_server) && !potential_prompt.is_empty() { + return Ok((potential_server.to_string(), potential_prompt.to_string())); + } + } + + // 没有前缀或前缀不匹配,在所有服务器中查找该提示词 + for (server_name, wrapper) in clients.iter() { + if let Some(service) = wrapper.running_service() { + // 尝试获取提示词列表并查找 + if let Ok(prompts) = service.list_all_prompts().await { + if prompts.iter().any(|p| p.name.as_str() == prompt_name) { + return Ok((server_name.clone(), prompt_name.to_string())); + } + } + } + } + + // 提示词未找到 + Err(McpError::ToolNotFound(format!( + "提示词不存在: {}", + prompt_name + ))) + } + + /// 转换 rmcp GetPromptResult 为 McpPromptResult + fn convert_get_prompt_result(result: rmcp::model::GetPromptResult) -> McpPromptResult { + let messages: Vec = result + .messages + .into_iter() + .map(|msg| Self::convert_prompt_message(msg)) + .collect(); + + McpPromptResult { + description: result.description.map(|s| s.to_string()), + messages, + } + } + + /// 转换 rmcp PromptMessage 为 McpPromptMessage + fn convert_prompt_message(msg: rmcp::model::PromptMessage) -> McpPromptMessage { + let role = match msg.role { + rmcp::model::PromptMessageRole::User => "user".to_string(), + rmcp::model::PromptMessageRole::Assistant => "assistant".to_string(), + }; + + let content = Self::convert_prompt_message_content(msg.content); + + McpPromptMessage { role, content } + } + + /// 转换 rmcp PromptMessageContent 为 McpContent + fn convert_prompt_message_content(content: rmcp::model::PromptMessageContent) -> McpContent { + match content { + rmcp::model::PromptMessageContent::Text { text } => McpContent::Text { text }, + rmcp::model::PromptMessageContent::Image { image } => McpContent::Image { + data: image.data.clone(), + mime_type: image.mime_type.clone(), + }, + rmcp::model::PromptMessageContent::Resource { resource } => { + let (uri, text, blob) = match &resource.resource { + rmcp::model::ResourceContents::TextResourceContents { uri, text, .. } => { + (uri.clone(), Some(text.clone()), None) + } + rmcp::model::ResourceContents::BlobResourceContents { uri, blob, .. } => { + (uri.clone(), None, Some(blob.clone())) + } + }; + McpContent::Resource { uri, text, blob } + } + rmcp::model::PromptMessageContent::ResourceLink { link } => McpContent::Resource { + uri: link.uri.clone(), + text: Some(link.name.clone()), + blob: None, + }, + } + } + + // ======================================================================== + // 资源管理方法 + // ======================================================================== + + /// 获取所有资源 + /// + /// 从所有运行中的服务器获取资源定义。 + /// + /// # Returns + /// + /// 返回所有可用资源的定义列表。 + /// + /// # 实现步骤(Task 4.5) + /// + /// 1. 遍历所有运行中的服务器 + /// 2. 检查服务器是否支持资源 + /// 3. 调用 list_all_resources 获取资源列表 + /// 4. 转换为 McpResourceDefinition 格式 + /// 5. 返回合并后的资源列表 + pub async fn list_resources(&self) -> Result, McpError> { + info!("获取所有 MCP 资源"); + + let mut all_resources: Vec = Vec::new(); + let clients = self.clients.read().await; + + for (server_name, wrapper) in clients.iter() { + // 检查服务器是否支持资源 + if let Some(ref info) = wrapper.server_info { + if !info.supports_resources { + debug!(server_name = %server_name, "服务器不支持资源,跳过"); + continue; + } + } + + // 获取 rmcp 服务 + let service = match wrapper.running_service() { + Some(s) => s, + None => { + warn!(server_name = %server_name, "服务器无运行服务,跳过"); + continue; + } + }; + + // 调用 list_all_resources 获取所有资源 + match service.list_all_resources().await { + Ok(resources) => { + debug!( + server_name = %server_name, + resource_count = resources.len(), + "获取服务器资源列表成功" + ); + for resource in resources { + all_resources.push(Self::convert_resource_to_definition( + resource, + server_name.clone(), + )); + } + } + Err(e) => { + warn!( + server_name = %server_name, + error = %e, + "获取服务器资源列表失败" + ); + // 继续处理其他服务器,不中断 + } + } + } + + info!(resource_count = all_resources.len(), "资源列表已获取"); + Ok(all_resources) + } + + /// 将 rmcp Resource 转换为 McpResourceDefinition + fn convert_resource_to_definition( + resource: rmcp::model::Resource, + server_name: String, + ) -> McpResourceDefinition { + McpResourceDefinition { + uri: resource.uri.clone(), + name: resource.name.clone(), + description: resource.description.clone(), + mime_type: resource.mime_type.clone(), + server_name, + } + } + + /// 读取资源内容 + /// + /// # Arguments + /// + /// * `uri` - 资源 URI + /// + /// # Returns + /// + /// 返回资源内容。 + /// + /// # 实现步骤(Task 4.5) + /// + /// 1. 解析资源 URI,确定目标服务器 + /// 2. 调用服务器的 read_resource 方法 + /// 3. 转换结果为 McpResourceContent + /// 4. 返回结果 + pub async fn read_resource(&self, uri: &str) -> Result { + info!(uri = %uri, "读取 MCP 资源"); + + // 1. 解析资源 URI,确定目标服务器 + let (server_name, _) = self.resolve_resource_target(uri).await?; + + debug!( + uri = %uri, + server_name = %server_name, + "解析资源目标" + ); + + // 2. 获取目标服务器的客户端 + let clients = self.clients.read().await; + let wrapper = clients + .get(&server_name) + .ok_or_else(|| McpError::ServerNotRunning(server_name.clone()))?; + + let service = wrapper + .running_service() + .ok_or_else(|| McpError::ServerNotRunning(server_name.clone()))?; + + // 3. 构建 read_resource 请求参数 + let read_param = rmcp::model::ReadResourceRequestParam { + uri: uri.to_string(), + }; + + // 4. 调用 read_resource + let result = service.read_resource(read_param).await.map_err(|e| { + error!( + uri = %uri, + server_name = %server_name, + error = %e, + "读取资源失败" + ); + McpError::ToolCallFailed(format!("读取资源失败: {}", e)) + })?; + + // 5. 转换结果为 McpResourceContent + let mcp_result = Self::convert_read_resource_result(uri, result); + + info!( + uri = %uri, + server_name = %server_name, + "资源读取完成" + ); + + Ok(mcp_result) + } + + /// 解析资源目标(服务器名称) + /// + /// # Arguments + /// + /// * `uri` - 资源 URI + /// + /// # Returns + /// + /// 返回 (服务器名称, 资源 URI) 元组。 + /// + /// # 解析逻辑 + /// + /// 遍历所有运行中的服务器,查找提供该资源的服务器。 + async fn resolve_resource_target(&self, uri: &str) -> Result<(String, String), McpError> { + let clients = self.clients.read().await; + + // 在所有服务器中查找该资源 + for (server_name, wrapper) in clients.iter() { + // 检查服务器是否支持资源 + if let Some(ref info) = wrapper.server_info { + if !info.supports_resources { + continue; + } + } + + if let Some(service) = wrapper.running_service() { + // 尝试获取资源列表并查找 + if let Ok(resources) = service.list_all_resources().await { + if resources.iter().any(|r| r.uri == uri) { + return Ok((server_name.clone(), uri.to_string())); + } + } + } + } + + // 资源未找到 + Err(McpError::ToolNotFound(format!("资源不存在: {}", uri))) + } + + /// 转换 rmcp ReadResourceResult 为 McpResourceContent + fn convert_read_resource_result( + uri: &str, + result: rmcp::model::ReadResourceResult, + ) -> McpResourceContent { + // 获取第一个内容(通常只有一个) + if let Some(content) = result.contents.into_iter().next() { + match content { + rmcp::model::ResourceContents::TextResourceContents { + uri: content_uri, + mime_type, + text, + .. + } => McpResourceContent { + uri: content_uri, + mime_type, + text: Some(text), + blob: None, + }, + rmcp::model::ResourceContents::BlobResourceContents { + uri: content_uri, + mime_type, + blob, + .. + } => McpResourceContent { + uri: content_uri, + mime_type, + text: None, + blob: Some(blob), + }, + } + } else { + // 如果没有内容,返回空的资源内容 + McpResourceContent { + uri: uri.to_string(), + mime_type: None, + text: None, + blob: None, + } + } + } +} + +/// Tauri 状态包装器 +pub type McpManagerState = Arc>; + +/// 创建 MCP 管理器状态 +pub fn create_mcp_manager_state(app_handle: Option) -> McpManagerState { + Arc::new(tokio::sync::Mutex::new(McpClientManager::new(app_handle))) +} + +// ============================================================================ +// 单元测试 +// ============================================================================ + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashMap; + + /// 创建测试用的服务器配置 + fn create_test_config() -> McpServerConfig { + McpServerConfig { + command: "test-command".to_string(), + args: vec!["--arg1".to_string(), "--arg2".to_string()], + env: HashMap::new(), + cwd: None, + timeout: 30, + } + } + + /// 创建测试用的客户端包装器 + fn create_test_client(name: &str) -> McpClientWrapper { + McpClientWrapper::new(name.to_string(), create_test_config(), None) + } + + #[test] + fn test_manager_creation() { + let manager = McpClientManager::new(None); + // 验证初始状态 + assert!(manager.app_handle.is_none()); + } + + #[tokio::test] + async fn test_initial_state() { + let manager = McpClientManager::new(None); + + // 验证连接池为空 + assert_eq!(manager.running_server_count().await, 0); + assert!(manager.get_running_servers().await.is_empty()); + + // 验证缓存无效 + assert!(!manager.is_tool_cache_valid().await); + assert!(manager.get_cached_tools().await.is_none()); + } + + #[tokio::test] + async fn test_add_client() { + let manager = McpClientManager::new(None); + let client = create_test_client("test-server"); + + // 添加客户端 + let result = manager.add_client("test-server".to_string(), client).await; + assert!(result.is_ok()); + + // 验证客户端已添加 + assert!(manager.is_server_running("test-server").await); + assert_eq!(manager.running_server_count().await, 1); + } + + #[tokio::test] + async fn test_add_duplicate_client() { + let manager = McpClientManager::new(None); + let client1 = create_test_client("test-server"); + let client2 = create_test_client("test-server"); + + // 添加第一个客户端 + manager + .add_client("test-server".to_string(), client1) + .await + .unwrap(); + + // 尝试添加重复的客户端 + let result = manager.add_client("test-server".to_string(), client2).await; + assert!(result.is_err()); + + // 验证错误类型 + match result { + Err(McpError::ServerAlreadyRunning(name)) => { + assert_eq!(name, "test-server"); + } + _ => panic!("Expected ServerAlreadyRunning error"), + } + } + + #[tokio::test] + async fn test_remove_client() { + let manager = McpClientManager::new(None); + let client = create_test_client("test-server"); + + // 添加客户端 + manager + .add_client("test-server".to_string(), client) + .await + .unwrap(); + + // 移除客户端 + let removed = manager.remove_client("test-server").await; + assert!(removed.is_some()); + + // 验证客户端已移除 + assert!(!manager.is_server_running("test-server").await); + assert_eq!(manager.running_server_count().await, 0); + } + + #[tokio::test] + async fn test_remove_nonexistent_client() { + let manager = McpClientManager::new(None); + + // 尝试移除不存在的客户端 + let removed = manager.remove_client("nonexistent").await; + assert!(removed.is_none()); + } + + #[tokio::test] + async fn test_has_client_and_get_config() { + let manager = McpClientManager::new(None); + let client = create_test_client("test-server"); + + // 添加客户端 + manager + .add_client("test-server".to_string(), client) + .await + .unwrap(); + + // 检查客户端是否存在 + assert!(manager.has_client("test-server").await); + assert!(!manager.has_client("nonexistent").await); + + // 获取客户端配置 + let config = manager.get_client_config("test-server").await; + assert!(config.is_some()); + assert_eq!(config.unwrap().command, "test-command"); + + // 获取不存在的客户端配置 + let nonexistent_config = manager.get_client_config("nonexistent").await; + assert!(nonexistent_config.is_none()); + } + + #[tokio::test] + async fn test_get_running_servers() { + let manager = McpClientManager::new(None); + + // 添加多个客户端 + manager + .add_client("server1".to_string(), create_test_client("server1")) + .await + .unwrap(); + manager + .add_client("server2".to_string(), create_test_client("server2")) + .await + .unwrap(); + manager + .add_client("server3".to_string(), create_test_client("server3")) + .await + .unwrap(); + + // 获取运行中的服务器列表 + let servers = manager.get_running_servers().await; + assert_eq!(servers.len(), 3); + assert!(servers.contains(&"server1".to_string())); + assert!(servers.contains(&"server2".to_string())); + assert!(servers.contains(&"server3".to_string())); + } + + #[tokio::test] + async fn test_tool_cache_operations() { + let manager = McpClientManager::new(None); + + // 初始状态:缓存无效 + assert!(!manager.is_tool_cache_valid().await); + assert!(manager.get_cached_tools().await.is_none()); + + // 更新缓存 + let tools = vec![ + McpToolDefinition { + name: "tool1".to_string(), + description: "Test tool 1".to_string(), + input_schema: serde_json::json!({}), + server_name: "server1".to_string(), + }, + McpToolDefinition { + name: "tool2".to_string(), + description: "Test tool 2".to_string(), + input_schema: serde_json::json!({}), + server_name: "server1".to_string(), + }, + ]; + manager.update_tool_cache(tools.clone()).await; + + // 验证缓存有效 + assert!(manager.is_tool_cache_valid().await); + let cached = manager.get_cached_tools().await; + assert!(cached.is_some()); + assert_eq!(cached.unwrap().len(), 2); + + // 失效缓存 + manager.invalidate_tool_cache().await; + + // 验证缓存已失效 + assert!(!manager.is_tool_cache_valid().await); + assert!(manager.get_cached_tools().await.is_none()); + } + + #[tokio::test] + async fn test_is_server_running() { + let manager = McpClientManager::new(None); + + // 初始状态:没有服务器运行 + assert!(!manager.is_server_running("test-server").await); + + // 添加客户端 + manager + .add_client("test-server".to_string(), create_test_client("test-server")) + .await + .unwrap(); + + // 验证服务器正在运行 + assert!(manager.is_server_running("test-server").await); + + // 移除客户端 + manager.remove_client("test-server").await; + + // 验证服务器不再运行 + assert!(!manager.is_server_running("test-server").await); + } + + #[test] + fn test_create_mcp_manager_state() { + let state = create_mcp_manager_state(None); + // 验证状态已创建 + assert!(Arc::strong_count(&state) >= 1); + } + + // ======================================================================== + // 服务器生命周期测试 + // ======================================================================== + + #[tokio::test] + async fn test_start_server_already_running() { + let manager = McpClientManager::new(None); + + // 先添加一个客户端模拟已运行的服务器 + let client = create_test_client("test-server"); + manager + .add_client("test-server".to_string(), client) + .await + .unwrap(); + + // 尝试启动已运行的服务器 + let config = create_test_config(); + let result = manager.start_server("test-server", &config).await; + + // 应该返回 ServerAlreadyRunning 错误 + assert!(result.is_err()); + match result { + Err(McpError::ServerAlreadyRunning(name)) => { + assert_eq!(name, "test-server"); + } + _ => panic!("Expected ServerAlreadyRunning error"), + } + } + + #[tokio::test] + async fn test_start_server_invalid_command() { + let manager = McpClientManager::new(None); + + // 使用不存在的命令 + let config = McpServerConfig { + command: "/nonexistent/command/that/does/not/exist".to_string(), + args: vec![], + env: HashMap::new(), + cwd: None, + timeout: 5, + }; + + let result = manager.start_server("test-server", &config).await; + + // 应该返回 ProcessSpawnFailed 错误 + assert!(result.is_err()); + match result { + Err(McpError::ProcessSpawnFailed(_)) => {} + Err(e) => panic!("Expected ProcessSpawnFailed error, got: {:?}", e), + Ok(_) => panic!("Expected error, but got Ok"), + } + } + + #[tokio::test] + async fn test_stop_server_not_running() { + let manager = McpClientManager::new(None); + + // 停止未运行的服务器(幂等操作,应该成功) + let result = manager.stop_server("nonexistent-server").await; + assert!(result.is_ok()); + } + + #[tokio::test] + async fn test_stop_server_removes_from_pool() { + let manager = McpClientManager::new(None); + + // 添加一个客户端 + let client = create_test_client("test-server"); + manager + .add_client("test-server".to_string(), client) + .await + .unwrap(); + + // 验证服务器在运行 + assert!(manager.is_server_running("test-server").await); + + // 停止服务器 + let result = manager.stop_server("test-server").await; + assert!(result.is_ok()); + + // 验证服务器已停止 + assert!(!manager.is_server_running("test-server").await); + } + + #[tokio::test] + async fn test_stop_server_invalidates_cache() { + let manager = McpClientManager::new(None); + + // 添加一个客户端 + let client = create_test_client("test-server"); + manager + .add_client("test-server".to_string(), client) + .await + .unwrap(); + + // 设置工具缓存 + let tools = vec![McpToolDefinition { + name: "tool1".to_string(), + description: "Test tool".to_string(), + input_schema: serde_json::json!({}), + server_name: "test-server".to_string(), + }]; + manager.update_tool_cache(tools).await; + assert!(manager.is_tool_cache_valid().await); + + // 停止服务器 + manager.stop_server("test-server").await.unwrap(); + + // 验证缓存已失效 + assert!(!manager.is_tool_cache_valid().await); + } + + #[tokio::test] + async fn test_restart_server_stops_then_starts() { + let manager = McpClientManager::new(None); + + // 添加一个客户端模拟已运行的服务器 + let client = create_test_client("test-server"); + manager + .add_client("test-server".to_string(), client) + .await + .unwrap(); + + // 使用无效命令重启(会失败在启动阶段) + let config = McpServerConfig { + command: "/nonexistent/command".to_string(), + args: vec![], + env: HashMap::new(), + cwd: None, + timeout: 5, + }; + + // 重启应该先停止成功,然后启动失败 + let result = manager.restart_server("test-server", &config).await; + assert!(result.is_err()); + + // 验证服务器已被停止(即使启动失败) + assert!(!manager.is_server_running("test-server").await); + } + + // ======================================================================== + // 工具名称冲突解决测试(Task 4.3) + // ======================================================================== + + #[test] + fn test_resolve_tool_name_conflicts_no_conflict() { + // 没有冲突的情况 + let tools = vec![ + McpToolDefinition { + name: "tool1".to_string(), + description: "Tool 1".to_string(), + input_schema: serde_json::json!({}), + server_name: "server1".to_string(), + }, + McpToolDefinition { + name: "tool2".to_string(), + description: "Tool 2".to_string(), + input_schema: serde_json::json!({}), + server_name: "server2".to_string(), + }, + ]; + + let resolved = McpClientManager::resolve_tool_name_conflicts(tools); + + // 名称应该保持不变 + assert_eq!(resolved.len(), 2); + assert!(resolved.iter().any(|t| t.name == "tool1")); + assert!(resolved.iter().any(|t| t.name == "tool2")); + } + + #[test] + fn test_resolve_tool_name_conflicts_with_conflict() { + // 有冲突的情况:两个服务器都提供 "read_file" 工具 + let tools = vec![ + McpToolDefinition { + name: "read_file".to_string(), + description: "Read file from server1".to_string(), + input_schema: serde_json::json!({}), + server_name: "server1".to_string(), + }, + McpToolDefinition { + name: "read_file".to_string(), + description: "Read file from server2".to_string(), + input_schema: serde_json::json!({}), + server_name: "server2".to_string(), + }, + McpToolDefinition { + name: "unique_tool".to_string(), + description: "Unique tool".to_string(), + input_schema: serde_json::json!({}), + server_name: "server1".to_string(), + }, + ]; + + let resolved = McpClientManager::resolve_tool_name_conflicts(tools); + + // 冲突的工具应该添加服务器前缀 + assert_eq!(resolved.len(), 3); + assert!(resolved.iter().any(|t| t.name == "server1_read_file")); + assert!(resolved.iter().any(|t| t.name == "server2_read_file")); + // 唯一的工具名称应该保持不变 + assert!(resolved.iter().any(|t| t.name == "unique_tool")); + } + + #[test] + fn test_resolve_tool_name_conflicts_multiple_conflicts() { + // 多个冲突的情况 + let tools = vec![ + McpToolDefinition { + name: "tool_a".to_string(), + description: "Tool A from server1".to_string(), + input_schema: serde_json::json!({}), + server_name: "server1".to_string(), + }, + McpToolDefinition { + name: "tool_a".to_string(), + description: "Tool A from server2".to_string(), + input_schema: serde_json::json!({}), + server_name: "server2".to_string(), + }, + McpToolDefinition { + name: "tool_a".to_string(), + description: "Tool A from server3".to_string(), + input_schema: serde_json::json!({}), + server_name: "server3".to_string(), + }, + ]; + + let resolved = McpClientManager::resolve_tool_name_conflicts(tools); + + // 所有冲突的工具都应该添加服务器前缀 + assert_eq!(resolved.len(), 3); + assert!(resolved.iter().any(|t| t.name == "server1_tool_a")); + assert!(resolved.iter().any(|t| t.name == "server2_tool_a")); + assert!(resolved.iter().any(|t| t.name == "server3_tool_a")); + } + + #[test] + fn test_resolve_tool_name_conflicts_empty_list() { + // 空列表的情况 + let tools: Vec = vec![]; + let resolved = McpClientManager::resolve_tool_name_conflicts(tools); + assert!(resolved.is_empty()); + } + + // ======================================================================== + // 工具列表缓存测试(Task 4.3) + // ======================================================================== + + #[tokio::test] + async fn test_list_tools_returns_cached_when_valid() { + let manager = McpClientManager::new(None); + + // 预先设置缓存 + let cached_tools = vec![McpToolDefinition { + name: "cached_tool".to_string(), + description: "Cached tool".to_string(), + input_schema: serde_json::json!({}), + server_name: "cached_server".to_string(), + }]; + manager.update_tool_cache(cached_tools.clone()).await; + + // 调用 list_tools 应该返回缓存的工具 + let result = manager.list_tools().await.unwrap(); + assert_eq!(result.len(), 1); + assert_eq!(result[0].name, "cached_tool"); + } + + #[tokio::test] + async fn test_list_tools_returns_empty_when_no_servers() { + let manager = McpClientManager::new(None); + + // 没有运行的服务器时,应该返回空列表 + let result = manager.list_tools().await.unwrap(); + assert!(result.is_empty()); + } + + // ======================================================================== + // 工具调用测试(Task 4.3) + // ======================================================================== + + #[tokio::test] + async fn test_call_tool_not_found() { + let manager = McpClientManager::new(None); + + // 调用不存在的工具 + let result = manager + .call_tool("nonexistent_tool", serde_json::json!({})) + .await; + + // 应该返回 ToolNotFound 错误 + assert!(result.is_err()); + match result { + Err(McpError::ToolNotFound(name)) => { + assert_eq!(name, "nonexistent_tool"); + } + _ => panic!("Expected ToolNotFound error"), + } + } + + #[tokio::test] + async fn test_call_tool_invalid_arguments() { + let manager = McpClientManager::new(None); + + // 添加一个客户端 + let client = create_test_client("test-server"); + manager + .add_client("test-server".to_string(), client) + .await + .unwrap(); + + // 使用非对象参数调用工具 + let result = manager + .call_tool("test-server_some_tool", serde_json::json!("invalid")) + .await; + + // 应该返回错误(参数必须是对象或 null) + assert!(result.is_err()); + } + + // ======================================================================== + // 内容转换测试(Task 4.3) + // ======================================================================== + + #[test] + fn test_convert_content_text() { + let content = rmcp::model::Content::text("Hello, World!"); + let mcp_content = McpClientManager::convert_content(content); + + match mcp_content { + McpContent::Text { text } => { + assert_eq!(text, "Hello, World!"); + } + _ => panic!("Expected Text content"), + } + } + + #[test] + fn test_convert_content_image() { + let content = rmcp::model::Content::image("base64data", "image/png"); + let mcp_content = McpClientManager::convert_content(content); + + match mcp_content { + McpContent::Image { data, mime_type } => { + assert_eq!(data, "base64data"); + assert_eq!(mime_type, "image/png"); + } + _ => panic!("Expected Image content"), + } + } + + // ======================================================================== + // 提示词管理测试(Task 4.4) + // ======================================================================== + + #[tokio::test] + async fn test_list_prompts_returns_empty_when_no_servers() { + let manager = McpClientManager::new(None); + + // 没有运行的服务器时,应该返回空列表 + let result = manager.list_prompts().await.unwrap(); + assert!(result.is_empty()); + } + + #[tokio::test] + async fn test_get_prompt_not_found() { + let manager = McpClientManager::new(None); + + // 获取不存在的提示词 + let result = manager + .get_prompt("nonexistent_prompt", serde_json::Map::new()) + .await; + + // 应该返回错误 + assert!(result.is_err()); + match result { + Err(McpError::ToolNotFound(msg)) => { + assert!(msg.contains("nonexistent_prompt")); + } + _ => panic!("Expected ToolNotFound error"), + } + } + + #[test] + fn test_convert_prompt_to_definition() { + // 创建一个 rmcp Prompt + let prompt = rmcp::model::Prompt { + name: "test_prompt".into(), + title: Some("Test Prompt Title".into()), + description: Some("A test prompt description".into()), + arguments: Some(vec![ + rmcp::model::PromptArgument { + name: "arg1".to_string(), + title: None, + description: Some("First argument".to_string()), + required: Some(true), + }, + rmcp::model::PromptArgument { + name: "arg2".to_string(), + title: None, + description: Some("Second argument".to_string()), + required: Some(false), + }, + ]), + icons: None, + }; + + // 转换为 McpPromptDefinition + let definition = + McpClientManager::convert_prompt_to_definition(prompt, "test_server".to_string()); + + // 验证转换结果 + assert_eq!(definition.name, "test_prompt"); + assert_eq!( + definition.description, + Some("A test prompt description".to_string()) + ); + assert_eq!(definition.server_name, "test_server"); + assert_eq!(definition.arguments.len(), 2); + + // 验证第一个参数 + assert_eq!(definition.arguments[0].name, "arg1"); + assert_eq!( + definition.arguments[0].description, + Some("First argument".to_string()) + ); + assert!(definition.arguments[0].required); + + // 验证第二个参数 + assert_eq!(definition.arguments[1].name, "arg2"); + assert_eq!( + definition.arguments[1].description, + Some("Second argument".to_string()) + ); + assert!(!definition.arguments[1].required); + } + + #[test] + fn test_convert_prompt_to_definition_no_arguments() { + // 创建一个没有参数的 rmcp Prompt + let prompt = rmcp::model::Prompt { + name: "simple_prompt".into(), + title: None, + description: None, + arguments: None, + icons: None, + }; + + // 转换为 McpPromptDefinition + let definition = + McpClientManager::convert_prompt_to_definition(prompt, "server1".to_string()); + + // 验证转换结果 + assert_eq!(definition.name, "simple_prompt"); + assert!(definition.description.is_none()); + assert_eq!(definition.server_name, "server1"); + assert!(definition.arguments.is_empty()); + } + + #[test] + fn test_convert_prompt_message_user() { + // 创建一个用户消息 + let msg = rmcp::model::PromptMessage::new_text( + rmcp::model::PromptMessageRole::User, + "Hello, assistant!", + ); + + // 转换为 McpPromptMessage + let mcp_msg = McpClientManager::convert_prompt_message(msg); + + // 验证转换结果 + assert_eq!(mcp_msg.role, "user"); + match mcp_msg.content { + McpContent::Text { text } => { + assert_eq!(text, "Hello, assistant!"); + } + _ => panic!("Expected Text content"), + } + } + + #[test] + fn test_convert_prompt_message_assistant() { + // 创建一个助手消息 + let msg = rmcp::model::PromptMessage::new_text( + rmcp::model::PromptMessageRole::Assistant, + "Hello, user!", + ); + + // 转换为 McpPromptMessage + let mcp_msg = McpClientManager::convert_prompt_message(msg); + + // 验证转换结果 + assert_eq!(mcp_msg.role, "assistant"); + match mcp_msg.content { + McpContent::Text { text } => { + assert_eq!(text, "Hello, user!"); + } + _ => panic!("Expected Text content"), + } + } + + #[test] + fn test_convert_prompt_message_content_text() { + // 创建文本内容 + let content = rmcp::model::PromptMessageContent::Text { + text: "Test text content".to_string(), + }; + + // 转换为 McpContent + let mcp_content = McpClientManager::convert_prompt_message_content(content); + + // 验证转换结果 + match mcp_content { + McpContent::Text { text } => { + assert_eq!(text, "Test text content"); + } + _ => panic!("Expected Text content"), + } + } + + #[test] + fn test_convert_get_prompt_result() { + // 创建 GetPromptResult + let result = rmcp::model::GetPromptResult { + description: Some("Test prompt result".into()), + messages: vec![ + rmcp::model::PromptMessage::new_text( + rmcp::model::PromptMessageRole::User, + "User message", + ), + rmcp::model::PromptMessage::new_text( + rmcp::model::PromptMessageRole::Assistant, + "Assistant response", + ), + ], + }; + + // 转换为 McpPromptResult + let mcp_result = McpClientManager::convert_get_prompt_result(result); + + // 验证转换结果 + assert_eq!( + mcp_result.description, + Some("Test prompt result".to_string()) + ); + assert_eq!(mcp_result.messages.len(), 2); + assert_eq!(mcp_result.messages[0].role, "user"); + assert_eq!(mcp_result.messages[1].role, "assistant"); + } + + // ======================================================================== + // 资源管理测试(Task 4.5) + // ======================================================================== + + #[tokio::test] + async fn test_list_resources_returns_empty_when_no_servers() { + let manager = McpClientManager::new(None); + + // 没有运行的服务器时,应该返回空列表 + let result = manager.list_resources().await.unwrap(); + assert!(result.is_empty()); + } + + #[tokio::test] + async fn test_read_resource_not_found() { + let manager = McpClientManager::new(None); + + // 读取不存在的资源 + let result = manager.read_resource("file:///nonexistent/resource").await; + + // 应该返回错误 + assert!(result.is_err()); + match result { + Err(McpError::ToolNotFound(msg)) => { + assert!(msg.contains("资源不存在")); + } + _ => panic!("Expected ToolNotFound error"), + } + } + + #[test] + fn test_convert_resource_to_definition() { + use rmcp::model::{AnnotateAble, RawResource}; + + // 创建一个 rmcp Resource + let raw_resource = RawResource { + uri: "file:///test/resource.txt".to_string(), + name: "resource.txt".to_string(), + title: Some("Test Resource".to_string()), + description: Some("A test resource".to_string()), + mime_type: Some("text/plain".to_string()), + size: Some(1024), + icons: None, + }; + let resource = raw_resource.no_annotation(); + + // 转换为 McpResourceDefinition + let definition = + McpClientManager::convert_resource_to_definition(resource, "test_server".to_string()); + + // 验证转换结果 + assert_eq!(definition.uri, "file:///test/resource.txt"); + assert_eq!(definition.name, "resource.txt"); + assert_eq!(definition.description, Some("A test resource".to_string())); + assert_eq!(definition.mime_type, Some("text/plain".to_string())); + assert_eq!(definition.server_name, "test_server"); + } + + #[test] + fn test_convert_resource_to_definition_minimal() { + use rmcp::model::{AnnotateAble, RawResource}; + + // 创建一个最小的 rmcp Resource(只有必需字段) + let raw_resource = + RawResource::new("file:///minimal.txt".to_string(), "minimal.txt".to_string()); + let resource = raw_resource.no_annotation(); + + // 转换为 McpResourceDefinition + let definition = + McpClientManager::convert_resource_to_definition(resource, "server1".to_string()); + + // 验证转换结果 + assert_eq!(definition.uri, "file:///minimal.txt"); + assert_eq!(definition.name, "minimal.txt"); + assert!(definition.description.is_none()); + assert!(definition.mime_type.is_none()); + assert_eq!(definition.server_name, "server1"); + } + + #[test] + fn test_convert_read_resource_result_text() { + // 创建文本资源内容 + let result = rmcp::model::ReadResourceResult { + contents: vec![rmcp::model::ResourceContents::text( + "Hello, World!", + "file:///test.txt", + )], + }; + + // 转换为 McpResourceContent + let mcp_content = + McpClientManager::convert_read_resource_result("file:///test.txt", result); + + // 验证转换结果 + assert_eq!(mcp_content.uri, "file:///test.txt"); + assert_eq!(mcp_content.mime_type, Some("text".to_string())); + assert_eq!(mcp_content.text, Some("Hello, World!".to_string())); + assert!(mcp_content.blob.is_none()); + } + + #[test] + fn test_convert_read_resource_result_blob() { + // 创建二进制资源内容 + let result = rmcp::model::ReadResourceResult { + contents: vec![rmcp::model::ResourceContents::BlobResourceContents { + uri: "file:///test.bin".to_string(), + mime_type: Some("application/octet-stream".to_string()), + blob: "base64encodeddata".to_string(), + meta: None, + }], + }; + + // 转换为 McpResourceContent + let mcp_content = + McpClientManager::convert_read_resource_result("file:///test.bin", result); + + // 验证转换结果 + assert_eq!(mcp_content.uri, "file:///test.bin"); + assert_eq!( + mcp_content.mime_type, + Some("application/octet-stream".to_string()) + ); + assert!(mcp_content.text.is_none()); + assert_eq!(mcp_content.blob, Some("base64encodeddata".to_string())); + } + + #[test] + fn test_convert_read_resource_result_empty() { + // 创建空的资源结果 + let result = rmcp::model::ReadResourceResult { contents: vec![] }; + + // 转换为 McpResourceContent + let mcp_content = + McpClientManager::convert_read_resource_result("file:///empty.txt", result); + + // 验证转换结果(应该返回空内容) + assert_eq!(mcp_content.uri, "file:///empty.txt"); + assert!(mcp_content.mime_type.is_none()); + assert!(mcp_content.text.is_none()); + assert!(mcp_content.blob.is_none()); + } +} diff --git a/src-tauri/src/mcp/mod.rs b/src-tauri/src/mcp/mod.rs new file mode 100644 index 000000000..7dcb75bb6 --- /dev/null +++ b/src-tauri/src/mcp/mod.rs @@ -0,0 +1,31 @@ +//! MCP(Model Context Protocol)模块 +//! +//! 本模块提供 MCP 协议的客户端实现,支持: +//! - MCP 服务器生命周期管理(启动、停止、状态监控) +//! - MCP 工具发现和调用 +//! - MCP 提示词和资源访问 +//! - 工具格式转换(OpenAI/Anthropic/Gemini) +//! +//! # 模块结构 +//! +//! - `types`: MCP 数据类型定义 +//! - `client`: MCP 客户端实现(rmcp ClientHandler) +//! - `manager`: MCP 客户端管理器(连接池、缓存) +//! - `tool_converter`: 工具格式转换器 + +pub mod client; +pub mod manager; +pub mod tool_converter; +pub mod types; + +// 显式导出,避免命名冲突 +pub use client::{McpClientWrapper, ProxyCastMcpClient}; +pub use manager::McpClientManager; +pub use tool_converter::ToolConverter; +pub use types::{ + McpContent, McpError, McpManagerState, McpPromptArgument, McpPromptDefinition, + McpPromptMessage, McpPromptResult, McpResourceContent, McpResourceDefinition, + McpServerCapabilities, McpServerConfig, McpServerErrorPayload, McpServerInfo, + McpServerStartedPayload, McpServerStoppedPayload, McpToolCall, McpToolDefinition, + McpToolResult, McpToolsUpdatedPayload, +}; diff --git a/src-tauri/src/mcp/tool_converter.rs b/src-tauri/src/mcp/tool_converter.rs new file mode 100644 index 000000000..1b5dabf6a --- /dev/null +++ b/src-tauri/src/mcp/tool_converter.rs @@ -0,0 +1,180 @@ +//! MCP 工具格式转换器 +//! +//! 本模块提供 MCP 工具定义与各 LLM Provider 格式之间的转换: +//! - OpenAI function calling 格式 +//! - Anthropic tool use 格式 +//! - Gemini function declaration 格式 + +#![allow(dead_code)] + +use serde::{Deserialize, Serialize}; + +use super::types::{McpToolCall, McpToolDefinition}; + +// ============================================================================ +// OpenAI 格式 +// ============================================================================ + +/// OpenAI 工具格式 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAITool { + #[serde(rename = "type")] + pub tool_type: String, + pub function: OpenAIFunction, +} + +/// OpenAI 函数定义 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIFunction { + pub name: String, + pub description: String, + pub parameters: serde_json::Value, +} + +/// OpenAI 工具调用 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIToolCall { + pub id: String, + #[serde(rename = "type")] + pub call_type: String, + pub function: OpenAIFunctionCall, +} + +/// OpenAI 函数调用 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIFunctionCall { + pub name: String, + pub arguments: String, +} + +// ============================================================================ +// Anthropic 格式 +// ============================================================================ + +/// Anthropic 工具格式 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AnthropicTool { + pub name: String, + pub description: String, + pub input_schema: serde_json::Value, +} + +/// Anthropic 工具使用 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AnthropicToolUse { + pub id: String, + pub name: String, + pub input: serde_json::Value, +} + +// ============================================================================ +// Gemini 格式 +// ============================================================================ + +/// Gemini 函数声明格式 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct GeminiFunctionDeclaration { + pub name: String, + pub description: String, + pub parameters: GeminiParameters, +} + +/// Gemini 参数定义 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct GeminiParameters { + #[serde(rename = "type")] + pub param_type: String, + pub properties: serde_json::Value, + pub required: Vec, +} + +// ============================================================================ +// 转换器实现 +// ============================================================================ + +/// MCP 工具格式转换器 +pub struct ToolConverter; + +impl ToolConverter { + /// 转换为 OpenAI function calling 格式 + pub fn to_openai(tools: &[McpToolDefinition]) -> Vec { + tools + .iter() + .map(|tool| OpenAITool { + tool_type: "function".to_string(), + function: OpenAIFunction { + name: tool.name.clone(), + description: tool.description.clone(), + parameters: tool.input_schema.clone(), + }, + }) + .collect() + } + + /// 转换为 Anthropic tool use 格式 + pub fn to_anthropic(tools: &[McpToolDefinition]) -> Vec { + tools + .iter() + .map(|tool| AnthropicTool { + name: tool.name.clone(), + description: tool.description.clone(), + input_schema: tool.input_schema.clone(), + }) + .collect() + } + + /// 转换为 Gemini function declaration 格式 + pub fn to_gemini(tools: &[McpToolDefinition]) -> Vec { + tools + .iter() + .map(|tool| { + // 从 input_schema 提取 properties 和 required + let properties = tool + .input_schema + .get("properties") + .cloned() + .unwrap_or(serde_json::json!({})); + + let required = tool + .input_schema + .get("required") + .and_then(|v| v.as_array()) + .map(|arr| { + arr.iter() + .filter_map(|v| v.as_str().map(String::from)) + .collect() + }) + .unwrap_or_default(); + + GeminiFunctionDeclaration { + name: tool.name.clone(), + description: tool.description.clone(), + parameters: GeminiParameters { + param_type: "object".to_string(), + properties, + required, + }, + } + }) + .collect() + } + + /// 从 OpenAI tool call 转换回 MCP 格式 + pub fn from_openai_call(call: &OpenAIToolCall) -> McpToolCall { + let arguments = + serde_json::from_str(&call.function.arguments).unwrap_or(serde_json::json!({})); + + McpToolCall { + name: call.function.name.clone(), + arguments, + } + } + + /// 从 Anthropic tool use 转换回 MCP 格式 + pub fn from_anthropic_use(use_: &AnthropicToolUse) -> McpToolCall { + McpToolCall { + name: use_.name.clone(), + arguments: use_.input.clone(), + } + } +} diff --git a/src-tauri/src/mcp/types.rs b/src-tauri/src/mcp/types.rs new file mode 100644 index 000000000..8ce6630de --- /dev/null +++ b/src-tauri/src/mcp/types.rs @@ -0,0 +1,244 @@ +//! MCP 类型定义 +//! +//! 本模块定义 MCP 协议相关的数据类型,包括: +//! - 服务器配置和状态 +//! - 工具定义、调用和结果 +//! - 提示词定义和结果 +//! - 资源定义和内容 +//! - 错误类型 +//! - Tauri 事件 Payload + +use serde::{Deserialize, Serialize}; +use std::collections::HashMap; + +// ============================================================================ +// 服务器配置和状态 +// ============================================================================ + +/// MCP 服务器配置 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpServerConfig { + /// 启动命令 + pub command: String, + /// 命令参数 + #[serde(default)] + pub args: Vec, + /// 环境变量 + #[serde(default)] + pub env: HashMap, + /// 工作目录 + pub cwd: Option, + /// 超时时间(秒) + #[serde(default = "default_timeout")] + pub timeout: u64, +} + +fn default_timeout() -> u64 { + 30 +} + +/// MCP 服务器信息(包含运行状态) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpServerInfo { + pub id: String, + pub name: String, + pub description: Option, + pub config: McpServerConfig, + pub is_running: bool, + pub server_info: Option, + pub enabled_proxycast: bool, + pub enabled_claude: bool, + pub enabled_codex: bool, + pub enabled_gemini: bool, +} + +/// MCP 服务器能力 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpServerCapabilities { + pub name: String, + pub version: String, + pub supports_tools: bool, + pub supports_prompts: bool, + pub supports_resources: bool, +} + +// ============================================================================ +// 工具类型 +// ============================================================================ + +/// MCP 工具定义 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpToolDefinition { + pub name: String, + pub description: String, + pub input_schema: serde_json::Value, + pub server_name: String, +} + +/// MCP 工具调用请求 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpToolCall { + pub name: String, + pub arguments: serde_json::Value, +} + +/// MCP 工具调用结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpToolResult { + pub content: Vec, + pub is_error: bool, +} + +/// MCP 内容类型 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type")] +pub enum McpContent { + #[serde(rename = "text")] + Text { text: String }, + #[serde(rename = "image")] + Image { data: String, mime_type: String }, + #[serde(rename = "resource")] + Resource { + uri: String, + text: Option, + blob: Option, + }, +} + +// ============================================================================ +// 提示词类型 +// ============================================================================ + +/// MCP 提示词定义 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpPromptDefinition { + pub name: String, + pub description: Option, + pub arguments: Vec, + pub server_name: String, +} + +/// MCP 提示词参数 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpPromptArgument { + pub name: String, + pub description: Option, + pub required: bool, +} + +/// MCP 提示词结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpPromptResult { + pub description: Option, + pub messages: Vec, +} + +/// MCP 提示词消息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpPromptMessage { + pub role: String, + pub content: McpContent, +} + +// ============================================================================ +// 资源类型 +// ============================================================================ + +/// MCP 资源定义 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpResourceDefinition { + pub uri: String, + pub name: String, + pub description: Option, + pub mime_type: Option, + pub server_name: String, +} + +/// MCP 资源内容 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpResourceContent { + pub uri: String, + pub mime_type: Option, + pub text: Option, + pub blob: Option, +} + +// ============================================================================ +// 错误类型 +// ============================================================================ + +/// MCP 错误类型 +#[derive(Debug, thiserror::Error)] +pub enum McpError { + #[error("服务器配置不存在: {0}")] + ConfigNotFound(String), + + #[error("服务器已在运行: {0}")] + ServerAlreadyRunning(String), + + #[error("服务器未运行: {0}")] + ServerNotRunning(String), + + #[error("无法启动服务器进程: {0}")] + ProcessSpawnFailed(String), + + #[error("MCP 连接失败: {0}")] + ConnectionFailed(String), + + #[error("工具不存在: {0}")] + ToolNotFound(String), + + #[error("工具调用失败: {0}")] + ToolCallFailed(String), + + #[error("操作超时")] + Timeout, + + #[error("数据库错误: {0}")] + DatabaseError(String), + + #[error("协议错误: {0}")] + ProtocolError(String), +} + +// ============================================================================ +// Tauri 事件 Payload +// ============================================================================ + +/// 服务器启动事件 +#[derive(Debug, Clone, Serialize)] +pub struct McpServerStartedPayload { + pub server_name: String, + pub server_info: Option, +} + +/// 服务器停止事件 +#[derive(Debug, Clone, Serialize)] +pub struct McpServerStoppedPayload { + pub server_name: String, +} + +/// 服务器错误事件 +#[derive(Debug, Clone, Serialize)] +pub struct McpServerErrorPayload { + pub server_name: String, + pub error: String, +} + +/// 工具列表更新事件 +#[derive(Debug, Clone, Serialize)] +pub struct McpToolsUpdatedPayload { + pub tools: Vec, +} + +// ============================================================================ +// Tauri 状态类型 +// ============================================================================ + +use std::sync::Arc; +use tokio::sync::Mutex; + +/// MCP 客户端管理器状态(Tauri 托管状态) +/// +/// 使用 Arc> 包装,支持跨线程共享和异步访问。 +pub type McpManagerState = Arc>; diff --git a/src-tauri/src/models/mcp_model.rs b/src-tauri/src/models/mcp_model.rs index 2a4044b5c..fe257cc13 100644 --- a/src-tauri/src/models/mcp_model.rs +++ b/src-tauri/src/models/mcp_model.rs @@ -1,5 +1,48 @@ use serde::{Deserialize, Serialize}; use serde_json::Value; +use std::collections::HashMap; + +/// MCP 服务器配置(类型化) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpServerConfigTyped { + /// 启动命令 + pub command: String, + /// 命令参数 + #[serde(default)] + pub args: Vec, + /// 环境变量 + #[serde(default)] + pub env: HashMap, + /// 工作目录 + #[serde(skip_serializing_if = "Option::is_none")] + pub cwd: Option, + /// 超时时间(秒) + #[serde(default = "default_timeout")] + pub timeout: u64, +} + +fn default_timeout() -> u64 { + 30 +} + +impl Default for McpServerConfigTyped { + fn default() -> Self { + Self { + command: String::new(), + args: Vec::new(), + env: HashMap::new(), + cwd: None, + timeout: 30, + } + } +} + +/// 配置验证错误 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ConfigValidationError { + pub field: String, + pub message: String, +} #[derive(Debug, Clone, Serialize, Deserialize)] pub struct McpServer { @@ -35,4 +78,104 @@ impl McpServer { created_at: Some(chrono::Utc::now().timestamp()), } } + + /// 解析 server_config 为类型化配置 + /// + /// 将 JSON Value 解析为 McpServerConfigTyped 结构。 + /// 如果解析失败,返回默认配置并尝试提取基本字段。 + pub fn parse_config(&self) -> McpServerConfigTyped { + serde_json::from_value(self.server_config.clone()).unwrap_or_else(|_| { + // 尝试手动提取字段 + McpServerConfigTyped { + command: self + .server_config + .get("command") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(), + args: self + .server_config + .get("args") + .and_then(|v| v.as_array()) + .map(|arr| { + arr.iter() + .filter_map(|v| v.as_str().map(|s| s.to_string())) + .collect() + }) + .unwrap_or_default(), + env: self + .server_config + .get("env") + .and_then(|v| v.as_object()) + .map(|obj| { + obj.iter() + .filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_string()))) + .collect() + }) + .unwrap_or_default(), + cwd: self + .server_config + .get("cwd") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()), + timeout: self + .server_config + .get("timeout") + .and_then(|v| v.as_u64()) + .unwrap_or(30), + } + }) + } + + /// 验证服务器配置 + /// + /// 检查配置是否有效,返回验证错误列表。 + /// 空列表表示配置有效。 + pub fn validate_config(&self) -> Vec { + let mut errors = Vec::new(); + let config = self.parse_config(); + + // 验证 command 不为空 + if config.command.trim().is_empty() { + errors.push(ConfigValidationError { + field: "command".to_string(), + message: "启动命令不能为空".to_string(), + }); + } + + // 验证 name 不为空 + if self.name.trim().is_empty() { + errors.push(ConfigValidationError { + field: "name".to_string(), + message: "服务器名称不能为空".to_string(), + }); + } + + // 验证 name 不包含特殊字符(用于工具名称前缀) + if !self + .name + .chars() + .all(|c| c.is_alphanumeric() || c == '-' || c == '_') + { + errors.push(ConfigValidationError { + field: "name".to_string(), + message: "服务器名称只能包含字母、数字、连字符和下划线".to_string(), + }); + } + + // 验证 timeout 在合理范围内 + if config.timeout == 0 || config.timeout > 300 { + errors.push(ConfigValidationError { + field: "timeout".to_string(), + message: "超时时间必须在 1-300 秒之间".to_string(), + }); + } + + errors + } + + /// 检查配置是否有效 + pub fn is_valid(&self) -> bool { + self.validate_config().is_empty() + } } diff --git a/src-tauri/src/models/project_model.rs b/src-tauri/src/models/project_model.rs index 6a42ab7a4..d264dc178 100644 --- a/src-tauri/src/models/project_model.rs +++ b/src-tauri/src/models/project_model.rs @@ -178,6 +178,7 @@ impl Default for MaterialType { } } +#[allow(dead_code)] impl MaterialType { pub fn as_str(&self) -> &'static str { match self { @@ -239,6 +240,7 @@ impl Default for ImageCategory { } } +#[allow(dead_code)] impl ImageCategory { pub fn as_str(&self) -> &'static str { match self { @@ -298,6 +300,7 @@ impl Default for LayoutCategory { } } +#[allow(dead_code)] impl LayoutCategory { pub fn as_str(&self) -> &'static str { match self { @@ -567,6 +570,7 @@ impl Default for Platform { } } +#[allow(dead_code)] impl Platform { pub fn as_str(&self) -> &'static str { match self { @@ -622,6 +626,7 @@ impl Default for EmojiUsage { } } +#[allow(dead_code)] impl EmojiUsage { pub fn as_str(&self) -> &'static str { match self { @@ -820,6 +825,7 @@ impl Default for BrandPersonality { } } +#[allow(dead_code)] impl BrandPersonality { pub fn as_str(&self) -> &'static str { match self { @@ -885,6 +891,7 @@ impl Default for DesignStyle { } } +#[allow(dead_code)] impl DesignStyle { pub fn as_str(&self) -> &'static str { match self { diff --git a/src-tauri/src/services/aster_session_store.rs b/src-tauri/src/services/aster_session_store.rs index 511f57b5d..a59baaabe 100644 --- a/src-tauri/src/services/aster_session_store.rs +++ b/src-tauri/src/services/aster_session_store.rs @@ -18,6 +18,7 @@ use aster::session::{ use async_trait::async_trait; use chrono::Utc; use std::collections::HashMap; +use std::fs; use std::path::PathBuf; /// ProxyCast 的 SessionStore 实现 @@ -44,6 +45,43 @@ impl ProxyCastSessionStore { "assistant".to_string() } } + + /// 解析会话 working_dir(优先默认 workspace,其次应用默认项目目录) + fn resolve_session_working_dir(conn: &rusqlite::Connection) -> PathBuf { + // 1) 优先使用默认 workspace(is_default = 1) + let default_workspace_path: Option = conn + .query_row( + "SELECT root_path FROM workspaces WHERE is_default = 1 LIMIT 1", + [], + |row| row.get(0), + ) + .ok(); + + if let Some(path) = default_workspace_path { + if !path.trim().is_empty() { + let pb = PathBuf::from(path); + return if pb.is_absolute() { + pb + } else { + std::env::current_dir() + .unwrap_or_else(|_| PathBuf::from(".")) + .join(pb) + }; + } + } + + // 2) 回退到 ~/.proxycast/projects/default + if let Some(home) = dirs::home_dir() { + let fallback = home.join(".proxycast").join("projects").join("default"); + if !fallback.exists() { + let _ = fs::create_dir_all(&fallback); + } + return fallback; + } + + // 3) 最终回退到进程当前目录 + std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")) + } } #[async_trait] @@ -154,6 +192,7 @@ impl SessionStore for ProxyCastSessionStore { .unwrap_or_else(|_| Utc::now()); let session_type = model.parse().unwrap_or(SessionType::User); + let working_dir = Self::resolve_session_working_dir(&conn); let conversation = if include_messages { Some(self.load_conversation(&conn, &id)?) @@ -165,7 +204,7 @@ impl SessionStore for ProxyCastSessionStore { Ok(Session { id: id.to_string(), - working_dir: PathBuf::from("."), + working_dir, name: title.unwrap_or_else(|| "未命名会话".to_string()), user_set_name: false, session_type, @@ -334,6 +373,7 @@ impl SessionStore for ProxyCastSessionStore { async fn list_sessions(&self) -> Result> { let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; + let default_working_dir = Self::resolve_session_working_dir(&conn); let mut stmt = conn.prepare( "SELECT id, model, system_prompt, title, created_at, updated_at @@ -362,7 +402,7 @@ impl SessionStore for ProxyCastSessionStore { Session { id, - working_dir: PathBuf::from("."), + working_dir: default_working_dir.clone(), name: title.unwrap_or_else(|| "未命名会话".to_string()), user_set_name: false, session_type, diff --git a/src-tauri/src/services/live_sync.rs b/src-tauri/src/services/live_sync.rs index 08ecd31cc..8808b1137 100644 --- a/src-tauri/src/services/live_sync.rs +++ b/src-tauri/src/services/live_sync.rs @@ -10,6 +10,10 @@ const ENV_BLOCK_END: &str = "# <<< ProxyCast Claude Config <<<"; /// 原子写入 JSON 文件,防止配置损坏 /// 参考 cc-switch 的实现:使用临时文件 + 重命名的原子操作 +/// +/// Windows 优化: +/// - 避免不必要的 flush() 调用(Windows 上 flush 会触发磁盘同步) +/// - 跳过验证步骤以减少文件读取 pub(crate) fn write_json_file_atomic( path: &std::path::Path, value: &Value, @@ -29,12 +33,21 @@ pub(crate) fn write_json_file_atomic( let content = serde_json::to_string_pretty(value)?; let mut temp_file = fs::File::create(&temp_path)?; temp_file.write_all(content.as_bytes())?; + + // Windows 优化:只在非 Windows 平台调用 flush + // Windows 上 flush() 会触发 FlushFileBuffers(),导致等待物理磁盘写入 + #[cfg(not(target_os = "windows"))] temp_file.flush()?; + drop(temp_file); // 确保文件句柄被释放 - // 验证 JSON 格式正确性 - let verify_content = fs::read_to_string(&temp_path)?; - let _: Value = serde_json::from_str(&verify_content)?; // 验证解析 + // Windows 优化:跳过验证步骤,减少一次文件读取 + // 验证主要是为了防止 JSON 序列化错误,但 serde_json 已经保证了正确性 + #[cfg(not(target_os = "windows"))] + { + let verify_content = fs::read_to_string(&temp_path)?; + let _: Value = serde_json::from_str(&verify_content)?; // 验证解析 + } // 原子性重命名 fs::rename(&temp_path, path)?; @@ -47,6 +60,11 @@ pub(crate) fn write_json_file_atomic( pub(crate) fn create_backup( path: &std::path::Path, ) -> Result<(), Box> { + if !should_create_backup() { + tracing::info!("Skip backup for: {}", path.display()); + return Ok(()); + } + if path.exists() { let backup_path = path.with_extension("bak"); std::fs::copy(path, &backup_path)?; @@ -55,6 +73,19 @@ pub(crate) fn create_backup( Ok(()) } +fn should_create_backup() -> bool { + if cfg!(target_os = "windows") { + return std::env::var("PROXYCAST_FORCE_BACKUP") + .map(|value| { + let value = value.to_lowercase(); + value == "1" || value == "true" || value == "yes" + }) + .unwrap_or(false); + } + + true +} + /// 获取当前 shell 配置文件路径 /// 优先级:zsh > bash fn get_shell_config_path() -> Result> { @@ -88,6 +119,8 @@ fn get_shell_config_path() -> Result Result<(), Box> { @@ -146,12 +179,15 @@ pub fn write_env_to_shell_config( new_content.push('\n'); } - // 创建备份 + // 创建备份(Windows 优化:异步或跳过备份可以进一步优化) create_backup(&config_path)?; // 写入文件 let mut file = fs::File::create(&config_path)?; file.write_all(new_content.as_bytes())?; + + // Windows 优化:只在非 Windows 平台调用 flush + #[cfg(not(target_os = "windows"))] file.flush()?; tracing::info!( @@ -338,18 +374,21 @@ fn sync_claude_settings( write_json_file_atomic(&config_path, &settings)?; tracing::info!("Claude 配置文件同步完成: {}", config_path.display()); - // 同时写入 shell 配置文件 + // 同时写入 shell 配置文件(后台任务,避免阻塞切换响应) if !env_vars_for_shell.is_empty() { - match write_env_to_shell_config(&env_vars_for_shell) { - Ok(_) => { - tracing::info!("Claude 环境变量已写入 shell 配置文件"); - tracing::info!("请重启终端或执行 'source ~/.zshrc' (或 ~/.bashrc) 使配置生效"); - } - Err(e) => { - tracing::warn!("写入 shell 配置文件失败: {}", e); - // 不中断流程,配置文件方式仍然可用 - } - } + let env_vars_for_shell = env_vars_for_shell; + std::thread::spawn( + move || match write_env_to_shell_config(&env_vars_for_shell) { + Ok(_) => { + tracing::info!("Claude 环境变量已写入 shell 配置文件"); + tracing::info!("请重启终端或执行 'source ~/.zshrc' (或 ~/.bashrc) 使配置生效"); + } + Err(e) => { + tracing::warn!("写入 shell 配置文件失败: {}", e); + // 不中断流程,配置文件方式仍然可用 + } + }, + ); } Ok(()) diff --git a/src-tauri/src/services/mcp_service.rs b/src-tauri/src/services/mcp_service.rs index ae7963ff5..399a44339 100644 --- a/src-tauri/src/services/mcp_service.rs +++ b/src-tauri/src/services/mcp_service.rs @@ -1,5 +1,6 @@ use crate::database::dao::mcp::McpDao; use crate::database::DbConnection; +use crate::models::mcp_model::ConfigValidationError; use crate::models::{AppType, McpServer}; use crate::services::mcp_sync; @@ -11,7 +12,58 @@ impl McpService { McpDao::get_all(&conn).map_err(|e| e.to_string()) } + /// 根据名称获取服务器 + pub fn get_by_name(db: &DbConnection, name: &str) -> Result, String> { + let conn = db.lock().map_err(|e| e.to_string())?; + let servers = McpDao::get_all(&conn).map_err(|e| e.to_string())?; + Ok(servers.into_iter().find(|s| s.name == name)) + } + + /// 检查名称是否已存在(排除指定 ID) + pub fn is_name_duplicate( + db: &DbConnection, + name: &str, + exclude_id: Option<&str>, + ) -> Result { + let conn = db.lock().map_err(|e| e.to_string())?; + let servers = McpDao::get_all(&conn).map_err(|e| e.to_string())?; + Ok(servers + .iter() + .any(|s| s.name == name && exclude_id.map_or(true, |id| s.id != id))) + } + + /// 验证服务器配置 + pub fn validate_server( + db: &DbConnection, + server: &McpServer, + is_update: bool, + ) -> Result, String> { + let mut errors = server.validate_config(); + + // 检查名称重复 + let exclude_id = if is_update { + Some(server.id.as_str()) + } else { + None + }; + if Self::is_name_duplicate(db, &server.name, exclude_id)? { + errors.push(ConfigValidationError { + field: "name".to_string(), + message: format!("服务器名称 '{}' 已存在", server.name), + }); + } + + Ok(errors) + } + pub fn add(db: &DbConnection, server: McpServer) -> Result<(), String> { + // 验证配置 + let errors = Self::validate_server(db, &server, false)?; + if !errors.is_empty() { + let error_msgs: Vec = errors.iter().map(|e| e.message.clone()).collect(); + return Err(format!("配置验证失败: {}", error_msgs.join("; "))); + } + let conn = db.lock().map_err(|e| e.to_string())?; McpDao::insert(&conn, &server).map_err(|e| e.to_string())?; @@ -23,6 +75,13 @@ impl McpService { } pub fn update(db: &DbConnection, server: McpServer) -> Result<(), String> { + // 验证配置 + let errors = Self::validate_server(db, &server, true)?; + if !errors.is_empty() { + let error_msgs: Vec = errors.iter().map(|e| e.message.clone()).collect(); + return Err(format!("配置验证失败: {}", error_msgs.join("; "))); + } + let conn = db.lock().map_err(|e| e.to_string())?; McpDao::update(&conn, &server).map_err(|e| e.to_string())?; diff --git a/src-tauri/src/services/mcp_sync.rs b/src-tauri/src/services/mcp_sync.rs index e392c35bf..ca5cd48c0 100644 --- a/src-tauri/src/services/mcp_sync.rs +++ b/src-tauri/src/services/mcp_sync.rs @@ -381,7 +381,7 @@ pub fn import_mcp_from_claude( name: id.clone(), server_config: config.clone(), description: None, - enabled_proxycast: false, + enabled_proxycast: true, enabled_claude: true, enabled_codex: false, enabled_gemini: false, @@ -426,7 +426,7 @@ pub fn import_mcp_from_codex( name: id.clone(), server_config: Value::Object(current_config.clone()), description: None, - enabled_proxycast: false, + enabled_proxycast: true, enabled_claude: false, enabled_codex: true, enabled_gemini: false, @@ -460,7 +460,7 @@ pub fn import_mcp_from_codex( name: id.clone(), server_config: Value::Object(current_config.clone()), description: None, - enabled_proxycast: false, + enabled_proxycast: true, enabled_claude: false, enabled_codex: true, enabled_gemini: false, @@ -508,7 +508,7 @@ pub fn import_mcp_from_codex( name: id.clone(), server_config: Value::Object(current_config), description: None, - enabled_proxycast: false, + enabled_proxycast: true, enabled_claude: false, enabled_codex: true, enabled_gemini: false, @@ -542,7 +542,7 @@ pub fn import_mcp_from_gemini( name: id.clone(), server_config: config.clone(), description: None, - enabled_proxycast: false, + enabled_proxycast: true, enabled_claude: false, enabled_codex: false, enabled_gemini: true, diff --git a/src-tauri/src/services/switch.rs b/src-tauri/src/services/switch.rs index b28f899af..bb9753e1a 100644 --- a/src-tauri/src/services/switch.rs +++ b/src-tauri/src/services/switch.rs @@ -2,9 +2,20 @@ use crate::database::dao::providers::ProviderDao; use crate::database::DbConnection; use crate::models::{AppType, Provider}; use crate::services::live_sync; +use once_cell::sync::Lazy; +use tokio::sync::Mutex; pub struct SwitchService; +static SWITCH_PROVIDER_LOCK: Lazy> = Lazy::new(|| Mutex::new(())); + +/// 用于在异步上下文中传递的切换数据 +struct SwitchContext { + target_provider: Provider, + current_provider: Option, + app_type_enum: AppType, +} + impl SwitchService { pub fn get_providers(db: &DbConnection, app_type: &str) -> Result, String> { let conn = db.lock().map_err(|e| e.to_string())?; @@ -181,6 +192,168 @@ impl SwitchService { Ok(()) } + /// 异步版本的 switch_provider,优化 Windows 性能 + /// + /// 优化策略: + /// 1. 减少数据库锁持有时间 - 先获取数据,释放锁,执行 I/O,再获取锁更新 + /// 2. 使用 spawn_blocking 将文件 I/O 移出主线程 + /// 3. 使用全局互斥锁确保切换流程串行化,避免并发写入 + pub async fn switch_provider_async( + db: &DbConnection, + app_type: &str, + id: &str, + ) -> Result<(), String> { + use tracing::{error, info, warn}; + + info!("开始切换 {} 配置到 provider: {} (异步)", app_type, id); + let _switch_guard = SWITCH_PROVIDER_LOCK.lock().await; + + // Step 1: 获取数据(短暂持有锁) + let ctx = { + let conn = db.lock().map_err(|e| e.to_string())?; + + // Get target provider + let target_provider = ProviderDao::get_by_id(&conn, app_type, id) + .map_err(|e| { + error!("查找目标 provider 失败: {}", e); + e.to_string() + })? + .ok_or_else(|| { + error!("目标 provider 不存在: {}", id); + format!("Provider not found: {id}") + })?; + + let app_type_enum = app_type.parse::().map_err(|e| { + error!("无效的 app_type: {} - {}", app_type, e); + e.to_string() + })?; + + // 获取当前 provider(用于回填和回滚) + let current_provider = if app_type_enum != AppType::ProxyCast { + ProviderDao::get_current(&conn, app_type).map_err(|e| { + error!("获取当前 provider 失败: {}", e); + e.to_string() + })? + } else { + None + }; + + // 锁在这里释放 + SwitchContext { + target_provider, + current_provider, + app_type_enum, + } + }; + + // Step 2: 执行文件 I/O(在后台线程,不持有锁) + if ctx.app_type_enum != AppType::ProxyCast { + let current_for_backfill = ctx.current_provider.clone(); + let app_type_for_sync = ctx.app_type_enum.clone(); + let target_id = id.to_string(); + + // 使用 spawn_blocking 将文件 I/O 移到后台线程 + let sync_result = tokio::task::spawn_blocking(move || { + // Step 2a: Backfill - 回填当前配置 + if let Some(ref current) = current_for_backfill { + if current.id != target_id { + info!("回填当前配置: {}", current.name); + match live_sync::read_live_settings(&app_type_for_sync) { + Ok(live_settings) => { + // 返回需要更新的 provider 数据 + Some((current.clone(), live_settings)) + } + Err(e) => { + warn!("读取当前配置失败,跳过回填: {}", e); + None + } + } + } else { + None + } + } else { + None + } + }) + .await + .map_err(|e| format!("后台任务失败: {e}"))?; + + // 如果需要回填,更新数据库(短暂持有锁) + if let Some((mut current, live_settings)) = sync_result { + let conn = db.lock().map_err(|e| e.to_string())?; + current.settings_config = live_settings; + if let Err(e) = ProviderDao::update(&conn, ¤t) { + warn!("回填配置失败,但继续执行: {}", e); + } else { + info!("回填配置完成"); + } + // 锁在这里释放 + } + + // Step 2b: 同步新配置(在后台线程) + let target_for_sync = ctx.target_provider.clone(); + let current_for_restore = ctx.current_provider.clone(); + let app_type_for_sync = ctx.app_type_enum.clone(); + + tokio::task::spawn_blocking(move || { + info!("验证目标配置可同步性"); + if let Err(sync_error) = + live_sync::sync_to_live(&app_type_for_sync, &target_for_sync) + { + error!("配置同步失败: {}", sync_error); + + // 尝试恢复原配置(如果有) + if let Some(ref current) = current_for_restore { + warn!("尝试恢复原配置: {}", current.name); + if let Err(restore_error) = + live_sync::sync_to_live(&app_type_for_sync, current) + { + error!("恢复原配置失败: {}", restore_error); + return Err(format!("切换失败且无法恢复原配置: {sync_error}")); + } + } + + return Err(format!("配置同步失败: {sync_error}")); + } + Ok(()) + }) + .await + .map_err(|e| format!("后台任务失败: {e}"))??; + } + + // Step 3: 更新数据库(短暂持有锁) + { + let conn = db.lock().map_err(|e| e.to_string())?; + info!("更新数据库中的当前 provider"); + if let Err(db_error) = ProviderDao::set_current(&conn, app_type, id) { + error!("数据库更新失败: {}", db_error); + + // 如果数据库更新失败,尝试恢复原配置文件 + if ctx.app_type_enum != AppType::ProxyCast { + if let Some(ref current) = ctx.current_provider { + warn!("数据库更新失败,尝试恢复原配置文件"); + let current_clone = current.clone(); + let app_type_clone = ctx.app_type_enum.clone(); + // 在后台线程恢复 + let _ = tokio::task::spawn_blocking(move || { + if let Err(restore_error) = + live_sync::sync_to_live(&app_type_clone, ¤t_clone) + { + error!("恢复配置文件失败: {}", restore_error); + } + }); + } + } + + return Err(db_error.to_string()); + } + // 锁在这里释放 + } + + info!("配置切换成功: {} -> {}", app_type, ctx.target_provider.name); + Ok(()) + } + /// Import current live config as a default provider pub fn import_default_config(db: &DbConnection, app_type: &str) -> Result { let conn = db.lock().map_err(|e| e.to_string())?; diff --git a/src-tauri/src/skills/README.md b/src-tauri/src/skills/README.md new file mode 100644 index 000000000..3dea5ad05 --- /dev/null +++ b/src-tauri/src/skills/README.md @@ -0,0 +1,96 @@ +# Skills 集成模块 + +本模块实现 aster-rust Skills 系统与 ProxyCast 的集成。 + +## 模块结构 + +| 文件 | 说明 | +|------|------| +| `mod.rs` | 模块导出 | +| `llm_provider.rs` | ProxyCastLlmProvider 实现 | +| `execution_callback.rs` | TauriExecutionCallback 实现 | + +## Skills 集成架构 + +### AI 自动调用 Skills(方案 A) + +ProxyCast 通过以下机制让 AI 能够自动发现和调用 Skills: + +1. **Agent 初始化时加载 Skills** + - `AsterAgentState::init_agent_with_db()` 调用 `load_proxycast_skills()` + - 从 `~/.proxycast/skills/` 目录加载所有 Skills + - 注册到 aster-rust 的 `global_registry` + +2. **SkillTool 自动注册** + - aster-rust 的 `register_default_tools()` 自动注册 `SkillTool` + - `SkillTool` 从 `global_registry` 读取可用 Skills + - AI 可以通过 `Skill` 工具调用任意已注册的 Skill + +3. **动态刷新** + - 安装/卸载 Skills 后调用 `AsterAgentState::reload_proxycast_skills()` + - 自动更新 `global_registry`,无需重启应用 + +### 数据流 + +``` +用户安装 Skill + ↓ +skill_cmd.rs::install_skill_for_app() + ↓ +AsterAgentState::reload_proxycast_skills() + ↓ +aster::skills::global_registry 更新 + ↓ +AI 通过 SkillTool 发现新 Skill + ↓ +用户对话时 AI 自动调用相关 Skill +``` + +## 核心组件 + +### ProxyCastLlmProvider + +使用 ProviderPoolService 选择凭证并调用 LLM API。 + +**功能**: +- 通过 ProviderPoolService 选择可用凭证 +- 支持指定 provider 类型和 model 参数 +- 智能降级到 API Key Provider + +### TauriExecutionCallback + +通过 Tauri 事件系统向前端发送 Skill 执行进度更新。 + +**事件类型**: +- `skill:step_start`: 步骤开始 +- `skill:step_complete`: 步骤完成 +- `skill:step_error`: 步骤错误 +- `skill:complete`: 执行完成 + +## 依赖关系 + +``` +agent/aster_state.rs +├── load_proxycast_skills() +│ ├── aster::skills::load_skills_from_directory() +│ └── aster::skills::global_registry() +└── reload_proxycast_skills() + +skills/ +├── llm_provider.rs +│ ├── ProviderPoolService (凭证池管理) +│ └── ApiKeyProviderService (API Key 服务) +└── execution_callback.rs + └── tauri::AppHandle (事件发送) + +commands/skill_cmd.rs +├── install_skill_for_app() +│ └── AsterAgentState::reload_proxycast_skills() +└── uninstall_skill_for_app() + └── AsterAgentState::reload_proxycast_skills() +``` + +## 相关文档 + +- 设计文档: `.kiro/specs/skills-integration/design.md` +- 需求文档: `.kiro/specs/skills-integration/requirements.md` diff --git a/src-tauri/src/skills/execution_callback.rs b/src-tauri/src/skills/execution_callback.rs new file mode 100644 index 000000000..4da4a535a --- /dev/null +++ b/src-tauri/src/skills/execution_callback.rs @@ -0,0 +1,320 @@ +//! Tauri 执行回调实现 +//! +//! 实现 aster-rust 的 ExecutionCallback trait,通过 Tauri 事件系统向前端发送进度更新。 +//! +//! ## 事件类型 +//! - `skill:step_start`: 步骤开始 +//! - `skill:step_complete`: 步骤完成 +//! - `skill:step_error`: 步骤错误 +//! - `skill:complete`: 执行完成 +//! +//! ## 使用示例 +//! ```ignore +//! let callback = TauriExecutionCallback::new(app_handle, "exec-123".to_string()); +//! callback.on_step_start("step-1", "数据处理", 1, 3); +//! ``` + +use serde::Serialize; +use std::sync::atomic::{AtomicUsize, Ordering}; +use tauri::{AppHandle, Emitter}; + +/// 步骤开始事件 Payload +#[derive(Debug, Clone, Serialize)] +pub struct StepStartPayload { + /// 执行 ID + pub execution_id: String, + /// 步骤 ID + pub step_id: String, + /// 步骤名称 + pub step_name: String, + /// 当前步骤序号(从 1 开始) + pub current_step: usize, + /// 总步骤数 + pub total_steps: usize, +} + +/// 步骤完成事件 Payload +#[derive(Debug, Clone, Serialize)] +pub struct StepCompletePayload { + /// 执行 ID + pub execution_id: String, + /// 步骤 ID + pub step_id: String, + /// 步骤输出 + pub output: String, +} + +/// 步骤错误事件 Payload +#[derive(Debug, Clone, Serialize)] +pub struct StepErrorPayload { + /// 执行 ID + pub execution_id: String, + /// 步骤 ID + pub step_id: String, + /// 错误信息 + pub error: String, + /// 是否会重试 + pub will_retry: bool, +} + +/// 执行完成事件 Payload +#[derive(Debug, Clone, Serialize)] +pub struct ExecutionCompletePayload { + /// 执行 ID + pub execution_id: String, + /// 是否成功 + pub success: bool, + /// 最终输出(成功时) + pub output: Option, + /// 错误信息(失败时) + pub error: Option, +} + +/// Tauri 事件名称常量 +pub mod events { + /// 步骤开始事件 + pub const STEP_START: &str = "skill:step_start"; + /// 步骤完成事件 + pub const STEP_COMPLETE: &str = "skill:step_complete"; + /// 步骤错误事件 + pub const STEP_ERROR: &str = "skill:step_error"; + /// 执行完成事件 + pub const COMPLETE: &str = "skill:complete"; +} + +/// ExecutionCallback Trait +/// +/// 定义 Skill 执行过程中的回调接口。 +/// 应用层需要实现此 trait 以接收执行进度更新。 +pub trait ExecutionCallback: Send + Sync { + /// 步骤开始回调 + /// + /// # 参数 + /// - `step_id`: 步骤 ID + /// - `step_name`: 步骤名称 + /// - `current_step`: 当前步骤序号(从 1 开始) + /// - `total_steps`: 总步骤数 + fn on_step_start( + &self, + step_id: &str, + step_name: &str, + current_step: usize, + total_steps: usize, + ); + + /// 步骤完成回调 + /// + /// # 参数 + /// - `step_id`: 步骤 ID + /// - `output`: 步骤输出 + fn on_step_complete(&self, step_id: &str, output: &str); + + /// 步骤错误回调 + /// + /// # 参数 + /// - `step_id`: 步骤 ID + /// - `error`: 错误信息 + /// - `will_retry`: 是否会重试 + fn on_step_error(&self, step_id: &str, error: &str, will_retry: bool); + + /// 执行完成回调 + /// + /// # 参数 + /// - `success`: 是否成功 + /// - `final_output`: 最终输出(成功时) + /// - `error`: 错误信息(失败时) + fn on_complete(&self, success: bool, final_output: Option<&str>, error: Option<&str>); +} + +/// Tauri 执行回调 +/// +/// 通过 Tauri 事件系统向前端发送 Skill 执行进度更新。 +/// 实现 aster-rust 定义的 ExecutionCallback trait。 +pub struct TauriExecutionCallback { + /// Tauri AppHandle + app_handle: AppHandle, + /// 执行 ID(用于区分多个并发执行) + execution_id: String, + /// 当前步骤计数器(用于跟踪步骤序号) + current_step: AtomicUsize, +} + +impl TauriExecutionCallback { + /// 创建新的 TauriExecutionCallback 实例 + /// + /// # Arguments + /// * `app_handle` - Tauri AppHandle + /// * `execution_id` - 执行 ID,用于区分多个并发执行 + pub fn new(app_handle: AppHandle, execution_id: String) -> Self { + Self { + app_handle, + execution_id, + current_step: AtomicUsize::new(0), + } + } + + /// 获取执行 ID + pub fn execution_id(&self) -> &str { + &self.execution_id + } + + /// 获取当前步骤序号 + pub fn current_step(&self) -> usize { + self.current_step.load(Ordering::SeqCst) + } +} + +/// ExecutionCallback trait 实现 +/// +/// 通过 Tauri 事件系统向前端发送进度更新。 +/// +/// # Requirements +/// - 2.2: on_step_start 发送 "skill:step_start" 事件 +/// - 2.3: on_step_complete 发送 "skill:step_complete" 事件 +/// - 2.4: on_step_error 发送 "skill:step_error" 事件 +/// - 2.5: on_complete 发送 "skill:complete" 事件 +impl ExecutionCallback for TauriExecutionCallback { + /// 步骤开始回调 + /// + /// 发送 "skill:step_start" Tauri 事件到前端。 + /// + /// # Requirements + /// - 2.2: WHEN on_step_start is called, emit a "skill:step_start" Tauri event + fn on_step_start( + &self, + step_id: &str, + step_name: &str, + current_step: usize, + total_steps: usize, + ) { + // 更新当前步骤计数器 + self.current_step.store(current_step, Ordering::SeqCst); + + let payload = StepStartPayload { + execution_id: self.execution_id.clone(), + step_id: step_id.to_string(), + step_name: step_name.to_string(), + current_step, + total_steps, + }; + + tracing::info!( + "[TauriExecutionCallback] 步骤开始: execution_id={}, step_id={}, step_name={}, {}/{}", + self.execution_id, + step_id, + step_name, + current_step, + total_steps + ); + + if let Err(e) = self.app_handle.emit(events::STEP_START, &payload) { + tracing::error!( + "[TauriExecutionCallback] 发送 {} 事件失败: {}", + events::STEP_START, + e + ); + } + } + + /// 步骤完成回调 + /// + /// 发送 "skill:step_complete" Tauri 事件到前端。 + /// + /// # Requirements + /// - 2.3: WHEN on_step_complete is called, emit a "skill:step_complete" Tauri event + fn on_step_complete(&self, step_id: &str, output: &str) { + let payload = StepCompletePayload { + execution_id: self.execution_id.clone(), + step_id: step_id.to_string(), + output: output.to_string(), + }; + + tracing::info!( + "[TauriExecutionCallback] 步骤完成: execution_id={}, step_id={}, output_len={}", + self.execution_id, + step_id, + output.len() + ); + + if let Err(e) = self.app_handle.emit(events::STEP_COMPLETE, &payload) { + tracing::error!( + "[TauriExecutionCallback] 发送 {} 事件失败: {}", + events::STEP_COMPLETE, + e + ); + } + } + + /// 步骤错误回调 + /// + /// 发送 "skill:step_error" Tauri 事件到前端。 + /// + /// # Requirements + /// - 2.4: WHEN on_step_error is called, emit a "skill:step_error" Tauri event + fn on_step_error(&self, step_id: &str, error: &str, will_retry: bool) { + let payload = StepErrorPayload { + execution_id: self.execution_id.clone(), + step_id: step_id.to_string(), + error: error.to_string(), + will_retry, + }; + + tracing::warn!( + "[TauriExecutionCallback] 步骤错误: execution_id={}, step_id={}, error={}, will_retry={}", + self.execution_id, + step_id, + error, + will_retry + ); + + if let Err(e) = self.app_handle.emit(events::STEP_ERROR, &payload) { + tracing::error!( + "[TauriExecutionCallback] 发送 {} 事件失败: {}", + events::STEP_ERROR, + e + ); + } + } + + /// 执行完成回调 + /// + /// 发送 "skill:complete" Tauri 事件到前端。 + /// + /// # Requirements + /// - 2.5: WHEN on_complete is called, emit a "skill:complete" Tauri event + fn on_complete(&self, success: bool, final_output: Option<&str>, error: Option<&str>) { + let payload = ExecutionCompletePayload { + execution_id: self.execution_id.clone(), + success, + output: final_output.map(|s| s.to_string()), + error: error.map(|s| s.to_string()), + }; + + if success { + tracing::info!( + "[TauriExecutionCallback] 执行完成: execution_id={}, success=true, output_len={}", + self.execution_id, + final_output.map(|s| s.len()).unwrap_or(0) + ); + } else { + tracing::warn!( + "[TauriExecutionCallback] 执行失败: execution_id={}, error={:?}", + self.execution_id, + error + ); + } + + if let Err(e) = self.app_handle.emit(events::COMPLETE, &payload) { + tracing::error!( + "[TauriExecutionCallback] 发送 {} 事件失败: {}", + events::COMPLETE, + e + ); + } + } +} + +#[cfg(test)] +mod tests { + // TODO: 在 Task 1.5 中添加属性测试 +} diff --git a/src-tauri/src/skills/llm_provider.rs b/src-tauri/src/skills/llm_provider.rs new file mode 100644 index 000000000..57f68cc41 --- /dev/null +++ b/src-tauri/src/skills/llm_provider.rs @@ -0,0 +1,597 @@ +//! ProxyCast LLM Provider 实现 +//! +//! 实现 aster-rust 的 LlmProvider trait,使用 ProviderPoolService 选择凭证并调用 LLM API。 +//! +//! ## 功能 +//! - 通过 ProviderPoolService 选择可用凭证 +//! - 支持指定 provider 类型和 model 参数 +//! - 智能降级到 API Key Provider +//! +//! ## 依赖 +//! - `ProviderPoolService`: 凭证池管理 +//! - `ApiKeyProviderService`: API Key 服务(降级使用) + +use std::sync::Arc; + +use async_trait::async_trait; +use serde::{Deserialize, Serialize}; + +use crate::database::DbConnection; +use crate::models::anthropic::AnthropicMessagesRequest; +#[cfg(test)] +use crate::models::provider_pool_model::PoolProviderType; +use crate::models::provider_pool_model::{CredentialData, ProviderCredential}; +use crate::providers::{ClaudeCustomProvider, KiroProvider, OpenAICustomProvider}; +use crate::services::api_key_provider_service::ApiKeyProviderService; +use crate::services::provider_pool_service::ProviderPoolService; + +/// Skill 执行错误类型 +/// +/// 用于 LlmProvider trait 的错误返回 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum SkillError { + /// Provider 错误(凭证不可用、API 调用失败等) + ProviderError(String), + /// 执行错误(Skill 执行过程中的错误) + ExecutionError(String), + /// 配置错误 + ConfigError(String), +} + +impl std::fmt::Display for SkillError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + SkillError::ProviderError(msg) => write!(f, "Provider error: {}", msg), + SkillError::ExecutionError(msg) => write!(f, "Execution error: {}", msg), + SkillError::ConfigError(msg) => write!(f, "Config error: {}", msg), + } + } +} + +impl std::error::Error for SkillError {} + +/// LLM Provider Trait +/// +/// 定义 Skill 执行引擎调用 LLM 的接口。 +/// 应用层需要实现此 trait 以提供 LLM 调用能力。 +#[async_trait] +pub trait LlmProvider: Send + Sync { + /// 调用 LLM 进行对话 + /// + /// # 参数 + /// - `system_prompt`: 系统提示词 + /// - `user_message`: 用户消息 + /// - `model`: 可选的模型名称 + /// + /// # 返回 + /// - `Ok(String)`: LLM 的响应文本 + /// - `Err(SkillError)`: 调用失败时的错误 + async fn chat( + &self, + system_prompt: &str, + user_message: &str, + model: Option<&str>, + ) -> Result; +} + +/// ProxyCast LLM Provider +/// +/// 使用 ProviderPoolService 选择凭证并调用 LLM API。 +/// 实现 aster-rust 定义的 LlmProvider trait。 +pub struct ProxyCastLlmProvider { + /// 凭证池服务 + pool_service: Arc, + /// API Key Provider 服务(用于智能降级) + api_key_service: Arc, + /// 数据库连接 + db: DbConnection, + /// 偏好的 Provider 类型(可选) + preferred_provider: Option, +} + +impl ProxyCastLlmProvider { + /// 创建新的 ProxyCastLlmProvider 实例 + /// + /// # Arguments + /// * `pool_service` - 凭证池服务 + /// * `api_key_service` - API Key 服务 + /// * `db` - 数据库连接 + pub fn new( + pool_service: Arc, + api_key_service: Arc, + db: DbConnection, + ) -> Self { + Self { + pool_service, + api_key_service, + db, + preferred_provider: None, + } + } + + /// 创建带有偏好 Provider 的实例 + /// + /// # Arguments + /// * `pool_service` - 凭证池服务 + /// * `api_key_service` - API Key 服务 + /// * `db` - 数据库连接 + /// * `preferred_provider` - 偏好的 Provider 类型 + pub fn with_preferred_provider( + pool_service: Arc, + api_key_service: Arc, + db: DbConnection, + preferred_provider: String, + ) -> Self { + Self { + pool_service, + api_key_service, + db, + preferred_provider: Some(preferred_provider), + } + } + + /// 设置偏好的 Provider 类型 + pub fn set_preferred_provider(&mut self, provider: Option) { + self.preferred_provider = provider; + } + + /// 获取偏好的 Provider 类型 + pub fn preferred_provider(&self) -> Option<&str> { + self.preferred_provider.as_deref() + } + + /// 将 Skill 的 provider 字段映射到 PoolProviderType + /// + /// # Arguments + /// * `provider` - Provider 名称字符串 + /// + /// # Returns + /// 对应的 PoolProviderType,未知类型返回 None + #[cfg(test)] + fn map_skill_provider_to_pool_type(provider: &str) -> Option { + match provider.to_lowercase().as_str() { + "openai" | "gpt" => Some(PoolProviderType::OpenAI), + "anthropic" | "claude" => Some(PoolProviderType::Claude), + "gemini" | "google" => Some(PoolProviderType::Gemini), + "kiro" | "codewhisperer" => Some(PoolProviderType::Kiro), + "vertex" => Some(PoolProviderType::Vertex), + "codex" => Some(PoolProviderType::Codex), + _ => None, + } + } + + /// 根据凭证调用 LLM API + /// + /// # Arguments + /// * `credential` - 选中的凭证 + /// * `system_prompt` - 系统提示词 + /// * `user_message` - 用户消息 + /// * `model` - 模型名称 + /// + /// # Returns + /// LLM 响应文本或错误 + async fn call_llm_with_credential( + &self, + credential: &ProviderCredential, + system_prompt: &str, + user_message: &str, + model: &str, + ) -> Result { + match &credential.credential { + CredentialData::KiroOAuth { creds_file_path } => { + self.call_kiro_api(creds_file_path, system_prompt, user_message, model) + .await + } + CredentialData::ClaudeKey { api_key, base_url } => { + self.call_claude_api( + api_key, + base_url.as_deref(), + system_prompt, + user_message, + model, + ) + .await + } + CredentialData::OpenAIKey { api_key, base_url } => { + self.call_openai_api( + api_key, + base_url.as_deref(), + system_prompt, + user_message, + model, + ) + .await + } + CredentialData::AnthropicKey { api_key, base_url } => { + // Anthropic API Key 使用 Claude API + self.call_claude_api( + api_key, + base_url.as_deref(), + system_prompt, + user_message, + model, + ) + .await + } + _ => Err(SkillError::ProviderError(format!( + "不支持的凭证类型: {:?}", + credential.provider_type + ))), + } + } + + /// 调用 Kiro API + async fn call_kiro_api( + &self, + creds_file_path: &str, + system_prompt: &str, + user_message: &str, + model: &str, + ) -> Result { + use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; + use crate::models::anthropic::AnthropicMessage; + use crate::providers::traits::CredentialProvider; + use crate::server_utils::parse_cw_response; + + let mut kiro = KiroProvider::new(); + kiro.load_credentials_from_path(creds_file_path) + .await + .map_err(|e| SkillError::ProviderError(format!("加载 Kiro 凭证失败: {}", e)))?; + + // 确保 Token 有效 + if !kiro.is_token_valid() || kiro.is_token_expiring_soon() { + kiro.refresh_token() + .await + .map_err(|e| SkillError::ProviderError(format!("刷新 Token 失败: {}", e)))?; + } + + // 构建 Anthropic 请求 + let request = AnthropicMessagesRequest { + model: model.to_string(), + max_tokens: Some(4096), + system: Some(serde_json::Value::String(system_prompt.to_string())), + messages: vec![AnthropicMessage { + role: "user".to_string(), + content: serde_json::Value::String(user_message.to_string()), + }], + stream: false, + temperature: None, + tools: None, + tool_choice: None, + }; + + // 转换为 OpenAI 格式并调用 + let openai_request = convert_anthropic_to_openai(&request); + let resp = kiro + .call_api(&openai_request) + .await + .map_err(|e| SkillError::ProviderError(format!("Kiro API 调用失败: {}", e)))?; + + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(SkillError::ProviderError(format!( + "Kiro API 返回错误: status={}, body={}", + status, body + ))); + } + + let bytes = resp + .bytes() + .await + .map_err(|e| SkillError::ProviderError(format!("读取响应失败: {}", e)))?; + let body = String::from_utf8_lossy(&bytes).to_string(); + let parsed = parse_cw_response(&body); + + Ok(parsed.content) + } + + /// 调用 Claude API + async fn call_claude_api( + &self, + api_key: &str, + base_url: Option<&str>, + system_prompt: &str, + user_message: &str, + model: &str, + ) -> Result { + use crate::models::anthropic::AnthropicMessage; + + let claude = + ClaudeCustomProvider::with_config(api_key.to_string(), base_url.map(|s| s.to_string())); + + // 构建 Anthropic 请求 + let request = AnthropicMessagesRequest { + model: model.to_string(), + max_tokens: Some(4096), + system: Some(serde_json::Value::String(system_prompt.to_string())), + messages: vec![AnthropicMessage { + role: "user".to_string(), + content: serde_json::Value::String(user_message.to_string()), + }], + stream: false, + temperature: None, + tools: None, + tool_choice: None, + }; + + let resp = claude + .call_api(&request) + .await + .map_err(|e| SkillError::ProviderError(format!("Claude API 调用失败: {}", e)))?; + + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(SkillError::ProviderError(format!( + "Claude API 返回错误: status={}, body={}", + status, body + ))); + } + + let body = resp + .text() + .await + .map_err(|e| SkillError::ProviderError(format!("读取响应失败: {}", e)))?; + + // 解析 Anthropic 响应 + let json: serde_json::Value = serde_json::from_str(&body) + .map_err(|e| SkillError::ProviderError(format!("解析响应失败: {}", e)))?; + + // 提取文本内容 + let content = json["content"] + .as_array() + .and_then(|arr| arr.first()) + .and_then(|block| block["text"].as_str()) + .unwrap_or(""); + + Ok(content.to_string()) + } + + /// 调用 OpenAI API + async fn call_openai_api( + &self, + api_key: &str, + base_url: Option<&str>, + system_prompt: &str, + user_message: &str, + model: &str, + ) -> Result { + use crate::models::openai::{ChatCompletionRequest, ChatMessage, MessageContent}; + + let openai = + OpenAICustomProvider::with_config(api_key.to_string(), base_url.map(|s| s.to_string())); + + // 构建 OpenAI 请求 + let request = ChatCompletionRequest { + model: model.to_string(), + messages: vec![ + ChatMessage { + role: "system".to_string(), + content: Some(MessageContent::Text(system_prompt.to_string())), + tool_calls: None, + tool_call_id: None, + reasoning_content: None, + }, + ChatMessage { + role: "user".to_string(), + content: Some(MessageContent::Text(user_message.to_string())), + tool_calls: None, + tool_call_id: None, + reasoning_content: None, + }, + ], + max_tokens: Some(4096), + stream: false, + temperature: None, + top_p: None, + tools: None, + tool_choice: None, + reasoning_effort: None, + }; + + let resp = openai + .call_api(&request) + .await + .map_err(|e| SkillError::ProviderError(format!("OpenAI API 调用失败: {}", e)))?; + + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(SkillError::ProviderError(format!( + "OpenAI API 返回错误: status={}, body={}", + status, body + ))); + } + + let body = resp + .text() + .await + .map_err(|e| SkillError::ProviderError(format!("读取响应失败: {}", e)))?; + + // 解析 OpenAI 响应 + let json: serde_json::Value = serde_json::from_str(&body) + .map_err(|e| SkillError::ProviderError(format!("解析响应失败: {}", e)))?; + + // 提取文本内容 + let content = json["choices"] + .as_array() + .and_then(|arr| arr.first()) + .and_then(|choice| choice["message"]["content"].as_str()) + .unwrap_or(""); + + Ok(content.to_string()) + } +} + +#[async_trait] +impl LlmProvider for ProxyCastLlmProvider { + /// 调用 LLM 进行对话 + /// + /// # 实现说明 + /// 1. 使用 ProviderPoolService.select_credential_with_fallback() 选择凭证 + /// 2. 如果指定了 preferred_provider,优先选择该类型的凭证 + /// 3. 如果指定了 model,传递给底层 provider + /// 4. 如果没有可用凭证,返回 ProviderError + /// + /// # Requirements + /// - 1.2: 使用 ProviderPoolService 选择可用凭证 + /// - 1.3: 优先选择指定 provider 类型的凭证 + /// - 1.4: 将 model 参数传递给底层 provider + /// - 1.5: 没有可用凭证时返回 ProviderError + async fn chat( + &self, + system_prompt: &str, + user_message: &str, + model: Option<&str>, + ) -> Result { + // 确定要使用的 provider 类型 + let provider_type = self.preferred_provider.as_deref().unwrap_or("claude"); // 默认使用 Claude + + // 确定要使用的模型 + let model_name = model.unwrap_or("claude-sonnet-4-5-20250514"); + + tracing::info!( + "[ProxyCastLlmProvider] chat 调用: provider_type={}, model={}", + provider_type, + model_name + ); + + // 使用 ProviderPoolService 选择凭证(Requirements 1.2, 1.3) + let credential = self + .pool_service + .select_credential_with_fallback( + &self.db, + &self.api_key_service, + provider_type, + Some(model_name), + None, // provider_id_hint + None, // client_type + ) + .await + .map_err(|e| SkillError::ProviderError(format!("选择凭证失败: {}", e)))? + .ok_or_else(|| { + // Requirements 1.5: 没有可用凭证时返回 ProviderError + SkillError::ProviderError(format!( + "没有可用的凭证: provider_type={}, model={}", + provider_type, model_name + )) + })?; + + tracing::info!( + "[ProxyCastLlmProvider] 选中凭证: uuid={}, type={:?}", + &credential.uuid[..8], + credential.provider_type + ); + + // 调用 LLM API(Requirements 1.4: 传递 model 参数) + let result = self + .call_llm_with_credential(&credential, system_prompt, user_message, model_name) + .await; + + // 记录使用情况 + match &result { + Ok(_) => { + let _ = self.pool_service.record_usage(&self.db, &credential.uuid); + let _ = + self.pool_service + .mark_healthy(&self.db, &credential.uuid, Some(model_name)); + } + Err(e) => { + let _ = self.pool_service.mark_unhealthy( + &self.db, + &credential.uuid, + Some(&e.to_string()), + ); + } + } + + result + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_map_skill_provider_openai() { + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("openai"), + Some(PoolProviderType::OpenAI) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("gpt"), + Some(PoolProviderType::OpenAI) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("OPENAI"), + Some(PoolProviderType::OpenAI) + ); + } + + #[test] + fn test_map_skill_provider_claude() { + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("claude"), + Some(PoolProviderType::Claude) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("anthropic"), + Some(PoolProviderType::Claude) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("CLAUDE"), + Some(PoolProviderType::Claude) + ); + } + + #[test] + fn test_map_skill_provider_gemini() { + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("gemini"), + Some(PoolProviderType::Gemini) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("google"), + Some(PoolProviderType::Gemini) + ); + } + + #[test] + fn test_map_skill_provider_kiro() { + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("kiro"), + Some(PoolProviderType::Kiro) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("codewhisperer"), + Some(PoolProviderType::Kiro) + ); + } + + #[test] + fn test_map_skill_provider_unknown() { + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("unknown_provider"), + None + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type(""), + None + ); + } + + #[test] + fn test_skill_error_display() { + let provider_err = SkillError::ProviderError("没有可用凭证".to_string()); + assert!(provider_err.to_string().contains("Provider error")); + assert!(provider_err.to_string().contains("没有可用凭证")); + + let exec_err = SkillError::ExecutionError("执行失败".to_string()); + assert!(exec_err.to_string().contains("Execution error")); + + let config_err = SkillError::ConfigError("配置错误".to_string()); + assert!(config_err.to_string().contains("Config error")); + } +} diff --git a/src-tauri/src/skills/mod.rs b/src-tauri/src/skills/mod.rs new file mode 100644 index 000000000..387695195 --- /dev/null +++ b/src-tauri/src/skills/mod.rs @@ -0,0 +1,32 @@ +//! Skills 集成模块 +//! +//! 本模块实现 aster-rust Skills 系统与 ProxyCast 的集成。 +//! +//! ## 模块结构 +//! - `llm_provider`: ProxyCastLlmProvider 实现,使用 ProviderPoolService 调用 LLM +//! - `execution_callback`: TauriExecutionCallback 实现,通过 Tauri 事件发送进度 +//! +//! ## 使用示例 +//! ```ignore +//! use proxycast::skills::{ProxyCastLlmProvider, TauriExecutionCallback}; +//! +//! let provider = ProxyCastLlmProvider::new(pool_service, api_key_service, db); +//! let callback = TauriExecutionCallback::new(app_handle, execution_id); +//! ``` + +mod execution_callback; +mod llm_provider; +mod skill_loader; + +pub use execution_callback::{ + events, ExecutionCallback, ExecutionCompletePayload, StepCompletePayload, StepErrorPayload, + StepStartPayload, TauriExecutionCallback, +}; +pub use llm_provider::{LlmProvider, ProxyCastLlmProvider, SkillError}; +pub(crate) use skill_loader::{ + find_skill_by_name, get_proxycast_skills_dir, load_skills_from_directory, +}; +#[cfg(test)] +pub(crate) use skill_loader::{ + load_skill_from_file, parse_allowed_tools, parse_boolean, parse_skill_frontmatter, +}; diff --git a/src-tauri/src/skills/skill_loader.rs b/src-tauri/src/skills/skill_loader.rs new file mode 100644 index 000000000..ac5cab1ae --- /dev/null +++ b/src-tauri/src/skills/skill_loader.rs @@ -0,0 +1,231 @@ +//! Skill 定义加载器 +//! +//! 负责从 `~/.proxycast/skills//SKILL.md` 加载并解析 Skill 定义。 +//! 命令层只负责编排执行,不再持有文件解析细节。 + +use std::path::{Path, PathBuf}; + +use serde::{Deserialize, Serialize}; + +/// Skill 前置元数据 +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub(crate) struct SkillFrontmatter { + /// Skill 名称 + pub name: Option, + /// Skill 描述 + pub description: Option, + /// 允许的工具 + #[serde(rename = "allowed-tools")] + pub allowed_tools: Option, + /// 参数提示 + #[serde(rename = "argument-hint")] + pub argument_hint: Option, + /// 使用场景 + #[serde(rename = "when-to-use")] + pub when_to_use: Option, + /// 版本 + pub version: Option, + /// 偏好模型 + pub model: Option, + /// 偏好 Provider + pub provider: Option, + /// 是否禁用模型调用 + #[serde(rename = "disable-model-invocation")] + pub disable_model_invocation: Option, + /// 执行模式 + #[serde(rename = "execution-mode")] + pub execution_mode: Option, +} + +/// 内部 Skill 定义(用于加载和执行) +#[derive(Debug, Clone)] +pub(crate) struct LoadedSkillDefinition { + /// Skill 名称 + pub skill_name: String, + /// 显示名称 + pub display_name: String, + /// 描述 + pub description: String, + /// Markdown 内容(System Prompt) + pub markdown_content: String, + /// 允许的工具 + pub allowed_tools: Option>, + /// 参数提示 + pub argument_hint: Option, + /// 使用场景 + pub when_to_use: Option, + /// 偏好模型 + pub model: Option, + /// 偏好 Provider + pub provider: Option, + /// 是否禁用模型调用 + pub disable_model_invocation: bool, + /// 执行模式 + pub execution_mode: String, +} + +/// 解析 Skill 文件的 frontmatter +pub(crate) fn parse_skill_frontmatter(content: &str) -> (SkillFrontmatter, String) { + let regex = regex::Regex::new(r"^---\s*\n([\s\S]*?)---\s*\n?").unwrap(); + + if let Some(captures) = regex.captures(content) { + let frontmatter_text = captures.get(1).map(|m| m.as_str()).unwrap_or(""); + let body_start = captures.get(0).map(|m| m.end()).unwrap_or(0); + let body = content.get(body_start..).unwrap_or("").to_string(); + + let mut frontmatter = SkillFrontmatter::default(); + + for line in frontmatter_text.lines() { + if let Some(colon_idx) = line.find(':') { + let key = line.get(..colon_idx).unwrap_or("").trim(); + let value = line.get(colon_idx + 1..).unwrap_or("").trim(); + let clean_value = value + .trim_start_matches('"') + .trim_end_matches('"') + .trim_start_matches('\'') + .trim_end_matches('\'') + .to_string(); + + match key { + "name" => frontmatter.name = Some(clean_value), + "description" => frontmatter.description = Some(clean_value), + "allowed-tools" => frontmatter.allowed_tools = Some(clean_value), + "argument-hint" => frontmatter.argument_hint = Some(clean_value), + "when-to-use" | "when_to_use" => frontmatter.when_to_use = Some(clean_value), + "version" => frontmatter.version = Some(clean_value), + "model" => frontmatter.model = Some(clean_value), + "provider" => frontmatter.provider = Some(clean_value), + "disable-model-invocation" => { + frontmatter.disable_model_invocation = Some(clean_value) + } + "execution-mode" => frontmatter.execution_mode = Some(clean_value), + _ => {} + } + } + } + + (frontmatter, body) + } else { + (SkillFrontmatter::default(), content.to_string()) + } +} + +/// 解析 allowed-tools 字段 +pub(crate) fn parse_allowed_tools(value: Option<&str>) -> Option> { + value.and_then(|v| { + if v.is_empty() { + return None; + } + if v.contains(',') { + Some( + v.split(',') + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()) + .collect(), + ) + } else { + Some(vec![v.trim().to_string()]) + } + }) +} + +/// 解析布尔值字段 +pub(crate) fn parse_boolean(value: Option<&str>, default: bool) -> bool { + value + .map(|v| { + let lower = v.to_lowercase(); + matches!(lower.as_str(), "true" | "1" | "yes") + }) + .unwrap_or(default) +} + +/// 从文件加载 Skill 定义 +pub(crate) fn load_skill_from_file( + skill_name: &str, + file_path: &Path, +) -> Result { + let content = + std::fs::read_to_string(file_path).map_err(|e| format!("读取 Skill 文件失败: {}", e))?; + + let (frontmatter, markdown_content) = parse_skill_frontmatter(&content); + + let display_name = frontmatter + .name + .clone() + .unwrap_or_else(|| skill_name.to_string()); + let description = frontmatter.description.clone().unwrap_or_default(); + let allowed_tools = parse_allowed_tools(frontmatter.allowed_tools.as_deref()); + let disable_model_invocation = + parse_boolean(frontmatter.disable_model_invocation.as_deref(), false); + let execution_mode = frontmatter + .execution_mode + .clone() + .unwrap_or_else(|| "prompt".to_string()); + + Ok(LoadedSkillDefinition { + skill_name: skill_name.to_string(), + display_name, + description, + markdown_content, + allowed_tools, + argument_hint: frontmatter.argument_hint, + when_to_use: frontmatter.when_to_use, + model: frontmatter.model, + provider: frontmatter.provider, + disable_model_invocation, + execution_mode, + }) +} + +/// 获取 ProxyCast Skills 目录 +pub(crate) fn get_proxycast_skills_dir() -> Option { + dirs::home_dir().map(|home| home.join(".proxycast").join("skills")) +} + +/// 从目录加载所有 Skills +pub(crate) fn load_skills_from_directory(dir_path: &Path) -> Vec { + let mut results = Vec::new(); + + if !dir_path.exists() { + return results; + } + + if let Ok(entries) = std::fs::read_dir(dir_path) { + for entry in entries.flatten() { + let path = entry.path(); + if !path.is_dir() { + continue; + } + + let skill_file = path.join("SKILL.md"); + if skill_file.exists() { + let skill_name = path + .file_name() + .and_then(|n| n.to_str()) + .unwrap_or("unknown") + .to_string(); + + if let Ok(skill) = load_skill_from_file(&skill_name, &skill_file) { + results.push(skill); + } + } + } + } + + results +} + +/// 根据名称查找 Skill +pub(crate) fn find_skill_by_name(skill_name: &str) -> Result { + let skills_dir = + get_proxycast_skills_dir().ok_or_else(|| "无法获取 Skills 目录".to_string())?; + + let skill_path = skills_dir.join(skill_name); + let skill_file = skill_path.join("SKILL.md"); + + if !skill_file.exists() { + return Err(format!("Skill 不存在: {}", skill_name)); + } + + load_skill_from_file(skill_name, &skill_file) +} diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 7fb8a624a..84daeaa7a 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "ProxyCast", - "version": "0.57.0", + "version": "0.58.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src-tauri/tests/api_key_provider_tests.proptest-regressions b/src-tauri/tests/api_key_provider_tests.proptest-regressions new file mode 100644 index 000000000..53a9b07e7 --- /dev/null +++ b/src-tauri/tests/api_key_provider_tests.proptest-regressions @@ -0,0 +1,13 @@ +# 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 399023994da7d3a7d7407cbcc210e9c88f626c7158b717607e01a3f767d1e0b6 # shrinks to num_providers = 2 +cc 9598ca62cf55f54ee97f00786a0e6029df29a5ede07c2f6ccd731a92ac2f1d6e # shrinks to num_keys = 2 +cc 8ecfedd6b97ec094a400ca1af4c6c011f39a60688dd76327247ca8a54ca3240c # shrinks to num_keys = 2 +cc d9e6f7a966ae7126d118843e3c99009616930f30348a08dfedaeab933fe9877b # shrinks to num_errors = 1 +cc dfc5e61afb3ab4ec5b6283e3b92fa88ae5149458c321170000e222b95bd499e4 # shrinks to name = "aaa", api_host = "https://aaa.aa/" +cc b35b5acac2443a80f05fd96b8f46cd2b80e38a54e73b4a09aaaf5c3b68af319c # shrinks to api_key = "a0a0___a-0a-aA_-A---" +cc 05448979dc0877ad4bffe94f37f10f79ba6243d3289e0bf7629616ac8901c292 # shrinks to api_key = "-A0a_aa0-A0a_-a-_Aaa", alias = None diff --git a/src/App.tsx b/src/App.tsx index 06566b330..855319bc9 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -19,8 +19,10 @@ import { ProviderPoolPage } from "./components/provider-pool"; import { ToolsPage } from "./components/tools/ToolsPage"; import { AgentChatPage } from "./components/agent"; import { PluginsPage } from "./components/plugins/PluginsPage"; +import { McpPanel } from "./components/mcp"; import { ImageGenPage } from "./components/image-gen"; import { ProjectsPage } from "./components/projects"; +import { VibePage } from "./components/vibe/VibePage"; import { ProjectDetailPage } from "./components/projects/ProjectDetailPage"; import { CreateProjectDialog } from "./components/projects/CreateProjectDialog"; import { ProjectType } from "./lib/api/project"; @@ -347,11 +349,21 @@ function AppContent() { + {/* MCP 页面 */} + + + + {/* Plugins 页面 */} + {/* Vibe 页面 */} + + + + {/* Settings 页面 */} diff --git a/src/components/AppSidebar.tsx b/src/components/AppSidebar.tsx index 9f79acb4a..ada433f56 100644 --- a/src/components/AppSidebar.tsx +++ b/src/components/AppSidebar.tsx @@ -25,6 +25,8 @@ import { Terminal, Image, FolderKanban, + Blocks, + Sparkles, LucideIcon, } from "lucide-react"; import * as LucideIcons from "lucide-react"; @@ -128,6 +130,8 @@ const mainMenuItems: { id: Page; label: string; icon: typeof Bot }[] = [ { id: "image-gen", label: "图片生成", icon: Image }, { id: "api-server", label: "API Server", icon: Globe }, { id: "provider-pool", label: "凭证池", icon: Database }, + { id: "mcp", label: "MCP 服务器", icon: Blocks }, + { id: "vibe", label: "Vibe Zone", icon: Sparkles }, { id: "terminal", label: "终端", icon: Terminal }, { id: "tools", label: "工具", icon: Wrench }, { id: "plugins", label: "插件中心", icon: Puzzle }, @@ -151,6 +155,8 @@ const DEFAULT_ENABLED_NAV_ITEMS = [ "image-gen", "api-server", "provider-pool", + "mcp", + "vibe", ]; export function AppSidebar({ currentPage, onNavigate }: AppSidebarProps) { @@ -178,9 +184,19 @@ export function AppSidebar({ currentPage, onNavigate }: AppSidebarProps) { const loadNavConfig = async () => { try { const config = await getConfig(); - setEnabledNavItems( - config.navigation?.enabled_items || DEFAULT_ENABLED_NAV_ITEMS, - ); + const saved = config.navigation?.enabled_items; + if (saved && saved.length > 0) { + // 自动补充新增的默认导航项(避免新功能不可见) + const merged = [...saved]; + for (const item of DEFAULT_ENABLED_NAV_ITEMS) { + if (!merged.includes(item)) { + merged.push(item); + } + } + setEnabledNavItems(merged); + } else { + setEnabledNavItems(DEFAULT_ENABLED_NAV_ITEMS); + } } catch (error) { console.error("加载导航配置失败:", error); } diff --git a/src/components/README.md b/src/components/README.md index f9bb61b83..3ec562cd2 100644 --- a/src/components/README.md +++ b/src/components/README.md @@ -17,7 +17,7 @@ React 组件层,包含 UI 组件和业务组件。 - `extensions/` - 扩展功能组件 - `flow-monitor/` - LLM 流量监控组件 - `general-chat/` - 通用对话功能组件(三栏布局:会话列表 + 聊天区域 + 画布) -- `mcp/` - MCP 服务器管理组件 +- `mcp/` - MCP 服务器管理组件(配置管理、运行时控制、工具/提示词/资源浏览与调用) - `plugins/` - 插件管理组件 - `prompts/` - Prompt 管理组件 - `provider-pool/` - Provider 凭证池管理组件 diff --git a/src/components/agent/chat/components/ChatSidebar.tsx b/src/components/agent/chat/components/ChatSidebar.tsx index 785d3eb17..a786287e6 100644 --- a/src/components/agent/chat/components/ChatSidebar.tsx +++ b/src/components/agent/chat/components/ChatSidebar.tsx @@ -256,7 +256,7 @@ export const ChatSidebar: React.FC = ({ const loadSkills = async () => { setLoadingSkills(true); try { - const allSkills = await skillsApi.getAll("claude"); + const allSkills = await skillsApi.getAll("proxycast"); setSkills(allSkills); } catch (error) { console.error("加载技能列表失败:", error); @@ -273,7 +273,7 @@ export const ChatSidebar: React.FC = ({ const handleInstall = async (skill: Skill) => { setActionLoading(skill.directory); try { - const result = await skillsApi.install(skill.directory, "claude"); + const result = await skillsApi.install(skill.directory, "proxycast"); if (result) { toast.success(`已安装: ${skill.name}`); await loadSkills(); @@ -292,7 +292,7 @@ export const ChatSidebar: React.FC = ({ const handleUninstall = async (skill: Skill) => { setActionLoading(skill.directory); try { - const result = await skillsApi.uninstall(skill.directory, "claude"); + const result = await skillsApi.uninstall(skill.directory, "proxycast"); if (result) { toast.success(`已卸载: ${skill.name}`); await loadSkills(); diff --git a/src/components/agent/chat/hooks/skillCommand.ts b/src/components/agent/chat/hooks/skillCommand.ts new file mode 100644 index 000000000..779facaec --- /dev/null +++ b/src/components/agent/chat/hooks/skillCommand.ts @@ -0,0 +1,531 @@ +import type { Dispatch, SetStateAction } from "react"; +import type { UnlistenFn } from "@tauri-apps/api/event"; +import { safeListen } from "@/lib/dev-bridge"; +import { parseStreamEvent, type StreamEvent } from "@/lib/api/agent"; +import { + skillExecutionApi, + type ExecutableSkillInfo, +} from "@/lib/api/skill-execution"; +import type { ActionRequired, Message } from "../types"; + +/** 解析 /skill-name args 命令 */ +export interface ParsedSkillCommand { + skillName: string; + userInput: string; +} + +/** Slash Skill 执行上下文 */ +export interface SlashSkillExecutionContext { + command: ParsedSkillCommand; + rawContent: string; + assistantMsgId: string; + providerType: string; + model?: string; + ensureSession: () => Promise; + setMessages: Dispatch>; + setIsSending: (value: boolean) => void; + setCurrentAssistantMsgId: (id: string | null) => void; + setStreamUnlisten: (unlisten: UnlistenFn | null) => void; + playTypewriterSound: () => void; + playToolcallSound: () => void; + onWriteFile?: (content: string, fileName: string) => void; +} + +const VALID_ACTION_TYPES = new Set([ + "tool_confirmation", + "ask_user", + "elicitation", +]); + +/** + * 解析 slash skill 命令。 + * + * 格式:`/skill-name` 或 `/skill-name args...` + */ +export function parseSkillSlashCommand( + content: string, +): ParsedSkillCommand | null { + const skillMatch = content.match(/^\/([a-zA-Z0-9_-]+)\s*([\s\S]*)$/); + if (!skillMatch) { + return null; + } + + const [, skillName, userInput] = skillMatch; + return { + skillName, + userInput: userInput?.trim() || "", + }; +} + +function resolveSkillProviderOverride( + providerType: string, + model: string | undefined, +): { providerOverride?: string; modelOverride?: string } { + const normalizedProvider = providerType.toLowerCase().trim(); + + if (!normalizedProvider) { + return {}; + } + + return { + providerOverride: providerType, + modelOverride: model, + }; +} + +function normalizeActionType(actionType: string): ActionRequired["actionType"] { + if (VALID_ACTION_TYPES.has(actionType as ActionRequired["actionType"])) { + return actionType as ActionRequired["actionType"]; + } + return "tool_confirmation"; +} + +function appendTextPart( + messages: Message[], + assistantMsgId: string, + textDelta: string, +) { + return messages.map((msg) => { + if (msg.id !== assistantMsgId) return msg; + + const nextParts = [...(msg.contentParts || [])]; + const lastPart = nextParts[nextParts.length - 1]; + + if (lastPart && lastPart.type === "text") { + nextParts[nextParts.length - 1] = { + type: "text", + text: lastPart.text + textDelta, + }; + } else { + nextParts.push({ type: "text", text: textDelta }); + } + + return { + ...msg, + content: (msg.content || "") + textDelta, + isThinking: false, + thinkingContent: undefined, + contentParts: nextParts, + }; + }); +} + +function appendThinkingPart( + messages: Message[], + assistantMsgId: string, + textDelta: string, +) { + return messages.map((msg) => { + if (msg.id !== assistantMsgId) return msg; + + const nextParts = [...(msg.contentParts || [])]; + const lastPart = nextParts[nextParts.length - 1]; + + if (lastPart && lastPart.type === "thinking") { + nextParts[nextParts.length - 1] = { + type: "thinking", + text: lastPart.text + textDelta, + }; + } else { + nextParts.push({ type: "thinking", text: textDelta }); + } + + return { + ...msg, + isThinking: true, + thinkingContent: (msg.thinkingContent || "") + textDelta, + contentParts: nextParts, + }; + }); +} + +function tryHandleToolWriteFile( + toolName: string, + toolArguments: string | undefined, + onWriteFile?: (content: string, fileName: string) => void, +) { + if (!onWriteFile || !toolArguments) { + return; + } + + const normalizedToolName = toolName.toLowerCase(); + const looksLikeWriteTool = + normalizedToolName.includes("write") || + normalizedToolName.includes("create"); + + if (!looksLikeWriteTool) { + return; + } + + try { + const parsed = JSON.parse(toolArguments) as Record; + const filePath = + (typeof parsed.path === "string" ? parsed.path : undefined) || + (typeof parsed.file_path === "string" ? parsed.file_path : undefined) || + (typeof parsed.filePath === "string" ? parsed.filePath : undefined); + + const fileContent = + (typeof parsed.content === "string" ? parsed.content : undefined) || + (typeof parsed.text === "string" ? parsed.text : undefined); + + if (filePath && fileContent) { + onWriteFile(fileContent, filePath); + } + } catch (error) { + console.warn("[SkillCommand] 解析 tool_start 参数失败:", error); + } +} + +async function findMatchedSkill( + skillName: string, +): Promise { + try { + const skills = await skillExecutionApi.listExecutableSkills(); + return skills.find((skill) => skill.name === skillName) || null; + } catch (error) { + console.warn("[SkillCommand] 获取可执行 Skills 失败,回退普通对话:", error); + return null; + } +} + +/** + * 尝试执行 slash skill 命令。 + * + * @returns true 表示已处理(包括执行成功或执行失败);false 表示非 Skill 命令或未命中技能。 + */ +export async function tryExecuteSlashSkillCommand( + ctx: SlashSkillExecutionContext, +): Promise { + const { + command, + rawContent, + assistantMsgId, + providerType, + model, + ensureSession, + setMessages, + setIsSending, + setCurrentAssistantMsgId, + setStreamUnlisten, + playTypewriterSound, + playToolcallSound, + onWriteFile, + } = ctx; + + const matchedSkill = await findMatchedSkill(command.skillName); + if (!matchedSkill) { + return false; + } + + const activeSessionId = await ensureSession(); + if (!activeSessionId) { + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...msg, + content: "Skill 执行失败:无法创建会话", + isThinking: false, + thinkingContent: undefined, + contentParts: [ + { type: "text" as const, text: "Skill 执行失败:无法创建会话" }, + ], + } + : msg, + ), + ); + setIsSending(false); + setCurrentAssistantMsgId(null); + return true; + } + + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...msg, + isThinking: true, + thinkingContent: `正在执行 Skill: ${matchedSkill.display_name}...`, + content: "", + contentParts: [], + } + : msg, + ), + ); + + const streamCounters = { + text_delta: 0, + thinking_delta: 0, + tool_start: 0, + tool_end: 0, + done: 0, + final_done: 0, + error: 0, + }; + + let accumulatedContent = ""; + let skillUnlisten: UnlistenFn | null = null; + + const cleanup = () => { + if (skillUnlisten) { + skillUnlisten(); + skillUnlisten = null; + } + setStreamUnlisten(null); + setIsSending(false); + setCurrentAssistantMsgId(null); + }; + + try { + const eventName = `skill-exec-${assistantMsgId}`; + + skillUnlisten = await safeListen(eventName, ({ payload }) => { + const streamEvent = parseStreamEvent(payload as unknown); + if (!streamEvent) return; + + switch (streamEvent.type) { + case "text_delta": { + streamCounters.text_delta += 1; + accumulatedContent += streamEvent.text; + playTypewriterSound(); + setMessages((prev) => + appendTextPart(prev, assistantMsgId, streamEvent.text), + ); + break; + } + case "thinking_delta": { + streamCounters.thinking_delta += 1; + setMessages((prev) => + appendThinkingPart(prev, assistantMsgId, streamEvent.text), + ); + break; + } + case "tool_start": { + streamCounters.tool_start += 1; + playToolcallSound(); + + tryHandleToolWriteFile( + streamEvent.tool_name, + streamEvent.arguments, + onWriteFile, + ); + + const newToolCall = { + id: streamEvent.tool_id, + name: streamEvent.tool_name, + arguments: streamEvent.arguments, + status: "running" as const, + startTime: new Date(), + }; + + setMessages((prev) => + prev.map((msg) => { + if (msg.id !== assistantMsgId) return msg; + const existing = msg.toolCalls?.find( + (tc) => tc.id === streamEvent.tool_id, + ); + if (existing) return msg; + + return { + ...msg, + toolCalls: [...(msg.toolCalls || []), newToolCall], + contentParts: [ + ...(msg.contentParts || []), + { type: "tool_use" as const, toolCall: newToolCall }, + ], + }; + }), + ); + break; + } + case "tool_end": { + streamCounters.tool_end += 1; + setMessages((prev) => + prev.map((msg) => { + if (msg.id !== assistantMsgId) return msg; + + const updatedToolCalls = (msg.toolCalls || []).map((tc) => + tc.id === streamEvent.tool_id + ? { + ...tc, + status: streamEvent.result.success + ? ("completed" as const) + : ("failed" as const), + result: streamEvent.result, + endTime: new Date(), + } + : tc, + ); + + const updatedParts = (msg.contentParts || []).map((part) => { + if ( + part.type === "tool_use" && + part.toolCall.id === streamEvent.tool_id + ) { + return { + ...part, + toolCall: { + ...part.toolCall, + status: streamEvent.result.success + ? ("completed" as const) + : ("failed" as const), + result: streamEvent.result, + endTime: new Date(), + }, + }; + } + return part; + }); + + return { + ...msg, + toolCalls: updatedToolCalls, + contentParts: updatedParts, + }; + }), + ); + break; + } + case "action_required": { + const actionRequired: ActionRequired = { + requestId: streamEvent.request_id, + actionType: normalizeActionType(streamEvent.action_type), + toolName: streamEvent.tool_name, + arguments: streamEvent.arguments, + prompt: streamEvent.prompt, + questions: streamEvent.questions, + requestedSchema: streamEvent.requested_schema, + }; + + setMessages((prev) => + prev.map((msg) => { + if (msg.id !== assistantMsgId) return msg; + const existing = msg.actionRequests?.find( + (item) => item.requestId === streamEvent.request_id, + ); + if (existing) return msg; + + return { + ...msg, + actionRequests: [...(msg.actionRequests || []), actionRequired], + contentParts: [ + ...(msg.contentParts || []), + { type: "action_required" as const, actionRequired }, + ], + }; + }), + ); + break; + } + case "done": { + streamCounters.done += 1; + break; + } + case "final_done": { + streamCounters.final_done += 1; + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...msg, + isThinking: false, + thinkingContent: undefined, + content: accumulatedContent || msg.content, + } + : msg, + ), + ); + break; + } + case "error": { + streamCounters.error += 1; + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...msg, + isThinking: false, + thinkingContent: undefined, + content: + accumulatedContent || `错误: ${streamEvent.message}`, + } + : msg, + ), + ); + break; + } + } + }); + + setStreamUnlisten(skillUnlisten); + + const { providerOverride, modelOverride } = resolveSkillProviderOverride( + providerType, + model, + ); + + const result = await skillExecutionApi.executeSkill({ + skillName: command.skillName, + userInput: command.userInput || rawContent, + providerOverride, + modelOverride, + executionId: assistantMsgId, + sessionId: activeSessionId, + }); + + console.log( + `[SkillCommand] 执行完成: name=${command.skillName}, success=${result.success}, output_len=${result.output?.length ?? 0}, stream_stats=${JSON.stringify(streamCounters)}`, + ); + + const hasStreamedContent = accumulatedContent.trim().length > 0; + const finalContent = hasStreamedContent + ? accumulatedContent + : result.output || result.error || "Skill 执行完成"; + + setMessages((prev) => + prev.map((msg) => { + if (msg.id !== assistantMsgId) return msg; + + const nextParts = [...(msg.contentParts || [])]; + if (nextParts.length === 0 && finalContent) { + nextParts.push({ type: "text", text: finalContent }); + } + + return { + ...msg, + content: finalContent, + isThinking: false, + thinkingContent: undefined, + contentParts: nextParts, + }; + }), + ); + + cleanup(); + return true; + } catch (error) { + console.error(`[SkillCommand] 执行失败: ${command.skillName}`, error); + + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...msg, + isThinking: false, + thinkingContent: undefined, + content: `Skill 执行失败: ${error instanceof Error ? error.message : String(error)}`, + contentParts: [ + { + type: "text", + text: `Skill 执行失败: ${error instanceof Error ? error.message : String(error)}`, + }, + ], + } + : msg, + ), + ); + + cleanup(); + return true; + } +} diff --git a/src/components/agent/chat/hooks/skillSettings.ts b/src/components/agent/chat/hooks/skillSettings.ts new file mode 100644 index 000000000..d52b860da --- /dev/null +++ b/src/components/agent/chat/hooks/skillSettings.ts @@ -0,0 +1,98 @@ +/** + * Skill 执行配置 + * + * 管理 slash skill 执行时的 Provider 覆盖策略与工具兼容 Provider 列表。 + */ + +export type SkillProviderOverrideMode = + | "compatible_only" + | "always_current" + | "auto_fallback"; + +const STORAGE_KEY_MODE = "proxycast_skill_provider_override_mode"; +const STORAGE_KEY_PROVIDERS = "proxycast_skill_tool_compatible_providers"; + +export const DEFAULT_SKILL_TOOL_COMPATIBLE_PROVIDERS = [ + "anthropic", + "claude", + "claude_oauth", + "openai", + "gemini", + "kiro", + "antigravity", +]; + +export const DEFAULT_SKILL_PROVIDER_OVERRIDE_MODE: SkillProviderOverrideMode = + "compatible_only"; + +export function isSkillProviderOverrideMode( + value: string, +): value is SkillProviderOverrideMode { + return ( + value === "compatible_only" || + value === "always_current" || + value === "auto_fallback" + ); +} + +export function getSkillProviderOverrideMode(): SkillProviderOverrideMode { + const stored = localStorage.getItem(STORAGE_KEY_MODE); + if (stored && isSkillProviderOverrideMode(stored)) { + return stored; + } + return DEFAULT_SKILL_PROVIDER_OVERRIDE_MODE; +} + +export function setSkillProviderOverrideMode( + mode: SkillProviderOverrideMode, +): void { + localStorage.setItem(STORAGE_KEY_MODE, mode); +} + +function normalizeProviderList(providers: string[]): string[] { + const normalized = providers + .map((provider) => provider.toLowerCase().trim()) + .filter((provider) => provider.length > 0); + return Array.from(new Set(normalized)); +} + +export function getSkillToolCompatibleProviders(): string[] { + const raw = localStorage.getItem(STORAGE_KEY_PROVIDERS); + if (!raw) { + return DEFAULT_SKILL_TOOL_COMPATIBLE_PROVIDERS; + } + + try { + const parsed = JSON.parse(raw); + if (Array.isArray(parsed)) { + const normalized = normalizeProviderList( + parsed.filter((item): item is string => typeof item === "string"), + ); + if (normalized.length > 0) { + return normalized; + } + } + } catch { + const normalized = normalizeProviderList(raw.split(",")); + if (normalized.length > 0) { + return normalized; + } + } + + return DEFAULT_SKILL_TOOL_COMPATIBLE_PROVIDERS; +} + +export function setSkillToolCompatibleProviders(providers: string[]): void { + const normalized = normalizeProviderList(providers); + const finalProviders = + normalized.length > 0 + ? normalized + : DEFAULT_SKILL_TOOL_COMPATIBLE_PROVIDERS; + + localStorage.setItem(STORAGE_KEY_PROVIDERS, JSON.stringify(finalProviders)); +} + +export function resetSkillProviderSettings(): void { + setSkillProviderOverrideMode(DEFAULT_SKILL_PROVIDER_OVERRIDE_MODE); + setSkillToolCompatibleProviders(DEFAULT_SKILL_TOOL_COMPATIBLE_PROVIDERS); +} diff --git a/src/components/agent/chat/hooks/useAgentChat.ts b/src/components/agent/chat/hooks/useAgentChat.ts index 7aefe15a9..0507d2d26 100644 --- a/src/components/agent/chat/hooks/useAgentChat.ts +++ b/src/components/agent/chat/hooks/useAgentChat.ts @@ -32,6 +32,10 @@ import { type ProviderConfigMap, } from "../types"; import { useArtifactParser } from "@/lib/artifact/hooks/useArtifactParser"; +import { + parseSkillSlashCommand, + tryExecuteSlashSkillCommand, +} from "./skillCommand"; /** 话题(会话)信息 */ export interface Topic { @@ -479,13 +483,44 @@ export function useAgentChat(options: UseAgentChatOptions = {}) { // 保存当前消息 ID 到 ref,用于停止时更新状态 currentAssistantMsgIdRef.current = assistantMsgId; - // 初始化 Artifact 解析器,开始新的解析会话 - startArtifactParsing(); - // 用于累积流式内容 let accumulatedContent = ""; let unlisten: UnlistenFn | null = null; + // === Skill 拦截逻辑 === + // 检测 /skill-name args 格式的输入,直接调用 execute_skill 命令 + // 绕过 aster_agent_chat_stream 路径(某些 Provider 如 Codex 不支持工具调用) + const parsedSkillCommand = parseSkillSlashCommand(content); + if (parsedSkillCommand) { + const skillHandled = await tryExecuteSlashSkillCommand({ + command: parsedSkillCommand, + rawContent: content, + assistantMsgId, + providerType, + model: model || undefined, + ensureSession: _ensureSession, + setMessages, + setIsSending, + setCurrentAssistantMsgId: (id) => { + currentAssistantMsgIdRef.current = id; + }, + setStreamUnlisten: (unlistenFn) => { + unlistenRef.current = unlistenFn; + }, + playTypewriterSound, + playToolcallSound, + onWriteFile, + }); + + if (skillHandled) { + return; + } + } + // === Skill 拦截结束 === + + // 初始化 Artifact 解析器,开始新的解析会话 + startArtifactParsing(); + /** * 辅助函数:更新 contentParts,支持交错显示 * - text_delta: 追加到最后一个 text 类型,或创建新的 text 类型 diff --git a/src/components/agent/chat/hooks/useAsterAgentChat.ts b/src/components/agent/chat/hooks/useAsterAgentChat.ts index d0c7b31b9..feef425fd 100644 --- a/src/components/agent/chat/hooks/useAsterAgentChat.ts +++ b/src/components/agent/chat/hooks/useAsterAgentChat.ts @@ -113,16 +113,16 @@ const mapProviderName = (providerType: string): string => { // Google google: "google", gemini: "google", - // DeepSeek - deepseek: "custom_deepseek", - "deepseek-reasoner": "custom_deepseek", + // DeepSeek(OpenAI 兼容) + deepseek: "deepseek", + "deepseek-reasoner": "deepseek", // Ollama ollama: "ollama", // OpenRouter openrouter: "openrouter", - // 其他 - groq: "groq", - mistral: "mistral", + // 其他(OpenAI 兼容) + groq: "openai", + mistral: "openai", }; return mapping[providerType.toLowerCase()] || providerType; }; @@ -172,9 +172,9 @@ export function useAsterAgentChat(options: UseAsterAgentChatOptions = {}) { id: s.id, title: s.name || - `话题 ${new Date(s.created_at).toLocaleDateString("zh-CN")}`, - createdAt: new Date(s.created_at), - messagesCount: s.messages_count, + `话题 ${new Date(s.created_at * 1000).toLocaleDateString("zh-CN")}`, + createdAt: new Date(s.created_at * 1000), + messagesCount: s.messages_count ?? 0, })); setTopics(topicList); } catch (err) { @@ -192,9 +192,9 @@ export function useAsterAgentChat(options: UseAsterAgentChatOptions = {}) { id: s.id, title: s.name || - `话题 ${new Date(s.created_at).toLocaleDateString("zh-CN")}`, - createdAt: new Date(s.created_at), - messagesCount: s.messages_count, + `话题 ${new Date(s.created_at * 1000).toLocaleDateString("zh-CN")}`, + createdAt: new Date(s.created_at * 1000), + messagesCount: s.messages_count ?? 0, })); setTopics(topicList); } catch (error) { @@ -571,13 +571,29 @@ export function useAsterAgentChat(options: UseAsterAgentChatOptions = {}) { try { const detail = await getAsterSession(topicId); - const loadedMessages: Message[] = detail.messages.map((msg, index) => ({ - id: `${topicId}-${index}`, - role: msg.role as "user" | "assistant", - content: msg.content, - timestamp: new Date(msg.timestamp), - isThinking: false, - })); + const loadedMessages: Message[] = detail.messages.map((msg, index) => { + // 从 TauriMessageContent 数组中提取文本和 contentParts + const contentParts: ContentPart[] = []; + const textParts: string[] = []; + + for (const part of msg.content) { + if (part.type === "text" && part.text) { + textParts.push(part.text); + contentParts.push({ type: "text", text: part.text }); + } else if (part.type === "thinking" && part.text) { + contentParts.push({ type: "thinking", text: part.text }); + } + } + + return { + id: `${topicId}-${index}`, + role: msg.role as "user" | "assistant", + content: textParts.join("\n"), + contentParts: contentParts.length > 0 ? contentParts : undefined, + timestamp: new Date(msg.timestamp * 1000), + isThinking: false, + }; + }); setMessages(loadedMessages); setSessionId(topicId); diff --git a/src/components/api-server/ApiServerPage.tsx b/src/components/api-server/ApiServerPage.tsx index 86293b429..380639ecd 100644 --- a/src/components/api-server/ApiServerPage.tsx +++ b/src/components/api-server/ApiServerPage.tsx @@ -10,7 +10,6 @@ import { import * as Select from "@radix-ui/react-select"; import { invoke } from "@tauri-apps/api/core"; import { LogsTab } from "./LogsTab"; -import { RoutesTab } from "./RoutesTab"; import { ProviderIcon } from "@/icons/providers"; import { startServer, @@ -56,7 +55,7 @@ interface TestState { httpStatus?: number; } -type TabId = "server" | "routes" | "logs"; +type TabId = "server" | "logs"; // Provider 到 API 类型的映射 type ApiType = "openai" | "anthropic" | "gemini"; @@ -67,6 +66,7 @@ const getProviderApiType = (provider: string): ApiType => { // OpenAI 兼容类型 if ( p === "codex" || + p === "codex_oauth" || p === "openai" || p === "openai-response" || p === "azure_openai" || @@ -108,6 +108,7 @@ const ALIAS_PROVIDERS = [ "antigravity", "kiro", "codex", + "codex_oauth", "gemini", "gemini_api_key", ]; @@ -115,6 +116,7 @@ const ALIAS_PROVIDERS = [ // 别名配置文件名映射(某些 Provider 共享同一个别名配置) const ALIAS_CONFIG_MAPPING: Record = { gemini_api_key: "gemini", + codex_oauth: "codex", }; // 可用的 Provider 信息(合并 OAuth 凭证池和 API Key Provider) @@ -518,7 +520,8 @@ export function ApiServerPage() { qwen: "Qwen", antigravity: "Antigravity", claude: "Claude", - codex: "Codex", + codex: "Codex API", + codex_oauth: "Codex OAuth", iflow: "iFlow", claude_oauth: "Claude OAuth", vertex: "Vertex AI", @@ -540,6 +543,7 @@ export function ApiServerPage() { antigravity: "gemini", claude: "claude", codex: "openai", + codex_oauth: "openai", iflow: "iflow", claude_oauth: "claude", vertex: "gemini", @@ -615,6 +619,7 @@ export function ApiServerPage() { }; // 合并 OAuth 凭证池和 API Key Provider,生成可用 Provider 列表 + // 注意:Codex 需要特殊处理,OAuth 和 API Key 分开显示 const buildAvailableProviders = () => { const providerMap = new Map(); @@ -624,7 +629,11 @@ export function ApiServerPage() { (c) => !c.is_disabled, ); if (enabledCredentials.length > 0) { - const id = overview.provider_type; + // Codex OAuth 使用特殊 ID,与 API Key 分开 + const id = + overview.provider_type === "codex" + ? "codex_oauth" + : overview.provider_type; const existing = providerMap.get(id); if (existing) { existing.oauthCount = enabledCredentials.length; @@ -634,10 +643,16 @@ export function ApiServerPage() { ? "both" : "oauth"; } else { + // Codex OAuth 使用特殊标签 + const label = + overview.provider_type === "codex" + ? "Codex OAuth" + : providerLabels[overview.provider_type] || + overview.provider_type; providerMap.set(id, { id, - label: providerLabels[id] || id, - iconType: providerIconMap[id] || "openai", + label, + iconType: providerIconMap[overview.provider_type] || "openai", source: "oauth", oauthCount: enabledCredentials.length, apiKeyCount: 0, @@ -751,7 +766,10 @@ export function ApiServerPage() { const handleSetDefaultProvider = async (providerId: string) => { try { - await setDefaultProvider(providerId); + // codex_oauth 在后端映射到 codex 凭证池 + const backendProviderId = + providerId === "codex_oauth" ? "codex" : providerId; + await setDefaultProvider(backendProviderId); setDefaultProviderState(providerId); // 获取最新的凭证池数据 @@ -1066,7 +1084,6 @@ export function ApiServerPage() {
{[ { id: "server" as TabId, name: "服务器控制" }, - { id: "routes" as TabId, name: "路由端点" }, { id: "logs" as TabId, name: "系统日志" }, ].map((tab) => (
)} - {/* Routes Tab */} - {activeTab === "routes" && } - {/* Logs Tab */} {activeTab === "logs" && } diff --git a/src/components/api-server/RoutesTab.tsx b/src/components/api-server/RoutesTab.tsx deleted file mode 100644 index 6cb104ec9..000000000 --- a/src/components/api-server/RoutesTab.tsx +++ /dev/null @@ -1,302 +0,0 @@ -import { useState, useEffect } from "react"; -import { Copy, Check, RefreshCw, Globe, Server, Tag } from "lucide-react"; -import { - routesApi, - RouteInfo, - RouteListResponse, - CurlExample, -} from "@/lib/api/routes"; - -export function RoutesTab() { - const [routes, setRoutes] = useState(null); - const [loading, setLoading] = useState(false); - const [error, setError] = useState(null); - const [expandedRoute, setExpandedRoute] = useState(null); - const [curlExamples, setCurlExamples] = useState< - Record - >({}); - const [copiedUrl, setCopiedUrl] = useState(null); - const [copiedCmd, setCopiedCmd] = useState(null); - - const fetchRoutes = async () => { - setLoading(true); - setError(null); - try { - const data = await routesApi.getAvailableRoutes(); - setRoutes(data); - } catch (e) { - setError(e instanceof Error ? e.message : String(e)); - } - setLoading(false); - }; - - useEffect(() => { - fetchRoutes(); - }, []); - - const fetchCurlExamples = async (selector: string) => { - if (curlExamples[selector]) return; - try { - const examples = await routesApi.getCurlExamples(selector); - setCurlExamples((prev) => ({ ...prev, [selector]: examples })); - } catch (e) { - console.error("Failed to fetch curl examples:", e); - } - }; - - const handleExpand = (selector: string) => { - if (expandedRoute === selector) { - setExpandedRoute(null); - } else { - setExpandedRoute(selector); - fetchCurlExamples(selector); - } - }; - - const copyToClipboard = (text: string, type: "url" | "cmd", id: string) => { - navigator.clipboard.writeText(text); - if (type === "url") { - setCopiedUrl(id); - setTimeout(() => setCopiedUrl(null), 2000); - } else { - setCopiedCmd(id); - setTimeout(() => setCopiedCmd(null), 2000); - } - }; - - const getProviderColor = (provider: string) => { - switch (provider) { - case "kiro": - return "bg-purple-100 text-purple-700 dark:bg-purple-900/30 dark:text-purple-400"; - case "gemini": - return "bg-blue-100 text-blue-700 dark:bg-blue-900/30 dark:text-blue-400"; - case "qwen": - return "bg-orange-100 text-orange-700 dark:bg-orange-900/30 dark:text-orange-400"; - case "openai": - return "bg-green-100 text-green-700 dark:bg-green-900/30 dark:text-green-400"; - case "claude": - return "bg-amber-100 text-amber-700 dark:bg-amber-900/30 dark:text-amber-400"; - case "antigravity": - return "bg-cyan-100 text-cyan-700 dark:bg-cyan-900/30 dark:text-cyan-400"; - default: - return "bg-gray-100 text-gray-700 dark:bg-gray-900/30 dark:text-gray-400"; - } - }; - - if (loading && !routes) { - return ( -
- -
- ); - } - - return ( -
-
-
-

可用路由端点

-

- 通过不同的 URL 路径访问不同的 Provider -

-
- -
- - {error && ( -
- {error} -
- )} - - {routes && ( -
- {/* Base URL Info */} -
-
- - 服务器地址: - - {routes.base_url} - -
-
- - {/* Routes List */} -
- {routes.routes.map((route) => ( - handleExpand(route.selector)} - curlExamples={curlExamples[route.selector]} - copiedUrl={copiedUrl} - copiedCmd={copiedCmd} - onCopyUrl={(url, id) => copyToClipboard(url, "url", id)} - onCopyCmd={(cmd, id) => copyToClipboard(cmd, "cmd", id)} - getProviderColor={getProviderColor} - /> - ))} -
-
- )} -
- ); -} - -interface RouteCardProps { - route: RouteInfo; - expanded: boolean; - onExpand: () => void; - curlExamples?: CurlExample[]; - copiedUrl: string | null; - copiedCmd: string | null; - onCopyUrl: (url: string, id: string) => void; - onCopyCmd: (cmd: string, id: string) => void; - getProviderColor: (provider: string) => string; -} - -function RouteCard({ - route, - expanded, - onExpand, - curlExamples, - copiedUrl, - copiedCmd, - onCopyUrl, - onCopyCmd, - getProviderColor, -}: RouteCardProps) { - return ( -
- {/* Header */} -
-
- -
-
- {route.selector} - - {route.provider_type} - - {route.tags.map((tag) => ( - - - {tag} - - ))} -
-
- {route.credential_count} 个凭证 - {!route.enabled && ( - (已禁用) - )} -
-
-
-
- {expanded ? "收起" : "展开"} -
-
- - {/* Expanded Content */} - {expanded && ( -
- {/* Endpoints */} -
-

端点地址

-
- {route.endpoints.map((endpoint, idx) => ( -
-
- - {endpoint.protocol.toUpperCase()} - - {endpoint.url} -
- -
- ))} -
-
- - {/* Curl Examples */} - {curlExamples && curlExamples.length > 0 && ( -
-

curl 示例

-
- {curlExamples.map((example, idx) => ( -
-
- - {example.description} - - -
-
-                      {example.command}
-                    
-
- ))} -
-
- )} -
- )} -
- ); -} diff --git a/src/components/content-creator/canvas/canvasUtils.test.ts b/src/components/content-creator/canvas/canvasUtils.test.ts index f470624fa..85f273179 100644 --- a/src/components/content-creator/canvas/canvasUtils.test.ts +++ b/src/components/content-creator/canvas/canvasUtils.test.ts @@ -24,9 +24,10 @@ describe("getCanvasTypeForTheme", () => { expect(getCanvasTypeForTheme("music")).toBe("music"); expect(getCanvasTypeForTheme("social-media")).toBe("document"); expect(getCanvasTypeForTheme("document")).toBe("document"); - expect(getCanvasTypeForTheme("general")).toBeNull(); - expect(getCanvasTypeForTheme("knowledge")).toBeNull(); - expect(getCanvasTypeForTheme("planning")).toBeNull(); + // 所有主题现在都支持 document 画布 + expect(getCanvasTypeForTheme("general")).toBe("document"); + expect(getCanvasTypeForTheme("knowledge")).toBe("document"); + expect(getCanvasTypeForTheme("planning")).toBe("document"); }); it("应该覆盖所有 ThemeType", () => { @@ -61,12 +62,13 @@ describe("isCanvasSupported", () => { expect(isCanvasSupported("music")).toBe(true); expect(isCanvasSupported("social-media")).toBe(true); expect(isCanvasSupported("document")).toBe(true); - expect(isCanvasSupported("general")).toBe(false); - expect(isCanvasSupported("knowledge")).toBe(false); - expect(isCanvasSupported("planning")).toBe(false); + // 所有主题现在都支持画布 + expect(isCanvasSupported("general")).toBe(true); + expect(isCanvasSupported("knowledge")).toBe(true); + expect(isCanvasSupported("planning")).toBe(true); }); - it("支持画布的主题数量应该是 6 种", () => { + it("所有 9 种主题都应该支持画布", () => { const allThemes: ThemeType[] = [ "general", "social-media", @@ -82,7 +84,7 @@ describe("isCanvasSupported", () => { const supportedCount = allThemes.filter((theme) => isCanvasSupported(theme), ).length; - expect(supportedCount).toBe(6); + expect(supportedCount).toBe(9); }); }); @@ -117,10 +119,19 @@ describe("createInitialCanvasState", () => { expect(socialState?.type).toBe("document"); }); - it("不支持画布的主题应该返回 null", () => { - expect(createInitialCanvasState("general", "test")).toBeNull(); - expect(createInitialCanvasState("knowledge", "test")).toBeNull(); - expect(createInitialCanvasState("planning", "test")).toBeNull(); + it("所有主题都应该返回有效的画布状态", () => { + // general、knowledge、planning 现在也支持 document 画布 + const generalState = createInitialCanvasState("general", "test"); + expect(generalState).not.toBeNull(); + expect(generalState?.type).toBe("document"); + + const knowledgeState = createInitialCanvasState("knowledge", "test"); + expect(knowledgeState).not.toBeNull(); + expect(knowledgeState?.type).toBe("document"); + + const planningState = createInitialCanvasState("planning", "test"); + expect(planningState).not.toBeNull(); + expect(planningState?.type).toBe("document"); }); it("应该正确处理空内容参数", () => { diff --git a/src/components/mcp/McpPage.tsx b/src/components/mcp/McpPage.tsx index 70f2fb845..931ad6c9c 100644 --- a/src/components/mcp/McpPage.tsx +++ b/src/components/mcp/McpPage.tsx @@ -85,6 +85,7 @@ export function McpPage({ hideHeader = false }: McpPageProps) { const [editName, setEditName] = useState(""); const [editDescription, setEditDescription] = useState(""); const [editConfig, setEditConfig] = useState(""); + const [enabledProxycast, setEnabledProxycast] = useState(true); const [enabledClaude, setEnabledClaude] = useState(true); const [enabledCodex, setEnabledCodex] = useState(true); const [enabledGemini, setEnabledGemini] = useState(true); @@ -125,6 +126,7 @@ export function McpPage({ hideHeader = false }: McpPageProps) { setEditName(server.name); setEditDescription(server.description || ""); setEditConfig(JSON.stringify(server.server_config, null, 2)); + setEnabledProxycast(server.enabled_proxycast); setEnabledClaude(server.enabled_claude); setEnabledCodex(server.enabled_codex); setEnabledGemini(server.enabled_gemini); @@ -138,6 +140,7 @@ export function McpPage({ hideHeader = false }: McpPageProps) { setEditName(""); setEditDescription(""); setEditConfig(defaultServerConfig); + setEnabledProxycast(true); setEnabledClaude(true); setEnabledCodex(true); setEnabledGemini(true); @@ -189,7 +192,7 @@ export function McpPage({ hideHeader = false }: McpPageProps) { name: editName.trim(), description: editDescription.trim() || undefined, server_config: serverConfig, - enabled_proxycast: false, + enabled_proxycast: enabledProxycast, enabled_claude: enabledClaude, enabled_codex: enabledCodex, enabled_gemini: enabledGemini, @@ -202,6 +205,7 @@ export function McpPage({ hideHeader = false }: McpPageProps) { name: editName.trim(), description: editDescription.trim() || undefined, server_config: serverConfig, + enabled_proxycast: enabledProxycast, enabled_claude: enabledClaude, enabled_codex: enabledCodex, enabled_gemini: enabledGemini, @@ -230,6 +234,7 @@ export function McpPage({ hideHeader = false }: McpPageProps) { // 获取启用的应用标签 const getEnabledApps = (server: McpServer) => { const apps: string[] = []; + if (server.enabled_proxycast) apps.push("ProxyCast"); if (server.enabled_claude) apps.push("Claude"); if (server.enabled_codex) apps.push("Codex"); if (server.enabled_gemini) apps.push("Gemini"); @@ -541,6 +546,15 @@ export function McpPage({ hideHeader = false }: McpPageProps) { 同步到: +