diff --git a/.gitignore b/.gitignore index 459ebc280..f990e0e79 100644 --- a/.gitignore +++ b/.gitignore @@ -3,6 +3,7 @@ node_modules/ # Build dist/ +.cargo/ src-tauri/.cargo/ *.exe *.pdb diff --git a/.tmp_perm_test b/.tmp_perm_test deleted file mode 100644 index 9daeafb98..000000000 --- a/.tmp_perm_test +++ /dev/null @@ -1 +0,0 @@ -test diff --git a/AGENTS.md b/AGENTS.md index a5445898c..aa9d529a8 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -24,6 +24,13 @@ 3. **谨慎新增子目录 AGENTS.md** - 仅当某个目录树存在长期稳定、只对该子树生效的规则时才新增;临时排障说明不要新增 AGENTS 4. **优先索引化而不是堆叠说明** - 根 AGENTS 更适合作为目录与约定入口,详细上下文交给专门文档 +## UI 全局指导 + +1. **界面改动先看视觉规范** - 涉及配色、渐变、卡片布局、设置页重排、工作台改版时,先读 `docs/aiprompts/design-language.md` +2. **宽度按页面类型选** - 表单页保持窄阅读宽度,卡片/工作台页面使用更宽的自适应内容区,不要整仓统一 `max-width` +3. **中文排版优先** - 避免过大英文 tracking、重复标题和挤压式统计卡文案 +4. **渐变只做氛围层** - 禁止用互相打架的多层渐变制造分割感,背景存在感必须弱于内容 + ## 详细文档 模块级详细文档位于 `docs/aiprompts/`: @@ -36,6 +43,7 @@ | [converter.md](docs/aiprompts/converter.md) | 协议转换 | | [server.md](docs/aiprompts/server.md) | HTTP 服务器 | | [components.md](docs/aiprompts/components.md) | 组件系统 | +| [design-language.md](docs/aiprompts/design-language.md) | 全局 UI 视觉语言 | | [hooks.md](docs/aiprompts/hooks.md) | React Hooks | | [services.md](docs/aiprompts/services.md) | 业务服务 | | [commands.md](docs/aiprompts/commands.md) | Tauri 命令 | @@ -93,6 +101,10 @@ npm run lint - 需要继续浏览器 E2E、复用现有 Playwright MCP 会话、排查 DevBridge/console 错误时,先读 `docs/aiprompts/playwright-e2e.md` - 如果只是仓库级规则,不要继续往本文件堆叠步骤说明 +## UI 设计入口 + +- 需要统一配色、修正渐变、调整页面宽度策略、重排卡片工作台时,先读 `docs/aiprompts/design-language.md` + ## 项目架构 ### 技术栈 diff --git a/RELEASE_NOTES.md b/RELEASE_NOTES.md index 187e56ba8..1d33e77a9 100644 --- a/RELEASE_NOTES.md +++ b/RELEASE_NOTES.md @@ -1,27 +1,59 @@ -## ProxyCast v0.85.0 +## ProxyCast v0.86.0 ### ✨ 新功能 -- 集成 AI 摘要到 SessionContextService,实现上下文智能管理 (621ab3ee) -- 新增 AI 摘要服务,用于上下文管理的 P0 阶段 1 实现 (75309bd2) -- 新增 Agent Timeline 服务,支持时间线视图 -- 新增 Chat History 服务,统一聊天历史管理 -- 新增多个 Agent 聊天相关组件(AgentPlanBlock、AgentRuntimeStrip、AgentThreadTimeline、SocialMediaHarnessCard) -- 新增社交媒体 Harness 工具集成 + +- **Aster 集成升级**: 升级到 aster-rust v0.17.1,带来更稳定的 Agent 运行时支持 +- **排队机制**: 新增 Agent 运行时排队系统,支持多轮对话请求的有序处理 +- **Artifact 自动预览**: 实现 Artifact 自动预览同步机制,提升内容创作体验 +- **搜索结果预览**: 新增搜索结果预览列表组件,优化 Web 搜索交互 +- **Skill 脚手架**: 新增 Skill 创建脚手架对话框,简化自定义 Skill 开发流程 +- **会话作用域存储**: 实现 Agent 会话级别的状态管理机制 ### 🔧 优化与重构 -- 移除 general-chat 相关的遗留代码和组件,完成向统一 Agent 系统的迁移 -- 清理 compat 兼容层代码(agentCompat、generalChatCompat) -- 重构数据库迁移结构,新增 migration_support 和 startup_migrations -- 优化 Agent 聊天 Hooks 架构,拆分为多个专职模块(agentChatActionState、agentChatCoreUtils、agentChatHistory 等) -- 完善测试覆盖率,新增 40+ 单元测试文件 -- 优化 Artifact 渲染器,新增 DocumentRenderer -- 优化内容创作工作流,新增社交媒体 Harness 测试 -### 📦 其他 -- 更新多个 AI 提示词文档 -- 新增 report-legacy-surfaces.mjs 脚本 -- 更新 ESLint 配置 +- **Request Tool Policy 重构**: 大幅重构请求工具策略模块(+1176 行),提升 Web 搜索预调用的可靠性 +- **Skill 服务增强**: 重构 Skill 加载器和匹配器,优化 Skill 发现和执行流程(+1064 行) +- **API Key Provider 优化**: 增强 API Key 提供商服务,改进凭证池管理(+433 行) +- **Skill Model 完善**: 新增 Skill 模型定义,规范化 Skill 元数据管理(+345 行) +- **Aster State 扩展**: 扩展 Agent 状态管理,支持更复杂的运行时场景(+367 行) +- **数据库 Schema 更新**: 新增 agent_runtime_queue 表,支持排队机制持久化 + +### 🐛 修复 + +- **Clippy 警告修复**: 修复多个 Rust clippy 警告,提升代码质量 + - 简化 `let...else` 为 `?` 操作符 + - 合并连续的字符串替换操作 + - 优化字符匹配模式 +- **类型安全改进**: 修复前端类型定义,增强 TypeScript 类型检查 + +### 📦 依赖更新 + +- 升级 aster-rust 到 v0.17.1 +- 更新相关 Rust 依赖包版本 + +### 🧪 测试 + +- 新增多个组件单元测试 + - Modal 组件测试 + - SearchResultPreviewList 测试 + - ArtifactRenderer UI 测试 + - SkillScaffoldDialog 测试 + - useArtifactAutoPreviewSync 测试 + - useArtifactDisplayState 测试 + - searchQueryGrouping 测试 + +### 📝 文档 + +- 新增设计语言文档 (design-language.md) +- 更新 AGENTS.md,补充跨平台兼容约束和 UI 指导 +- 完善 aiprompts 文档索引 + +### 🛠️ 开发体验 + +- 新增本地 aster-rust 覆盖脚本 (setup-local-aster-override.mjs) +- 改进 legacy surfaces 报告脚本 +- 优化开发工具链配置 --- -**完整变更**: v0.84.0...v0.85.0 +**完整变更**: v0.85.0...v0.86.0 diff --git a/docs/README.md b/docs/README.md index ea19815f3..fe12bf917 100644 --- a/docs/README.md +++ b/docs/README.md @@ -24,6 +24,7 @@ - `develop/execution-tracker-p0-acceptance-report.md`:统一执行追踪 P0 验收报告 - `develop/execution-tracker-p1-p2-roadmap.md`:统一执行追踪后续路线(P1/P2) - `develop/scheduler-task-governance-p1.md`:调度任务治理 P1(连续失败、自动停用、冷却恢复) +- `roadmap/proxycast-skills-standardization-roadmap.md`:Skills 标准化与产品化路线图 - `ops.md`:运维与发布说明 - `app.config.ts` / `nuxt.config.ts` / `package.json`:文档站配置 diff --git a/docs/aiprompts/README.md b/docs/aiprompts/README.md index 063c29321..9f8d60125 100644 --- a/docs/aiprompts/README.md +++ b/docs/aiprompts/README.md @@ -21,6 +21,7 @@ AI Agent 专用文档目录,提供模块级别的详细说明。 ### 前端模块 - `components.md` - React 组件系统 +- `design-language.md` - 全局 UI 视觉语言(配色、渐变、排版、宽度策略) - `hooks.md` - 自定义 React Hooks - `lib.md` - 工具库和 API 封装 @@ -54,6 +55,9 @@ AI Agent 在处理特定模块时,应先阅读对应的 aiprompts 文档: # 处理 Provider 相关任务 → 先读 docs/aiprompts/providers.md +# 处理 UI 配色、页面重排、视觉统一 +→ 先读 docs/aiprompts/design-language.md + # 处理新旧并存、迁移、重构、架构收口 → 先读 docs/aiprompts/governance.md → 再执行 npm run governance:legacy-report diff --git a/docs/aiprompts/design-language.md b/docs/aiprompts/design-language.md new file mode 100644 index 000000000..9ef7e1798 --- /dev/null +++ b/docs/aiprompts/design-language.md @@ -0,0 +1,218 @@ +# ProxyCast UI 视觉语言 + +## 适用范围 + +本规范适用于 ProxyCast 全部前端界面,包括: + +- 设置页 +- 功能工作台 +- 卡片型列表页 +- 弹窗、面板、侧边栏 +- 新增页面与既有页面重构 + +目标不是“做得更花”,而是统一全局视觉判断标准,减少页面之间的配色漂移、布局割裂和局部过度设计。 + +## 视觉目标 + +ProxyCast 的整体界面应当接近以下气质: + +- 轻盈、清晰、专业,而不是炫技 +- 信息优先,装饰只服务层级 +- 桌面应用感强,避免网页营销风格 +- 用少量颜色建立辨识度,而不是靠大面积饱和色 +- 大多数页面默认明亮、安静、可长时间使用 + +## 配色原则 + +### 1. 基础底色 + +- 主背景以浅白、浅灰、浅青为主 +- 常用基底建议接近: + - `slate-50` + - `white` + - `emerald-50` + - `sky-50` +- 不要默认使用纯白大平面配纯黑文本;应保留轻微层次 + +### 2. 主强调色 + +- 主操作优先使用深色实心按钮,而不是高饱和彩色按钮 +- 当前推荐主强调色方向: + - 深海军蓝 / 深石板蓝:主按钮、关键操作 + - 柔和绿色:成功、可用、健康状态 + - 天空蓝:辅助信息、候选状态 + - 琥珀色:提醒、兼容、注意事项 + +### 3. 状态色映射 + +- 成功 / 可用:`emerald` +- 提醒 / 兼容:`amber` +- 信息 / 次级操作:`sky` 或 `slate` +- 错误 / 危险:`rose` 或 `red` + +不要出现同一语义在不同页面频繁换色。 + +## 渐变与氛围层 + +### 1. 渐变使用规则 + +- 渐变只能作为“气氛层”,不能承担信息表达 +- 优先使用单一连续渐变,不要做左右两段式硬切 +- 渐变对比度必须低,不能影响文字可读性 + +推荐做法: + +- 一个大底渐变 +- 叠加 2 到 3 个低透明模糊色块 + +避免做法: + +- 两层方向冲突的渐变叠加 +- 右侧单独再盖一层灰蓝或灰黑渐变 +- 大面积高饱和紫色、蓝紫色背景 + +### 2. 背景层级 + +- 背景先轻,再让卡片浮出来 +- 如果背景已经有气氛层,卡片本身就要更克制 +- 当页面出现“背景比内容更显眼”的情况,优先减背景,不要继续加组件装饰 + +## 容器与布局 + +### 1. 内容宽度策略 + +根据页面类型决定宽度,不要所有页面统一 `max-width`: + +- 表单型设置页:约 `800px` +- 信息面板页:约 `960px` 到 `1200px` +- 卡片工作台 / 列表工作台:约 `1200px` 到 `1440px` + +原则: + +- 不盲目铺满 +- 也不要把卡片型页面锁死在窄列中 +- 应根据容器尺寸自适应,而不是写死单一宽度 + +### 2. 信息区块结构 + +推荐页面结构: + +1. 页面标题区 +2. 摘要区 / 工作台区 +3. 搜索筛选区 +4. 主内容区 + +每一层都要有独立职责,不要把标题、筛选、统计、主操作全挤进一行。 + +### 3. 卡片型列表页 + +- 卡片页优先考虑桌面端横向效率 +- 列数根据容器宽度自适应 +- 卡片之间的间距要比普通表单更大 +- 同一行卡片高度尽量对齐 + +## 文字排版 + +### 1. 中文优先 + +- 中文标题不要使用过大的字间距 +- 英文标签可以少量增加 tracking,但只能用于辅助标签 +- 不要让英文全大写成为页面主标题的唯一信息承载 + +### 2. 层级规则 + +- 一个页面只保留一个主标题中心 +- 子页面不要重复出现“标题 + 同标题卡片标题”双重堆叠 +- 标题负责定性,说明负责补充,不要写两句意思相同的话 + +### 3. 统计卡文字 + +统计卡推荐结构: + +- 图标 +- 指标名 +- 简短说明 +- 大数字 + +避免: + +- 把中文标题挤成竖排感 +- 说明文字过长导致卡片内断裂 +- 英文 tracking 过大压过中文阅读节奏 + +## 卡片与表面 + +### 1. 卡片表面 + +- 卡片底色优先白色或轻微染色白 +- 边框比阴影更重要 +- 阴影要浅,重点靠层级和留白,而不是重投影 + +推荐方向: + +- `bg-white` 或 `bg-white/90` +- `border-slate-200/80` +- `shadow-sm shadow-slate-950/5` + +### 2. 组件圆角 + +- 工作台容器:大圆角 +- 卡片:中大圆角 +- 小型状态标签:胶囊圆角 + +不要在同一页面混用太多种圆角尺度。 + +## 交互元素 + +### 1. 按钮 + +- 主按钮:深色实心 +- 次按钮:白底描边 +- 状态按钮:按语义色做浅底+描边 + +避免所有按钮都长得一样,也避免每个按钮一个颜色体系。 + +### 2. 筛选器 + +- 筛选器应独立成区,不要贴在大标题下直接散排 +- 激活态和未激活态差异必须明确 +- 数量徽标可以保留,但颜色不要喧宾夺主 + +## 跨页面一致性要求 + +涉及以下改动时,应优先遵守本规范: + +- 页面重排 +- 新增工作台 +- 设置页视觉升级 +- 卡片列表改版 +- 主题与配色统一 + +如果已有页面明显更成熟,应向成熟页面靠拢,而不是重新发明一套风格。 + +## 实施清单 + +改 UI 前先检查: + +1. 这是表单页、信息页还是卡片页? +2. 当前页面是否被错误限制在过窄宽度? +3. 是否存在重复标题或重复说明? +4. 背景是否比内容更显眼? +5. 中文排版是否被英文 tracking 或窄卡片破坏? +6. 按钮层级是否清晰? +7. 状态色是否遵守全局映射? + +## 本次沉淀的直接经验 + +这次 `skills` 页面调整形成了以下可复用结论: + +- 卡片工作台页不应沿用普通设置页的 `800px` 窄列宽度 +- 中文统计卡要优先保证标题和说明的横向可读性 +- 氛围背景应采用连续浅渐变,不要做分段叠色 +- 工作台主操作应与统计信息分栏,而不是挤成一行 +- 分组标题应使用“中文主标题 + 英文辅助标签”的组合,而不是反过来 + +## 关联文档 + +- [components.md](components.md) +- [overview.md](overview.md) diff --git a/package.json b/package.json index e250f046d..4d1ea1dba 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.85.0", + "version": "0.86.0", "type": "module", "engines": { "node": ">=22.0.0" @@ -13,7 +13,7 @@ "homepage": "https://github.com/aiclientproxy/proxycast", "scripts": { "predev": "npm run verify:app-version && node scripts/ensure-dev-port.mjs", - "dev": "npx vite", + "dev": "vite", "build": "npm run verify:app-version && tsc && vite build", "preview": "vite preview", "tauri": "tauri", @@ -38,7 +38,8 @@ "bridge:health": "node scripts/check-dev-bridge-health.mjs", "smoke:social-workbench": "node scripts/social-workbench-e2e-smoke.mjs", "dev:web-bridge": "node scripts/start-web-bridge-dev.mjs", - "governance:legacy-report": "node scripts/report-legacy-surfaces.mjs" + "governance:legacy-report": "node scripts/report-legacy-surfaces.mjs", + "setup:local-aster": "node scripts/setup-local-aster-override.mjs" }, "dependencies": { "@babel/standalone": "^7.29.0", diff --git a/scripts/report-legacy-surfaces.mjs b/scripts/report-legacy-surfaces.mjs index 0666613e9..26448cd71 100644 --- a/scripts/report-legacy-surfaces.mjs +++ b/scripts/report-legacy-surfaces.mjs @@ -226,6 +226,17 @@ const rustTextSurfaceMonitors = [ includePathPrefixes: ["src-tauri/src"], allowedPaths: [], }, + { + id: "rust-request-tool-policy-compat-service", + classification: "deprecated", + description: "request_tool_policy 旧服务壳回流", + patterns: [ + "crate::services::request_tool_policy_prompt_service::", + "proxycast_lib::services::request_tool_policy_prompt_service::", + "services::request_tool_policy_prompt_service::", + ], + allowedPaths: [], + }, { id: "rust-migration-setting-key-leak", classification: "deprecated", diff --git a/scripts/setup-local-aster-override.mjs b/scripts/setup-local-aster-override.mjs new file mode 100644 index 000000000..41ce59e32 --- /dev/null +++ b/scripts/setup-local-aster-override.mjs @@ -0,0 +1,150 @@ +#!/usr/bin/env node + +import fs from "node:fs"; +import path from "node:path"; +import process from "node:process"; + +const repoRoot = path.resolve(process.cwd()); +const cargoConfigDir = path.join(repoRoot, ".cargo"); +const cargoConfigPath = path.join(cargoConfigDir, "config.toml"); +const defaultAsterRepo = path.resolve(repoRoot, "..", "..", "astercloud", "aster-rust"); +const blockStart = "# >>> proxycast local aster override >>>"; +const blockEnd = "# <<< proxycast local aster override <<<"; + +function normalizePath(filePath) { + return filePath.split(path.sep).join("/"); +} + +function escapeRegExp(text) { + return text.replace(/[.*+?^${}()|[\]\\]/g, "\\$&"); +} + +function printUsage() { + console.log("用法:"); + console.log( + " npm run setup:local-aster -- [aster-rust 仓库路径] 生成仓库根 .cargo/config.toml 覆盖配置", + ); + console.log(" npm run setup:local-aster -- --clear 删除本地 Cargo 覆盖配置"); +} + +function ensureDirectory(dirPath) { + fs.mkdirSync(dirPath, { recursive: true }); +} + +function resolveAsterRepoPath() { + const arg = process.argv[2]; + if (!arg) { + return defaultAsterRepo; + } + return path.resolve(repoRoot, arg); +} + +function validateAsterRepo(asterRepoPath) { + const crates = [ + path.join(asterRepoPath, "crates", "aster", "Cargo.toml"), + path.join(asterRepoPath, "crates", "aster-models", "Cargo.toml"), + ]; + + for (const cratePath of crates) { + if (!fs.existsSync(cratePath)) { + console.error(`[proxycast] 未找到 Aster crate: ${cratePath}`); + process.exit(1); + } + } +} + +function buildConfigContent(asterRepoPath) { + const asterPath = normalizePath( + path.relative(cargoConfigDir, path.join(asterRepoPath, "crates", "aster")), + ); + const asterModelsPath = normalizePath( + path.relative(cargoConfigDir, path.join(asterRepoPath, "crates", "aster-models")), + ); + + return `${blockStart} +# 本地 Aster 覆盖配置 +# 由 scripts/setup-local-aster-override.mjs 生成。 +# 该文件已被 .gitignore 忽略,不会影响 CI/CD。 + +[patch."https://github.com/astercloud/aster-rust"] +aster-core = { path = "${asterPath}" } +aster-models = { path = "${asterModelsPath}" } +${blockEnd} +`; +} + +function readExistingConfig() { + if (!fs.existsSync(cargoConfigPath)) { + return ""; + } + + return fs.readFileSync(cargoConfigPath, "utf8"); +} + +function upsertManagedBlock(existingContent, managedBlock) { + if (!existingContent.trim()) { + return managedBlock; + } + + const blockPattern = new RegExp( + `${escapeRegExp(blockStart)}[\\s\\S]*?${escapeRegExp(blockEnd)}\\n?`, + ); + + if (blockPattern.test(existingContent)) { + return existingContent.replace(blockPattern, `${managedBlock}\n`); + } + + return `${managedBlock}\n${existingContent}`; +} + +function removeManagedBlock(existingContent) { + if (!existingContent.trim()) { + return ""; + } + + const blockPattern = new RegExp( + `${escapeRegExp(blockStart)}[\\s\\S]*?${escapeRegExp(blockEnd)}\\n?`, + ); + + return existingContent + .replace(blockPattern, "") + .replace(/^\s+/, "") + .replace(/\n{3,}/g, "\n\n") + .trim(); +} + +if (process.argv.includes("--help") || process.argv.includes("-h")) { + printUsage(); + process.exit(0); +} + +if (process.argv.includes("--clear")) { + const existingContent = readExistingConfig(); + if (!existingContent) { + console.log("[proxycast] 本地 Aster 覆盖配置不存在,无需删除。"); + process.exit(0); + } + + const nextContent = removeManagedBlock(existingContent); + if (nextContent) { + fs.writeFileSync(cargoConfigPath, `${nextContent}\n`, "utf8"); + console.log(`[proxycast] 已移除本地 Aster 覆盖区块: ${cargoConfigPath}`); + } else { + fs.rmSync(cargoConfigPath); + console.log(`[proxycast] 已删除本地 Aster 覆盖配置: ${cargoConfigPath}`); + } + process.exit(0); +} + +const asterRepoPath = resolveAsterRepoPath(); +validateAsterRepo(asterRepoPath); +ensureDirectory(cargoConfigDir); +const existingContent = readExistingConfig(); +const nextContent = upsertManagedBlock( + existingContent, + buildConfigContent(asterRepoPath), +); +fs.writeFileSync(cargoConfigPath, nextContent, "utf8"); + +console.log(`[proxycast] 已生成本地 Aster 覆盖配置: ${cargoConfigPath}`); +console.log(`[proxycast] Aster 仓库: ${asterRepoPath}`); diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 98e21b314..5eace9f7e 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -369,7 +369,7 @@ checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" [[package]] name = "aster-core" -version = "0.17.0" +version = "0.17.1" dependencies = [ "ahash", "anyhow", @@ -461,7 +461,7 @@ dependencies = [ [[package]] name = "aster-models" -version = "0.17.0" +version = "0.17.1" dependencies = [ "serde", "serde_json", @@ -2399,7 +2399,7 @@ dependencies = [ "dtoa-short", "itoa", "matches", - "phf 0.8.0", + "phf 0.10.1", "proc-macro2", "quote", "smallvec", @@ -2415,7 +2415,7 @@ dependencies = [ "cssparser-macros", "dtoa-short", "itoa", - "phf 0.8.0", + "phf 0.11.3", "smallvec", ] @@ -4336,7 +4336,7 @@ dependencies = [ "js-sys", "log", "wasm-bindgen", - "windows-core 0.56.0", + "windows-core 0.57.0", ] [[package]] @@ -5715,7 +5715,7 @@ version = "0.7.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ff32365de1b6743cb203b710788263c44a03de03802daf96092f2da4fe6ba4d7" dependencies = [ - "proc-macro-crate 1.3.1", + "proc-macro-crate 2.0.2", "proc-macro2", "quote", "syn 2.0.117", @@ -6442,9 +6442,7 @@ 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]] @@ -6453,7 +6451,9 @@ 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]] @@ -6557,12 +6557,12 @@ dependencies = [ [[package]] name = "phf_macros" -version = "0.8.0" +version = "0.10.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7f6fde18ff429ffc8fe78e2bf7f8b7a5a5a6e2a8b58bc5a9ac69198bbda9189c" +checksum = "58fdf3184dd560f160dd73922bea2d5cd6e8f064bf4b13110abd81b03697b4e0" dependencies = [ - "phf_generator 0.8.0", - "phf_shared 0.8.0", + "phf_generator 0.10.0", + "phf_shared 0.10.0", "proc-macro-hack", "proc-macro2", "quote", @@ -6974,7 +6974,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8a56d757972c98b346a9b766e3f02746cde6dd1cd1d1d563472929fdd74bec4d" dependencies = [ "anyhow", - "itertools 0.12.1", + "itertools 0.13.0", "proc-macro2", "quote", "syn 2.0.117", @@ -7084,7 +7084,7 @@ dependencies = [ [[package]] name = "proxycast-agent" -version = "0.85.0" +version = "0.86.0" dependencies = [ "aster-core", "async-trait", @@ -7109,7 +7109,7 @@ dependencies = [ [[package]] name = "proxycast-config" -version = "0.85.0" +version = "0.86.0" dependencies = [ "async-trait", "parking_lot", @@ -7125,7 +7125,7 @@ dependencies = [ [[package]] name = "proxycast-core" -version = "0.85.0" +version = "0.86.0" dependencies = [ "aster-models", "async-trait", @@ -7165,7 +7165,7 @@ dependencies = [ [[package]] name = "proxycast-credential" -version = "0.85.0" +version = "0.86.0" dependencies = [ "axum 0.7.9", "base64 0.22.1", @@ -7200,7 +7200,7 @@ dependencies = [ [[package]] name = "proxycast-gateway" -version = "0.85.0" +version = "0.86.0" dependencies = [ "axum 0.7.9", "chrono", @@ -7221,7 +7221,7 @@ dependencies = [ [[package]] name = "proxycast-infra" -version = "0.85.0" +version = "0.86.0" dependencies = [ "chrono", "dashmap 5.5.3", @@ -7241,7 +7241,7 @@ dependencies = [ [[package]] name = "proxycast-mcp" -version = "0.85.0" +version = "0.86.0" dependencies = [ "async-trait", "dirs 5.0.1", @@ -7273,7 +7273,7 @@ dependencies = [ [[package]] name = "proxycast-processor" -version = "0.85.0" +version = "0.86.0" dependencies = [ "async-trait", "parking_lot", @@ -7292,7 +7292,7 @@ dependencies = [ [[package]] name = "proxycast-providers" -version = "0.85.0" +version = "0.86.0" dependencies = [ "anyhow", "async-stream", @@ -7346,7 +7346,7 @@ dependencies = [ [[package]] name = "proxycast-server" -version = "0.85.0" +version = "0.86.0" dependencies = [ "aster-core", "async-stream", @@ -7391,7 +7391,7 @@ dependencies = [ [[package]] name = "proxycast-server-utils" -version = "0.85.0" +version = "0.86.0" dependencies = [ "axum 0.7.9", "futures", @@ -7406,7 +7406,7 @@ dependencies = [ [[package]] name = "proxycast-services" -version = "0.85.0" +version = "0.86.0" dependencies = [ "anyhow", "aster-core", @@ -7448,7 +7448,7 @@ dependencies = [ [[package]] name = "proxycast-skills" -version = "0.85.0" +version = "0.86.0" dependencies = [ "async-trait", "dirs 5.0.1", @@ -7459,12 +7459,14 @@ dependencies = [ "regex", "serde", "serde_json", + "serde_yaml", + "tempfile", "tracing", ] [[package]] name = "proxycast-terminal" -version = "0.85.0" +version = "0.86.0" dependencies = [ "async-trait", "base64 0.22.1", @@ -7491,7 +7493,7 @@ dependencies = [ [[package]] name = "proxycast-websocket" -version = "0.85.0" +version = "0.86.0" dependencies = [ "axum 0.7.9", "chrono", @@ -8970,7 +8972,7 @@ version = "3.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8b1fdf65dd6331831494dd616b30351c38e96e45921a27745cf98490458b90bb" dependencies = [ - "dirs 4.0.0", + "dirs 6.0.0", ] [[package]] diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 63e39f14c..124abce9f 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -3,7 +3,7 @@ members = ["crates/*"] resolver = "2" [workspace.package] -version = "0.85.0" +version = "0.86.0" edition = "2021" authors = ["coso"] repository = "https://github.com/aiclientproxy/proxycast" @@ -122,13 +122,12 @@ if-addrs = "0.13" 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.17.0" -# aster = { package = "aster-core", path = "../../../astercloud/aster-rust/crates/aster" } -aster = { package = "aster-core", git = "https://github.com/astercloud/aster-rust", tag = "v0.17.0" } -# 本地开发: aster-models = { path = "../../../astercloud/aster-rust/crates/aster-models" } -aster-models = { git = "https://github.com/astercloud/aster-rust", tag = "v0.17.0" } +# 默认固定到远程 Git tag,避免 CI/CD 与其他开发环境依赖本地绝对路径。 +# 如需联调本地 aster-rust,请运行: +# npm run setup:local-aster -- /path/to/aster-rust +# 脚本会在仓库根 .cargo/config.toml 写入本地 patch 覆盖;该文件已被 .gitignore 忽略。 +aster = { package = "aster-core", git = "https://github.com/astercloud/aster-rust", tag = "v0.17.1" } +aster-models = { git = "https://github.com/astercloud/aster-rust", tag = "v0.17.1" } # MCP (Model Context Protocol) rmcp = { version = "0.12.0", features = ["client", "transport-io", "transport-child-process"] } diff --git a/src-tauri/crates/agent/src/aster_state.rs b/src-tauri/crates/agent/src/aster_state.rs index ee8163af3..6f7df496c 100644 --- a/src-tauri/crates/agent/src/aster_state.rs +++ b/src-tauri/crates/agent/src/aster_state.rs @@ -26,6 +26,7 @@ use aster::agents::{Agent, SessionConfig}; use aster::model::ModelConfig; #[cfg(test)] use aster::skills::{global_registry, load_skills_from_directory, SkillSource}; +use aster::tools::{create_shared_history, EditTool, WriteTool}; use std::path::PathBuf; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; @@ -33,9 +34,250 @@ use tokio::sync::RwLock; use tokio_util::sync::CancellationToken; use crate::credential_bridge::{create_aster_provider, AsterProviderConfig, CredentialBridge}; +use crate::queued_turn::QueuedTurnSnapshot; use proxycast_core::database::DbConnection; use proxycast_services::aster_session_store::ProxyCastSessionStore; +use std::collections::{HashMap, VecDeque}; +use std::sync::Mutex; + +async fn configure_proxycast_native_file_tools(agent: &Agent) { + let shared_history = create_shared_history(); + let registry_arc = agent.tool_registry().clone(); + let mut registry = registry_arc.write().await; + registry.register(Box::new( + WriteTool::new(shared_history.clone()).with_require_read_before_overwrite(false), + )); + registry.register(Box::new( + EditTool::new(shared_history).with_require_read_before_edit(false), + )); +} + +/// 会话级 turn 排队任务 +#[derive(Debug, Clone)] +pub struct QueuedTurnTask { + pub queued_turn_id: String, + pub session_id: String, + pub event_name: String, + pub message_preview: String, + pub message_text: String, + pub created_at: i64, + pub image_count: usize, + pub payload: T, +} + +#[derive(Debug, Clone)] +struct ActiveTurnMeta { + #[cfg_attr(not(test), allow(dead_code))] + queued_turn_id: String, +} + +#[derive(Debug)] +struct SessionTurnQueueState { + active: Option, + pending: VecDeque>, +} + +impl Default for SessionTurnQueueState { + fn default() -> Self { + Self { + active: None, + pending: VecDeque::new(), + } + } +} + +impl QueuedTurnTask { + fn active_meta(&self) -> ActiveTurnMeta { + ActiveTurnMeta { + queued_turn_id: self.queued_turn_id.clone(), + } + } + + fn snapshot(&self, position: usize) -> QueuedTurnSnapshot { + QueuedTurnSnapshot { + queued_turn_id: self.queued_turn_id.clone(), + message_preview: self.message_preview.clone(), + message_text: self.message_text.clone(), + created_at: self.created_at, + image_count: self.image_count, + position, + } + } +} + +#[derive(Debug)] +pub enum QueueInsertResult { + StartNow(QueuedTurnTask), + Enqueued { + event_name: String, + snapshot: QueuedTurnSnapshot, + }, +} + +/// 会话级 turn 队列 +#[derive(Debug, Clone)] +pub struct SessionTurnQueueManager { + inner: Arc>>>, +} + +impl Default for SessionTurnQueueManager { + fn default() -> Self { + Self { + inner: Arc::new(Mutex::new(HashMap::new())), + } + } +} + +impl SessionTurnQueueManager { + pub fn has_session_state(&self, session_id: &str) -> bool { + let sessions = match self.inner.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + sessions.contains_key(session_id) + } + + pub fn restore_pending(&self, session_id: &str, tasks: Vec>) { + if tasks.is_empty() { + return; + } + + let mut sessions = match self.inner.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + let state = sessions + .entry(session_id.to_string()) + .or_insert_with(SessionTurnQueueState::default); + + if state.active.is_some() || !state.pending.is_empty() { + return; + } + + state.pending = tasks.into_iter().collect(); + } + + pub fn start_or_enqueue(&self, task: QueuedTurnTask) -> QueueInsertResult { + let mut sessions = match self.inner.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + let state = sessions + .entry(task.session_id.clone()) + .or_insert_with(SessionTurnQueueState::default); + + if state.active.is_none() { + state.active = Some(task.active_meta()); + return QueueInsertResult::StartNow(task); + } + + let position = state.pending.len() + 1; + let event_name = task.event_name.clone(); + let snapshot = task.snapshot(position); + state.pending.push_back(task); + + QueueInsertResult::Enqueued { + event_name, + snapshot, + } + } + + pub fn finish_and_take_next(&self, session_id: &str) -> Option> { + let mut sessions = match self.inner.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + let state = sessions.get_mut(session_id)?; + state.active = None; + let next = state.pending.pop_front(); + if let Some(task) = next.as_ref() { + state.active = Some(task.active_meta()); + } + if state.active.is_none() && state.pending.is_empty() { + sessions.remove(session_id); + } + next + } + + pub fn remove_queued( + &self, + session_id: &str, + queued_turn_id: &str, + ) -> Option> { + let mut sessions = match self.inner.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + let state = sessions.get_mut(session_id)?; + let index = state + .pending + .iter() + .position(|task| task.queued_turn_id == queued_turn_id)?; + let removed = state.pending.remove(index); + if state.active.is_none() && state.pending.is_empty() { + sessions.remove(session_id); + } + removed + } + + pub fn clear_pending(&self, session_id: &str) -> Vec> { + let mut sessions = match self.inner.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + let Some(state) = sessions.get_mut(session_id) else { + return Vec::new(); + }; + let cleared = state.pending.drain(..).collect::>(); + if state.active.is_none() && state.pending.is_empty() { + sessions.remove(session_id); + } + cleared + } + + pub fn snapshot(&self, session_id: &str) -> Vec { + let sessions = match self.inner.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + sessions + .get(session_id) + .map(|state| { + state + .pending + .iter() + .enumerate() + .map(|(index, task)| task.snapshot(index + 1)) + .collect::>() + }) + .unwrap_or_default() + } + + pub fn has_active(&self, session_id: &str) -> bool { + let sessions = match self.inner.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + sessions + .get(session_id) + .and_then(|state| state.active.as_ref()) + .is_some() + } + + #[cfg(test)] + fn active_queued_turn_id(&self, session_id: &str) -> Option { + let sessions = match self.inner.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + sessions + .get(session_id) + .and_then(|state| state.active.as_ref()) + .map(|meta| meta.queued_turn_id.clone()) + } +} + /// Provider 配置信息 #[derive(Debug, Clone)] pub struct ProviderConfig { @@ -67,6 +309,22 @@ pub struct AsterAgentState { initialized_cache: Arc, /// Provider 配置状态缓存(避免每次都获取锁) provider_configured_cache: Arc, + /// 会话级 turn 队列 + turn_queue: SessionTurnQueueManager, +} + +impl Clone for AsterAgentState { + fn clone(&self) -> Self { + Self { + agent: self.agent.clone(), + cancel_tokens: self.cancel_tokens.clone(), + current_provider_config: self.current_provider_config.clone(), + credential_bridge: CredentialBridge::new(), + initialized_cache: self.initialized_cache.clone(), + provider_configured_cache: self.provider_configured_cache.clone(), + turn_queue: self.turn_queue.clone(), + } + } } impl Default for AsterAgentState { @@ -85,6 +343,7 @@ impl AsterAgentState { credential_bridge: CredentialBridge::new(), initialized_cache: Arc::new(AtomicBool::new(false)), provider_configured_cache: Arc::new(AtomicBool::new(false)), + turn_queue: SessionTurnQueueManager::default(), } } @@ -174,6 +433,7 @@ impl AsterAgentState { // 使用异步方法设置 ProxyCast 专属身份 let identity = crate::create_proxycast_identity(); agent.set_identity(identity).await; + configure_proxycast_native_file_tools(&agent).await; // 加载 ProxyCast Skills 到 aster-rust 的 global_registry crate::reload_proxycast_skills(); @@ -492,6 +752,11 @@ impl AsterAgentState { self.agent.clone() } + /// 获取会话级 turn 队列管理器 + pub fn turn_queue(&self) -> SessionTurnQueueManager { + self.turn_queue.clone() + } + /// 创建新的取消令牌 pub async fn create_cancel_token(&self, session_id: &str) -> CancellationToken { let token = CancellationToken::new(); @@ -650,6 +915,108 @@ mod tests { assert!(!state.cancel_session(session_id).await); } + #[test] + fn test_session_turn_queue_manager() { + let manager = SessionTurnQueueManager::default(); + + let first = QueuedTurnTask { + queued_turn_id: "turn-1".to_string(), + session_id: "session-queue".to_string(), + event_name: "event-1".to_string(), + message_preview: "first".to_string(), + message_text: "first body".to_string(), + created_at: 1_700_000_000_000, + image_count: 0, + payload: serde_json::json!({ "message": "first" }), + }; + let second = QueuedTurnTask { + queued_turn_id: "turn-2".to_string(), + session_id: "session-queue".to_string(), + event_name: "event-2".to_string(), + message_preview: "second".to_string(), + message_text: "second body".to_string(), + created_at: 1_700_000_000_001, + image_count: 1, + payload: serde_json::json!({ "message": "second" }), + }; + + match manager.start_or_enqueue(first) { + QueueInsertResult::StartNow(task) => { + assert_eq!(task.queued_turn_id, "turn-1"); + } + QueueInsertResult::Enqueued { .. } => panic!("首条 turn 不应进入队列"), + } + + match manager.start_or_enqueue(second) { + QueueInsertResult::Enqueued { snapshot, .. } => { + assert_eq!(snapshot.queued_turn_id, "turn-2"); + assert_eq!(snapshot.message_text, "second body"); + assert_eq!(snapshot.position, 1); + } + QueueInsertResult::StartNow(_) => panic!("第二条 turn 应进入队列"), + } + + assert_eq!( + manager.active_queued_turn_id("session-queue").as_deref(), + Some("turn-1") + ); + assert_eq!(manager.snapshot("session-queue").len(), 1); + + let promoted = manager + .finish_and_take_next("session-queue") + .expect("应提升下一条 turn"); + assert_eq!(promoted.queued_turn_id, "turn-2"); + assert_eq!( + manager.active_queued_turn_id("session-queue").as_deref(), + Some("turn-2") + ); + } + + #[test] + fn test_session_turn_queue_manager_restore_pending() { + let manager = SessionTurnQueueManager::default(); + + manager.restore_pending( + "session-restore", + vec![ + QueuedTurnTask { + queued_turn_id: "turn-restore-1".to_string(), + session_id: "session-restore".to_string(), + event_name: "event-restore-1".to_string(), + message_preview: "restore-1".to_string(), + message_text: "restore body 1".to_string(), + created_at: 1_700_000_000_000, + image_count: 0, + payload: serde_json::json!({ "message": "restore-1" }), + }, + QueuedTurnTask { + queued_turn_id: "turn-restore-2".to_string(), + session_id: "session-restore".to_string(), + event_name: "event-restore-2".to_string(), + message_preview: "restore-2".to_string(), + message_text: "restore body 2".to_string(), + created_at: 1_700_000_000_001, + image_count: 0, + payload: serde_json::json!({ "message": "restore-2" }), + }, + ], + ); + + assert!(manager.has_session_state("session-restore")); + let snapshot = manager.snapshot("session-restore"); + assert_eq!(snapshot.len(), 2); + assert_eq!(snapshot[0].queued_turn_id, "turn-restore-1"); + + let promoted = manager + .finish_and_take_next("session-restore") + .expect("应从恢复队列中取出首条任务"); + assert_eq!(promoted.queued_turn_id, "turn-restore-1"); + assert_eq!( + manager.active_queued_turn_id("session-restore").as_deref(), + Some("turn-restore-1") + ); + } + // ========================================================================= // Skills 集成测试 // ========================================================================= diff --git a/src-tauri/crates/agent/src/event_converter.rs b/src-tauri/crates/agent/src/event_converter.rs index 9a79cd167..d347df4c1 100644 --- a/src-tauri/crates/agent/src/event_converter.rs +++ b/src-tauri/crates/agent/src/event_converter.rs @@ -10,6 +10,7 @@ use regex::Regex; use serde::{Deserialize, Serialize}; use crate::tool_io_offload::{maybe_offload_tool_arguments, maybe_offload_tool_result_payload}; +use crate::QueuedTurnSnapshot; const JSON_RECURSION_LIMIT: usize = 50; const JSON_TRAVERSAL_NODE_LIMIT: usize = 4_096; @@ -576,6 +577,10 @@ pub enum TauriAgentEvent { result: TauriToolResult, }, + /// 文件产物快照 + #[serde(rename = "artifact_snapshot")] + ArtifactSnapshot { artifact: TauriArtifactSnapshot }, + /// 需要用户操作(权限确认、用户输入等) #[serde(rename = "action_required")] ActionRequired { @@ -592,6 +597,38 @@ pub enum TauriAgentEvent { #[serde(rename = "context_trace")] ContextTrace { steps: Vec }, + /// 当前回合运行态摘要 + #[serde(rename = "runtime_status")] + RuntimeStatus { status: TauriRuntimeStatus }, + + /// 队列新增 + #[serde(rename = "queue_added")] + QueueAdded { + session_id: String, + queued_turn: QueuedTurnSnapshot, + }, + + /// 队列项移除 + #[serde(rename = "queue_removed")] + QueueRemoved { + session_id: String, + queued_turn_id: String, + }, + + /// 队列项开始执行 + #[serde(rename = "queue_started")] + QueueStarted { + session_id: String, + queued_turn_id: String, + }, + + /// 队列被清空 + #[serde(rename = "queue_cleared")] + QueueCleared { + session_id: String, + queued_turn_ids: Vec, + }, + /// 完成(单次响应完成) #[serde(rename = "done")] Done { @@ -646,6 +683,18 @@ pub struct TauriToolResult { pub metadata: Option>, } +/// 文件产物快照 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct TauriArtifactSnapshot { + pub artifact_id: String, + pub file_path: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub metadata: Option>, +} + /// Token 使用量 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TauriTokenUsage { @@ -660,6 +709,15 @@ pub struct TauriContextTraceStep { pub detail: String, } +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TauriRuntimeStatus { + pub phase: String, + pub title: String, + pub detail: String, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub checkpoints: Vec, +} + /// 简化的消息结构 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TauriMessage { diff --git a/src-tauri/crates/agent/src/lib.rs b/src-tauri/crates/agent/src/lib.rs index 2a4a02037..3e325522e 100644 --- a/src-tauri/crates/agent/src/lib.rs +++ b/src-tauri/crates/agent/src/lib.rs @@ -20,6 +20,7 @@ pub mod hooks; pub mod lsp_bridge; pub mod mcp_bridge; pub mod prompt; +pub mod queued_turn; pub mod request_tool_policy; pub mod session_store; pub mod shell_security; @@ -27,9 +28,11 @@ pub mod subagent_scheduler; pub mod tool_io_offload; pub mod tool_permissions; pub mod tools; +mod write_artifact_events; pub use ask_bridge::{create_ask_callback, extract_response as extract_ask_response}; pub use aster_state::{AsterAgentState, ProviderConfig}; +pub use aster_state::{QueueInsertResult, QueuedTurnTask, SessionTurnQueueManager}; pub use aster_state_support::{ build_project_system_prompt, create_proxycast_identity, create_proxycast_tool_config, create_session_config_with_project, message_helpers, reload_proxycast_skills, @@ -43,13 +46,19 @@ pub use durable_memory_fs::{ resolve_virtual_memory_path, to_virtual_memory_path, virtual_memory_relative_path, DURABLE_MEMORY_ROOT_ENV, DURABLE_MEMORY_VIRTUAL_ROOT, }; -pub use event_converter::{convert_agent_event, convert_to_tauri_message, TauriAgentEvent}; +pub use event_converter::{ + convert_agent_event, convert_to_tauri_message, TauriAgentEvent, TauriArtifactSnapshot, + TauriRuntimeStatus, +}; pub use lsp_bridge::create_lsp_callback; pub use prompt::SystemPromptBuilder; +pub use queued_turn::QueuedTurnSnapshot; pub use request_tool_policy::{ execute_web_search_preflight_if_needed, merge_system_prompt_with_request_tool_policy, - resolve_request_tool_policy, stream_reply_with_policy, ReplyAttemptError, RequestToolPolicy, - StreamReplyExecution, WebSearchExecutionTracker, REQUEST_TOOL_POLICY_MARKER, + merge_system_prompt_with_web_search_preflight_context, message_suggests_news_expansion, + resolve_request_tool_policy, resolve_request_tool_policy_with_mode, stream_reply_with_policy, + ReplyAttemptError, RequestToolPolicy, RequestToolPolicyMode, StreamReplyExecution, + WebSearchExecutionTracker, REQUEST_TOOL_POLICY_MARKER, }; pub use session_store::{ create_session_sync, get_session_sync, list_sessions_sync, SessionDetail, SessionInfo, @@ -61,3 +70,4 @@ pub use subagent_scheduler::{ }; pub use tool_permissions::{DynamicPermissionCheck, PermissionBehavior}; pub use tools::{BrowserAction, BrowserTool, BrowserToolError, BrowserToolResult}; +pub use write_artifact_events::WriteArtifactEventEmitter; diff --git a/src-tauri/crates/agent/src/queued_turn.rs b/src-tauri/crates/agent/src/queued_turn.rs new file mode 100644 index 000000000..25b4bbbb1 --- /dev/null +++ b/src-tauri/crates/agent/src/queued_turn.rs @@ -0,0 +1,12 @@ +use serde::{Deserialize, Serialize}; + +/// 会话内排队 turn 快照 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +pub struct QueuedTurnSnapshot { + pub queued_turn_id: String, + pub message_preview: String, + pub message_text: String, + pub created_at: i64, + pub image_count: usize, + pub position: usize, +} diff --git a/src-tauri/crates/agent/src/request_tool_policy.rs b/src-tauri/crates/agent/src/request_tool_policy.rs index b7f0ac109..dabaacc45 100644 --- a/src-tauri/crates/agent/src/request_tool_policy.rs +++ b/src-tauri/crates/agent/src/request_tool_policy.rs @@ -3,17 +3,25 @@ //! 该模块沉淀“请求级工具策略(例如联网搜索)”与统一流式执行逻辑, //! 供 aster_agent_cmd、scheduler、gateway 等入口复用同一条执行主链。 -use crate::event_converter::{convert_agent_event, TauriAgentEvent, TauriToolResult}; +use crate::event_converter::{ + convert_agent_event, TauriAgentEvent, TauriRuntimeStatus, TauriToolResult, +}; +use crate::write_artifact_events::WriteArtifactEventEmitter; use aster::agents::{Agent, AgentEvent}; use aster::conversation::message::Message; use aster::tools::ToolContext; -use futures::StreamExt; -use std::collections::HashMap; +use chrono::{Datelike, Local, NaiveDate}; +use futures::{stream, StreamExt}; +use regex::Regex; +use serde::{Deserialize, Serialize}; +use std::collections::{HashMap, HashSet}; use std::path::Path; use tokio_util::sync::CancellationToken; use uuid::Uuid; pub const REQUEST_TOOL_POLICY_MARKER: &str = "【请求级工具策略】"; +pub const WEB_SEARCH_PREFETCH_CONTEXT_MARKER: &str = "【联网预检索上下文】"; +pub const WEB_SEARCH_SYNTHESIS_MARKER: &str = "【预检索后输出要求】"; const DEFAULT_REQUIRED_TOOLS: &[&str] = &["WebSearch"]; const DEFAULT_ALLOWED_TOOLS: &[&str] = &["WebSearch", "WebFetch"]; @@ -24,9 +32,44 @@ const WEB_SEARCH_PREFLIGHT_ENABLED_ENV: &str = "PROXYCAST_WEB_SEARCH_PREFLIGHT_E const STREAM_EVENT_DIAG_WARN_TEXT_DELTA_CHARS: usize = 2_000; const STREAM_EVENT_DIAG_WARN_TOOL_OUTPUT_CHARS: usize = 8_000; const STREAM_EVENT_DIAG_WARN_CONTEXT_STEPS: usize = 24; +const NEWS_PREFLIGHT_QUERY_LIMIT: usize = 4; +const NEWS_PREFLIGHT_QUERY_PARALLELISM: usize = 4; +const NEWS_PREFLIGHT_QUERY_OUTPUT_CHAR_LIMIT: usize = 1_600; +const NEWS_PREFLIGHT_CONTEXT_CHAR_LIMIT: usize = 6_000; +const NEWS_PREFLIGHT_RESULT_LINES: usize = 18; +const WEB_SEARCH_EMPTY_REPLY_RETRY_PROMPT: &str = "请继续。你已经完成本回合所需的 WebSearch 预检索,现在必须直接给出最终答复,不要再次调用 WebSearch 或 WebFetch。请至少输出:1. 结论摘要;2. 主题归纳;3. 关键信息;4. 如有分歧,说明来源差异。"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum RequestToolPolicyMode { + #[default] + Disabled, + Allowed, + Required, +} + +impl RequestToolPolicyMode { + pub fn enables_web_search(self) -> bool { + !matches!(self, Self::Disabled) + } + + pub fn requires_web_search(self) -> bool { + matches!(self, Self::Required) + } + + pub fn as_str(self) -> &'static str { + match self { + Self::Disabled => "disabled", + Self::Allowed => "allowed", + Self::Required => "required", + } + } +} #[derive(Debug, Clone, PartialEq, Eq)] pub struct RequestToolPolicy { + /// 本次请求的联网搜索语义 + pub search_mode: RequestToolPolicyMode, /// 本次请求是否开启联网搜索策略 pub effective_web_search: bool, /// 必须至少成功一次的工具(默认 WebSearch) @@ -100,7 +143,7 @@ impl WebSearchExecutionTracker { &self, policy: &RequestToolPolicy, ) -> Result<(), String> { - if !policy.effective_web_search { + if !policy.requires_web_search() { return Ok(()); } @@ -177,14 +220,42 @@ impl WebSearchExecutionTracker { #[derive(Debug, Clone)] pub struct PreflightToolExecution { pub events: Vec, + pub planned_queries: Vec, + pub system_prompt_appendix: Option, + pub coverage_summary: Option, + pub expanded_news_search: bool, } impl PreflightToolExecution { fn none() -> Self { - Self { events: Vec::new() } + Self { + events: Vec::new(), + planned_queries: Vec::new(), + system_prompt_appendix: None, + coverage_summary: None, + expanded_news_search: false, + } } } +#[derive(Debug, Clone)] +struct PlannedWebSearchQuery { + index: usize, + query: String, + tool_id: String, + arguments: Option, +} + +#[derive(Debug, Clone)] +struct PreflightSearchOutcome { + index: usize, + query: String, + tool_id: String, + success: bool, + output: String, + error: Option, +} + #[derive(Debug, Clone)] pub struct ReplyAttemptError { pub message: String, @@ -262,6 +333,14 @@ pub struct StreamReplyExecution { } impl RequestToolPolicy { + pub fn allows_web_search(&self) -> bool { + self.search_mode.enables_web_search() + } + + pub fn requires_web_search(&self) -> bool { + self.search_mode.requires_web_search() + } + pub fn matches_any_required_tool(&self, tool_name: &str) -> bool { matches_tool_list(tool_name, &self.required_tools) } @@ -283,7 +362,22 @@ pub fn resolve_request_tool_policy( request_web_search: Option, mode_default: bool, ) -> RequestToolPolicy { - let effective_web_search = request_web_search.unwrap_or(mode_default); + resolve_request_tool_policy_with_mode(request_web_search, None, mode_default) +} + +pub fn resolve_request_tool_policy_with_mode( + request_web_search: Option, + request_search_mode: Option, + mode_default: bool, +) -> RequestToolPolicy { + let search_mode = match (request_web_search, request_search_mode) { + (Some(false), _) => RequestToolPolicyMode::Disabled, + (_, Some(mode)) => mode, + (Some(true), None) => RequestToolPolicyMode::Allowed, + (None, None) if mode_default => RequestToolPolicyMode::Allowed, + _ => RequestToolPolicyMode::Disabled, + }; + let effective_web_search = search_mode.enables_web_search(); let required_tools = parse_tool_list_env(WEB_SEARCH_REQUIRED_TOOLS_ENV, DEFAULT_REQUIRED_TOOLS); let mut allowed_tools = parse_tool_list_env(WEB_SEARCH_ALLOWED_TOOLS_ENV, DEFAULT_ALLOWED_TOOLS); @@ -299,6 +393,7 @@ pub fn resolve_request_tool_policy( } RequestToolPolicy { + search_mode, effective_web_search, required_tools, allowed_tools, @@ -314,7 +409,7 @@ pub fn merge_system_prompt_with_request_tool_policy( base_prompt: Option, policy: &RequestToolPolicy, ) -> Option { - if !policy.effective_web_search { + if !policy.allows_web_search() { return base_prompt; } @@ -324,17 +419,32 @@ pub fn merge_system_prompt_with_request_tool_policy( policy.disallowed_tools.join(", ") }; - let policy_prompt = format!( - "{REQUEST_TOOL_POLICY_MARKER}\n\ -- 用户在本次请求中已开启“联网搜索”开关。\n\ + let policy_prompt = match policy.search_mode { + RequestToolPolicyMode::Disabled => return base_prompt, + RequestToolPolicyMode::Allowed => format!( + "{REQUEST_TOOL_POLICY_MARKER}\n\ +- 用户在本次请求中允许你使用联网搜索,但这不代表本回合必须联网。\n\ +- 你必须先理解用户意图,优先判断应该直接回答、深度思考、规划、后台任务、多代理,还是联网核实。\n\ +- 只有在用户明确要求搜索,或问题涉及最新、实时、价格、政策、规则、版本、新闻、日期敏感信息,或高风险信息需要核实时,才调用 {}(必要时再调用 WebFetch)。\n\ +- 若无需联网即可可靠完成,就直接回答,不要为了展示工具能力而搜索。\n\ +- 允许工具: {}\n\ +- 禁止工具: {}", + policy.required_tools.join(", "), + policy.allowed_tools.join(", "), + disallowed_line + ), + RequestToolPolicyMode::Required => format!( + "{REQUEST_TOOL_POLICY_MARKER}\n\ +- 用户在本次请求中已明确要求联网搜索。\n\ - 必须先调用 {} 至少一次(必要时再调用 WebFetch),再输出最终答复。\n\ - 若工具调用失败,必须返回失败原因与尝试记录;不要在未完成必需工具调用前直接给最终结论。\n\ - 允许工具: {}\n\ - 禁止工具: {}", - policy.required_tools.join(", "), - policy.allowed_tools.join(", "), - disallowed_line - ); + policy.required_tools.join(", "), + policy.allowed_tools.join(", "), + disallowed_line + ), + }; match base_prompt { Some(base) => { @@ -403,223 +513,547 @@ fn normalize_tool_name(value: &str) -> String { .collect::() } -fn extract_inline_agent_provider_error(message: &Message) -> Option { - let text = message.as_concat_text(); - let text = text.trim(); - if text.is_empty() { - return None; - } - if !text.contains("Ran into this error:") { - return None; - } - if !text.contains("Please retry if you think this is a transient or recoverable error.") { - return None; +pub fn message_suggests_news_expansion(message: &str) -> bool { + let trimmed = message.trim(); + if trimmed.is_empty() { + return false; } - let after_prefix = text.split_once("Ran into this error:")?.1.trim(); - let detail = after_prefix - .split_once("\n\nPlease retry if you think this is a transient or recoverable error.") - .map(|(left, _)| left.trim()) - .unwrap_or(after_prefix) - .trim_end_matches('.'); - - if detail.is_empty() { - return Some("Agent provider execution failed".to_string()); + let normalized = trimmed.to_ascii_lowercase(); + let has_news_keyword = [ + "新闻", + "快讯", + "头条", + "要闻", + "news", + "headline", + "headlines", + "briefing", + "roundup", + "recap", + "digest", + ] + .iter() + .any(|keyword| normalized.contains(keyword)); + if !has_news_keyword { + return false; } - Some(format!("Agent provider execution failed: {detail}")) + let has_time_keyword = [ + "今天", + "今日", + "昨天", + "昨晚", + "最新", + "实时", + "本周", + "这周", + "today", + "latest", + "breaking", + "march", + "april", + "may", + "june", + "july", + "august", + "september", + "october", + "november", + "december", + "january", + "february", + ] + .iter() + .any(|keyword| normalized.contains(keyword)); + let has_summary_keyword = [ + "汇总", + "综述", + "盘点", + "整理", + "总结", + "写一篇", + "写成", + "简报", + "日报", + "报道", + "summary", + "summarize", + "wrap up", + "report", + "brief", + "briefing", + ] + .iter() + .any(|keyword| normalized.contains(keyword)); + let has_explicit_date = Regex::new(r"\d{1,2}月\d{1,2}日") + .ok() + .map(|re| re.is_match(trimmed)) + .unwrap_or(false) + || Regex::new( + r"\b(?:jan|feb|mar|apr|may|jun|jul|aug|sep|sept|oct|nov|dec)[a-z]*\s+\d{1,2}\b", + ) + .ok() + .map(|re| re.is_match(&normalized)) + .unwrap_or(false); + + has_time_keyword || has_summary_keyword || has_explicit_date } -/// 当开启联网搜索时,在正式回复前执行一次 WebSearch 预调用。 -/// -/// 目标: -/// - 通过执行层保证至少一次 WebSearch 调用(而非仅依赖提示词) -/// - 统一生成 tool_start/tool_end 事件,供前端落地 -/// - 若预调用失败,返回失败原因并由上层中断本次回答 -pub async fn execute_web_search_preflight_if_needed( - agent: &Agent, - session_id: &str, - message_text: &str, - working_directory: Option<&Path>, - cancel_token: Option, - policy: &RequestToolPolicy, - tracker: &mut WebSearchExecutionTracker, -) -> Result { - if !policy.effective_web_search || !is_web_search_preflight_enabled() { - return Ok(PreflightToolExecution::none()); - } - - let registry_arc = agent.tool_registry().clone(); - let registry = registry_arc.read().await; - let available_tools = registry.get_definitions(); - let preflight_tool = available_tools - .iter() - .find(|definition| { - policy.matches_any_required_tool(&definition.name) - && normalize_tool_name(&definition.name).contains("websearch") - }) - .ok_or_else(|| { - format!( - "联网搜索已开启,但未找到可执行的必需工具定义。required_tools={}, available_tools={}", - policy.required_tools.join(", "), - available_tools - .iter() - .map(|definition| definition.name.clone()) - .collect::>() - .join(", ") - ) - })?; - - let query = derive_preflight_query(message_text); - let params = serde_json::json!({ "query": query }); - let arguments = serde_json::to_string(¶ms).ok(); - let tool_id = format!("preflight-websearch-{}", Uuid::new_v4()); - tracker.record_tool_start(policy, &tool_id, &preflight_tool.name); - - let mut context = ToolContext::new( - working_directory - .map(Path::to_path_buf) - .or_else(|| std::env::current_dir().ok()) - .unwrap_or_default(), - ) - .with_session_id(session_id.to_string()); - if let Some(token) = cancel_token { - context = context.with_cancellation_token(token); - } - - let mut events = vec![TauriAgentEvent::ToolStart { - tool_name: preflight_tool.name.clone(), - tool_id: tool_id.clone(), - arguments, - }]; - - let result = registry - .execute(&preflight_tool.name, params, &context, None) - .await - .map_err(|error| format!("执行 WebSearch 预调用失败: {}", error.to_string())); - - match result { - Ok(tool_result) => { - tracker.record_tool_end( - policy, - &tool_id, - tool_result.success, - tool_result.error.as_deref(), - ); - let event = TauriAgentEvent::ToolEnd { - tool_id, - result: TauriToolResult { - success: tool_result.success, - output: tool_result.output.unwrap_or_default(), - error: tool_result.error, - images: None, - metadata: None, - }, - }; - events.push(event); - - if events - .last() - .and_then(|event| match event { - TauriAgentEvent::ToolEnd { result, .. } => Some(result.success), - _ => None, - }) - .unwrap_or(false) - { - Ok(PreflightToolExecution { events }) +pub fn merge_system_prompt_with_web_search_preflight_context( + base_prompt: Option, + appendix: Option, +) -> Option { + match (base_prompt, appendix) { + (Some(base), Some(extra)) => { + if base.contains(WEB_SEARCH_PREFETCH_CONTEXT_MARKER) { + Some(base) + } else if base.trim().is_empty() { + Some(extra) } else { - let failure = events.last().and_then(|event| match event { - TauriAgentEvent::ToolEnd { result, .. } => result.error.clone(), - _ => None, - }); - Err(format!( - "联网搜索预调用失败: {}", - failure.unwrap_or_else(|| "unknown".to_string()) - )) + Some(format!("{base}\n\n{extra}")) } } - Err(error) => { - tracker.record_tool_end(policy, &tool_id, false, Some(error.as_str())); - events.push(TauriAgentEvent::ToolEnd { - tool_id, - result: TauriToolResult { - success: false, - output: String::new(), - error: Some(error.clone()), - images: None, - metadata: None, - }, - }); - Err(error) - } + (Some(base), None) => Some(base), + (None, Some(extra)) => Some(extra), + (None, None) => None, } } -/// 统一流式执行器:执行 preflight + reply 流,并复用统一的策略校验。 -pub async fn stream_reply_with_policy( - agent: &Agent, +fn should_run_web_search_preflight(policy: &RequestToolPolicy, message_text: &str) -> bool { + if !is_web_search_preflight_enabled() { + return false; + } + + policy.requires_web_search() + || (policy.allows_web_search() && message_suggests_news_expansion(message_text)) +} + +fn split_before_followup_clause(message_text: &str) -> String { + let mut candidate = message_text.trim().replace('\n', " "); + for delimiter in [ + ",并", + ",并", + "并将", + "并且", + "然后", + "再把", + "再将", + "再帮我", + " afterwards ", + " and then ", + " then ", + ] { + if let Some((head, _)) = candidate.split_once(delimiter) { + candidate = head.trim().to_string(); + break; + } + } + candidate +} + +fn sanitize_news_search_clause(message_text: &str) -> String { + let mut candidate = split_before_followup_clause(message_text); + for prefix in [ + "请帮我", + "帮我", + "麻烦你", + "请你", + "请", + "帮忙", + "替我", + "能否", + "可以", + ] { + candidate = candidate.trim_start_matches(prefix).trim().to_string(); + } + for verb in [ + "找一下", + "搜一下", + "搜索", + "查一下", + "查找", + "检索", + "收集", + "整理", + "找", + "搜", + "查", + ] { + candidate = candidate.trim_start_matches(verb).trim().to_string(); + } + + let collapsed = candidate + .replace(['?', '?', '。', ',', ','], " ") + .split_whitespace() + .collect::>() + .join(" "); + if collapsed.is_empty() { + derive_preflight_query(message_text) + } else { + collapsed + } +} + +fn month_name(month: u32) -> &'static str { + match month { + 1 => "January", + 2 => "February", + 3 => "March", + 4 => "April", + 5 => "May", + 6 => "June", + 7 => "July", + 8 => "August", + 9 => "September", + 10 => "October", + 11 => "November", + 12 => "December", + _ => "March", + } +} + +fn resolve_topic_labels(message_text: &str) -> (&'static str, &'static str) { + let normalized = message_text.to_ascii_lowercase(); + if normalized.contains("国际") + || normalized.contains("international") + || normalized.contains("world") + { + ("国际新闻", "international news") + } else if normalized.contains("国内") || normalized.contains("china") { + ("国内新闻", "china news") + } else if normalized.contains("科技") + || normalized.contains("ai ") + || normalized.contains("ai新闻") + { + ("科技新闻", "technology news") + } else { + ("新闻", "news") + } +} + +fn resolve_absolute_news_date(message_text: &str, today: NaiveDate) -> Option<(String, String)> { + let normalized = message_text.to_ascii_lowercase(); + if normalized.contains("今天") || normalized.contains("今日") || normalized.contains("today") + { + return Some(( + format!("{}年{}月{}日", today.year(), today.month(), today.day()), + format!( + "{} {} {}", + month_name(today.month()), + today.day(), + today.year() + ), + )); + } + + if let Ok(re) = Regex::new(r"(?P\d{1,2})月(?P\d{1,2})日") { + if let Some(captures) = re.captures(message_text) { + let month = captures + .name("month") + .and_then(|value| value.as_str().parse::().ok())?; + let day = captures + .name("day") + .and_then(|value| value.as_str().parse::().ok())?; + if (1..=12).contains(&month) && (1..=31).contains(&day) { + return Some(( + format!("{}年{}月{}日", today.year(), month, day), + format!("{} {} {}", month_name(month), day, today.year()), + )); + } + } + } + + None +} + +fn dedup_queries(values: Vec) -> Vec { + let mut seen = HashSet::new(); + let mut result = Vec::new(); + for value in values { + let normalized = value.trim(); + if normalized.is_empty() { + continue; + } + let key = normalized.to_ascii_lowercase(); + if seen.insert(key) { + result.push(normalized.to_string()); + } + } + result +} + +fn build_news_preflight_queries_with_reference( message_text: &str, - working_directory: Option<&Path>, + today: NaiveDate, +) -> Vec { + let base_clause = sanitize_news_search_clause(message_text); + let (zh_topic, en_topic) = resolve_topic_labels(message_text); + let mut queries = vec![derive_preflight_query(&base_clause)]; + + if let Some((zh_date, en_date)) = resolve_absolute_news_date(message_text, today) { + queries.push(format!("{zh_date} {zh_topic}")); + queries.push(format!("{en_date} {en_topic}")); + queries.push(format!("{en_date} world headlines")); + } else { + queries.push(format!("{base_clause} {zh_topic}")); + queries.push(format!("{base_clause} {en_topic}")); + queries.push(format!("{base_clause} latest headlines")); + } + + dedup_queries(queries) + .into_iter() + .take(NEWS_PREFLIGHT_QUERY_LIMIT) + .collect() +} + +fn build_preflight_queries(message_text: &str, policy: &RequestToolPolicy) -> Vec { + if message_suggests_news_expansion(message_text) && policy.allows_web_search() { + return build_news_preflight_queries_with_reference( + message_text, + Local::now().date_naive(), + ); + } + + vec![derive_preflight_query(message_text)] +} + +fn normalize_url_candidate(raw_url: &str) -> String { + raw_url + .trim() + .trim_end_matches([',', '.', ';', ')', ']', '>']) + .to_string() +} + +fn extract_urls_from_output(output: &str) -> Vec { + let mut urls = Vec::new(); + let mut seen = HashSet::new(); + if let Ok(re) = Regex::new(r#"https?://[^\s<>"')\]]+"#) { + for capture in re.find_iter(output) { + let url = normalize_url_candidate(capture.as_str()); + if !url.is_empty() && seen.insert(url.clone()) { + urls.push(url); + } + } + } + urls +} + +fn extract_domain(url: &str) -> String { + let without_protocol = url + .trim() + .trim_start_matches("https://") + .trim_start_matches("http://"); + without_protocol + .split(['/', '?', '#']) + .next() + .unwrap_or(without_protocol) + .trim_start_matches("www.") + .to_string() +} + +fn truncate_output_for_context(output: &str, max_chars: usize) -> String { + let normalized = output + .lines() + .map(str::trim_end) + .filter(|line| !line.trim().is_empty()) + .take(NEWS_PREFLIGHT_RESULT_LINES) + .collect::>() + .join("\n") + .trim() + .to_string(); + if normalized.chars().count() <= max_chars { + normalized + } else { + normalized.chars().take(max_chars).collect::() + "…" + } +} + +fn build_coverage_summary( + planned_queries: &[String], + outcomes: &[PreflightSearchOutcome], +) -> Option { + if planned_queries.is_empty() { + return None; + } + + let successful = outcomes.iter().filter(|item| item.success).count(); + let mut unique_urls = HashSet::new(); + let mut unique_domains = HashSet::new(); + for outcome in outcomes { + for url in extract_urls_from_output(&outcome.output) { + unique_domains.insert(extract_domain(&url)); + unique_urls.insert(url); + } + } + + Some(format!( + "已并发预检索 {} 组查询,成功 {} 组,提取 {} 条去重链接,覆盖 {} 个站点。", + planned_queries.len(), + successful, + unique_urls.len(), + unique_domains.len() + )) +} + +fn build_preflight_prompt_appendix( + planned_queries: &[String], + outcomes: &[PreflightSearchOutcome], +) -> Option { + let successful = outcomes + .iter() + .filter(|item| item.success && !item.output.trim().is_empty()) + .collect::>(); + if successful.is_empty() { + return None; + } + + let mut sections = vec![ + WEB_SEARCH_PREFETCH_CONTEXT_MARKER.to_string(), + "本回合已先使用统一的 WebSearch 工具完成预检索。请优先基于以下结果做主题聚类、交叉验证和来源整合,不要退回到一次浅层搜索。".to_string(), + "除非这些结果明显不足以回答用户问题,否则不要再次调用 WebSearch 或 WebFetch,也不要重复同一组查询;下一步应直接输出最终总结,而不是停留在工具轨迹。".to_string(), + ]; + if let Some(summary) = build_coverage_summary(planned_queries, outcomes) { + sections.push(summary); + } + sections.push("整理要求:先归纳主题,再写结论;优先采用多来源一致信息;若只来自单一来源,要在回答里显式说明。".to_string()); + + let mut remaining_chars = NEWS_PREFLIGHT_CONTEXT_CHAR_LIMIT; + for outcome in successful { + if remaining_chars == 0 { + break; + } + let excerpt_limit = remaining_chars.min(NEWS_PREFLIGHT_QUERY_OUTPUT_CHAR_LIMIT); + let excerpt = truncate_output_for_context(&outcome.output, excerpt_limit); + if excerpt.trim().is_empty() { + continue; + } + remaining_chars = remaining_chars.saturating_sub(excerpt.chars().count()); + sections.push(format!( + "### Query {}: {}\n{}", + outcome.index + 1, + outcome.query, + excerpt + )); + } + + Some(sections.join("\n\n")) +} + +fn merge_system_prompt_with_web_search_synthesis_instruction( + base_prompt: Option, +) -> Option { + let synthesis_prompt = format!( + "{WEB_SEARCH_SYNTHESIS_MARKER}\n\ +- 你已经完成本回合所需的 WebSearch 预检索。\n\ +- 现在必须直接输出最终答复,不要再次调用 WebSearch 或 WebFetch。\n\ +- 至少给出:结论摘要、主题归纳、关键信息、来源分歧说明。\n\ +- 绝不能只停留在搜索轨迹或工具状态。" + ); + + match base_prompt { + Some(base) => { + if base.contains(WEB_SEARCH_SYNTHESIS_MARKER) { + Some(base) + } else if base.trim().is_empty() { + Some(synthesis_prompt) + } else { + Some(format!("{base}\n\n{synthesis_prompt}")) + } + } + None => Some(synthesis_prompt), + } +} + +fn build_web_search_synthesis_runtime_status(coverage_summary: Option<&str>) -> TauriRuntimeStatus { + let mut checkpoints = vec![ + "已完成 WebSearch 预检索".to_string(), + "正在把检索结果整理为最终答复".to_string(), + "本阶段不再重复执行搜索".to_string(), + ]; + if let Some(summary) = coverage_summary + .map(str::trim) + .filter(|value| !value.is_empty()) + { + checkpoints.push(summary.to_string()); + } + + TauriRuntimeStatus { + phase: "synthesizing".to_string(), + title: "正在整理联网结果".to_string(), + detail: "已完成前置扩搜,正在基于已有 WebSearch 结果输出最终总结,不再重复检索。" + .to_string(), + checkpoints, + } +} + +fn duplicate_session_config(config: &aster::agents::SessionConfig) -> aster::agents::SessionConfig { + aster::agents::SessionConfig { + id: config.id.clone(), + schedule_id: config.schedule_id.clone(), + max_turns: config.max_turns, + retry_config: config.retry_config.clone(), + system_prompt: config.system_prompt.clone(), + include_context_trace: config.include_context_trace, + } +} + +fn should_retry_after_empty_reply( + preflight_execution: &PreflightToolExecution, + current_text_output: &str, + tracker: &WebSearchExecutionTracker, +) -> bool { + if !current_text_output.trim().is_empty() { + return false; + } + + preflight_execution.system_prompt_appendix.is_some() + || preflight_execution.expanded_news_search + || !tracker.ordered_tool_ids.is_empty() +} + +#[allow(clippy::too_many_arguments)] +async fn stream_agent_reply_once( + agent: &Agent, + user_message: Message, session_config: aster::agents::SessionConfig, cancel_token: Option, request_tool_policy: &RequestToolPolicy, - mut on_event: F, -) -> Result + web_search_tracker: &mut WebSearchExecutionTracker, + write_artifact_emitter: &mut WriteArtifactEventEmitter, + emitted_any: &mut bool, + text_chunks: &mut Vec, + event_errors: &mut Vec, + diagnostics: &mut StreamEventDiagnostics, + on_event: &mut F, +) -> Result<(), ReplyAttemptError> where F: FnMut(&TauriAgentEvent), { - let mut web_search_tracker = WebSearchExecutionTracker::default(); - let preflight = execute_web_search_preflight_if_needed( - agent, - &session_config.id, - message_text, - working_directory, - cancel_token.clone(), - request_tool_policy, - &mut web_search_tracker, - ) - .await; - match preflight { - Ok(preflight_execution) => { - for event in preflight_execution.events { - on_event(&event); - } - } - Err(error) => { - return Err(ReplyAttemptError { - message: format!( - "{error}\n尝试记录: {}", - web_search_tracker.format_attempts() - ), - emitted_any: false, - }); - } - } - - let user_message = Message::user().with_text(message_text); let mut stream = agent .reply(user_message, session_config, cancel_token) .await .map_err(|e| ReplyAttemptError { message: format!("Agent error: {e}"), - emitted_any: false, + emitted_any: *emitted_any, })?; - let mut emitted_any = false; - let mut text_chunks: Vec = Vec::new(); - let mut event_errors: Vec = Vec::new(); - let mut diagnostics = StreamEventDiagnostics::default(); - while let Some(event_result) = stream.next().await { match event_result { Ok(agent_event) => { - emitted_any = true; + *emitted_any = true; let inline_provider_error = match &agent_event { AgentEvent::Message(message) => extract_inline_agent_provider_error(message), _ => None, }; let tauri_events = convert_agent_event(agent_event); - for tauri_event in tauri_events { + for mut tauri_event in tauri_events { + let extra_events = write_artifact_emitter.process_event(&mut tauri_event); + for extra_event in &extra_events { + update_stream_event_diagnostics(diagnostics, extra_event); + on_event(extra_event); + } + match &tauri_event { TauriAgentEvent::TextDelta { text } => { if !text.is_empty() { @@ -648,7 +1082,7 @@ where } _ => {} } - update_stream_event_diagnostics(&mut diagnostics, &tauri_event); + update_stream_event_diagnostics(diagnostics, &tauri_event); on_event(&tauri_event); } if let Some(message) = inline_provider_error { @@ -661,12 +1095,314 @@ where Err(e) => { return Err(ReplyAttemptError { message: format!("Stream error: {e}"), - emitted_any, + emitted_any: *emitted_any, }); } } } + Ok(()) +} + +fn extract_inline_agent_provider_error(message: &Message) -> Option { + let text = message.as_concat_text(); + let text = text.trim(); + if text.is_empty() { + return None; + } + if !text.contains("Ran into this error:") { + return None; + } + if !text.contains("Please retry if you think this is a transient or recoverable error.") { + return None; + } + + let after_prefix = text.split_once("Ran into this error:")?.1.trim(); + let detail = after_prefix + .split_once("\n\nPlease retry if you think this is a transient or recoverable error.") + .map(|(left, _)| left.trim()) + .unwrap_or(after_prefix) + .trim_end_matches('.'); + + if detail.is_empty() { + return Some("Agent provider execution failed".to_string()); + } + + Some(format!("Agent provider execution failed: {detail}")) +} + +/// 当开启联网搜索时,在正式回复前执行 WebSearch 预检索。 +/// +/// 目标: +/// - 在需要时通过执行层主动完成新闻类扩搜,而不是只依赖模型自己多次调用搜索。 +/// - 统一生成 tool_start/tool_end 事件,供前端 harness 展示。 +/// - 将预检索结果压缩注入 system prompt,帮助模型做更深的事实整合。 +/// - 若本回合被明确要求必须先搜索,且预检索全部失败,则由上层中断本次回答。 +pub async fn execute_web_search_preflight_if_needed( + agent: &Agent, + session_id: &str, + message_text: &str, + working_directory: Option<&Path>, + cancel_token: Option, + policy: &RequestToolPolicy, + tracker: &mut WebSearchExecutionTracker, +) -> Result { + if !should_run_web_search_preflight(policy, message_text) { + return Ok(PreflightToolExecution::none()); + } + + let registry_arc = agent.tool_registry().clone(); + let registry = registry_arc.read().await; + let available_tools = registry.get_definitions(); + let preflight_tool = available_tools + .iter() + .find(|definition| { + policy.matches_any_required_tool(&definition.name) + && normalize_tool_name(&definition.name).contains("websearch") + }) + .ok_or_else(|| { + format!( + "联网搜索已开启,但未找到可执行的必需工具定义。required_tools={}, available_tools={}", + policy.required_tools.join(", "), + available_tools + .iter() + .map(|definition| definition.name.clone()) + .collect::>() + .join(", ") + ) + })?; + let preflight_tool_name = preflight_tool.name.clone(); + drop(registry); + + let planned_queries = build_preflight_queries(message_text, policy) + .into_iter() + .enumerate() + .map(|(index, query)| { + let params = serde_json::json!({ "query": query }); + PlannedWebSearchQuery { + index, + query, + tool_id: format!("preflight-websearch-{}-{}", index + 1, Uuid::new_v4()), + arguments: serde_json::to_string(¶ms).ok(), + } + }) + .collect::>(); + let expanded_news_search = planned_queries.len() > 1; + + let working_directory = working_directory + .map(Path::to_path_buf) + .or_else(|| std::env::current_dir().ok()) + .unwrap_or_default(); + let mut events = Vec::new(); + for planned in &planned_queries { + tracker.record_tool_start(policy, &planned.tool_id, &preflight_tool_name); + events.push(TauriAgentEvent::ToolStart { + tool_name: preflight_tool_name.clone(), + tool_id: planned.tool_id.clone(), + arguments: planned.arguments.clone(), + }); + } + + #[allow(clippy::redundant_iter_cloned)] + let mut outcomes = stream::iter(planned_queries.iter().cloned().map(|planned| { + let registry_arc = registry_arc.clone(); + let preflight_tool_name = preflight_tool_name.clone(); + let session_id = session_id.to_string(); + let working_directory = working_directory.clone(); + let cancel_token = cancel_token.clone(); + async move { + let query = planned.query.clone(); + let params = serde_json::json!({ "query": query }); + let mut context = ToolContext::new(working_directory).with_session_id(session_id); + if let Some(token) = cancel_token { + context = context.with_cancellation_token(token); + } + let result = { + let registry = registry_arc.read().await; + registry + .execute(&preflight_tool_name, params, &context, None) + .await + }; + match result { + Ok(tool_result) => PreflightSearchOutcome { + index: planned.index, + query: planned.query, + tool_id: planned.tool_id, + success: tool_result.success, + output: tool_result.output.unwrap_or_default(), + error: tool_result.error, + }, + Err(error) => PreflightSearchOutcome { + index: planned.index, + query: planned.query, + tool_id: planned.tool_id, + success: false, + output: String::new(), + error: Some(format!("执行 WebSearch 预调用失败: {}", error)), + }, + } + } + })) + .buffer_unordered(NEWS_PREFLIGHT_QUERY_PARALLELISM) + .collect::>() + .await; + outcomes.sort_by_key(|item| item.index); + + for outcome in &outcomes { + tracker.record_tool_end( + policy, + &outcome.tool_id, + outcome.success, + outcome.error.as_deref(), + ); + events.push(TauriAgentEvent::ToolEnd { + tool_id: outcome.tool_id.clone(), + result: TauriToolResult { + success: outcome.success, + output: outcome.output.clone(), + error: outcome.error.clone(), + images: None, + metadata: None, + }, + }); + } + + let planned_query_texts = planned_queries + .iter() + .map(|item| item.query.clone()) + .collect::>(); + let successful_required = outcomes.iter().any(|item| item.success); + let coverage_summary = build_coverage_summary(&planned_query_texts, &outcomes); + let system_prompt_appendix = build_preflight_prompt_appendix(&planned_query_texts, &outcomes); + + if policy.requires_web_search() && !successful_required { + let failure_details = outcomes + .iter() + .map(|item| { + format!( + "{} => {}", + item.query, + item.error.clone().unwrap_or_else(|| "unknown".to_string()) + ) + }) + .collect::>() + .join(" | "); + Err(format!("联网搜索预调用失败: {failure_details}")) + } else { + Ok(PreflightToolExecution { + events, + planned_queries: planned_query_texts, + system_prompt_appendix, + coverage_summary, + expanded_news_search, + }) + } +} + +/// 统一流式执行器:执行 preflight + reply 流,并复用统一的策略校验。 +pub async fn stream_reply_with_policy( + agent: &Agent, + message_text: &str, + working_directory: Option<&Path>, + mut session_config: aster::agents::SessionConfig, + cancel_token: Option, + request_tool_policy: &RequestToolPolicy, + mut on_event: F, +) -> Result +where + F: FnMut(&TauriAgentEvent), +{ + let mut web_search_tracker = WebSearchExecutionTracker::default(); + let preflight = execute_web_search_preflight_if_needed( + agent, + &session_config.id, + message_text, + working_directory, + cancel_token.clone(), + request_tool_policy, + &mut web_search_tracker, + ) + .await; + let preflight_execution = match preflight { + Ok(preflight_execution) => { + session_config.system_prompt = merge_system_prompt_with_web_search_preflight_context( + session_config.system_prompt.take(), + preflight_execution.system_prompt_appendix.clone(), + ); + for event in &preflight_execution.events { + on_event(event); + } + preflight_execution + } + Err(error) => { + return Err(ReplyAttemptError { + message: format!( + "{error}\n尝试记录: {}", + web_search_tracker.format_attempts() + ), + emitted_any: false, + }); + } + }; + + let mut write_artifact_emitter = WriteArtifactEventEmitter::new(session_config.id.clone()); + let mut emitted_any = false; + let mut text_chunks: Vec = Vec::new(); + let mut event_errors: Vec = Vec::new(); + let mut diagnostics = StreamEventDiagnostics::default(); + stream_agent_reply_once( + agent, + Message::user().with_text(message_text), + duplicate_session_config(&session_config), + cancel_token.clone(), + request_tool_policy, + &mut web_search_tracker, + &mut write_artifact_emitter, + &mut emitted_any, + &mut text_chunks, + &mut event_errors, + &mut diagnostics, + &mut on_event, + ) + .await?; + + let current_text_output = text_chunks.join(""); + if should_retry_after_empty_reply( + &preflight_execution, + ¤t_text_output, + &web_search_tracker, + ) { + tracing::warn!( + "[AsterAgent][WebSearchPrefetch] empty final text after preflight, retrying synthesis: session={}, attempts={}", + session_config.id, + web_search_tracker.format_attempts() + ); + let status = TauriAgentEvent::RuntimeStatus { + status: build_web_search_synthesis_runtime_status( + preflight_execution.coverage_summary.as_deref(), + ), + }; + on_event(&status); + session_config.system_prompt = merge_system_prompt_with_web_search_synthesis_instruction( + session_config.system_prompt.take(), + ); + stream_agent_reply_once( + agent, + Message::user().with_text(WEB_SEARCH_EMPTY_REPLY_RETRY_PROMPT), + duplicate_session_config(&session_config), + cancel_token, + request_tool_policy, + &mut web_search_tracker, + &mut write_artifact_emitter, + &mut emitted_any, + &mut text_chunks, + &mut event_errors, + &mut diagnostics, + &mut on_event, + ) + .await?; + } + if let Err(validation_error) = web_search_tracker.validate_web_search_requirement(request_tool_policy) { @@ -688,8 +1424,19 @@ where diagnostics.max_context_trace_steps ); + let final_text_output = text_chunks.join(""); + if final_text_output.trim().is_empty() { + return Err(ReplyAttemptError { + message: format!( + "已完成当前回合的工具执行,但模型未输出最终答复。\n尝试记录: {}", + web_search_tracker.format_attempts() + ), + emitted_any, + }); + } + Ok(StreamReplyExecution { - text_output: text_chunks.join(""), + text_output: final_text_output, event_errors, emitted_any, attempts_summary: web_search_tracker.format_attempts(), @@ -731,18 +1478,33 @@ mod tests { fn resolves_effective_web_search_with_request_override() { let policy = resolve_request_tool_policy(Some(false), true); assert!(!policy.effective_web_search); + assert_eq!(policy.search_mode, RequestToolPolicyMode::Disabled); let policy = resolve_request_tool_policy(Some(true), false); assert!(policy.effective_web_search); + assert_eq!(policy.search_mode, RequestToolPolicyMode::Allowed); } #[test] fn resolves_effective_web_search_with_mode_default() { let policy = resolve_request_tool_policy(None, true); assert!(policy.effective_web_search); + assert_eq!(policy.search_mode, RequestToolPolicyMode::Allowed); let policy = resolve_request_tool_policy(None, false); assert!(!policy.effective_web_search); + assert_eq!(policy.search_mode, RequestToolPolicyMode::Disabled); + } + + #[test] + fn resolves_required_mode_when_explicitly_requested() { + let policy = resolve_request_tool_policy_with_mode( + Some(true), + Some(RequestToolPolicyMode::Required), + false, + ); + assert!(policy.effective_web_search); + assert!(policy.requires_web_search()); } #[test] @@ -762,10 +1524,24 @@ mod tests { merge_system_prompt_with_request_tool_policy(Some("base".to_string()), &policy) .expect("merged prompt should exist"); assert!(merged.contains(REQUEST_TOOL_POLICY_MARKER)); - assert!(merged.contains("必须先调用")); + assert!(merged.contains("不代表本回合必须联网")); + assert!(merged.contains("先理解用户意图")); assert!(merged.contains("WebSearch")); } + #[test] + fn appends_required_policy_prompt_when_required() { + let policy = resolve_request_tool_policy_with_mode( + Some(true), + Some(RequestToolPolicyMode::Required), + false, + ); + let merged = + merge_system_prompt_with_request_tool_policy(Some("base".to_string()), &policy) + .expect("merged prompt should exist"); + assert!(merged.contains("必须先调用")); + } + #[test] fn no_duplicate_when_marker_exists() { let base = Some(format!("{REQUEST_TOOL_POLICY_MARKER}\nexists")); @@ -777,20 +1553,21 @@ mod tests { } #[test] - fn tracker_requires_websearch_when_enabled() { + fn tracker_does_not_require_websearch_when_only_allowed() { let policy = resolve_request_tool_policy(Some(true), false); let mut tracker = WebSearchExecutionTracker::default(); tracker.record_tool_start(&policy, "tool-1", "WebFetch"); tracker.record_tool_end(&policy, "tool-1", true, None); - let err = tracker - .validate_web_search_requirement(&policy) - .expect_err("missing web search should fail"); - assert!(err.contains("未检测到必需工具调用")); + assert!(tracker.validate_web_search_requirement(&policy).is_ok()); } #[test] - fn tracker_accepts_successful_websearch() { - let policy = resolve_request_tool_policy(Some(true), false); + fn tracker_accepts_successful_required_websearch() { + let policy = resolve_request_tool_policy_with_mode( + Some(true), + Some(RequestToolPolicyMode::Required), + false, + ); let mut tracker = WebSearchExecutionTracker::default(); tracker.record_tool_start(&policy, "tool-1", "WebSearch"); tracker.record_tool_end(&policy, "tool-1", true, None); @@ -799,7 +1576,11 @@ mod tests { #[test] fn tracker_reports_failure_record() { - let policy = resolve_request_tool_policy(Some(true), false); + let policy = resolve_request_tool_policy_with_mode( + Some(true), + Some(RequestToolPolicyMode::Required), + false, + ); let mut tracker = WebSearchExecutionTracker::default(); tracker.record_tool_start(&policy, "tool-1", "WebSearch"); tracker.record_tool_end(&policy, "tool-1", false, Some("network timeout")); @@ -809,4 +1590,57 @@ mod tests { assert!(err.contains("network timeout")); assert!(err.contains("尝试记录")); } + + #[test] + fn detects_news_expansion_for_daily_news_summary_requests() { + assert!(message_suggests_news_expansion("帮我汇总3月13日国际新闻")); + assert!(message_suggests_news_expansion( + "Please summarize the latest world news for March 13" + )); + assert!(!message_suggests_news_expansion("帮我解释一下牛顿第二定律")); + } + + #[test] + fn builds_news_preflight_queries_with_absolute_date_variants() { + let queries = build_news_preflight_queries_with_reference( + "帮我汇总3月13日国际新闻", + NaiveDate::from_ymd_opt(2026, 3, 13).expect("valid date"), + ); + + assert_eq!(queries[0], "汇总3月13日国际新闻"); + assert!(queries.contains(&"2026年3月13日 国际新闻".to_string())); + assert!(queries.contains(&"March 13 2026 international news".to_string())); + assert!(queries.contains(&"March 13 2026 world headlines".to_string())); + } + + #[test] + fn merges_web_search_preflight_context_without_duplication() { + let merged = merge_system_prompt_with_web_search_preflight_context( + Some("base".to_string()), + Some(format!("{WEB_SEARCH_PREFETCH_CONTEXT_MARKER}\ncontext")), + ) + .expect("merged prompt should exist"); + assert!(merged.contains(WEB_SEARCH_PREFETCH_CONTEXT_MARKER)); + + let preserved = merge_system_prompt_with_web_search_preflight_context( + Some(merged.clone()), + Some(format!("{WEB_SEARCH_PREFETCH_CONTEXT_MARKER}\nother")), + ) + .expect("prompt should be preserved"); + assert_eq!(preserved, merged); + } + + #[test] + fn appends_synthesis_instruction_without_duplication() { + let merged = + merge_system_prompt_with_web_search_synthesis_instruction(Some("base".to_string())) + .expect("merged prompt should exist"); + assert!(merged.contains(WEB_SEARCH_SYNTHESIS_MARKER)); + assert!(merged.contains("不要再次调用 WebSearch")); + + let preserved = + merge_system_prompt_with_web_search_synthesis_instruction(Some(merged.clone())) + .expect("prompt should be preserved"); + assert_eq!(preserved, merged); + } } diff --git a/src-tauri/crates/agent/src/subagent_scheduler.rs b/src-tauri/crates/agent/src/subagent_scheduler.rs index 326d3ea9b..4a1f4dc14 100644 --- a/src-tauri/crates/agent/src/subagent_scheduler.rs +++ b/src-tauri/crates/agent/src/subagent_scheduler.rs @@ -393,6 +393,12 @@ pub struct SubAgentProgressEvent { pub failed: usize, /// 运行中数 pub running: usize, + /// 等待中数 + pub pending: usize, + /// 已跳过数 + pub skipped: usize, + /// 是否已取消 + pub cancelled: bool, /// 进度百分比 pub percentage: f64, /// 当前任务 @@ -408,6 +414,9 @@ impl From for SubAgentProgressEvent { completed: progress.completed, failed: progress.failed, running: progress.running, + pending: progress.pending, + skipped: progress.skipped, + cancelled: progress.cancelled, percentage: progress.percentage, current_tasks: progress.current_tasks, role: None, @@ -542,6 +551,9 @@ mod tests { completed: 1, failed: 0, running: 1, + pending: 1, + skipped: 0, + cancelled: false, percentage: 33.3, current_tasks: vec!["task-1".to_string()], role: None, diff --git a/src-tauri/crates/agent/src/tool_io_offload.rs b/src-tauri/crates/agent/src/tool_io_offload.rs index 9e262f5e6..ab2893a58 100644 --- a/src-tauri/crates/agent/src/tool_io_offload.rs +++ b/src-tauri/crates/agent/src/tool_io_offload.rs @@ -658,6 +658,7 @@ pub fn build_history_tool_io_eviction_plan_for_model( #[cfg(test)] mod tests { use super::*; + use chrono::Utc; use proxycast_core::agent::types::{AgentMessage, FunctionCall, MessageContent, ToolCall}; use std::ffi::OsString; use std::sync::{Mutex, OnceLock}; diff --git a/src-tauri/crates/agent/src/write_artifact_events.rs b/src-tauri/crates/agent/src/write_artifact_events.rs new file mode 100644 index 000000000..3c401b7d0 --- /dev/null +++ b/src-tauri/crates/agent/src/write_artifact_events.rs @@ -0,0 +1,956 @@ +use crate::event_converter::{TauriAgentEvent, TauriArtifactSnapshot, TauriToolResult}; +use regex::Regex; +use serde_json::Value; +use std::collections::HashMap; +use std::sync::OnceLock; + +const PREVIEW_TEXT_MAX_CHARS: usize = 480; +const LATEST_CHUNK_MAX_CHARS: usize = 240; +const WRITE_FILE_CLOSE_TAG: &str = ""; + +#[derive(Debug, Clone, PartialEq, Eq)] +struct ParsedWriteBlock { + key: String, + path: String, + content: String, + is_complete: bool, +} + +#[derive(Debug, Clone)] +struct TrackedArtifactState { + artifact_id: String, + file_path: String, + content: String, + closed: bool, +} + +#[derive(Debug, Clone)] +struct ToolEndArtifact { + artifact_id: String, + file_path: String, + content: String, +} + +#[derive(Debug, Default)] +pub struct WriteArtifactEventEmitter { + scope_id: String, + accumulated_text: String, + next_artifact_seq: usize, + inline_artifact_ids: HashMap, + tool_artifact_ids: HashMap>, + tracked_artifacts: HashMap, +} + +impl WriteArtifactEventEmitter { + pub fn new(scope_id: impl Into) -> Self { + Self { + scope_id: scope_id.into(), + ..Self::default() + } + } + + pub fn process_event(&mut self, event: &mut TauriAgentEvent) -> Vec { + match event { + TauriAgentEvent::ToolStart { + tool_name, + tool_id, + arguments, + } => self.handle_tool_start(tool_name, tool_id, arguments.as_deref()), + TauriAgentEvent::TextDelta { text } => self.handle_text_delta(text), + TauriAgentEvent::ToolEnd { tool_id, result } => self.handle_tool_end(tool_id, result), + _ => Vec::new(), + } + } + + fn handle_tool_start( + &mut self, + tool_name: &str, + tool_id: &str, + arguments: Option<&str>, + ) -> Vec { + let Some(arguments_value) = parse_json_str(arguments) else { + return Vec::new(); + }; + let patch_text = extract_candidate_patch_text(&arguments_value); + if !is_write_like_tool(tool_name) + && !patch_text + .as_deref() + .map(contains_patch_file_directive) + .unwrap_or(false) + { + return Vec::new(); + } + + let paths = extract_candidate_paths(&arguments_value); + if paths.is_empty() { + return Vec::new(); + } + + let content = extract_candidate_content(&arguments_value).unwrap_or_default(); + let base_metadata = extract_embedded_metadata(&arguments_value); + let mut events = Vec::new(); + + for path in paths { + let artifact_id = self.ensure_tool_artifact(tool_id, path.as_str(), content.as_str()); + let phase = if content.trim().is_empty() { + "preparing" + } else { + "streaming" + }; + let metadata = build_snapshot_metadata( + base_metadata.as_ref(), + "tool_start", + phase, + false, + content.as_str(), + None, + ); + events.push(build_artifact_snapshot_event( + artifact_id, + path.as_str(), + content.as_str(), + metadata, + )); + } + + events + } + + fn handle_text_delta(&mut self, text: &str) -> Vec { + if text.is_empty() { + return Vec::new(); + } + + self.accumulated_text.push_str(text); + let blocks = parse_write_file_blocks(&self.accumulated_text); + let mut events = Vec::new(); + + for block in blocks { + let artifact_id = + self.resolve_inline_artifact_id(block.key.as_str(), block.path.as_str()); + let previous = self.tracked_artifacts.get(&artifact_id).cloned(); + let changed = previous + .as_ref() + .map(|state| state.content != block.content || state.closed != block.is_complete) + .unwrap_or(true); + if !changed { + continue; + } + + self.tracked_artifacts.insert( + artifact_id.clone(), + TrackedArtifactState { + artifact_id: artifact_id.clone(), + file_path: block.path.clone(), + content: block.content.clone(), + closed: block.is_complete, + }, + ); + + let phase = if block.is_complete { + "persisted" + } else if block.content.trim().is_empty() { + "preparing" + } else { + "streaming" + }; + let metadata = build_snapshot_metadata( + None, + "message_content", + phase, + block.is_complete, + block.content.as_str(), + None, + ); + events.push(build_artifact_snapshot_event( + artifact_id, + block.path.as_str(), + block.content.as_str(), + metadata, + )); + } + + events + } + + fn handle_tool_end( + &mut self, + tool_id: &str, + result: &mut TauriToolResult, + ) -> Vec { + let artifacts = self.collect_tool_end_artifacts(tool_id, result.metadata.as_ref()); + if artifacts.is_empty() { + return Vec::new(); + } + + annotate_tool_result_metadata(result, &artifacts); + + let mut events = Vec::new(); + for artifact in &artifacts { + if let Some(state) = self.tracked_artifacts.get_mut(&artifact.artifact_id) { + state.file_path = artifact.file_path.clone(); + state.content = artifact.content.clone(); + state.closed = true; + } else { + self.tracked_artifacts.insert( + artifact.artifact_id.clone(), + TrackedArtifactState { + artifact_id: artifact.artifact_id.clone(), + file_path: artifact.file_path.clone(), + content: artifact.content.clone(), + closed: true, + }, + ); + } + + if result.success { + let metadata = build_snapshot_metadata( + result.metadata.as_ref(), + "tool_result", + "completed", + true, + artifact.content.as_str(), + result.error.as_deref(), + ); + events.push(build_artifact_snapshot_event( + artifact.artifact_id.clone(), + artifact.file_path.as_str(), + artifact.content.as_str(), + metadata, + )); + } + } + + events + } + + fn resolve_inline_artifact_id(&mut self, block_key: &str, file_path: &str) -> String { + if let Some(existing) = self.inline_artifact_ids.get(block_key) { + return existing.clone(); + } + + let artifact_id = self + .find_active_artifact_id_by_path(file_path) + .unwrap_or_else(|| self.new_artifact_id(file_path)); + self.inline_artifact_ids + .insert(block_key.to_string(), artifact_id.clone()); + artifact_id + } + + fn ensure_tool_artifact(&mut self, tool_id: &str, file_path: &str, content: &str) -> String { + if let Some(existing) = self + .tool_artifact_ids + .get(tool_id) + .and_then(|artifact_ids| { + artifact_ids.iter().find_map(|artifact_id| { + self.tracked_artifacts.get(artifact_id).and_then(|state| { + if state.file_path == file_path { + Some(state.artifact_id.clone()) + } else { + None + } + }) + }) + }) + { + return existing; + } + + let artifact_id = self + .find_active_artifact_id_by_path(file_path) + .unwrap_or_else(|| self.new_artifact_id(file_path)); + + self.tool_artifact_ids + .entry(tool_id.to_string()) + .or_default() + .push(artifact_id.clone()); + self.tracked_artifacts.insert( + artifact_id.clone(), + TrackedArtifactState { + artifact_id: artifact_id.clone(), + file_path: file_path.to_string(), + content: content.to_string(), + closed: false, + }, + ); + artifact_id + } + + fn collect_tool_end_artifacts( + &mut self, + tool_id: &str, + metadata: Option<&HashMap>, + ) -> Vec { + let metadata_artifacts = extract_artifacts_from_metadata(metadata); + let tracked_ids = self + .tool_artifact_ids + .get(tool_id) + .cloned() + .unwrap_or_default(); + let mut artifacts = Vec::new(); + + for artifact_id in tracked_ids { + let Some(state) = self.tracked_artifacts.get(&artifact_id).cloned() else { + continue; + }; + + let metadata_match = metadata_artifacts + .iter() + .find(|artifact| artifact.file_path == state.file_path); + artifacts.push(ToolEndArtifact { + artifact_id: state.artifact_id.clone(), + file_path: metadata_match + .map(|artifact| artifact.file_path.clone()) + .unwrap_or_else(|| state.file_path.clone()), + content: state.content, + }); + } + + for metadata_artifact in metadata_artifacts { + if artifacts + .iter() + .any(|artifact| artifact.file_path == metadata_artifact.file_path) + { + continue; + } + + let artifact_id = metadata_artifact + .artifact_id + .clone() + .or_else(|| { + self.find_known_artifact_id_by_path(metadata_artifact.file_path.as_str()) + }) + .unwrap_or_else(|| self.new_artifact_id(metadata_artifact.file_path.as_str())); + let content = self + .tracked_artifacts + .get(&artifact_id) + .map(|state| state.content.clone()) + .unwrap_or_default(); + artifacts.push(ToolEndArtifact { + artifact_id, + file_path: metadata_artifact.file_path, + content, + }); + } + + artifacts + } + + fn find_active_artifact_id_by_path(&self, file_path: &str) -> Option { + self.tracked_artifacts.values().find_map(|state| { + if state.file_path == file_path && !state.closed { + Some(state.artifact_id.clone()) + } else { + None + } + }) + } + + fn find_known_artifact_id_by_path(&self, file_path: &str) -> Option { + self.tracked_artifacts.values().find_map(|state| { + if state.file_path == file_path { + Some(state.artifact_id.clone()) + } else { + None + } + }) + } + + fn new_artifact_id(&mut self, file_path: &str) -> String { + self.next_artifact_seq += 1; + format!( + "artifact:{}:{}:{:08x}", + self.scope_id, + self.next_artifact_seq, + stable_hash(file_path), + ) + } +} + +#[derive(Debug, Clone)] +struct MetadataArtifact { + artifact_id: Option, + file_path: String, +} + +fn build_artifact_snapshot_event( + artifact_id: impl Into, + file_path: &str, + content: &str, + metadata: HashMap, +) -> TauriAgentEvent { + TauriAgentEvent::ArtifactSnapshot { + artifact: TauriArtifactSnapshot { + artifact_id: artifact_id.into(), + file_path: file_path.to_string(), + content: Some(content.to_string()), + metadata: if metadata.is_empty() { + None + } else { + Some(metadata) + }, + }, + } +} + +fn annotate_tool_result_metadata(result: &mut TauriToolResult, artifacts: &[ToolEndArtifact]) { + let metadata = result.metadata.get_or_insert_with(HashMap::new); + metadata.insert("artifact_streamed".to_string(), Value::Bool(true)); + + if artifacts.len() == 1 { + metadata.insert( + "artifact_id".to_string(), + Value::String(artifacts[0].artifact_id.clone()), + ); + metadata.insert( + "artifact_path".to_string(), + Value::String(artifacts[0].file_path.clone()), + ); + metadata + .entry("path".to_string()) + .or_insert_with(|| Value::String(artifacts[0].file_path.clone())); + metadata + .entry("file_path".to_string()) + .or_insert_with(|| Value::String(artifacts[0].file_path.clone())); + } else { + metadata.insert( + "artifact_ids".to_string(), + Value::Array( + artifacts + .iter() + .map(|artifact| Value::String(artifact.artifact_id.clone())) + .collect(), + ), + ); + metadata + .entry("artifact_paths".to_string()) + .or_insert_with(|| { + Value::Array( + artifacts + .iter() + .map(|artifact| Value::String(artifact.file_path.clone())) + .collect(), + ) + }); + } +} + +fn build_snapshot_metadata( + base: Option<&HashMap>, + source: &str, + phase: &str, + complete: bool, + content: &str, + error: Option<&str>, +) -> HashMap { + let mut metadata = base.cloned().unwrap_or_default(); + metadata.insert("complete".to_string(), Value::Bool(complete)); + metadata.insert("writePhase".to_string(), Value::String(phase.to_string())); + metadata.insert("isPartial".to_string(), Value::Bool(!complete)); + metadata.insert( + "lastUpdateSource".to_string(), + Value::String(source.to_string()), + ); + + if let Some(preview) = truncate_chars(content, PREVIEW_TEXT_MAX_CHARS) { + metadata.insert("previewText".to_string(), Value::String(preview)); + } + if let Some(chunk) = take_last_chars(content, LATEST_CHUNK_MAX_CHARS) { + metadata.insert("latestChunk".to_string(), Value::String(chunk)); + } + if let Some(message) = error.map(str::trim).filter(|value| !value.is_empty()) { + metadata.insert("error".to_string(), Value::String(message.to_string())); + } + + metadata +} + +fn truncate_chars(value: &str, limit: usize) -> Option { + let trimmed = value.trim(); + if trimmed.is_empty() { + return None; + } + + let collected = trimmed.chars().take(limit).collect::(); + Some(collected) +} + +fn take_last_chars(value: &str, limit: usize) -> Option { + let trimmed = value.trim(); + if trimmed.is_empty() { + return None; + } + + let chars = trimmed.chars().collect::>(); + let start = chars.len().saturating_sub(limit); + Some(chars[start..].iter().collect()) +} + +fn stable_hash(input: &str) -> u32 { + let mut hash: u32 = 0x811c9dc5; + for byte in input.as_bytes() { + hash ^= u32::from(*byte); + hash = hash.wrapping_mul(0x01000193); + } + hash +} + +fn is_write_like_tool(tool_name: &str) -> bool { + let normalized = tool_name + .chars() + .filter(|ch| ch.is_ascii_alphanumeric()) + .flat_map(|ch| ch.to_lowercase()) + .collect::(); + normalized.contains("write") + || normalized.contains("create") + || normalized.contains("save") + || normalized.contains("output") + || normalized.contains("edit") + || normalized.contains("patch") + || normalized.contains("update") + || normalized.contains("replace") +} + +fn parse_json_str(raw: Option<&str>) -> Option { + let text = raw?.trim(); + if text.is_empty() { + return None; + } + serde_json::from_str::(text).ok() +} + +fn extract_candidate_paths(value: &Value) -> Vec { + let Some(object) = value.as_object() else { + return Vec::new(); + }; + + let mut paths = Vec::new(); + for key in [ + "path", + "file_path", + "filePath", + "target_path", + "targetPath", + "output_path", + "outputPath", + "artifact_path", + "artifactPath", + "artifact_paths", + "artifactPaths", + ] { + if let Some(candidate) = object.get(key) { + push_paths_from_value(&mut paths, candidate); + } + } + + if paths.is_empty() { + if let Some(patch_text) = extract_candidate_patch_text(value) { + push_paths_from_patch_text(&mut paths, patch_text.as_str()); + } + } + + paths +} + +fn extract_candidate_content(value: &Value) -> Option { + let object = value.as_object()?; + for key in ["content", "text", "contents", "body"] { + let Some(candidate) = object.get(key).and_then(Value::as_str) else { + continue; + }; + return Some(candidate.to_string()); + } + None +} + +fn extract_candidate_patch_text(value: &Value) -> Option { + let object = value.as_object()?; + for key in ["patch", "command", "cmd", "script"] { + let Some(candidate) = object.get(key) else { + continue; + }; + if let Some(text) = value_to_text(candidate) { + return Some(text); + } + } + None +} + +fn value_to_text(value: &Value) -> Option { + match value { + Value::String(text) => Some(text.to_string()), + Value::Array(items) => { + let parts = items + .iter() + .filter_map(Value::as_str) + .map(str::trim) + .filter(|part| !part.is_empty()) + .collect::>(); + if parts.is_empty() { + None + } else { + Some(parts.join("\n")) + } + } + _ => None, + } +} + +fn contains_patch_file_directive(text: &str) -> bool { + text.lines().any(|line| { + let trimmed = line.trim(); + trimmed.starts_with("*** Add File:") + || trimmed.starts_with("*** Update File:") + || trimmed.starts_with("*** Delete File:") + || trimmed.starts_with("*** Move to:") + }) +} + +fn push_paths_from_patch_text(target: &mut Vec, patch_text: &str) { + for line in patch_text.lines() { + let trimmed = line.trim(); + for prefix in [ + "*** Add File:", + "*** Update File:", + "*** Delete File:", + "*** Move to:", + ] { + if let Some(path) = trimmed.strip_prefix(prefix) { + if let Some(normalized) = normalize_path(path.trim()) { + if !target.iter().any(|item| item == &normalized) { + target.push(normalized); + } + } + } + } + } +} + +fn extract_embedded_metadata(value: &Value) -> Option> { + let object = value.as_object()?; + for key in ["metadata", "meta"] { + let Some(candidate) = object.get(key).and_then(Value::as_object) else { + continue; + }; + return Some( + candidate + .iter() + .map(|(key, value)| (key.clone(), value.clone())) + .collect(), + ); + } + None +} + +fn push_paths_from_value(target: &mut Vec, value: &Value) { + match value { + Value::String(path) => { + if let Some(normalized) = normalize_path(path) { + if !target.iter().any(|item| item == &normalized) { + target.push(normalized); + } + } + } + Value::Array(values) => { + for nested in values { + push_paths_from_value(target, nested); + } + } + _ => {} + } +} + +fn normalize_path(raw: &str) -> Option { + let trimmed = raw.trim(); + if trimmed.is_empty() { + None + } else { + Some(trimmed.replace('\\', "/")) + } +} + +fn extract_artifacts_from_metadata( + metadata: Option<&HashMap>, +) -> Vec { + let Some(metadata) = metadata else { + return Vec::new(); + }; + + let mut paths = Vec::new(); + for key in [ + "artifact_paths", + "artifact_path", + "path", + "absolute_path", + "output_file", + "file_path", + "output_path", + "filePath", + "outputPath", + "article_path", + "cover_meta_path", + "publish_path", + ] { + if let Some(value) = metadata.get(key) { + push_paths_from_value(&mut paths, value); + } + } + + if paths.is_empty() { + return Vec::new(); + } + + let ids = metadata + .get("artifact_ids") + .and_then(Value::as_array) + .map(|items| { + items + .iter() + .filter_map(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .collect::>() + }) + .unwrap_or_default(); + let single_id = metadata + .get("artifact_id") + .or_else(|| metadata.get("artifactId")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string); + + paths + .into_iter() + .enumerate() + .map(|(index, file_path)| MetadataArtifact { + artifact_id: ids.get(index).cloned().or_else(|| { + if index == 0 { + single_id.clone() + } else { + None + } + }), + file_path, + }) + .collect() +} + +fn parse_write_file_blocks(text: &str) -> Vec { + let regex = write_file_open_tag_regex(); + let lower = text.to_ascii_lowercase(); + let mut search_offset = 0usize; + let mut order = 0usize; + let mut blocks = Vec::new(); + + while let Some(captures) = regex.captures(&text[search_offset..]) { + let Some(full_match) = captures.get(0) else { + break; + }; + let Some(path_match) = captures.get(1) else { + search_offset += full_match.end(); + continue; + }; + + let open_start = search_offset + full_match.start(); + let open_end = search_offset + full_match.end(); + let Some(path) = normalize_path(path_match.as_str()) else { + search_offset = open_end; + continue; + }; + + let remainder = &lower[open_end..]; + if let Some(close_offset) = remainder.find(WRITE_FILE_CLOSE_TAG) { + let content_end = open_end + close_offset; + blocks.push(ParsedWriteBlock { + key: format!("{order}:{path}"), + path: path.clone(), + content: text[open_end..content_end].to_string(), + is_complete: true, + }); + search_offset = content_end + WRITE_FILE_CLOSE_TAG.len(); + } else { + blocks.push(ParsedWriteBlock { + key: format!("{order}:{path}"), + path: path.clone(), + content: text[open_end..].to_string(), + is_complete: false, + }); + break; + } + order += 1; + if search_offset <= open_start { + break; + } + } + + blocks +} + +fn write_file_open_tag_regex() -> &'static Regex { + static RE: OnceLock = OnceLock::new(); + RE.get_or_init(|| { + Regex::new(r#"(?is)"#) + .expect("write_file open tag regex should be valid") + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn assert_snapshot( + event: &TauriAgentEvent, + expected_path: &str, + expected_content: &str, + expected_complete: bool, + ) -> String { + match event { + TauriAgentEvent::ArtifactSnapshot { artifact } => { + assert_eq!(artifact.file_path, expected_path); + assert_eq!(artifact.content.as_deref(), Some(expected_content)); + assert_eq!( + artifact + .metadata + .as_ref() + .and_then(|metadata| metadata.get("complete")) + .and_then(Value::as_bool), + Some(expected_complete) + ); + artifact.artifact_id.clone() + } + _ => panic!("expected artifact snapshot"), + } + } + + #[test] + fn tool_start_with_path_only_emits_preparing_snapshot() { + let mut emitter = WriteArtifactEventEmitter::new("session-1"); + let mut event = TauriAgentEvent::ToolStart { + tool_name: "write_file".to_string(), + tool_id: "tool-1".to_string(), + arguments: Some(r#"{"path":"drafts/demo.md"}"#.to_string()), + }; + + let extras = emitter.process_event(&mut event); + assert_eq!(extras.len(), 1); + let artifact_id = assert_snapshot(&extras[0], "drafts/demo.md", "", false); + + match &extras[0] { + TauriAgentEvent::ArtifactSnapshot { artifact } => { + assert_eq!( + artifact + .metadata + .as_ref() + .and_then(|metadata| metadata.get("writePhase")) + .and_then(Value::as_str), + Some("preparing") + ); + } + _ => panic!("expected artifact snapshot"), + } + assert!(artifact_id.starts_with("artifact:session-1:")); + } + + #[test] + fn tool_start_apply_patch_emits_preparing_snapshot_for_target_file() { + let mut emitter = WriteArtifactEventEmitter::new("session-1"); + let mut event = TauriAgentEvent::ToolStart { + tool_name: "apply_patch".to_string(), + tool_id: "tool-patch-1".to_string(), + arguments: Some( + r#"{"patch":"*** Begin Patch\n*** Update File: drafts/demo.md\n@@\n-old\n+new\n*** End Patch\n"}"# + .to_string(), + ), + }; + + let extras = emitter.process_event(&mut event); + + assert_eq!(extras.len(), 1); + assert_snapshot(&extras[0], "drafts/demo.md", "", false); + } + + #[test] + fn shell_apply_patch_command_emits_preparing_snapshot_for_target_file() { + let mut emitter = WriteArtifactEventEmitter::new("session-1"); + let mut event = TauriAgentEvent::ToolStart { + tool_name: "bash".to_string(), + tool_id: "tool-shell-patch-1".to_string(), + arguments: Some( + r#"{"command":"apply_patch <<'PATCH'\n*** Begin Patch\n*** Add File: notes/live.md\n+hello\n*** End Patch\nPATCH\n"}"# + .to_string(), + ), + }; + + let extras = emitter.process_event(&mut event); + + assert_eq!(extras.len(), 1); + assert_snapshot(&extras[0], "notes/live.md", "", false); + } + + #[test] + fn text_delta_write_file_stream_emits_incremental_snapshots() { + let mut emitter = WriteArtifactEventEmitter::new("session-1"); + let mut first = TauriAgentEvent::TextDelta { + text: "开始 Hello".to_string(), + }; + let mut second = TauriAgentEvent::TextDelta { + text: " world 完成".to_string(), + }; + + let first_extras = emitter.process_event(&mut first); + let second_extras = emitter.process_event(&mut second); + + assert_eq!(first_extras.len(), 1); + assert_eq!(second_extras.len(), 1); + let first_id = assert_snapshot(&first_extras[0], "notes/demo.md", "Hello", false); + let second_id = assert_snapshot(&second_extras[0], "notes/demo.md", "Hello world", true); + assert_eq!(first_id, second_id); + } + + #[test] + fn tool_end_emits_completed_snapshot_and_backfills_metadata() { + let mut emitter = WriteArtifactEventEmitter::new("session-1"); + let mut tool_start = TauriAgentEvent::ToolStart { + tool_name: "write_file".to_string(), + tool_id: "tool-1".to_string(), + arguments: Some(r##"{"path":"drafts/demo.md","content":"# 标题"}"##.to_string()), + }; + emitter.process_event(&mut tool_start); + + let mut tool_end = TauriAgentEvent::ToolEnd { + tool_id: "tool-1".to_string(), + result: TauriToolResult { + success: true, + output: "写入完成".to_string(), + error: None, + images: None, + metadata: None, + }, + }; + + let extras = emitter.process_event(&mut tool_end); + assert_eq!(extras.len(), 1); + let artifact_id = assert_snapshot(&extras[0], "drafts/demo.md", "# 标题", true); + + match &tool_end { + TauriAgentEvent::ToolEnd { result, .. } => { + let metadata = result.metadata.as_ref().expect("tool_end metadata"); + assert_eq!( + metadata.get("artifact_id").and_then(Value::as_str), + Some(artifact_id.as_str()) + ); + assert_eq!( + metadata.get("path").and_then(Value::as_str), + Some("drafts/demo.md") + ); + assert_eq!( + metadata.get("artifact_streamed").and_then(Value::as_bool), + Some(true) + ); + } + _ => panic!("expected tool_end"), + } + } +} diff --git a/src-tauri/crates/core/src/app_paths.rs b/src-tauri/crates/core/src/app_paths.rs index 0dfed819d..7f2578076 100644 --- a/src-tauri/crates/core/src/app_paths.rs +++ b/src-tauri/crates/core/src/app_paths.rs @@ -62,6 +62,21 @@ pub fn resolve_skills_dir() -> Result { resolve_runtime_subdir("skills") } +pub fn resolve_project_skills_dir() -> Option { + std::env::current_dir() + .ok() + .map(|cwd| resolve_project_skills_dir_from_cwd(&cwd)) +} + +pub fn resolve_proxycast_skill_roots() -> Result, String> { + let mut roots = Vec::new(); + if let Some(project_dir) = resolve_project_skills_dir() { + roots.push(project_dir); + } + roots.push(resolve_skills_dir()?); + Ok(roots) +} + pub fn resolve_user_memory_path() -> Result { with_app_roots(resolve_user_memory_path_from_roots) } @@ -98,6 +113,10 @@ fn fallback_runtime_subdir(subdir: &str) -> PathBuf { fallback_app_data_dir().join(subdir) } +fn resolve_project_skills_dir_from_cwd(cwd: &Path) -> PathBuf { + cwd.join(".agents").join("skills") +} + fn fallback_app_data_dir() -> PathBuf { std::env::temp_dir().join(APP_DATA_DIR_NAME) } @@ -561,6 +580,13 @@ mod tests { ); } + #[test] + fn resolve_project_skills_dir_from_cwd_builds_agents_skills_path() { + let cwd = Path::new("/tmp/workspace"); + let resolved = resolve_project_skills_dir_from_cwd(cwd); + assert_eq!(resolved, cwd.join(".agents").join("skills")); + } + #[test] fn resolve_user_memory_path_copies_legacy_agents_file() { let temp = tempdir().unwrap(); diff --git a/src-tauri/crates/core/src/database/dao/agent_runtime_queue.rs b/src-tauri/crates/core/src/database/dao/agent_runtime_queue.rs new file mode 100644 index 000000000..e15e61ec5 --- /dev/null +++ b/src-tauri/crates/core/src/database/dao/agent_runtime_queue.rs @@ -0,0 +1,183 @@ +//! 统一运行时排队 turn 持久化 DAO +//! +//! 用于在应用重启后恢复会话级排队请求。 + +use rusqlite::{params, Connection}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct AgentRuntimeQueuedTurnRecord { + pub id: i64, + pub queued_turn_id: String, + pub session_id: String, + pub event_name: String, + pub message_preview: String, + pub message_text: String, + pub payload_json: String, + pub image_count: usize, + pub created_at: i64, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct NewAgentRuntimeQueuedTurnRecord { + pub queued_turn_id: String, + pub session_id: String, + pub event_name: String, + pub message_preview: String, + pub message_text: String, + pub payload_json: String, + pub image_count: usize, + pub created_at: i64, +} + +pub struct AgentRuntimeQueuedTurnDao; + +impl AgentRuntimeQueuedTurnDao { + pub fn insert( + conn: &Connection, + record: &NewAgentRuntimeQueuedTurnRecord, + ) -> Result<(), rusqlite::Error> { + conn.execute( + "INSERT INTO agent_runtime_queued_turns ( + queued_turn_id, + session_id, + event_name, + message_preview, + message_text, + payload_json, + image_count, + created_at + ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", + params![ + record.queued_turn_id, + record.session_id, + record.event_name, + record.message_preview, + record.message_text, + record.payload_json, + record.image_count as i64, + record.created_at, + ], + )?; + Ok(()) + } + + pub fn remove(conn: &Connection, queued_turn_id: &str) -> Result { + let changed = conn.execute( + "DELETE FROM agent_runtime_queued_turns WHERE queued_turn_id = ?1", + params![queued_turn_id], + )?; + Ok(changed > 0) + } + + pub fn list_by_session( + conn: &Connection, + session_id: &str, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT + id, + queued_turn_id, + session_id, + event_name, + message_preview, + message_text, + payload_json, + image_count, + created_at + FROM agent_runtime_queued_turns + WHERE session_id = ?1 + ORDER BY id ASC", + )?; + + let rows = stmt.query_map(params![session_id], |row| { + Ok(AgentRuntimeQueuedTurnRecord { + id: row.get(0)?, + queued_turn_id: row.get(1)?, + session_id: row.get(2)?, + event_name: row.get(3)?, + message_preview: row.get(4)?, + message_text: row.get(5)?, + payload_json: row.get(6)?, + image_count: row.get::<_, i64>(7)? as usize, + created_at: row.get(8)?, + }) + })?; + + rows.collect() + } + + pub fn list_distinct_session_ids(conn: &Connection) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT DISTINCT session_id + FROM agent_runtime_queued_turns + ORDER BY session_id ASC", + )?; + + let rows = stmt.query_map([], |row| row.get::<_, String>(0))?; + rows.collect() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn setup_conn() -> Connection { + let conn = Connection::open_in_memory().unwrap(); + conn.execute( + "CREATE TABLE agent_runtime_queued_turns ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + queued_turn_id TEXT NOT NULL UNIQUE, + session_id TEXT NOT NULL, + event_name TEXT NOT NULL, + message_preview TEXT NOT NULL, + message_text TEXT NOT NULL, + payload_json TEXT NOT NULL, + image_count INTEGER NOT NULL DEFAULT 0, + created_at INTEGER NOT NULL + )", + [], + ) + .unwrap(); + conn + } + + #[test] + fn should_insert_list_and_remove_queued_turn() { + let conn = setup_conn(); + let first = NewAgentRuntimeQueuedTurnRecord { + queued_turn_id: "queued-1".to_string(), + session_id: "session-1".to_string(), + event_name: "event-1".to_string(), + message_preview: "preview-1".to_string(), + message_text: "body-1".to_string(), + payload_json: "{\"message\":\"body-1\"}".to_string(), + image_count: 0, + created_at: 1, + }; + let second = NewAgentRuntimeQueuedTurnRecord { + queued_turn_id: "queued-2".to_string(), + session_id: "session-1".to_string(), + event_name: "event-2".to_string(), + message_preview: "preview-2".to_string(), + message_text: "body-2".to_string(), + payload_json: "{\"message\":\"body-2\"}".to_string(), + image_count: 2, + created_at: 2, + }; + + AgentRuntimeQueuedTurnDao::insert(&conn, &first).unwrap(); + AgentRuntimeQueuedTurnDao::insert(&conn, &second).unwrap(); + + let rows = AgentRuntimeQueuedTurnDao::list_by_session(&conn, "session-1").unwrap(); + assert_eq!(rows.len(), 2); + assert_eq!(rows[0].queued_turn_id, "queued-1"); + assert_eq!(rows[1].message_text, "body-2"); + + let session_ids = AgentRuntimeQueuedTurnDao::list_distinct_session_ids(&conn).unwrap(); + assert_eq!(session_ids, vec!["session-1".to_string()]); + + assert!(AgentRuntimeQueuedTurnDao::remove(&conn, "queued-1").unwrap()); + assert!(!AgentRuntimeQueuedTurnDao::remove(&conn, "missing").unwrap()); + } +} diff --git a/src-tauri/crates/core/src/database/dao/mod.rs b/src-tauri/crates/core/src/database/dao/mod.rs index 43a8ff816..10081c157 100644 --- a/src-tauri/crates/core/src/database/dao/mod.rs +++ b/src-tauri/crates/core/src/database/dao/mod.rs @@ -1,6 +1,7 @@ pub mod a2ui_form_dao; pub mod agent; pub mod agent_run; +pub mod agent_runtime_queue; pub mod agent_timeline; pub mod api_key_provider; pub mod brand_persona_dao; diff --git a/src-tauri/crates/core/src/database/schema.rs b/src-tauri/crates/core/src/database/schema.rs index 41961e6f0..2fa76aff9 100644 --- a/src-tauri/crates/core/src/database/schema.rs +++ b/src-tauri/crates/core/src/database/schema.rs @@ -543,6 +543,30 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> { [], )?; + // 统一运行时排队 turn 表 + // 持久化 pending 队列,用于应用重启后恢复会话级排队请求 + conn.execute( + "CREATE TABLE IF NOT EXISTS agent_runtime_queued_turns ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + queued_turn_id TEXT NOT NULL UNIQUE, + session_id TEXT NOT NULL, + event_name TEXT NOT NULL, + message_preview TEXT NOT NULL, + message_text TEXT NOT NULL, + payload_json TEXT NOT NULL, + image_count INTEGER NOT NULL DEFAULT 0, + created_at INTEGER NOT NULL, + FOREIGN KEY (session_id) REFERENCES agent_sessions(id) ON DELETE CASCADE + )", + [], + )?; + + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_agent_runtime_queued_turns_session + ON agent_runtime_queued_turns(session_id, id)", + [], + )?; + // ============================================================================ // General Chat 相关表 // ============================================================================ diff --git a/src-tauri/crates/core/src/models/mod.rs b/src-tauri/crates/core/src/models/mod.rs index 0993deec6..dccbeb7ea 100644 --- a/src-tauri/crates/core/src/models/mod.rs +++ b/src-tauri/crates/core/src/models/mod.rs @@ -38,11 +38,13 @@ pub use provider_model::Provider; pub use provider_pool_model::*; pub use provider_type::ProviderType; pub use skill_model::{ - resolve_skill_source_kind, Skill, SkillMetadata, SkillRepo, SkillSourceKind, SkillState, - SkillStates, BROADCAST_GENERATE_SKILL_DIRECTORY, COVER_GENERATE_SKILL_DIRECTORY, - DEFAULT_PROXYCAST_SKILL_DIRECTORIES, IMAGE_GENERATE_SKILL_DIRECTORY, LIBRARY_SKILL_DIRECTORY, - MODAL_RESOURCE_SEARCH_SKILL_DIRECTORY, RESEARCH_SKILL_DIRECTORY, - SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY, TYPESETTING_SKILL_DIRECTORY, URL_PARSE_SKILL_DIRECTORY, - VIDEO_GENERATE_SKILL_DIRECTORY, + parse_skill_manifest_from_content, resolve_skill_source_kind, split_skill_frontmatter, + summarize_skill_resources_dir, ParsedSkillManifest, Skill, SkillCatalogSource, SkillMetadata, + SkillPackageInspection, SkillRepo, SkillResourceSummary, SkillSourceKind, + SkillStandardCompliance, SkillState, SkillStates, BROADCAST_GENERATE_SKILL_DIRECTORY, + COVER_GENERATE_SKILL_DIRECTORY, DEFAULT_PROXYCAST_SKILL_DIRECTORIES, + IMAGE_GENERATE_SKILL_DIRECTORY, LIBRARY_SKILL_DIRECTORY, MODAL_RESOURCE_SEARCH_SKILL_DIRECTORY, + RESEARCH_SKILL_DIRECTORY, SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY, TYPESETTING_SKILL_DIRECTORY, + URL_PARSE_SKILL_DIRECTORY, VIDEO_GENERATE_SKILL_DIRECTORY, }; pub use vertex_model::{VertexApiKeyEntry, VertexModelAlias}; diff --git a/src-tauri/crates/core/src/models/skill_model.rs b/src-tauri/crates/core/src/models/skill_model.rs index 141c6f8e6..7b0cc8784 100644 --- a/src-tauri/crates/core/src/models/skill_model.rs +++ b/src-tauri/crates/core/src/models/skill_model.rs @@ -2,6 +2,24 @@ use super::app_type::AppType; use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; +use std::path::Path; + +const SKILL_FRONTMATTER_NAME: &str = "name"; +const SKILL_FRONTMATTER_DESCRIPTION: &str = "description"; +const SKILL_FRONTMATTER_LICENSE: &str = "license"; +const SKILL_FRONTMATTER_METADATA: &str = "metadata"; +const SKILL_FRONTMATTER_ALLOWED_TOOLS: &str = "allowed-tools"; +const SKILL_FRONTMATTER_ALLOWED_TOOLS_ALIAS: &str = "allowed_tools"; +const LEGACY_PROXYCAST_TOP_LEVEL_FIELDS: &[&str] = &[ + "argument-hint", + "argument_hint", + "when-to-use", + "when_to_use", + "execution-mode", + "steps-json", + "provider", + "disable-model-invocation", +]; pub const VIDEO_GENERATE_SKILL_DIRECTORY: &str = "video_generate"; pub const BROADCAST_GENERATE_SKILL_DIRECTORY: &str = "broadcast_generate"; @@ -34,6 +52,14 @@ pub enum SkillSourceKind { Other, } +#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "lowercase")] +pub enum SkillCatalogSource { + Project, + User, + Remote, +} + #[derive(Debug, Clone, Serialize, Deserialize)] pub struct Skill { pub key: String, @@ -45,12 +71,28 @@ pub struct Skill { pub installed: bool, #[serde(rename = "sourceKind")] pub source_kind: SkillSourceKind, + #[serde(rename = "catalogSource")] + pub catalog_source: SkillCatalogSource, #[serde(rename = "repoOwner", skip_serializing_if = "Option::is_none")] pub repo_owner: Option, #[serde(rename = "repoName", skip_serializing_if = "Option::is_none")] pub repo_name: Option, #[serde(rename = "repoBranch", skip_serializing_if = "Option::is_none")] pub repo_branch: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub license: Option, + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub metadata: HashMap, + #[serde( + rename = "allowedTools", + default, + skip_serializing_if = "Vec::is_empty" + )] + pub allowed_tools: Vec, + #[serde(rename = "resourceSummary", skip_serializing_if = "Option::is_none")] + pub resource_summary: Option, + #[serde(rename = "standardCompliance", skip_serializing_if = "Option::is_none")] + pub standard_compliance: Option, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -71,6 +113,309 @@ pub struct SkillState { pub struct SkillMetadata { pub name: Option, pub description: Option, + pub license: Option, + #[serde(default)] + pub metadata: HashMap, + #[serde(default)] + pub allowed_tools: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)] +pub struct SkillResourceSummary { + #[serde(rename = "hasScripts")] + pub has_scripts: bool, + #[serde(rename = "hasReferences")] + pub has_references: bool, + #[serde(rename = "hasAssets")] + pub has_assets: bool, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)] +pub struct SkillStandardCompliance { + #[serde(rename = "isStandard")] + pub is_standard: bool, + #[serde( + rename = "validationErrors", + default, + skip_serializing_if = "Vec::is_empty" + )] + pub validation_errors: Vec, + #[serde( + rename = "deprecatedFields", + default, + skip_serializing_if = "Vec::is_empty" + )] + pub deprecated_fields: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)] +pub struct SkillPackageInspection { + pub content: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub license: Option, + #[serde(default)] + pub metadata: HashMap, + #[serde(rename = "allowedTools", default)] + pub allowed_tools: Vec, + #[serde(rename = "resourceSummary")] + pub resource_summary: SkillResourceSummary, + #[serde(rename = "standardCompliance")] + pub standard_compliance: SkillStandardCompliance, +} + +#[derive(Debug, Clone)] +pub struct ParsedSkillManifest { + pub metadata: SkillMetadata, + pub compliance: SkillStandardCompliance, + pub raw_frontmatter: serde_yaml::Value, +} + +impl ParsedSkillManifest { + pub fn metadata_value(&self, key: &str) -> Option<&str> { + self.metadata.metadata.get(key).map(|value| value.as_str()) + } + + pub fn raw_string(&self, key: &str) -> Option { + let mapping = self.raw_frontmatter.as_mapping()?; + yaml_mapping_get(mapping, key).and_then(yaml_scalar_to_string) + } + + pub fn raw_bool(&self, key: &str) -> Option { + let mapping = self.raw_frontmatter.as_mapping()?; + yaml_mapping_get(mapping, key).and_then(yaml_scalar_to_bool) + } +} + +pub fn split_skill_frontmatter(content: &str) -> Option<(&str, &str)> { + let content = content.trim_start_matches('\u{feff}'); + let regex = regex::Regex::new(r"(?s)\A---\s*\n(?P.*?)\n---\s*(?:\n|$)").ok()?; + let captures = regex.captures(content)?; + let frontmatter = captures.name("frontmatter")?.as_str(); + let body_start = captures.get(0)?.end(); + let body = content.get(body_start..).unwrap_or(""); + Some((frontmatter, body)) +} + +pub fn parse_skill_manifest_from_content(content: &str) -> Result { + let Some((frontmatter, _body)) = split_skill_frontmatter(content) else { + return Ok(ParsedSkillManifest { + metadata: SkillMetadata { + name: None, + description: None, + license: None, + metadata: HashMap::new(), + allowed_tools: Vec::new(), + }, + compliance: SkillStandardCompliance { + is_standard: false, + validation_errors: vec!["缺少以 --- 包裹的 YAML frontmatter".to_string()], + deprecated_fields: Vec::new(), + }, + raw_frontmatter: serde_yaml::Value::Null, + }); + }; + + let raw_frontmatter = serde_yaml::from_str::(frontmatter) + .map_err(|error| format!("解析 YAML frontmatter 失败: {error}"))?; + let mapping = raw_frontmatter + .as_mapping() + .ok_or_else(|| "YAML frontmatter 顶层必须是对象".to_string())?; + + let mut validation_errors = Vec::new(); + let mut deprecated_fields = Vec::new(); + + let name = required_string_field( + mapping, + SKILL_FRONTMATTER_NAME, + &mut validation_errors, + "缺少必填字段 `name`", + ); + let description = required_string_field( + mapping, + SKILL_FRONTMATTER_DESCRIPTION, + &mut validation_errors, + "缺少必填字段 `description`", + ); + let license = optional_string_field(mapping, SKILL_FRONTMATTER_LICENSE, &mut validation_errors); + let allowed_tools = parse_allowed_tools_field(mapping, &mut validation_errors); + let metadata = parse_metadata_field(mapping, &mut validation_errors); + + for field in LEGACY_PROXYCAST_TOP_LEVEL_FIELDS { + if yaml_mapping_get(mapping, field).is_some() { + deprecated_fields.push((*field).to_string()); + } + } + + deprecated_fields.sort(); + deprecated_fields.dedup(); + + let compliance = SkillStandardCompliance { + is_standard: validation_errors.is_empty(), + validation_errors, + deprecated_fields, + }; + + Ok(ParsedSkillManifest { + metadata: SkillMetadata { + name, + description, + license, + metadata, + allowed_tools, + }, + compliance, + raw_frontmatter, + }) +} + +pub fn summarize_skill_resources_dir(skill_dir: &Path) -> SkillResourceSummary { + SkillResourceSummary { + has_scripts: skill_dir.join("scripts").is_dir(), + has_references: skill_dir.join("references").is_dir(), + has_assets: skill_dir.join("assets").is_dir(), + } +} + +fn required_string_field( + mapping: &serde_yaml::Mapping, + key: &str, + validation_errors: &mut Vec, + missing_error: &str, +) -> Option { + match yaml_mapping_get(mapping, key) { + Some(value) => match yaml_scalar_to_string(value) { + Some(parsed) if !parsed.trim().is_empty() => Some(parsed), + Some(_) => { + validation_errors.push(format!("字段 `{key}` 不能为空")); + None + } + None => { + validation_errors.push(format!("字段 `{key}` 必须是字符串")); + None + } + }, + None => { + validation_errors.push(missing_error.to_string()); + None + } + } +} + +fn optional_string_field( + mapping: &serde_yaml::Mapping, + key: &str, + validation_errors: &mut Vec, +) -> Option { + let value = yaml_mapping_get(mapping, key)?; + + match yaml_scalar_to_string(value) { + Some(parsed) if !parsed.trim().is_empty() => Some(parsed), + Some(_) => None, + None => { + validation_errors.push(format!("字段 `{key}` 必须是字符串")); + None + } + } +} + +fn parse_allowed_tools_field( + mapping: &serde_yaml::Mapping, + validation_errors: &mut Vec, +) -> Vec { + let Some(value) = yaml_mapping_get(mapping, SKILL_FRONTMATTER_ALLOWED_TOOLS) + .or_else(|| yaml_mapping_get(mapping, SKILL_FRONTMATTER_ALLOWED_TOOLS_ALIAS)) + else { + return Vec::new(); + }; + + match value { + serde_yaml::Value::String(single) => split_allowed_tools_csv(single), + serde_yaml::Value::Sequence(values) => { + let mut tools = Vec::new(); + for item in values { + match yaml_scalar_to_string(item) { + Some(tool) if !tool.trim().is_empty() => tools.push(tool), + _ => validation_errors + .push("字段 `allowed-tools` 只能包含字符串条目".to_string()), + } + } + tools + } + _ => { + validation_errors.push("字段 `allowed-tools` 必须是字符串或字符串数组".to_string()); + Vec::new() + } + } +} + +fn split_allowed_tools_csv(value: &str) -> Vec { + value + .split(',') + .map(|item| item.trim()) + .filter(|item| !item.is_empty()) + .map(ToString::to_string) + .collect() +} + +fn parse_metadata_field( + mapping: &serde_yaml::Mapping, + validation_errors: &mut Vec, +) -> HashMap { + let Some(value) = yaml_mapping_get(mapping, SKILL_FRONTMATTER_METADATA) else { + return HashMap::new(); + }; + + let Some(meta_mapping) = value.as_mapping() else { + validation_errors.push("字段 `metadata` 必须是键值对象".to_string()); + return HashMap::new(); + }; + + let mut metadata = HashMap::new(); + for (raw_key, raw_value) in meta_mapping { + let Some(key) = yaml_scalar_to_string(raw_key) else { + validation_errors.push("字段 `metadata` 的 key 必须是字符串".to_string()); + continue; + }; + + match yaml_scalar_to_string(raw_value) { + Some(value) => { + metadata.insert(key, value); + } + None => { + validation_errors.push(format!("字段 `metadata.{key}` 必须是字符串、数字或布尔值")) + } + } + } + + metadata +} + +fn yaml_mapping_get<'a>( + mapping: &'a serde_yaml::Mapping, + key: &str, +) -> Option<&'a serde_yaml::Value> { + mapping.get(serde_yaml::Value::String(key.to_string())) +} + +fn yaml_scalar_to_string(value: &serde_yaml::Value) -> Option { + match value { + serde_yaml::Value::String(value) => Some(value.clone()), + serde_yaml::Value::Number(value) => Some(value.to_string()), + serde_yaml::Value::Bool(value) => Some(value.to_string()), + _ => None, + } +} + +fn yaml_scalar_to_bool(value: &serde_yaml::Value) -> Option { + match value { + serde_yaml::Value::Bool(value) => Some(*value), + serde_yaml::Value::String(value) => match value.trim().to_ascii_lowercase().as_str() { + "true" | "1" | "yes" | "enabled" => Some(true), + "false" | "0" | "no" | "disabled" => Some(false), + _ => None, + }, + _ => None, + } } impl Default for SkillRepo { diff --git a/src-tauri/crates/gateway/src/discord.rs b/src-tauri/crates/gateway/src/discord.rs index 040eafdb5..5fca438ce 100644 --- a/src-tauri/crates/gateway/src/discord.rs +++ b/src-tauri/crates/gateway/src/discord.rs @@ -1408,7 +1408,7 @@ async fn handle_plain_text( logs.write().await.add( "info", &format!( - "[DiscordGateway] account={} 收到文本消息: guild={:?} channel={} sender={} messageId={} session={} model={} web_search=true", + "[DiscordGateway] account={} 收到文本消息: guild={:?} channel={} sender={} messageId={} session={} model={} search_mode=allowed", account.account_id, inbound.guild_id, inbound.channel_id, @@ -1430,6 +1430,7 @@ async fn handle_plain_text( "stream": false, "model": account.default_model.clone(), "web_search": true, + "search_mode": "allowed", })), }) .await; @@ -1698,7 +1699,8 @@ fn build_rpc_request(command: DiscordCommand) -> Result ( diff --git a/src-tauri/crates/gateway/src/feishu.rs b/src-tauri/crates/gateway/src/feishu.rs index 616896bf6..006b56a87 100644 --- a/src-tauri/crates/gateway/src/feishu.rs +++ b/src-tauri/crates/gateway/src/feishu.rs @@ -1353,7 +1353,7 @@ async fn handle_plain_text( logs.write().await.add( "info", &format!( - "[FeishuGateway] account={} 收到文本消息: chat={} sender={:?} messageId={} session={} model={} web_search=true", + "[FeishuGateway] account={} 收到文本消息: chat={} sender={:?} messageId={} session={} model={} search_mode=allowed", account.account_id, inbound.chat_id, inbound.sender_id, @@ -1374,6 +1374,7 @@ async fn handle_plain_text( "stream": false, "model": account.default_model.clone(), "web_search": true, + "search_mode": "allowed", })), }) .await; @@ -1678,7 +1679,8 @@ fn build_rpc_request(command: FeishuCommand) -> Result ( diff --git a/src-tauri/crates/gateway/src/telegram.rs b/src-tauri/crates/gateway/src/telegram.rs index aac854b85..7fc23e7c9 100644 --- a/src-tauri/crates/gateway/src/telegram.rs +++ b/src-tauri/crates/gateway/src/telegram.rs @@ -1119,7 +1119,7 @@ async fn handle_plain_text_with_mode( logs.write().await.add( "info", &format!( - "[TelegramGateway] account={} 收到文本消息: chat={} sender={:?} messageId={} session={} streaming={:?} model={} web_search=true", + "[TelegramGateway] account={} 收到文本消息: chat={} sender={:?} messageId={} session={} streaming={:?} model={} search_mode=allowed", account.account_id, inbound.chat_id, inbound.sender_id, @@ -1168,6 +1168,7 @@ async fn handle_plain_text_with_mode( "stream": false, "model": account.default_model.clone(), "web_search": true, + "search_mode": "allowed", })), }) .await; @@ -1597,7 +1598,8 @@ fn build_rpc_request(command: TelegramCommand) -> Result ( diff --git a/src-tauri/crates/services/src/api_key_provider_service.rs b/src-tauri/crates/services/src/api_key_provider_service.rs index 76d06a2fd..c0c4eda99 100644 --- a/src-tauri/crates/services/src/api_key_provider_service.rs +++ b/src-tauri/crates/services/src/api_key_provider_service.rs @@ -46,6 +46,7 @@ mod tests { use proxycast_core::database::dao::api_key_provider::ApiProviderType; use proxycast_core::database::init_database; use rusqlite::OptionalExtension; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; fn resolve_real_codex_provider_id( db: &proxycast_core::database::DbConnection, @@ -115,14 +116,14 @@ data: [DONE]\n"; let with_explicit = ApiKeyProviderService::pick_test_model( Some("explicit-model".to_string()), &["custom-model".to_string()], - &["fallback-model".to_string()], + &["fallback-model".to_string(), "fallback-mini".to_string()], ); assert_eq!(with_explicit.as_deref(), Some("explicit-model")); let with_custom = ApiKeyProviderService::pick_test_model( None, &["custom-model".to_string()], - &["fallback-model".to_string()], + &["fallback-model".to_string(), "fallback-mini".to_string()], ); assert_eq!(with_custom.as_deref(), Some("custom-model")); @@ -130,10 +131,239 @@ data: [DONE]\n"; ApiKeyProviderService::pick_test_model(None, &[], &["fallback-model".to_string()]); assert_eq!(with_local_fallback.as_deref(), Some("fallback-model")); + let with_preferred_fallback = ApiKeyProviderService::pick_test_model( + None, + &[], + &[ + "gpt-5.2-pro".to_string(), + "gpt-5-nano".to_string(), + "gpt-4.1-mini".to_string(), + ], + ); + assert_eq!(with_preferred_fallback.as_deref(), Some("gpt-5-nano")); + let none = ApiKeyProviderService::pick_test_model(None, &[], &[]); assert!(none.is_none()); } + #[test] + fn test_parse_openai_responses_content_prefers_output_text() { + let body = serde_json::json!({ + "id": "resp_test", + "output_text": "hello from responses", + "output": [{ + "type": "message", + "content": [{"type": "output_text", "text": "fallback text"}] + }] + }) + .to_string(); + + let content = ApiKeyProviderService::parse_openai_responses_content(&body) + .expect("应解析 output_text"); + assert_eq!(content, "hello from responses"); + } + + #[test] + fn test_parse_openai_responses_content_reads_output_blocks() { + let body = serde_json::json!({ + "id": "resp_test", + "output": [{ + "type": "message", + "content": [ + {"type": "output_text", "text": "hello"}, + {"type": "output_text", "text": " world"} + ] + }] + }) + .to_string(); + + let content = ApiKeyProviderService::parse_openai_responses_content(&body) + .expect("应解析 output.content"); + assert_eq!(content, "hello world"); + } + + async fn spawn_single_response_server( + response_body: serde_json::Value, + ) -> (String, tokio::task::JoinHandle<(String, String)>) { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("绑定测试服务失败"); + let addr = listener.local_addr().expect("读取测试地址失败"); + + let server = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("接受连接失败"); + let mut buffer = Vec::new(); + let mut header_end = None; + + loop { + let mut chunk = [0u8; 1024]; + let read = stream.read(&mut chunk).await.expect("读取请求失败"); + if read == 0 { + break; + } + buffer.extend_from_slice(&chunk[..read]); + if let Some(pos) = buffer.windows(4).position(|w| w == b"\r\n\r\n") { + header_end = Some(pos + 4); + break; + } + } + + let header_end = header_end.expect("请求头未结束"); + let header_text = String::from_utf8_lossy(&buffer[..header_end]).to_string(); + let content_length = header_text + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + if name.eq_ignore_ascii_case("content-length") { + value.trim().parse::().ok() + } else { + None + } + }) + .unwrap_or(0); + + while buffer.len() < header_end + content_length { + let mut chunk = vec![0u8; content_length.max(1024)]; + let read = stream.read(&mut chunk).await.expect("补全请求体失败"); + if read == 0 { + break; + } + buffer.extend_from_slice(&chunk[..read]); + } + + let request_line = header_text.lines().next().expect("缺少请求行").to_string(); + let request_path = request_line + .split_whitespace() + .nth(1) + .expect("缺少请求路径") + .to_string(); + let request_body = + String::from_utf8_lossy(&buffer[header_end..header_end + content_length]) + .to_string(); + let response_text = response_body.to_string(); + let response = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + response_text.len(), + response_text + ); + stream + .write_all(response.as_bytes()) + .await + .expect("写回响应失败"); + + (request_path, request_body) + }); + + (format!("http://{}", addr), server) + } + + #[tokio::test] + #[ignore = "当前沙箱禁止本地 TCP 监听,需在本机放行后运行"] + async fn test_openai_response_chat_uses_responses_endpoint_and_fallback_model() { + let db = init_database().expect("初始化数据库失败"); + let service = ApiKeyProviderService::new(); + let response_body = serde_json::json!({ + "id": "resp_test", + "output_text": "OK" + }); + let (base_url, server) = spawn_single_response_server(response_body).await; + + let provider = service + .add_custom_provider( + &db, + "OpenAI Responses Test".to_string(), + ApiProviderType::OpenaiResponse, + base_url, + None, + None, + None, + None, + ) + .expect("创建 Provider 失败"); + service + .add_api_key(&db, &provider.id, "sk-test-key", Some("test".to_string())) + .expect("添加 API Key 失败"); + + let result = service + .test_chat_with_fallback_models( + &db, + &provider.id, + None, + "hello".to_string(), + vec!["gpt-4.1-mini".to_string()], + ) + .await + .expect("对话测试调用失败"); + + assert!(result.success, "结果应成功: {:?}", result.error); + assert_eq!(result.content.as_deref(), Some("OK")); + + let (request_path, request_body) = server.await.expect("等待测试服务失败"); + assert_eq!(request_path, "/v1/responses"); + let request_json: serde_json::Value = + serde_json::from_str(&request_body).expect("请求体应为 JSON"); + assert_eq!(request_json["model"].as_str(), Some("gpt-4.1-mini")); + assert_eq!(request_json["stream"].as_bool(), Some(false)); + assert_eq!( + request_json["input"][0]["content"][0]["text"].as_str(), + Some("hello") + ); + } + + #[tokio::test] + #[ignore = "当前沙箱禁止本地 TCP 监听,需在本机放行后运行"] + async fn test_openai_response_connection_prefers_responses_endpoint() { + let db = init_database().expect("初始化数据库失败"); + let service = ApiKeyProviderService::new(); + let response_body = serde_json::json!({ + "id": "resp_test", + "output_text": "OK" + }); + let (base_url, server) = spawn_single_response_server(response_body).await; + + let provider = service + .add_custom_provider( + &db, + "OpenAI Responses Test".to_string(), + ApiProviderType::OpenaiResponse, + base_url, + None, + None, + None, + None, + ) + .expect("创建 Provider 失败"); + service + .add_api_key(&db, &provider.id, "sk-test-key", Some("test".to_string())) + .expect("添加 API Key 失败"); + + let result = service + .test_connection_with_fallback_models( + &db, + &provider.id, + None, + vec!["gpt-4.1-mini".to_string()], + ) + .await + .expect("连接测试调用失败"); + + assert!(result.success, "连接测试应成功: {:?}", result.error); + assert_eq!( + result.models.as_deref(), + Some(&["gpt-4.1-mini".to_string()][..]) + ); + + let (request_path, request_body) = server.await.expect("等待测试服务失败"); + assert_eq!(request_path, "/v1/responses"); + let request_json: serde_json::Value = + serde_json::from_str(&request_body).expect("请求体应为 JSON"); + assert_eq!(request_json["model"].as_str(), Some("gpt-4.1-mini")); + assert_eq!( + request_json["input"][0]["content"][0]["text"].as_str(), + Some("hi") + ); + } + #[tokio::test] #[ignore = "真实联网测试:设置 PROXYCAST_REAL_API_TEST=1 后执行"] async fn test_real_codex_provider_chat_gpt_5_3_codex() { @@ -298,6 +528,18 @@ impl ApiKeyProviderService { provider_id: &str, model_name: Option, prompt: String, + ) -> Result { + self.test_chat_with_fallback_models(db, provider_id, model_name, prompt, Vec::new()) + .await + } + + pub async fn test_chat_with_fallback_models( + &self, + db: &DbConnection, + provider_id: &str, + model_name: Option, + prompt: String, + fallback_models: Vec, ) -> Result { use std::time::Instant; @@ -311,9 +553,9 @@ impl ApiKeyProviderService { .get_next_api_key(db, provider_id)? .ok_or_else(|| "没有可用的 API Key".to_string())?; - let test_model = model_name.or_else(|| provider.custom_models.first().cloned()); let test_model = - test_model.ok_or_else(|| "缺少模型名称:请在自定义模型中填写一个模型名".to_string())?; + Self::pick_test_model(model_name, &provider.custom_models, &fallback_models) + .ok_or_else(|| "缺少模型名称:请在自定义模型中填写一个模型名".to_string())?; let start = Instant::now(); @@ -329,6 +571,11 @@ impl ApiKeyProviderService { ) .await } + // OpenAI Responses API 走 /v1/responses + provider_type if Self::uses_openai_responses_protocol(provider_type) => { + self.test_openai_responses_once(&api_key, &provider.api_host, &test_model, &prompt) + .await + } // Anthropic / AnthropicCompatible 统一走 /v1/messages provider_type if Self::uses_anthropic_protocol(provider_type) => { self.test_anthropic_chat_once(&api_key, &provider.api_host, &test_model, &prompt) @@ -365,6 +612,27 @@ impl ApiKeyProviderService { provider_type.is_anthropic_protocol() } + #[inline] + fn uses_openai_responses_protocol(provider_type: ApiProviderType) -> bool { + matches!(provider_type, ApiProviderType::OpenaiResponse) + } + + fn pick_preferred_fallback_model(fallback_models: &[String]) -> Option { + const PREFERRED_MARKERS: &[&str] = + &["nano", "mini", "flash-lite", "flash", "haiku", "lite"]; + + for marker in PREFERRED_MARKERS { + if let Some(model) = fallback_models + .iter() + .find(|model| model.to_ascii_lowercase().contains(marker)) + { + return Some(model.clone()); + } + } + + fallback_models.first().cloned() + } + fn pick_test_model( model_name: Option, custom_models: &[String], @@ -372,7 +640,7 @@ impl ApiKeyProviderService { ) -> Option { model_name .or_else(|| custom_models.first().cloned()) - .or_else(|| fallback_models.first().cloned()) + .or_else(|| Self::pick_preferred_fallback_model(fallback_models)) } async fn test_openai_chat_once( @@ -442,7 +710,7 @@ impl ApiKeyProviderService { let body2 = resp2.text().await.unwrap_or_default(); if !status2.is_success() { - return Err(format!("API 返回错误: {status2} - {body2}")); + return Err(Self::format_http_api_error(status2, &body2)); } let content = Self::parse_chat_completions_sse_content(&body2); @@ -456,7 +724,7 @@ impl ApiKeyProviderService { .await; } - Err(format!("API 返回错误: {status} - {body}")) + Err(Self::format_http_api_error(status, &body)) } async fn test_anthropic_chat_once( @@ -486,7 +754,7 @@ impl ApiKeyProviderService { let body = resp.text().await.unwrap_or_default(); if !status.is_success() { - return Err(format!("API 返回错误: {status} - {body}")); + return Err(Self::format_http_api_error(status, &body)); } let parsed: serde_json::Value = @@ -530,7 +798,11 @@ impl ApiKeyProviderService { out } - fn build_codex_responses_request(model: &str, prompt: &str) -> serde_json::Value { + fn build_openai_responses_request( + model: &str, + prompt: &str, + stream: bool, + ) -> serde_json::Value { serde_json::json!({ "model": model, "input": [ @@ -544,11 +816,122 @@ impl ApiKeyProviderService { ] } ], - "stream": true, + "stream": stream, "max_output_tokens": 64 }) } + fn parse_openai_responses_content(body: &str) -> Result { + let parsed: serde_json::Value = + serde_json::from_str(body).map_err(|e| format!("解析响应失败: {e} - {body}"))?; + + if let Some(output_text) = parsed["output_text"].as_str() { + return Ok(output_text.to_string()); + } + + let mut content = String::new(); + if let Some(output) = parsed["output"].as_array() { + for output_item in output { + if output_item["type"].as_str() != Some("message") { + continue; + } + + if let Some(content_items) = output_item["content"].as_array() { + for item in content_items { + let item_type = item["type"].as_str().unwrap_or_default(); + if matches!(item_type, "output_text" | "text") { + if let Some(text) = item["text"].as_str() { + content.push_str(text); + } + } + } + } + } + } + + Ok(content) + } + + fn extract_json_error_message(body: &str) -> Option { + let parsed: serde_json::Value = serde_json::from_str(body).ok()?; + let error = parsed.get("error")?; + + if let Some(message) = error.get("message").and_then(|value| value.as_str()) { + return Some(message.to_string()); + } + + error.as_str().map(|value| value.to_string()) + } + + fn format_http_api_error(status: reqwest::StatusCode, body: &str) -> String { + match Self::extract_json_error_message(body) { + Some(message) if body.contains("insufficient_quota") => { + format!("API 返回错误: {status} - OpenAI 账户配额不足或未开通计费:{message}") + } + Some(message) => format!("API 返回错误: {status} - {message}"), + None => format!("API 返回错误: {status} - {body}"), + } + } + + fn build_codex_responses_request(model: &str, prompt: &str) -> serde_json::Value { + Self::build_openai_responses_request(model, prompt, true) + } + + async fn test_openai_responses_once( + &self, + api_key: &str, + api_host: &str, + model: &str, + prompt: &str, + ) -> Result<(String, String), String> { + use proxycast_providers::providers::codex::CodexProvider; + + let url = CodexProvider::build_responses_url(api_host); + let request_body = Self::build_openai_responses_request(model, prompt, false); + + let client = reqwest::Client::new(); + let resp = client + .post(&url) + .header("Authorization", format!("Bearer {api_key}")) + .header("Content-Type", "application/json") + .json(&request_body) + .send() + .await + .map_err(|e| format!("API 调用失败: {e}"))?; + + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + + if status.is_success() { + let content = Self::parse_openai_responses_content(&body)?; + return Ok((content, body)); + } + + if status.as_u16() == 400 && body.contains("Stream must be set to true") { + let streaming_request = Self::build_openai_responses_request(model, prompt, true); + let resp2 = client + .post(&url) + .header("Authorization", format!("Bearer {api_key}")) + .header("Content-Type", "application/json") + .header("Accept", "text/event-stream") + .json(&streaming_request) + .send() + .await + .map_err(|e| format!("API 调用失败: {e}"))?; + + let status2 = resp2.status(); + let body2 = resp2.text().await.unwrap_or_default(); + if !status2.is_success() { + return Err(Self::format_http_api_error(status2, &body2)); + } + + let content = Self::parse_codex_responses_sse_content(&body2); + return Ok((content, body2)); + } + + Err(Self::format_http_api_error(status, &body)) + } + /// 测试 Codex /responses 端点(用于不支持 messages 参数的上游) async fn test_codex_responses_endpoint( &self, @@ -557,15 +940,9 @@ impl ApiKeyProviderService { model: &str, prompt: &str, ) -> Result<(String, String), String> { - // 构建 /responses 端点 URL - let base = api_host.trim_end_matches('/'); - let url = if base.ends_with("/v1") { - format!("{base}/responses") - } else if base.ends_with("/openai") { - format!("{base}/v1/responses") - } else { - format!("{base}/v1/responses") - }; + use proxycast_providers::providers::codex::CodexProvider; + + let url = CodexProvider::build_responses_url(api_host); // Codex Responses 格式请求体(input 必须是列表) let request_body = Self::build_codex_responses_request(model, prompt); @@ -584,7 +961,7 @@ impl ApiKeyProviderService { let body = resp.text().await.unwrap_or_default(); if !status.is_success() { - return Err(format!("API 返回错误: {status} - {body}")); + return Err(Self::format_http_api_error(status, &body)); } // 解析 Codex SSE 响应 @@ -1789,6 +2166,22 @@ impl ApiKeyProviderService { .await .map(|_| vec![test_model]) } + provider_type if Self::uses_openai_responses_protocol(provider_type) => { + let test_model = Self::pick_test_model( + model_name.clone(), + &provider.custom_models, + &fallback_models, + ); + + if let Some(test_model) = test_model { + self.test_openai_responses_once(&api_key, &provider.api_host, &test_model, "hi") + .await + .map(|_| vec![test_model]) + } else { + self.test_openai_models_endpoint(&api_key, &provider.api_host) + .await + } + } _ => { // OpenAI 兼容类型,优先使用 /models 端点 eprintln!("[TEST_CONNECTION] model_name param: {model_name:?}"); diff --git a/src-tauri/crates/services/src/skill_service.rs b/src-tauri/crates/services/src/skill_service.rs index 827b483e7..7f91b36e6 100644 --- a/src-tauri/crates/services/src/skill_service.rs +++ b/src-tauri/crates/services/src/skill_service.rs @@ -3,13 +3,17 @@ use parking_lot::{Mutex, RwLock}; use reqwest::Client; use std::collections::HashMap; use std::fs; -use std::path::{Path, PathBuf}; +use std::path::{Component, Path, PathBuf}; use std::sync::Arc; use std::time::{Duration, Instant}; use tokio::time::timeout; +use proxycast_core::app_paths; use proxycast_core::models::{ - resolve_skill_source_kind, AppType, Skill, SkillMetadata, SkillRepo, SkillState, + parse_skill_manifest_from_content as parse_manifest_content, resolve_skill_source_kind, + summarize_skill_resources_dir, AppType, ParsedSkillManifest, Skill, SkillCatalogSource, + SkillPackageInspection, SkillRepo, SkillResourceSummary, SkillSourceKind, + SkillStandardCompliance, SkillState, }; const DOWNLOAD_TIMEOUT: Duration = Duration::from_secs(60); @@ -70,6 +74,20 @@ impl RepoCacheEntry { } } +struct RemoteSkillArchiveEntry { + directory: String, + content: String, + readme_parent: String, + files: HashMap>, + resource_summary: SkillResourceSummary, +} + +#[derive(Debug, Clone)] +struct CatalogSkillRoot { + source: SkillCatalogSource, + path: PathBuf, +} + pub struct SkillService { client: Client, repo_cache: RwLock>, @@ -90,76 +108,63 @@ impl SkillService { }) } - /// 获取技能安装目录 fn get_skills_dir(app_type: &AppType) -> Result { - let home = dirs::home_dir().ok_or_else(|| anyhow!("Failed to get home directory"))?; - let skills_dir = match app_type { - AppType::Claude => home.join(".claude").join("skills"), - AppType::Codex => home.join(".codex").join("skills"), - AppType::Gemini => home.join(".gemini").join("skills"), - AppType::ProxyCast => home.join(".proxycast").join("skills"), + AppType::ProxyCast => app_paths::resolve_skills_dir().map_err(|e| anyhow!(e))?, + AppType::Claude => dirs::home_dir() + .ok_or_else(|| anyhow!("Failed to get home directory"))? + .join(".claude") + .join("skills"), + AppType::Codex => dirs::home_dir() + .ok_or_else(|| anyhow!("Failed to get home directory"))? + .join(".codex") + .join("skills"), + AppType::Gemini => dirs::home_dir() + .ok_or_else(|| anyhow!("Failed to get home directory"))? + .join(".gemini") + .join("skills"), }; Ok(skills_dir) } - /// 仅列出内置 + 本地技能(不访问远程仓库,速度快) + fn get_catalog_roots(app_type: &AppType) -> Result> { + match app_type { + AppType::ProxyCast => { + let mut roots = Vec::new(); + if let Some(project_dir) = app_paths::resolve_project_skills_dir() { + roots.push(CatalogSkillRoot { + source: SkillCatalogSource::Project, + path: project_dir, + }); + } + roots.push(CatalogSkillRoot { + source: SkillCatalogSource::User, + path: app_paths::resolve_skills_dir().map_err(|e| anyhow!(e))?, + }); + Ok(roots) + } + _ => Ok(vec![CatalogSkillRoot { + source: SkillCatalogSource::User, + path: Self::get_skills_dir(app_type)?, + }]), + } + } + pub fn list_local_skills( &self, app_type: &AppType, _installed_states: &HashMap, ) -> Result> { let mut all_skills: HashMap = HashMap::new(); - - // 扫描本地目录 - let skills_dir = Self::get_skills_dir(app_type)?; - if skills_dir.exists() { - if let Ok(entries) = fs::read_dir(&skills_dir) { - for entry in entries.flatten() { - if entry.path().is_dir() { - let directory = entry.file_name().to_string_lossy().to_string(); - let key = format!("local:{directory}"); - let skill_md = entry.path().join("SKILL.md"); - let (name, description) = if skill_md.exists() { - self.parse_skill_metadata(&skill_md) - .map(|m| { - ( - m.name.unwrap_or_else(|| directory.clone()), - m.description.unwrap_or_default(), - ) - }) - .unwrap_or_else(|_| (directory.clone(), String::new())) - } else { - (directory.clone(), String::new()) - }; - - all_skills.insert( - key.clone(), - Skill { - key, - name, - description, - directory: directory.clone(), - readme_url: None, - installed: true, - source_kind: resolve_skill_source_kind(app_type, &directory), - repo_owner: None, - repo_name: None, - repo_branch: None, - }, - ); - } - } - } - } + let roots = Self::get_catalog_roots(app_type)?; + self.collect_local_skills(app_type, &roots, &mut all_skills)?; let mut skills: Vec = all_skills.into_values().collect(); skills.sort_by(|a, b| a.name.cmp(&b.name)); Ok(skills) } - /// 列出所有技能 pub async fn list_skills( &self, app_type: &AppType, @@ -167,14 +172,19 @@ impl SkillService { installed_states: &HashMap, ) -> Result> { let mut all_skills: HashMap = HashMap::new(); + let roots = Self::get_catalog_roots(app_type)?; + self.collect_local_skills(app_type, &roots, &mut all_skills)?; - // 1. 从启用的仓库获取技能 - let enabled_repos: Vec<_> = repos.iter().filter(|r| r.enabled).collect(); - - for repo in enabled_repos { + for repo in repos.iter().filter(|repo| repo.enabled) { match timeout(DOWNLOAD_TIMEOUT, self.fetch_skills_from_repo_cached(repo)).await { Ok(Ok(remote_skills)) => { for mut skill in remote_skills { + if all_skills + .values() + .any(|existing| existing.directory == skill.directory) + { + continue; + } let app_key = format!( "{}:{}", app_type.to_string().to_lowercase(), @@ -187,75 +197,83 @@ impl SkillService { all_skills.insert(skill.key.clone(), skill); } } - Ok(Err(e)) => { + Ok(Err(error)) => { tracing::warn!( "Failed to fetch skills from {}/{}: {}", repo.owner, repo.name, - e + error ); } Err(_) => { - tracing::warn!("Timeout fetching skills from {}/{}", repo.owner, repo.name); + tracing::warn!("Timeout fetching skills from {}/{}", repo.owner, repo.name) } } } - // 2. 添加本地已安装但不在任何仓库中的技能 - let skills_dir = Self::get_skills_dir(app_type)?; - if skills_dir.exists() { - if let Ok(entries) = fs::read_dir(&skills_dir) { - for entry in entries.flatten() { - if entry.path().is_dir() { - let directory = entry.file_name().to_string_lossy().to_string(); - - // 检查是否已有相同 directory 的 skill(按 directory 去重) - let already_exists = all_skills.values().any(|s| s.directory == directory); - - if !already_exists { - let key = format!("local:{directory}"); - let skill_md = entry.path().join("SKILL.md"); - let (name, description) = if skill_md.exists() { - self.parse_skill_metadata(&skill_md) - .map(|m| { - ( - m.name.unwrap_or_else(|| directory.clone()), - m.description.unwrap_or_default(), - ) - }) - .unwrap_or_else(|_| (directory.clone(), String::new())) - } else { - (directory.clone(), String::new()) - }; - - all_skills.insert( - key.clone(), - Skill { - key, - name, - description, - directory: directory.clone(), - readme_url: None, - installed: true, - source_kind: resolve_skill_source_kind(app_type, &directory), - repo_owner: None, - repo_name: None, - repo_branch: None, - }, - ); - } - } - } - } - } - - // 3. 排序并返回 let mut skills: Vec = all_skills.into_values().collect(); skills.sort_by(|a, b| a.name.cmp(&b.name)); - Ok(skills) } + fn collect_local_skills( + &self, + app_type: &AppType, + roots: &[CatalogSkillRoot], + all_skills: &mut HashMap, + ) -> Result<()> { + for root in roots { + if !root.path.exists() { + continue; + } + + for entry in fs::read_dir(&root.path) + .with_context(|| { + format!("Failed to read skills directory {}", root.path.display()) + })? + .flatten() + { + if !entry.path().is_dir() { + continue; + } + + let directory = entry.file_name().to_string_lossy().to_string(); + if all_skills + .values() + .any(|skill| skill.directory == directory) + { + continue; + } + + let skill_md = entry.path().join("SKILL.md"); + if !skill_md.is_file() { + continue; + } + + let key = format!("local:{directory}"); + let resource_summary = summarize_skill_resources_dir(&entry.path()); + let source_kind = resolve_skill_source_kind(app_type, &directory); + + let skill = self.build_skill_from_file( + &skill_md, + key, + directory, + true, + source_kind, + root.source, + None, + None, + None, + None, + resource_summary, + )?; + all_skills.insert(skill.key.clone(), skill); + } + } + + Ok(()) + } + async fn fetch_skills_from_repo_cached(&self, repo: &SkillRepo) -> Result> { let cache_key = RepoCacheKey::from(repo); @@ -320,7 +338,6 @@ impl SkillService { }) } - /// 从仓库获取技能列表 async fn fetch_skills_from_repo_uncached(&self, repo: &SkillRepo) -> Result> { let mut last_error = None; @@ -354,6 +371,37 @@ impl SkillService { } async fn fetch_skills_from_branch(&self, repo: &SkillRepo, branch: &str) -> Result> { + let manifests = self.load_remote_skill_archive_entries(repo, branch).await?; + let mut skills = Vec::new(); + let repo_key_prefix = format!("{}/{}:", repo.owner, repo.name); + for entry in manifests.into_values() { + if entry.content.is_empty() { + continue; + } + let key = format!("{repo_key_prefix}{}", entry.directory); + let readme_url = Some(format!( + "https://github.com/{}/{}/blob/{}/{}/SKILL.md", + repo.owner, repo.name, branch, entry.readme_parent + )); + let skill = self.build_skill_from_remote_archive_entry( + entry, + key, + readme_url, + Some(repo.owner.clone()), + Some(repo.name.clone()), + Some(branch.to_string()), + ); + skills.push(skill); + } + + Ok(skills) + } + + async fn load_remote_skill_archive_entries( + &self, + repo: &SkillRepo, + branch: &str, + ) -> Result> { let zip_url = format!( "https://github.com/{}/{}/archive/refs/heads/{}.zip", repo.owner, repo.name, branch @@ -374,57 +422,113 @@ impl SkillService { let cursor = std::io::Cursor::new(bytes); let mut archive = zip::ZipArchive::new(cursor).context("Failed to open ZIP archive")?; - let mut skills = Vec::new(); - let repo_key_prefix = format!("{}/{}:", repo.owner, repo.name); + let mut manifests: HashMap = HashMap::new(); + for index in 0..archive.len() { + let mut file = archive + .by_index(index) + .context("Failed to read ZIP entry")?; + let archive_path = Path::new(file.name()); - for i in 0..archive.len() { - let mut file = archive.by_index(i).context("Failed to read ZIP entry")?; - let file_path = file.name().to_string(); - - if file_path.ends_with("/SKILL.md") || file_path.ends_with("\\SKILL.md") { - let path = Path::new(&file_path); - let directory = path - .parent() - .and_then(|p| p.file_name()) - .and_then(|n| n.to_str()) - .unwrap_or("unknown") - .to_string(); - - let mut content = String::new(); - use std::io::Read; - file.read_to_string(&mut content) - .context("Failed to read SKILL.md")?; - - let metadata = self.parse_skill_metadata_from_content(&content)?; - let name = metadata.name.unwrap_or_else(|| directory.clone()); - let description = metadata.description.unwrap_or_default(); - let key = format!("{repo_key_prefix}{directory}"); - let readme_url = path.parent().map(|parent| { - format!( - "https://github.com/{}/{}/blob/{}/{}/SKILL.md", - repo.owner, - repo.name, - branch, - parent.to_str().unwrap_or("") - ) - }); - - skills.push(Skill { - key, - name, - description, - directory, - readme_url, - installed: false, - source_kind: proxycast_core::models::SkillSourceKind::Other, - repo_owner: Some(repo.owner.clone()), - repo_name: Some(repo.name.clone()), - repo_branch: Some(branch.to_string()), - }); + let Some((directory, relative_path)) = + Self::skill_file_marker_from_archive_path(archive_path) + else { + continue; + }; + let readme_parent = archive_path + .parent() + .map(|parent| parent.to_string_lossy().to_string()) + .unwrap_or_default(); + let entry = + manifests + .entry(directory.clone()) + .or_insert_with(|| RemoteSkillArchiveEntry { + directory: directory.clone(), + content: String::new(), + readme_parent: readme_parent.clone(), + files: HashMap::new(), + resource_summary: SkillResourceSummary::default(), + }); + if entry.readme_parent.is_empty() { + entry.readme_parent = readme_parent; } + if let Some(resource_dir) = + relative_path + .components() + .next() + .and_then(|component| match component { + Component::Normal(value) => Some(value.to_string_lossy().to_string()), + _ => None, + }) + { + match resource_dir.as_str() { + "scripts" => entry.resource_summary.has_scripts = true, + "references" => entry.resource_summary.has_references = true, + "assets" => entry.resource_summary.has_assets = true, + _ => {} + } + } + + if file.is_dir() { + continue; + } + + use std::io::Read; + let mut bytes = Vec::new(); + file.read_to_end(&mut bytes) + .context("Failed to read skill archive entry")?; + + if relative_path == Path::new("SKILL.md") { + entry.content = String::from_utf8(bytes).context("Failed to read SKILL.md")?; + continue; + } + + entry.files.insert(relative_path, bytes); } - Ok(skills) + Ok(manifests) + } + + fn skill_file_marker_from_archive_path(path: &Path) -> Option<(String, PathBuf)> { + let components: Vec = path + .components() + .filter_map(|component| match component { + Component::Normal(value) => Some(value.to_string_lossy().to_string()), + _ => None, + }) + .collect(); + + if components.len() < 3 { + return None; + } + + let directory = components.get(1)?.clone(); + let mut relative_path = PathBuf::new(); + for component in components.iter().skip(2) { + relative_path.push(component); + } + if relative_path.as_os_str().is_empty() { + return None; + } + + Some((directory, relative_path)) + } + + #[cfg(test)] + fn skill_resource_marker_from_archive_path(path: &Path) -> Option<(String, String)> { + let (directory, relative_path) = Self::skill_file_marker_from_archive_path(path)?; + let resource_dir = + relative_path + .components() + .next() + .and_then(|component| match component { + Component::Normal(value) => Some(value.to_string_lossy().to_string()), + _ => None, + })?; + if !matches!(resource_dir.as_str(), "scripts" | "references" | "assets") { + return None; + } + + Some((directory, resource_dir)) } fn build_branch_candidates(branch: &str) -> Vec { @@ -438,7 +542,6 @@ impl SkillService { } } - /// 安装技能 pub async fn install_skill( &self, app_type: &AppType, @@ -455,7 +558,6 @@ impl SkillService { fs::remove_dir_all(&target_dir).context("Failed to remove existing skill")?; } - // 尝试多个分支 let branches = if repo_branch == "main" { vec!["main", "master"] } else { @@ -463,7 +565,6 @@ impl SkillService { }; let mut last_error = None; - for branch in branches { let zip_url = format!( "https://github.com/{repo_owner}/{repo_name}/archive/refs/heads/{branch}.zip" @@ -473,18 +574,71 @@ impl SkillService { .download_and_extract(&zip_url, &target_dir, directory) .await { - Ok(_) => return Ok(()), - Err(e) => { - last_error = Some(e); - continue; - } + Ok(()) => match self.validate_installed_skill_dir(&target_dir) { + Ok(()) => return Ok(()), + Err(error) => { + let _ = fs::remove_dir_all(&target_dir); + last_error = Some(error); + } + }, + Err(error) => last_error = Some(error), } } Err(last_error.unwrap_or_else(|| anyhow!("Failed to install skill"))) } - /// 下载并解压技能 + pub async fn inspect_remote_skill( + &self, + repo_owner: &str, + repo_name: &str, + repo_branch: &str, + directory: &str, + ) -> Result { + for branch in Self::build_branch_candidates(repo_branch) { + let repo = SkillRepo { + owner: repo_owner.to_string(), + name: repo_name.to_string(), + branch: branch.clone(), + enabled: true, + }; + + match self.load_remote_skill_archive_entries(&repo, &branch).await { + Ok(entries) => { + let entry = entries + .into_values() + .find(|entry| entry.directory == directory) + .ok_or_else(|| anyhow!("Skill directory not found in archive"))?; + + return Ok(Self::inspect_remote_skill_package( + &entry.content, + entry.resource_summary, + &entry.files, + )); + } + Err(error) => { + if branch == repo_branch { + continue; + } + tracing::warn!( + "[SkillService] inspect remote skill fallback {} -> {} failed: {}", + repo_branch, + branch, + error + ); + } + } + } + + Err(anyhow!( + "Failed to inspect remote skill {}/{}@{}:{}", + repo_owner, + repo_name, + repo_branch, + directory + )) + } + async fn download_and_extract( &self, zip_url: &str, @@ -506,36 +660,39 @@ impl SkillService { let cursor = std::io::Cursor::new(bytes); let mut archive = zip::ZipArchive::new(cursor).context("Failed to open ZIP")?; - // 查找技能目录 let skill_prefix = format!("/{directory}/"); let mut found = false; - for i in 0..archive.len() { - let mut file = archive.by_index(i)?; + for index in 0..archive.len() { + let mut file = archive.by_index(index)?; let file_path = file.name().to_string(); - if file_path.contains(&skill_prefix) { - found = true; - let relative_path = file_path - .split(&skill_prefix) - .nth(1) - .unwrap_or("") - .to_string(); - - if !relative_path.is_empty() { - let output_path = target_dir.join(&relative_path); - - if file.is_dir() { - fs::create_dir_all(&output_path)?; - } else { - if let Some(parent) = output_path.parent() { - fs::create_dir_all(parent)?; - } - let mut output_file = fs::File::create(&output_path)?; - std::io::copy(&mut file, &mut output_file)?; - } - } + if !file_path.contains(&skill_prefix) { + continue; } + + found = true; + let relative_path = file_path + .split(&skill_prefix) + .nth(1) + .unwrap_or("") + .to_string(); + + if relative_path.is_empty() { + continue; + } + + let output_path = target_dir.join(&relative_path); + if file.is_dir() { + fs::create_dir_all(&output_path)?; + continue; + } + + if let Some(parent) = output_path.parent() { + fs::create_dir_all(parent)?; + } + let mut output_file = fs::File::create(&output_path)?; + std::io::copy(&mut file, &mut output_file)?; } if !found { @@ -545,7 +702,6 @@ impl SkillService { Ok(()) } - /// 卸载技能 pub fn uninstall_skill(app_type: &AppType, directory: &str) -> Result<()> { let skills_dir = Self::get_skills_dir(app_type)?; let target_dir = skills_dir.join(directory); @@ -557,32 +713,313 @@ impl SkillService { Ok(()) } - /// 解析技能元数据 - fn parse_skill_metadata(&self, path: &Path) -> Result { - let content = fs::read_to_string(path).context("Failed to read SKILL.md")?; - self.parse_skill_metadata_from_content(&content) + fn parse_skill_manifest_from_content(&self, content: &str) -> Result { + parse_manifest_content(content).map_err(|error| anyhow!(error)) } - /// 从内容解析技能元数据 - fn parse_skill_metadata_from_content(&self, content: &str) -> Result { - let content = content.trim_start_matches('\u{feff}'); - let parts: Vec<&str> = content.splitn(3, "---").collect(); + fn validate_installed_skill_dir(&self, skill_dir: &Path) -> Result<()> { + let inspection = Self::inspect_skill_dir(skill_dir)?; + if !inspection.standard_compliance.validation_errors.is_empty() { + return Err(anyhow!( + "Skill package is not Agent Skills compliant: {}", + inspection.standard_compliance.validation_errors.join("; ") + )); + } + Ok(()) + } - if parts.len() < 3 { - return Ok(SkillMetadata { - name: None, - description: None, - }); + pub fn inspect_skill_dir(skill_dir: &Path) -> Result { + let skill_md = skill_dir.join("SKILL.md"); + if !skill_md.is_file() { + return Err(anyhow!("Skill package missing SKILL.md")); } - let front_matter = parts[1].trim(); - let meta: SkillMetadata = - serde_yaml::from_str(front_matter).context("Failed to parse YAML front matter")?; + let content = fs::read_to_string(&skill_md).context("Failed to read SKILL.md")?; + let resource_summary = summarize_skill_resources_dir(skill_dir); - Ok(meta) + Ok(Self::build_skill_inspection( + &content, + resource_summary, + |relative_path| Self::read_skill_relative_file(skill_dir, relative_path), + )) + } + + fn inspect_remote_skill_package( + content: &str, + resource_summary: SkillResourceSummary, + files: &HashMap>, + ) -> SkillPackageInspection { + Self::build_skill_inspection(content, resource_summary, |relative_path| { + Self::read_archive_skill_relative_file(files, relative_path) + }) + } + + fn build_skill_inspection( + content: &str, + resource_summary: SkillResourceSummary, + mut read_relative_file: F, + ) -> SkillPackageInspection + where + F: FnMut(&str) -> Result, + { + let (license, metadata, allowed_tools, mut standard_compliance) = + match parse_manifest_content(content) { + Ok(manifest) => ( + manifest.metadata.license, + manifest.metadata.metadata, + manifest.metadata.allowed_tools, + manifest.compliance, + ), + Err(error) => ( + None, + HashMap::new(), + Vec::new(), + SkillStandardCompliance { + is_standard: false, + validation_errors: vec![error], + deprecated_fields: Vec::new(), + }, + ), + }; + + Self::validate_proxycast_skill_metadata( + &metadata, + &mut standard_compliance.validation_errors, + |relative_path| read_relative_file(relative_path), + ); + standard_compliance.validation_errors.sort(); + standard_compliance.validation_errors.dedup(); + standard_compliance.is_standard = standard_compliance.validation_errors.is_empty(); + + SkillPackageInspection { + content: content.to_string(), + license, + metadata, + allowed_tools, + resource_summary, + standard_compliance, + } + } + + fn validate_proxycast_skill_metadata( + metadata: &HashMap, + validation_errors: &mut Vec, + mut read_relative_file: impl FnMut(&str) -> Result, + ) { + let Some(workflow_ref) = metadata.get("proxycast_workflow_ref") else { + return; + }; + + let workflow_ref = workflow_ref.trim(); + if workflow_ref.is_empty() { + validation_errors.push("字段 `metadata.proxycast_workflow_ref` 不能为空".to_string()); + return; + } + + match read_relative_file(workflow_ref) { + Ok(content) => { + if let Err(error) = Self::validate_workflow_content(workflow_ref, &content) { + validation_errors.push(format!( + "字段 `metadata.proxycast_workflow_ref` 校验失败: {error}" + )); + } + } + Err(error) => { + validation_errors.push(format!( + "字段 `metadata.proxycast_workflow_ref` 校验失败: {error}" + )); + } + } + } + + fn validate_workflow_content(workflow_ref: &str, content: &str) -> Result<()> { + let parsed = serde_yaml::from_str::(content) + .map_err(|error| anyhow!("`{workflow_ref}` 无法解析为 JSON/YAML: {error}"))?; + + let is_valid = parsed.is_sequence() + || parsed + .as_mapping() + .and_then(|mapping| { + mapping + .get(serde_yaml::Value::String("steps".to_string())) + .and_then(|value| value.as_sequence()) + }) + .is_some(); + + if !is_valid { + return Err(anyhow!( + "`{workflow_ref}` 必须是数组,或包含数组字段 `steps` 的对象" + )); + } + + Ok(()) + } + + fn normalize_skill_relative_path(relative_path: &str) -> Result { + let relative_path = Path::new(relative_path); + if relative_path.is_absolute() { + return Err(anyhow!("必须引用 skill 包内的相对路径")); + } + + let mut normalized = PathBuf::new(); + for component in relative_path.components() { + match component { + Component::Normal(value) => normalized.push(value), + Component::CurDir => {} + _ => return Err(anyhow!("不能引用 skill 包外路径")), + } + } + + if normalized.as_os_str().is_empty() { + return Err(anyhow!("引用文件路径不能为空")); + } + + Ok(normalized) + } + + fn read_skill_relative_file(skill_dir: &Path, relative_path: &str) -> Result { + let workflow_path = Self::resolve_skill_relative_file(skill_dir, relative_path)?; + fs::read_to_string(&workflow_path) + .with_context(|| format!("无法读取 workflow 引用文件 `{}`", workflow_path.display())) + } + + fn read_archive_skill_relative_file( + files: &HashMap>, + relative_path: &str, + ) -> Result { + let normalized = Self::normalize_skill_relative_path(relative_path)?; + let bytes = files + .get(&normalized) + .ok_or_else(|| anyhow!("引用文件不存在: {}", normalized.display()))?; + String::from_utf8(bytes.clone()).map_err(|error| { + anyhow!( + "无法读取 workflow 引用文件 `{}`: {error}", + normalized.display() + ) + }) + } + + fn resolve_skill_relative_file(skill_dir: &Path, relative_path: &str) -> Result { + let normalized = Self::normalize_skill_relative_path(relative_path)?; + let candidate = skill_dir.join(&normalized); + if !candidate.is_file() { + return Err(anyhow!("引用文件不存在: {}", normalized.display())); + } + + let canonical_skill_dir = skill_dir.canonicalize().context("无法解析 skill 包目录")?; + let canonical_candidate = candidate + .canonicalize() + .with_context(|| format!("无法解析引用文件: {}", candidate.display()))?; + if !canonical_candidate.starts_with(&canonical_skill_dir) { + return Err(anyhow!("不能引用 skill 包外路径")); + } + + Ok(canonical_candidate) + } + + fn build_skill_from_file( + &self, + skill_md: &Path, + key: String, + directory: String, + installed: bool, + source_kind: SkillSourceKind, + catalog_source: SkillCatalogSource, + readme_url: Option, + repo_owner: Option, + repo_name: Option, + repo_branch: Option, + _resource_summary: SkillResourceSummary, + ) -> Result { + let skill_dir = skill_md + .parent() + .ok_or_else(|| anyhow!("Failed to resolve skill directory"))?; + let inspection = Self::inspect_skill_dir(skill_dir)?; + Ok(self.build_skill_from_inspection( + inspection, + key, + directory, + installed, + source_kind, + catalog_source, + readme_url, + repo_owner, + repo_name, + repo_branch, + )) + } + + fn build_skill_from_remote_archive_entry( + &self, + entry: RemoteSkillArchiveEntry, + key: String, + readme_url: Option, + repo_owner: Option, + repo_name: Option, + repo_branch: Option, + ) -> Skill { + let inspection = Self::inspect_remote_skill_package( + &entry.content, + entry.resource_summary, + &entry.files, + ); + self.build_skill_from_inspection( + inspection, + key, + entry.directory, + false, + SkillSourceKind::Other, + SkillCatalogSource::Remote, + readme_url, + repo_owner, + repo_name, + repo_branch, + ) + } + + fn build_skill_from_inspection( + &self, + inspection: SkillPackageInspection, + key: String, + directory: String, + installed: bool, + source_kind: SkillSourceKind, + catalog_source: SkillCatalogSource, + readme_url: Option, + repo_owner: Option, + repo_name: Option, + repo_branch: Option, + ) -> Skill { + let parsed_manifest = self + .parse_skill_manifest_from_content(&inspection.content) + .ok(); + + Skill { + key, + name: parsed_manifest + .as_ref() + .and_then(|manifest| manifest.metadata.name.clone()) + .unwrap_or_else(|| directory.clone()), + description: parsed_manifest + .as_ref() + .and_then(|manifest| manifest.metadata.description.clone()) + .unwrap_or_default(), + directory, + readme_url, + installed, + source_kind, + catalog_source, + repo_owner, + repo_name, + repo_branch, + license: inspection.license, + metadata: inspection.metadata, + allowed_tools: inspection.allowed_tools, + resource_summary: Some(inspection.resource_summary), + standard_compliance: Some(inspection.standard_compliance), + } } - /// 清空技能仓库缓存 pub fn refresh_cache(&self) { self.repo_cache.write().clear(); } @@ -590,7 +1027,11 @@ impl SkillService { #[cfg(test)] mod tests { - use super::SkillService; + use super::{CatalogSkillRoot, SkillService}; + use proxycast_core::models::{AppType, SkillCatalogSource}; + use std::collections::HashMap; + use std::path::{Path, PathBuf}; + use tempfile::TempDir; #[test] fn build_branch_candidates_should_include_main_master_fallback() { @@ -607,4 +1048,211 @@ mod tests { vec!["release".to_string()] ); } + + #[test] + fn resource_marker_should_extract_skill_dir_and_resource_dir() { + let path = Path::new("repo-main/social_post_with_cover/references/workflow.json"); + let marker = SkillService::skill_resource_marker_from_archive_path(path); + assert_eq!( + marker, + Some(( + "social_post_with_cover".to_string(), + "references".to_string() + )) + ); + } + + #[test] + fn inspect_skill_dir_should_collect_standard_metadata_and_workflow_state() { + let temp_dir = TempDir::new().unwrap(); + let skill_dir = temp_dir.path().join("social_post_with_cover"); + let references_dir = skill_dir.join("references"); + std::fs::create_dir_all(&references_dir).unwrap(); + std::fs::write( + skill_dir.join("SKILL.md"), + r#"--- +name: Social Post +description: Generate social posts +license: MIT +metadata: + proxycast_workflow_ref: references/workflow.json + proxycast_category: social +allowed-tools: + - web.search +--- + +# Social Post +"#, + ) + .unwrap(); + std::fs::write( + references_dir.join("workflow.json"), + r#"{"steps":[{"id":"draft","title":"起草"}]}"#, + ) + .unwrap(); + + let inspection = SkillService::inspect_skill_dir(&skill_dir).unwrap(); + + assert_eq!(inspection.license.as_deref(), Some("MIT")); + assert_eq!( + inspection + .metadata + .get("proxycast_category") + .map(String::as_str), + Some("social") + ); + assert_eq!(inspection.allowed_tools, vec!["web.search".to_string()]); + assert!(inspection.resource_summary.has_references); + assert!(inspection.standard_compliance.is_standard); + assert!(inspection.standard_compliance.validation_errors.is_empty()); + } + + #[test] + fn inspect_skill_dir_should_report_invalid_workflow_reference() { + let temp_dir = TempDir::new().unwrap(); + let skill_dir = temp_dir.path().join("broken_skill"); + std::fs::create_dir_all(&skill_dir).unwrap(); + std::fs::write( + skill_dir.join("SKILL.md"), + r#"--- +name: Broken +description: Broken workflow +metadata: + proxycast_workflow_ref: ../outside.yaml +--- +"#, + ) + .unwrap(); + + let inspection = SkillService::inspect_skill_dir(&skill_dir).unwrap(); + + assert!(!inspection.standard_compliance.is_standard); + assert!(inspection + .standard_compliance + .validation_errors + .iter() + .any(|error| error.contains("不能引用 skill 包外路径"))); + } + + #[test] + fn inspect_remote_skill_package_should_report_invalid_workflow_reference() { + let mut files = HashMap::new(); + files.insert( + PathBuf::from("references").join("workflow.json"), + br#"{"title":"missing steps"}"#.to_vec(), + ); + + let inspection = SkillService::inspect_remote_skill_package( + r#"--- +name: Remote Broken +description: Broken workflow +metadata: + proxycast_workflow_ref: references/workflow.json +--- +"#, + proxycast_core::models::SkillResourceSummary { + has_references: true, + ..Default::default() + }, + &files, + ); + + assert!(!inspection.standard_compliance.is_standard); + assert!(inspection + .standard_compliance + .validation_errors + .iter() + .any(|error| error.contains("必须是数组,或包含数组字段 `steps` 的对象"))); + } + + #[test] + fn collect_local_skills_should_prefer_project_root_and_mark_catalog_source() { + let service = SkillService::new().unwrap(); + let temp_dir = TempDir::new().unwrap(); + let project_root = temp_dir + .path() + .join("project") + .join(".agents") + .join("skills"); + let user_root = temp_dir.path().join("user").join("skills"); + let project_skill_dir = project_root.join("shared-skill"); + let user_skill_dir = user_root.join("shared-skill"); + + std::fs::create_dir_all(&project_skill_dir).unwrap(); + std::fs::create_dir_all(&user_skill_dir).unwrap(); + std::fs::write( + project_skill_dir.join("SKILL.md"), + "---\nname: Project Skill\ndescription: from project\n---\n", + ) + .unwrap(); + std::fs::write( + user_skill_dir.join("SKILL.md"), + "---\nname: User Skill\ndescription: from user\n---\n", + ) + .unwrap(); + + let mut all_skills = HashMap::new(); + service + .collect_local_skills( + &AppType::ProxyCast, + &[ + CatalogSkillRoot { + source: SkillCatalogSource::Project, + path: project_root, + }, + CatalogSkillRoot { + source: SkillCatalogSource::User, + path: user_root, + }, + ], + &mut all_skills, + ) + .unwrap(); + + let skill = all_skills.get("local:shared-skill").unwrap(); + assert_eq!(skill.name, "Project Skill"); + assert_eq!(skill.catalog_source, SkillCatalogSource::Project); + } + + #[test] + fn collect_local_skills_should_surface_workflow_validation_errors() { + let service = SkillService::new().unwrap(); + let temp_dir = TempDir::new().unwrap(); + let user_root = temp_dir.path().join("user").join("skills"); + let broken_skill_dir = user_root.join("broken-workflow"); + std::fs::create_dir_all(&broken_skill_dir).unwrap(); + std::fs::write( + broken_skill_dir.join("SKILL.md"), + r#"--- +name: Broken Workflow +description: local validation +metadata: + proxycast_workflow_ref: references/workflow.json +--- +"#, + ) + .unwrap(); + + let mut all_skills = HashMap::new(); + service + .collect_local_skills( + &AppType::ProxyCast, + &[CatalogSkillRoot { + source: SkillCatalogSource::User, + path: user_root, + }], + &mut all_skills, + ) + .unwrap(); + + let skill = all_skills.get("local:broken-workflow").unwrap(); + assert!(!skill.standard_compliance.as_ref().unwrap().is_standard); + assert!(skill + .standard_compliance + .as_ref() + .unwrap() + .validation_errors + .iter() + .any(|error| error.contains("引用文件不存在"))); + } } diff --git a/src-tauri/crates/skills/Cargo.toml b/src-tauri/crates/skills/Cargo.toml index 878fc6da6..e54f40660 100644 --- a/src-tauri/crates/skills/Cargo.toml +++ b/src-tauri/crates/skills/Cargo.toml @@ -13,7 +13,11 @@ proxycast-server-utils.workspace = true serde.workspace = true serde_json.workspace = true +serde_yaml.workspace = true async-trait.workspace = true tracing.workspace = true regex.workspace = true dirs.workspace = true + +[dev-dependencies] +tempfile.workspace = true diff --git a/src-tauri/crates/skills/src/lib.rs b/src-tauri/crates/skills/src/lib.rs index 98163a677..b92144f5e 100644 --- a/src-tauri/crates/skills/src/lib.rs +++ b/src-tauri/crates/skills/src/lib.rs @@ -21,8 +21,9 @@ pub use execution_callback::{ pub use llm_provider::{LlmProvider, SkillError}; pub use proxycast_llm_provider::ProxyCastLlmProvider; pub use skill_loader::{ - find_skill_by_name, get_proxycast_skills_dir, load_skill_from_file, load_skills_from_directory, - parse_allowed_tools, parse_boolean, parse_skill_frontmatter, parse_workflow_steps, - LoadedSkillDefinition, SkillFrontmatter, SkillTriggerConfig, WorkflowStep, + find_skill_by_name, get_project_skills_dir, get_proxycast_skills_dir, get_skill_roots, + load_skill_from_file, load_skills_from_directory, parse_allowed_tools, parse_boolean, + parse_skill_frontmatter, parse_workflow_steps, LoadedSkillDefinition, SkillFrontmatter, + SkillTriggerConfig, WorkflowStep, }; pub use skill_matcher::{SkillMatch, SkillMatcher}; diff --git a/src-tauri/crates/skills/src/skill_loader.rs b/src-tauri/crates/skills/src/skill_loader.rs index 4288d0dbf..63ff1ca76 100644 --- a/src-tauri/crates/skills/src/skill_loader.rs +++ b/src-tauri/crates/skills/src/skill_loader.rs @@ -1,9 +1,16 @@ //! Skill 定义加载器 //! -//! 负责从 `~/.proxycast/skills//SKILL.md` 加载并解析 Skill 定义。 +//! 负责从标准 Agent Skills 包中加载并解析 Skill 定义。 +use std::collections::HashMap; use std::path::{Path, PathBuf}; +use proxycast_core::app_paths; +use proxycast_core::models::{ + parse_skill_manifest_from_content, split_skill_frontmatter, ParsedSkillManifest, + SkillStandardCompliance, +}; +use proxycast_services::skill_service::SkillService; use serde::{Deserialize, Serialize}; /// Skill 自动触发条件配置 @@ -44,22 +51,23 @@ fn default_step_execution_mode() -> String { pub struct SkillFrontmatter { pub name: Option, pub description: Option, - #[serde(rename = "allowed-tools")] - pub allowed_tools: Option, - #[serde(rename = "argument-hint")] + pub license: Option, + #[serde(default)] + pub metadata: HashMap, + pub allowed_tools: Option>, pub argument_hint: Option, - #[serde(rename = "when-to-use")] pub when_to_use: Option, pub version: Option, pub model: Option, pub provider: Option, - #[serde(rename = "disable-model-invocation")] pub disable_model_invocation: Option, - #[serde(rename = "execution-mode")] pub execution_mode: Option, - /// Workflow 步骤定义(JSON 格式) - #[serde(rename = "steps-json")] pub steps_json: Option, + pub workflow_ref: Option, + #[serde(default)] + pub deprecated_fields: Vec, + #[serde(default)] + pub validation_errors: Vec, } /// 内部 Skill 定义(用于加载和执行) @@ -69,6 +77,8 @@ pub struct LoadedSkillDefinition { pub display_name: String, pub description: String, pub markdown_content: String, + pub license: Option, + pub metadata: HashMap, pub allowed_tools: Option>, pub argument_hint: Option, pub when_to_use: Option, @@ -78,58 +88,102 @@ pub struct LoadedSkillDefinition { pub provider: Option, pub disable_model_invocation: bool, pub execution_mode: String, + pub workflow_ref: Option, /// Workflow 步骤定义(仅 execution_mode == "workflow" 时有效) pub workflow_steps: Vec, + pub standard_compliance: SkillStandardCompliance, +} + +#[derive(Debug, Deserialize)] +struct WorkflowDocument { + #[serde(default)] + steps: Vec, } -/// 解析 Skill 文件的 frontmatter pub fn parse_skill_frontmatter(content: &str) -> (SkillFrontmatter, String) { - let regex = regex::Regex::new(r"^---\s*\n([\s\S]*?)---\s*\n?").unwrap(); + let Some((_frontmatter, body)) = split_skill_frontmatter(content) else { + return (SkillFrontmatter::default(), content.to_string()); + }; - 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), - "steps-json" => frontmatter.steps_json = Some(clean_value), - _ => {} - } - } + match parse_skill_manifest_from_content(content) { + Ok(parsed) => { + let frontmatter = build_skill_frontmatter_from_manifest(&parsed); + (frontmatter, body.to_string()) } + Err(error) => ( + SkillFrontmatter { + validation_errors: vec![error], + ..SkillFrontmatter::default() + }, + body.to_string(), + ), + } +} - (frontmatter, body) - } else { - (SkillFrontmatter::default(), content.to_string()) +fn build_skill_frontmatter_from_manifest(parsed: &ParsedSkillManifest) -> SkillFrontmatter { + let metadata = parsed.metadata.metadata.clone(); + let version = metadata + .get("proxycast_version") + .cloned() + .or_else(|| parsed.raw_string("version")); + let argument_hint = metadata + .get("proxycast_argument_hint") + .cloned() + .or_else(|| parsed.raw_string("argument-hint")) + .or_else(|| parsed.raw_string("argument_hint")); + let when_to_use = metadata + .get("proxycast_when_to_use") + .cloned() + .or_else(|| parsed.raw_string("when-to-use")) + .or_else(|| parsed.raw_string("when_to_use")); + let model = metadata + .get("proxycast_model_preference") + .cloned() + .or_else(|| parsed.raw_string("model")); + let provider = metadata + .get("proxycast_provider_preference") + .cloned() + .or_else(|| parsed.raw_string("provider")); + let workflow_ref = metadata + .get("proxycast_workflow_ref") + .cloned() + .filter(|value| !value.trim().is_empty()); + let execution_mode = metadata + .get("proxycast_execution_mode") + .cloned() + .or_else(|| parsed.raw_string("execution-mode")) + .or_else(|| workflow_ref.as_ref().map(|_| "workflow".to_string())); + + let disable_model_invocation = metadata + .get("proxycast_disable_model_invocation") + .cloned() + .or_else(|| { + parsed + .raw_bool("disable-model-invocation") + .map(|value| value.to_string()) + }); + + SkillFrontmatter { + name: parsed.metadata.name.clone(), + description: parsed.metadata.description.clone(), + license: parsed.metadata.license.clone(), + metadata, + allowed_tools: (!parsed.metadata.allowed_tools.is_empty()) + .then(|| parsed.metadata.allowed_tools.clone()), + argument_hint, + when_to_use, + version, + model, + provider, + disable_model_invocation, + execution_mode, + steps_json: parsed.raw_string("steps-json"), + workflow_ref, + deprecated_fields: parsed.compliance.deprecated_fields.clone(), + validation_errors: parsed.compliance.validation_errors.clone(), } } -/// 解析 allowed-tools 字段 pub fn parse_allowed_tools(value: Option<&str>) -> Option> { value.and_then(|v| { if v.is_empty() { @@ -148,30 +202,28 @@ pub fn parse_allowed_tools(value: Option<&str>) -> Option> { }) } -/// 解析布尔值字段 pub fn parse_boolean(value: Option<&str>, default: bool) -> bool { value - .map(|v| { - let lower = v.to_lowercase(); - matches!(lower.as_str(), "true" | "1" | "yes") + .map(|v| match v.trim().to_ascii_lowercase().as_str() { + "true" | "1" | "yes" | "enabled" => true, + "false" | "0" | "no" | "disabled" => false, + _ => default, }) .unwrap_or(default) } -/// 解析 workflow steps -/// -/// 支持两种来源: -/// 1. frontmatter 中的 `steps-json` 字段(单行 JSON 数组) -/// 2. markdown body 中的 `` 注释块 pub fn parse_workflow_steps(steps_json: Option<&str>, markdown_content: &str) -> Vec { - // 优先使用 frontmatter 中的 steps-json if let Some(json) = steps_json { if let Ok(steps) = serde_json::from_str::>(json) { return steps; } + if let Ok(document) = serde_json::from_str::(json) { + if !document.steps.is_empty() { + return document.steps; + } + } } - // 回退:从 markdown body 中解析 let re = regex::Regex::new(r"").unwrap(); if let Some(captures) = re.captures(markdown_content) { if let Some(json_match) = captures.get(1) { @@ -185,7 +237,42 @@ pub fn parse_workflow_steps(steps_json: Option<&str>, markdown_content: &str) -> Vec::new() } -/// 从文件加载 Skill 定义 +fn parse_workflow_steps_from_reference( + base_dir: &Path, + workflow_ref: Option<&str>, +) -> Vec { + let Some(workflow_ref) = workflow_ref else { + return Vec::new(); + }; + if workflow_ref.trim().is_empty() { + return Vec::new(); + } + + let workflow_path = base_dir.join(workflow_ref); + let Ok(canonical_base) = base_dir.canonicalize() else { + return Vec::new(); + }; + let Ok(canonical_workflow) = workflow_path.canonicalize() else { + return Vec::new(); + }; + if !canonical_workflow.starts_with(&canonical_base) { + return Vec::new(); + } + + let Ok(content) = std::fs::read_to_string(&canonical_workflow) else { + return Vec::new(); + }; + + if let Ok(steps) = serde_yaml::from_str::>(&content) { + return steps; + } + if let Ok(document) = serde_yaml::from_str::(&content) { + return document.steps; + } + + Vec::new() +} + pub fn load_skill_from_file( skill_name: &str, file_path: &Path, @@ -193,41 +280,67 @@ pub fn load_skill_from_file( let content = std::fs::read_to_string(file_path).map_err(|e| format!("读取 Skill 文件失败: {}", e))?; - let (frontmatter, markdown_content) = parse_skill_frontmatter(&content); + let (mut frontmatter, markdown_content) = parse_skill_frontmatter(&content); + let base_dir = file_path + .parent() + .ok_or_else(|| "Skill 文件缺少父目录".to_string())?; + let inspection = SkillService::inspect_skill_dir(base_dir) + .map_err(|e| format!("检查 Skill 包失败: {}", e))?; + + frontmatter.license = inspection.license.clone(); + frontmatter.metadata = inspection.metadata.clone(); + if !inspection.allowed_tools.is_empty() { + frontmatter.allowed_tools = Some(inspection.allowed_tools.clone()); + } + frontmatter.deprecated_fields = inspection.standard_compliance.deprecated_fields.clone(); + frontmatter.validation_errors = inspection.standard_compliance.validation_errors.clone(); 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 allowed_tools = frontmatter.allowed_tools.clone().or_else(|| { + parse_allowed_tools( + frontmatter + .metadata + .get("allowed_tools") + .map(|value| value.as_str()), + ) + }); let disable_model_invocation = parse_boolean(frontmatter.disable_model_invocation.as_deref(), false); - let execution_mode = frontmatter + let mut execution_mode = frontmatter .execution_mode .clone() .unwrap_or_else(|| "prompt".to_string()); - let workflow_steps = parse_workflow_steps(frontmatter.steps_json.as_deref(), &markdown_content); - - // 如果有 steps 但 execution_mode 未显式设置,自动升级为 workflow - let execution_mode = if !workflow_steps.is_empty() && execution_mode == "prompt" { - "workflow".to_string() - } else { - execution_mode + let workflow_steps = { + let referenced = + parse_workflow_steps_from_reference(base_dir, frontmatter.workflow_ref.as_deref()); + if referenced.is_empty() { + parse_workflow_steps(frontmatter.steps_json.as_deref(), &markdown_content) + } else { + referenced + } }; - // 尝试将 when_to_use 解析为 JSON 格式的 SkillTriggerConfig + if !workflow_steps.is_empty() && execution_mode == "prompt" { + execution_mode = "workflow".to_string(); + } + let when_to_use_config = frontmatter .when_to_use .as_deref() - .and_then(|v| serde_json::from_str::(v).ok()); + .and_then(|value| serde_json::from_str::(value).ok()); Ok(LoadedSkillDefinition { skill_name: skill_name.to_string(), display_name, description, markdown_content, + license: inspection.license, + metadata: frontmatter.metadata, allowed_tools, argument_hint: frontmatter.argument_hint, when_to_use: frontmatter.when_to_use, @@ -236,16 +349,33 @@ pub fn load_skill_from_file( provider: frontmatter.provider, disable_model_invocation, execution_mode, + workflow_ref: frontmatter.workflow_ref, workflow_steps, + standard_compliance: inspection.standard_compliance, }) } -/// 获取 ProxyCast Skills 目录 pub fn get_proxycast_skills_dir() -> Option { - dirs::home_dir().map(|home| home.join(".proxycast").join("skills")) + app_paths::resolve_skills_dir().ok() +} + +pub fn get_project_skills_dir() -> Option { + app_paths::resolve_project_skills_dir() +} + +pub fn get_skill_roots() -> Vec { + app_paths::resolve_proxycast_skill_roots().unwrap_or_else(|_| { + let mut roots = Vec::new(); + if let Some(project_dir) = get_project_skills_dir() { + roots.push(project_dir); + } + if let Some(user_dir) = get_proxycast_skills_dir() { + roots.push(user_dir); + } + roots + }) } -/// 从目录加载所有 Skills pub fn load_skills_from_directory(dir_path: &Path) -> Vec { let mut results = Vec::new(); @@ -261,15 +391,25 @@ pub fn load_skills_from_directory(dir_path: &Path) -> Vec } 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 !skill_file.exists() { + continue; + } - if let Ok(skill) = load_skill_from_file(&skill_name, &skill_file) { + let skill_name = path + .file_name() + .and_then(|name| name.to_str()) + .unwrap_or("unknown") + .to_string(); + + if let Ok(skill) = load_skill_from_file(&skill_name, &skill_file) { + if skill.standard_compliance.validation_errors.is_empty() { results.push(skill); + } else { + tracing::warn!( + "[load_skills_from_directory] 跳过无效 Skill: name={}, errors={}", + skill.skill_name, + skill.standard_compliance.validation_errors.join("; ") + ); } } } @@ -278,17 +418,91 @@ pub fn load_skills_from_directory(dir_path: &Path) -> Vec results } -/// 根据名称查找 Skill pub 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)); + for skills_dir in get_skill_roots() { + let skill_file = skills_dir.join(skill_name).join("SKILL.md"); + if !skill_file.exists() { + continue; + } + return load_skill_from_file(skill_name, &skill_file); } - load_skill_from_file(skill_name, &skill_file) + Err(format!("Skill 不存在: {}", skill_name)) +} + +#[cfg(test)] +mod tests { + use super::{load_skill_from_file, load_skills_from_directory}; + use tempfile::TempDir; + + #[test] + fn load_skill_from_file_should_surface_invalid_workflow_reference() { + let temp_dir = TempDir::new().unwrap(); + let skill_dir = temp_dir.path().join("workflow-skill"); + std::fs::create_dir(&skill_dir).unwrap(); + + let skill_file = skill_dir.join("SKILL.md"); + std::fs::write( + &skill_file, + r#"--- +name: workflow-skill +description: Workflow skill +metadata: + proxycast_workflow_ref: references/missing.json +--- + +# Workflow Skill +"#, + ) + .unwrap(); + + let skill = load_skill_from_file("workflow-skill", &skill_file).unwrap(); + + assert!(!skill.standard_compliance.is_standard); + assert!(skill + .standard_compliance + .validation_errors + .iter() + .any(|error| error.contains("metadata.proxycast_workflow_ref"))); + assert!(skill.workflow_steps.is_empty()); + } + + #[test] + fn load_skills_from_directory_should_skip_invalid_skill_packages() { + let temp_dir = TempDir::new().unwrap(); + let skills_dir = temp_dir.path(); + + let valid_dir = skills_dir.join("skill-valid"); + std::fs::create_dir(&valid_dir).unwrap(); + std::fs::write( + valid_dir.join("SKILL.md"), + r#"--- +name: skill-valid +description: Valid skill +--- +Valid content +"#, + ) + .unwrap(); + + let invalid_dir = skills_dir.join("skill-invalid"); + std::fs::create_dir(&invalid_dir).unwrap(); + std::fs::write( + invalid_dir.join("SKILL.md"), + r#"--- +name: skill-invalid +description: Invalid skill +metadata: + proxycast_workflow_ref: references/missing.json +--- +Invalid content +"#, + ) + .unwrap(); + + let skills = load_skills_from_directory(skills_dir); + + assert_eq!(skills.len(), 1); + assert_eq!(skills[0].skill_name, "skill-valid"); + } } diff --git a/src-tauri/crates/skills/src/skill_matcher.rs b/src-tauri/crates/skills/src/skill_matcher.rs index 2143c8689..051a7a5e0 100644 --- a/src-tauri/crates/skills/src/skill_matcher.rs +++ b/src-tauri/crates/skills/src/skill_matcher.rs @@ -247,6 +247,8 @@ mod tests { display_name: name.to_string(), description: String::new(), markdown_content: String::new(), + license: None, + metadata: std::collections::HashMap::new(), allowed_tools: None, argument_hint: None, when_to_use: None, @@ -258,7 +260,13 @@ mod tests { provider: None, disable_model_invocation: false, execution_mode: "prompt".to_string(), + workflow_ref: None, workflow_steps: Vec::new(), + standard_compliance: proxycast_core::models::SkillStandardCompliance { + is_standard: true, + validation_errors: Vec::new(), + deprecated_fields: Vec::new(), + }, } } diff --git a/src-tauri/resources/default-skills/broadcast_generate/SKILL.md b/src-tauri/resources/default-skills/broadcast_generate/SKILL.md index 3a3a99e86..56e3ff339 100644 --- a/src-tauri/resources/default-skills/broadcast_generate/SKILL.md +++ b/src-tauri/resources/default-skills/broadcast_generate/SKILL.md @@ -2,10 +2,13 @@ name: broadcast_generate description: 将文章整理为可转播客音频的源文本(下游负责真实音频合成)。 allowed-tools: proxycast_create_broadcast_generation_task -argument-hint: 输入原文、目标听众、语气、预计时长、重点段落。 -when-to-use: 用户希望把现有文稿转成播客内容,但不要求你直接写主持稿。 -version: 1.0.1 -execution-mode: prompt +metadata: + proxycast_argument_hint: 输入原文、目标听众、语气、预计时长、重点段落。 + proxycast_when_to_use: 用户希望把现有文稿转成播客内容,但不要求你直接写主持稿。 + proxycast_version: 1.1.0 + proxycast_execution_mode: prompt + proxycast_surface: creator + proxycast_category: media --- 你是 ProxyCast 的播客内容整理助手。 diff --git a/src-tauri/resources/default-skills/cover_generate/SKILL.md b/src-tauri/resources/default-skills/cover_generate/SKILL.md index f6a7c559d..6c67224d1 100644 --- a/src-tauri/resources/default-skills/cover_generate/SKILL.md +++ b/src-tauri/resources/default-skills/cover_generate/SKILL.md @@ -2,10 +2,13 @@ name: cover_generate description: 为文章或视频生成平台封面图,并写回主稿(封面场景优先使用本技能)。 allowed-tools: social_generate_cover_image, proxycast_create_cover_generation_task -argument-hint: 输入平台、标题、受众、视觉风格、尺寸要求。 -when-to-use: 用户明确要求“封面图”时使用,不要被普通配图任务替代。 -version: 1.0.1 -execution-mode: prompt +metadata: + proxycast_argument_hint: 输入平台、标题、受众、视觉风格、尺寸要求。 + proxycast_when_to_use: 用户明确要求“封面图”时使用,不要被普通配图任务替代。 + proxycast_version: 1.1.0 + proxycast_execution_mode: prompt + proxycast_surface: creator + proxycast_category: media --- 你是 ProxyCast 的封面生成助手。 diff --git a/src-tauri/resources/default-skills/image_generate/SKILL.md b/src-tauri/resources/default-skills/image_generate/SKILL.md index d5427691e..5dd578727 100644 --- a/src-tauri/resources/default-skills/image_generate/SKILL.md +++ b/src-tauri/resources/default-skills/image_generate/SKILL.md @@ -2,10 +2,13 @@ name: image_generate description: 根据文本描述生成配图素材(非封面场景)。 allowed-tools: proxycast_create_image_generation_task -argument-hint: 输入主题、画面主体、风格、构图、数量、尺寸。 -when-to-use: 用户需要普通配图、插图或概念图时使用;封面需求优先交给 cover_generate。 -version: 1.0.1 -execution-mode: prompt +metadata: + proxycast_argument_hint: 输入主题、画面主体、风格、构图、数量、尺寸。 + proxycast_when_to_use: 用户需要普通配图、插图或概念图时使用;封面需求优先交给 cover_generate。 + proxycast_version: 1.1.0 + proxycast_execution_mode: prompt + proxycast_surface: creator + proxycast_category: media --- 你是 ProxyCast 的通用配图助手。 diff --git a/src-tauri/resources/default-skills/library/SKILL.md b/src-tauri/resources/default-skills/library/SKILL.md index c6da6b091..2427b8512 100644 --- a/src-tauri/resources/default-skills/library/SKILL.md +++ b/src-tauri/resources/default-skills/library/SKILL.md @@ -2,10 +2,13 @@ name: library description: 【外部资产库】读取项目参考资料(/project)或风格参考(/styles)。 allowed-tools: list_directory, read_file -argument-hint: 输入要读取的目录、文件路径、目标主题与提取重点。 -when-to-use: 需要读取项目内参考资料,或提炼风格样例时使用。 -version: 1.0.1 -execution-mode: prompt +metadata: + proxycast_argument_hint: 输入要读取的目录、文件路径、目标主题与提取重点。 + proxycast_when_to_use: 需要读取项目内参考资料,或提炼风格样例时使用。 + proxycast_version: 1.1.0 + proxycast_execution_mode: prompt + proxycast_surface: chat + proxycast_category: research --- 你是 ProxyCast 的资料库读取助手。 diff --git a/src-tauri/resources/default-skills/modal_resource_search/SKILL.md b/src-tauri/resources/default-skills/modal_resource_search/SKILL.md index 20940e593..3ca11bef9 100644 --- a/src-tauri/resources/default-skills/modal_resource_search/SKILL.md +++ b/src-tauri/resources/default-skills/modal_resource_search/SKILL.md @@ -2,10 +2,13 @@ name: modal_resource_search description: 提交资源检索任务(图片、背景音乐、音效等),供前端资源面板消费。 allowed-tools: proxycast_create_modal_resource_search_task -argument-hint: 输入资源类型、关键词、风格、用途、数量与限制条件。 -when-to-use: 用户需要为当前内容补充外部素材资源时使用。 -version: 1.0.1 -execution-mode: prompt +metadata: + proxycast_argument_hint: 输入资源类型、关键词、风格、用途、数量与限制条件。 + proxycast_when_to_use: 用户需要为当前内容补充外部素材资源时使用。 + proxycast_version: 1.1.0 + proxycast_execution_mode: prompt + proxycast_surface: creator + proxycast_category: media --- 你是 ProxyCast 的资源检索编排助手。 diff --git a/src-tauri/resources/default-skills/research/SKILL.md b/src-tauri/resources/default-skills/research/SKILL.md index 7b1798d5f..527579653 100644 --- a/src-tauri/resources/default-skills/research/SKILL.md +++ b/src-tauri/resources/default-skills/research/SKILL.md @@ -2,10 +2,13 @@ name: research description: 联网信息检索与趋势调研(优先产出可引用结论,而非原始片段堆砌)。 allowed-tools: search_query -argument-hint: 输入调研主题、目标平台、时间范围、输出深度与关注维度。 -when-to-use: 用户需要事实核验、最新信息补充、行业/平台趋势调研时使用。 -version: 1.0.1 -execution-mode: prompt +metadata: + proxycast_argument_hint: 输入调研主题、目标平台、时间范围、输出深度与关注维度。 + proxycast_when_to_use: 用户需要事实核验、最新信息补充、行业/平台趋势调研时使用。 + proxycast_version: 1.1.0 + proxycast_execution_mode: prompt + proxycast_surface: chat + proxycast_category: research --- 你是 ProxyCast 的调研助手。 diff --git a/src-tauri/resources/default-skills/social_post_with_cover/SKILL.md b/src-tauri/resources/default-skills/social_post_with_cover/SKILL.md index 81a70407b..6a195878e 100644 --- a/src-tauri/resources/default-skills/social_post_with_cover/SKILL.md +++ b/src-tauri/resources/default-skills/social_post_with_cover/SKILL.md @@ -2,11 +2,14 @@ name: social_post_with_cover description: 生成可直接发布的社媒成稿(默认公众号风格)并自动生成 1 张头图,最终以 write_file 落盘。 allowed-tools: social_generate_cover_image, search_query -argument-hint: 输入主题、平台(如公众号/小红书)、目标受众、语气、字数、转化目标和已知素材。 -when-to-use: 用户需要“社媒文章 + 封面图”一体化输出,且希望直接复制发布。 -version: 1.3.1 -execution-mode: prompt -steps-json: '[{"id":"research","name":"阅读项目素材并检索资料","prompt":"当前任务是信息收集阶段。\n1. 仔细阅读用户提供的全部上下文素材([生效上下文]、[历史内容]、链接等)。\n2. 如果上下文不足,使用 search_query 进行 2-4 次检索,覆盖核心主题、目标受众关注点、最新案例。\n3. 将收集到的信息整理为结构化素材摘要,格式:【主题定位】【关键信息点】【目标受众洞察】【可用素材来源】。\n4. 不要撰写正文,只输出素材摘要,供后续步骤使用。","execution_mode":"prompt"},{"id":"write","name":"撰写社媒主稿","prompt":"当前任务是文稿撰写阶段。\n基于前序步骤提供的素材摘要,按照 skill 中的文案生成规则撰写完整社媒文章:\n1. 输出完整文章:标题、导语、正文(分节)、结尾 CTA。\n2. 严格匹配用户指定平台语气,未指定时默认公众号长文风格。\n3. 不要调用 social_generate_cover_image,不要输出 write_file,只输出文章正文 Markdown。\n4. 在文章末尾另起一行,输出封面图提示词建议:【封面图提示词建议】xxx(用于下一步生成)。","execution_mode":"prompt"},{"id":"cover","name":"生成封面图并输出主稿","prompt":"当前任务是封面图生成与最终落盘阶段。\n基于前序步骤提供的完整文稿:\n1. 提取文章主题与核心视觉元素,调用 social_generate_cover_image 生成 1 张封面图(尺寸 1024x1024)。\n2. 将文章与封面图整合,严格按照 skill 规定的 write_file 格式输出最终主稿文件(配图说明不得放入 write_file 内)。\n3. 如封面图生成失败,使用 【img:multimodel:{你准备的封面图提示词}】 作为占位 URL(例如 ![封面图](【img:multimodel:科技感实验室,蓝色调】)),继续完成 write_file 输出。\n4. 在 write_file 之后另起一行输出配图说明(提示词/尺寸/状态/备注),不放在文件内容里。","execution_mode":"prompt"}]' +metadata: + proxycast_argument_hint: 输入主题、平台(如公众号/小红书)、目标受众、语气、字数、转化目标和已知素材。 + proxycast_when_to_use: 用户需要“社媒文章 + 封面图”一体化输出,且希望直接复制发布。 + proxycast_version: 1.4.0 + proxycast_execution_mode: workflow + proxycast_workflow_ref: references/workflow.json + proxycast_surface: creator + proxycast_category: social --- diff --git a/src-tauri/resources/default-skills/social_post_with_cover/references/workflow.json b/src-tauri/resources/default-skills/social_post_with_cover/references/workflow.json new file mode 100644 index 000000000..928762068 --- /dev/null +++ b/src-tauri/resources/default-skills/social_post_with_cover/references/workflow.json @@ -0,0 +1,20 @@ +[ + { + "id": "research", + "name": "阅读项目素材并检索资料", + "prompt": "当前任务是信息收集阶段。\n1. 仔细阅读用户提供的全部上下文素材([生效上下文]、[历史内容]、链接等)。\n2. 如果上下文不足,使用 search_query 进行 2-4 次检索,覆盖核心主题、目标受众关注点、最新案例。\n3. 将收集到的信息整理为结构化素材摘要,格式:【主题定位】【关键信息点】【目标受众洞察】【可用素材来源】。\n4. 不要撰写正文,只输出素材摘要,供后续步骤使用。", + "execution_mode": "prompt" + }, + { + "id": "write", + "name": "撰写社媒主稿", + "prompt": "当前任务是文稿撰写阶段。\n基于前序步骤提供的素材摘要,按照 skill 中的文案生成规则撰写完整社媒文章:\n1. 输出完整文章:标题、导语、正文(分节)、结尾 CTA。\n2. 严格匹配用户指定平台语气,未指定时默认公众号长文风格。\n3. 不要调用 social_generate_cover_image,不要输出 write_file,只输出文章正文 Markdown。\n4. 在文章末尾另起一行,输出封面图提示词建议:【封面图提示词建议】xxx(用于下一步生成)。", + "execution_mode": "prompt" + }, + { + "id": "cover", + "name": "生成封面图并输出主稿", + "prompt": "当前任务是封面图生成与最终落盘阶段。\n基于前序步骤提供的完整文稿:\n1. 提取文章主题与核心视觉元素,调用 social_generate_cover_image 生成 1 张封面图(尺寸 1024x1024)。\n2. 将文章与封面图整合,严格按照 skill 规定的 write_file 格式输出最终主稿文件(配图说明不得放入 write_file 内)。\n3. 如封面图生成失败,使用 【img:multimodel:{你准备的封面图提示词}】 作为占位 URL(例如 ![封面图](【img:multimodel:科技感实验室,蓝色调】)),继续完成 write_file 输出。\n4. 在 write_file 之后另起一行输出配图说明(提示词/尺寸/状态/备注),不放在文件内容里。", + "execution_mode": "prompt" + } +] diff --git a/src-tauri/resources/default-skills/typesetting/SKILL.md b/src-tauri/resources/default-skills/typesetting/SKILL.md index 88432a8d0..c35c72c19 100644 --- a/src-tauri/resources/default-skills/typesetting/SKILL.md +++ b/src-tauri/resources/default-skills/typesetting/SKILL.md @@ -2,10 +2,13 @@ name: typesetting description: 优化文稿排版与可读性,不改变原始事实与核心表达。 allowed-tools: proxycast_create_typesetting_task -argument-hint: 输入目标平台、语气要求、段落长度偏好、标题层级规范。 -when-to-use: 用户希望提升文本可读性、结构清晰度、发布观感时使用。 -version: 1.0.1 -execution-mode: prompt +metadata: + proxycast_argument_hint: 输入目标平台、语气要求、段落长度偏好、标题层级规范。 + proxycast_when_to_use: 用户希望提升文本可读性、结构清晰度、发布观感时使用。 + proxycast_version: 1.1.0 + proxycast_execution_mode: prompt + proxycast_surface: creator + proxycast_category: writing --- 你是 ProxyCast 的排版优化助手。 diff --git a/src-tauri/resources/default-skills/url_parse/SKILL.md b/src-tauri/resources/default-skills/url_parse/SKILL.md index e73f73f32..dcf8d8664 100644 --- a/src-tauri/resources/default-skills/url_parse/SKILL.md +++ b/src-tauri/resources/default-skills/url_parse/SKILL.md @@ -2,10 +2,13 @@ name: url_parse description: 解析外部 URL 内容,并沉淀为可阅读的文本结果。 allowed-tools: proxycast_create_url_parse_task -argument-hint: 输入 URL、抽取目标(摘要/要点/全文清洗)、输出格式要求。 -when-to-use: 用户提供链接并希望抽取正文、要点或可引用信息时使用。 -version: 1.0.1 -execution-mode: prompt +metadata: + proxycast_argument_hint: 输入 URL、抽取目标(摘要/要点/全文清洗)、输出格式要求。 + proxycast_when_to_use: 用户提供链接并希望抽取正文、要点或可引用信息时使用。 + proxycast_version: 1.1.0 + proxycast_execution_mode: prompt + proxycast_surface: chat + proxycast_category: research --- 你是 ProxyCast 的链接解析助手。 diff --git a/src-tauri/resources/default-skills/video_generate/SKILL.md b/src-tauri/resources/default-skills/video_generate/SKILL.md index e351b2bb8..56cdf1d57 100644 --- a/src-tauri/resources/default-skills/video_generate/SKILL.md +++ b/src-tauri/resources/default-skills/video_generate/SKILL.md @@ -2,10 +2,13 @@ name: video_generate description: 提交视频生成任务,并触发前端视频生成流程。 allowed-tools: proxycast_create_video_generation_task -argument-hint: 输入主题、受众、平台、时长、画幅、风格、素材来源。 -when-to-use: 用户要求生成视频,或将现有文稿改编为短视频。 -version: 1.0.1 -execution-mode: prompt +metadata: + proxycast_argument_hint: 输入主题、受众、平台、时长、画幅、风格、素材来源。 + proxycast_when_to_use: 用户要求生成视频,或将现有文稿改编为短视频。 + proxycast_version: 1.1.0 + proxycast_execution_mode: prompt + proxycast_surface: creator + proxycast_category: media --- 你是 ProxyCast 的视频任务编排助手。 diff --git a/src-tauri/src/agent/aster_agent.rs b/src-tauri/src/agent/aster_agent.rs index a90d80b5a..8e39ba7f2 100644 --- a/src-tauri/src/agent/aster_agent.rs +++ b/src-tauri/src/agent/aster_agent.rs @@ -7,7 +7,7 @@ use crate::agent::aster_state::{AsterAgentState, SessionConfigBuilder}; use crate::database::DbConnection; use aster::conversation::message::Message; use futures::StreamExt; -use proxycast_agent::{convert_agent_event, TauriAgentEvent}; +use proxycast_agent::{convert_agent_event, TauriAgentEvent, WriteArtifactEventEmitter}; use tauri::{AppHandle, Emitter}; pub use proxycast_agent::session_store::{SessionDetail, SessionInfo}; @@ -56,6 +56,7 @@ impl AsterAgentWrapper { let stream_result = agent .reply(user_message, session_config, Some(cancel_token.clone())) .await; + let mut write_artifact_emitter = WriteArtifactEventEmitter::new(session_id.clone()); match stream_result { Ok(mut stream) => { @@ -63,7 +64,17 @@ impl AsterAgentWrapper { match event_result { Ok(agent_event) => { let tauri_events = convert_agent_event(agent_event); - for tauri_event in tauri_events { + for mut tauri_event in tauri_events { + let extra_events = + write_artifact_emitter.process_event(&mut tauri_event); + for extra_event in &extra_events { + if let Err(error) = app.emit(&event_name, extra_event) { + tracing::error!( + "[AsterAgentWrapper] 发送补充事件失败: {}", + error + ); + } + } if let Err(error) = app.emit(&event_name, &tauri_event) { tracing::error!("[AsterAgentWrapper] 发送事件失败: {}", error); } diff --git a/src-tauri/src/agent/mod.rs b/src-tauri/src/agent/mod.rs index dcdcb1aa9..0a49e9426 100644 --- a/src-tauri/src/agent/mod.rs +++ b/src-tauri/src/agent/mod.rs @@ -24,7 +24,10 @@ pub use credential_bridge::{ create_aster_provider, AsterProviderConfig, CredentialBridge, CredentialBridgeError, }; pub use heartbeat_service_adapter::HeartbeatServiceAdapter; -pub use proxycast_agent::{convert_agent_event, convert_to_tauri_message, TauriAgentEvent}; +pub use proxycast_agent::{ + convert_agent_event, convert_to_tauri_message, QueueInsertResult, QueuedTurnSnapshot, + QueuedTurnTask, SessionTurnQueueManager, TauriAgentEvent, +}; pub use subagent_scheduler::{ ProxyCastScheduler, ProxyCastSubAgentExecutor, SubAgentProgressEvent, SubAgentRole, }; diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index 2977d58f8..bf6ffa25c 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -295,6 +295,85 @@ pub fn run() { tracing::info!("[启动] PluginManager 任务事件发射器已设置"); } + let startup_runtime_resume = { + let aster_agent_state = app.try_state::(); + let db_state = app.try_state::(); + let api_key_provider_service = + app.try_state::(); + let log_state = app.try_state::(); + let config_manager = app.try_state::(); + let mcp_manager = app.try_state::(); + let heartbeat_state = + app.try_state::(); + + match ( + aster_agent_state, + db_state, + api_key_provider_service, + log_state, + config_manager, + mcp_manager, + heartbeat_state, + ) { + ( + Some(aster_agent_state), + Some(db_state), + Some(api_key_provider_service), + Some(log_state), + Some(config_manager), + Some(mcp_manager), + Some(heartbeat_state), + ) => Some(( + app.handle().clone(), + aster_agent_state.inner().clone(), + db_state.inner().clone(), + crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState( + api_key_provider_service.0.clone(), + ), + log_state.inner().clone(), + crate::config::GlobalConfigManagerState( + config_manager.0.clone(), + ), + mcp_manager.inner().clone(), + heartbeat_state.inner().clone(), + )), + _ => None, + } + }; + + if let Some(( + app_handle, + state, + db, + api_key_provider_service, + logs, + config_manager, + mcp_manager, + heartbeat_state, + )) = startup_runtime_resume + { + match crate::commands::aster_agent_cmd::resume_persisted_runtime_queues_on_startup( + app_handle, + &state, + &db, + &api_key_provider_service, + &logs, + &config_manager, + &mcp_manager, + &heartbeat_state, + ) { + Ok(resumed) if resumed > 0 => { + tracing::info!("[启动] 已恢复 {} 个会话的排队执行", resumed); + } + Ok(_) => { + tracing::debug!("[启动] 无需恢复持久化排队执行"); + } + Err(error) => { + tracing::warn!("[启动] 恢复持久化排队执行失败: {}", error); + } + } + } + #[cfg(debug_assertions)] { let app_handle = app.handle().clone(); @@ -1068,7 +1147,9 @@ pub fn run() { commands::skill_cmd::remove_skill_repo, commands::skill_cmd::refresh_skill_cache, commands::skill_cmd::get_installed_proxycast_skills, - commands::skill_cmd::get_local_skill_content, + commands::skill_cmd::inspect_local_skill_for_app, + commands::skill_cmd::create_skill_scaffold_for_app, + commands::skill_cmd::inspect_remote_skill, // Skill Execution commands commands::skill_exec_cmd::execute_skill, commands::skill_exec_cmd::list_executable_skills, @@ -1288,6 +1369,7 @@ pub fn run() { commands::aster_agent_cmd::aster_agent_stop, commands::aster_agent_cmd::agent_runtime_submit_turn, commands::aster_agent_cmd::agent_runtime_interrupt_turn, + commands::aster_agent_cmd::agent_runtime_remove_queued_turn, commands::aster_agent_cmd::aster_session_create, commands::aster_agent_cmd::aster_session_set_execution_strategy, commands::aster_agent_cmd::aster_session_list, diff --git a/src-tauri/src/commands/api_key_provider_cmd.rs b/src-tauri/src/commands/api_key_provider_cmd.rs index a7526b319..88fedd00c 100644 --- a/src-tauri/src/commands/api_key_provider_cmd.rs +++ b/src-tauri/src/commands/api_key_provider_cmd.rs @@ -541,12 +541,34 @@ pub async fn test_api_key_provider_connection( pub async fn test_api_key_provider_chat( db: State<'_, DbConnection>, service: State<'_, ApiKeyProviderServiceState>, + model_registry_state: State<'_, crate::commands::model_registry_cmd::ModelRegistryState>, provider_id: String, model_name: Option, prompt: String, ) -> Result { + let provider = service + .0 + .get_provider(&db, &provider_id)? + .ok_or_else(|| format!("Provider 不存在: {provider_id}"))?; + + let fallback_models = { + let guard = model_registry_state.read().await; + if let Some(model_registry) = guard.as_ref() { + model_registry + .get_local_fallback_model_ids_with_hints( + &provider_id, + &provider.provider.api_host, + Some(provider.provider.provider_type), + &provider.provider.custom_models, + ) + .await + } else { + Vec::new() + } + }; + service .0 - .test_chat(&db, &provider_id, model_name, prompt) + .test_chat_with_fallback_models(&db, &provider_id, model_name, prompt, fallback_models) .await } diff --git a/src-tauri/src/commands/aster_agent_cmd.rs b/src-tauri/src/commands/aster_agent_cmd.rs index bc907a924..fe29b81b0 100644 --- a/src-tauri/src/commands/aster_agent_cmd.rs +++ b/src-tauri/src/commands/aster_agent_cmd.rs @@ -6,8 +6,9 @@ use crate::agent::aster_state::{ProviderConfig, SessionConfigBuilder}; use crate::agent::{ - AsterAgentState, AsterAgentWrapper, HeartbeatServiceAdapter, ProxyCastScheduler, SessionDetail, - SessionInfo, SubAgentRole, TauriAgentEvent, + AsterAgentState, AsterAgentWrapper, HeartbeatServiceAdapter, ProxyCastScheduler, + QueueInsertResult, QueuedTurnSnapshot, QueuedTurnTask, SessionDetail, SessionInfo, + SubAgentRole, TauriAgentEvent, }; use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState; use crate::commands::webview_cmd::{ @@ -15,6 +16,9 @@ use crate::commands::webview_cmd::{ }; use crate::config::{GlobalConfigManager, GlobalConfigManagerState}; use crate::database::dao::agent::AgentDao; +use crate::database::dao::agent_runtime_queue::{ + AgentRuntimeQueuedTurnDao, NewAgentRuntimeQueuedTurnRecord, +}; use crate::database::DbConnection; use crate::mcp::{McpManagerState, McpServerConfig}; use crate::services::agent_timeline_service::{ @@ -53,12 +57,13 @@ use futures::StreamExt; #[cfg(test)] use proxycast_agent::request_tool_policy::REQUEST_TOOL_POLICY_MARKER; use proxycast_agent::request_tool_policy::{ - merge_system_prompt_with_request_tool_policy, resolve_request_tool_policy, - stream_reply_with_policy, ReplyAttemptError, RequestToolPolicy, + merge_system_prompt_with_request_tool_policy, resolve_request_tool_policy_with_mode, + stream_reply_with_policy, ReplyAttemptError, RequestToolPolicy, RequestToolPolicyMode, }; use proxycast_agent::{ - durable_memory_permission_pattern, is_virtual_memory_path, resolve_virtual_memory_path, - virtual_memory_relative_path, DURABLE_MEMORY_VIRTUAL_ROOT, + durable_memory_permission_pattern, is_virtual_memory_path, message_suggests_news_expansion, + resolve_virtual_memory_path, virtual_memory_relative_path, TauriRuntimeStatus, + DURABLE_MEMORY_VIRTUAL_ROOT, }; use proxycast_services::api_key_provider_service::ApiKeyProviderService; use proxycast_services::mcp_service::McpService; @@ -72,6 +77,7 @@ use std::sync::{Arc, Mutex, OnceLock}; use std::time::Duration; use tauri::{AppHandle, Emitter, State}; use tokio_util::sync::CancellationToken; +use uuid::Uuid; const DEFAULT_BASH_TIMEOUT_SECS: u64 = 300; const MAX_BASH_TIMEOUT_SECS: u64 = 1800; @@ -193,7 +199,7 @@ pub struct AsterAgentStatus { } /// Provider 配置请求 -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct ConfigureProviderRequest { #[serde(default)] pub provider_id: Option, @@ -346,7 +352,7 @@ pub async fn aster_agent_reset( } /// 发送消息请求参数 -#[derive(Debug, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct AsterChatRequest { pub message: String, #[serde(alias = "sessionId")] @@ -368,6 +374,9 @@ pub struct AsterChatRequest { /// 是否强制开启联网搜索工具策略 #[serde(default, alias = "webSearch")] pub web_search: Option, + /// 联网搜索模式(disabled / allowed / required) + #[serde(default, alias = "searchMode")] + pub search_mode: Option, /// 执行策略(react / code_orchestrated / auto) #[serde(default, alias = "executionStrategy")] pub execution_strategy: Option, @@ -380,9 +389,15 @@ pub struct AsterChatRequest { /// 请求级元数据(可选,用于 harness / 主题工作台状态对齐) #[serde(default)] pub metadata: Option, + /// 会话忙时是否进入后端队列 + #[serde(default, alias = "queueIfBusy")] + pub queue_if_busy: Option, + /// 队列项 ID(由前端或后端生成) + #[serde(default, alias = "queuedTurnId")] + pub queued_turn_id: Option, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct AgentTurnConfigSnapshot { #[serde(default, alias = "providerConfig")] pub provider_config: Option, @@ -390,6 +405,8 @@ pub struct AgentTurnConfigSnapshot { pub execution_strategy: Option, #[serde(default, alias = "webSearch")] pub web_search: Option, + #[serde(default, alias = "searchMode")] + pub search_mode: Option, #[serde(default, alias = "autoContinue")] pub auto_continue: Option, #[serde(default, alias = "systemPrompt")] @@ -398,7 +415,7 @@ pub struct AgentTurnConfigSnapshot { pub metadata: Option, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct AgentRuntimeSubmitTurnRequest { pub message: String, #[serde(alias = "sessionId")] @@ -414,6 +431,10 @@ pub struct AgentRuntimeSubmitTurnRequest { #[serde(default, alias = "turnId")] #[allow(dead_code)] pub turn_id: Option, + #[serde(default, alias = "queueIfBusy")] + pub queue_if_busy: Option, + #[serde(default, alias = "queuedTurnId")] + pub queued_turn_id: Option, } impl From for AsterChatRequest { @@ -430,6 +451,7 @@ impl From for AsterChatRequest { project_id: None, workspace_id: request.workspace_id, web_search: turn_config.as_ref().and_then(|config| config.web_search), + search_mode: turn_config.as_ref().and_then(|config| config.search_mode), execution_strategy: turn_config .as_ref() .and_then(|config| config.execution_strategy), @@ -440,6 +462,8 @@ impl From for AsterChatRequest { .as_ref() .and_then(|config| config.system_prompt.clone()), metadata: turn_config.and_then(|config| config.metadata), + queue_if_busy: request.queue_if_busy, + queued_turn_id: request.queued_turn_id, } } } @@ -453,6 +477,46 @@ pub struct AgentRuntimeInterruptTurnRequest { pub turn_id: Option, } +#[derive(Debug, Deserialize)] +pub struct AgentRuntimeRemoveQueuedTurnRequest { + #[serde(alias = "sessionId")] + pub session_id: String, + #[serde(alias = "queuedTurnId")] + pub queued_turn_id: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AgentRuntimeSessionDetail { + pub id: String, + pub name: String, + pub created_at: i64, + pub updated_at: i64, + pub thread_id: String, + pub messages: Vec, + pub execution_strategy: Option, + pub turns: Vec, + pub items: Vec, + #[serde(default)] + pub queued_turns: Vec, +} + +impl AgentRuntimeSessionDetail { + fn from_session_detail(detail: SessionDetail, queued_turns: Vec) -> Self { + Self { + id: detail.id, + name: detail.name, + created_at: detail.created_at, + updated_at: detail.updated_at, + thread_id: detail.thread_id, + messages: detail.messages, + execution_strategy: detail.execution_strategy, + turns: detail.turns, + items: detail.items, + queued_turns, + } + } +} + #[derive(Debug, Clone, Copy, Deserialize, PartialEq, Eq)] #[serde(rename_all = "snake_case")] pub enum AgentRuntimeActionType { @@ -639,6 +703,13 @@ impl ChatRunObservation { } } } + TauriAgentEvent::ArtifactSnapshot { artifact } => { + if let Some(path) = + normalize_metadata_path(artifact.file_path.as_str(), workspace_root) + { + self.record_artifact_path(path, request_metadata); + } + } _ => {} } } @@ -800,6 +871,16 @@ fn extract_harness_string( .map(str::to_string) } +fn extract_harness_bool( + request_metadata: Option<&serde_json::Value>, + keys: &[&str], +) -> Option { + let harness = extract_harness_object(request_metadata)?; + keys.iter() + .filter_map(|key| harness.get(*key)) + .find_map(serde_json::Value::as_bool) +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum RuntimeChatMode { Agent, @@ -826,6 +907,282 @@ fn default_web_search_enabled_for_chat_mode(_chat_mode: RuntimeChatMode) -> bool false } +fn execution_strategy_label(strategy: AsterExecutionStrategy) -> &'static str { + match strategy { + AsterExecutionStrategy::React => "对话执行优先", + AsterExecutionStrategy::CodeOrchestrated => "代码编排执行", + AsterExecutionStrategy::Auto => "自动路由执行", + } +} + +fn model_supports_reasoning(model_name: Option<&str>) -> bool { + let Some(model_name) = model_name.map(str::trim).filter(|value| !value.is_empty()) else { + return false; + }; + let normalized = model_name.to_ascii_lowercase(); + normalized.contains("thinking") + || normalized.contains("reason") + || normalized.contains("r1") + || normalized.contains("o1") + || normalized.contains("o3") + || normalized.contains("o4") + || normalized.contains("gpt-5") + || normalized.contains("2.5") +} + +fn message_suggests_live_search(message: &str) -> bool { + let normalized = message.to_ascii_lowercase(); + [ + "搜索", + "搜一下", + "查一下", + "查一查", + "检索", + "上网查", + "联网查", + "最新", + "今天", + "刚刚", + "实时", + "新闻", + "股价", + "汇率", + "天气", + "政策", + "法规", + "版本", + "价格", + "热搜", + "上线", + "发布", + "search", + "look up", + "google", + "browse", + "now", + "today", + "latest", + "recent", + "price", + "version", + "news", + "weather", + ] + .iter() + .any(|keyword| normalized.contains(keyword)) +} + +fn message_suggests_planning(message: &str) -> bool { + let normalized = message.to_ascii_lowercase(); + [ + "计划", + "规划", + "roadmap", + "拆解", + "分步骤", + "执行方案", + "实施方案", + "阶段", + "里程碑", + "todo", + ] + .iter() + .any(|keyword| normalized.contains(keyword)) +} + +fn message_suggests_task(message: &str) -> bool { + let normalized = message.to_ascii_lowercase(); + [ + "后台", + "稍后", + "异步", + "排队", + "持续生成", + "长时间", + "继续跑", + "持续跑", + ] + .iter() + .any(|keyword| normalized.contains(keyword)) +} + +fn message_suggests_subagent(message: &str) -> bool { + let normalized = message.to_ascii_lowercase(); + [ + "并行", + "多代理", + "分工", + "分别分析", + "从多个角度", + "parallel", + "subagent", + "delegate", + ] + .iter() + .any(|keyword| normalized.contains(keyword)) +} + +fn build_turn_runtime_statuses( + request: &AsterChatRequest, + effective_strategy: AsterExecutionStrategy, + request_tool_policy: &RequestToolPolicy, + model_name: Option<&str>, +) -> (TauriRuntimeStatus, TauriRuntimeStatus) { + let thinking_enabled = extract_harness_bool( + request.metadata.as_ref(), + &["thinking_enabled", "thinkingEnabled"], + ) + .unwrap_or(false); + let task_enabled = extract_harness_bool( + request.metadata.as_ref(), + &["task_mode_enabled", "taskModeEnabled"], + ) + .unwrap_or(false); + let subagent_enabled = extract_harness_bool( + request.metadata.as_ref(), + &["subagent_mode_enabled", "subagentModeEnabled"], + ) + .unwrap_or(false); + let reasoning_supported = model_supports_reasoning(model_name); + let news_expansion_needed = request_tool_policy.allows_web_search() + && message_suggests_news_expansion(&request.message); + + let initial_checkpoints = vec![ + execution_strategy_label(effective_strategy).to_string(), + if request_tool_policy.requires_web_search() { + "本回合必须先联网核实".to_string() + } else if news_expansion_needed { + "已识别新闻综述类输入,将先并发 WebSearch 扩搜".to_string() + } else if request_tool_policy.allows_web_search() { + "联网搜索仅作为候选能力待命".to_string() + } else { + "默认直接回答优先".to_string() + }, + if thinking_enabled && reasoning_supported { + "模型支持深度思考,先进入推理判定".to_string() + } else if thinking_enabled { + "当前模型不支持显式 thinking,改走轻量意图理解".to_string() + } else { + "先做轻量意图理解".to_string() + }, + if task_enabled { + "后台任务能力已待命".to_string() + } else { + "默认不升级后台任务".to_string() + }, + if subagent_enabled { + "多代理能力已待命".to_string() + } else { + "默认由单 Agent 先判断".to_string() + }, + ]; + + let decided = if request_tool_policy.requires_web_search() { + ( + "已决定:先联网检索".to_string(), + "当前回合已被明确指定为先搜索后答复,会先完成联网核实再继续生成。".to_string(), + vec![ + "用户明确要求联网搜索".to_string(), + "搜索结果返回后再形成最终答复".to_string(), + ], + ) + } else if news_expansion_needed { + ( + "已决定:先联网扩搜".to_string(), + "当前输入属于新闻/最新动态综述类请求,会先并发执行多组 WebSearch,再基于结果做主题聚类与交叉验证。" + .to_string(), + vec![ + "统一使用 WebSearch 执行多组扩搜".to_string(), + "完成来源整合后再组织最终答复".to_string(), + ], + ) + } else if subagent_enabled && message_suggests_subagent(&request.message) { + ( + "已决定:优先拆分为多代理".to_string(), + "用户输入更适合并行分工处理,先按多代理路径组织执行。".to_string(), + vec![ + "检测到并行/多角度需求".to_string(), + "主线程先承担协调职责".to_string(), + ], + ) + } else if task_enabled && message_suggests_task(&request.message) { + ( + "已决定:升级为后台任务".to_string(), + "用户输入更接近耗时或异步推进场景,优先走后台任务链路。".to_string(), + vec![ + "检测到排队/持续执行诉求".to_string(), + "先建立任务,再回传过程与产出".to_string(), + ], + ) + } else if thinking_enabled && reasoning_supported { + ( + "已决定:先深度思考".to_string(), + "当前模型支持 reasoning,先做更充分的意图理解与方案判断,再决定是否调用搜索或工具。" + .to_string(), + vec![ + "thinking 已开启".to_string(), + "搜索与工具保持候选状态,不默认触发".to_string(), + ], + ) + } else if thinking_enabled { + ( + "已决定:轻量理解后回答".to_string(), + "当前模型不支持显式 reasoning,先做轻量意图理解,再决定是否需要搜索或其他能力。" + .to_string(), + vec![ + "thinking 已开启".to_string(), + "当前模型回退为轻量推理".to_string(), + ], + ) + } else if request_tool_policy.allows_web_search() + && message_suggests_live_search(&request.message) + { + ( + "已决定:先联网核实".to_string(), + "问题包含明显时效性或实时性特征,先搜索核实再回答更稳妥。".to_string(), + vec![ + "已检测到最新/实时信息需求".to_string(), + "搜索完成后继续组织答复".to_string(), + ], + ) + } else if message_suggests_planning(&request.message) { + ( + "已决定:先规划再输出".to_string(), + "当前请求更像计划或方案拆解,会先整理执行路径和关键步骤。".to_string(), + vec![ + "检测到计划/拆解需求".to_string(), + "优先输出结构化行动路径".to_string(), + ], + ) + } else { + ( + "已决定:直接回答优先".to_string(), + "当前请求无需默认升级为搜索或任务,先直接给出结果,必要时再调用工具。".to_string(), + vec![ + "默认保持单回合直接回答".to_string(), + "只有证据不足或时效性要求出现时才升级".to_string(), + ], + ) + }; + + ( + TauriRuntimeStatus { + phase: "preparing".to_string(), + title: "正在理解意图".to_string(), + detail: + "正在判断当前回合应该直接回答、深度思考、规划、联网核实,还是升级为任务/多代理。" + .to_string(), + checkpoints: initial_checkpoints, + }, + TauriRuntimeStatus { + phase: "routing".to_string(), + title: decided.0, + detail: decided.1, + checkpoints: decided.2, + }, + ) +} + fn extend_map_with_harness_fields( target: &mut serde_json::Map, request_metadata: Option<&serde_json::Value>, @@ -893,6 +1250,10 @@ fn build_chat_run_metadata_base( "web_search_enabled".to_string(), serde_json::json!(request_tool_policy.effective_web_search), ); + metadata.insert( + "web_search_mode".to_string(), + serde_json::json!(request_tool_policy.search_mode.as_str()), + ); metadata.insert( "auto_continue_enabled".to_string(), serde_json::json!(auto_continue_enabled), @@ -4514,23 +4875,22 @@ async fn apply_workspace_sandbox_permissions( /// 图片输入 #[allow(dead_code)] -#[derive(Debug, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct ImageInput { pub data: String, pub media_type: String, } -/// 发送消息并获取流式响应 -#[tauri::command] -pub async fn aster_agent_chat_stream( - app: AppHandle, - state: State<'_, AsterAgentState>, - db: State<'_, DbConnection>, - api_key_provider_service: State<'_, ApiKeyProviderServiceState>, - logs: State<'_, LogState>, - config_manager: State<'_, GlobalConfigManagerState>, - mcp_manager: State<'_, McpManagerState>, - heartbeat_state: State<'_, HeartbeatServiceState>, +/// 执行单个 turn 的流式响应 +async fn execute_aster_chat_request( + app: &AppHandle, + state: &AsterAgentState, + db: &DbConnection, + api_key_provider_service: &ApiKeyProviderServiceState, + logs: &LogState, + config_manager: &GlobalConfigManagerState, + mcp_manager: &McpManagerState, + heartbeat_state: &HeartbeatServiceState, request: AsterChatRequest, ) -> Result<(), String> { tracing::info!( @@ -4544,7 +4904,7 @@ pub async fn aster_agent_chat_stream( tracing::warn!("[AsterAgent] Agent 初始化状态: {}", is_init); if !is_init { tracing::warn!("[AsterAgent] Agent 未初始化,开始初始化..."); - state.init_agent_with_db(&db).await?; + state.init_agent_with_db(db).await?; tracing::warn!("[AsterAgent] Agent 初始化完成"); } else { tracing::warn!("[AsterAgent] Agent 已初始化,检查 session_store..."); @@ -4556,7 +4916,7 @@ pub async fn aster_agent_chat_stream( tracing::warn!("[AsterAgent] session_store 存在: {}", has_store); } } - ensure_social_image_tool_registered(state.inner(), config_manager.inner()).await?; + ensure_social_image_tool_registered(state, config_manager).await?; // 直接使用前端传递的 session_id // ProxyCastSessionStore 会在 add_message 时自动创建不存在的 session @@ -4572,7 +4932,7 @@ pub async fn aster_agent_chat_stream( return Err(message); } - let manager = WorkspaceManager::new(db.inner().clone()); + let manager = WorkspaceManager::new(db.clone()); let workspace = match manager.get(&workspace_id) { Ok(Some(workspace)) => workspace, Ok(None) => { @@ -4665,7 +5025,7 @@ pub async fn aster_agent_chat_stream( } // 启动并注入 MCP extensions 到 Aster Agent - let (_start_ok, start_fail) = ensure_proxycast_mcp_servers_running(&db, &mcp_manager).await; + let (_start_ok, start_fail) = ensure_proxycast_mcp_servers_running(db, mcp_manager).await; if start_fail > 0 { tracing::warn!( "[AsterAgent] 部分 MCP server 自动启动失败 ({} 失败),后续可用工具可能不完整", @@ -4673,7 +5033,7 @@ pub async fn aster_agent_chat_stream( ); } - let (_mcp_ok, mcp_fail) = inject_mcp_extensions(&state, &mcp_manager).await; + let (_mcp_ok, mcp_fail) = inject_mcp_extensions(state, mcp_manager).await; if mcp_fail > 0 { tracing::warn!( "[AsterAgent] 部分 MCP extension 注入失败 ({} 失败),Agent 可能无法使用某些 MCP 工具", @@ -4684,16 +5044,23 @@ pub async fn aster_agent_chat_stream( let runtime_chat_mode = resolve_runtime_chat_mode(request.metadata.as_ref()); let mode_default_web_search = default_web_search_enabled_for_chat_mode(runtime_chat_mode); - // 构建请求级工具策略:默认不强制联网搜索,仅在用户显式开启开关时把搜索升级为必需步骤。 - let request_tool_policy = - resolve_request_tool_policy(request.web_search, mode_default_web_search); + // 构建请求级工具策略: + // - web_search=true 默认只表示“允许搜索” + // - 仅显式 search_mode=required 时才强制预搜索 + let request_tool_policy = resolve_request_tool_policy_with_mode( + request.web_search, + request.search_mode, + mode_default_web_search, + ); tracing::info!( - "[AsterAgent][WebSearchGuard] session={}, chat_mode={:?}, request_web_search={:?}, mode_default_web_search={}, effective_web_search={}", + "[AsterAgent][WebSearchGuard] session={}, chat_mode={:?}, request_web_search={:?}, request_search_mode={:?}, mode_default_web_search={}, effective_web_search={}, search_mode={}", session_id, runtime_chat_mode, request.web_search, + request.search_mode, mode_default_web_search, - request_tool_policy.effective_web_search + request_tool_policy.effective_web_search, + request_tool_policy.search_mode.as_str() ); // 构建 system_prompt:优先使用项目上下文,其次使用 session 的 system_prompt @@ -4709,7 +5076,7 @@ pub async fn aster_agent_chat_stream( // 1. 如果提供了 project_id,构建项目上下文 let project_prompt = if let Some(ref project_id) = request.project_id { - match AsterAgentState::build_project_system_prompt(&db, project_id) { + match AsterAgentState::build_project_system_prompt(db, project_id) { Ok(prompt) => { tracing::info!( "[AsterAgent] 已加载项目上下文: project_id={}, prompt_len={}", @@ -4832,7 +5199,7 @@ pub async fn aster_agent_chat_stream( }; // 如果前端提供了 api_key,直接使用;否则从凭证池选择凭证 if provider_config.api_key.is_some() { - state.configure_provider(config, session_id, &db).await?; + state.configure_provider(config, session_id, db).await?; } else { // 没有 api_key,使用凭证池(优先 provider_id,其次 provider_name) let provider_selector = provider_config @@ -4841,7 +5208,7 @@ pub async fn aster_agent_chat_stream( .unwrap_or(&provider_config.provider_name); state .configure_provider_from_pool( - &db, + db, provider_selector, &provider_config.model_name, session_id, @@ -4856,12 +5223,12 @@ pub async fn aster_agent_chat_stream( } let sandbox_outcome = apply_workspace_sandbox_permissions( - &state, - config_manager.inner(), - db.inner(), - api_key_provider_service.inner(), - heartbeat_state.inner(), - &app, + state, + config_manager, + db, + api_key_provider_service, + heartbeat_state, + app, &workspace_root, requested_strategy, ) @@ -4903,7 +5270,7 @@ pub async fn aster_agent_chat_stream( } } - let tracker = ExecutionTracker::new(db.inner().clone()); + let tracker = ExecutionTracker::new(db.clone()); let cancel_token = state.create_cancel_token(session_id).await; let auto_continue_metadata = auto_continue_config.clone(); let request_metadata = request.metadata.clone(); @@ -4919,7 +5286,7 @@ pub async fn aster_agent_chat_stream( let run_observation_for_finalize = run_observation.clone(); let run_start_metadata_for_finalize = run_start_metadata.clone(); let timeline_recorder = Arc::new(Mutex::new(AgentTimelineRecorder::create( - db.inner().clone(), + db.clone(), session_id.to_string(), request.message.clone(), )?)); @@ -4929,7 +5296,32 @@ pub async fn aster_agent_chat_stream( Ok(guard) => guard, Err(error) => error.into_inner(), }; - recorder.emit_start(&app, &request.event_name)?; + recorder.emit_start(app, &request.event_name)?; + } + + let (initial_runtime_status, decided_runtime_status) = build_turn_runtime_statuses( + &request, + effective_strategy, + &request_tool_policy, + request + .provider_config + .as_ref() + .map(|config| config.model_name.as_str()), + ); + for status in [initial_runtime_status, decided_runtime_status] { + let event = TauriAgentEvent::RuntimeStatus { status }; + if let Err(error) = app.emit(&request.event_name, &event) { + tracing::warn!("[AsterAgent] 发送 runtime_status 失败: {}", error); + } + let mut recorder = match timeline_recorder.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + if let Err(error) = + recorder.record_legacy_event(app, &request.event_name, &event, workspace_root.as_str()) + { + tracing::warn!("[AsterAgent] 记录 runtime_status 失败: {}", error); + } } // 获取 Agent Arc 并保持 guard 在整个流处理期间存活 @@ -4963,7 +5355,7 @@ pub async fn aster_agent_chat_stream( let primary_result = stream_reply_once( agent, - &app, + app, &request.event_name, &request.message, Some(Path::new(&workspace_root)), @@ -5137,7 +5529,7 @@ pub async fn aster_agent_chat_stream( Ok(guard) => guard, Err(error) => error.into_inner(), }; - if let Err(error) = recorder.complete_turn_success(&app, &request.event_name) { + if let Err(error) = recorder.complete_turn_success(app, &request.event_name) { tracing::warn!("[AsterAgent] 完成 turn 时间线失败(已降级继续): {}", error); } } @@ -5152,7 +5544,7 @@ pub async fn aster_agent_chat_stream( Ok(guard) => guard, Err(error) => error.into_inner(), }; - if let Err(timeline_error) = recorder.fail_turn(&app, &request.event_name, &e) { + if let Err(timeline_error) = recorder.fail_turn(app, &request.event_name, &e) { tracing::warn!( "[AsterAgent] 记录失败 turn 时间线失败(已降级继续): {}", timeline_error @@ -5174,6 +5566,420 @@ pub async fn aster_agent_chat_stream( Ok(()) } +/// 发送消息并获取流式响应 +#[tauri::command] +pub async fn aster_agent_chat_stream( + app: AppHandle, + state: State<'_, AsterAgentState>, + db: State<'_, DbConnection>, + api_key_provider_service: State<'_, ApiKeyProviderServiceState>, + logs: State<'_, LogState>, + config_manager: State<'_, GlobalConfigManagerState>, + mcp_manager: State<'_, McpManagerState>, + heartbeat_state: State<'_, HeartbeatServiceState>, + request: AsterChatRequest, +) -> Result<(), String> { + execute_aster_chat_request( + &app, + state.inner(), + db.inner(), + api_key_provider_service.inner(), + logs.inner(), + config_manager.inner(), + mcp_manager.inner(), + heartbeat_state.inner(), + request, + ) + .await +} + +struct AgentRuntimeExecutionContext { + app: AppHandle, + state: AsterAgentState, + db: DbConnection, + api_key_provider_service: ApiKeyProviderServiceState, + logs: LogState, + config_manager: GlobalConfigManagerState, + mcp_manager: McpManagerState, + heartbeat_state: HeartbeatServiceState, +} + +impl AgentRuntimeExecutionContext { + fn from_states( + app: AppHandle, + state: &AsterAgentState, + db: &DbConnection, + api_key_provider_service: &ApiKeyProviderServiceState, + logs: &LogState, + config_manager: &GlobalConfigManagerState, + mcp_manager: &McpManagerState, + heartbeat_state: &HeartbeatServiceState, + ) -> Self { + Self { + app, + state: state.clone(), + db: db.clone(), + api_key_provider_service: ApiKeyProviderServiceState( + api_key_provider_service.0.clone(), + ), + logs: logs.clone(), + config_manager: GlobalConfigManagerState(config_manager.0.clone()), + mcp_manager: mcp_manager.clone(), + heartbeat_state: heartbeat_state.clone(), + } + } +} + +impl Clone for AgentRuntimeExecutionContext { + fn clone(&self) -> Self { + Self { + app: self.app.clone(), + state: self.state.clone(), + db: self.db.clone(), + api_key_provider_service: ApiKeyProviderServiceState( + self.api_key_provider_service.0.clone(), + ), + logs: self.logs.clone(), + config_manager: GlobalConfigManagerState(self.config_manager.0.clone()), + mcp_manager: self.mcp_manager.clone(), + heartbeat_state: self.heartbeat_state.clone(), + } + } +} + +fn build_queued_turn_preview(message: &str) -> String { + let compact = message.split_whitespace().collect::>().join(" "); + if compact.is_empty() { + return "空白输入".to_string(); + } + + let preview = compact.chars().take(80).collect::(); + if compact.chars().count() > 80 { + format!("{preview}...") + } else { + preview + } +} + +fn build_queued_turn_task( + mut request: AsterChatRequest, +) -> Result, String> { + let queued_turn_id = request + .queued_turn_id + .clone() + .unwrap_or_else(|| Uuid::new_v4().to_string()); + request.queued_turn_id = Some(queued_turn_id.clone()); + + let image_count = request + .images + .as_ref() + .map(|images| images.len()) + .unwrap_or(0); + let payload = + serde_json::to_value(&request).map_err(|e| format!("序列化排队 turn 失败: {e}"))?; + + Ok(QueuedTurnTask { + queued_turn_id, + session_id: request.session_id.clone(), + event_name: request.event_name.clone(), + message_preview: build_queued_turn_preview(&request.message), + message_text: request.message.clone(), + created_at: chrono::Utc::now().timestamp_millis(), + image_count, + payload, + }) +} + +fn deserialize_queued_turn_request(payload: serde_json::Value) -> Result { + serde_json::from_value(payload).map_err(|e| format!("反序列化排队 turn 失败: {e}")) +} + +fn persist_runtime_queued_turn( + db: &DbConnection, + task: &QueuedTurnTask, +) -> Result<(), String> { + let payload_json = serde_json::to_string(&task.payload) + .map_err(|e| format!("序列化排队 turn 持久化 payload 失败: {e}"))?; + let conn = crate::database::lock_db(db)?; + AgentRuntimeQueuedTurnDao::insert( + &conn, + &NewAgentRuntimeQueuedTurnRecord { + queued_turn_id: task.queued_turn_id.clone(), + session_id: task.session_id.clone(), + event_name: task.event_name.clone(), + message_preview: task.message_preview.clone(), + message_text: task.message_text.clone(), + payload_json, + image_count: task.image_count, + created_at: task.created_at, + }, + ) + .map_err(|e| format!("持久化排队 turn 失败: {e}"))?; + Ok(()) +} + +fn remove_persisted_runtime_queued_turn( + db: &DbConnection, + queued_turn_id: &str, +) -> Result { + let conn = crate::database::lock_db(db)?; + AgentRuntimeQueuedTurnDao::remove(&conn, queued_turn_id) + .map_err(|e| format!("删除持久化排队 turn 失败: {e}")) +} + +fn list_persisted_runtime_queue_session_ids(db: &DbConnection) -> Result, String> { + let conn = crate::database::lock_db(db)?; + AgentRuntimeQueuedTurnDao::list_distinct_session_ids(&conn) + .map_err(|e| format!("读取排队会话列表失败: {e}")) +} + +fn load_persisted_runtime_queue_tasks( + db: &DbConnection, + session_id: &str, +) -> Result>, String> { + let conn = crate::database::lock_db(db)?; + let records = AgentRuntimeQueuedTurnDao::list_by_session(&conn, session_id) + .map_err(|e| format!("读取持久化排队 turn 失败: {e}"))?; + + let mut tasks = Vec::with_capacity(records.len()); + let mut invalid_ids = Vec::new(); + for record in records { + match serde_json::from_str::(&record.payload_json) { + Ok(payload) => tasks.push(QueuedTurnTask { + queued_turn_id: record.queued_turn_id, + session_id: record.session_id, + event_name: record.event_name, + message_preview: record.message_preview, + message_text: record.message_text, + created_at: record.created_at, + image_count: record.image_count, + payload, + }), + Err(error) => { + tracing::warn!( + "[AsterAgent][Queue] 跳过损坏的持久化排队 turn: session_id={}, queued_turn_id={}, error={}", + session_id, + record.queued_turn_id, + error + ); + invalid_ids.push(record.queued_turn_id); + } + } + } + + for queued_turn_id in invalid_ids { + if let Err(error) = AgentRuntimeQueuedTurnDao::remove(&conn, &queued_turn_id) { + tracing::warn!( + "[AsterAgent][Queue] 删除损坏的持久化排队 turn 失败: queued_turn_id={}, error={}", + queued_turn_id, + error + ); + } + } + + Ok(tasks) +} + +fn ensure_runtime_queue_loaded( + state: &AsterAgentState, + db: &DbConnection, + session_id: &str, +) -> Result<(), String> { + if state.turn_queue().has_session_state(session_id) { + return Ok(()); + } + + let tasks = load_persisted_runtime_queue_tasks(db, session_id)?; + if tasks.is_empty() { + return Ok(()); + } + + tracing::info!( + "[AsterAgent][Queue] 从持久化存储恢复会话排队: session_id={}, count={}", + session_id, + tasks.len() + ); + state.turn_queue().restore_pending(session_id, tasks); + Ok(()) +} + +fn emit_runtime_queue_event(app: &AppHandle, event_name: &str, event: &TauriAgentEvent) { + if let Err(error) = app.emit(event_name, event) { + tracing::warn!( + "[AsterAgent][Queue] 发送队列事件失败: event_name={}, error={}", + event_name, + error + ); + } +} + +fn schedule_next_runtime_turn(context: AgentRuntimeExecutionContext, session_id: String) { + let queue = context.state.turn_queue(); + + loop { + let Some(next_task) = queue.finish_and_take_next(&session_id) else { + return; + }; + + if let Err(error) = + remove_persisted_runtime_queued_turn(&context.db, &next_task.queued_turn_id) + { + tracing::warn!( + "[AsterAgent][Queue] 删除已启动的持久化排队 turn 失败: queued_turn_id={}, error={}", + next_task.queued_turn_id, + error + ); + } + + emit_runtime_queue_event( + &context.app, + &next_task.event_name, + &TauriAgentEvent::QueueStarted { + session_id: session_id.clone(), + queued_turn_id: next_task.queued_turn_id.clone(), + }, + ); + + let next_request = match deserialize_queued_turn_request(next_task.payload) { + Ok(request) => request, + Err(error) => { + emit_runtime_queue_event( + &context.app, + &next_task.event_name, + &TauriAgentEvent::Error { message: error }, + ); + continue; + } + }; + + tokio::spawn(async move { + if let Err(error) = + execute_runtime_turn_and_continue_queue(context.clone(), next_request).await + { + tracing::warn!("[AsterAgent][Queue] 队列任务执行失败: {}", error); + } + }); + return; + } +} + +async fn execute_runtime_turn_and_continue_queue( + context: AgentRuntimeExecutionContext, + request: AsterChatRequest, +) -> Result<(), String> { + let session_id = request.session_id.clone(); + let result = execute_aster_chat_request( + &context.app, + &context.state, + &context.db, + &context.api_key_provider_service, + &context.logs, + &context.config_manager, + &context.mcp_manager, + &context.heartbeat_state, + request, + ) + .await; + + schedule_next_runtime_turn(context, session_id); + result +} + +fn resume_runtime_queue_if_needed( + context: AgentRuntimeExecutionContext, + session_id: String, +) -> Result { + ensure_runtime_queue_loaded(&context.state, &context.db, &session_id)?; + + if context.state.turn_queue().has_active(&session_id) { + return Ok(false); + } + + if context.state.turn_queue().snapshot(&session_id).is_empty() { + return Ok(false); + } + + schedule_next_runtime_turn(context, session_id); + Ok(true) +} + +fn clear_pending_runtime_queue( + app: &AppHandle, + state: &AsterAgentState, + db: &DbConnection, + session_id: &str, +) -> Vec> { + let cleared = state.turn_queue().clear_pending(session_id); + if cleared.is_empty() { + return cleared; + } + + let queued_turn_ids = cleared + .iter() + .map(|task| task.queued_turn_id.clone()) + .collect::>(); + for queued_turn_id in &queued_turn_ids { + if let Err(error) = remove_persisted_runtime_queued_turn(db, queued_turn_id) { + tracing::warn!( + "[AsterAgent][Queue] 删除已清空的持久化排队 turn 失败: queued_turn_id={}, error={}", + queued_turn_id, + error + ); + } + } + for task in &cleared { + emit_runtime_queue_event( + app, + &task.event_name, + &TauriAgentEvent::QueueCleared { + session_id: session_id.to_string(), + queued_turn_ids: queued_turn_ids.clone(), + }, + ); + } + + cleared +} + +pub fn resume_persisted_runtime_queues_on_startup( + app: AppHandle, + state: &AsterAgentState, + db: &DbConnection, + api_key_provider_service: &ApiKeyProviderServiceState, + logs: &LogState, + config_manager: &GlobalConfigManagerState, + mcp_manager: &McpManagerState, + heartbeat_state: &HeartbeatServiceState, +) -> Result { + let session_ids = list_persisted_runtime_queue_session_ids(db)?; + if session_ids.is_empty() { + return Ok(0); + } + + let mut resumed = 0usize; + for session_id in session_ids { + let context = AgentRuntimeExecutionContext::from_states( + app.clone(), + state, + db, + api_key_provider_service, + logs, + config_manager, + mcp_manager, + heartbeat_state, + ); + if resume_runtime_queue_if_needed(context, session_id.clone())? { + resumed += 1; + tracing::info!( + "[AsterAgent][Queue] 启动阶段已恢复会话排队执行: session_id={}", + session_id + ); + } + } + + Ok(resumed) +} + /// 停止当前会话 #[tauri::command] pub async fn aster_agent_stop( @@ -5197,27 +6003,68 @@ pub async fn agent_runtime_submit_turn( heartbeat_state: State<'_, HeartbeatServiceState>, request: AgentRuntimeSubmitTurnRequest, ) -> Result<(), String> { - aster_agent_chat_stream( + let runtime_request: AsterChatRequest = request.into(); + let queue_if_busy = runtime_request.queue_if_busy.unwrap_or(false); + let queued_task = build_queued_turn_task(runtime_request)?; + let session_id = queued_task.session_id.clone(); + ensure_runtime_queue_loaded(state.inner(), db.inner(), &session_id)?; + let context = AgentRuntimeExecutionContext::from_states( app, - state, - db, - api_key_provider_service, - logs, - config_manager, - mcp_manager, - heartbeat_state, - request.into(), - ) - .await + state.inner(), + db.inner(), + api_key_provider_service.inner(), + logs.inner(), + config_manager.inner(), + mcp_manager.inner(), + heartbeat_state.inner(), + ); + + let _ = resume_runtime_queue_if_needed(context.clone(), session_id.clone())?; + + if !queue_if_busy && state.inner().turn_queue().has_active(&session_id) { + return Err("当前会话仍在生成,无法立即开始执行".to_string()); + } + + match state + .inner() + .turn_queue() + .start_or_enqueue(queued_task.clone()) + { + QueueInsertResult::StartNow(task) => { + let request = deserialize_queued_turn_request(task.payload)?; + execute_runtime_turn_and_continue_queue(context, request).await + } + QueueInsertResult::Enqueued { + event_name, + snapshot, + } => { + persist_runtime_queued_turn(db.inner(), &queued_task)?; + emit_runtime_queue_event( + &context.app, + &event_name, + &TauriAgentEvent::QueueAdded { + session_id, + queued_turn: snapshot, + }, + ); + Ok(()) + } + } } /// 统一运行时:中断当前 turn。 #[tauri::command] pub async fn agent_runtime_interrupt_turn( + app: AppHandle, state: State<'_, AsterAgentState>, + db: State<'_, DbConnection>, request: AgentRuntimeInterruptTurnRequest, ) -> Result { - aster_agent_stop(state, request.session_id).await + let session_id = request.session_id; + ensure_runtime_queue_loaded(state.inner(), db.inner(), &session_id)?; + let cancelled = state.cancel_session(&session_id).await; + let cleared = clear_pending_runtime_queue(&app, state.inner(), db.inner(), &session_id); + Ok(cancelled || !cleared.is_empty()) } /// 创建新会话 @@ -5329,10 +6176,77 @@ pub async fn aster_session_get( /// 统一运行时:获取会话详情。 #[tauri::command] pub async fn agent_runtime_get_session( + app: AppHandle, + state: State<'_, AsterAgentState>, db: State<'_, DbConnection>, + api_key_provider_service: State<'_, ApiKeyProviderServiceState>, + logs: State<'_, LogState>, + config_manager: State<'_, GlobalConfigManagerState>, + mcp_manager: State<'_, McpManagerState>, + heartbeat_state: State<'_, HeartbeatServiceState>, session_id: String, -) -> Result { - aster_session_get(db, session_id).await +) -> Result { + ensure_runtime_queue_loaded(state.inner(), db.inner(), &session_id)?; + let detail = AsterAgentWrapper::get_session_sync(db.inner(), &session_id)?; + let queued_turns = state.inner().turn_queue().snapshot(&session_id); + if !queued_turns.is_empty() && !state.inner().turn_queue().has_active(&session_id) { + let context = AgentRuntimeExecutionContext::from_states( + app, + state.inner(), + db.inner(), + api_key_provider_service.inner(), + logs.inner(), + config_manager.inner(), + mcp_manager.inner(), + heartbeat_state.inner(), + ); + if let Err(error) = resume_runtime_queue_if_needed(context, session_id.clone()) { + tracing::warn!( + "[AsterAgent][Queue] 获取会话后恢复排队执行失败: session_id={}, error={}", + session_id, + error + ); + } + } + Ok(AgentRuntimeSessionDetail::from_session_detail( + detail, + queued_turns, + )) +} + +/// 统一运行时:移除单个排队 turn。 +#[tauri::command] +pub async fn agent_runtime_remove_queued_turn( + app: AppHandle, + state: State<'_, AsterAgentState>, + db: State<'_, DbConnection>, + request: AgentRuntimeRemoveQueuedTurnRequest, +) -> Result { + let session_id = request.session_id.trim().to_string(); + let queued_turn_id = request.queued_turn_id.trim().to_string(); + if session_id.is_empty() || queued_turn_id.is_empty() { + return Ok(false); + } + + ensure_runtime_queue_loaded(state.inner(), db.inner(), &session_id)?; + let removed = state + .inner() + .turn_queue() + .remove_queued(&session_id, &queued_turn_id); + if let Some(task) = removed { + remove_persisted_runtime_queued_turn(db.inner(), &queued_turn_id)?; + emit_runtime_queue_event( + &app, + &task.event_name, + &TauriAgentEvent::QueueRemoved { + session_id, + queued_turn_id, + }, + ); + return Ok(true); + } + + Ok(false) } /// 重命名会话 @@ -5389,10 +6303,15 @@ pub async fn aster_session_delete( /// 统一运行时:删除会话。 #[tauri::command] pub async fn agent_runtime_delete_session( + app: AppHandle, + state: State<'_, AsterAgentState>, db: State<'_, DbConnection>, session_id: String, ) -> Result<(), String> { - aster_session_delete(db, session_id).await + let trimmed_session_id = session_id.trim().to_string(); + let _ = state.cancel_session(&trimmed_session_id).await; + let _ = clear_pending_runtime_queue(&app, state.inner(), db.inner(), &trimmed_session_id); + aster_session_delete(db, trimmed_session_id).await } /// 确认权限请求 @@ -5579,6 +6498,7 @@ pub async fn aster_agent_submit_elicitation_response( mod tests { use super::*; use async_trait::async_trait; + use proxycast_agent::request_tool_policy::resolve_request_tool_policy; use regex::Regex; use std::ffi::OsString; use std::path::{Path, PathBuf}; @@ -5669,6 +6589,19 @@ mod tests { assert_eq!(request.auto_continue, None); } + #[test] + fn test_message_suggests_live_search_accepts_explicit_search_verbs() { + assert!(message_suggests_live_search( + "请帮我搜一下哥德尔不完备定理的历史背景" + )); + assert!(message_suggests_live_search( + "please look up kyoto travel tips" + )); + assert!(!message_suggests_live_search( + "帮我解释一下什么是向量数据库" + )); + } + #[test] fn test_aster_chat_request_deserialize_with_execution_strategy() { let json = r#"{ @@ -5950,6 +6883,7 @@ mod tests { project_id: Some("project-1".to_string()), workspace_id: "workspace-1".to_string(), web_search: Some(false), + search_mode: None, execution_strategy: Some(AsterExecutionStrategy::React), auto_continue: None, system_prompt: None, @@ -5959,10 +6893,13 @@ mod tests { "gate_key": "write_mode" } })), + queue_if_busy: None, + queued_turn_id: None, }, "workspace-1", AsterExecutionStrategy::React, &RequestToolPolicy { + search_mode: RequestToolPolicyMode::Disabled, effective_web_search: false, required_tools: vec![], allowed_tools: vec![], diff --git a/src-tauri/src/commands/skill_cmd.rs b/src-tauri/src/commands/skill_cmd.rs index e57a21e2c..c966fee25 100644 --- a/src-tauri/src/commands/skill_cmd.rs +++ b/src-tauri/src/commands/skill_cmd.rs @@ -2,10 +2,14 @@ use crate::agent::aster_state::AsterAgentState; use crate::database::dao::skills::SkillDao; use crate::database::DbConnection; use crate::models::app_type::AppType; -use crate::models::skill_model::{Skill, SkillRepo, SkillState}; +use crate::models::skill_model::{ + Skill, SkillCatalogSource, SkillPackageInspection, SkillRepo, SkillState, +}; use chrono::Utc; use proxycast_core::app_paths; use proxycast_services::skill_service::SkillService; +use serde::Serialize; +use std::fs; use std::path::{Component, Path, PathBuf}; use std::sync::Arc; use tauri::State; @@ -58,6 +62,13 @@ fn get_skills_dir(app_type: &AppType) -> Result { } } +fn get_skill_lookup_roots(app_type: &AppType) -> Result, String> { + match app_type { + AppType::ProxyCast => app_paths::resolve_proxycast_skill_roots(), + _ => Ok(vec![get_skills_dir(app_type)?]), + } +} + fn validate_skill_directory(directory: &str) -> Result<(), String> { if directory.trim().is_empty() { return Err("Skill directory is required".to_string()); @@ -82,11 +93,14 @@ fn validate_skill_directory(directory: &str) -> Result<(), String> { } } -fn read_local_skill_content(skills_dir: &Path, directory: &str) -> Result { +fn try_resolve_local_skill_dir( + skills_dir: &Path, + directory: &str, +) -> Result, String> { validate_skill_directory(directory)?; if !skills_dir.exists() { - return Err("Skills directory not found".to_string()); + return Ok(None); } let canonical_skills_dir = skills_dir @@ -95,7 +109,7 @@ fn read_local_skill_content(skills_dir: &Path, directory: &str) -> Result Result Result { + validate_skill_directory(directory)?; + + for root in skill_roots { + if let Some(skill_dir) = try_resolve_local_skill_dir(root, directory)? { + return Ok(skill_dir); + } + } + + Err(format!("Skill not found: {directory}")) +} + +fn inspect_local_skill( + skill_roots: &[PathBuf], + directory: &str, +) -> Result { + let skill_dir = resolve_local_skill_dir(skill_roots, directory)?; + SkillService::inspect_skill_dir(&skill_dir).map_err(|e| e.to_string()) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum SkillScaffoldTarget { + Project, + User, +} + +impl SkillScaffoldTarget { + fn parse(value: &str) -> Result { + match value.trim().to_ascii_lowercase().as_str() { + "project" => Ok(Self::Project), + "user" => Ok(Self::User), + _ => Err(format!("Unsupported scaffold target: {value}")), + } + } +} + +#[derive(Serialize)] +struct SkillScaffoldFrontmatter<'a> { + name: &'a str, + description: &'a str, +} + +fn resolve_skill_scaffold_root( + app_type: &AppType, + target: SkillScaffoldTarget, +) -> Result { + match target { + SkillScaffoldTarget::User => get_skills_dir(app_type), + SkillScaffoldTarget::Project => match app_type { + AppType::ProxyCast => app_paths::resolve_project_skills_dir() + .ok_or_else(|| "Failed to resolve project skills directory".to_string()), + _ => Err("Project skill scaffold is only supported for proxycast".to_string()), + }, + } +} + +fn build_skill_scaffold_content(name: &str, description: &str) -> Result { + let frontmatter = serde_yaml::to_string(&SkillScaffoldFrontmatter { name, description }) + .map_err(|e| format!("Failed to build skill frontmatter: {e}"))?; + let frontmatter = frontmatter.strip_prefix("---\n").unwrap_or(&frontmatter); + + Ok(format!( + "---\n{frontmatter}---\n\n# {name}\n\n## 何时使用\n- 描述该 Skill 的适用场景。\n\n## 输入\n- 说明用户需要提供的上下文、约束和素材。\n\n## 执行要求\n1. 先明确目标、边界和输出格式。\n2. 如需引用资料,请将文件放到 `references/` 目录。\n3. 如需脚本或素材,请分别放到 `scripts/` 与 `assets/` 目录。\n\n## 输出\n- 说明最终交付物及验收标准。\n" + )) +} + +fn create_skill_scaffold_in_root( + skills_root: &Path, + directory: &str, + name: &str, + description: &str, +) -> Result { + validate_skill_directory(directory)?; + + let name = name.trim(); + if name.is_empty() { + return Err("Skill name is required".to_string()); + } + + let description = description.trim(); + if description.is_empty() { + return Err("Skill description is required".to_string()); + } + + fs::create_dir_all(skills_root).map_err(|e| { + format!( + "Failed to create skills root {}: {e}", + skills_root.display() + ) + })?; + + let skill_dir = skills_root.join(directory); + if skill_dir.exists() { + return Err(format!("Skill directory already exists: {directory}")); + } + + fs::create_dir_all(&skill_dir).map_err(|e| { + format!( + "Failed to create skill directory {}: {e}", + skill_dir.display() + ) + })?; + + let skill_md_content = build_skill_scaffold_content(name, description)?; + let skill_md_path = skill_dir.join("SKILL.md"); + if let Err(error) = fs::write(&skill_md_path, skill_md_content) { + let _ = fs::remove_dir_all(&skill_dir); + return Err(format!( + "Failed to write scaffold file {}: {error}", + skill_md_path.display() + )); + } + + match SkillService::inspect_skill_dir(&skill_dir) { + Ok(inspection) => Ok(inspection), + Err(error) => { + let _ = fs::remove_dir_all(&skill_dir); + Err(format!("Created scaffold failed inspection: {error}")) + } + } } /// 获取已安装的 ProxyCast Skills 目录列表 @@ -128,22 +264,70 @@ pub async fn get_installed_proxycast_skills() -> Result, String> { Ok(scan_installed_skills(&skills_dir)) } -/// 获取本地已安装 Skill 的 SKILL.md 内容 +/// 获取本地已安装 Skill 的标准检查结果 /// -/// 仅支持读取本地 Skills 目录下的文件,包含目录合法性与路径穿越防护。 +/// 仅支持读取本地 Skills 目录下的文件,包含目录合法性、路径穿越防护、 +/// Agent Skills 标准检查和 ProxyCast 扩展引用校验。 /// /// # Arguments /// - `app`: 应用类型(proxycast/claude/codex/gemini) /// - `directory`: Skill 目录名 /// /// # Returns -/// - `Ok(String)`: SKILL.md 的文本内容 +/// - `Ok(SkillPackageInspection)`: Skill 检查结果与原始内容 /// - `Err(String)`: 错误信息 #[tauri::command] -pub fn get_local_skill_content(app: String, directory: String) -> Result { +pub fn inspect_local_skill_for_app( + app: String, + directory: String, +) -> Result { let app_type: AppType = app.parse().map_err(|e: String| e)?; - let skills_dir = get_skills_dir(&app_type)?; - read_local_skill_content(&skills_dir, &directory) + let skill_roots = get_skill_lookup_roots(&app_type)?; + inspect_local_skill(&skill_roots, &directory) +} + +/// 创建标准 Skill 脚手架 +/// +/// 在项目级或用户级 Skills root 下创建一个最小 Agent Skills 标准包, +/// 并返回创建后的 inspection 结果,供 UI 立即预览。 +#[tauri::command] +pub fn create_skill_scaffold_for_app( + app: String, + target: String, + directory: String, + name: String, + description: String, +) -> Result { + let app_type: AppType = app.parse().map_err(|e: String| e)?; + let target = SkillScaffoldTarget::parse(&target)?; + let skills_root = resolve_skill_scaffold_root(&app_type, target)?; + let inspection = create_skill_scaffold_in_root(&skills_root, &directory, &name, &description)?; + + if matches!(app_type, AppType::ProxyCast) { + AsterAgentState::reload_proxycast_skills(); + } + + Ok(inspection) +} + +/// 获取远程 Skill 包的标准检查结果 +/// +/// 直接从远程仓库读取目标 Skill 目录,返回标准检查结果与原始 SKILL.md, +/// 用于安装前预检和 workflow/reference 可见性。 +#[tauri::command] +pub async fn inspect_remote_skill( + skill_service: State<'_, SkillServiceState>, + owner: String, + name: String, + branch: String, + directory: String, +) -> Result { + validate_skill_directory(&directory)?; + skill_service + .0 + .inspect_remote_skill(&owner, &name, &branch, &directory) + .await + .map_err(|e| e.to_string()) } pub struct SkillServiceState(pub Arc); @@ -211,7 +395,7 @@ pub async fn get_skills_for_app( let existing_states = SkillDao::get_skills(&conn).map_err(|e| e.to_string())?; for skill in &skills { - if skill.installed { + if skill.installed && skill.catalog_source != SkillCatalogSource::Project { let key = get_skill_key(&app_type, &skill.directory); if !existing_states.contains_key(&key) { let state = SkillState { @@ -421,7 +605,7 @@ mod tests { // **Feature: skills-platform-mvp, Property 2: Installed Skills Discovery** // **Validates: Requirements 2.1, 2.2, 2.3** // - // *For any* valid ~/.proxycast/skills/ directory containing subdirectories + // *For any* valid skills 目录 containing subdirectories // with SKILL.md files, calling `scan_installed_skills()` SHALL return a list // containing exactly those subdirectory names. proptest! { @@ -512,42 +696,61 @@ mod tests { } #[test] - fn test_read_local_skill_content_success() { + fn test_inspect_local_skill_success() { let temp_dir = TempDir::new().unwrap(); let skills_dir = temp_dir.path().join("skills"); let skill_dir = skills_dir.join("demo-skill"); - std::fs::create_dir_all(&skill_dir).unwrap(); - std::fs::write(skill_dir.join("SKILL.md"), "# Demo Skill\ncontent").unwrap(); + let references_dir = skill_dir.join("references"); + std::fs::create_dir_all(&references_dir).unwrap(); + std::fs::write( + skill_dir.join("SKILL.md"), + r#"--- +name: Demo Skill +description: Inspect me +metadata: + proxycast_workflow_ref: references/workflow.yaml +--- - let content = read_local_skill_content(&skills_dir, "demo-skill").unwrap(); - assert!(content.contains("# Demo Skill")); - assert!(content.contains("content")); +# Demo Skill +content"#, + ) + .unwrap(); + std::fs::write( + references_dir.join("workflow.yaml"), + "- id: draft\n title: 起草\n", + ) + .unwrap(); + + let inspection = inspect_local_skill(&[skills_dir.clone()], "demo-skill").unwrap(); + assert!(inspection.content.contains("# Demo Skill")); + assert!(inspection.resource_summary.has_references); + assert!(inspection.standard_compliance.validation_errors.is_empty()); } #[test] - fn test_read_local_skill_content_rejects_traversal() { + fn test_inspect_local_skill_rejects_traversal() { let temp_dir = TempDir::new().unwrap(); let skills_dir = temp_dir.path().join("skills"); std::fs::create_dir_all(&skills_dir).unwrap(); - let err = read_local_skill_content(&skills_dir, "../outside").unwrap_err(); + let err = inspect_local_skill(&[skills_dir.clone()], "../outside").unwrap_err(); assert!(err.contains("Invalid skill directory")); } #[test] - fn test_read_local_skill_content_missing_skill_md() { + fn test_inspect_local_skill_missing_skill_md() { let temp_dir = TempDir::new().unwrap(); let skills_dir = temp_dir.path().join("skills"); let skill_dir = skills_dir.join("no-skill-md"); std::fs::create_dir_all(&skill_dir).unwrap(); - let err = read_local_skill_content(&skills_dir, "no-skill-md").unwrap_err(); - assert!(err.contains("SKILL.md not found")); + let err = inspect_local_skill(&[skills_dir.clone()], "no-skill-md").unwrap_err(); + assert!(err.contains("Skill not found")); } #[cfg(unix)] #[test] - fn test_read_local_skill_content_rejects_symlink_escape() { + fn test_inspect_local_skill_rejects_symlink_escape() { use std::os::unix::fs::symlink; let temp_dir = TempDir::new().unwrap(); @@ -561,7 +764,79 @@ mod tests { let symlink_dir = skills_dir.join("escape-skill"); symlink(&outside_dir, &symlink_dir).unwrap(); - let err = read_local_skill_content(&skills_dir, "escape-skill").unwrap_err(); + let err = inspect_local_skill(&[skills_dir.clone()], "escape-skill").unwrap_err(); assert!(err.contains("Invalid skill directory path")); } + + #[test] + fn test_inspect_local_skill_prefers_project_root_order() { + let temp_dir = TempDir::new().unwrap(); + let project_skills_dir = temp_dir + .path() + .join("project") + .join(".agents") + .join("skills"); + let user_skills_dir = temp_dir.path().join("user-skills"); + let project_skill_dir = project_skills_dir.join("demo-skill"); + let user_skill_dir = user_skills_dir.join("demo-skill"); + + std::fs::create_dir_all(&project_skill_dir).unwrap(); + std::fs::create_dir_all(&user_skill_dir).unwrap(); + std::fs::write( + project_skill_dir.join("SKILL.md"), + "---\nname: Project Skill\ndescription: project\n---\n", + ) + .unwrap(); + std::fs::write( + user_skill_dir.join("SKILL.md"), + "---\nname: User Skill\ndescription: user\n---\n", + ) + .unwrap(); + + let inspection = inspect_local_skill( + &[project_skills_dir.clone(), user_skills_dir.clone()], + "demo-skill", + ) + .unwrap(); + + assert!(inspection.content.contains("Project Skill")); + assert!(!inspection.content.contains("User Skill")); + } + + #[test] + fn test_create_skill_scaffold_in_root_creates_standard_package() { + let temp_dir = TempDir::new().unwrap(); + let skills_dir = temp_dir.path().join("skills"); + + let inspection = create_skill_scaffold_in_root( + &skills_dir, + "draft-skill", + "Draft Skill", + "Create a new draft", + ) + .unwrap(); + + let skill_md = skills_dir.join("draft-skill").join("SKILL.md"); + assert!(skill_md.is_file()); + assert!(inspection.standard_compliance.is_standard); + assert!(inspection.content.contains("name: Draft Skill")); + assert!(inspection.content.contains("# Draft Skill")); + } + + #[test] + fn test_create_skill_scaffold_in_root_rejects_existing_directory() { + let temp_dir = TempDir::new().unwrap(); + let skills_dir = temp_dir.path().join("skills"); + std::fs::create_dir_all(skills_dir.join("draft-skill")).unwrap(); + + let err = create_skill_scaffold_in_root( + &skills_dir, + "draft-skill", + "Draft Skill", + "Create a new draft", + ) + .unwrap_err(); + + assert!(err.contains("already exists")); + } } diff --git a/src-tauri/src/commands/skill_exec_cmd.rs b/src-tauri/src/commands/skill_exec_cmd.rs index 9f799e9b1..a1f5d2490 100644 --- a/src-tauri/src/commands/skill_exec_cmd.rs +++ b/src-tauri/src/commands/skill_exec_cmd.rs @@ -40,9 +40,13 @@ use crate::database::DbConnection; use crate::services::execution_tracker_service::{ExecutionTracker, RunFinishDecision, RunSource}; use crate::services::memory_profile_prompt_service::build_memory_profile_prompt; use crate::skills::TauriExecutionCallback; -use proxycast_agent::event_converter::{convert_agent_event, TauriToolResult}; +use proxycast_agent::event_converter::{ + convert_agent_event, TauriArtifactSnapshot, TauriToolResult, +}; +use proxycast_agent::WriteArtifactEventEmitter; use proxycast_skills::{ - find_skill_by_name, get_proxycast_skills_dir, load_skills_from_directory, ExecutionCallback, + find_skill_by_name, get_skill_roots, load_skills_from_directory, ExecutionCallback, + LoadedSkillDefinition, }; #[cfg(test)] use proxycast_skills::{ @@ -135,6 +139,18 @@ pub struct SkillExecutionResult { pub steps_completed: Vec, } +fn invalid_skill_message(skill: &LoadedSkillDefinition) -> Option { + if skill.standard_compliance.validation_errors.is_empty() { + return None; + } + + Some(format!( + "Skill '{}' 未通过标准校验: {}", + skill.skill_name, + skill.standard_compliance.validation_errors.join("; ") + )) +} + const SOCIAL_POST_WITH_COVER_SKILL_NAME: &str = "social_post_with_cover"; const SOCIAL_POST_OUTPUT_DIR: &str = "social-posts"; const SOCIAL_POST_WRITE_TOOL_NAME: &str = "write_file"; @@ -484,11 +500,37 @@ fn emit_social_write_file_events( ) { let event_name = format!("skill-exec-{execution_id}"); let tool_id = build_social_tool_event_id(execution_id, file_path); + let artifact_id = format!("{tool_id}:artifact"); let arguments = serde_json::json!({ "path": file_path, "content": file_content, }) .to_string(); + let preview_text = file_content.trim().chars().take(480).collect::(); + let latest_chunk = file_content + .trim() + .chars() + .rev() + .take(240) + .collect::>() + .into_iter() + .rev() + .collect::(); + let mut artifact_metadata = std::collections::HashMap::from([ + ("complete".to_string(), serde_json::json!(true)), + ("writePhase".to_string(), serde_json::json!("persisted")), + ("isPartial".to_string(), serde_json::json!(false)), + ( + "lastUpdateSource".to_string(), + serde_json::json!("tool_result"), + ), + ]); + if !preview_text.is_empty() { + artifact_metadata.insert("previewText".to_string(), serde_json::json!(preview_text)); + } + if !latest_chunk.is_empty() { + artifact_metadata.insert("latestChunk".to_string(), serde_json::json!(latest_chunk)); + } let tool_start = TauriAgentEvent::ToolStart { tool_name: SOCIAL_POST_WRITE_TOOL_NAME.to_string(), @@ -499,6 +541,24 @@ fn emit_social_write_file_events( tracing::warn!("[execute_skill] 发送社媒写入工具开始事件失败: {}", err); } + let artifact_snapshot = TauriAgentEvent::ArtifactSnapshot { + artifact: TauriArtifactSnapshot { + artifact_id: artifact_id.clone(), + file_path: file_path.to_string(), + content: Some(file_content.to_string()), + metadata: Some(artifact_metadata.clone()), + }, + }; + if let Err(err) = app_handle.emit(&event_name, &artifact_snapshot) { + tracing::warn!("[execute_skill] 发送社媒产物快照事件失败: {}", err); + } + + let mut tool_end_metadata = artifact_metadata; + tool_end_metadata.insert("artifact_streamed".to_string(), serde_json::json!(true)); + tool_end_metadata.insert("artifact_id".to_string(), serde_json::json!(artifact_id)); + tool_end_metadata.insert("artifact_path".to_string(), serde_json::json!(file_path)); + tool_end_metadata.insert("path".to_string(), serde_json::json!(file_path)); + tool_end_metadata.insert("file_path".to_string(), serde_json::json!(file_path)); let tool_end = TauriAgentEvent::ToolEnd { tool_id, result: TauriToolResult { @@ -506,7 +566,7 @@ fn emit_social_write_file_events( output: format!("写入社媒文稿: {file_path}"), error: None, images: None, - metadata: None, + metadata: Some(tool_end_metadata), }, }; if let Err(err) = app_handle.emit(&event_name, &tool_end) { @@ -584,6 +644,10 @@ pub async fn execute_skill( // 1. 从 registry 加载 skill(Requirements 3.2) let skill = find_skill_by_name(&skill_name).map_err(map_find_skill_error)?; + if let Some(message) = invalid_skill_message(&skill) { + return Err(format_skill_error(SKILL_ERR_EXECUTE_FAILED, message)); + } + // 检查是否禁用了模型调用 if skill.disable_model_invocation { return Err(format_skill_error( @@ -851,6 +915,7 @@ async fn execute_skill_prompt( let mut has_error = false; let mut error_message: Option = None; let event_name = format!("skill-exec-{execution_id}"); + let mut write_artifact_emitter = WriteArtifactEventEmitter::new(session_id.to_string()); match stream_result { Ok(mut stream) => { @@ -858,7 +923,14 @@ async fn execute_skill_prompt( match event_result { Ok(agent_event) => { let tauri_events = convert_agent_event(agent_event); - for tauri_event in tauri_events { + for mut tauri_event in tauri_events { + let extra_events = + write_artifact_emitter.process_event(&mut tauri_event); + for extra_event in &extra_events { + if let Err(e) = app_handle.emit(&event_name, extra_event) { + tracing::error!("[execute_skill] 发送补充事件失败: {}", e); + } + } if let TauriAgentEvent::TextDelta { ref text } = tauri_event { final_output.push_str(text); } @@ -1042,6 +1114,7 @@ async fn execute_skill_workflow( let mut step_output = String::new(); let mut step_error: Option = None; + let mut write_artifact_emitter = WriteArtifactEventEmitter::new(step_session_id.clone()); match stream_result { Ok(mut stream) => { @@ -1049,7 +1122,17 @@ async fn execute_skill_workflow( match event_result { Ok(agent_event) => { let tauri_events = convert_agent_event(agent_event); - for tauri_event in tauri_events { + for mut tauri_event in tauri_events { + let extra_events = + write_artifact_emitter.process_event(&mut tauri_event); + for extra_event in &extra_events { + if let Err(e) = app_handle.emit(&event_name, extra_event) { + tracing::error!( + "[execute_skill_workflow] 发送补充事件失败: {}", + e + ); + } + } if let TauriAgentEvent::TextDelta { ref text } = tauri_event { step_output.push_str(text); } @@ -1147,7 +1230,8 @@ async fn execute_skill_workflow( /// 列出可执行的 Skills /// -/// 返回所有可以执行的 Skills 列表,过滤掉 disable_model_invocation=true 的 Skills。 +/// 返回所有可以执行的 Skills 列表,过滤掉无效 Skill 包和 +/// disable_model_invocation=true 的 Skills。 /// /// # Returns /// * `Ok(Vec)` - 可执行的 Skills 列表 @@ -1158,13 +1242,26 @@ async fn execute_skill_workflow( /// - 4.2: 包含 name, description, execution_mode /// - 4.3: 指示是否有 workflow 定义 /// - 4.4: 过滤 disable_model_invocation=true 的 skills +/// - 4.5: 过滤未通过标准校验的 skills #[tauri::command] pub async fn list_executable_skills() -> Result, String> { - let skills_dir = get_proxycast_skills_dir() - .ok_or_else(|| format_skill_error(SKILL_ERR_CATALOG_UNAVAILABLE, "无法获取 Skills 目录"))?; + let skill_roots = get_skill_roots(); + if skill_roots.is_empty() { + return Err(format_skill_error( + SKILL_ERR_CATALOG_UNAVAILABLE, + "无法获取 Skills 目录", + )); + } - // 加载所有 skills - let all_skills = load_skills_from_directory(&skills_dir); + let mut all_skills = Vec::new(); + let mut seen = std::collections::HashSet::new(); + for skill_root in skill_roots { + for skill in load_skills_from_directory(&skill_root) { + if seen.insert(skill.skill_name.clone()) { + all_skills.push(skill); + } + } + } // 过滤掉 disable_model_invocation=true 的 skills(Requirements 4.4) let executable_skills: Vec = all_skills @@ -1210,6 +1307,9 @@ pub async fn list_executable_skills() -> Result, String pub async fn get_skill_detail(skill_name: String) -> Result { // 查找 skill(Requirements 5.1, 5.4) let skill = find_skill_by_name(&skill_name).map_err(map_find_skill_error)?; + if let Some(message) = invalid_skill_message(&skill) { + return Err(format_skill_error(SKILL_ERR_EXECUTE_FAILED, message)); + } // 转换为 SkillDetailInfo(Requirements 5.2, 5.3) let detail = SkillDetailInfo { @@ -1333,8 +1433,9 @@ mod tests { let content = r#"--- name: test-skill description: A test skill -model: claude-sonnet-4-5-20250514 -provider: claude +metadata: + proxycast_model_preference: claude-sonnet-4-5-20250514 + proxycast_provider_preference: claude --- # Test Skill @@ -1514,8 +1615,9 @@ Body name: my-skill description: Test skill description allowed-tools: tool1, tool2 -model: gpt-4 -provider: openai +metadata: + proxycast_model_preference: gpt-4 + proxycast_provider_preference: openai --- # My Skill @@ -1538,6 +1640,41 @@ Instructions here. assert_eq!(skill.provider, Some("openai".to_string())); assert!(!skill.disable_model_invocation); assert_eq!(skill.execution_mode, "prompt"); + assert!(skill.standard_compliance.is_standard); + } + + #[test] + fn test_load_skill_from_file_should_surface_invalid_workflow_reference() { + use tempfile::TempDir; + + let temp_dir = TempDir::new().unwrap(); + let skill_dir = temp_dir.path().join("workflow-skill"); + std::fs::create_dir(&skill_dir).unwrap(); + + let skill_file = skill_dir.join("SKILL.md"); + std::fs::write( + &skill_file, + r#"--- +name: workflow-skill +description: Workflow skill +metadata: + proxycast_workflow_ref: references/missing.json +--- + +# Workflow Skill +"#, + ) + .unwrap(); + + let skill = load_skill_from_file("workflow-skill", &skill_file).unwrap(); + + assert!(!skill.standard_compliance.is_standard); + assert!(skill + .standard_compliance + .validation_errors + .iter() + .any(|error| error.contains("metadata.proxycast_workflow_ref"))); + assert!(skill.workflow_steps.is_empty()); } #[test] @@ -1588,6 +1725,47 @@ Content 2 assert!(skill_two.disable_model_invocation); } + #[test] + fn test_load_skills_from_directory_should_skip_invalid_skill_packages() { + use tempfile::TempDir; + + let temp_dir = TempDir::new().unwrap(); + let skills_dir = temp_dir.path(); + + let valid_dir = skills_dir.join("skill-valid"); + std::fs::create_dir(&valid_dir).unwrap(); + std::fs::write( + valid_dir.join("SKILL.md"), + r#"--- +name: skill-valid +description: Valid skill +--- +Valid content +"#, + ) + .unwrap(); + + let invalid_dir = skills_dir.join("skill-invalid"); + std::fs::create_dir(&invalid_dir).unwrap(); + std::fs::write( + invalid_dir.join("SKILL.md"), + r#"--- +name: skill-invalid +description: Invalid skill +metadata: + proxycast_workflow_ref: references/missing.json +--- +Invalid content +"#, + ) + .unwrap(); + + let skills = load_skills_from_directory(skills_dir); + + assert_eq!(skills.len(), 1); + assert_eq!(skills[0].skill_name, "skill-valid"); + } + #[test] fn test_load_skills_from_nonexistent_directory() { let skills = load_skills_from_directory(std::path::Path::new("/nonexistent/path")); @@ -1605,6 +1783,10 @@ Content 2 assert_eq!(skill.skill_name, "social_post_with_cover"); assert_eq!(skill.execution_mode, "workflow"); + assert_eq!( + skill.workflow_ref, + Some("references/workflow.json".to_string()) + ); assert_eq!( skill.allowed_tools, Some(vec![ @@ -1614,5 +1796,6 @@ Content 2 ); assert!(content.contains(" Result ( RpcMethod::AgentRun, - Some(json!({ "message": message, "stream": false })), + Some(json!({ + "message": message, + "stream": false, + "web_search": true, + "search_mode": "allowed" + })), ), TelegramCommand::Status(run_id) => ( RpcMethod::AgentWait, diff --git a/src-tauri/src/commands/theme_context_cmd.rs b/src-tauri/src/commands/theme_context_cmd.rs index 70a1ca377..dbfe5931e 100644 --- a/src-tauri/src/commands/theme_context_cmd.rs +++ b/src-tauri/src/commands/theme_context_cmd.rs @@ -14,7 +14,8 @@ use crate::services::web_search_runtime_service::apply_web_search_runtime_env; use crate::services::workspace_health_service::ensure_workspace_ready_with_auto_relocate; use crate::workspace::WorkspaceManager; use proxycast_agent::{ - resolve_request_tool_policy, stream_reply_with_policy, SessionConfigBuilder, + resolve_request_tool_policy_with_mode, stream_reply_with_policy, RequestToolPolicyMode, + SessionConfigBuilder, }; use serde::{Deserialize, Serialize}; use std::path::Path; @@ -380,7 +381,11 @@ pub async fn aster_agent_theme_context_search( } }); - let request_tool_policy = resolve_request_tool_policy(Some(true), false); + let request_tool_policy = resolve_request_tool_policy_with_mode( + Some(true), + Some(RequestToolPolicyMode::Required), + false, + ); let working_dir = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from(".")); let system_prompt = proxycast_agent::merge_system_prompt_with_request_tool_policy( merge_system_prompt_with_web_search( diff --git a/src-tauri/src/commands/unified_chat_cmd.rs b/src-tauri/src/commands/unified_chat_cmd.rs index 652baf395..46fb28907 100644 --- a/src-tauri/src/commands/unified_chat_cmd.rs +++ b/src-tauri/src/commands/unified_chat_cmd.rs @@ -22,16 +22,17 @@ use crate::database::DbConnection; use crate::services::memory_profile_prompt_service::{ merge_system_prompt_with_memory_profile, merge_system_prompt_with_memory_sources, }; -use crate::services::request_tool_policy_prompt_service::{ - execute_web_search_preflight_if_needed, merge_system_prompt_with_request_tool_policy, - resolve_request_tool_policy, RequestToolPolicy, WebSearchExecutionTracker, -}; use crate::services::web_search_prompt_service::merge_system_prompt_with_web_search; use crate::services::web_search_runtime_service::apply_web_search_runtime_env; use aster::agents::extension::ExtensionConfig; use aster::conversation::message::Message; use futures::StreamExt; -use proxycast_agent::event_converter::convert_agent_event; +use proxycast_agent::{ + convert_agent_event, execute_web_search_preflight_if_needed, + merge_system_prompt_with_request_tool_policy, + merge_system_prompt_with_web_search_preflight_context, resolve_request_tool_policy_with_mode, + RequestToolPolicy, RequestToolPolicyMode, WebSearchExecutionTracker, WriteArtifactEventEmitter, +}; use serde::{Deserialize, Serialize}; use tauri::{AppHandle, Emitter, State}; @@ -75,6 +76,9 @@ pub struct SendMessageRequest { /// 请求级联网搜索开关 #[serde(default, alias = "webSearch")] pub web_search: Option, + /// 联网搜索模式(disabled / allowed / required) + #[serde(default, alias = "searchMode")] + pub search_mode: Option, } /// 图片输入 @@ -384,16 +388,21 @@ pub async fn chat_send_message( &config, ); - let mode_default_web_search = matches!(session.mode, ChatMode::General); - let request_tool_policy = - resolve_request_tool_policy(request.web_search, mode_default_web_search); + let mode_default_web_search = false; + let request_tool_policy = resolve_request_tool_policy_with_mode( + request.web_search, + request.search_mode, + mode_default_web_search, + ); tracing::info!( - "[UnifiedChat][WebSearchGuard] session={}, mode={:?}, request_web_search={:?}, mode_default_web_search={}, effective_web_search={}", + "[UnifiedChat][WebSearchGuard] session={}, mode={:?}, request_web_search={:?}, request_search_mode={:?}, mode_default_web_search={}, effective_web_search={}, search_mode={}", request.session_id, session.mode, request.web_search, + request.search_mode, mode_default_web_search, - request_tool_policy.effective_web_search + request_tool_policy.effective_web_search, + request_tool_policy.search_mode.as_str() ); let result = send_message_with_aster( @@ -471,7 +480,7 @@ async fn send_message_with_aster( if let Some(prompt) = effective_system_prompt { session_config_builder = session_config_builder.system_prompt(prompt); } - let session_config = session_config_builder + let mut session_config = session_config_builder .include_context_trace(include_context_trace) .build(); @@ -481,7 +490,7 @@ async fn send_message_with_aster( let agent = guard.as_ref().ok_or("Agent 未初始化")?; let mut removed_extension: Option = None; - if request_tool_policy.effective_web_search { + if request_tool_policy.requires_web_search() { let extension_configs = agent.get_extension_configs().await; if let Some(extension) = extension_configs .into_iter() @@ -527,6 +536,18 @@ async fn send_message_with_aster( .await; match preflight { Ok(preflight_execution) => { + session_config.system_prompt = merge_system_prompt_with_web_search_preflight_context( + session_config.system_prompt.take(), + preflight_execution.system_prompt_appendix.clone(), + ); + if let Some(summary) = preflight_execution.coverage_summary.as_deref() { + tracing::info!( + "[UnifiedChat][WebSearchPrefetch] session={}, expanded_news_search={}, summary={}", + session_id, + preflight_execution.expanded_news_search, + summary + ); + } for event in preflight_execution.events { if let Err(error) = app.emit(event_name, &event) { tracing::error!("[UnifiedChat] 发送预调用事件失败: {}", error); @@ -562,6 +583,8 @@ async fn send_message_with_aster( let mut first_chunk_time: Option = None; let mut chunk_count = 0; let mut stream_error: Option = None; + let mut text_output = String::new(); + let mut write_artifact_emitter = WriteArtifactEventEmitter::new(session_id); match stream_result { Ok(mut stream) => { @@ -577,8 +600,20 @@ async fn send_message_with_aster( chunk_count += 1; let tauri_events = convert_agent_event(agent_event); - for tauri_event in tauri_events { + for mut tauri_event in tauri_events { + let extra_events = + write_artifact_emitter.process_event(&mut tauri_event); + for extra_event in &extra_events { + if let Err(e) = app.emit(event_name, extra_event) { + tracing::error!("[UnifiedChat] 发送补充事件失败: {}", e); + } + } match &tauri_event { + TauriAgentEvent::TextDelta { text } => { + if !text.is_empty() { + text_output.push_str(text); + } + } TauriAgentEvent::ToolStart { tool_name, tool_id, .. } => web_search_tracker.record_tool_start( @@ -624,6 +659,18 @@ async fn send_message_with_aster( } } + if stream_error.is_none() && text_output.trim().is_empty() { + let message = format!( + "已完成当前回合的工具执行,但模型未输出最终答复。\n尝试记录: {}", + web_search_tracker.format_attempts() + ); + let error_event = TauriAgentEvent::Error { + message: message.clone(), + }; + let _ = app.emit(event_name, &error_event); + stream_error = Some(message); + } + if stream_error.is_none() { // 发送完成事件 let done_event = TauriAgentEvent::FinalDone { usage: None }; @@ -708,7 +755,7 @@ pub async fn chat_configure_provider( #[cfg(test)] mod tests { use super::*; - use crate::services::request_tool_policy_prompt_service::resolve_request_tool_policy; + use proxycast_agent::resolve_request_tool_policy; #[test] fn test_send_message_request_deserialize_web_search_camel_case() { diff --git a/src-tauri/src/dev_bridge/dispatcher.rs b/src-tauri/src/dev_bridge/dispatcher.rs index a2a716149..880224a63 100644 --- a/src-tauri/src/dev_bridge/dispatcher.rs +++ b/src-tauri/src/dev_bridge/dispatcher.rs @@ -570,6 +570,55 @@ pub async fn handle_command( } } + "inspect_local_skill_for_app" => { + let args = args.unwrap_or_default(); + let app = args + .get("app") + .and_then(|value| value.as_str()) + .unwrap_or("proxycast") + .to_string(); + let directory = get_string_arg(&args, "directory", "directory")?; + let inspection = crate::commands::skill_cmd::inspect_local_skill_for_app(app, directory) + .map_err(|e| format!("检查本地 Skill 失败: {e}"))?; + Ok(serde_json::to_value(inspection)?) + } + + "create_skill_scaffold_for_app" => { + let args = args.unwrap_or_default(); + let app = args + .get("app") + .and_then(|value| value.as_str()) + .unwrap_or("proxycast") + .to_string(); + let target = get_string_arg(&args, "target", "target")?; + let directory = get_string_arg(&args, "directory", "directory")?; + let name = get_string_arg(&args, "name", "name")?; + let description = get_string_arg(&args, "description", "description")?; + let inspection = crate::commands::skill_cmd::create_skill_scaffold_for_app( + app, + target, + directory, + name, + description, + ) + .map_err(|e| format!("创建 Skill 脚手架失败: {e}"))?; + Ok(serde_json::to_value(inspection)?) + } + + "inspect_remote_skill" => { + let args = args.unwrap_or_default(); + let owner = get_string_arg(&args, "owner", "owner")?; + let name = get_string_arg(&args, "name", "name")?; + let branch = get_string_arg(&args, "branch", "branch")?; + let directory = get_string_arg(&args, "directory", "directory")?; + let inspection = state + .skill_service + .inspect_remote_skill(&owner, &name, &branch, &directory) + .await + .map_err(|e| format!("检查远程 Skill 失败: {e}"))?; + Ok(serde_json::to_value(inspection)?) + } + "test_api" => { // 测试 API 连接 // 从 args 获取 provider diff --git a/src-tauri/src/services/agent_timeline_service.rs b/src-tauri/src/services/agent_timeline_service.rs index 492d9b466..6edad00b4 100644 --- a/src-tauri/src/services/agent_timeline_service.rs +++ b/src-tauri/src/services/agent_timeline_service.rs @@ -13,6 +13,25 @@ use uuid::Uuid; const PROPOSED_PLAN_OPEN: &str = ""; const PROPOSED_PLAN_CLOSE: &str = ""; +fn format_runtime_status_text(title: &str, detail: &str, checkpoints: &[String]) -> String { + let mut lines = Vec::new(); + let trimmed_title = title.trim(); + if !trimmed_title.is_empty() { + lines.push(trimmed_title.to_string()); + } + let trimmed_detail = detail.trim(); + if !trimmed_detail.is_empty() { + lines.push(trimmed_detail.to_string()); + } + for checkpoint in checkpoints { + let trimmed = checkpoint.trim(); + if !trimmed.is_empty() { + lines.push(format!("• {trimmed}")); + } + } + lines.join("\n") +} + fn emit_event(app: &AppHandle, event_name: &str, event: &TauriAgentEvent) { if let Err(error) = app.emit(event_name, event) { tracing::error!("[AgentTimeline] 发送事件失败: {}", error); @@ -65,7 +84,45 @@ fn extract_command_text(arguments: Option<&Value>) -> Option { ) } -fn extract_file_paths(arguments: Option<&Value>, metadata: Option<&Value>) -> Vec { +#[derive(Debug, Clone)] +struct ExtractedFileArtifact { + path: String, + artifact_id: Option, +} + +fn push_unique_file_path(target: &mut Vec, raw: &str) { + let trimmed = raw.trim(); + if trimmed.is_empty() || target.iter().any(|item| item == trimmed) { + return; + } + target.push(trimmed.to_string()); +} + +fn collect_string_values(value: &Value) -> Vec { + match value { + Value::String(text) => { + let trimmed = text.trim(); + if trimmed.is_empty() { + Vec::new() + } else { + vec![trimmed.to_string()] + } + } + Value::Array(items) => items + .iter() + .filter_map(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .collect(), + _ => Vec::new(), + } +} + +fn extract_file_artifacts( + arguments: Option<&Value>, + metadata: Option<&Value>, +) -> Vec { let mut paths = Vec::new(); for source in [arguments, metadata] { let Some(object) = source.and_then(as_object) else { @@ -86,28 +143,69 @@ fn extract_file_paths(arguments: Option<&Value>, metadata: Option<&Value>) -> Ve let Some(value) = object.get(key) else { continue; }; - match value { - Value::String(text) => { - let trimmed = text.trim(); - if !trimmed.is_empty() && !paths.iter().any(|item| item == trimmed) { - paths.push(trimmed.to_string()); - } - } - Value::Array(items) => { - for item in items { - if let Some(text) = item.as_str() { - let trimmed = text.trim(); - if !trimmed.is_empty() && !paths.iter().any(|entry| entry == trimmed) { - paths.push(trimmed.to_string()); - } - } - } - } - _ => {} + for path in collect_string_values(value) { + push_unique_file_path(&mut paths, path.as_str()); } } } + + let metadata_object = metadata.and_then(as_object); + let artifact_ids = metadata_object + .and_then(|object| object.get("artifact_ids")) + .map(collect_string_values) + .unwrap_or_default(); + let single_artifact_id = metadata_object + .and_then(|object| { + object + .get("artifact_id") + .or_else(|| object.get("artifactId")) + }) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string); + paths + .into_iter() + .enumerate() + .map(|(index, path)| ExtractedFileArtifact { + path, + artifact_id: artifact_ids.get(index).cloned().or_else(|| { + if index == 0 { + single_artifact_id.clone() + } else { + None + } + }), + }) + .collect() +} + +fn resolve_artifact_item_status(metadata: Option<&Value>) -> AgentThreadItemStatus { + let write_phase = metadata + .and_then(|value| value.get("writePhase")) + .and_then(Value::as_str); + if matches!(write_phase, Some("failed")) { + return AgentThreadItemStatus::Failed; + } + + match metadata + .and_then(|value| value.get("complete")) + .and_then(Value::as_bool) + { + Some(false) => AgentThreadItemStatus::InProgress, + _ => AgentThreadItemStatus::Completed, + } +} + +fn resolve_artifact_item_source(metadata: Option<&Value>) -> String { + metadata + .and_then(|value| value.get("lastUpdateSource")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .unwrap_or_else(|| "artifact_snapshot".to_string()) } fn extract_proposed_plan_block(text: &str) -> Option { @@ -212,6 +310,7 @@ pub struct AgentTimelineRecorder { assistant_text: String, reasoning_text: String, plan_text: Option, + turn_summary_text: Option, } impl AgentTimelineRecorder { @@ -252,6 +351,7 @@ impl AgentTimelineRecorder { assistant_text: String::new(), reasoning_text: String::new(), plan_text: None, + turn_summary_text: None, }) } @@ -338,6 +438,20 @@ impl AgentTimelineRecorder { ); self.persist_and_emit_item(app, event_name, item)?; } + TauriAgentEvent::RuntimeStatus { status } => { + let text = + format_runtime_status_text(&status.title, &status.detail, &status.checkpoints); + if !text.is_empty() { + self.turn_summary_text = Some(text.clone()); + let item = self.build_item( + format!("turn_summary:{}", self.turn_id), + AgentThreadItemStatus::InProgress, + None, + AgentThreadItemPayload::TurnSummary { text }, + ); + self.persist_and_emit_item(app, event_name, item)?; + } + } TauriAgentEvent::ToolStart { tool_name, tool_id, @@ -439,19 +553,27 @@ impl AgentTimelineRecorder { let item = self.build_item( tool_id.clone(), - status, + status.clone(), Some(Utc::now().to_rfc3339()), payload, ); self.persist_and_emit_item(app, event_name, item)?; - for path in extract_file_paths(None, metadata_value.as_ref()) { + for artifact in extract_file_artifacts(None, metadata_value.as_ref()) { + let artifact_path = artifact.path.clone(); let file_item = self.build_item( - format!("artifact:{}:{}", tool_id, path), - AgentThreadItemStatus::Completed, - Some(Utc::now().to_rfc3339()), + artifact + .artifact_id + .clone() + .unwrap_or_else(|| format!("artifact:{}:{}", tool_id, artifact_path)), + status.clone(), + if matches!(status, AgentThreadItemStatus::InProgress) { + None + } else { + Some(Utc::now().to_rfc3339()) + }, AgentThreadItemPayload::FileArtifact { - path, + path: artifact_path, source: "tool_result".to_string(), content: None, metadata: metadata_value.clone(), @@ -460,6 +582,29 @@ impl AgentTimelineRecorder { self.persist_and_emit_item(app, event_name, file_item)?; } } + TauriAgentEvent::ArtifactSnapshot { artifact } => { + let metadata_value = artifact + .metadata + .as_ref() + .and_then(|metadata| serde_json::to_value(metadata).ok()); + let status = resolve_artifact_item_status(metadata_value.as_ref()); + let item = self.build_item( + artifact.artifact_id.clone(), + status.clone(), + if matches!(status, AgentThreadItemStatus::InProgress) { + None + } else { + Some(Utc::now().to_rfc3339()) + }, + AgentThreadItemPayload::FileArtifact { + path: artifact.file_path.clone(), + source: resolve_artifact_item_source(metadata_value.as_ref()), + content: artifact.content.clone(), + metadata: metadata_value, + }, + ); + self.persist_and_emit_item(app, event_name, item)?; + } TauriAgentEvent::ActionRequired { request_id, action_type, @@ -650,13 +795,25 @@ impl AgentTimelineRecorder { if let Some(plan_text) = self.plan_text.clone() { let item = self.build_item( format!("plan:{}", self.turn_id), - status, + status.clone(), Some(Utc::now().to_rfc3339()), AgentThreadItemPayload::Plan { text: plan_text }, ); self.persist_and_emit_item(app, event_name, item)?; } + if let Some(turn_summary_text) = self.turn_summary_text.clone() { + let item = self.build_item( + format!("turn_summary:{}", self.turn_id), + status, + Some(Utc::now().to_rfc3339()), + AgentThreadItemPayload::TurnSummary { + text: turn_summary_text, + }, + ); + self.persist_and_emit_item(app, event_name, item)?; + } + Ok(()) } diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index 3247213f2..d9f3d0dbb 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -18,7 +18,6 @@ pub mod memory_rules_loader_service; pub mod memory_source_resolver_service; pub mod novel_service; pub mod openclaw_service; -pub mod request_tool_policy_prompt_service; pub mod sysinfo_service; pub mod update_check_service; pub mod update_window; diff --git a/src-tauri/src/services/request_tool_policy_prompt_service.rs b/src-tauri/src/services/request_tool_policy_prompt_service.rs deleted file mode 100644 index c8a50e197..000000000 --- a/src-tauri/src/services/request_tool_policy_prompt_service.rs +++ /dev/null @@ -1,10 +0,0 @@ -//! 请求级工具策略服务(兼容层) -//! -//! 真实实现已迁移到 `proxycast-agent::request_tool_policy`, -//! 此处保留旧路径仅用于兼容主 crate 既有调用与测试引用。 - -pub use proxycast_agent::request_tool_policy::{ - execute_web_search_preflight_if_needed, merge_system_prompt_with_request_tool_policy, - resolve_request_tool_policy, stream_reply_with_policy, ReplyAttemptError, RequestToolPolicy, - StreamReplyExecution, WebSearchExecutionTracker, REQUEST_TOOL_POLICY_MARKER, -}; diff --git a/src-tauri/src/skills/README.md b/src-tauri/src/skills/README.md index bbdf70d4b..7c982b1f1 100644 --- a/src-tauri/src/skills/README.md +++ b/src-tauri/src/skills/README.md @@ -12,13 +12,14 @@ ## Skills 集成架构 -### AI 自动调用 Skills(方案 A) +### AI 自动调用 Skills(标准化后) ProxyCast 通过以下机制让 AI 能够自动发现和调用 Skills: 1. **Agent 初始化时加载 Skills** - `AsterAgentState::init_agent_with_db()` 调用 `load_proxycast_skills()` - - 从 `~/.proxycast/skills/` 目录加载所有 Skills + - 技能包以 Agent Skills 标准 `SKILL.md` 为主格式 + - 默认从应用级 Skills 目录加载,并支持项目级 `./.agents/skills` - 注册到 aster-rust 的 `global_registry` 2. **SkillTool 自动注册** @@ -95,3 +96,13 @@ commands/skill_cmd.rs - 设计文档: `.kiro/specs/skills-integration/design.md` - 需求文档: `.kiro/specs/skills-integration/requirements.md` +- 路线图: `docs/roadmap/proxycast-skills-standardization-roadmap.md` + +## 当前标准约定 + +- Agent Skills 是唯一标准格式 +- ProxyCast 私有能力统一写入 `metadata.proxycast_*` +- Workflow 不再推荐使用 `steps-json` 内联,优先通过 `metadata.proxycast_workflow_ref` 指向 `references/` 下文件 +- 服务层和执行层共用 `SkillService::inspect_*` inspection 结果作为标准合规事实源,并向前端暴露标准合规状态与资源摘要 +- 无效 Skill 仍可在管理页中看到检查结果,但不会进入运行时自动加载和可执行列表 +- 管理链路支持创建最小标准 Skill 脚手架,新建结果会立即经过统一 inspection 校验 diff --git a/src-tauri/src/skills/default_skills.rs b/src-tauri/src/skills/default_skills.rs index f6e655e94..04a3306df 100644 --- a/src-tauri/src/skills/default_skills.rs +++ b/src-tauri/src/skills/default_skills.rs @@ -5,6 +5,7 @@ use std::path::Path; use std::path::PathBuf; use proxycast_core::app_paths; +use proxycast_core::models::parse_skill_manifest_from_content; use proxycast_core::models::{ BROADCAST_GENERATE_SKILL_DIRECTORY, COVER_GENERATE_SKILL_DIRECTORY, IMAGE_GENERATE_SKILL_DIRECTORY, LIBRARY_SKILL_DIRECTORY, MODAL_RESOURCE_SEARCH_SKILL_DIRECTORY, @@ -41,27 +42,79 @@ const TYPESETTING_SKILL_CONTENT: &str = const SOCIAL_POST_WITH_COVER_SKILL_CONTENT: &str = include_str!("../../resources/default-skills/social_post_with_cover/SKILL.md"); -fn default_skills() -> [(&'static str, &'static str); 10] { +const SOCIAL_POST_WITH_COVER_WORKFLOW_CONTENT: &str = + include_str!("../../resources/default-skills/social_post_with_cover/references/workflow.json"); + +#[derive(Clone, Copy)] +struct BundledSkillFile { + relative_path: &'static str, + content: &'static str, +} + +#[derive(Clone, Copy)] +struct BundledSkillDefinition { + directory: &'static str, + skill_content: &'static str, + extra_files: &'static [BundledSkillFile], +} + +const SOCIAL_POST_WITH_COVER_EXTRA_FILES: &[BundledSkillFile] = &[BundledSkillFile { + relative_path: "references/workflow.json", + content: SOCIAL_POST_WITH_COVER_WORKFLOW_CONTENT, +}]; + +fn default_skills() -> [BundledSkillDefinition; 10] { [ - (VIDEO_GENERATE_SKILL_DIRECTORY, VIDEO_GENERATE_SKILL_CONTENT), - ( - BROADCAST_GENERATE_SKILL_DIRECTORY, - BROADCAST_GENERATE_SKILL_CONTENT, - ), - (COVER_GENERATE_SKILL_DIRECTORY, COVER_GENERATE_SKILL_CONTENT), - ( - MODAL_RESOURCE_SEARCH_SKILL_DIRECTORY, - MODAL_RESOURCE_SEARCH_SKILL_CONTENT, - ), - (IMAGE_GENERATE_SKILL_DIRECTORY, IMAGE_GENERATE_SKILL_CONTENT), - (LIBRARY_SKILL_DIRECTORY, LIBRARY_SKILL_CONTENT), - (URL_PARSE_SKILL_DIRECTORY, URL_PARSE_SKILL_CONTENT), - (RESEARCH_SKILL_DIRECTORY, RESEARCH_SKILL_CONTENT), - (TYPESETTING_SKILL_DIRECTORY, TYPESETTING_SKILL_CONTENT), - ( - SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY, - SOCIAL_POST_WITH_COVER_SKILL_CONTENT, - ), + BundledSkillDefinition { + directory: VIDEO_GENERATE_SKILL_DIRECTORY, + skill_content: VIDEO_GENERATE_SKILL_CONTENT, + extra_files: &[], + }, + BundledSkillDefinition { + directory: BROADCAST_GENERATE_SKILL_DIRECTORY, + skill_content: BROADCAST_GENERATE_SKILL_CONTENT, + extra_files: &[], + }, + BundledSkillDefinition { + directory: COVER_GENERATE_SKILL_DIRECTORY, + skill_content: COVER_GENERATE_SKILL_CONTENT, + extra_files: &[], + }, + BundledSkillDefinition { + directory: MODAL_RESOURCE_SEARCH_SKILL_DIRECTORY, + skill_content: MODAL_RESOURCE_SEARCH_SKILL_CONTENT, + extra_files: &[], + }, + BundledSkillDefinition { + directory: IMAGE_GENERATE_SKILL_DIRECTORY, + skill_content: IMAGE_GENERATE_SKILL_CONTENT, + extra_files: &[], + }, + BundledSkillDefinition { + directory: LIBRARY_SKILL_DIRECTORY, + skill_content: LIBRARY_SKILL_CONTENT, + extra_files: &[], + }, + BundledSkillDefinition { + directory: URL_PARSE_SKILL_DIRECTORY, + skill_content: URL_PARSE_SKILL_CONTENT, + extra_files: &[], + }, + BundledSkillDefinition { + directory: RESEARCH_SKILL_DIRECTORY, + skill_content: RESEARCH_SKILL_CONTENT, + extra_files: &[], + }, + BundledSkillDefinition { + directory: TYPESETTING_SKILL_DIRECTORY, + skill_content: TYPESETTING_SKILL_CONTENT, + extra_files: &[], + }, + BundledSkillDefinition { + directory: SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY, + skill_content: SOCIAL_POST_WITH_COVER_SKILL_CONTENT, + extra_files: SOCIAL_POST_WITH_COVER_EXTRA_FILES, + }, ] } @@ -72,18 +125,19 @@ fn skills_root_from_base(base_dir: &Path) -> PathBuf { /// 从 SKILL.md 内容中提取版本号,返回 (major, minor, patch) fn parse_skill_version(content: &str) -> Option<(u32, u32, u32)> { - for line in content.lines() { - let trimmed = line.trim(); - if trimmed.starts_with("version:") { - let version_str = trimmed.split_once(':')?.1.trim(); - let parts: Vec<&str> = version_str.split('.').collect(); - if parts.len() == 3 { - let major = parts[0].trim().parse::().ok()?; - let minor = parts[1].trim().parse::().ok()?; - let patch = parts[2].trim().parse::().ok()?; - return Some((major, minor, patch)); - } - } + let manifest = parse_skill_manifest_from_content(content).ok()?; + let version_str = manifest + .metadata + .metadata + .get("proxycast_version") + .cloned() + .or_else(|| manifest.raw_string("version"))?; + let parts: Vec<&str> = version_str.split('.').collect(); + if parts.len() == 3 { + let major = parts[0].trim().parse::().ok()?; + let minor = parts[1].trim().parse::().ok()?; + let patch = parts[2].trim().parse::().ok()?; + return Some((major, minor, patch)); } None } @@ -93,7 +147,9 @@ fn ensure_default_local_skills_in_dir(skills_root: &Path) -> Result, .map_err(|e| format!("创建技能目录失败 {}: {e}", skills_root.display()))?; let mut installed = Vec::new(); - for (skill_name, skill_content) in default_skills() { + for bundled_skill in default_skills() { + let skill_name = bundled_skill.directory; + let skill_content = bundled_skill.skill_content; let skill_dir = skills_root.join(skill_name); let skill_md_path = skill_dir.join("SKILL.md"); if skill_md_path.exists() { @@ -107,9 +163,13 @@ fn ensure_default_local_skills_in_dir(skills_root: &Path) -> Result, fs::write(&skill_md_path, skill_content).map_err(|e| { format!("升级默认技能失败 {}: {e}", skill_md_path.display()) })?; + sync_bundled_skill_files(&skill_dir, bundled_skill.extra_files)?; installed.push(skill_name.to_string()); } - _ => continue, // 版本相同或无法比较,跳过 + _ => { + sync_bundled_skill_files(&skill_dir, bundled_skill.extra_files)?; + continue; + } // 版本相同或无法比较,跳过 } continue; } @@ -118,11 +178,28 @@ fn ensure_default_local_skills_in_dir(skills_root: &Path) -> Result, .map_err(|e| format!("创建默认技能目录失败 {}: {e}", skill_dir.display()))?; fs::write(&skill_md_path, skill_content) .map_err(|e| format!("写入默认技能失败 {}: {e}", skill_md_path.display()))?; + sync_bundled_skill_files(&skill_dir, bundled_skill.extra_files)?; installed.push(skill_name.to_string()); } Ok(installed) } +fn sync_bundled_skill_files( + skill_dir: &Path, + extra_files: &[BundledSkillFile], +) -> Result<(), String> { + for extra_file in extra_files { + let target_path = skill_dir.join(extra_file.relative_path); + if let Some(parent) = target_path.parent() { + fs::create_dir_all(parent) + .map_err(|e| format!("创建技能资源目录失败 {}: {e}", parent.display()))?; + } + fs::write(&target_path, extra_file.content) + .map_err(|e| format!("写入技能资源失败 {}: {e}", target_path.display()))?; + } + Ok(()) +} + pub fn ensure_default_local_skills() -> Result, String> { let skills_root = app_paths::resolve_skills_dir()?; ensure_default_local_skills_in_dir(&skills_root) @@ -186,8 +263,8 @@ mod tests { let current_content = fs::read_to_string(&skill_md_path).expect("read skill"); assert_ne!(current_content, old_content, "旧版本内容应被替换"); assert!( - current_content.contains("steps-json"), - "升级后应包含 steps-json 字段" + current_content.contains("proxycast_workflow_ref"), + "升级后应包含 workflow 引用字段" ); } @@ -210,6 +287,8 @@ mod tests { .contains("allowed-tools: social_generate_cover_image, search_query")); assert!(SOCIAL_POST_WITH_COVER_SKILL_CONTENT.contains("**配图说明**")); assert!(SOCIAL_POST_WITH_COVER_SKILL_CONTENT.contains("状态:{成功/失败}")); + assert!(SOCIAL_POST_WITH_COVER_SKILL_CONTENT.contains("proxycast_workflow_ref")); + assert!(SOCIAL_POST_WITH_COVER_WORKFLOW_CONTENT.contains("\"id\": \"research\"")); } #[test] @@ -224,4 +303,19 @@ mod tests { assert!(RESEARCH_SKILL_CONTENT.contains("name: research")); assert!(TYPESETTING_SKILL_CONTENT.contains("name: typesetting")); } + + #[test] + fn should_sync_extra_files_for_social_post_skill() { + let temp = tempfile::tempdir().expect("create temp dir"); + let skills_root = skills_root_from_base(temp.path()); + ensure_default_local_skills_in_dir(&skills_root).expect("install"); + + let workflow_path = skills_root + .join(SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY) + .join("references") + .join("workflow.json"); + assert!(workflow_path.exists()); + let workflow_content = fs::read_to_string(workflow_path).expect("read workflow"); + assert!(workflow_content.contains("\"cover\"")); + } } diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 55454c98d..e979cfd3d 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.85.0", + "version": "0.86.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src-tauri/tests/real_web_search_policy.rs b/src-tauri/tests/real_web_search_policy.rs index dc0712d59..b487577a3 100644 --- a/src-tauri/tests/real_web_search_policy.rs +++ b/src-tauri/tests/real_web_search_policy.rs @@ -1,13 +1,10 @@ use futures::StreamExt; use proxycast_agent::{ - convert_agent_event, AsterAgentState, SessionConfigBuilder, TauriAgentEvent, + convert_agent_event, merge_system_prompt_with_request_tool_policy, resolve_request_tool_policy, + AsterAgentState, SessionConfigBuilder, TauriAgentEvent, WebSearchExecutionTracker, }; use proxycast_core::database::dao::api_key_provider::ApiProviderType; use proxycast_core::database::init_database; -use proxycast_lib::services::request_tool_policy_prompt_service::{ - merge_system_prompt_with_request_tool_policy, resolve_request_tool_policy, - WebSearchExecutionTracker, -}; use proxycast_services::api_key_provider_service::ApiKeyProviderService; use uuid::Uuid; diff --git a/src-tauri/tests/real_web_search_preflight_short_input.rs b/src-tauri/tests/real_web_search_preflight_short_input.rs index 3d45bea23..c4bfc6057 100644 --- a/src-tauri/tests/real_web_search_preflight_short_input.rs +++ b/src-tauri/tests/real_web_search_preflight_short_input.rs @@ -1,9 +1,9 @@ -use proxycast_agent::AsterAgentState; +use proxycast_agent::{ + execute_web_search_preflight_if_needed, resolve_request_tool_policy_with_mode, AsterAgentState, + RequestToolPolicyMode, WebSearchExecutionTracker, +}; use proxycast_core::database::dao::api_key_provider::ApiProviderType; use proxycast_core::database::init_database; -use proxycast_lib::services::request_tool_policy_prompt_service::{ - execute_web_search_preflight_if_needed, resolve_request_tool_policy, WebSearchExecutionTracker, -}; use proxycast_services::api_key_provider_service::ApiKeyProviderService; use uuid::Uuid; @@ -84,7 +84,11 @@ async fn test_real_web_search_preflight_short_input_continue() { let guard = agent_arc.read().await; let agent = guard.as_ref().expect("Agent 未初始化"); - let policy = resolve_request_tool_policy(Some(true), false); + let policy = resolve_request_tool_policy_with_mode( + Some(true), + Some(RequestToolPolicyMode::Required), + false, + ); let mut tracker = WebSearchExecutionTracker::default(); let execution = execute_web_search_preflight_if_needed( agent, diff --git a/src/components/Modal.test.tsx b/src/components/Modal.test.tsx new file mode 100644 index 000000000..71670063f --- /dev/null +++ b/src/components/Modal.test.tsx @@ -0,0 +1,97 @@ +import { act } from "react"; +import { createRoot, type Root } from "react-dom/client"; +import { afterEach, beforeEach, describe, expect, it } from "vitest"; +import { Modal } from "./Modal"; + +interface MountedRoot { + container: HTMLDivElement; + root: Root; +} + +const mountedRoots: MountedRoot[] = []; + +function renderModal() { + const container = document.createElement("div"); + document.body.appendChild(container); + const root = createRoot(container); + + act(() => { + root.render( + {}} + draggable={true} + dragHandleSelector='[data-drag-handle="true"]' + > +
+
拖拽头部
+
弹窗内容
+
+
, + ); + }); + + const mounted = { container, root }; + mountedRoots.push(mounted); + return mounted; +} + +beforeEach(() => { + ( + globalThis as typeof globalThis & { + IS_REACT_ACT_ENVIRONMENT?: boolean; + } + ).IS_REACT_ACT_ENVIRONMENT = true; +}); + +afterEach(() => { + while (mountedRoots.length > 0) { + const mounted = mountedRoots.pop(); + if (!mounted) { + break; + } + act(() => { + mounted.root.unmount(); + }); + mounted.container.remove(); + } +}); + +describe("Modal", () => { + it("启用 draggable 时应支持通过手柄拖动弹窗", () => { + renderModal(); + + const dragHandle = document.body.querySelector( + '[data-drag-handle="true"]', + ) as HTMLDivElement | null; + const modalSurface = document.body.querySelector( + '[data-draggable="true"]', + ) as HTMLDivElement | null; + + act(() => { + dragHandle?.dispatchEvent( + new MouseEvent("mousedown", { + bubbles: true, + clientX: 20, + clientY: 30, + }), + ); + }); + + act(() => { + window.dispatchEvent( + new MouseEvent("mousemove", { + bubbles: true, + clientX: 70, + clientY: 90, + }), + ); + }); + + expect(modalSurface?.style.transform).toBe("translate(50px, 60px)"); + + act(() => { + window.dispatchEvent(new MouseEvent("mouseup", { bubbles: true })); + }); + }); +}); diff --git a/src/components/Modal.tsx b/src/components/Modal.tsx index ccfc4dadb..b403bcb73 100644 --- a/src/components/Modal.tsx +++ b/src/components/Modal.tsx @@ -1,4 +1,10 @@ -import { useEffect, type ReactNode, type MouseEvent } from "react"; +import { + useEffect, + useRef, + useState, + type ReactNode, + type MouseEvent, +} from "react"; import { createPortal } from "react-dom"; import { X } from "lucide-react"; @@ -14,6 +20,10 @@ interface ModalProps { closeOnOverlayClick?: boolean; /** 内容区最大宽度,默认 max-w-lg */ maxWidth?: string; + /** 是否允许拖拽弹窗 */ + draggable?: boolean; + /** 指定拖拽手柄选择器,仅命中该区域时才允许拖拽 */ + dragHandleSelector?: string; } export function Modal({ @@ -24,7 +34,17 @@ export function Modal({ showCloseButton = true, closeOnOverlayClick = true, maxWidth = "max-w-lg", + draggable = false, + dragHandleSelector, }: ModalProps) { + const [dragOffset, setDragOffset] = useState({ x: 0, y: 0 }); + const dragStateRef = useRef<{ + startX: number; + startY: number; + originX: number; + originY: number; + } | null>(null); + // ESC 键关闭 useEffect(() => { if (!isOpen) return; @@ -51,6 +71,12 @@ export function Modal({ }; }, [isOpen]); + useEffect(() => { + if (!isOpen) { + setDragOffset({ x: 0, y: 0 }); + } + }, [isOpen]); + if (!isOpen) return null; const handleOverlayClick = (e: MouseEvent) => { @@ -59,6 +85,65 @@ export function Modal({ } }; + const handleDragStart = (e: MouseEvent) => { + if (!draggable) { + return; + } + + const target = e.target as HTMLElement | null; + if ( + dragHandleSelector && + (!target || !target.closest(dragHandleSelector)) + ) { + return; + } + + const insideInteractive = Boolean( + target?.closest( + 'button, a, input, textarea, select, [role="button"], [role="link"]', + ), + ); + const insideHandle = dragHandleSelector + ? Boolean(target?.closest(dragHandleSelector)) + : true; + + if (insideInteractive && !insideHandle) { + return; + } + + dragStateRef.current = { + startX: e.clientX, + startY: e.clientY, + originX: dragOffset.x, + originY: dragOffset.y, + }; + + const originalUserSelect = document.body.style.userSelect; + document.body.style.userSelect = "none"; + + const handleMouseMove = (moveEvent: globalThis.MouseEvent) => { + const state = dragStateRef.current; + if (!state) { + return; + } + + setDragOffset({ + x: state.originX + (moveEvent.clientX - state.startX), + y: state.originY + (moveEvent.clientY - state.startY), + }); + }; + + const handleMouseUp = () => { + dragStateRef.current = null; + document.body.style.userSelect = originalUserSelect; + window.removeEventListener("mousemove", handleMouseMove); + window.removeEventListener("mouseup", handleMouseUp); + }; + + window.addEventListener("mousemove", handleMouseMove); + window.addEventListener("mouseup", handleMouseUp); + }; + return createPortal(
{showCloseButton && (
)} -

已提交,等待助手继续执行...

+

+ {isQueued + ? "答案已记录,等待系统请求 ID 就绪后会自动提交。" + : "已提交,等待助手继续执行..."} +

); @@ -614,9 +621,7 @@ export function DecisionPanel({ request, onSubmit }: DecisionPanelProps) { option.label, ); const shouldAutoSubmit = - questions.length === 1 && - !q.multiSelect && - !isFallbackAskPending; + questions.length === 1 && !q.multiSelect; return ( + ) : null} + +
+ + + {results.length > 0 && resultsExpanded ? ( + + ) : !results.length && signal.preview ? ( +
+ +
+ ) : null} + + ); +} + +function SearchOutputBatchCard({ + signals, + onOpenUrl, + onOpenDetail, +}: { + signals: HarnessOutputSignal[]; + onOpenUrl: (url: string) => void | Promise; + onOpenDetail: (signal: HarnessOutputSignal) => void; +}) { + const [expanded, setExpanded] = useState(true); + const semanticSummaries = useMemo( + () => summarizeSearchQuerySemantics(signals.map((signal) => signal.summary)), + [signals], + ); + const preview = signals + .slice(0, 2) + .map((signal) => signal.summary) + .join(" · "); + const hiddenCount = Math.max(signals.length - 2, 0); + + return ( +
+ + {semanticSummaries.length > 0 ? ( +
+ {semanticSummaries.map((item) => ( + + {item.label} {item.count} + + ))} +
+ ) : null} + + {expanded ? ( +
+ {signals.map((signal) => ( + onOpenDetail(signal)} + /> + ))} +
+ ) : null} +
+ ); +} + function SummaryCard({ title, value, hint, icon: Icon, onClick, + compact = false, }: { title: string; value: string; hint: string; icon: LucideIcon; onClick?: () => void; + compact?: boolean; }) { const cardContent = (
{title}
-
{value}
+
+ {value} +
{hint}
@@ -365,7 +866,10 @@ function SummaryCard({ return ( + ))} +
+ + ) : null} + + {harnessState.outputSignals.length > 0 ? ( +
+
+
+ {outputFilterOptions.map((option) => { + const count = + option.value === "all" + ? harnessState.outputSignals.length + : harnessState.outputSignals.filter((signal) => + matchesOutputFilter(signal, option.value), + ).length; + const active = option.value === outputFilter; + + return ( + + ); + })} +
+ {filteredOutputSignals.length > 0 ? ( + groupedOutputEntries.map((entry) => { + if (entry.type === "search_batch") { + if (entry.signals.length === 1) { + const signal = entry.signals[0]; + return ( + + void openPreview({ + title: signal.title, + description: signal.summary, + path: getSignalPath(signal), + content: signal.content, + preview: signal.preview, + }) + } + /> + ); + } + + return ( + signal.id).join("|")} + signals={entry.signals} + onOpenUrl={handleOpenExternalLink} + onOpenDetail={(signal) => + void openPreview({ + title: signal.title, + description: signal.summary, + path: getSignalPath(signal), + content: signal.content, + preview: signal.preview, + }) + } + /> + ); + } + + const signal = entry.signal; + const signalPath = getSignalPath(signal); + const signalUrl = findFirstUrl( + signal.summary, + signal.content, + signal.preview, + signal.title, + ); + const canOpenPreview = Boolean( + signalPath || signal.content || signal.preview, + ); + const canOpenUrl = !canOpenPreview && Boolean(signalUrl); + + return ( + + ); + }) + ) : ( +
+ 当前筛选条件下暂无记录。 +
+ )} +
+
+ ) : null} + + {harnessState.pendingApprovals.length > 0 ? ( +
+
+ {harnessState.pendingApprovals.map((item) => ( +
+
+ + +
+ {describeApproval(item) ? ( + + ) : null} +
+ 请求 ID:{item.requestId} +
+
+ ))} +
+
+ ) : null} + {harnessState.recentFileEvents.length > 0 ? (
-
- {group.path} -
+
{group.count} 次活动 @@ -1083,9 +2001,14 @@ export function HarnessStatusPanel({ {group.actionSummary}
{latestEvent.preview ? ( -
-                                  {latestEvent.preview}
-                                
+
+ +
) : null} ); @@ -1117,9 +2040,12 @@ export function HarnessStatusPanel({ {event.displayName} -
- {event.path} -
+
@@ -1137,9 +2063,14 @@ export function HarnessStatusPanel({ {event.sourceToolName}
{event.preview ? ( -
-                                  {event.preview}
-                                
+
+ +
) : null} ); @@ -1154,175 +2085,6 @@ export function HarnessStatusPanel({ ) : null} - {harnessState.outputSignals.length > 0 ? ( -
-
-
- {outputFilterOptions.map((option) => { - const count = - option.value === "all" - ? harnessState.outputSignals.length - : harnessState.outputSignals.filter((signal) => - matchesOutputFilter(signal, option.value), - ).length; - const active = option.value === outputFilter; - - return ( - - ); - })} -
- {filteredOutputSignals.length > 0 ? ( - filteredOutputSignals.map((signal) => { - const signalPath = getSignalPath(signal); - return ( - - ); - }) - ) : ( -
- 当前筛选条件下暂无记录。 -
- )} -
-
- ) : null} - - {harnessState.runtimeStatus ? ( -
-
-
-
- - {harnessState.runtimeStatus.title} -
-
- {harnessState.runtimeStatus.detail} -
-
- - {harnessState.runtimeStatus.checkpoints && - harnessState.runtimeStatus.checkpoints.length > 0 ? ( -
- {harnessState.runtimeStatus.checkpoints.map( - (checkpoint, index) => ( - - {checkpoint} - - ), - )} -
- ) : null} -
-
- ) : null} - - {harnessState.pendingApprovals.length > 0 ? ( -
-
- {harnessState.pendingApprovals.map((item) => ( -
-
- - {item.prompt || "等待用户确认"} -
- {describeApproval(item) ? ( -
- {describeApproval(item)} -
- ) : null} -
- 请求 ID:{item.requestId} -
-
- ))} -
-
- ) : null} - {harnessState.plan.phase !== "idle" || harnessState.plan.items.length > 0 ? (
-
- {item.content} -
+ 0 ? (
当前任务: - {subAgentRuntime.progress.currentTasks.join("、")} +
) : null} @@ -1441,9 +2208,11 @@ export function HarnessStatusPanel({ {task.model ? 模型:{task.model} : null} {task.summary ? ( -
- {task.summary} -
+ ) : null} - {summarizeSchedulerEvent(event)} + ))} @@ -1485,13 +2257,19 @@ export function HarnessStatusPanel({ {subAgentRuntime.error ? (
- {subAgentRuntime.error} +
) : null} {subAgentRuntime.result?.mergedSummary ? (
- {subAgentRuntime.result.mergedSummary} +
) : null} @@ -1515,9 +2293,11 @@ export function HarnessStatusPanel({ {step.stage} -
- {step.detail} -
+ ))} @@ -1534,9 +2314,13 @@ export function HarnessStatusPanel({
{environment.skillNames.length > 0 ? ( environment.skillNames.map((name) => ( - - {name} - + )) ) : ( @@ -1548,9 +2332,13 @@ export function HarnessStatusPanel({
{environment.memorySignals.length > 0 ? ( environment.memorySignals.map((signal) => ( - - {signal} - + )) ) : ( @@ -1565,7 +2353,20 @@ export function HarnessStatusPanel({ {environment.contextItemsCount}
{environment.contextItemNames.length > 0 ? ( -
活跃上下文:{environment.contextItemNames.join("、")}
+
+
活跃上下文:
+
+ {environment.contextItemNames.map((item) => ( + + ))} +
+
) : null}
@@ -1609,12 +2410,18 @@ export function HarnessStatusPanel({ {previewDialog.title} {previewDialog.description ? ( - {previewDialog.description} + ) : null} {previewDialog.path ? ( - - {previewDialog.path} - + ) : null} @@ -1650,9 +2457,13 @@ export function HarnessStatusPanel({ {previewDialog.error} ) : previewDialog.content ? ( -
-                  {previewDialog.content}
-                
+
+ +
) : (
diff --git a/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx b/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx index d7deb4fc3..b89302bb8 100644 --- a/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx +++ b/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx @@ -2,6 +2,7 @@ import React from "react"; import type { ChatInputAdapter } from "@/components/input-kit/adapters/types"; import type { Character } from "@/lib/api/memory"; import type { Skill } from "@/lib/api/skills"; +import type { QueuedTurnSnapshot } from "@/lib/api/agentRuntime"; import type { MessageImage } from "../../../types"; import { CharacterMention } from "./CharacterMention"; import { InputbarCore } from "./InputbarCore"; @@ -43,6 +44,8 @@ interface InputbarComposerSectionProps { strategy: "react" | "code_orchestrated" | "auto", ) => void; topExtra?: React.ReactNode; + queuedTurns: QueuedTurnSnapshot[]; + onRemoveQueuedTurn?: (queuedTurnId: string) => void | Promise; } export const InputbarComposerSection: React.FC< @@ -74,6 +77,8 @@ export const InputbarComposerSection: React.FC< onManageProviders, setExecutionStrategy, topExtra, + queuedTurns, + onRemoveQueuedTurn, }) => { if (renderThemeWorkbenchGeneratingPanel) { return ( @@ -141,6 +146,8 @@ export const InputbarComposerSection: React.FC< visualVariant={isThemeWorkbenchVariant ? "floating" : "default"} topExtra={topExtra} activeTheme={activeTheme} + queuedTurns={queuedTurns} + onRemoveQueuedTurn={onRemoveQueuedTurn} leftExtra={ ({ })); vi.mock("@/components/ui/tooltip", () => ({ - TooltipProvider: ({ children }: { children: React.ReactNode }) => <>{children}, + TooltipProvider: ({ children }: { children: React.ReactNode }) => ( + <>{children} + ), Tooltip: ({ children }: { children: React.ReactNode }) => <>{children}, - TooltipTrigger: ({ children }: { children: React.ReactNode }) => <>{children}, - TooltipContent: ({ children }: { children: React.ReactNode }) => <>{children}, + TooltipTrigger: ({ children }: { children: React.ReactNode }) => ( + <>{children} + ), + TooltipContent: ({ children }: { children: React.ReactNode }) => ( + <>{children} + ), })); const mountedRoots: Array<{ root: Root; container: HTMLDivElement }> = []; @@ -37,7 +43,9 @@ afterEach(() => { vi.clearAllMocks(); }); -const renderInputbarCore = () => { +const renderInputbarCore = ( + props?: Partial>, +) => { const container = document.createElement("div"); document.body.appendChild(container); const root = createRoot(container); @@ -53,6 +61,7 @@ const renderInputbarCore = () => { showTranslate={false} toolMode="attach-only" visualVariant="floating" + {...props} />, ); }); @@ -64,12 +73,18 @@ const renderInputbarCore = () => { describe("InputbarCore", () => { it("主题工作台未聚焦时应使用单行紧凑态,点击展开,移出后收起", () => { const container = renderInputbarCore(); - const textarea = container.querySelector("textarea") as HTMLTextAreaElement | null; - const inputBar = container.querySelector('[data-testid="inputbar-core-container"]') as HTMLDivElement | null; + const textarea = container.querySelector( + "textarea", + ) as HTMLTextAreaElement | null; + const inputBar = container.querySelector( + '[data-testid="inputbar-core-container"]', + ) as HTMLDivElement | null; expect(textarea).toBeTruthy(); expect(inputBar).toBeTruthy(); expect(textarea?.className).toContain("floating-collapsed"); - expect(container.querySelector('[data-testid="inputbar-tools"]')).toBeNull(); + expect( + container.querySelector('[data-testid="inputbar-tools"]'), + ).toBeNull(); act(() => { inputBar?.dispatchEvent(new MouseEvent("mousedown", { bubbles: true })); @@ -77,25 +92,91 @@ describe("InputbarCore", () => { }); expect(textarea?.className).not.toContain("floating-collapsed"); - expect(container.querySelector('[data-testid="inputbar-tools"]')).toBeTruthy(); + expect( + container.querySelector('[data-testid="inputbar-tools"]'), + ).toBeTruthy(); act(() => { inputBar?.dispatchEvent( - new MouseEvent("mouseout", { bubbles: true, relatedTarget: document.body }), + new MouseEvent("mouseout", { + bubbles: true, + relatedTarget: document.body, + }), ); }); expect(textarea?.className).not.toContain("floating-collapsed"); - expect(container.querySelector('[data-testid="inputbar-tools"]')).toBeTruthy(); + expect( + container.querySelector('[data-testid="inputbar-tools"]'), + ).toBeTruthy(); act(() => { textarea?.blur(); inputBar?.dispatchEvent( - new MouseEvent("mouseout", { bubbles: true, relatedTarget: document.body }), + new MouseEvent("mouseout", { + bubbles: true, + relatedTarget: document.body, + }), ); }); expect(textarea?.className).toContain("floating-collapsed"); - expect(container.querySelector('[data-testid="inputbar-tools"]')).toBeNull(); + expect( + container.querySelector('[data-testid="inputbar-tools"]'), + ).toBeNull(); + }); + + it("生成中应显示排队与停止按钮,并渲染排队列表", () => { + const onSend = vi.fn(); + const onStop = vi.fn(); + const container = renderInputbarCore({ + text: "下一条需求", + onSend, + onStop, + isLoading: true, + queuedTurns: [ + { + queued_turn_id: "queued-1", + message_preview: "本周复盘摘要", + message_text: "这里是完整的排队输入内容,点击后应展开查看。", + created_at: 1700000000000, + image_count: 0, + position: 1, + }, + ], + }); + + const queueButton = Array.from(container.querySelectorAll("button")).find( + (button) => button.textContent?.includes("排队"), + ); + const stopButton = Array.from(container.querySelectorAll("button")).find( + (button) => button.textContent?.includes("停止"), + ); + + expect(queueButton).toBeTruthy(); + expect(stopButton).toBeTruthy(); + expect(container.textContent).toContain("已排队 1"); + expect(container.textContent).not.toContain("这里是完整的排队输入内容"); + + const queueCard = Array.from(container.querySelectorAll("button")).find( + (button) => button.textContent?.includes("本周复盘摘要"), + ); + expect(queueCard).toBeTruthy(); + + act(() => { + queueCard?.dispatchEvent(new MouseEvent("click", { bubbles: true })); + }); + + expect(container.textContent).toContain("这里是完整的排队输入内容"); + + act(() => { + queueButton?.dispatchEvent(new MouseEvent("click", { bubbles: true })); + }); + act(() => { + stopButton?.dispatchEvent(new MouseEvent("click", { bubbles: true })); + }); + + expect(onSend).toHaveBeenCalledTimes(1); + expect(onStop).toHaveBeenCalledTimes(1); }); }); diff --git a/src/components/agent/chat/components/Inputbar/components/InputbarCore.tsx b/src/components/agent/chat/components/Inputbar/components/InputbarCore.tsx index 6ee5f0f22..a9d4ce024 100644 --- a/src/components/agent/chat/components/Inputbar/components/InputbarCore.tsx +++ b/src/components/agent/chat/components/Inputbar/components/InputbarCore.tsx @@ -1,5 +1,6 @@ import React, { useCallback, useRef, useState } from "react"; import { + ActionButtonGroup, Container, InputBarContainer, StyledTextarea, @@ -7,6 +8,7 @@ import { LeftSection, RightSection, SendButton, + SecondaryActionButton, DragHandle, ImagePreviewContainer, ImagePreviewItem, @@ -24,6 +26,8 @@ import { TooltipTrigger, } from "@/components/ui/tooltip"; import type { MessageImage } from "../../../types"; +import type { QueuedTurnSnapshot } from "@/lib/api/agentRuntime"; +import { QueuedTurnsPanel } from "./QueuedTurnsPanel"; const INTERACTIVE_TARGET_SELECTOR = "button, a, input, textarea, select, option, [role='button'], [contenteditable=''], [contenteditable='true'], [contenteditable='plaintext-only']"; @@ -72,6 +76,8 @@ interface InputbarCoreProps { /** 视觉风格 */ visualVariant?: "default" | "floating"; activeTheme?: string; + queuedTurns?: QueuedTurnSnapshot[]; + onRemoveQueuedTurn?: (queuedTurnId: string) => void | Promise; } export const InputbarCore: React.FC = ({ @@ -100,6 +106,8 @@ export const InputbarCore: React.FC = ({ showDragHandle = true, visualVariant = "default", activeTheme, + queuedTurns = [], + onRemoveQueuedTurn, }) => { const [isComposerExpanded, setIsComposerExpanded] = useState(false); const inputBarContainerRef = useRef(null); @@ -108,7 +116,8 @@ export const InputbarCore: React.FC = ({ isFloatingVariant && toolMode === "attach-only" && !isComposerExpanded && - pendingImages.length === 0; + pendingImages.length === 0 && + queuedTurns.length === 0; const shouldUseCompactFloatingComposer = shouldCollapseFloatingTools && !topExtra; const containerClassName = [ @@ -188,6 +197,7 @@ export const InputbarCore: React.FC = ({ maxAutoHeight={isFloatingVariant ? 160 : 300} textareaRef={externalTextareaRef} onEscape={() => onToolClick("fullscreen")} + allowSendWhileLoading placeholder={ placeholder || (isFullscreen @@ -241,6 +251,10 @@ export const InputbarCore: React.FC = ({ )} {topExtra} + = ({ ) : null} - + {isLoading ? ( - - ) : ( - - )} - + + + 停止 + + ) : null} + + + {isLoading ? 排队 : null} + + diff --git a/src/components/agent/chat/components/Inputbar/components/QueuedTurnsPanel.tsx b/src/components/agent/chat/components/Inputbar/components/QueuedTurnsPanel.tsx new file mode 100644 index 000000000..df921c851 --- /dev/null +++ b/src/components/agent/chat/components/Inputbar/components/QueuedTurnsPanel.tsx @@ -0,0 +1,103 @@ +import React, { useEffect, useState } from "react"; +import { X } from "lucide-react"; +import type { QueuedTurnSnapshot } from "@/lib/api/agentRuntime"; + +interface QueuedTurnsPanelProps { + queuedTurns: QueuedTurnSnapshot[]; + onRemoveQueuedTurn?: (queuedTurnId: string) => void | Promise; +} + +export const QueuedTurnsPanel: React.FC = ({ + queuedTurns, + onRemoveQueuedTurn, +}) => { + const [expandedTurnId, setExpandedTurnId] = useState(null); + + useEffect(() => { + if ( + expandedTurnId && + !queuedTurns.some((item) => item.queued_turn_id === expandedTurnId) + ) { + setExpandedTurnId(null); + } + }, [expandedTurnId, queuedTurns]); + + if (queuedTurns.length === 0) { + return null; + } + + return ( +
+
+ 已排队 {queuedTurns.length} + 按顺序执行 +
+
+ {queuedTurns.map((item) => { + const messageText = item.message_text.trim() + ? item.message_text + : item.message_preview || "空白输入"; + const title = item.message_preview.trim() + ? item.message_preview + : messageText; + const isExpanded = expandedTurnId === item.queued_turn_id; + const detailId = `queued-turn-detail-${item.queued_turn_id}`; + + return ( +
+ + +
+ ); + })} +
+
+ ); +}; diff --git a/src/components/agent/chat/components/Inputbar/index.tsx b/src/components/agent/chat/components/Inputbar/index.tsx index 948c5245f..97a9f5ee0 100644 --- a/src/components/agent/chat/components/Inputbar/index.tsx +++ b/src/components/agent/chat/components/Inputbar/index.tsx @@ -2,6 +2,7 @@ import React from "react"; import type { MessageImage } from "../../types"; import type { Character } from "@/lib/api/memory"; import type { Skill } from "@/lib/api/skills"; +import type { QueuedTurnSnapshot } from "@/lib/api/agentRuntime"; import type { TaskFile } from "../TaskFiles"; import { InputbarComposerSection } from "./components/InputbarComposerSection"; import { InputbarOverlayShell } from "./components/InputbarOverlayShell"; @@ -71,6 +72,8 @@ interface InputbarProps { onA2UISubmit?: (formData: A2UIFormData) => void; /** A2UI 表单已提交提示 */ a2uiSubmissionNotice?: A2UISubmissionNoticeData | null; + queuedTurns?: QueuedTurnSnapshot[]; + onRemoveQueuedTurn?: (queuedTurnId: string) => void | Promise; } export const Inputbar: React.FC = ({ @@ -109,6 +112,8 @@ export const Inputbar: React.FC = ({ pendingA2UIForm, onA2UISubmit, a2uiSubmissionNotice, + queuedTurns = [], + onRemoveQueuedTurn, }) => { const { textareaRef, @@ -216,6 +221,8 @@ export const Inputbar: React.FC = ({ onManageProviders={onManageProviders} setExecutionStrategy={setExecutionStrategy} topExtra={topExtra} + queuedTurns={queuedTurns} + onRemoveQueuedTurn={onRemoveQueuedTurn} /> ); diff --git a/src/components/agent/chat/components/Inputbar/styles.ts b/src/components/agent/chat/components/Inputbar/styles.ts index e8d2ab36d..29d8cc6ac 100644 --- a/src/components/agent/chat/components/Inputbar/styles.ts +++ b/src/components/agent/chat/components/Inputbar/styles.ts @@ -233,6 +233,12 @@ export const RightSection = styled.div` } `; +export const ActionButtonGroup = styled.div` + display: flex; + align-items: center; + gap: 8px; +`; + // --- InputbarTools Styles --- export const ToolButton = styled.button` @@ -272,13 +278,16 @@ export const Divider = styled.div` margin: 0 4px; `; -export const SendButton = styled.button<{ $isStop?: boolean }>` +export const SendButton = styled.button<{ $isStop?: boolean; $hasLabel?: boolean }>` display: flex; align-items: center; justify-content: center; - width: 30px; + gap: ${({ $hasLabel }) => ($hasLabel ? "6px" : "0")}; + width: ${({ $hasLabel }) => ($hasLabel ? "auto" : "30px")}; + min-width: ${({ $hasLabel }) => ($hasLabel ? "68px" : "30px")}; height: 30px; - border-radius: 50%; + padding: ${({ $hasLabel }) => ($hasLabel ? "0 12px" : "0")}; + border-radius: ${({ $hasLabel }) => ($hasLabel ? "999px" : "50%")}; background-color: ${({ $isStop }) => $isStop ? "hsl(var(--destructive))" : "transparent"}; color: ${({ $isStop }) => ($isStop ? "white" : "hsl(var(--primary))")}; @@ -301,6 +310,32 @@ export const SendButton = styled.button<{ $isStop?: boolean }>` } `; +export const SecondaryActionButton = styled.button` + display: inline-flex; + align-items: center; + justify-content: center; + gap: 6px; + min-width: 68px; + height: 30px; + padding: 0 12px; + border-radius: 999px; + border: 1px solid hsl(var(--border)); + background: hsl(var(--background)); + color: hsl(var(--destructive)); + transition: all 0.2s; + + &:hover:not(:disabled) { + border-color: hsl(var(--destructive) / 0.4); + background: hsl(var(--destructive) / 0.06); + } + + &:disabled { + cursor: default; + color: hsl(var(--muted-foreground)); + opacity: 0.6; + } +`; + // --- Image Preview Styles --- export const ImagePreviewContainer = styled.div` diff --git a/src/components/agent/chat/components/MessageList.tsx b/src/components/agent/chat/components/MessageList.tsx index aa4ca34e9..3459feed0 100644 --- a/src/components/agent/chat/components/MessageList.tsx +++ b/src/components/agent/chat/components/MessageList.tsx @@ -28,6 +28,11 @@ import { MarkdownRenderer } from "./MarkdownRenderer"; import { StreamingRenderer } from "./StreamingRenderer"; import { TokenUsageDisplay } from "./TokenUsageDisplay"; import { AgentThreadTimeline } from "./AgentThreadTimeline"; +import { + formatArtifactWritePhaseLabel, + resolveArtifactPreviewText, + resolveArtifactWritePhase, +} from "../utils/messageArtifacts"; import { Message, type AgentThreadItem, @@ -204,12 +209,9 @@ const MessageListInner: React.FC = ({ typeof artifact.meta.filePath === "string" ? artifact.meta.filePath : artifact.meta.filename || artifact.title; - const statusLabel = - artifact.status === "streaming" - ? "生成中" - : artifact.status === "error" - ? "失败" - : "已生成"; + const writePhase = resolveArtifactWritePhase(artifact); + const statusLabel = formatArtifactWritePhaseLabel(writePhase); + const previewText = resolveArtifactPreviewText(artifact, 180); return (
- - {statusLabel} -
diff --git a/src/components/agent/chat/components/SearchResultPreviewList.test.tsx b/src/components/agent/chat/components/SearchResultPreviewList.test.tsx new file mode 100644 index 000000000..e67c0e4e1 --- /dev/null +++ b/src/components/agent/chat/components/SearchResultPreviewList.test.tsx @@ -0,0 +1,93 @@ +import { act } from "react"; +import { createRoot, type Root } from "react-dom/client"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { SearchResultPreviewList } from "./SearchResultPreviewList"; + +interface RenderResult { + container: HTMLDivElement; + root: Root; +} + +const mountedRoots: RenderResult[] = []; + +function renderList() { + const container = document.createElement("div"); + document.body.appendChild(container); + const root = createRoot(container); + + const items = Array.from({ length: 6 }, (_, index) => ({ + id: `result-${index + 1}`, + title: `结果 ${index + 1}`, + url: `https://example.com/${index + 1}`, + hostname: "example.com", + snippet: `摘要 ${index + 1}`, + })); + + act(() => { + root.render( + , + ); + }); + + const rendered = { container, root }; + mountedRoots.push(rendered); + return rendered; +} + +beforeEach(() => { + ( + globalThis as typeof globalThis & { + IS_REACT_ACT_ENVIRONMENT?: boolean; + } + ).IS_REACT_ACT_ENVIRONMENT = true; +}); + +afterEach(() => { + while (mountedRoots.length > 0) { + const mounted = mountedRoots.pop(); + if (!mounted) break; + act(() => { + mounted.root.unmount(); + }); + mounted.container.remove(); + } +}); + +describe("SearchResultPreviewList", () => { + it("搜索结果应默认折叠,并支持展开与收起", () => { + const { container } = renderList(); + + expect(container.textContent).toContain("结果 1"); + expect(container.textContent).toContain("结果 4"); + expect(container.textContent).not.toContain("结果 5"); + expect(container.textContent).toContain("展开其余 2 条结果"); + + const toggleButton = container.querySelector( + 'button[aria-label="展开搜索结果"]', + ) as HTMLButtonElement | null; + + act(() => { + toggleButton?.click(); + }); + + expect(container.textContent).toContain("结果 5"); + expect(container.textContent).toContain("结果 6"); + expect(container.textContent).toContain("收起结果"); + + const collapseButton = container.querySelector( + 'button[aria-label="收起搜索结果"]', + ) as HTMLButtonElement | null; + + act(() => { + collapseButton?.click(); + }); + + expect(container.textContent).not.toContain("结果 5"); + expect(container.textContent).not.toContain("结果 6"); + expect(container.textContent).toContain("展开其余 2 条结果"); + }); +}); diff --git a/src/components/agent/chat/components/SearchResultPreviewList.tsx b/src/components/agent/chat/components/SearchResultPreviewList.tsx new file mode 100644 index 000000000..7f4499ce6 --- /dev/null +++ b/src/components/agent/chat/components/SearchResultPreviewList.tsx @@ -0,0 +1,184 @@ +import { useCallback, useEffect, useMemo, useRef, useState } from "react"; +import { + ChevronDown, + ChevronRight, + ExternalLink, + Globe, + Search, +} from "lucide-react"; +import { Popover, PopoverContent, PopoverTrigger } from "@/components/ui/popover"; +import { cn } from "@/lib/utils"; +import type { SearchResultPreviewItem } from "../utils/searchResultPreview"; + +function SearchResultHoverCard({ + item, + onOpenUrl, + popoverSide = "right", + popoverAlign = "start", +}: { + item: SearchResultPreviewItem; + onOpenUrl: (url: string) => void | Promise; + popoverSide?: "top" | "right" | "bottom" | "left"; + popoverAlign?: "start" | "center" | "end"; +}) { + const [open, setOpen] = useState(false); + const closeTimerRef = useRef(null); + + const clearCloseTimer = useCallback(() => { + if (closeTimerRef.current !== null && typeof window !== "undefined") { + window.clearTimeout(closeTimerRef.current); + closeTimerRef.current = null; + } + }, []); + + const handleOpenPreview = useCallback(() => { + clearCloseTimer(); + setOpen(true); + }, [clearCloseTimer]); + + const handleScheduleClose = useCallback(() => { + clearCloseTimer(); + if (typeof window === "undefined") { + setOpen(false); + return; + } + closeTimerRef.current = window.setTimeout(() => { + setOpen(false); + closeTimerRef.current = null; + }, 120); + }, [clearCloseTimer]); + + useEffect(() => () => clearCloseTimer(), [clearCloseTimer]); + + return ( + + + + + +
+
+
+ +
+
+
+ {item.title} +
+
+ + {item.hostname} +
+
+
+
+ {item.snippet || "暂无摘要,点击可直接打开来源。"} +
+ +
+
+
+ ); +} + +export function SearchResultPreviewList({ + items, + onOpenUrl, + popoverSide = "right", + popoverAlign = "start", + className, + collapsedCount = 4, +}: { + items: SearchResultPreviewItem[]; + onOpenUrl: (url: string) => void | Promise; + popoverSide?: "top" | "right" | "bottom" | "left"; + popoverAlign?: "start" | "center" | "end"; + className?: string; + collapsedCount?: number; +}) { + const [expanded, setExpanded] = useState(false); + const identityKey = useMemo( + () => items.map((item) => item.id).join("|"), + [items], + ); + + useEffect(() => { + setExpanded(false); + }, [identityKey]); + + if (items.length === 0) { + return null; + } + + const shouldCollapse = items.length > collapsedCount; + const visibleItems = + shouldCollapse && !expanded ? items.slice(0, collapsedCount) : items; + const hiddenCount = items.length - visibleItems.length; + + return ( +
+ {visibleItems.map((item) => ( + + ))} + {shouldCollapse ? ( + + ) : null} +
+ ); +} diff --git a/src/components/agent/chat/components/StreamingRenderer.test.tsx b/src/components/agent/chat/components/StreamingRenderer.test.tsx index fdccb20e1..cb7d245c7 100644 --- a/src/components/agent/chat/components/StreamingRenderer.test.tsx +++ b/src/components/agent/chat/components/StreamingRenderer.test.tsx @@ -3,7 +3,11 @@ import { act } from "react"; import { createRoot, type Root } from "react-dom/client"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { StreamingRenderer } from "./StreamingRenderer"; -import type { AgentRuntimeStatus, ContentPart } from "../types"; +import type { + AgentRuntimeStatus, + ContentPart, + WriteArtifactContext, +} from "../types"; const parseAIResponseMock = vi.fn(); @@ -89,6 +93,12 @@ function renderHarness(props: { contentParts?: ContentPart[]; renderA2UIInline?: boolean; runtimeStatus?: AgentRuntimeStatus; + showRuntimeStatusInline?: boolean; + onWriteFile?: ( + content: string, + fileName: string, + context?: WriteArtifactContext, + ) => void; }) { const container = document.createElement("div"); document.body.appendChild(container); @@ -185,6 +195,43 @@ describe("StreamingRenderer", () => { expect(container.textContent).toContain("请先补充以下信息:"); }); + it("pending_write_file 应触发流式 onWriteFile 回调", () => { + const onWriteFile = vi.fn(); + parseAIResponseMock.mockReturnValue({ + parts: [ + { + type: "pending_write_file", + content: "# 草稿\n正在生成中", + filePath: "notes/live.md", + }, + ], + hasA2UI: false, + hasWriteFile: true, + hasPending: true, + }); + + renderHarness({ + content: '# 草稿\n正在生成中', + isStreaming: true, + onWriteFile, + }); + + expect(onWriteFile).toHaveBeenCalledTimes(1); + expect(onWriteFile).toHaveBeenCalledWith( + "# 草稿\n正在生成中", + "notes/live.md", + expect.objectContaining({ + source: "message_content", + status: "streaming", + metadata: expect.objectContaining({ + writePhase: "streaming", + lastUpdateSource: "message_content", + isPartial: true, + }), + }), + ); + }); + it("应将 proposed_plan 片段渲染为独立计划卡片", () => { const { container } = renderHarness({ content: @@ -210,6 +257,7 @@ describe("StreamingRenderer", () => { detail: "正在理解请求并准备回合。", checkpoints: ["对话优先执行", "等待首个事件"], }, + showRuntimeStatusInline: true, }); expect(container.textContent).toContain("Agent 正在准备执行"); diff --git a/src/components/agent/chat/components/StreamingRenderer.tsx b/src/components/agent/chat/components/StreamingRenderer.tsx index 3cb208aed..f8d7b9ed3 100644 --- a/src/components/agent/chat/components/StreamingRenderer.tsx +++ b/src/components/agent/chat/components/StreamingRenderer.tsx @@ -511,6 +511,7 @@ interface StreamingRendererProps { /** 代码块点击回调(用于在画布中显示) */ onCodeBlockClick?: (language: string, code: string) => void; runtimeStatus?: AgentRuntimeStatus; + showRuntimeStatusInline?: boolean; } const RUNTIME_PHASE_LABELS: Record = { @@ -579,6 +580,7 @@ export const StreamingRenderer: React.FC = memo( collapseCodeBlocks, onCodeBlockClick, runtimeStatus, + showRuntimeStatusInline = false, }) => { // 判断是否使用交错显示模式 const useInterleavedMode = contentParts && contentParts.length > 0; @@ -626,26 +628,103 @@ export const StreamingRenderer: React.FC = memo( return result; }, [parsedVisibleText, isStreaming, useInterleavedMode]); - // 处理文件写入 - 使用 ref 来追踪已处理的内容 - const processedWriteFilesRef = useRef>(new Set()); + const interleavedParsedContent = useMemo(() => { + if (!useInterleavedMode) { + return []; + } + + return (contentParts || []).map((part) => { + if (part.type !== "text") { + return EMPTY_PARSE_RESULT; + } + + return getCachedStructuredParse(parseCacheRef, part.text, isStreaming); + }); + }, [contentParts, isStreaming, useInterleavedMode]); + + // 处理文件写入 - 使用 ref 追踪同一路径的最新阶段与内容签名 + const processedWriteFilesRef = useRef>(new Map()); + + const emitWriteFile = React.useCallback( + (part: ParsedMessageContent, signatureKey: string) => { + if ( + !onWriteFile || + (part.type !== "write_file" && part.type !== "pending_write_file") || + !part.filePath + ) { + return; + } + + const contentValue = + typeof part.content === "string" ? part.content : ""; + const signature = `${part.type}:${signatureKey}:${contentValue}`; + const previousSignature = processedWriteFilesRef.current.get( + part.filePath, + ); + if (previousSignature === signature) { + return; + } + + processedWriteFilesRef.current.set(part.filePath, signature); + const metadata: WriteArtifactContext["metadata"] = { + writePhase: + part.type === "pending_write_file" + ? "streaming" + : isStreaming + ? "streaming" + : "completed", + previewText: contentValue.trim() + ? contentValue.slice(0, 480).trim() + : undefined, + latestChunk: contentValue.trim() + ? contentValue.slice(-240).trim() + : undefined, + isPartial: part.type === "pending_write_file" || isStreaming, + lastUpdateSource: "message_content", + }; + + onWriteFile(contentValue, part.filePath, { + source: "message_content", + status: + part.type === "pending_write_file" || isStreaming + ? "streaming" + : "complete", + metadata, + }); + }, + [isStreaming, onWriteFile], + ); useEffect(() => { if (!onWriteFile) return; - for (const part of parsedContent.parts) { + const writeCandidates = useInterleavedMode + ? interleavedParsedContent.flatMap((parsed, index) => + parsed.parts.map((part, partIndex) => ({ + part, + signatureKey: `interleaved:${index}:${partIndex}`, + })), + ) + : parsedContent.parts.map((part, index) => ({ + part, + signatureKey: `standard:${index}`, + })); + + for (const candidate of writeCandidates) { if ( - part.type === "write_file" && - part.filePath && - typeof part.content === "string" + candidate.part.type === "write_file" || + candidate.part.type === "pending_write_file" ) { - const key = `${part.filePath}:${part.content.length}`; - if (!processedWriteFilesRef.current.has(key)) { - processedWriteFilesRef.current.add(key); - onWriteFile(part.content, part.filePath); - } + emitWriteFile(candidate.part, candidate.signatureKey); } } - }, [parsedContent.parts, onWriteFile]); + }, [ + emitWriteFile, + interleavedParsedContent, + onWriteFile, + parsedContent.parts, + useInterleavedMode, + ]); // 使用外部提供的思考内容或解析出的内容 const finalThinking = externalThinking || thinkingText; @@ -689,11 +768,8 @@ export const StreamingRenderer: React.FC = memo( if (!partText) return null; // 解析 write_file 标签 - const partParsed = getCachedStructuredParse( - parseCacheRef, - partText, - isStreaming, - ); + const partParsed = + interleavedParsedContent[index] || EMPTY_PARSE_RESULT; const isLastPart = index === contentParts.length - 1; // 添加调试日志 @@ -708,23 +784,6 @@ export const StreamingRenderer: React.FC = memo( ); } - // 处理文件写入回调 - if (onWriteFile) { - for (const p of partParsed.parts) { - if ( - p.type === "write_file" && - p.filePath && - typeof p.content === "string" - ) { - const key = `interleaved-${p.filePath}:${p.content.length}`; - if (!processedWriteFilesRef.current.has(key)) { - processedWriteFilesRef.current.add(key); - onWriteFile(p.content, p.filePath); - } - } - } - } - // 如果包含 write_file,按部分渲染 if (partParsed.hasWriteFile) { return ( @@ -742,7 +801,6 @@ export const StreamingRenderer: React.FC = memo( className="flex items-center gap-2 px-3 py-2 bg-muted/50 rounded-lg text-sm text-muted-foreground cursor-pointer hover:bg-muted/70 transition-colors" onClick={() => p.filePath && - fileContent && onFileClick?.(p.filePath, fileContent) } > @@ -858,6 +916,7 @@ export const StreamingRenderer: React.FC = memo( const hasToolCalls = toolCalls && toolCalls.length > 0; const hasActionRequests = actionRequests && actionRequests.length > 0; const shouldShowRuntimeStatus = + showRuntimeStatusInline && Boolean(runtimeStatus) && isStreaming && !hasVisibleContent && @@ -899,9 +958,7 @@ export const StreamingRenderer: React.FC = memo( key={`write-${index}`} className="flex items-center gap-2 px-3 py-2 bg-muted/50 rounded-lg text-sm text-muted-foreground cursor-pointer hover:bg-muted/70 transition-colors" onClick={() => - part.filePath && - fileContent && - onFileClick?.(part.filePath, fileContent) + part.filePath && onFileClick?.(part.filePath, fileContent) } > diff --git a/src/components/agent/chat/components/ThemeWorkbenchSidebar.tsx b/src/components/agent/chat/components/ThemeWorkbenchSidebar.tsx index f8400349d..950e72824 100644 --- a/src/components/agent/chat/components/ThemeWorkbenchSidebar.tsx +++ b/src/components/agent/chat/components/ThemeWorkbenchSidebar.tsx @@ -2356,7 +2356,8 @@ function ThemeWorkbenchSidebarComponent({ if (n === "load_skill") return "加载技能"; if (n.includes("write_file") || n.includes("create_file")) return "创建文件"; if (n.includes("read_file")) return "读取文件"; - if (n.includes("search_query") || n.includes("web_search") || n === "search") return "网络检索"; + if (n.includes("websearch")) return "网络检索"; + if (n.includes("webfetch")) return "网页抓取"; if (n.includes("social_generate_cover") || n.includes("generate_image")) return "生成封面图"; if (n.includes("execute") || n.includes("bash")) return "执行命令"; if (n.includes("context") || n.includes("retrieve")) return "检索上下文"; diff --git a/src/components/agent/chat/components/ToolCallDisplay.test.tsx b/src/components/agent/chat/components/ToolCallDisplay.test.tsx index 94e072bcb..9db6bdef1 100644 --- a/src/components/agent/chat/components/ToolCallDisplay.test.tsx +++ b/src/components/agent/chat/components/ToolCallDisplay.test.tsx @@ -1,16 +1,46 @@ -import React from "react"; import { act } from "react"; import { createRoot, type Root } from "react-dom/client"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { ToolCallDisplay, ToolCallList } from "./ToolCallDisplay"; import type { ToolCallState } from "@/lib/api/agentStream"; -import { ToolCallDisplay } from "./ToolCallDisplay"; -interface MountedHarness { +vi.mock("@tauri-apps/plugin-shell", () => ({ + open: vi.fn().mockResolvedValue(undefined), +})); + +interface RenderResult { container: HTMLDivElement; root: Root; } -const mountedRoots: MountedHarness[] = []; +const mountedRoots: RenderResult[] = []; + +function renderTool(toolCall: ToolCallState): RenderResult { + const container = document.createElement("div"); + document.body.appendChild(container); + const root = createRoot(container); + + act(() => { + root.render(); + }); + + const rendered = { container, root }; + mountedRoots.push(rendered); + return rendered; +} + +afterEach(() => { + while (mountedRoots.length > 0) { + const mounted = mountedRoots.pop(); + if (!mounted) { + break; + } + act(() => { + mounted.root.unmount(); + }); + mounted.container.remove(); + } +}); beforeEach(() => { ( @@ -20,188 +50,115 @@ beforeEach(() => { ).IS_REACT_ACT_ENVIRONMENT = true; }); -afterEach(() => { - while (mountedRoots.length > 0) { - const mounted = mountedRoots.pop(); - if (!mounted) break; - act(() => { - mounted.root.unmount(); - }); - mounted.container.remove(); - } -}); - -function render( - toolCall: ToolCallState, - options: { - onFileClick?: (fileName: string, content: string) => void; - } = {}, -): HTMLDivElement { - const container = document.createElement("div"); - document.body.appendChild(container); - const root = createRoot(container); - - act(() => { - root.render( - , - ); - }); - - mountedRoots.push({ container, root }); - return container; -} - describe("ToolCallDisplay", () => { - it("工具结果包含图片时应渲染缩略图预览", () => { - const toolCall: ToolCallState = { - id: "tool-image-1", - name: "Read", + it("WebSearch 工具结果应在 AI 对话区展示搜索列表并支持悬浮预览", async () => { + renderTool({ + id: "tool-search-1", + name: "WebSearch", + arguments: JSON.stringify({ query: "3月13日国际新闻" }), status: "completed", - startTime: new Date(), - endTime: new Date(), result: { success: true, - output: "图片已生成", - images: [ - { src: "data:image/png;base64,aGVsbG8=", mimeType: "image/png" }, - ], + output: [ + "Xinhua world news summary at 0030 GMT, March 13", + "https://example.com/xinhua", + "全球要闻摘要,覆盖国际局势与市场动态。", + "", + "Friday morning news: March 13, 2026 | WORLD - wng.org", + "https://example.com/wng", + "补充国际动态与区域冲突更新。", + ].join("\n"), }, - }; - - const container = render(toolCall); - const previewImage = container.querySelector( - 'img[alt="工具结果图片预览"]', - ) as HTMLImageElement | null; - expect(previewImage).not.toBeNull(); - expect(previewImage?.src).toContain("data:image/png;base64,aGVsbG8="); - }); - - it("点击缩略图后应显示大图预览层", () => { - const toolCall: ToolCallState = { - id: "tool-image-2", - name: "Read", - status: "completed", - startTime: new Date(), - endTime: new Date(), - result: { - success: true, - output: "图片已生成", - images: [ - { src: "data:image/png;base64,aGVsbG8=", mimeType: "image/png" }, - ], - }, - }; - - const container = render(toolCall); - const thumbnail = container.querySelector( - 'img[alt="工具结果图片预览"]', - ) as HTMLImageElement | null; - expect(thumbnail).not.toBeNull(); - - act(() => { - thumbnail?.click(); + startTime: new Date("2026-03-13T12:00:00.000Z"), + endTime: new Date("2026-03-13T12:00:02.000Z"), }); - const enlargedImage = document.querySelector( - 'img[alt="工具结果图片大图"]', - ) as HTMLImageElement | null; - expect(enlargedImage).not.toBeNull(); - }); - - it("工具结果包含 metadata 时应渲染执行摘要", () => { - const toolCall: ToolCallState = { - id: "tool-meta-1", - name: "Bash", - status: "failed", - startTime: new Date(), - endTime: new Date(), - result: { - success: false, - output: "命令执行失败", - metadata: { - exit_code: 1, - stdout_length: 120, - stderr_length: 32, - sandboxed: true, - output_file: "/tmp/aster_tasks/task-1.log", - }, - }, - }; - - const container = render(toolCall); - expect(container.textContent).toContain("退出码 1"); - expect(container.textContent).toContain("stdout 120"); - expect(container.textContent).toContain("已隔离执行"); - expect(container.textContent).toContain( - "输出文件: /tmp/aster_tasks/task-1.log", + expect(document.body.textContent).toContain( + "Xinhua world news summary at 0030 GMT, March 13", ); - }); - - it("工具结果完成 offload 转存时应显示转存摘要与文件路径", () => { - const toolCall: ToolCallState = { - id: "tool-offload-1", - name: "Write", - status: "completed", - startTime: new Date(), - endTime: new Date(), - result: { - success: true, - output: - "preview line\n\n[ProxyCast Offload] 完整输出已转存到文件:/tmp/proxycast/harness/tool-io/results/tool-offload-1.json", - metadata: { - proxycast_offloaded: true, - offload_file: - "/tmp/proxycast/harness/tool-io/results/tool-offload-1.json", - offload_original_chars: 18234, - offload_original_tokens: 4521, - offload_trigger: "token_limit_before_evict", - }, - }, - }; - - const container = render(toolCall); - expect(container.textContent).toContain("完整输出已转存"); - expect(container.textContent).toContain("原始 18234 字符"); - expect(container.textContent).toContain("约 4521 tokens"); - expect(container.textContent).toContain("token 阈值触发"); - expect(container.textContent).toContain( - "转存文件: /tmp/proxycast/harness/tool-io/results/tool-offload-1.json", + expect(document.body.textContent).toContain( + "Friday morning news: March 13, 2026 | WORLD - wng.org", ); - }); - it("存在文件路径时应显示打开图标,并可直接送入画布", () => { - const onFileClick = vi.fn(); - const toolCall: ToolCallState = { - id: "tool-open-file-1", - name: "Write", - status: "completed", - startTime: new Date(), - endTime: new Date(), - result: { - success: true, - output: "文件已生成", - metadata: { - output_file: "/tmp/workspace/summary.md", - }, - }, - }; - - const container = render(toolCall, { onFileClick }); - const openButton = container.querySelector( - 'button[aria-label="在画布中打开-/tmp/workspace/summary.md"]', + const firstSearchResult = document.body.querySelector( + '[aria-label="预览搜索结果:Xinhua world news summary at 0030 GMT, March 13"]', ) as HTMLButtonElement | null; - expect(openButton).not.toBeNull(); - - act(() => { - openButton?.click(); + await act(async () => { + firstSearchResult?.dispatchEvent( + new MouseEvent("mouseover", { bubbles: true }), + ); + await Promise.resolve(); }); - expect(onFileClick).toHaveBeenCalledWith("/tmp/workspace/summary.md", ""); + expect(document.body.textContent).toContain( + "全球要闻摘要,覆盖国际局势与市场动态。", + ); + expect(document.body.textContent).toContain("https://example.com/xinhua"); + + const collapseButton = document.body.querySelector( + 'button[title="收起详情"]', + ) as HTMLButtonElement | null; + + act(() => { + collapseButton?.click(); + }); + + expect(document.body.textContent).not.toContain( + "Xinhua world news summary at 0030 GMT, March 13", + ); + + const expandButton = document.body.querySelector( + 'button[title="展开详情"]', + ) as HTMLButtonElement | null; + + act(() => { + expandButton?.click(); + }); + + expect(document.body.textContent).toContain( + "Xinhua world news summary at 0030 GMT, March 13", + ); + }); + + it("连续多次 WebSearch 应在对话区按搜索批次分组展示", () => { + const container = document.createElement("div"); + document.body.appendChild(container); + const root = createRoot(container); + + act(() => { + root.render( + , + ); + }); + + mountedRoots.push({ container, root }); + + expect(container.textContent).toContain("已搜索 2 组查询"); + expect(container.textContent).toContain("3月13日国际新闻"); + expect(container.textContent).toContain("March 13 2026 world headlines"); + expect(container.textContent).toContain("中文日期检索"); + expect(container.textContent).toContain("头条检索"); }); }); diff --git a/src/components/agent/chat/components/ToolCallDisplay.tsx b/src/components/agent/chat/components/ToolCallDisplay.tsx index 09e844766..f56c9bed5 100644 --- a/src/components/agent/chat/components/ToolCallDisplay.tsx +++ b/src/components/agent/chat/components/ToolCallDisplay.tsx @@ -12,7 +12,9 @@ import React, { useMemo, useCallback, } from "react"; +import { open as openExternal } from "@tauri-apps/plugin-shell"; import { + ChevronDown, Terminal, FileText, Edit3, @@ -31,6 +33,15 @@ import { import { cn } from "@/lib/utils"; import type { ToolCallState, ToolResultImage } from "@/lib/api/agentStream"; import { MarkdownRenderer } from "./MarkdownRenderer"; +import { SearchResultPreviewList } from "./SearchResultPreviewList"; +import { + isUnifiedWebSearchToolName, + resolveSearchResultPreviewItemsFromText, +} from "../utils/searchResultPreview"; +import { + classifySearchQuerySemantic, + summarizeSearchQuerySemantics, +} from "../utils/searchQueryGrouping"; // ============ 类型定义 ============ @@ -132,6 +143,26 @@ const getToolIcon = (toolName: string) => { const normalizeToolNameKey = (value: string): string => value.replace(/[\s_-]+/g, "").trim().toLowerCase(); +const extractSearchQueryLabel = (toolCall: ToolCallState): string => { + try { + const args = toolCall.arguments ? JSON.parse(toolCall.arguments) : {}; + const record = + args && typeof args === "object" && !Array.isArray(args) + ? (args as Record) + : {}; + for (const key of ["query", "q", "pattern", "search", "url"]) { + const value = record[key]; + if (typeof value === "string" && value.trim()) { + return value.trim(); + } + } + } catch { + // ignore parse failure + } + + return toolCall.name; +}; + const PLANNING_TOOL_KEYS = new Set([ "todowrite", "writetodos", @@ -623,6 +654,7 @@ export const ToolCallDisplay: React.FC = ({ }) => { const [isExpanded, setIsExpanded] = useState(defaultExpanded); const [previewImageSrc, setPreviewImageSrc] = useState(null); + const hasUserToggledExpandedRef = useRef(false); // 解析参数 const parsedArgs = useMemo(() => { @@ -777,17 +809,46 @@ export const ToolCallDisplay: React.FC = ({ () => resultPath?.value || filePath, [filePath, resultPath?.value], ); + const searchResultItems = useMemo(() => { + if (!isUnifiedWebSearchToolName(toolCall.name)) { + return []; + } + + return resolveSearchResultPreviewItemsFromText(toolCall.result?.output); + }, [toolCall.name, toolCall.result?.output]); + const searchSemantic = useMemo( + () => classifySearchQuerySemantic(extractSearchQueryLabel(toolCall)), + [toolCall], + ); const hasResultImages = resultImages.length > 0; + const hasSearchResults = searchResultItems.length > 0; + + const handleOpenExternalUrl = useCallback(async (url: string) => { + try { + await openExternal(url); + } catch { + if (typeof window !== "undefined" && typeof window.open === "function") { + window.open(url, "_blank"); + return; + } + throw new Error("当前环境不支持打开外部链接"); + } + }, []); useEffect(() => { - if (isMessageStreaming && (isRunning || hasResult || hasResultImages)) { + if ( + isMessageStreaming && + (isRunning || hasResult || hasResultImages || hasSearchResults) + ) { setIsExpanded(true); - return; } - if (!isMessageStreaming && !isRunning) { - setIsExpanded(false); + }, [isMessageStreaming, isRunning, hasResult, hasResultImages, hasSearchResults]); + + useEffect(() => { + if (hasSearchResults && !hasUserToggledExpandedRef.current) { + setIsExpanded(true); } - }, [isMessageStreaming, isRunning, hasResult, hasResultImages]); + }, [hasSearchResults]); // 处理点击事件 - 如果是文件写入工具,打开右边栏 const handleOpenFile = useCallback(() => { @@ -796,6 +857,11 @@ export const ToolCallDisplay: React.FC = ({ } }, [fileContent, onFileClick, openableFilePath]); + const handleToggleExpanded = useCallback(() => { + hasUserToggledExpandedRef.current = true; + setIsExpanded((prev) => !prev); + }, []); + // 简洁模式:单行显示 - Claude 风格 return (
@@ -863,9 +929,9 @@ export const ToolCallDisplay: React.FC = ({ )} {/* 展开/折叠按钮 */} - {hasResult && ( + {(hasResult || hasSearchResults) && (
)} + {hasSearchResults && isExpanded && ( +
+
+ + {searchSemantic.label} + +
+ +
+ )} + {/* 展开的详情 - Claude 风格 */} {isExpanded && hasResult && (
@@ -971,20 +1053,147 @@ export const ToolCallList: React.FC = ({ }) => { if (!toolCalls || toolCalls.length === 0) return null; + const groups: Array< + | { + type: "search"; + id: string; + items: ToolCallState[]; + } + | { + type: "single"; + id: string; + item: ToolCallState; + } + > = []; + + for (const toolCall of toolCalls) { + const isSearch = isUnifiedWebSearchToolName(toolCall.name); + const lastGroup = groups[groups.length - 1]; + if ( + isSearch && + lastGroup && + lastGroup.type === "search" + ) { + lastGroup.items.push(toolCall); + continue; + } + + if (isSearch) { + groups.push({ + type: "search", + id: `search-group:${toolCall.id}`, + items: [toolCall], + }); + continue; + } + + groups.push({ + type: "single", + id: toolCall.id, + item: toolCall, + }); + } + return (
- {toolCalls.map((tc) => ( - - ))} + {groups.map((group) => { + if (group.type === "single") { + return ( + + ); + } + + return ( + + ); + })}
); }; +function SearchToolCallGroup({ + toolCalls, + isMessageStreaming, + onFileClick, +}: { + toolCalls: ToolCallState[]; + isMessageStreaming: boolean; + onFileClick?: (fileName: string, content: string) => void; +}) { + const [expanded, setExpanded] = useState(true); + const semanticSummaries = summarizeSearchQuerySemantics( + toolCalls.map(extractSearchQueryLabel), + ); + const queryPreview = toolCalls + .slice(0, 2) + .map(extractSearchQueryLabel) + .join(" · "); + const hiddenCount = Math.max(toolCalls.length - 2, 0); + + return ( +
+ + {semanticSummaries.length > 0 ? ( +
+ {semanticSummaries.map((item) => ( + + {item.label} {item.count} + + ))} +
+ ) : null} + {expanded ? ( +
+ {toolCalls.map((toolCall) => ( + + ))} +
+ ) : null} +
+ ); +} + // 导出别名,用于交错显示模式 export const ToolCallItem = ToolCallDisplay; diff --git a/src/components/agent/chat/hooks/agentRuntimeAdapter.ts b/src/components/agent/chat/hooks/agentRuntimeAdapter.ts index f45d7c731..c397005d7 100644 --- a/src/components/agent/chat/hooks/agentRuntimeAdapter.ts +++ b/src/components/agent/chat/hooks/agentRuntimeAdapter.ts @@ -6,12 +6,14 @@ import { getAgentRuntimeSession, initAsterAgent, interruptAgentRuntimeTurn, + removeAgentRuntimeQueuedTurn, listAgentRuntimeSessions, respondAgentRuntimeAction, submitAgentRuntimeTurn, updateAgentRuntimeSession, type AsterExecutionStrategy, type AsterProviderConfig, + type AgentSearchMode, type AsterSessionDetail, type AsterSessionInfo, type AutoContinueRequestPayload, @@ -28,9 +30,12 @@ export interface AgentRuntimeTurnRequest { providerConfig?: AsterProviderConfig; executionStrategy?: AsterExecutionStrategy; webSearch?: boolean; + searchMode?: AgentSearchMode; autoContinue?: AutoContinueRequestPayload; systemPrompt?: string; metadata?: Record; + queueIfBusy?: boolean; + queuedTurnId?: string; } export interface AgentRuntimeActionResponse { @@ -59,6 +64,7 @@ export interface AgentRuntimeAdapter { ): Promise; submitTurn(request: AgentRuntimeTurnRequest): Promise; interruptTurn(sessionId: string): Promise; + removeQueuedTurn(sessionId: string, queuedTurnId: string): Promise; respondToAction(request: AgentRuntimeActionResponse): Promise; listenToTurnEvents( eventName: string, @@ -105,10 +111,13 @@ export const defaultAgentRuntimeAdapter: AgentRuntimeAdapter = { provider_config: request.providerConfig, execution_strategy: request.executionStrategy, web_search: request.webSearch, + search_mode: request.searchMode, auto_continue: request.autoContinue, system_prompt: request.systemPrompt, metadata: request.metadata, }, + queue_if_busy: request.queueIfBusy, + queued_turn_id: request.queuedTurnId, }); }, async interruptTurn(sessionId) { @@ -116,6 +125,12 @@ export const defaultAgentRuntimeAdapter: AgentRuntimeAdapter = { session_id: sessionId, }); }, + async removeQueuedTurn(sessionId, queuedTurnId) { + return removeAgentRuntimeQueuedTurn({ + session_id: sessionId, + queued_turn_id: queuedTurnId, + }); + }, async respondToAction(request) { await respondAgentRuntimeAction({ session_id: request.sessionId, diff --git a/src/components/agent/chat/hooks/agentSessionScopedStorage.ts b/src/components/agent/chat/hooks/agentSessionScopedStorage.ts new file mode 100644 index 000000000..8b884e225 --- /dev/null +++ b/src/components/agent/chat/hooks/agentSessionScopedStorage.ts @@ -0,0 +1,26 @@ +import { getScopedStorageKey } from "./agentChatShared"; + +export interface AgentSessionScopedKeys { + currentSessionKey: string; + messagesKey: string; + persistedSessionKey: string; + turnsKey: string; + itemsKey: string; + currentTurnKey: string; +} + +export function getAgentSessionScopedKeys( + workspaceId: string, +): AgentSessionScopedKeys { + return { + currentSessionKey: getScopedStorageKey(workspaceId, "aster_curr_sessionId"), + messagesKey: getScopedStorageKey(workspaceId, "aster_messages"), + persistedSessionKey: getScopedStorageKey( + workspaceId, + "aster_last_sessionId", + ), + turnsKey: getScopedStorageKey(workspaceId, "aster_thread_turns"), + itemsKey: getScopedStorageKey(workspaceId, "aster_thread_items"), + currentTurnKey: getScopedStorageKey(workspaceId, "aster_curr_turnId"), + }; +} diff --git a/src/components/agent/chat/hooks/agentStreamEventProcessor.ts b/src/components/agent/chat/hooks/agentStreamEventProcessor.ts index e6c989a32..bd75e047a 100644 --- a/src/components/agent/chat/hooks/agentStreamEventProcessor.ts +++ b/src/components/agent/chat/hooks/agentStreamEventProcessor.ts @@ -9,11 +9,7 @@ import type { StreamEventToolStart, } from "@/lib/api/agentStream"; import type { Artifact } from "@/lib/artifact/types"; -import type { - ActionRequired, - Message, - WriteArtifactContext, -} from "../types"; +import type { ActionRequired, Message, WriteArtifactContext } from "../types"; import { activityLogger } from "@/components/content-creator/utils/activityLogger"; import { isAskToolName, @@ -34,6 +30,7 @@ import { import { buildArtifactFromWrite, extractArtifactPathsFromMetadata, + findMessageArtifact, upsertMessageArtifact, } from "../utils/messageArtifacts"; import type { AgentRuntimeAdapter } from "./agentRuntimeAdapter"; @@ -60,16 +57,202 @@ interface ToolTrackingContext { toolNameByToolId: Map; } -function upsertAssistantArtifact( - messages: Message[], - assistantMsgId: string, - artifact: Artifact, -): Message[] { - return messages.map((message) => - message.id === assistantMsgId - ? upsertMessageArtifact(message, artifact) - : message, +function normalizeToolNameForFileMutation(value: string): string { + return value + .toLowerCase() + .replace(/[^a-z0-9]+/g, "") + .trim(); +} + +function isFileMutationToolName(toolName: string): boolean { + const normalized = normalizeToolNameForFileMutation(toolName); + return [ + "write", + "create", + "save", + "output", + "edit", + "patch", + "update", + "replace", + ].some((keyword) => normalized.includes(keyword)); +} + +function extractPatchPath(rawText?: string): string | undefined { + if (!rawText) { + return undefined; + } + + for (const line of rawText.split(/\r?\n/)) { + const trimmed = line.trim(); + for (const prefix of [ + "*** Add File:", + "*** Update File:", + "*** Delete File:", + "*** Move to:", + ]) { + if (trimmed.startsWith(prefix)) { + const path = trimmed.slice(prefix.length).trim(); + if (path) { + return path.replace(/\\/g, "/"); + } + } + } + } + + return undefined; +} + +function extractPatchText( + toolArgs: Record | null, +): string | undefined { + if (!toolArgs) { + return undefined; + } + + for (const key of ["patch", "command", "cmd", "script"]) { + const value = toolArgs[key]; + if (typeof value === "string" && value.trim()) { + return value; + } + if (Array.isArray(value)) { + const text = value + .filter((item): item is string => typeof item === "string") + .map((item) => item.trim()) + .filter(Boolean) + .join("\n"); + if (text) { + return text; + } + } + } + + return undefined; +} + +function extractToolArgPath( + toolArgs: Record | null, +): string | undefined { + if (!toolArgs) { + return undefined; + } + + for (const key of [ + "path", + "file_path", + "filePath", + "target_path", + "targetPath", + "output_path", + "outputPath", + ]) { + const value = toolArgs[key]; + if (typeof value === "string" && value.trim()) { + return value.trim(); + } + } + + return extractPatchPath(extractPatchText(toolArgs)); +} + +function extractWriteLikeContent( + toolArgs: Record | null, +): string | undefined { + const directContent = extractToolArgContent(toolArgs); + if (directContent !== undefined) { + return directContent; + } + + return undefined; +} + +function extractToolArgContent( + toolArgs: Record | null, +): string | undefined { + if (!toolArgs) { + return undefined; + } + + for (const key of ["content", "text", "contents", "body"]) { + const value = toolArgs[key]; + if (typeof value === "string") { + return value; + } + } + + return undefined; +} + +function buildWriteMetadata( + baseMetadata: Record | undefined, + options: { + source: WriteArtifactContext["source"]; + phase: "preparing" | "streaming" | "persisted" | "completed" | "failed"; + content: string; + isPartial: boolean; + }, +): WriteArtifactContext["metadata"] { + const previewText = options.content.trim() + ? options.content.slice(0, 480).trim() + : undefined; + const latestChunk = options.content.trim() + ? options.content.slice(-240).trim() + : undefined; + + return { + ...(baseMetadata || {}), + writePhase: options.phase, + previewText, + latestChunk, + isPartial: options.isPartial, + lastUpdateSource: options.source, + }; +} + +function upsertAssistantWriteArtifact({ + assistantMsgId, + setMessages, + filePath, + content, + context, +}: { + assistantMsgId: string; + setMessages: Dispatch>; + filePath: string; + content: string; + context: Omit; +}): Artifact | null { + let nextArtifact: Artifact | null = null; + + setMessages((prev) => + prev.map((message) => { + if (message.id !== assistantMsgId) { + return message; + } + + const existingArtifact = findMessageArtifact(message, { + artifactId: context.artifactId, + filePath, + }); + const nextContent = + content.length > 0 || !existingArtifact + ? content + : existingArtifact.content; + nextArtifact = buildArtifactFromWrite({ + filePath, + content: nextContent, + context: { + ...context, + artifact: existingArtifact, + artifactId: existingArtifact?.id || context.artifactId, + }, + }); + + return upsertMessageArtifact(message, nextArtifact); + }), ); + + return nextArtifact; } export function handleToolStartEvent({ @@ -120,42 +303,51 @@ export function handleToolStartEvent({ const toolArgs = parseJsonObject(data.arguments); const toolName = data.tool_name.toLowerCase(); - if (toolName.includes("write") || toolName.includes("create")) { - const filePath = toolArgs?.path || toolArgs?.file_path || toolArgs?.filePath; - const fileContent = toolArgs?.content || toolArgs?.text || ""; - if ( - typeof filePath === "string" && - typeof fileContent === "string" && - filePath && - fileContent - ) { - const nextArtifact = buildArtifactFromWrite({ - filePath, - content: fileContent, - context: { - artifactId: `artifact:${assistantMsgId}:${filePath}`, - source: "tool_start", - sourceMessageId: assistantMsgId, - status: "streaming", - metadata: - toolArgs?.metadata && typeof toolArgs.metadata === "object" - ? (toolArgs.metadata as Record) - : {}, - }, - }); - - setMessages((prev) => - upsertAssistantArtifact(prev, assistantMsgId, nextArtifact), - ); - - onWriteFile?.(fileContent, filePath, { - artifact: nextArtifact, - artifactId: nextArtifact.id, + if (isFileMutationToolName(toolName)) { + const filePath = extractToolArgPath(toolArgs); + const fileContent = extractWriteLikeContent(toolArgs) || ""; + if (filePath) { + const baseMetadata = + toolArgs?.metadata && typeof toolArgs.metadata === "object" + ? (toolArgs.metadata as Record) + : undefined; + const writeContext: WriteArtifactContext = { + artifactId: `artifact:${assistantMsgId}:${filePath}`, source: "tool_start", sourceMessageId: assistantMsgId, - status: nextArtifact.status, - metadata: nextArtifact.meta, + status: "streaming", + metadata: buildWriteMetadata(baseMetadata, { + source: "tool_start", + phase: fileContent.trim() ? "streaming" : "preparing", + content: fileContent, + isPartial: true, + }), + }; + const nextArtifact = upsertAssistantWriteArtifact({ + assistantMsgId, + setMessages, + filePath, + content: fileContent, + context: writeContext, }); + const emittedArtifact = + nextArtifact || + buildArtifactFromWrite({ + filePath, + content: fileContent, + context: writeContext, + }); + + if (emittedArtifact) { + onWriteFile?.(fileContent, filePath, { + artifact: emittedArtifact, + artifactId: emittedArtifact.id, + source: "tool_start", + sourceMessageId: assistantMsgId, + status: emittedArtifact.status, + metadata: emittedArtifact.meta, + }); + } } } @@ -171,7 +363,6 @@ export function handleToolStartEvent({ return { ...message, - runtimeStatus: undefined, toolCalls: [...(message.toolCalls || []), newToolCall], contentParts: [ ...(message.contentParts || []), @@ -195,14 +386,13 @@ export function handleToolStartEvent({ toolArgs?.options || toolArgs?.choices || toolArgs?.enum, ); const explicitRequestId = requestIdFromArgs?.trim(); - const normalizedQuestions = - questionList ?? [ - { - question, - options: askOptions, - multiSelect: false, - }, - ]; + const normalizedQuestions = questionList ?? [ + { + question, + options: askOptions, + multiSelect: false, + }, + ]; const fallbackAction: ActionRequired = { requestId: @@ -237,11 +427,16 @@ export function handleToolEndEvent({ ToolTrackingContext & { data: StreamEventToolEnd; }) { - const normalizedOutput = extractProxycastToolMetadataBlock(data.result?.output); + const normalizedOutput = extractProxycastToolMetadataBlock( + data.result?.output, + ); const normalizedResult = { ...data.result, output: normalizedOutput.text, - images: normalizeToolResultImages(data.result?.images, normalizedOutput.text), + images: normalizeToolResultImages( + data.result?.images, + normalizedOutput.text, + ), metadata: normalizeToolResultMetadata( data.result?.metadata, data.result?.output, @@ -314,43 +509,57 @@ export function handleToolEndEvent({ return { ...message, - runtimeStatus: undefined, toolCalls: updatedToolCalls, contentParts: updatedContentParts, }; }), ); - const artifactPaths = extractArtifactPathsFromMetadata(normalizedResult.metadata); + const artifactPaths = extractArtifactPathsFromMetadata( + normalizedResult.metadata, + ); if (artifactPaths.length === 0) { return; } for (const artifactPath of artifactPaths) { - const nextArtifact = buildArtifactFromWrite({ - filePath: artifactPath, - content: "", - context: { - artifactId: `artifact:${assistantMsgId}:${artifactPath}`, - source: "tool_result", - sourceMessageId: assistantMsgId, - status: isSuccess ? "complete" : "error", - metadata: normalizedResult.metadata, - }, - }); - - setMessages((prev) => - upsertAssistantArtifact(prev, assistantMsgId, nextArtifact), - ); - - onWriteFile?.("", artifactPath, { - artifact: nextArtifact, - artifactId: nextArtifact.id, + const writeContext: WriteArtifactContext = { + artifactId: `artifact:${assistantMsgId}:${artifactPath}`, source: "tool_result", sourceMessageId: assistantMsgId, - status: nextArtifact.status, - metadata: nextArtifact.meta, + status: isSuccess ? "complete" : "error", + metadata: buildWriteMetadata(normalizedResult.metadata, { + source: "tool_result", + phase: isSuccess ? "completed" : "failed", + content: "", + isPartial: false, + }), + }; + const nextArtifact = upsertAssistantWriteArtifact({ + assistantMsgId, + setMessages, + filePath: artifactPath, + content: "", + context: writeContext, }); + const emittedArtifact = + nextArtifact || + buildArtifactFromWrite({ + filePath: artifactPath, + content: "", + context: writeContext, + }); + + if (emittedArtifact) { + onWriteFile?.(emittedArtifact.content, artifactPath, { + artifact: emittedArtifact, + artifactId: emittedArtifact.id, + source: "tool_result", + sourceMessageId: assistantMsgId, + status: emittedArtifact.status, + metadata: emittedArtifact.meta, + }); + } } } @@ -369,30 +578,46 @@ export function handleArtifactSnapshotEvent({ } const metadata = data.artifact.metadata; - const nextArtifact = buildArtifactFromWrite({ - filePath: artifactPath, - content: - typeof data.artifact.content === "string" ? data.artifact.content : "", - context: { - artifactId: - data.artifact.artifactId || `artifact:${assistantMsgId}:${artifactPath}`, - source: "artifact_snapshot", - sourceMessageId: assistantMsgId, - status: metadata?.complete === false ? "streaming" : "complete", - metadata, - }, - }); - - setMessages((prev) => upsertAssistantArtifact(prev, assistantMsgId, nextArtifact)); - - onWriteFile?.(nextArtifact.content, artifactPath, { - artifact: nextArtifact, - artifactId: nextArtifact.id, + const snapshotContent = + typeof data.artifact.content === "string" ? data.artifact.content : ""; + const writeContext: WriteArtifactContext = { + artifactId: + data.artifact.artifactId || `artifact:${assistantMsgId}:${artifactPath}`, source: "artifact_snapshot", sourceMessageId: assistantMsgId, - status: nextArtifact.status, - metadata: nextArtifact.meta, + status: "streaming", + metadata: buildWriteMetadata(metadata, { + source: "artifact_snapshot", + phase: metadata?.complete === false ? "streaming" : "persisted", + content: snapshotContent, + isPartial: metadata?.complete === false, + }), + }; + const nextArtifact = upsertAssistantWriteArtifact({ + assistantMsgId, + setMessages, + filePath: artifactPath, + content: snapshotContent, + context: writeContext, }); + const emittedArtifact = + nextArtifact || + buildArtifactFromWrite({ + filePath: artifactPath, + content: snapshotContent, + context: writeContext, + }); + + if (emittedArtifact) { + onWriteFile?.(emittedArtifact.content, artifactPath, { + artifact: emittedArtifact, + artifactId: emittedArtifact.id, + source: "artifact_snapshot", + sourceMessageId: assistantMsgId, + status: emittedArtifact.status, + metadata: emittedArtifact.meta, + }); + } } export function handleActionRequiredEvent({ @@ -500,7 +725,9 @@ export function handleContextTraceEvent({ } const seen = new Set( - (message.contextTrace || []).map((step) => `${step.stage}::${step.detail}`), + (message.contextTrace || []).map( + (step) => `${step.stage}::${step.detail}`, + ), ); const nextSteps = [...(message.contextTrace || [])]; diff --git a/src/components/agent/chat/hooks/agentStreamRuntimeHandler.ts b/src/components/agent/chat/hooks/agentStreamRuntimeHandler.ts new file mode 100644 index 000000000..5dd7bcb86 --- /dev/null +++ b/src/components/agent/chat/hooks/agentStreamRuntimeHandler.ts @@ -0,0 +1,446 @@ +import { toast } from "sonner"; +import type { Dispatch, MutableRefObject, SetStateAction } from "react"; +import type { + AsterExecutionStrategy, + QueuedTurnSnapshot, +} from "@/lib/api/agentRuntime"; +import type { + AgentThreadItem, + AgentThreadTurn, + StreamEvent, +} from "@/lib/api/agentStream"; +import { activityLogger } from "@/components/content-creator/utils/activityLogger"; +import type { ActionRequired, Message } from "../types"; +import { appendTextToParts } from "./agentChatHistory"; +import { updateMessageArtifactsStatus } from "../utils/messageArtifacts"; +import { WORKSPACE_PATH_AUTO_CREATED_WARNING_CODE } from "./agentChatCoreUtils"; +import { + removeThreadItemState, + removeThreadTurnState, + upsertThreadItemState, + upsertThreadTurnState, +} from "./agentThreadState"; +import { + handleActionRequiredEvent, + handleArtifactSnapshotEvent, + handleContextTraceEvent, + handleToolEndEvent, + handleToolStartEvent, +} from "./agentStreamEventProcessor"; +import type { AgentRuntimeAdapter } from "./agentRuntimeAdapter"; + +type MessageParts = NonNullable; + +interface StreamObserver { + onTextDelta?: (delta: string, accumulated: string) => void; + onComplete?: (content: string) => void; + onError?: (message: string) => void; +} + +interface StreamRequestState { + accumulatedContent: string; + queuedTurnId: string | null; + requestLogId: string | null; + requestStartedAt: number; + requestFinished: boolean; +} + +interface StreamLifecycleCallbacks { + activateStream: () => void; + isStreamActivated: () => boolean; + clearOptimisticItem: () => void; + clearOptimisticTurn: () => void; + disposeListener: () => void; + removeQueuedDraftMessages: () => void; + clearActiveStreamIfMatch: (eventName: string) => boolean; + upsertQueuedTurn: (queuedTurn: QueuedTurnSnapshot) => void; + removeQueuedTurnState: (queuedTurnIds: string[]) => void; + playToolcallSound: () => void; + playTypewriterSound: () => void; + appendThinkingToParts: ( + parts: MessageParts, + textDelta: string, + ) => MessageParts; +} + +interface HandleTurnStreamEventOptions { + data: StreamEvent; + requestState: StreamRequestState; + callbacks: StreamLifecycleCallbacks; + observer?: StreamObserver; + eventName: string; + optimisticTurnId: string; + optimisticItemId: string; + assistantMsgId: string; + activeSessionId: string; + resolvedWorkspaceId: string; + effectiveExecutionStrategy: AsterExecutionStrategy; + runtime: AgentRuntimeAdapter; + warnedKeysRef: MutableRefObject>; + actionLoggedKeys: Set; + toolLogIdByToolId: Map; + toolStartedAtByToolId: Map; + toolNameByToolId: Map; + onWriteFile?: ( + content: string, + fileName: string, + context?: import("../types").WriteArtifactContext, + ) => void; + setMessages: Dispatch>; + setPendingActions: Dispatch>; + setThreadItems: Dispatch>; + setThreadTurns: Dispatch>; + setCurrentTurnId: Dispatch>; +} + +function finishRequestLog( + requestState: StreamRequestState, + payload: { + eventType: "chat_request_complete" | "chat_request_error"; + status: "success" | "error"; + description?: string; + error?: string; + }, +) { + if (!requestState.requestLogId || requestState.requestFinished) { + return; + } + + requestState.requestFinished = true; + activityLogger.updateLog(requestState.requestLogId, { + eventType: payload.eventType, + status: payload.status, + duration: Date.now() - requestState.requestStartedAt, + description: payload.description, + error: payload.error, + }); +} + +export function handleTurnStreamEvent({ + data, + requestState, + callbacks, + observer, + eventName, + optimisticTurnId, + optimisticItemId, + assistantMsgId, + activeSessionId, + resolvedWorkspaceId, + effectiveExecutionStrategy, + runtime, + warnedKeysRef, + actionLoggedKeys, + toolLogIdByToolId, + toolStartedAtByToolId, + toolNameByToolId, + onWriteFile, + setMessages, + setPendingActions, + setThreadItems, + setThreadTurns, + setCurrentTurnId, +}: HandleTurnStreamEventOptions): void { + const { + activateStream, + isStreamActivated, + clearOptimisticItem, + clearOptimisticTurn, + disposeListener, + removeQueuedDraftMessages, + clearActiveStreamIfMatch, + upsertQueuedTurn, + removeQueuedTurnState, + playToolcallSound, + playTypewriterSound, + appendThinkingToParts, + } = callbacks; + + switch (data.type) { + case "thread_started": + break; + + case "queue_added": + requestState.queuedTurnId = data.queued_turn.queued_turn_id; + upsertQueuedTurn(data.queued_turn); + break; + + case "queue_removed": + removeQueuedTurnState([data.queued_turn_id]); + if ( + !isStreamActivated() && + (!requestState.queuedTurnId || + requestState.queuedTurnId === data.queued_turn_id) + ) { + disposeListener(); + removeQueuedDraftMessages(); + } + break; + + case "queue_started": + requestState.queuedTurnId = data.queued_turn_id; + removeQueuedTurnState([data.queued_turn_id]); + activateStream(); + break; + + case "queue_cleared": + removeQueuedTurnState(data.queued_turn_ids); + if ( + !isStreamActivated() && + (!requestState.queuedTurnId || + data.queued_turn_ids.includes(requestState.queuedTurnId)) + ) { + disposeListener(); + removeQueuedDraftMessages(); + } + break; + + case "turn_started": + activateStream(); + setCurrentTurnId(data.turn.id); + setThreadTurns((prev) => + upsertThreadTurnState( + removeThreadTurnState(prev, optimisticTurnId), + data.turn, + ), + ); + clearOptimisticItem(); + break; + + case "item_started": + case "item_updated": + case "item_completed": + activateStream(); + setThreadItems((prev) => + upsertThreadItemState( + removeThreadItemState(prev, optimisticItemId), + data.item, + ), + ); + break; + + case "turn_completed": + case "turn_failed": + activateStream(); + clearOptimisticItem(); + setThreadTurns((prev) => + upsertThreadTurnState( + removeThreadTurnState(prev, optimisticTurnId), + data.turn, + ), + ); + setCurrentTurnId(data.turn.id); + break; + + case "runtime_status": + activateStream(); + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...msg, + runtimeStatus: data.status, + } + : msg, + ), + ); + break; + + case "thinking_delta": + activateStream(); + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...msg, + isThinking: true, + thinkingContent: (msg.thinkingContent || "") + data.text, + contentParts: appendThinkingToParts( + msg.contentParts || [], + data.text, + ), + } + : msg, + ), + ); + break; + + case "text_delta": + activateStream(); + clearOptimisticItem(); + requestState.accumulatedContent += data.text; + observer?.onTextDelta?.(data.text, requestState.accumulatedContent); + playTypewriterSound(); + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...msg, + content: requestState.accumulatedContent, + thinkingContent: undefined, + contentParts: appendTextToParts(msg.contentParts || [], data.text), + } + : msg, + ), + ); + break; + + case "tool_start": + activateStream(); + clearOptimisticItem(); + playToolcallSound(); + handleToolStartEvent({ + data, + setPendingActions, + onWriteFile, + toolLogIdByToolId, + toolStartedAtByToolId, + toolNameByToolId, + assistantMsgId, + activeSessionId, + resolvedWorkspaceId, + setMessages, + }); + break; + + case "tool_end": + activateStream(); + clearOptimisticItem(); + handleToolEndEvent({ + data, + onWriteFile, + toolLogIdByToolId, + toolStartedAtByToolId, + toolNameByToolId, + assistantMsgId, + activeSessionId, + resolvedWorkspaceId, + setMessages, + }); + break; + + case "artifact_snapshot": + activateStream(); + clearOptimisticItem(); + handleArtifactSnapshotEvent({ + data, + onWriteFile, + assistantMsgId, + activeSessionId, + resolvedWorkspaceId, + setMessages, + }); + break; + + case "action_required": + activateStream(); + clearOptimisticItem(); + handleActionRequiredEvent({ + data, + actionLoggedKeys, + effectiveExecutionStrategy, + runtime, + setPendingActions, + assistantMsgId, + activeSessionId, + resolvedWorkspaceId, + setMessages, + }); + break; + + case "context_trace": + activateStream(); + clearOptimisticItem(); + handleContextTraceEvent({ + data, + assistantMsgId, + activeSessionId, + resolvedWorkspaceId, + setMessages, + }); + break; + + case "final_done": { + clearOptimisticItem(); + clearOptimisticTurn(); + removeQueuedTurnState(requestState.queuedTurnId ? [requestState.queuedTurnId] : []); + finishRequestLog(requestState, { + eventType: "chat_request_complete", + status: "success", + description: `请求完成,工具调用 ${toolLogIdByToolId.size} 次`, + }); + const finalContent = + requestState.accumulatedContent.trim() || + "已完成工具执行,但模型未输出最终答复,请重试。"; + if (!requestState.accumulatedContent.trim()) { + toast.error("已完成工具执行,但模型未输出最终答复,请重试"); + } + observer?.onComplete?.(finalContent); + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...updateMessageArtifactsStatus(msg, "complete"), + isThinking: false, + content: finalContent, + runtimeStatus: undefined, + } + : msg, + ), + ); + clearActiveStreamIfMatch(eventName); + disposeListener(); + break; + } + + case "error": + clearOptimisticItem(); + clearOptimisticTurn(); + removeQueuedTurnState(requestState.queuedTurnId ? [requestState.queuedTurnId] : []); + finishRequestLog(requestState, { + eventType: "chat_request_error", + status: "error", + error: data.message, + }); + observer?.onError?.(data.message); + if ( + data.message.includes("429") || + data.message.toLowerCase().includes("rate limit") + ) { + toast.warning("请求过于频繁,请稍后重试"); + } else { + toast.error(`响应错误: ${data.message}`); + } + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...updateMessageArtifactsStatus(msg, "error"), + isThinking: false, + content: + requestState.accumulatedContent || `错误: ${data.message}`, + runtimeStatus: undefined, + } + : msg, + ), + ); + clearActiveStreamIfMatch(eventName); + disposeListener(); + break; + + case "warning": { + if (data.code === WORKSPACE_PATH_AUTO_CREATED_WARNING_CODE) { + break; + } + const warningKey = `${activeSessionId}:${data.code || data.message}`; + if (!warnedKeysRef.current.has(warningKey)) { + warnedKeysRef.current.add(warningKey); + toast.warning(data.message); + } + break; + } + + default: + break; + } +} diff --git a/src/components/agent/chat/hooks/index.ts b/src/components/agent/chat/hooks/index.ts index a031a5934..33fd734ad 100644 --- a/src/components/agent/chat/hooks/index.ts +++ b/src/components/agent/chat/hooks/index.ts @@ -5,6 +5,7 @@ */ import { useAsterAgentChat } from "./useAsterAgentChat"; +export { useArtifactAutoPreviewSync } from "./useArtifactAutoPreviewSync"; export type { Topic } from "./useAgentChat"; diff --git a/src/components/agent/chat/hooks/skillCommand.test.ts b/src/components/agent/chat/hooks/skillCommand.test.ts index a6a870e2e..186c35100 100644 --- a/src/components/agent/chat/hooks/skillCommand.test.ts +++ b/src/components/agent/chat/hooks/skillCommand.test.ts @@ -213,6 +213,78 @@ describe("tryExecuteSlashSkillCommand 社媒主链路", () => { expect(onWriteFile).not.toHaveBeenCalled(); }); + it("收到 artifact_snapshot 时应立刻透传给 onWriteFile", async () => { + const store = createMessageStore([buildBaseMessage()]); + const onWriteFile = vi.fn(); + let streamHandler: ((event: { payload: unknown }) => void) | null = null; + + mockSafeListen.mockImplementation(async (_eventName, handler) => { + streamHandler = handler as (event: { payload: unknown }) => void; + return () => { + streamHandler = null; + }; + }); + + mockExecuteSkill.mockImplementation(async () => { + streamHandler?.({ + payload: { + type: "artifact_snapshot", + artifact: { + artifactId: "artifact-1", + filePath: "social-posts/live.md", + content: "# 实时稿", + metadata: { + complete: false, + writePhase: "streaming", + lastUpdateSource: "artifact_snapshot", + }, + }, + }, + }); + streamHandler?.({ payload: { type: "final_done" } }); + + return { + success: true, + output: + '\n# 实时稿\n', + steps_completed: [], + }; + }); + + const handled = await tryExecuteSlashSkillCommand({ + command: { + skillName: "social_post_with_cover", + userInput: "实时写作", + }, + rawContent: "/social_post_with_cover 实时写作", + assistantMsgId: "assistant-1", + providerType: "anthropic", + model: "claude-sonnet-4-20250514", + ensureSession: async () => "session-1", + setMessages: store.setMessages, + setIsSending: vi.fn(), + setCurrentAssistantMsgId: vi.fn(), + setStreamUnlisten: vi.fn(), + setActiveSessionIdForStop: vi.fn(), + isExecutionCancelled: () => false, + playTypewriterSound: vi.fn(), + playToolcallSound: vi.fn(), + onWriteFile, + }); + + expect(handled).toBe(true); + expect(onWriteFile).toHaveBeenCalledTimes(1); + expect(onWriteFile).toHaveBeenCalledWith( + "# 实时稿", + "social-posts/live.md", + expect.objectContaining({ + artifactId: "artifact-1", + source: "artifact_snapshot", + status: "streaming", + }), + ); + }); + it("当社媒结果无 write_file 时应走前端兜底写入", async () => { const store = createMessageStore([buildBaseMessage()]); const onWriteFile = vi.fn(); diff --git a/src/components/agent/chat/hooks/skillCommand.ts b/src/components/agent/chat/hooks/skillCommand.ts index c532982e6..cc50e4d23 100644 --- a/src/components/agent/chat/hooks/skillCommand.ts +++ b/src/components/agent/chat/hooks/skillCommand.ts @@ -228,6 +228,20 @@ function tryHandleToolWriteFile( } } +function resolveSnapshotStatus( + metadata: Record | undefined, +): WriteArtifactContext["status"] { + const writePhase = + typeof metadata?.writePhase === "string" ? metadata.writePhase : undefined; + if (writePhase === "failed") { + return "error"; + } + if (metadata?.complete === false) { + return "streaming"; + } + return "complete"; +} + interface MatchedSkillResult { matchedSkill: ExecutableSkillInfo | null; catalogLoadFailed: boolean; @@ -354,6 +368,7 @@ export async function tryExecuteSlashSkillCommand( thinking_delta: 0, tool_start: 0, tool_end: 0, + artifact_snapshot: 0, done: 0, final_done: 0, error: 0, @@ -510,6 +525,28 @@ export async function tryExecuteSlashSkillCommand( ); break; } + case "artifact_snapshot": { + streamCounters.artifact_snapshot += 1; + const filePath = streamEvent.artifact.filePath; + if (!filePath) { + break; + } + + onWriteFile?.( + typeof streamEvent.artifact.content === "string" + ? streamEvent.artifact.content + : "", + filePath, + { + artifactId: streamEvent.artifact.artifactId, + source: "artifact_snapshot", + sourceMessageId: assistantMsgId, + status: resolveSnapshotStatus(streamEvent.artifact.metadata), + metadata: streamEvent.artifact.metadata, + }, + ); + break; + } case "action_required": { const actionRequired: ActionRequired = { requestId: streamEvent.request_id, diff --git a/src/components/agent/chat/hooks/useAgentSession.ts b/src/components/agent/chat/hooks/useAgentSession.ts index a31edc030..9fbc95d62 100644 --- a/src/components/agent/chat/hooks/useAgentSession.ts +++ b/src/components/agent/chat/hooks/useAgentSession.ts @@ -1,19 +1,23 @@ import { useCallback, useEffect, + useMemo, useRef, useState, type MutableRefObject, } from "react"; import { toast } from "sonner"; -import type { AsterExecutionStrategy } from "@/lib/api/agentRuntime"; +import type { + AsterExecutionStrategy, + QueuedTurnSnapshot, +} from "@/lib/api/agentRuntime"; +import { normalizeQueuedTurnSnapshots } from "@/lib/api/queuedTurn"; import { isAsterSessionNotFoundError, resolveRestorableSessionId, } from "@/lib/asterSessionRecovery"; import type { AgentThreadItem, AgentThreadTurn, Message } from "../types"; import { - getScopedStorageKey, mapSessionToTopic, type ClearMessagesOptions, type SessionModelPreference, @@ -23,6 +27,7 @@ import { hydrateSessionDetailMessages, normalizeHistoryMessages, } from "./agentChatHistory"; +import { getAgentSessionScopedKeys } from "./agentSessionScopedStorage"; import { loadPersisted, loadPersistedString, @@ -80,29 +85,8 @@ export function useAgentSession(options: UseAgentSessionOptions) { filterSessionsByWorkspace, setExecutionStrategyState, } = options; - - const getScopedSessionKey = useCallback( - () => getScopedStorageKey(workspaceId, "aster_curr_sessionId"), - [workspaceId], - ); - const getScopedMessagesKey = useCallback( - () => getScopedStorageKey(workspaceId, "aster_messages"), - [workspaceId], - ); - const getScopedPersistedSessionKey = useCallback( - () => getScopedStorageKey(workspaceId, "aster_last_sessionId"), - [workspaceId], - ); - const getScopedTurnsKey = useCallback( - () => getScopedStorageKey(workspaceId, "aster_thread_turns"), - [workspaceId], - ); - const getScopedItemsKey = useCallback( - () => getScopedStorageKey(workspaceId, "aster_thread_items"), - [workspaceId], - ); - const getScopedCurrentTurnKey = useCallback( - () => getScopedStorageKey(workspaceId, "aster_curr_turnId"), + const scopedKeys = useMemo( + () => getAgentSessionScopedKeys(workspaceId), [workspaceId], ); @@ -153,6 +137,7 @@ export function useAgentSession(options: UseAgentSessionOptions) { ) : null, ); + const [queuedTurns, setQueuedTurns] = useState([]); const [topics, setTopics] = useState([]); const [topicsReady, setTopicsReady] = useState(false); @@ -180,8 +165,8 @@ export function useAgentSession(options: UseAgentSessionOptions) { return; } - const scopedSessionKey = getScopedSessionKey(); - const scopedPersistedSessionKey = getScopedPersistedSessionKey(); + const scopedSessionKey = scopedKeys.currentSessionKey; + const scopedPersistedSessionKey = scopedKeys.persistedSessionKey; saveTransient(scopedSessionKey, sessionId); savePersisted(scopedPersistedSessionKey, sessionId); @@ -204,40 +189,35 @@ export function useAgentSession(options: UseAgentSessionOptions) { savePersisted(sessionWorkspaceKey, resolvedWorkspaceId); } } - }, [ - getScopedPersistedSessionKey, - getScopedSessionKey, - sessionId, - workspaceId, - ]); + }, [scopedKeys, sessionId, workspaceId]); useEffect(() => { if (!workspaceId?.trim()) { return; } - saveTransient(getScopedMessagesKey(), messages); - }, [getScopedMessagesKey, messages, workspaceId]); + saveTransient(scopedKeys.messagesKey, messages); + }, [messages, scopedKeys, workspaceId]); useEffect(() => { if (!workspaceId?.trim()) { return; } - saveTransient(getScopedTurnsKey(), threadTurns); - }, [getScopedTurnsKey, threadTurns, workspaceId]); + saveTransient(scopedKeys.turnsKey, threadTurns); + }, [scopedKeys, threadTurns, workspaceId]); useEffect(() => { if (!workspaceId?.trim()) { return; } - saveTransient(getScopedItemsKey(), threadItems); - }, [getScopedItemsKey, threadItems, workspaceId]); + saveTransient(scopedKeys.itemsKey, threadItems); + }, [scopedKeys, threadItems, workspaceId]); useEffect(() => { if (!workspaceId?.trim()) { return; } - saveTransient(getScopedCurrentTurnKey(), currentTurnId); - }, [currentTurnId, getScopedCurrentTurnKey, workspaceId]); + saveTransient(scopedKeys.currentTurnKey, currentTurnId); + }, [currentTurnId, scopedKeys, workspaceId]); useEffect(() => { if (!workspaceId?.trim()) { @@ -246,6 +226,7 @@ export function useAgentSession(options: UseAgentSessionOptions) { setThreadTurns([]); setThreadItems([]); setCurrentTurnId(null); + setQueuedTurns([]); resetPendingActions(); resetStreamingRefs(); restoredWorkspaceRef.current = null; @@ -255,14 +236,20 @@ export function useAgentSession(options: UseAgentSessionOptions) { } const scopedSessionId = - loadTransient(getScopedSessionKey(), null) ?? - loadPersisted(getScopedPersistedSessionKey(), null); + loadTransient(scopedKeys.currentSessionKey, null) ?? + loadPersisted(scopedKeys.persistedSessionKey, null); - const scopedMessages = loadTransient(getScopedMessagesKey(), []); - const scopedTurns = loadTransient(getScopedTurnsKey(), []); - const scopedItems = loadTransient(getScopedItemsKey(), []); + const scopedMessages = loadTransient(scopedKeys.messagesKey, []); + const scopedTurns = loadTransient( + scopedKeys.turnsKey, + [], + ); + const scopedItems = loadTransient( + scopedKeys.itemsKey, + [], + ); const scopedCurrentTurnId = loadTransient( - getScopedCurrentTurnKey(), + scopedKeys.currentTurnKey, null, ); @@ -271,20 +258,16 @@ export function useAgentSession(options: UseAgentSessionOptions) { setThreadTurns(scopedTurns); setThreadItems(scopedItems); setCurrentTurnId(scopedCurrentTurnId); + setQueuedTurns([]); resetPendingActions(); resetStreamingRefs(); restoredWorkspaceRef.current = null; hydratedSessionRef.current = null; skipAutoRestoreRef.current = false; }, [ - getScopedMessagesKey, - getScopedItemsKey, - getScopedTurnsKey, - getScopedCurrentTurnKey, - getScopedPersistedSessionKey, - getScopedSessionKey, resetPendingActions, resetStreamingRefs, + scopedKeys, workspaceId, ]); @@ -369,6 +352,7 @@ export function useAgentSession(options: UseAgentSessionOptions) { setThreadTurns([]); setThreadItems([]); setCurrentTurnId(null); + setQueuedTurns([]); setTopics((prev) => [ { id: newSessionId, @@ -391,12 +375,12 @@ export function useAgentSession(options: UseAgentSessionOptions) { providerTypeRef.current, modelRef.current, ); - saveTransient(getScopedSessionKey(), newSessionId); - savePersisted(getScopedPersistedSessionKey(), newSessionId); - saveTransient(getScopedMessagesKey(), []); - saveTransient(getScopedTurnsKey(), []); - saveTransient(getScopedItemsKey(), []); - saveTransient(getScopedCurrentTurnKey(), null); + saveTransient(scopedKeys.currentSessionKey, newSessionId); + savePersisted(scopedKeys.persistedSessionKey, newSessionId); + saveTransient(scopedKeys.messagesKey, []); + saveTransient(scopedKeys.turnsKey, []); + saveTransient(scopedKeys.itemsKey, []); + saveTransient(scopedKeys.currentTurnKey, null); void loadTopics(); return newSessionId; @@ -408,12 +392,6 @@ export function useAgentSession(options: UseAgentSessionOptions) { }, [ executionStrategy, - getScopedMessagesKey, - getScopedItemsKey, - getScopedTurnsKey, - getScopedCurrentTurnKey, - getScopedPersistedSessionKey, - getScopedSessionKey, loadTopics, modelRef, persistSessionModelPreference, @@ -421,6 +399,7 @@ export function useAgentSession(options: UseAgentSessionOptions) { resetPendingActions, resetStreamingRefs, runtime, + scopedKeys, workspaceId, ], ); @@ -434,14 +413,15 @@ export function useAgentSession(options: UseAgentSessionOptions) { (options: ClearMessagesOptions = {}) => { const { showToast = true, toastMessage = "新话题已创建" } = options; - const scopedSessionKey = getScopedSessionKey(); - const scopedPersistedSessionKey = getScopedPersistedSessionKey(); - const scopedMessagesKey = getScopedMessagesKey(); + const scopedSessionKey = scopedKeys.currentSessionKey; + const scopedPersistedSessionKey = scopedKeys.persistedSessionKey; + const scopedMessagesKey = scopedKeys.messagesKey; setMessages([]); setThreadTurns([]); setThreadItems([]); setCurrentTurnId(null); + setQueuedTurns([]); setSessionId(null); resetPendingActions(); restoredWorkspaceRef.current = null; @@ -452,23 +432,18 @@ export function useAgentSession(options: UseAgentSessionOptions) { saveTransient(scopedSessionKey, null); savePersisted(scopedPersistedSessionKey, null); saveTransient(scopedMessagesKey, []); - saveTransient(getScopedTurnsKey(), []); - saveTransient(getScopedItemsKey(), []); - saveTransient(getScopedCurrentTurnKey(), null); + saveTransient(scopedKeys.turnsKey, []); + saveTransient(scopedKeys.itemsKey, []); + saveTransient(scopedKeys.currentTurnKey, null); if (showToast) { toast.success(toastMessage); } }, [ - getScopedMessagesKey, - getScopedItemsKey, - getScopedTurnsKey, - getScopedCurrentTurnKey, - getScopedPersistedSessionKey, - getScopedSessionKey, resetPendingActions, resetStreamingRefs, + scopedKeys, ], ); @@ -484,6 +459,36 @@ export function useAgentSession(options: UseAgentSessionOptions) { ); }, []); + const applySessionDetail = useCallback( + ( + topicId: string, + detail: Awaited>, + options?: { syncSessionId?: boolean }, + ) => { + setMessages(hydrateSessionDetailMessages(detail, topicId)); + setThreadTurns(detail.turns || []); + setThreadItems(detail.items || []); + setQueuedTurns(normalizeQueuedTurnSnapshots(detail.queued_turns)); + setCurrentTurnId( + detail.turns && detail.turns.length > 0 + ? detail.turns[detail.turns.length - 1]?.id || null + : null, + ); + + const selectedTopic = topics.find((topic) => topic.id === topicId); + setExecutionStrategyState( + normalizeExecutionStrategy( + detail.execution_strategy || selectedTopic?.executionStrategy, + ), + ); + + if (options?.syncSessionId) { + setSessionId(topicId); + } + }, + [setExecutionStrategyState, topics], + ); + const switchTopic = useCallback( async (topicId: string) => { if (topicId === sessionIdRef.current && messages.length > 0) return; @@ -502,21 +507,7 @@ export function useAgentSession(options: UseAgentSessionOptions) { const detail = await runtime.getSession(topicId); const topicPreference = loadSessionModelPreference(topicId); - setMessages(hydrateSessionDetailMessages(detail, topicId)); - setThreadTurns(detail.turns || []); - setThreadItems(detail.items || []); - setCurrentTurnId( - detail.turns && detail.turns.length > 0 - ? detail.turns[detail.turns.length - 1]?.id || null - : null, - ); - const selectedTopic = topics.find((topic) => topic.id === topicId); - setExecutionStrategyState( - normalizeExecutionStrategy( - detail.execution_strategy || selectedTopic?.executionStrategy, - ), - ); - setSessionId(topicId); + applySessionDetail(topicId, detail, { syncSessionId: true }); if (topicPreference) { applySessionModelPreference(topicId, topicPreference); @@ -529,9 +520,10 @@ export function useAgentSession(options: UseAgentSessionOptions) { setThreadTurns([]); setThreadItems([]); setCurrentTurnId(null); + setQueuedTurns([]); setSessionId(null); - saveTransient(getScopedSessionKey(), null); - savePersisted(getScopedPersistedSessionKey(), null); + saveTransient(scopedKeys.currentSessionKey, null); + savePersisted(scopedKeys.persistedSessionKey, null); void loadTopics(); return; } @@ -539,18 +531,18 @@ export function useAgentSession(options: UseAgentSessionOptions) { setThreadTurns([]); setThreadItems([]); setCurrentTurnId(null); + setQueuedTurns([]); setSessionId(null); - saveTransient(getScopedSessionKey(), null); - savePersisted(getScopedPersistedSessionKey(), null); + saveTransient(scopedKeys.currentSessionKey, null); + savePersisted(scopedKeys.persistedSessionKey, null); toast.error( `加载对话历史失败: ${error instanceof Error ? error.message : String(error)}`, ); } }, [ + applySessionDetail, applySessionModelPreference, - getScopedPersistedSessionKey, - getScopedSessionKey, loadSessionModelPreference, loadTopics, messages.length, @@ -558,12 +550,33 @@ export function useAgentSession(options: UseAgentSessionOptions) { persistSessionModelPreference, providerTypeRef, runtime, - setExecutionStrategyState, + scopedKeys, sessionIdRef, - topics, ], ); + const refreshSessionDetail = useCallback( + async (targetSessionId?: string) => { + const resolvedSessionId = targetSessionId || sessionIdRef.current; + if (!resolvedSessionId?.trim()) { + return false; + } + + try { + const detail = await runtime.getSession(resolvedSessionId); + if (sessionIdRef.current !== resolvedSessionId) { + return false; + } + applySessionDetail(resolvedSessionId, detail); + return true; + } catch (error) { + console.warn("[AsterChat] 刷新会话详情失败:", error); + return false; + } + }, + [applySessionDetail, runtime, sessionIdRef], + ); + useEffect(() => { const resolvedWorkspaceId = workspaceId?.trim(); if (!resolvedWorkspaceId) return; @@ -577,8 +590,8 @@ export function useAgentSession(options: UseAgentSessionOptions) { restoredWorkspaceRef.current = resolvedWorkspaceId; const scopedCandidate = - loadTransient(getScopedSessionKey(), null) || - loadPersisted(getScopedPersistedSessionKey(), null); + loadTransient(scopedKeys.currentSessionKey, null) || + loadPersisted(scopedKeys.persistedSessionKey, null); const targetSessionId = resolveRestorableSessionId({ candidateSessionId: scopedCandidate, sessions: topics.map((topic) => ({ @@ -593,14 +606,13 @@ export function useAgentSession(options: UseAgentSessionOptions) { switchTopic(targetSessionId).catch((error) => { console.warn("[AsterChat] 自动恢复会话失败:", error); - saveTransient(getScopedSessionKey(), null); - savePersisted(getScopedPersistedSessionKey(), null); + saveTransient(scopedKeys.currentSessionKey, null); + savePersisted(scopedKeys.persistedSessionKey, null); }); }, [ - getScopedPersistedSessionKey, - getScopedSessionKey, isInitialized, sessionId, + scopedKeys, switchTopic, topics, topicsReady, @@ -623,13 +635,17 @@ export function useAgentSession(options: UseAgentSessionOptions) { setThreadTurns([]); setThreadItems([]); setCurrentTurnId(null); - saveTransient(getScopedSessionKey(), null); - savePersisted(getScopedPersistedSessionKey(), null); + setQueuedTurns([]); + saveTransient(scopedKeys.currentSessionKey, null); + savePersisted(scopedKeys.persistedSessionKey, null); hydratedSessionRef.current = null; return; } - if (messages.length > 0 && (threadTurns.length > 0 || threadItems.length > 0)) { + if ( + messages.length > 0 && + (threadTurns.length > 0 || threadItems.length > 0) + ) { hydratedSessionRef.current = sessionId; return; } @@ -645,9 +661,8 @@ export function useAgentSession(options: UseAgentSessionOptions) { hydratedSessionRef.current = null; }); }, [ - getScopedPersistedSessionKey, - getScopedSessionKey, messages.length, + scopedKeys, sessionId, switchTopic, threadItems.length, @@ -668,15 +683,16 @@ export function useAgentSession(options: UseAgentSessionOptions) { setThreadTurns([]); setThreadItems([]); setCurrentTurnId(null); + setQueuedTurns([]); resetPendingActions(); resetStreamingRefs(); hydratedSessionRef.current = null; restoredWorkspaceRef.current = null; - saveTransient(getScopedSessionKey(), null); - savePersisted(getScopedPersistedSessionKey(), null); - saveTransient(getScopedTurnsKey(), []); - saveTransient(getScopedItemsKey(), []); - saveTransient(getScopedCurrentTurnKey(), null); + saveTransient(scopedKeys.currentSessionKey, null); + savePersisted(scopedKeys.persistedSessionKey, null); + saveTransient(scopedKeys.turnsKey, []); + saveTransient(scopedKeys.itemsKey, []); + saveTransient(scopedKeys.currentTurnKey, null); } toast.success("话题已删除"); @@ -686,15 +702,11 @@ export function useAgentSession(options: UseAgentSessionOptions) { } }, [ - getScopedPersistedSessionKey, - getScopedSessionKey, - getScopedItemsKey, - getScopedTurnsKey, - getScopedCurrentTurnKey, loadTopics, resetPendingActions, resetStreamingRefs, runtime, + scopedKeys, sessionIdRef, ], ); @@ -719,7 +731,10 @@ export function useAgentSession(options: UseAgentSessionOptions) { ); const updateTopicExecutionStrategy = useCallback( - (targetSessionId: string, nextExecutionStrategy: AsterExecutionStrategy) => { + ( + targetSessionId: string, + nextExecutionStrategy: AsterExecutionStrategy, + ) => { setTopics((prev) => prev.map((topic) => topic.id === targetSessionId @@ -742,6 +757,8 @@ export function useAgentSession(options: UseAgentSessionOptions) { setThreadItems, currentTurnId, setCurrentTurnId, + queuedTurns, + setQueuedTurns, topics, setTopics, topicsReady, @@ -751,6 +768,7 @@ export function useAgentSession(options: UseAgentSessionOptions) { switchTopic, deleteTopic, renameTopic, + refreshSessionDetail, clearMessages, deleteMessage, editMessage, diff --git a/src/components/agent/chat/hooks/useAgentStream.ts b/src/components/agent/chat/hooks/useAgentStream.ts index ea13f5e15..8746d60cd 100644 --- a/src/components/agent/chat/hooks/useAgentStream.ts +++ b/src/components/agent/chat/hooks/useAgentStream.ts @@ -11,18 +11,14 @@ import { toast } from "sonner"; import type { AsterExecutionStrategy, AutoContinueRequestPayload, + QueuedTurnSnapshot, } from "@/lib/api/agentRuntime"; import { parseStreamEvent, type AgentThreadItem, type AgentThreadTurn, - type StreamEvent, } from "@/lib/api/agentStream"; -import type { - ActionRequired, - Message, - MessageImage, -} from "../types"; +import type { ActionRequired, Message, MessageImage } from "../types"; import { activityLogger } from "@/components/content-creator/utils/activityLogger"; import { parseSkillSlashCommand, @@ -31,13 +27,9 @@ import { import { isWorkspacePathErrorMessage, mapProviderName, - WORKSPACE_PATH_AUTO_CREATED_WARNING_CODE, } from "./agentChatCoreUtils"; -import { appendTextToParts } from "./agentChatHistory"; import { playToolcallSound, playTypewriterSound } from "./agentChatStorage"; -import { - updateMessageArtifactsStatus, -} from "../utils/messageArtifacts"; +import { updateMessageArtifactsStatus } from "../utils/messageArtifacts"; import type { SendMessageOptions, WorkspacePathMissingState, @@ -53,13 +45,39 @@ import { buildInitialAgentRuntimeStatus, buildWaitingAgentRuntimeStatus, } from "../utils/agentRuntimeStatus"; -import { - handleActionRequiredEvent, - handleArtifactSnapshotEvent, - handleContextTraceEvent, - handleToolEndEvent, - handleToolStartEvent, -} from "./agentStreamEventProcessor"; +import { handleTurnStreamEvent } from "./agentStreamRuntimeHandler"; + +function buildQueuedMessagePreview(content: string): string { + const compact = content.split(/\s+/).filter(Boolean).join(" "); + if (!compact) { + return "空白输入"; + } + + const preview = Array.from(compact).slice(0, 80).join(""); + return compact.length > preview.length ? `${preview}...` : preview; +} + +function appendThinkingToParts( + parts: NonNullable, + textDelta: string, +): NonNullable { + const nextParts = [...parts]; + const lastPart = nextParts[nextParts.length - 1]; + + if (lastPart?.type === "thinking") { + nextParts[nextParts.length - 1] = { + type: "thinking", + text: lastPart.text + textDelta, + }; + return nextParts; + } + + nextParts.push({ + type: "thinking", + text: textDelta, + }); + return nextParts; +} interface UseAgentStreamOptions { runtime: AgentRuntimeAdapter; @@ -85,6 +103,8 @@ interface UseAgentStreamOptions { setThreadItems: Dispatch>; setThreadTurns: Dispatch>; setCurrentTurnId: Dispatch>; + queuedTurns: QueuedTurnSnapshot[]; + setQueuedTurns: Dispatch>; setPendingActions: Dispatch>; } @@ -107,21 +127,78 @@ export function useAgentStream(options: UseAgentStreamOptions) { setThreadItems, setThreadTurns, setCurrentTurnId, + queuedTurns, + setQueuedTurns, setPendingActions, } = options; const [isSending, setIsSending] = useState(false); - const unlistenRef = useRef<(() => void) | null>(null); + const listenerMapRef = useRef(new Map void>()); + const activeStreamRef = useRef<{ + assistantMsgId: string; + eventName: string; + sessionId: string; + } | null>(null); useEffect(() => { + const listenerMap = listenerMapRef.current; return () => { - if (unlistenRef.current) { - unlistenRef.current(); - unlistenRef.current = null; + for (const unlisten of listenerMap.values()) { + unlisten(); } + listenerMap.clear(); }; }, []); + const setActiveStream = useCallback( + ( + nextActive: { + assistantMsgId: string; + eventName: string; + sessionId: string; + } | null, + ) => { + activeStreamRef.current = nextActive; + currentAssistantMsgIdRef.current = nextActive?.assistantMsgId ?? null; + currentStreamingSessionIdRef.current = nextActive?.sessionId ?? null; + setIsSending(Boolean(nextActive)); + }, + [currentAssistantMsgIdRef, currentStreamingSessionIdRef], + ); + + const clearActiveStreamIfMatch = useCallback( + (eventName: string) => { + if (activeStreamRef.current?.eventName !== eventName) { + return false; + } + setActiveStream(null); + return true; + }, + [setActiveStream], + ); + + const buildQueuedRuntimeStatus = useCallback( + ( + currentExecutionStrategy: AsterExecutionStrategy, + content: string, + webSearch?: boolean, + ) => ({ + phase: "routing" as const, + title: "已加入排队列表", + detail: `当前会话仍在执行中,本条消息会在前一条完成后自动开始。待处理内容:${buildQueuedMessagePreview(content)}`, + checkpoints: [ + "已创建待执行回合", + webSearch ? "联网搜索能力待命" : "直接回答优先", + currentExecutionStrategy === "code_orchestrated" + ? "代码编排待命" + : currentExecutionStrategy === "react" + ? "对话执行待命" + : "自动路由待命", + ], + }), + [], + ); + const buildRuntimeStatusSummary = useCallback( (status?: Message["runtimeStatus"]): string => { if (!status?.title) { @@ -157,8 +234,11 @@ export function useAgentStream(options: UseAgentStreamOptions) { const observer = options?.observer; const requestMetadata = options?.requestMetadata; const messagePurpose = options?.purpose; + const expectingQueue = + Boolean(activeStreamRef.current) || queuedTurns.length > 0; const assistantMsgId = crypto.randomUUID(); + const userMsgId = skipUserMessage ? null : crypto.randomUUID(); const assistantMsg: Message = { id: assistantMsgId, role: "assistant", @@ -166,12 +246,18 @@ export function useAgentStream(options: UseAgentStreamOptions) { timestamp: new Date(), isThinking: true, contentParts: [], - runtimeStatus: buildInitialAgentRuntimeStatus({ - executionStrategy: effectiveExecutionStrategy, - webSearch, - thinking: _thinking, - skipUserMessage, - }), + runtimeStatus: expectingQueue + ? buildQueuedRuntimeStatus( + effectiveExecutionStrategy, + content, + webSearch, + ) + : buildInitialAgentRuntimeStatus({ + executionStrategy: effectiveExecutionStrategy, + webSearch, + thinking: _thinking, + skipUserMessage, + }), purpose: messagePurpose, }; @@ -179,7 +265,7 @@ export function useAgentStream(options: UseAgentStreamOptions) { setMessages((prev) => [...prev, assistantMsg]); } else { const userMsg: Message = { - id: crypto.randomUUID(), + id: userMsgId as string, role: "user", content, images: images.length > 0 ? images : undefined, @@ -188,12 +274,20 @@ export function useAgentStream(options: UseAgentStreamOptions) { }; setMessages((prev) => [...prev, userMsg, assistantMsg]); } - setIsSending(true); - currentAssistantMsgIdRef.current = assistantMsgId; - if (!skipUserMessage) { + if (!expectingQueue) { + setIsSending(true); + } + + if (!skipUserMessage && !expectingQueue) { const parsedSkillCommand = parseSkillSlashCommand(content); if (parsedSkillCommand) { + const skillEventName = `skill-exec-${assistantMsgId}`; + setActiveStream({ + assistantMsgId, + eventName: skillEventName, + sessionId: sessionIdRef.current || "", + }); const skillHandled = await tryExecuteSlashSkillCommand({ command: parsedSkillCommand, rawContent: content, @@ -204,16 +298,43 @@ export function useAgentStream(options: UseAgentStreamOptions) { setMessages, setIsSending, setCurrentAssistantMsgId: (id) => { - currentAssistantMsgIdRef.current = id; + if (!id) { + clearActiveStreamIfMatch(skillEventName); + return; + } + setActiveStream({ + assistantMsgId: id, + eventName: skillEventName, + sessionId: + activeStreamRef.current?.sessionId || + sessionIdRef.current || + "", + }); }, setStreamUnlisten: (unlistenFn) => { - unlistenRef.current = unlistenFn; + const previous = listenerMapRef.current.get(skillEventName); + if (previous) { + previous(); + listenerMapRef.current.delete(skillEventName); + } + if (unlistenFn) { + listenerMapRef.current.set(skillEventName, unlistenFn); + } }, setActiveSessionIdForStop: (sessionIdForStop) => { - currentStreamingSessionIdRef.current = sessionIdForStop; + if (!sessionIdForStop) { + clearActiveStreamIfMatch(skillEventName); + return; + } + setActiveStream({ + assistantMsgId: + activeStreamRef.current?.assistantMsgId || assistantMsgId, + eventName: skillEventName, + sessionId: sessionIdForStop, + }); }, isExecutionCancelled: () => - currentAssistantMsgIdRef.current !== assistantMsgId, + activeStreamRef.current?.assistantMsgId !== assistantMsgId, playTypewriterSound, playToolcallSound, onWriteFile, @@ -222,14 +343,20 @@ export function useAgentStream(options: UseAgentStreamOptions) { if (skillHandled) { return; } + + clearActiveStreamIfMatch(skillEventName); } } - let accumulatedContent = ""; let unlisten: (() => void) | null = null; - let requestLogId: string | null = null; - let requestStartedAt = 0; - let requestFinished = false; + const requestState = { + accumulatedContent: "", + requestLogId: null as string | null, + requestStartedAt: 0, + requestFinished: false, + queuedTurnId: null as string | null, + }; + let streamActivated = false; const optimisticStartedAt = assistantMsg.timestamp.toISOString(); const optimisticTurnId = `local-turn:${assistantMsgId}`; const optimisticItemId = `local-item:${assistantMsgId}:turn-summary`; @@ -240,91 +367,165 @@ export function useAgentStream(options: UseAgentStreamOptions) { const toolNameByToolId = new Map(); const actionLoggedKeys = new Set(); + const upsertQueuedTurn = (nextQueuedTurn: QueuedTurnSnapshot) => { + setQueuedTurns((prev) => + [ + ...prev.filter( + (item) => item.queued_turn_id !== nextQueuedTurn.queued_turn_id, + ), + nextQueuedTurn, + ].sort((left, right) => { + if (left.position !== right.position) { + return left.position - right.position; + } + return left.created_at - right.created_at; + }), + ); + }; + + const removeQueuedTurnState = (queuedTurnIds: string[]) => { + if (queuedTurnIds.length === 0) { + return; + } + setQueuedTurns((prev) => { + const idSet = new Set(queuedTurnIds); + return prev + .filter((item) => !idSet.has(item.queued_turn_id)) + .map((item, index) => ({ + ...item, + position: index + 1, + })); + }); + }; + + const removeQueuedDraftMessages = () => { + setMessages((prev) => + prev.filter( + (msg) => + msg.id !== assistantMsgId && + (userMsgId ? msg.id !== userMsgId : true), + ), + ); + }; + const clearOptimisticItem = () => { + if (expectingQueue) { + return; + } setThreadItems((prev) => removeThreadItemState(prev, optimisticItemId)); }; const clearOptimisticTurn = () => { + if (expectingQueue) { + return; + } setThreadTurns((prev) => removeThreadTurnState(prev, optimisticTurnId)); - setCurrentTurnId((prev) => - prev === optimisticTurnId ? null : prev, - ); + setCurrentTurnId((prev) => (prev === optimisticTurnId ? null : prev)); }; - setThreadTurns((prev) => - upsertThreadTurnState(prev, { - id: optimisticTurnId, - thread_id: optimisticThreadId, - prompt_text: content, - status: "running", - started_at: optimisticStartedAt, - created_at: optimisticStartedAt, - updated_at: optimisticStartedAt, - }), - ); - setThreadItems((prev) => - upsertThreadItemState(prev, { - id: optimisticItemId, - thread_id: optimisticThreadId, - turn_id: optimisticTurnId, - sequence: 0, - status: "in_progress", - started_at: optimisticStartedAt, - updated_at: optimisticStartedAt, - type: "turn_summary", - text: buildRuntimeStatusSummary(assistantMsg.runtimeStatus), - }), - ); - setCurrentTurnId(optimisticTurnId); + const disposeListener = () => { + const registered = listenerMapRef.current.get(eventName); + if (registered) { + registered(); + listenerMapRef.current.delete(eventName); + } else if (unlisten) { + unlisten(); + } + unlisten = null; + }; + + if (!expectingQueue) { + setThreadTurns((prev) => + upsertThreadTurnState(prev, { + id: optimisticTurnId, + thread_id: optimisticThreadId, + prompt_text: content, + status: "running", + started_at: optimisticStartedAt, + created_at: optimisticStartedAt, + updated_at: optimisticStartedAt, + }), + ); + setThreadItems((prev) => + upsertThreadItemState(prev, { + id: optimisticItemId, + thread_id: optimisticThreadId, + turn_id: optimisticTurnId, + sequence: 0, + status: "in_progress", + started_at: optimisticStartedAt, + updated_at: optimisticStartedAt, + type: "turn_summary", + text: buildRuntimeStatusSummary(assistantMsg.runtimeStatus), + }), + ); + setCurrentTurnId(optimisticTurnId); + } + + const eventName = `aster_stream_${assistantMsgId}`; try { const activeSessionId = await ensureSession(); if (!activeSessionId) throw new Error("无法创建会话"); - currentStreamingSessionIdRef.current = activeSessionId; const resolvedWorkspaceId = getRequiredWorkspaceId(); const waitingRuntimeStatus = buildWaitingAgentRuntimeStatus({ executionStrategy: effectiveExecutionStrategy, webSearch, thinking: _thinking, }); - setMessages((prev) => - prev.map((msg) => - msg.id === assistantMsgId - ? { - ...msg, - runtimeStatus: waitingRuntimeStatus, - } - : msg, - ), - ); - setThreadTurns((prev) => - upsertThreadTurnState(prev, { - id: optimisticTurnId, - thread_id: activeSessionId, - prompt_text: content, - status: "running", - started_at: optimisticStartedAt, - created_at: optimisticStartedAt, - updated_at: new Date().toISOString(), - }), - ); - setThreadItems((prev) => - upsertThreadItemState(prev, { - id: optimisticItemId, - thread_id: activeSessionId, - turn_id: optimisticTurnId, - sequence: 0, - status: "in_progress", - started_at: optimisticStartedAt, - updated_at: new Date().toISOString(), - type: "turn_summary", - text: buildRuntimeStatusSummary(waitingRuntimeStatus), - }), - ); - const eventName = `aster_stream_${assistantMsgId}`; - requestStartedAt = Date.now(); - requestLogId = activityLogger.log({ + const activateStream = () => { + if (streamActivated) { + return; + } + streamActivated = true; + setActiveStream({ + assistantMsgId, + eventName, + sessionId: activeSessionId, + }); + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...msg, + runtimeStatus: waitingRuntimeStatus, + } + : msg, + ), + ); + }; + + if (!expectingQueue) { + activateStream(); + setThreadTurns((prev) => + upsertThreadTurnState(prev, { + id: optimisticTurnId, + thread_id: activeSessionId, + prompt_text: content, + status: "running", + started_at: optimisticStartedAt, + created_at: optimisticStartedAt, + updated_at: new Date().toISOString(), + }), + ); + setThreadItems((prev) => + upsertThreadItemState(prev, { + id: optimisticItemId, + thread_id: activeSessionId, + turn_id: optimisticTurnId, + sequence: 0, + status: "in_progress", + started_at: optimisticStartedAt, + updated_at: new Date().toISOString(), + type: "turn_summary", + text: buildRuntimeStatusSummary(waitingRuntimeStatus), + }), + ); + } + + requestState.requestStartedAt = Date.now(); + requestState.requestLogId = activityLogger.log({ eventType: "chat_request_start", status: "pending", title: skipUserMessage ? "系统引导请求" : "发送请求", @@ -340,247 +541,60 @@ export function useAgentStream(options: UseAgentStreamOptions) { skipUserMessage, autoContinueEnabled: autoContinue?.enabled ?? false, autoContinue: autoContinue?.enabled ? autoContinue : undefined, + queuedSubmission: expectingQueue, }, }); unlisten = await runtime.listenToTurnEvents( eventName, - (event: { payload: StreamEvent | unknown }) => { + (event: { payload: unknown }) => { const data = parseStreamEvent(event.payload); - if (!data) return; - - switch (data.type) { - case "thread_started": - break; - - case "turn_started": - setCurrentTurnId(data.turn.id); - setThreadTurns((prev) => - upsertThreadTurnState( - removeThreadTurnState(prev, optimisticTurnId), - data.turn, - ), - ); - clearOptimisticItem(); - break; - - case "item_started": - case "item_updated": - case "item_completed": - setThreadItems((prev) => - upsertThreadItemState( - removeThreadItemState(prev, optimisticItemId), - data.item, - ), - ); - break; - - case "turn_completed": - case "turn_failed": - clearOptimisticItem(); - setThreadTurns((prev) => - upsertThreadTurnState( - removeThreadTurnState(prev, optimisticTurnId), - data.turn, - ), - ); - setCurrentTurnId(data.turn.id); - break; - - case "text_delta": - clearOptimisticItem(); - accumulatedContent += data.text; - observer?.onTextDelta?.(data.text, accumulatedContent); - playTypewriterSound(); - setMessages((prev) => - prev.map((msg) => - msg.id === assistantMsgId - ? { - ...msg, - content: accumulatedContent, - thinkingContent: undefined, - runtimeStatus: undefined, - contentParts: appendTextToParts( - msg.contentParts || [], - data.text, - ), - } - : msg, - ), - ); - break; - - case "tool_start": { - clearOptimisticItem(); - playToolcallSound(); - handleToolStartEvent({ - data, - setPendingActions, - onWriteFile, - toolLogIdByToolId, - toolStartedAtByToolId, - toolNameByToolId, - assistantMsgId, - activeSessionId, - resolvedWorkspaceId, - setMessages, - }); - break; - } - - case "tool_end": { - clearOptimisticItem(); - handleToolEndEvent({ - data, - onWriteFile, - toolLogIdByToolId, - toolStartedAtByToolId, - toolNameByToolId, - assistantMsgId, - activeSessionId, - resolvedWorkspaceId, - setMessages, - }); - break; - } - - case "artifact_snapshot": { - clearOptimisticItem(); - handleArtifactSnapshotEvent({ - data, - onWriteFile, - assistantMsgId, - activeSessionId, - resolvedWorkspaceId, - setMessages, - }); - break; - } - - case "action_required": { - clearOptimisticItem(); - handleActionRequiredEvent({ - data, - actionLoggedKeys, - effectiveExecutionStrategy, - runtime, - setPendingActions, - assistantMsgId, - activeSessionId, - resolvedWorkspaceId, - setMessages, - }); - break; - } - - case "context_trace": - clearOptimisticItem(); - handleContextTraceEvent({ - data, - assistantMsgId, - activeSessionId, - resolvedWorkspaceId, - setMessages, - }); - break; - - case "final_done": { - clearOptimisticItem(); - clearOptimisticTurn(); - if (requestLogId && !requestFinished) { - requestFinished = true; - activityLogger.updateLog(requestLogId, { - eventType: "chat_request_complete", - status: "success", - duration: Date.now() - requestStartedAt, - description: `请求完成,工具调用 ${toolLogIdByToolId.size} 次`, - }); - } - const finalContent = accumulatedContent || "(无响应)"; - observer?.onComplete?.(finalContent); - setMessages((prev) => - prev.map((msg) => - msg.id === assistantMsgId - ? { - ...updateMessageArtifactsStatus(msg, "complete"), - isThinking: false, - content: finalContent, - runtimeStatus: undefined, - } - : msg, - ), - ); - setIsSending(false); - unlistenRef.current = null; - currentAssistantMsgIdRef.current = null; - currentStreamingSessionIdRef.current = null; - if (unlisten) { - unlisten(); - unlisten = null; - } - break; - } - - case "error": - clearOptimisticItem(); - clearOptimisticTurn(); - if (requestLogId && !requestFinished) { - requestFinished = true; - activityLogger.updateLog(requestLogId, { - eventType: "chat_request_error", - status: "error", - duration: Date.now() - requestStartedAt, - error: data.message, - }); - } - observer?.onError?.(data.message); - if ( - data.message.includes("429") || - data.message.toLowerCase().includes("rate limit") - ) { - toast.warning("请求过于频繁,请稍后重试"); - } else { - toast.error(`响应错误: ${data.message}`); - } - setMessages((prev) => - prev.map((msg) => - msg.id === assistantMsgId - ? { - ...updateMessageArtifactsStatus(msg, "error"), - isThinking: false, - content: accumulatedContent || `错误: ${data.message}`, - runtimeStatus: undefined, - } - : msg, - ), - ); - setIsSending(false); - currentStreamingSessionIdRef.current = null; - if (unlisten) { - unlisten(); - unlisten = null; - } - break; - - case "warning": { - if (data.code === WORKSPACE_PATH_AUTO_CREATED_WARNING_CODE) { - break; - } - const warningKey = `${activeSessionId}:${data.code || data.message}`; - if (!warnedKeysRef.current.has(warningKey)) { - warnedKeysRef.current.add(warningKey); - toast.warning(data.message); - } - break; - } - - default: - break; + if (!data) { + return; } + + handleTurnStreamEvent({ + data, + requestState, + callbacks: { + activateStream, + isStreamActivated: () => streamActivated, + clearOptimisticItem, + clearOptimisticTurn, + disposeListener, + removeQueuedDraftMessages, + clearActiveStreamIfMatch, + upsertQueuedTurn, + removeQueuedTurnState, + playToolcallSound, + playTypewriterSound, + appendThinkingToParts, + }, + observer, + eventName, + optimisticTurnId, + optimisticItemId, + assistantMsgId, + activeSessionId, + resolvedWorkspaceId, + effectiveExecutionStrategy, + runtime, + warnedKeysRef, + actionLoggedKeys, + toolLogIdByToolId, + toolStartedAtByToolId, + toolNameByToolId, + onWriteFile, + setMessages, + setPendingActions, + setThreadItems, + setThreadTurns, + setCurrentTurnId, + }); }, ); - unlistenRef.current = unlisten; + listenerMapRef.current.set(eventName, unlisten); const imagesToSend = images.length > 0 @@ -605,17 +619,19 @@ export function useAgentStream(options: UseAgentStreamOptions) { providerConfig, executionStrategy: effectiveExecutionStrategy, webSearch, + searchMode: webSearch ? "allowed" : "disabled", autoContinue, systemPrompt, metadata: requestMetadata, + queueIfBusy: true, }); } catch (error) { - if (requestLogId && !requestFinished) { - requestFinished = true; - activityLogger.updateLog(requestLogId, { + if (requestState.requestLogId && !requestState.requestFinished) { + requestState.requestFinished = true; + activityLogger.updateLog(requestState.requestLogId, { eventType: "chat_request_error", status: "error", - duration: Date.now() - requestStartedAt, + duration: Date.now() - requestState.requestStartedAt, error: error instanceof Error ? error.message : String(error), }); } @@ -634,45 +650,61 @@ export function useAgentStream(options: UseAgentStreamOptions) { } clearOptimisticItem(); clearOptimisticTurn(); - setMessages((prev) => prev.filter((msg) => msg.id !== assistantMsgId)); - setIsSending(false); - currentStreamingSessionIdRef.current = null; - if (unlisten) { - unlisten(); + removeQueuedTurnState( + requestState.queuedTurnId ? [requestState.queuedTurnId] : [], + ); + setMessages((prev) => + prev.filter( + (msg) => + msg.id !== assistantMsgId && + (!expectingQueue || !userMsgId || msg.id !== userMsgId), + ), + ); + clearActiveStreamIfMatch(eventName); + disposeListener(); + if (!expectingQueue && !activeStreamRef.current) { + setIsSending(false); } } }, [ - currentAssistantMsgIdRef, - currentStreamingSessionIdRef, + activeStreamRef, + buildQueuedRuntimeStatus, + buildRuntimeStatusSummary, + clearActiveStreamIfMatch, ensureSession, executionStrategy, getRequiredWorkspaceId, modelRef, onWriteFile, providerTypeRef, + queuedTurns.length, runtime, + sessionIdRef, + setActiveStream, + setCurrentTurnId, setMessages, + setPendingActions, + setQueuedTurns, setThreadItems, setThreadTurns, - setCurrentTurnId, - setPendingActions, setWorkspacePathMissing, systemPrompt, warnedKeysRef, - sessionIdRef, - buildRuntimeStatusSummary, ], ); const stopSending = useCallback(async () => { - if (unlistenRef.current) { - unlistenRef.current(); - unlistenRef.current = null; + const activeStream = activeStreamRef.current; + if (activeStream) { + const activeUnlisten = listenerMapRef.current.get(activeStream.eventName); + if (activeUnlisten) { + activeUnlisten(); + listenerMapRef.current.delete(activeStream.eventName); + } } - const activeSessionId = - currentStreamingSessionIdRef.current || sessionIdRef.current; + const activeSessionId = activeStream?.sessionId || sessionIdRef.current; if (activeSessionId) { try { await runtime.interruptTurn(activeSessionId); @@ -681,15 +713,17 @@ export function useAgentStream(options: UseAgentStreamOptions) { } } - if (currentAssistantMsgIdRef.current) { - const optimisticTurnId = `local-turn:${currentAssistantMsgIdRef.current}`; - const optimisticItemId = `${`local-item:${currentAssistantMsgIdRef.current}`}:turn-summary`; + setQueuedTurns([]); + + if (activeStream?.assistantMsgId) { + const optimisticTurnId = `local-turn:${activeStream.assistantMsgId}`; + const optimisticItemId = `${`local-item:${activeStream.assistantMsgId}`}:turn-summary`; setThreadItems((prev) => removeThreadItemState(prev, optimisticItemId)); setThreadTurns((prev) => removeThreadTurnState(prev, optimisticTurnId)); setCurrentTurnId((prev) => (prev === optimisticTurnId ? null : prev)); setMessages((prev) => prev.map((msg) => - msg.id === currentAssistantMsgIdRef.current + msg.id === activeStream.assistantMsgId ? { ...updateMessageArtifactsStatus(msg, "complete"), isThinking: false, @@ -699,26 +733,57 @@ export function useAgentStream(options: UseAgentStreamOptions) { : msg, ), ); - currentAssistantMsgIdRef.current = null; } - currentStreamingSessionIdRef.current = null; - setIsSending(false); + setActiveStream(null); toast.info("已停止生成"); }, [ - currentAssistantMsgIdRef, - currentStreamingSessionIdRef, runtime, sessionIdRef, + setActiveStream, + setCurrentTurnId, setMessages, + setQueuedTurns, setThreadItems, setThreadTurns, - setCurrentTurnId, ]); + const removeQueuedTurn = useCallback( + async (queuedTurnId: string) => { + const activeSessionId = sessionIdRef.current; + if (!activeSessionId || !queuedTurnId.trim()) { + return false; + } + + try { + const removed = await runtime.removeQueuedTurn( + activeSessionId, + queuedTurnId, + ); + if (removed) { + setQueuedTurns((prev) => + prev + .filter((item) => item.queued_turn_id !== queuedTurnId) + .map((item, index) => ({ + ...item, + position: index + 1, + })), + ); + } + return removed; + } catch (error) { + console.error("[AsterChat] 移除排队消息失败:", error); + toast.error("移除排队消息失败"); + return false; + } + }, + [runtime, sessionIdRef, setQueuedTurns], + ); + return { isSending, sendMessage, stopSending, + removeQueuedTurn, }; } diff --git a/src/components/agent/chat/hooks/useAgentTools.ts b/src/components/agent/chat/hooks/useAgentTools.ts index e55055c6a..21a29d8e5 100644 --- a/src/components/agent/chat/hooks/useAgentTools.ts +++ b/src/components/agent/chat/hooks/useAgentTools.ts @@ -1,5 +1,6 @@ import { useCallback, + useEffect, useRef, useState, type Dispatch, @@ -36,6 +37,14 @@ export function useAgentTools(options: UseAgentToolsOptions) { const [pendingActions, setPendingActions] = useState([]); const warnedKeysRef = useRef>(new Set()); + const queuedFallbackResponsesRef = useRef< + Map< + string, + Omit & { + requestId: string; + } + > + >(new Map()); const confirmAction = useCallback( async (response: ConfirmResponse) => { @@ -94,7 +103,55 @@ export function useAgentTools(options: UseAgentToolsOptions) { }); if (!resolvedAction) { - throw new Error("Ask 请求 ID 尚未就绪,请稍后再试"); + queuedFallbackResponsesRef.current.set(fallbackPromptKey, { + ...response, + actionType, + requestId: pendingAction.requestId, + userData, + }); + setPendingActions((prev) => + prev.map((item) => + item.requestId === pendingAction.requestId + ? { + ...item, + status: "queued", + submittedResponse: normalizedResponse || undefined, + submittedUserData, + } + : item, + ), + ); + setMessages((prev) => + prev.map((msg) => ({ + ...msg, + actionRequests: msg.actionRequests?.map((item) => + item.requestId === pendingAction.requestId + ? { + ...item, + status: "queued" as const, + submittedResponse: normalizedResponse || undefined, + submittedUserData, + } + : item, + ), + contentParts: msg.contentParts?.map((part) => + part.type === "action_required" && + part.actionRequired.requestId === pendingAction.requestId + ? { + ...part, + actionRequired: { + ...part.actionRequired, + status: "queued" as const, + submittedResponse: normalizedResponse || undefined, + submittedUserData, + }, + } + : part, + ), + })), + ); + toast.info("已记录你的回答,等待系统请求就绪后自动提交"); + return; } effectiveRequestId = resolvedAction.requestId; @@ -190,6 +247,37 @@ export function useAgentTools(options: UseAgentToolsOptions) { ], ); + useEffect(() => { + for (const pendingAction of pendingActions) { + if ( + pendingAction.isFallback || + pendingAction.status === "submitted" || + (pendingAction.actionType !== "ask_user" && + pendingAction.actionType !== "elicitation") + ) { + continue; + } + + const promptKey = resolveActionPromptKey(pendingAction); + if (!promptKey) { + continue; + } + + const queuedResponse = queuedFallbackResponsesRef.current.get(promptKey); + if (!queuedResponse) { + continue; + } + + queuedFallbackResponsesRef.current.delete(promptKey); + void confirmAction({ + ...queuedResponse, + requestId: pendingAction.requestId, + actionType: pendingAction.actionType, + }); + break; + } + }, [confirmAction, pendingActions]); + const handlePermissionResponse = useCallback( async (response: ConfirmResponse) => { await confirmAction(response); diff --git a/src/components/agent/chat/hooks/useArtifactAutoPreviewSync.test.ts b/src/components/agent/chat/hooks/useArtifactAutoPreviewSync.test.ts new file mode 100644 index 000000000..290ac56ec --- /dev/null +++ b/src/components/agent/chat/hooks/useArtifactAutoPreviewSync.test.ts @@ -0,0 +1,91 @@ +import { describe, expect, it } from "vitest"; +import type { Artifact } from "@/lib/artifact/types"; +import { + mergePreviewContentIntoArtifact, + shouldAutoSyncArtifactPreview, +} from "./useArtifactAutoPreviewSync"; + +function createArtifact(overrides: Partial = {}): Artifact { + const content = overrides.content ?? ""; + return { + id: overrides.id ?? "artifact-1", + type: overrides.type ?? "document", + title: overrides.title ?? "demo.md", + content, + status: overrides.status ?? "pending", + meta: { + filePath: overrides.meta?.filePath ?? "workspace/demo.md", + filename: overrides.meta?.filename ?? "demo.md", + ...overrides.meta, + }, + position: overrides.position ?? { start: 0, end: content.length }, + createdAt: overrides.createdAt ?? 1, + updatedAt: overrides.updatedAt ?? 1, + error: overrides.error, + }; +} + +describe("useArtifactAutoPreviewSync helpers", () => { + it("空内容的 pending artifact 应触发自动预览同步", () => { + const artifact = createArtifact({ + status: "pending", + meta: { + filePath: "workspace/demo.md", + writePhase: "preparing", + }, + }); + + expect(shouldAutoSyncArtifactPreview(artifact)).toBe(true); + }); + + it("已经完成且有内容的 artifact 不应继续轮询", () => { + const artifact = createArtifact({ + status: "complete", + content: "# 已完成", + meta: { + filePath: "workspace/demo.md", + writePhase: "completed", + }, + }); + + expect(shouldAutoSyncArtifactPreview(artifact)).toBe(false); + }); + + it("读取到文件内容后应把 preview 合并回 artifact", () => { + const artifact = createArtifact({ + status: "streaming", + meta: { + filePath: "workspace/demo.md", + writePhase: "streaming", + }, + }); + + const merged = mergePreviewContentIntoArtifact(artifact, { + path: "workspace/demo.md", + content: "# 标题\n\n第一段", + }); + + expect(merged).not.toBeNull(); + expect(merged?.content).toContain("第一段"); + expect(merged?.status).toBe("streaming"); + expect(merged?.meta.writePhase).toBe("streaming"); + }); + + it("已有更长内容时不应被更短的 preview 回退覆盖", () => { + const artifact = createArtifact({ + status: "streaming", + content: "# 标题\n\n第一段\n第二段", + meta: { + filePath: "workspace/demo.md", + writePhase: "streaming", + }, + }); + + const merged = mergePreviewContentIntoArtifact(artifact, { + path: "workspace/demo.md", + content: "# 标题\n\n第一段", + }); + + expect(merged).toBeNull(); + }); +}); diff --git a/src/components/agent/chat/hooks/useArtifactAutoPreviewSync.ts b/src/components/agent/chat/hooks/useArtifactAutoPreviewSync.ts new file mode 100644 index 000000000..10f74a00b --- /dev/null +++ b/src/components/agent/chat/hooks/useArtifactAutoPreviewSync.ts @@ -0,0 +1,227 @@ +import { useEffect } from "react"; +import type { Artifact } from "@/lib/artifact/types"; +import type { + ArtifactWriteMetadata, + WriteArtifactContext, +} from "../types"; +import { + buildArtifactFromWrite, + resolveArtifactWritePhase, +} from "../utils/messageArtifacts"; + +export interface ArtifactAutoPreviewResult { + path?: string; + content?: string | null; + isBinary?: boolean; + error?: string | null; +} + +interface UseArtifactAutoPreviewSyncOptions { + enabled: boolean; + artifact: Artifact | null; + loadPreview: (path: string) => Promise; + onSyncArtifact: (artifact: Artifact) => void; +} + +const STREAM_SYNC_POLL_INTERVAL_MS = 280; +const EMPTY_COMPLETE_SYNC_TIMEOUT_MS = 8000; +const PREVIEW_TEXT_MAX_CHARS = 480; +const LATEST_CHUNK_MAX_CHARS = 240; + +function resolveArtifactFilePath(artifact: Pick): string { + if (typeof artifact.meta.filePath === "string" && artifact.meta.filePath.trim()) { + return artifact.meta.filePath.trim(); + } + if (typeof artifact.meta.filename === "string" && artifact.meta.filename.trim()) { + return artifact.meta.filename.trim(); + } + return artifact.title; +} + +function normalizePreviewText(value: string, maxChars: number): string { + const trimmed = value.trim(); + if (trimmed.length <= maxChars) { + return trimmed; + } + return `${trimmed.slice(0, maxChars).trimEnd()}…`; +} + +export function shouldAutoSyncArtifactPreview(artifact: Artifact | null): boolean { + if (!artifact) { + return false; + } + + const artifactPath = resolveArtifactFilePath(artifact); + if (!artifactPath.trim()) { + return false; + } + + const writePhase = resolveArtifactWritePhase(artifact); + if (!artifact.content.trim()) { + return ( + artifact.status === "pending" || + artifact.status === "streaming" || + artifact.status === "complete" || + writePhase === "preparing" || + writePhase === "streaming" || + writePhase === "persisted" || + writePhase === "completed" + ); + } + + return artifact.status === "streaming" || writePhase === "streaming"; +} + +export function mergePreviewContentIntoArtifact( + artifact: Artifact, + preview: ArtifactAutoPreviewResult, +): Artifact | null { + if (preview.isBinary || preview.error) { + return null; + } + + const nextContent = + typeof preview.content === "string" ? preview.content : artifact.content; + const nextPath = preview.path?.trim() || resolveArtifactFilePath(artifact); + const currentContent = artifact.content; + + if (!nextContent.trim() && currentContent.trim()) { + return null; + } + + if ( + currentContent.trim() && + nextContent.length < currentContent.length && + currentContent.startsWith(nextContent) + ) { + return null; + } + + if (nextContent === currentContent && nextPath === resolveArtifactFilePath(artifact)) { + return null; + } + + const currentWritePhase = resolveArtifactWritePhase(artifact); + const nextStatus = + artifact.status === "complete" || currentWritePhase === "completed" + ? "complete" + : artifact.status === "error" || currentWritePhase === "failed" + ? "error" + : nextContent.trim() + ? "streaming" + : artifact.status; + const nextWritePhase: WriteArtifactContext["metadata"] = { + ...(artifact.meta as ArtifactWriteMetadata), + writePhase: + nextStatus === "complete" + ? "completed" + : nextStatus === "error" + ? "failed" + : nextContent.trim() + ? "streaming" + : currentWritePhase || undefined, + previewText: nextContent.trim() + ? normalizePreviewText(nextContent, PREVIEW_TEXT_MAX_CHARS) + : (artifact.meta.previewText as string | undefined), + latestChunk: nextContent.trim() + ? normalizePreviewText( + nextContent.slice(-LATEST_CHUNK_MAX_CHARS), + LATEST_CHUNK_MAX_CHARS, + ) + : (artifact.meta.latestChunk as string | undefined), + isPartial: nextStatus !== "complete" && nextStatus !== "error", + lastUpdateSource: + (artifact.meta.lastUpdateSource as WriteArtifactContext["source"]) || + "artifact_snapshot", + }; + + return buildArtifactFromWrite({ + filePath: nextPath, + content: nextContent, + context: { + artifact, + artifactId: artifact.id, + source: + (artifact.meta.lastUpdateSource as WriteArtifactContext["source"]) || + "artifact_snapshot", + sourceMessageId: + typeof artifact.meta.sourceMessageId === "string" + ? artifact.meta.sourceMessageId + : undefined, + status: nextStatus, + metadata: nextWritePhase, + }, + }); +} + +export function useArtifactAutoPreviewSync({ + enabled, + artifact, + loadPreview, + onSyncArtifact, +}: UseArtifactAutoPreviewSyncOptions): void { + useEffect(() => { + if (!enabled || !artifact || !shouldAutoSyncArtifactPreview(artifact)) { + return; + } + + const artifactPath = resolveArtifactFilePath(artifact); + if (!artifactPath.trim()) { + return; + } + + let disposed = false; + let timer: number | null = null; + let inFlight = false; + const startedAt = Date.now(); + + const scheduleNext = () => { + if (disposed) { + return; + } + timer = window.setTimeout(runSync, STREAM_SYNC_POLL_INTERVAL_MS); + }; + + const runSync = async () => { + if (disposed || inFlight) { + return; + } + + const currentWritePhase = resolveArtifactWritePhase(artifact); + const shouldStopOnTimeout = + !artifact.content.trim() && + (artifact.status === "complete" || currentWritePhase === "completed") && + Date.now() - startedAt >= EMPTY_COMPLETE_SYNC_TIMEOUT_MS; + if (shouldStopOnTimeout) { + return; + } + + inFlight = true; + try { + const preview = await loadPreview(artifactPath); + if (disposed) { + return; + } + + const nextArtifact = mergePreviewContentIntoArtifact(artifact, preview); + if (nextArtifact) { + onSyncArtifact(nextArtifact); + } + } catch { + // 预览同步只做兜底,不影响主流程。 + } finally { + inFlight = false; + scheduleNext(); + } + }; + + void runSync(); + + return () => { + disposed = true; + if (timer !== null) { + window.clearTimeout(timer); + } + }; + }, [artifact, enabled, loadPreview, onSyncArtifact]); +} diff --git a/src/components/agent/chat/hooks/useArtifactDisplayState.test.ts b/src/components/agent/chat/hooks/useArtifactDisplayState.test.ts new file mode 100644 index 000000000..5953225db --- /dev/null +++ b/src/components/agent/chat/hooks/useArtifactDisplayState.test.ts @@ -0,0 +1,162 @@ +import { describe, expect, it } from "vitest"; +import type { Artifact } from "@/lib/artifact/types"; +import { resolveArtifactDisplayState } from "./useArtifactDisplayState"; + +function createArtifact(overrides: Partial = {}): Artifact { + const content = overrides.content ?? "# 内容"; + return { + id: overrides.id ?? "artifact-1", + type: overrides.type ?? "document", + title: overrides.title ?? "demo.md", + content, + status: overrides.status ?? "complete", + meta: { + filePath: overrides.meta?.filePath ?? "workspace/demo.md", + filename: overrides.meta?.filename ?? "demo.md", + ...overrides.meta, + }, + position: overrides.position ?? { start: 0, end: content.length }, + createdAt: overrides.createdAt ?? 1, + updatedAt: overrides.updatedAt ?? 1, + error: overrides.error, + }; +} + +describe("resolveArtifactDisplayState", () => { + it("有内容时应直接展示 live artifact", () => { + const liveArtifact = createArtifact({ + id: "artifact-live", + content: "# 最新版本", + status: "streaming", + meta: { + filePath: "workspace/live.md", + writePhase: "streaming", + }, + }); + + const state = resolveArtifactDisplayState({ + liveArtifact, + artifacts: [liveArtifact], + }); + + expect(state.mode).toBe("content"); + expect(state.displayArtifact?.id).toBe("artifact-live"); + expect(state.overlay).toBeNull(); + expect(state.showPreviousVersionBadge).toBe(false); + }); + + it("新 artifact 为空且存在上一版本时应保留旧内容并显示 overlay", () => { + const previousArtifact = createArtifact({ + id: "artifact-prev", + title: "old.md", + content: "# 旧版本", + status: "complete", + meta: { + filePath: "workspace/old.md", + }, + }); + const liveArtifact = createArtifact({ + id: "artifact-live", + title: "new.md", + content: "", + status: "streaming", + meta: { + filePath: "workspace/new.md", + writePhase: "streaming", + }, + }); + + const state = resolveArtifactDisplayState({ + liveArtifact, + artifacts: [previousArtifact, liveArtifact], + previousRenderableArtifact: previousArtifact, + }); + + expect(state.mode).toBe("overlay-on-previous"); + expect(state.displayArtifact?.id).toBe("artifact-prev"); + expect(state.overlay?.phase).toBe("streaming_content"); + expect(state.overlay?.displayName).toBe("new.md"); + expect(state.showPreviousVersionBadge).toBe(true); + }); + + it("没有上一版本时应回退到类型化 skeleton", () => { + const liveArtifact = createArtifact({ + id: "artifact-live", + title: "index.ts", + type: "code", + content: "", + status: "pending", + meta: { + filePath: "workspace/index.ts", + writePhase: "preparing", + language: "typescript", + }, + }); + + const state = resolveArtifactDisplayState({ + liveArtifact, + artifacts: [liveArtifact], + }); + + expect(state.mode).toBe("typed-skeleton"); + expect(state.displayArtifact?.id).toBe("artifact-live"); + expect(state.overlay).toBeNull(); + expect(state.showPreviousVersionBadge).toBe(false); + }); + + it("完成但仍无内容时应保留上一版本并提示完成态", () => { + const previousArtifact = createArtifact({ + id: "artifact-prev", + title: "report.md", + content: "# 上一版报告", + status: "complete", + meta: { + filePath: "workspace/report.md", + }, + }); + const liveArtifact = createArtifact({ + id: "artifact-live", + title: "summary.md", + content: "", + status: "complete", + meta: { + filePath: "workspace/summary.md", + writePhase: "completed", + }, + }); + + const state = resolveArtifactDisplayState({ + liveArtifact, + artifacts: [previousArtifact, liveArtifact], + previousRenderableArtifact: previousArtifact, + }); + + expect(state.mode).toBe("overlay-on-previous"); + expect(state.displayArtifact?.id).toBe("artifact-prev"); + expect(state.overlay?.phase).toBe("finalized_empty"); + expect(state.overlay?.title).toContain("写入已结束"); + }); + + it("写入失败且没有上一版本时应展示错误态空画布", () => { + const liveArtifact = createArtifact({ + id: "artifact-live", + title: "broken.md", + content: "", + status: "error", + error: "磁盘写入失败", + meta: { + filePath: "workspace/broken.md", + writePhase: "failed", + }, + }); + + const state = resolveArtifactDisplayState({ + liveArtifact, + artifacts: [liveArtifact], + }); + + expect(state.mode).toBe("error"); + expect(state.displayArtifact?.id).toBe("artifact-live"); + expect(state.overlay).toBeNull(); + }); +}); diff --git a/src/components/agent/chat/hooks/useArtifactDisplayState.ts b/src/components/agent/chat/hooks/useArtifactDisplayState.ts new file mode 100644 index 000000000..3692ec890 --- /dev/null +++ b/src/components/agent/chat/hooks/useArtifactDisplayState.ts @@ -0,0 +1,325 @@ +import { useEffect, useMemo, useRef, useState } from "react"; +import type { Artifact } from "@/lib/artifact/types"; +import { resolveArtifactWritePhase } from "../utils/messageArtifacts"; + +export type ArtifactDisplayMode = + | "content" + | "overlay-on-previous" + | "typed-skeleton" + | "empty-finished" + | "error"; + +export type ArtifactOverlayPhase = + | "creating" + | "streaming_content" + | "updating_file" + | "finalized_empty" + | "failed"; + +export interface ArtifactDisplayOverlayState { + phase: ArtifactOverlayPhase; + phaseLabel: string; + title: string; + detail: string; + displayName: string; + filePath: string; + showProgress: boolean; +} + +export interface ArtifactDisplayState { + liveArtifact: Artifact | null; + displayArtifact: Artifact | null; + mode: ArtifactDisplayMode; + overlay: ArtifactDisplayOverlayState | null; + showPreviousVersionBadge: boolean; +} + +export interface ResolveArtifactDisplayStateOptions { + liveArtifact: Artifact | null; + artifacts: Artifact[]; + previousRenderableArtifact?: Artifact | null; + isSlowTransition?: boolean; +} + +const SLOW_TRANSITION_THRESHOLD_MS = 900; + +function hasRenderableArtifactContent(artifact: Artifact | null | undefined): boolean { + return Boolean(artifact?.content.trim()); +} + +function resolveArtifactPath(artifact: Pick): string { + if (typeof artifact.meta.filePath === "string" && artifact.meta.filePath.trim()) { + return artifact.meta.filePath.trim(); + } + if (typeof artifact.meta.filename === "string" && artifact.meta.filename.trim()) { + return artifact.meta.filename.trim(); + } + return artifact.title; +} + +function resolveArtifactDisplayName(path: string): string { + const normalized = path.replace(/\\/g, "/"); + const segments = normalized.split("/"); + return segments[segments.length - 1] || path; +} + +function artifactStillExists( + artifact: Artifact | null | undefined, + artifacts: Artifact[], +): artifact is Artifact { + if (!artifact) { + return false; + } + return artifacts.some((candidate) => candidate.id === artifact.id); +} + +function findPreviousRenderableArtifact( + liveArtifact: Artifact, + artifacts: Artifact[], + preferred: Artifact | null | undefined, +): Artifact | null { + if ( + preferred && + preferred.id !== liveArtifact.id && + artifactStillExists(preferred, artifacts) && + hasRenderableArtifactContent(preferred) + ) { + return preferred; + } + + for (let index = artifacts.length - 1; index >= 0; index -= 1) { + const candidate = artifacts[index]; + if (candidate.id === liveArtifact.id) { + continue; + } + if (hasRenderableArtifactContent(candidate)) { + return candidate; + } + } + + return null; +} + +function buildOverlayState( + artifact: Artifact, + phase: ArtifactOverlayPhase, + options: { isSlowTransition: boolean }, +): ArtifactDisplayOverlayState { + const filePath = resolveArtifactPath(artifact); + const displayName = resolveArtifactDisplayName(filePath); + + switch (phase) { + case "creating": + return { + phase, + phaseLabel: "准备写入", + title: "正在创建文件", + detail: options.isSlowTransition + ? "文件已创建,正在生成首段内容。" + : "正在准备首段内容,画布会在内容到达后立即切换。", + displayName, + filePath, + showProgress: true, + }; + case "streaming_content": + return { + phase, + phaseLabel: "正在写入", + title: "正在生成新版本", + detail: options.isSlowTransition + ? "内容还在流式生成中,当前先保留上一份可见版本。" + : "新的内容片段正在写入,首段到达后会切换到最新版本。", + displayName, + filePath, + showProgress: true, + }; + case "updating_file": + return { + phase, + phaseLabel: "已落盘", + title: "正在同步最新内容", + detail: "文件已经落盘,正在等待可渲染内容同步到画布。", + displayName, + filePath, + showProgress: true, + }; + case "finalized_empty": + return { + phase, + phaseLabel: "已完成", + title: "写入已结束", + detail: "文件已经完成,但当前还没有可直接渲染的内容,暂时保留上一版本。", + displayName, + filePath, + showProgress: false, + }; + case "failed": + return { + phase, + phaseLabel: "失败", + title: "写入未完成", + detail: + artifact.error?.trim() || + "文件写入过程中出现异常,当前先保留上一份可见内容。", + displayName, + filePath, + showProgress: false, + }; + } +} + +export function resolveArtifactDisplayState({ + liveArtifact, + artifacts, + previousRenderableArtifact, + isSlowTransition = false, +}: ResolveArtifactDisplayStateOptions): ArtifactDisplayState { + if (!liveArtifact) { + return { + liveArtifact: null, + displayArtifact: null, + mode: "content", + overlay: null, + showPreviousVersionBadge: false, + }; + } + + if (hasRenderableArtifactContent(liveArtifact)) { + return { + liveArtifact, + displayArtifact: liveArtifact, + mode: "content", + overlay: null, + showPreviousVersionBadge: false, + }; + } + + const writePhase = resolveArtifactWritePhase(liveArtifact); + const previousArtifact = findPreviousRenderableArtifact( + liveArtifact, + artifacts, + previousRenderableArtifact, + ); + + if (liveArtifact.status === "error" || writePhase === "failed") { + return { + liveArtifact, + displayArtifact: previousArtifact || liveArtifact, + mode: "error", + overlay: previousArtifact + ? buildOverlayState(liveArtifact, "failed", { isSlowTransition }) + : null, + showPreviousVersionBadge: Boolean(previousArtifact), + }; + } + + if ( + liveArtifact.status === "complete" || + writePhase === "completed" || + writePhase === "persisted" + ) { + return { + liveArtifact, + displayArtifact: previousArtifact || liveArtifact, + mode: previousArtifact ? "overlay-on-previous" : "empty-finished", + overlay: previousArtifact + ? buildOverlayState(liveArtifact, "finalized_empty", { + isSlowTransition, + }) + : null, + showPreviousVersionBadge: Boolean(previousArtifact), + }; + } + + if (liveArtifact.status === "pending" || writePhase === "preparing") { + return { + liveArtifact, + displayArtifact: previousArtifact || liveArtifact, + mode: previousArtifact ? "overlay-on-previous" : "typed-skeleton", + overlay: previousArtifact + ? buildOverlayState(liveArtifact, "creating", { isSlowTransition }) + : null, + showPreviousVersionBadge: Boolean(previousArtifact), + }; + } + + if (liveArtifact.status === "streaming" || writePhase === "streaming") { + return { + liveArtifact, + displayArtifact: previousArtifact || liveArtifact, + mode: previousArtifact ? "overlay-on-previous" : "typed-skeleton", + overlay: previousArtifact + ? buildOverlayState(liveArtifact, "streaming_content", { + isSlowTransition, + }) + : null, + showPreviousVersionBadge: Boolean(previousArtifact), + }; + } + + return { + liveArtifact, + displayArtifact: previousArtifact || liveArtifact, + mode: previousArtifact ? "overlay-on-previous" : "typed-skeleton", + overlay: previousArtifact + ? buildOverlayState(liveArtifact, "updating_file", { isSlowTransition }) + : null, + showPreviousVersionBadge: Boolean(previousArtifact), + }; +} + +export function useArtifactDisplayState( + liveArtifact: Artifact | null, + artifacts: Artifact[], +): ArtifactDisplayState { + const lastRenderableArtifactRef = useRef(null); + const [isSlowTransition, setIsSlowTransition] = useState(false); + + useEffect(() => { + if (!liveArtifact && artifacts.length === 0) { + lastRenderableArtifactRef.current = null; + setIsSlowTransition(false); + return; + } + + if (hasRenderableArtifactContent(liveArtifact)) { + lastRenderableArtifactRef.current = liveArtifact; + setIsSlowTransition(false); + return; + } + + const writePhase = liveArtifact ? resolveArtifactWritePhase(liveArtifact) : null; + const isPendingTransition = + liveArtifact && + !hasRenderableArtifactContent(liveArtifact) && + (liveArtifact.status === "pending" || + liveArtifact.status === "streaming" || + writePhase === "preparing" || + writePhase === "streaming"); + + if (!isPendingTransition) { + setIsSlowTransition(false); + return; + } + + setIsSlowTransition(false); + const timer = window.setTimeout(() => { + setIsSlowTransition(true); + }, SLOW_TRANSITION_THRESHOLD_MS); + + return () => { + window.clearTimeout(timer); + }; + }, [artifacts.length, liveArtifact]); + + return useMemo( + () => + resolveArtifactDisplayState({ + liveArtifact, + artifacts, + previousRenderableArtifact: lastRenderableArtifactRef.current, + isSlowTransition, + }), + [artifacts, isSlowTransition, liveArtifact], + ); +} diff --git a/src/components/agent/chat/hooks/useAsterAgentChat.test.tsx b/src/components/agent/chat/hooks/useAsterAgentChat.test.tsx index 84aa8c6f9..08381f057 100644 --- a/src/components/agent/chat/hooks/useAsterAgentChat.test.tsx +++ b/src/components/agent/chat/hooks/useAsterAgentChat.test.tsx @@ -1,6 +1,7 @@ import { act } from "react"; import { createRoot } from "react-dom/client"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import type { WriteArtifactContext } from "../types"; const { mockInitAsterAgent, @@ -11,6 +12,7 @@ const { mockUpdateAgentRuntimeSession, mockDeleteAsterSession, mockInterruptAgentRuntimeTurn, + mockRemoveAgentRuntimeQueuedTurn, mockRespondAgentRuntimeAction, mockParseStreamEvent, mockSafeListen, @@ -26,6 +28,7 @@ const { mockUpdateAgentRuntimeSession: vi.fn(), mockDeleteAsterSession: vi.fn(), mockInterruptAgentRuntimeTurn: vi.fn(), + mockRemoveAgentRuntimeQueuedTurn: vi.fn(), mockRespondAgentRuntimeAction: vi.fn(), mockParseStreamEvent: vi.fn((payload: unknown) => payload), mockSafeListen: vi.fn(), @@ -56,6 +59,7 @@ vi.mock("@/lib/api/agentRuntime", () => ({ updateAgentRuntimeSession: mockUpdateAgentRuntimeSession, deleteAgentRuntimeSession: mockDeleteAsterSession, interruptAgentRuntimeTurn: mockInterruptAgentRuntimeTurn, + removeAgentRuntimeQueuedTurn: mockRemoveAgentRuntimeQueuedTurn, respondAgentRuntimeAction: mockRespondAgentRuntimeAction, })); @@ -83,7 +87,16 @@ interface HookHarness { unmount: () => void; } -function mountHook(workspaceId = "ws-test"): HookHarness { +function mountHook( + workspaceId = "ws-test", + currentOptions: { + onWriteFile?: ( + content: string, + fileName: string, + context?: WriteArtifactContext, + ) => void; + } = {}, +): HookHarness { const container = document.createElement("div"); document.body.appendChild(container); const root = createRoot(container); @@ -91,7 +104,10 @@ function mountHook(workspaceId = "ws-test"): HookHarness { let hookValue: ReturnType | null = null; function TestComponent() { - hookValue = useAsterAgentChat({ workspaceId }); + hookValue = useAsterAgentChat({ + workspaceId, + onWriteFile: currentOptions.onWriteFile, + }); return null; } @@ -161,6 +177,7 @@ beforeEach(() => { mockUpdateAgentRuntimeSession.mockResolvedValue(undefined); mockDeleteAsterSession.mockResolvedValue(undefined); mockInterruptAgentRuntimeTurn.mockResolvedValue(undefined); + mockRemoveAgentRuntimeQueuedTurn.mockResolvedValue(true); mockRespondAgentRuntimeAction.mockResolvedValue(undefined); mockSafeListen.mockResolvedValue(() => {}); mockParseSkillSlashCommand.mockReturnValue(null); @@ -306,6 +323,57 @@ describe("useAsterAgentChat.confirmAction", () => { }); }); +describe("useAsterAgentChat queue hydration", () => { + it("切换话题时应恢复后端返回的排队项", async () => { + mockListAsterSessions.mockResolvedValue([ + { + id: "session-queue", + name: "带队列的话题", + created_at: 1, + updated_at: 2, + }, + ]); + mockGetAsterSession.mockResolvedValue({ + id: "session-queue", + messages: [], + turns: [], + items: [], + queued_turns: [ + { + queuedTurnId: "queued-1", + messagePreview: "继续补充 PRD", + messageText: "继续补充 PRD,并补一版里程碑拆解", + createdAt: 1700000000000, + imageCount: 0, + position: 1, + }, + ], + }); + + const harness = mountHook("ws-queue-hydration"); + + try { + await flushEffects(); + await act(async () => { + await harness.getValue().switchTopic("session-queue"); + }); + + expect(harness.getValue().queuedTurns).toEqual([ + { + queued_turn_id: "queued-1", + message_preview: "继续补充 PRD", + message_text: "继续补充 PRD,并补一版里程碑拆解", + created_at: 1700000000000, + image_count: 0, + position: 1, + }, + ]); + } finally { + harness.unmount(); + } + }); +}); + describe("useAsterAgentChat thread timeline", () => { it("sendMessage 后在首个流事件前应先注入本地回合占位", async () => { const workspaceId = "ws-thread-optimistic"; @@ -468,6 +536,185 @@ describe("useAsterAgentChat thread timeline", () => { }); }); +describe("useAsterAgentChat runtime routing", () => { + it("开启搜索能力时应提交 allowed 模式,而不是强制 required", async () => { + const workspaceId = "ws-search-mode-allowed"; + seedSession(workspaceId, "session-search-mode-allowed"); + const harness = mountHook(workspaceId); + + try { + await flushEffects(); + + await act(async () => { + await harness + .getValue() + .sendMessage("帮我看看今天的黄金价格", [], true, false, false, "react"); + }); + + expect(mockSubmitAgentRuntimeTurn).toHaveBeenCalledTimes(1); + expect(mockSubmitAgentRuntimeTurn).toHaveBeenCalledWith( + expect.objectContaining({ + message: "帮我看看今天的黄金价格", + turn_config: expect.objectContaining({ + web_search: true, + search_mode: "allowed", + }), + queue_if_busy: true, + }), + ); + } finally { + harness.unmount(); + } + }); + + it("runtime_status 与 thinking_delta 应在 final_done 前持续保留", async () => { + const workspaceId = "ws-runtime-status-stream"; + seedSession(workspaceId, "session-runtime-status-stream"); + const harness = mountHook(workspaceId); + + let streamHandler: ((event: { payload: unknown }) => void) | null = null; + mockSafeListen.mockImplementationOnce(async (_eventName, handler) => { + streamHandler = handler as (event: { payload: unknown }) => void; + return () => { + streamHandler = null; + }; + }); + + try { + await flushEffects(); + + await act(async () => { + await harness + .getValue() + .sendMessage("请先分析,再决定要不要搜索", [], true, true, false, "react"); + }); + + act(() => { + streamHandler?.({ + payload: { + type: "runtime_status", + status: { + phase: "routing", + title: "已决定:先深度思考", + detail: "先做更充分的意图理解,再决定是否调用搜索。", + checkpoints: ["thinking 已开启", "搜索与工具保持候选状态"], + }, + }, + }); + streamHandler?.({ + payload: { + type: "thinking_delta", + text: "先判断任务是直接回答还是需要联网。", + }, + }); + streamHandler?.({ + payload: { + type: "text_delta", + text: "我会先分析你的诉求。", + }, + }); + }); + + let assistantMessage = [...harness.getValue().messages] + .reverse() + .find((msg) => msg.role === "assistant"); + + expect(assistantMessage?.runtimeStatus).toMatchObject({ + phase: "routing", + title: "已决定:先深度思考", + }); + expect( + assistantMessage?.contentParts?.some( + (part) => + part.type === "thinking" && + part.text.includes("先判断任务是直接回答还是需要联网"), + ), + ).toBe(true); + expect(assistantMessage?.content).toContain("我会先分析你的诉求。"); + + act(() => { + streamHandler?.({ + payload: { + type: "final_done", + }, + }); + }); + + assistantMessage = [...harness.getValue().messages] + .reverse() + .find((msg) => msg.role === "assistant"); + + expect(assistantMessage?.runtimeStatus).toBeUndefined(); + expect(assistantMessage?.isThinking).toBe(false); + } finally { + harness.unmount(); + } + }); + + it("final_done 前未收到正文时应给出明确失败提示,而不是静默无响应", async () => { + const workspaceId = "ws-empty-final-response"; + seedSession(workspaceId, "session-empty-final-response"); + const harness = mountHook(workspaceId); + + let streamHandler: ((event: { payload: unknown }) => void) | null = null; + mockSafeListen.mockImplementationOnce(async (_eventName, handler) => { + streamHandler = handler as (event: { payload: unknown }) => void; + return () => { + streamHandler = null; + }; + }); + + try { + await flushEffects(); + + await act(async () => { + await harness + .getValue() + .sendMessage("帮我汇总今天的国际新闻", [], true, false, false, "react"); + }); + + act(() => { + streamHandler?.({ + payload: { + type: "tool_start", + tool_name: "WebSearch", + tool_id: "tool-search-1", + arguments: JSON.stringify({ query: "今天的国际新闻" }), + }, + }); + streamHandler?.({ + payload: { + type: "tool_end", + tool_id: "tool-search-1", + result: { + success: true, + output: "https://example.com/world-news", + }, + }, + }); + streamHandler?.({ + payload: { + type: "final_done", + }, + }); + }); + + const assistantMessage = [...harness.getValue().messages] + .reverse() + .find((msg) => msg.role === "assistant"); + + expect(assistantMessage?.content).toContain( + "已完成工具执行,但模型未输出最终答复,请重试。", + ); + expect(mockToast.error).toHaveBeenCalledWith( + "已完成工具执行,但模型未输出最终答复,请重试", + ); + } finally { + harness.unmount(); + } + }); +}); + describe("useAsterAgentChat slash skill 执行链路", () => { it("命中 slash skill 时应走 execute_skill 分支而非 chat_stream", async () => { const workspaceId = "ws-slash-skill"; @@ -831,7 +1078,7 @@ describe("useAsterAgentChat action_required 渲染链路", () => { } }); - it("fallback ask 在真实 request_id 未就绪前不应提交,避免卡住", async () => { + it("fallback ask 在真实 request_id 未就绪前应先记录答案,并在真实 request_id 到达后自动提交", async () => { const workspaceId = "ws-ask-fallback-pending"; seedSession(workspaceId, "session-ask-fallback-pending"); const harness = mountHook(workspaceId); @@ -876,10 +1123,58 @@ describe("useAsterAgentChat action_required 渲染链路", () => { }); }); + let assistantMessage = [...harness.getValue().messages] + .reverse() + .find((msg) => msg.role === "assistant"); + expect(mockRespondAgentRuntimeAction).not.toHaveBeenCalled(); - expect(mockToast.error).toHaveBeenCalledWith( - "Ask 请求 ID 尚未就绪,请稍后再试", + expect(mockToast.info).toHaveBeenCalledWith( + "已记录你的回答,等待系统请求就绪后自动提交", ); + expect(assistantMessage?.actionRequests?.[0]).toMatchObject({ + requestId: "fallback:tool-fallback-only", + status: "queued", + submittedResponse: '{"answer":"网络矩阵"}', + submittedUserData: { answer: "网络矩阵" }, + }); + + act(() => { + streamHandler?.({ + payload: { + type: "action_required", + request_id: "req-ask-real-1", + action_type: "ask_user", + prompt: "请选择您喜欢的科技风格类型", + questions: [ + { + question: "请选择您喜欢的科技风格类型", + options: ["网络矩阵", "极简未来"], + }, + ], + }, + }); + }); + + await flushEffects(); + + assistantMessage = [...harness.getValue().messages] + .reverse() + .find((msg) => msg.role === "assistant"); + + expect(mockRespondAgentRuntimeAction).toHaveBeenCalledWith({ + session_id: "session-ask-fallback-pending", + request_id: "req-ask-real-1", + action_type: "ask_user", + confirmed: true, + response: '{"answer":"网络矩阵"}', + user_data: { answer: "网络矩阵" }, + }); + expect( + assistantMessage?.actionRequests?.some( + (item) => + item.requestId === "req-ask-real-1" && item.status === "submitted", + ), + ).toBe(true); } finally { harness.unmount(); } @@ -1128,6 +1423,144 @@ describe("useAsterAgentChat action_required 渲染链路", () => { } }); + it("write_file 工具启动时即使没有内容也应立即创建 preparing artifact 并触发 onWriteFile", async () => { + const workspaceId = "ws-artifact-tool-start-preparing"; + seedSession(workspaceId, "session-artifact-tool-start-preparing"); + const onWriteFile = vi.fn(); + const harness = mountHook(workspaceId, { onWriteFile }); + + let streamHandler: ((event: { payload: unknown }) => void) | null = null; + mockSafeListen.mockImplementationOnce(async (_eventName, handler) => { + streamHandler = handler as (event: { payload: unknown }) => void; + return () => { + streamHandler = null; + }; + }); + + try { + await flushEffects(); + + await act(async () => { + await harness + .getValue() + .sendMessage("准备写入空文件", [], false, false, false, "react"); + }); + + act(() => { + streamHandler?.({ + payload: { + type: "tool_start", + tool_id: "tool-write-prepare-1", + tool_name: "write_file", + arguments: JSON.stringify({ + path: "notes/preparing.md", + }), + }, + }); + }); + + const assistantMessage = [...harness.getValue().messages] + .reverse() + .find((msg) => msg.role === "assistant"); + + expect(assistantMessage?.artifacts).toHaveLength(1); + expect(assistantMessage?.artifacts?.[0]).toMatchObject({ + title: "preparing.md", + content: "", + status: "streaming", + meta: expect.objectContaining({ + filePath: "notes/preparing.md", + writePhase: "preparing", + source: "tool_start", + }), + }); + expect(onWriteFile).toHaveBeenCalledWith( + "", + "notes/preparing.md", + expect.objectContaining({ + source: "tool_start", + status: "streaming", + metadata: expect.objectContaining({ + writePhase: "preparing", + lastUpdateSource: "tool_start", + }), + }), + ); + } finally { + harness.unmount(); + } + }); + + it("apply_patch 工具启动时应立即暴露目标文件,避免工作台空白等待", async () => { + const workspaceId = "ws-artifact-apply-patch"; + seedSession(workspaceId, "session-artifact-apply-patch"); + const onWriteFile = vi.fn(); + const harness = mountHook(workspaceId, { onWriteFile }); + + let streamHandler: ((event: { payload: unknown }) => void) | null = null; + mockSafeListen.mockImplementationOnce(async (_eventName, handler) => { + streamHandler = handler as (event: { payload: unknown }) => void; + return () => { + streamHandler = null; + }; + }); + + try { + await flushEffects(); + + await act(async () => { + await harness + .getValue() + .sendMessage("补丁更新文档", [], false, false, false, "react"); + }); + + act(() => { + streamHandler?.({ + payload: { + type: "tool_start", + tool_id: "tool-apply-patch-1", + tool_name: "apply_patch", + arguments: JSON.stringify({ + patch: [ + "*** Begin Patch", + "*** Update File: notes/patched.md", + "@@", + "-old", + "+new", + "*** End Patch", + ].join("\n"), + }), + }, + }); + }); + + const assistantMessage = [...harness.getValue().messages] + .reverse() + .find((msg) => msg.role === "assistant"); + + expect(assistantMessage?.artifacts?.[0]).toMatchObject({ + title: "patched.md", + content: "", + status: "streaming", + meta: expect.objectContaining({ + filePath: "notes/patched.md", + writePhase: "preparing", + source: "tool_start", + }), + }); + expect(onWriteFile).toHaveBeenCalledWith( + "", + "notes/patched.md", + expect.objectContaining({ + source: "tool_start", + status: "streaming", + }), + ); + } finally { + harness.unmount(); + } + }); + it("artifact_snapshot 完成后应在 final_done 时将 artifact 标记为 complete", async () => { const workspaceId = "ws-artifact-snapshot"; seedSession(workspaceId, "session-artifact-snapshot"); @@ -1198,6 +1631,80 @@ describe("useAsterAgentChat action_required 渲染链路", () => { harness.unmount(); } }); + + it("artifact_snapshot 到来时应复用同路径 artifact 而不是重复新增", async () => { + const workspaceId = "ws-artifact-snapshot-reuse"; + seedSession(workspaceId, "session-artifact-snapshot-reuse"); + const harness = mountHook(workspaceId); + + let streamHandler: ((event: { payload: unknown }) => void) | null = null; + mockSafeListen.mockImplementationOnce(async (_eventName, handler) => { + streamHandler = handler as (event: { payload: unknown }) => void; + return () => { + streamHandler = null; + }; + }); + + try { + await flushEffects(); + + await act(async () => { + await harness + .getValue() + .sendMessage("生成复用快照", [], false, false, false, "react"); + }); + + act(() => { + streamHandler?.({ + payload: { + type: "tool_start", + tool_id: "tool-write-reuse-1", + tool_name: "write_file", + arguments: JSON.stringify({ + path: "notes/reuse.md", + }), + }, + }); + }); + + const initialAssistantMessage = [...harness.getValue().messages] + .reverse() + .find((msg) => msg.role === "assistant"); + const initialArtifactId = initialAssistantMessage?.artifacts?.[0]?.id; + + act(() => { + streamHandler?.({ + payload: { + type: "artifact_snapshot", + artifact: { + artifactId: "server-artifact-id-1", + filePath: "notes/reuse.md", + content: "# Reused\n\nsnapshot body", + metadata: { + complete: false, + }, + }, + }, + }); + }); + + const assistantMessage = [...harness.getValue().messages] + .reverse() + .find((msg) => msg.role === "assistant"); + + expect(assistantMessage?.artifacts).toHaveLength(1); + expect(assistantMessage?.artifacts?.[0]?.id).toBe(initialArtifactId); + expect(assistantMessage?.artifacts?.[0]).toMatchObject({ + content: "# Reused\n\nsnapshot body", + meta: expect.objectContaining({ + writePhase: "streaming", + source: "artifact_snapshot", + }), + }); + } finally { + harness.unmount(); + } + }); }); describe("useAsterAgentChat 偏好持久化", () => { diff --git a/src/components/agent/chat/hooks/useAsterAgentChat.ts b/src/components/agent/chat/hooks/useAsterAgentChat.ts index b4c107369..9601d5d28 100644 --- a/src/components/agent/chat/hooks/useAsterAgentChat.ts +++ b/src/components/agent/chat/hooks/useAsterAgentChat.ts @@ -12,7 +12,10 @@ import { useAgentContext } from "./useAgentContext"; import { useAgentSession } from "./useAgentSession"; import { useAgentTools } from "./useAgentTools"; import { useAgentStream } from "./useAgentStream"; -import type { SendMessageFn, UseAsterAgentChatOptions } from "./agentChatShared"; +import type { + SendMessageFn, + UseAsterAgentChatOptions, +} from "./agentChatShared"; export type { Topic } from "./agentChatShared"; @@ -27,7 +30,8 @@ export function useAsterAgentChat(options: UseAsterAgentChatOptions) { const sendMessageRef = useRef(null); const resetPendingActionsRef = useRef<(() => void) | null>(null); const topicsUpdaterRef = useRef< - ((sessionId: string, executionStrategy: AsterExecutionStrategy) => void) | null + | ((sessionId: string, executionStrategy: AsterExecutionStrategy) => void) + | null >(null); const resetPendingActions = useCallback(() => { @@ -88,6 +92,8 @@ export function useAsterAgentChat(options: UseAsterAgentChatOptions) { setThreadItems: session.setThreadItems, setThreadTurns: session.setThreadTurns, setCurrentTurnId: session.setCurrentTurnId, + queuedTurns: session.queuedTurns, + setQueuedTurns: session.setQueuedTurns, setPendingActions: tools.setPendingActions, }); @@ -111,6 +117,37 @@ export function useAsterAgentChat(options: UseAsterAgentChatOptions) { init(); }, [runtime]); + useEffect(() => { + const refreshSessionDetail = session.refreshSessionDetail; + const activeSessionId = session.sessionId; + const queuedTurnCount = session.queuedTurns.length; + const threadTurns = session.threadTurns; + + if (!activeSessionId || stream.isSending) { + return; + } + + const hasRecoveredQueueWork = + queuedTurnCount > 0 || threadTurns.some((turn) => turn.status === "running"); + if (!hasRecoveredQueueWork) { + return; + } + + const timer = window.setInterval(() => { + void refreshSessionDetail(activeSessionId); + }, 1500); + + return () => { + window.clearInterval(timer); + }; + }, [ + session.queuedTurns.length, + session.refreshSessionDetail, + session.sessionId, + session.threadTurns, + stream.isSending, + ]); + const handleStartProcess = async () => { // Aster 不需要显式启动独立进程,初始化在 effect 中完成。 }; @@ -138,9 +175,11 @@ export function useAsterAgentChat(options: UseAsterAgentChatOptions) { currentTurnId: session.currentTurnId, turns: session.threadTurns, threadItems: session.threadItems, + queuedTurns: session.queuedTurns, isSending: stream.isSending, sendMessage: stream.sendMessage, stopSending: stream.stopSending, + removeQueuedTurn: stream.removeQueuedTurn, clearMessages: session.clearMessages, deleteMessage: session.deleteMessage, editMessage: session.editMessage, diff --git a/src/components/agent/chat/index.test.tsx b/src/components/agent/chat/index.test.tsx index 22bafd829..406e2d429 100644 --- a/src/components/agent/chat/index.test.tsx +++ b/src/components/agent/chat/index.test.tsx @@ -4,6 +4,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; const { mockUseAgentChatUnified, + mockUseArtifactAutoPreviewSync, mockUseThemeContextWorkspace, mockUseTopicBranchBoard, mockGetProject, @@ -29,6 +30,7 @@ const { mockSkillExecutionGetDetail, } = vi.hoisted(() => ({ mockUseAgentChatUnified: vi.fn(), + mockUseArtifactAutoPreviewSync: vi.fn(), mockUseThemeContextWorkspace: vi.fn(), mockUseTopicBranchBoard: vi.fn(), mockGetProject: vi.fn(), @@ -71,6 +73,7 @@ vi.mock("sonner", () => ({ vi.mock("./hooks", () => ({ useAgentChatUnified: mockUseAgentChatUnified, + useArtifactAutoPreviewSync: mockUseArtifactAutoPreviewSync, useThemeContextWorkspace: mockUseThemeContextWorkspace, useTopicBranchBoard: mockUseTopicBranchBoard, })); @@ -1919,8 +1922,8 @@ describe("AgentChatPage 自动引导", () => { toolCalls: [ { id: "tool-search-1", - name: "search_query", - arguments: JSON.stringify({ q: "Rokid Glasses 最新功能" }), + name: "WebSearch", + arguments: JSON.stringify({ query: "Rokid Glasses 最新功能" }), status: "completed", startTime: new Date("2026-03-06T11:00:01.500Z"), endTime: new Date("2026-03-06T11:00:02.000Z"), diff --git a/src/components/agent/chat/index.tsx b/src/components/agent/chat/index.tsx index a631869d4..5da63a3e4 100644 --- a/src/components/agent/chat/index.tsx +++ b/src/components/agent/chat/index.tsx @@ -25,9 +25,11 @@ import { readFilePreview } from "@/lib/api/fileBrowser"; import { uploadImageToSession, importDocument } from "@/lib/api/session-files"; import { useAgentChatUnified, + useArtifactAutoPreviewSync, useThemeContextWorkspace, useTopicBranchBoard, } from "./hooks"; +import { useArtifactDisplayState } from "./hooks/useArtifactDisplayState"; import type { SidebarActivityLog } from "./hooks/useThemeContextWorkspace"; import type { TopicBranchStatus } from "./hooks/useTopicBranchBoard"; import { useSessionFiles } from "./hooks/useSessionFiles"; @@ -79,7 +81,11 @@ import { selectedArtifactAtom, selectedArtifactIdAtom, } from "@/lib/artifact/store"; -import { ArtifactRenderer, ArtifactToolbar } from "@/components/artifact"; +import { + ArtifactCanvasOverlay, + ArtifactRenderer, + ArtifactToolbar, +} from "@/components/artifact"; import type { Artifact } from "@/lib/artifact/types"; import { useAtomValue, useSetAtom } from "jotai"; import { createInitialMusicState } from "@/components/content-creator/canvas/music/types"; @@ -119,7 +125,6 @@ import { skillsApi, type Skill } from "@/lib/api/skills"; import { buildHomeAgentParams } from "@/lib/workspace/navigation"; import { useConfiguredProviders } from "@/hooks/useConfiguredProviders"; import { useSubAgentScheduler } from "@/hooks/useSubAgentScheduler"; -import { LatestRunStatusBadge } from "@/components/execution/LatestRunStatusBadge"; import { executionRunGet, executionRunGetThemeWorkbenchState, @@ -260,9 +265,11 @@ function mergeMessageArtifactsIntoStore( } const shouldReuseExistingContent = - artifact.content.length === 0 && - artifact.meta.source === "tool_result" && - existing.content.length > 0; + existing.content.length > 0 && + (artifact.content.length === 0 || + (artifact.status === "streaming" && + artifact.content.length < existing.content.length && + existing.content.startsWith(artifact.content))); return { ...existing, @@ -621,12 +628,7 @@ function resolveThemeWorkbenchToolTaskTitle(toolCall: ToolCallState): string { ? `写入 ${getThemeWorkbenchFileLabel(pathValue)}` : "写入主稿文件"; } - if ( - normalized.includes("websearch") || - normalized.includes("search_query") || - normalized.includes("web_search") || - normalized.includes("search") - ) { + if (normalized.includes("websearch")) { return queryValue ? `检索 ${truncateThemeWorkbenchLabel(queryValue)}` : "检索参考资料"; @@ -2045,6 +2047,15 @@ export function AgentChatPage({ const selectedArtifact = useAtomValue(selectedArtifactAtom); const setArtifacts = useSetAtom(artifactsAtom); const setSelectedArtifactId = useSetAtom(selectedArtifactIdAtom); + const liveArtifact = useMemo( + () => + selectedArtifact || + (artifacts.length > 0 ? artifacts[artifacts.length - 1] : null), + [artifacts, selectedArtifact], + ); + const artifactDisplayState = useArtifactDisplayState(liveArtifact, artifacts); + const currentCanvasArtifact = artifactDisplayState.liveArtifact; + const displayedCanvasArtifact = artifactDisplayState.displayArtifact; // Artifact 预览状态 const [artifactViewMode, setArtifactViewMode] = useState< @@ -2376,9 +2387,11 @@ export function AgentChatPage({ currentTurnId, turns, threadItems, + queuedTurns = [], isSending, sendMessage, stopSending, + removeQueuedTurn = async () => false, clearMessages, deleteMessage, editMessage, @@ -2492,11 +2505,11 @@ export function AgentChatPage({ }, [activeTheme, artifacts, selectedArtifact, setSelectedArtifactId]); useEffect(() => { - if (activeTheme !== "general" || !selectedArtifact) { + if (activeTheme !== "general" || !displayedCanvasArtifact) { return; } - setArtifactViewMode(resolveDefaultArtifactViewMode(selectedArtifact)); - }, [activeTheme, selectedArtifact]); + setArtifactViewMode(resolveDefaultArtifactViewMode(displayedCanvasArtifact)); + }, [activeTheme, displayedCanvasArtifact]); useEffect(() => { savePersistedBoolean(HARNESS_PANEL_VISIBILITY_KEY, harnessPanelVisible); @@ -4906,12 +4919,36 @@ export function AgentChatPage({ // General 主题使用专门的画布处理 if (activeTheme === "general") { + const existingArtifact = artifacts.find((artifact) => { + if (context?.artifactId && artifact.id === context.artifactId) { + return true; + } + + if (context?.artifact?.id && artifact.id === context.artifact.id) { + return true; + } + + return ( + typeof artifact.meta.filePath === "string" && + artifact.meta.filePath === fileName + ); + }); + const nextContent = + content.length > 0 + ? content + : context?.artifact?.content || existingArtifact?.content || ""; const nextArtifact = context?.artifact ? { + ...(existingArtifact || {}), ...context.artifact, - content: content || context.artifact.content, - status: context.status || context.artifact.status, + content: nextContent, + status: + context.status || + context.artifact.status || + existingArtifact?.status || + "pending", meta: { + ...(existingArtifact?.meta || {}), ...context.artifact.meta, ...(context.metadata || {}), }, @@ -4919,15 +4956,18 @@ export function AgentChatPage({ } : buildArtifactFromWrite({ filePath: fileName, - content, + content: nextContent, context: { ...context, - status: context?.status || (content.length > 0 ? "complete" : "pending"), + artifact: existingArtifact, + status: + context?.status || + (nextContent.length > 0 ? "complete" : "pending"), }, }); - if (content.length > 0) { - saveSessionFile(fileName, content).catch((error) => { + if (nextContent.length > 0) { + saveSessionFile(fileName, nextContent).catch((error) => { console.error("[AgentChatPage] 持久化 artifact 失败:", error); }); } @@ -5310,6 +5350,7 @@ export function AgentChatPage({ }, [ activeTheme, // 添加 activeTheme 依赖 + artifacts, setArtifactViewMode, setSelectedArtifactId, currentGate.key, @@ -5399,6 +5440,13 @@ export function AgentChatPage({ [readSessionFile, sessionFiles, taskFiles], ); + useArtifactAutoPreviewSync({ + enabled: activeTheme === "general", + artifact: currentCanvasArtifact, + loadPreview: handleHarnessLoadFilePreview, + onSyncArtifact: upsertGeneralArtifact, + }); + const openArtifactInWorkbench = useCallback( async (artifact: Artifact) => { let nextArtifact = artifact; @@ -5953,8 +6001,10 @@ export function AgentChatPage({ return ( 0} providerType={providerType} setProviderType={setProviderType} model={model} @@ -6332,6 +6382,8 @@ export function AgentChatPage({ onToolStatesChange={setChatToolPreferences} onSelectCharacter={handleSelectCharacter} onNavigateToSettings={handleNavigateToSkillSettings} + queuedTurns={queuedTurns} + onRemoveQueuedTurn={removeQueuedTurn} /> ), [ @@ -6348,6 +6400,7 @@ export function AgentChatPage({ handleToggleCanvas, handleToggleTaskFiles, input, + queuedTurns, isSending, isThemeWorkbench, layoutMode, @@ -6355,6 +6408,7 @@ export function AgentChatPage({ projectId, projectMemory?.characters, providerType, + removeQueuedTurn, setExecutionStrategy, setInput, setModel, @@ -6381,31 +6435,33 @@ export function AgentChatPage({ return ( - - } - onLoadFilePreview={handleHarnessLoadFilePreview} - onOpenFile={handleFileClick} - /> + + } + onLoadFilePreview={handleHarnessLoadFilePreview} + onOpenFile={handleFileClick} + /> ); @@ -6697,32 +6753,44 @@ export function AgentChatPage({ ) as ThemeType; // 如果有 artifact,优先使用 ArtifactRenderer 渲染 - const currentArtifact = - selectedArtifact || - (artifacts.length > 0 ? artifacts[artifacts.length - 1] : null); - if (renderCanvasTheme === "general" && currentArtifact) { + if ( + renderCanvasTheme === "general" && + currentCanvasArtifact && + displayedCanvasArtifact + ) { return (
-
+
+ {artifactDisplayState.overlay ? ( + + ) : null}
@@ -6811,8 +6879,10 @@ export function AgentChatPage({ } return null; }, [ - artifacts, - selectedArtifact, + currentCanvasArtifact, + displayedCanvasArtifact, + artifactDisplayState.overlay, + artifactDisplayState.showPreviousVersionBadge, generalCanvasState, canvasState, resolvedCanvasState, @@ -6890,14 +6960,6 @@ export function AgentChatPage({ } /> - {!isThemeWorkbench ? ( - - ) : null} - {!isThemeWorkbench && contentId && syncStatus !== "idle" && (
{ + writePhase?: ArtifactWritePhase; + previewText?: string; + latestChunk?: string; + isPartial?: boolean; + lastUpdateSource?: ArtifactWriteSource; +} + export interface WriteArtifactContext { artifact?: Artifact; artifactId?: string; - source?: "tool_start" | "artifact_snapshot" | "tool_result" | "message_content"; + source?: ArtifactWriteSource; sourceMessageId?: string; status?: ArtifactStatus; - metadata?: Record; + metadata?: ArtifactWriteMetadata; } export interface AgentRuntimeStatus { diff --git a/src/components/agent/chat/utils/agentRuntimeStatus.ts b/src/components/agent/chat/utils/agentRuntimeStatus.ts index 0d46e36a9..4a6ee689b 100644 --- a/src/components/agent/chat/utils/agentRuntimeStatus.ts +++ b/src/components/agent/chat/utils/agentRuntimeStatus.ts @@ -23,7 +23,7 @@ export function buildInitialAgentRuntimeStatus(options: { }): AgentRuntimeStatus { const checkpoints = [ buildExecutionLabel(options.executionStrategy), - options.webSearch ? "已允许联网检索" : "优先本地直接回答", + options.webSearch ? "联网搜索仅作为候选能力待命" : "优先本地直接回答", options.thinking ? "必要时启用深度思考" : "先走轻量推理", options.skipUserMessage ? "系统引导请求" : "用户请求已入队", ]; @@ -44,7 +44,7 @@ export function buildWaitingAgentRuntimeStatus(options: { const checkpoints = [ "会话已建立", buildExecutionLabel(options.executionStrategy), - options.webSearch ? "检索能力待命" : "直接回答优先", + options.webSearch ? "先理解意图,再决定是否联网" : "直接回答优先", options.thinking ? "推理增强已待命" : "等待首个模型事件", ]; diff --git a/src/components/agent/chat/utils/generalAgentPrompt.test.ts b/src/components/agent/chat/utils/generalAgentPrompt.test.ts index 62f51784a..2d541940d 100644 --- a/src/components/agent/chat/utils/generalAgentPrompt.test.ts +++ b/src/components/agent/chat/utils/generalAgentPrompt.test.ts @@ -41,6 +41,8 @@ describe("generalAgentPrompt", () => { expect(prompt).toContain("执行车道"); expect(prompt).toContain("后台任务:已开启"); expect(prompt).toContain("多代理:已开启"); + expect(prompt).toContain("统一使用 WebSearch"); + expect(prompt).toContain("不要混用 search/search_query/tool_search"); }); it("知识主题 Prompt 应强调事实与时效性", () => { @@ -49,6 +51,7 @@ describe("generalAgentPrompt", () => { expect(prompt).toContain("知识探索"); expect(prompt).toContain("区分事实、推断与不确定性"); expect(prompt).toContain("优先核对时间与来源"); + expect(prompt).toContain("3-4 组 WebSearch 扩搜"); }); it("计划主题 Prompt 应强调执行节奏与风险", () => { diff --git a/src/components/agent/chat/utils/generalAgentPrompt.ts b/src/components/agent/chat/utils/generalAgentPrompt.ts index 3674c973e..0ca9d6bda 100644 --- a/src/components/agent/chat/utils/generalAgentPrompt.ts +++ b/src/components/agent/chat/utils/generalAgentPrompt.ts @@ -1,7 +1,11 @@ import type { ThemeType } from "@/components/content-creator/types"; import type { ChatToolPreferences } from "./chatToolPreferences"; -const GENERAL_AGENT_THEMES = new Set(["general", "knowledge", "planning"]); +const GENERAL_AGENT_THEMES = new Set([ + "general", + "knowledge", + "planning", +]); const GENERAL_THEME_LABELS: Record = { general: "通用对话", @@ -71,11 +75,7 @@ export function buildGeneralAgentSystemPrompt( theme: ThemeType | string = "general", options: GeneralAgentPromptOptions = {}, ): string { - const { - now = new Date(), - toolPreferences, - harness, - } = options; + const { now = new Date(), toolPreferences, harness } = options; const normalizedTheme = theme.trim().toLowerCase(); const themeLabel = GENERAL_THEME_LABELS[normalizedTheme] || GENERAL_THEME_LABELS.general; @@ -120,23 +120,24 @@ ${toolPreferenceLines} 执行车道: - 直接回答:适用于多数问答、改写、总结、解释、比较、建议、轻量规划。 -- 联网检索:适用于用户明确要求搜索,或问题涉及最新、实时、价格、政策、规则、版本、新闻、日期敏感信息。 +- 联网检索:适用于用户明确要求搜索,或问题涉及最新、实时、价格、政策、规则、版本、新闻、日期敏感信息。统一使用 WebSearch 作为检索入口;需要打开具体页面时再使用 WebFetch。 - 深度思考:适用于复杂推理、强约束规划、多方案取舍、高风险判断。 - 后台任务:适用于耗时生成、需要排队、异步产出或用户明确要求后台推进。 - 多代理:适用于任务天然可拆分、需要并行探索,或主线程上下文会显著过载。 工具使用原则: 1. 不要为了显得像 agent 而强行调用工具;直接回答更合适时,就直接回答。 -2. 只有在以下场景才主动联网核实:用户明确要求搜索;问题涉及今天、最新、价格、政策、法律、版本、新闻、实时数据;或者高风险信息需要校验。 +2. 只有在以下场景才主动联网核实:用户明确要求搜索;问题涉及今天、最新、价格、政策、法律、版本、新闻、实时数据;或者高风险信息需要校验。联网时统一使用 WebSearch,不要混用 search/search_query/tool_search 之类别名。 3. 深度思考默认只用于复杂推理、多方案比较、严格规划或高风险判断;简单问答、轻量改写、普通说明不要默认进入长链路推理。 4. 只有当任务长耗时、异步生成、跨模态产出、需要排队执行,或用户明确要求后台推进时,再升级为 task;否则优先在当前回合直接完成。 5. 只有当问题天然可拆分、需要并行探索/规划/执行、或主线程上下文会显著过载时,再使用 subagent;否则优先单 agent 完成。 6. 用户明确要求读取、修改、创建、保存项目或工作区内容时,再使用文件或工作区能力;否则默认以对话结果为主,不主动落盘。 -7. 如果用户开启联网搜索或明确要求检索,先查再答;如果搜索结果不足,要明确说明不足点,不要假装确定。 -8. 遇到 ask_user、elicitation、权限确认等 action_required 流程时,暂停推进并请求最小必要信息,不要伪装成已完成。 -9. 不输出原始思维链路;只输出结论、关键依据、必要假设、来源时间和可执行下一步。 -10. 如果用户只是要一个答案、草稿、提纲、比较或总结,不要擅自把问题升级成项目制流程。 -11. 如果要给出计划,默认同时给优先级、阶段划分、约束、风险和下一步动作,而不是抽象口号。 +7. 如果用户明确要求检索,或问题涉及最新、实时、价格、政策、法律、版本、新闻、日期敏感信息,先核实再答;仅仅开启联网搜索能力不等于必须联网。 +8. 新闻、最新动态、某日综述、热点盘点类请求,不要只做一次浅搜;至少围绕原始 query、中文日期/主题 query、英文等价 query、headlines/roundup query 做 3-4 组 WebSearch 扩搜,再按主题聚类总结。 +9. 遇到 ask_user、elicitation、权限确认等 action_required 流程时,暂停推进并请求最小必要信息,不要伪装成已完成。 +10. 不输出原始思维链路;只输出结论、关键依据、必要假设、来源时间和可执行下一步。 +11. 如果用户只是要一个答案、草稿、提纲、比较或总结,不要擅自把问题升级成项目制流程。 +12. 如果要给出计划,默认同时给优先级、阶段划分、约束、风险和下一步动作,而不是抽象口号。 行为协议: - 先判断应该走哪条车道:直接回答 / 联网检索 / 深度思考 / 后台任务 / 多代理。 diff --git a/src/components/agent/chat/utils/harnessState.test.ts b/src/components/agent/chat/utils/harnessState.test.ts index 9ec6a8289..af327e203 100644 --- a/src/components/agent/chat/utils/harnessState.test.ts +++ b/src/components/agent/chat/utils/harnessState.test.ts @@ -78,4 +78,98 @@ describe("deriveHarnessSessionState", () => { expect(state.outputSignals[0]?.artifactPath).toBe("workspace/plan.md"); expect(state.latestContextTrace).toHaveLength(1); }); + + it("应从消息 artifacts 提取当前文件写入状态", () => { + const messages = [ + createMessage({ + artifacts: [ + { + id: "artifact-live-1", + type: "document", + title: "live.md", + content: "# 实时草稿\n\n正在写入第二段", + status: "streaming", + meta: { + filePath: "workspace/live.md", + writePhase: "streaming", + previewText: "# 实时草稿\n\n正在写入第二段", + latestChunk: "正在写入第二段", + lastUpdateSource: "artifact_snapshot", + }, + position: { start: 0, end: 12 }, + createdAt: Date.now() - 1000, + updatedAt: Date.now(), + }, + ], + }), + ]; + + const state = deriveHarnessSessionState(messages, []); + + expect(state.activeFileWrites).toHaveLength(1); + expect(state.activeFileWrites[0]).toMatchObject({ + path: "workspace/live.md", + displayName: "live.md", + phase: "streaming", + source: "artifact_snapshot", + }); + expect(state.activeFileWrites[0]?.preview).toContain("实时草稿"); + }); + + it("应为搜索工具调用生成工作台可消费的搜索输出信号", () => { + const messages = [ + createMessage({ + toolCalls: [ + { + id: "tool-search-1", + name: "WebSearch", + arguments: JSON.stringify({ query: "3月13日国际新闻" }), + status: "completed", + result: { + success: true, + output: [ + "Xinhua world news summary at 0030 GMT, March 13", + "https://example.com/xinhua", + "全球要闻摘要,覆盖国际局势与市场动态。", + ].join("\n"), + }, + startTime: new Date("2026-03-13T12:00:00.000Z"), + endTime: new Date("2026-03-13T12:00:03.000Z"), + }, + ], + }), + ]; + + const state = deriveHarnessSessionState(messages, []); + + expect(state.outputSignals).toHaveLength(1); + expect(state.outputSignals[0]).toMatchObject({ + title: "联网检索摘要", + summary: "3月13日国际新闻", + }); + expect(state.outputSignals[0]?.preview).toContain("Xinhua world news summary"); + expect(state.outputSignals[0]?.content).toContain("https://example.com/xinhua"); + }); + + it("应保留最近 8 条输出信号以承载多组 WebSearch 扩搜", () => { + const toolCalls = Array.from({ length: 9 }, (_, index) => ({ + id: `tool-search-${index + 1}`, + name: "WebSearch", + arguments: JSON.stringify({ query: `query-${index + 1}` }), + status: "completed" as const, + result: { + success: true, + output: `结果 ${index + 1}\nhttps://example.com/${index + 1}`, + }, + startTime: new Date(`2026-03-13T12:00:0${Math.min(index, 8)}.000Z`), + endTime: new Date(`2026-03-13T12:00:1${Math.min(index, 8)}.000Z`), + })); + const messages = [createMessage({ toolCalls })]; + + const state = deriveHarnessSessionState(messages, []); + + expect(state.outputSignals).toHaveLength(8); + expect(state.outputSignals[0]?.summary).toBe("query-9"); + expect(state.outputSignals[7]?.summary).toBe("query-2"); + }); }); diff --git a/src/components/agent/chat/utils/harnessState.ts b/src/components/agent/chat/utils/harnessState.ts index 733171ae4..672334239 100644 --- a/src/components/agent/chat/utils/harnessState.ts +++ b/src/components/agent/chat/utils/harnessState.ts @@ -3,7 +3,12 @@ import type { ContextTraceStep, ToolCallState, } from "@/lib/api/agentStream"; +import type { ArtifactStatus } from "@/lib/artifact/types"; import type { ActionRequired, AgentRuntimeStatus, Message } from "../types"; +import { + resolveArtifactPreviewText, + resolveArtifactWritePhase, +} from "./messageArtifacts"; export type HarnessTodoStatus = "pending" | "in_progress" | "completed"; export type HarnessPlanPhase = "idle" | "planning" | "ready"; @@ -47,6 +52,7 @@ export interface HarnessOutputSignal { title: string; summary: string; preview?: string; + content?: string; outputFile?: string; offloadFile?: string; artifactPath?: string; @@ -90,6 +96,19 @@ export interface HarnessFileEvent { clickable: boolean; } +export interface HarnessActiveFileWrite { + id: string; + path: string; + displayName: string; + phase: NonNullable>; + status: ArtifactStatus; + source?: string; + updatedAt?: Date; + preview?: string; + latestChunk?: string; + content?: string; +} + export interface HarnessSessionState { runtimeStatus: AgentRuntimeStatus | null; pendingApprovals: ActionRequired[]; @@ -98,6 +117,7 @@ export interface HarnessSessionState { activity: HarnessToolActivity; delegatedTasks: HarnessDelegatedTask[]; outputSignals: HarnessOutputSignal[]; + activeFileWrites: HarnessActiveFileWrite[]; recentFileEvents: HarnessFileEvent[]; hasSignals: boolean; } @@ -145,6 +165,7 @@ const FILESYSTEM_TOOL_NAMES = new Set([ ]); const WEB_TOOL_RE = /^(websearch|webfetch)|browser|playwright/i; +const HARNESS_OUTPUT_SIGNAL_LIMIT = 8; const SKILL_TOOL_NAMES = new Set(["skill", "threestageworkflow"]); const PROXYCAST_TOOL_METADATA_BEGIN = "[ProxyCast 工具元数据开始]"; const PROXYCAST_TOOL_METADATA_END = "[ProxyCast 工具元数据结束]"; @@ -397,6 +418,92 @@ function maybeKeepTextContent(raw?: string): string | undefined { return normalized; } +function extractActiveFileWrites( + messages: Message[], +): HarnessActiveFileWrite[] { + const activeWrites = new Map(); + + for (const message of messages) { + for (const artifact of message.artifacts || []) { + const phase = resolveArtifactWritePhase(artifact); + if (!phase || phase === "completed") { + continue; + } + + const path = + typeof artifact.meta.filePath === "string" && + artifact.meta.filePath.trim() + ? artifact.meta.filePath.trim() + : typeof artifact.meta.filename === "string" && + artifact.meta.filename.trim() + ? artifact.meta.filename.trim() + : artifact.title; + if (!path) { + continue; + } + + const updatedAt = + Number.isFinite(artifact.updatedAt) && artifact.updatedAt > 0 + ? new Date(artifact.updatedAt) + : undefined; + const preview = buildTextPreview( + typeof artifact.meta.previewText === "string" + ? artifact.meta.previewText + : resolveArtifactPreviewText(artifact), + { + maxLines: 4, + maxChars: 240, + }, + ); + const latestChunk = buildTextPreview( + typeof artifact.meta.latestChunk === "string" + ? artifact.meta.latestChunk + : undefined, + { + maxLines: 3, + maxChars: 180, + }, + ); + const nextWrite: HarnessActiveFileWrite = { + id: artifact.id, + path, + displayName: fileNameFromPath(path), + phase, + status: artifact.status, + source: + typeof artifact.meta.lastUpdateSource === "string" + ? artifact.meta.lastUpdateSource + : typeof artifact.meta.source === "string" + ? artifact.meta.source + : undefined, + updatedAt, + preview, + latestChunk, + content: maybeKeepTextContent(artifact.content), + }; + const previous = activeWrites.get(nextWrite.id); + if (!previous) { + activeWrites.set(nextWrite.id, nextWrite); + continue; + } + + const previousTime = previous.updatedAt?.getTime() ?? 0; + const nextTime = nextWrite.updatedAt?.getTime() ?? 0; + if (nextTime >= previousTime) { + activeWrites.set(nextWrite.id, nextWrite); + } + } + } + + return Array.from(activeWrites.values()) + .sort((left, right) => { + const leftTime = left.updatedAt?.getTime() ?? 0; + const rightTime = right.updatedAt?.getTime() ?? 0; + return rightTime - leftTime; + }) + .slice(0, 5); +} + function extractMetadata( toolCall: ToolCallState, ): Record | null { @@ -473,7 +580,9 @@ function pickFirstPath(value: unknown): string | undefined { return undefined; } -function extractPathFromRecord(record: Record | null): string | undefined { +function extractPathFromRecord( + record: Record | null, +): string | undefined { if (!record) { return undefined; } @@ -529,6 +638,31 @@ function extractContentFromRecord( return undefined; } +function extractSearchQuery( + record: Record | null, +): string | undefined { + if (!record) { + return undefined; + } + + for (const key of [ + "q", + "query", + "question", + "search", + "keywords", + "keyword", + "url", + ]) { + const value = record[key]; + if (typeof value === "string" && value.trim()) { + return value.trim(); + } + } + + return undefined; +} + function resolveFileKind( path: string, preferred?: HarnessFileKind, @@ -577,16 +711,9 @@ function resolveFileKind( } if ( - [ - "md", - "markdown", - "txt", - "pdf", - "doc", - "docx", - "csv", - "rtf", - ].includes(extension) + ["md", "markdown", "txt", "pdf", "doc", "docx", "csv", "rtf"].includes( + extension, + ) ) { return "document"; } @@ -610,6 +737,8 @@ function extractOutputSignal( if (!toolCall.result) return null; const metadata = extractMetadata(toolCall); + const argumentsRecord = asRecord(parseJsonValue(toolCall.arguments)); + const normalizedName = normalizeToolName(toolCall.name); const output = toolCall.result.output; const outputFile = normalizeString(metadata?.output_file) || @@ -656,6 +785,11 @@ function extractOutputSignal( ); const offloadTrigger = normalizeString(metadata?.offload_trigger); const preview = buildTextPreview(output); + const content = maybeKeepTextContent(output); + const searchQuery = + extractSearchQuery(argumentsRecord) || + extractSearchQuery(metadata) || + extractRegexValue(/^(?:query|q|搜索词|检索词):\s*(.+)$/im, output); if ( !outputFile && @@ -668,6 +802,19 @@ function extractOutputSignal( !truncated && !offloaded ) { + if (WEB_TOOL_RE.test(normalizedName) && (preview || content)) { + const queryLabel = searchQuery || toolCall.name; + const searchLike = !/^https?:\/\//i.test(queryLabel); + return { + id: `${toolCall.id}:output-signal`, + toolCallId: toolCall.id, + toolName: toolCall.name, + title: searchLike ? "联网检索摘要" : "网页访问摘要", + summary: queryLabel, + preview, + content, + }; + } return null; } @@ -745,6 +892,7 @@ function extractOutputSignal( title, summary: summaryParts.join(" / ") || "存在可观测输出信号", preview, + content, outputFile, offloadFile, artifactPath, @@ -781,13 +929,14 @@ function extractFileEventFromToolCall( } const timestamp = - normalizeDate(toolCall.endTime) ?? normalizeDate(toolCall.startTime) ?? undefined; - const action: HarnessFileAction = - normalizedName.startsWith("read") - ? "read" - : normalizedName.includes("edit") - ? "edit" - : "write"; + normalizeDate(toolCall.endTime) ?? + normalizeDate(toolCall.startTime) ?? + undefined; + const action: HarnessFileAction = normalizedName.startsWith("read") + ? "read" + : normalizedName.includes("edit") + ? "edit" + : "write"; const sourceContent = action === "read" ? toolCall.result?.output @@ -817,7 +966,9 @@ function extractFileEventsFromOutputSignal( toolCall: ToolCallState, ): HarnessFileEvent[] { const timestamp = - normalizeDate(toolCall.endTime) ?? normalizeDate(toolCall.startTime) ?? undefined; + normalizeDate(toolCall.endTime) ?? + normalizeDate(toolCall.startTime) ?? + undefined; const events: HarnessFileEvent[] = []; if (signal.outputFile) { @@ -958,9 +1109,7 @@ function parsePlanTextToTodoItems(text: string): HarnessTodoItem[] { .split(/\r?\n/) .map((line) => line.trim()) .filter(Boolean) - .filter((line) => - /^(\d+\.\s+|[-*]\s+|\[[ xX]\]\s+)/.test(line), - ) + .filter((line) => /^(\d+\.\s+|[-*]\s+|\[[ xX]\]\s+)/.test(line)) .map((line, index) => { const completed = /^\[[xX]\]/.test(line); return { @@ -1036,7 +1185,9 @@ function deriveHarnessSessionStateFromItems( const safePendingApprovals = Array.isArray(pendingApprovals) ? pendingApprovals : []; - const sortedItems = [...items].sort((left, right) => itemTimestamp(left) - itemTimestamp(right)); + const sortedItems = [...items].sort( + (left, right) => itemTimestamp(left) - itemTimestamp(right), + ); const latestContextTrace = [...messages] .reverse() @@ -1046,6 +1197,7 @@ function deriveHarnessSessionStateFromItems( message.contextTrace.length > 0, )?.contextTrace || []; const runtimeStatus = extractLatestRuntimeStatus(messages); + const activeFileWrites = extractActiveFileWrites(messages); const activity: HarnessToolActivity = { planning: 0, @@ -1077,7 +1229,8 @@ function deriveHarnessSessionStateFromItems( kind: resolveFileKind(item.path, "artifact"), action: "persist", sourceToolName: "Artifact", - timestamp: normalizeDate(item.completed_at || item.updated_at) ?? undefined, + timestamp: + normalizeDate(item.completed_at || item.updated_at) ?? undefined, preview: buildTextPreview(item.content), content: maybeKeepTextContent(item.content), clickable: true, @@ -1089,6 +1242,7 @@ function deriveHarnessSessionStateFromItems( title: "产物已写入", summary: fileNameFromPath(item.path), preview: buildTextPreview(item.content), + content: maybeKeepTextContent(item.content), artifactPath: item.path, }); break; @@ -1102,6 +1256,7 @@ function deriveHarnessSessionStateFromItems( title: "命令执行摘要", summary: item.command, preview: buildTextPreview(item.aggregated_output), + content: maybeKeepTextContent(item.aggregated_output), exitCode: item.exit_code, }); break; @@ -1114,19 +1269,41 @@ function deriveHarnessSessionStateFromItems( title: "联网检索摘要", summary: item.query || "联网检索", preview: buildTextPreview(item.output), + content: maybeKeepTextContent(item.output), + }); + break; + case "turn_summary": + outputSignals.push({ + id: `${item.id}:summary`, + toolCallId: item.id, + toolName: "turn_summary", + title: "回合决策摘要", + summary: item.text.split(/\r?\n/)[0] || "回合决策", + preview: buildTextPreview(item.text), + content: maybeKeepTextContent(item.text), }); break; case "tool_call": { const normalizedName = normalizeToolName(item.tool_name); classifyToolActivity(activity, normalizedName); const artifactPath = pickItemPath(item); + const argumentRecord = asRecord(item.arguments); + const queryLabel = extractSearchQuery(argumentRecord); + const searchLike = WEB_TOOL_RE.test(normalizedName); outputSignals.push({ id: `${item.id}:tool`, toolCallId: item.id, toolName: item.tool_name, - title: artifactPath ? "产物已写入" : "工具执行摘要", - summary: artifactPath || item.tool_name, + title: artifactPath + ? "产物已写入" + : searchLike + ? /^https?:\/\//i.test(queryLabel || "") + ? "网页访问摘要" + : "联网检索摘要" + : "工具执行摘要", + summary: artifactPath || queryLabel || item.tool_name, preview: buildTextPreview(item.output), + content: maybeKeepTextContent(item.output), artifactPath, }); if (artifactPath) { @@ -1138,7 +1315,8 @@ function deriveHarnessSessionStateFromItems( kind: resolveFileKind(artifactPath, "artifact"), action: "persist", sourceToolName: item.tool_name, - timestamp: normalizeDate(item.completed_at || item.updated_at) ?? undefined, + timestamp: + normalizeDate(item.completed_at || item.updated_at) ?? undefined, preview: buildTextPreview(item.output), content: maybeKeepTextContent(item.output), clickable: true, @@ -1187,12 +1365,11 @@ function deriveHarnessSessionStateFromItems( } } - const planPhase: HarnessPlanPhase = - !latestPlanItem - ? "idle" - : latestPlanItem.status === "completed" - ? "ready" - : "planning"; + const planPhase: HarnessPlanPhase = !latestPlanItem + ? "idle" + : latestPlanItem.status === "completed" + ? "ready" + : "planning"; const hasSignals = runtimeStatus !== null || mergedApprovals.length > 0 || @@ -1200,6 +1377,7 @@ function deriveHarnessSessionStateFromItems( planItems.length > 0 || delegatedTasks.length > 0 || outputSignals.length > 0 || + activeFileWrites.length > 0 || recentFileEvents.length > 0 || Object.values(activity).some((count) => count > 0); @@ -1214,7 +1392,8 @@ function deriveHarnessSessionStateFromItems( }, activity, delegatedTasks: delegatedTasks.slice(-5).reverse(), - outputSignals: outputSignals.slice(-5).reverse(), + outputSignals: outputSignals.slice(-HARNESS_OUTPUT_SIGNAL_LIMIT).reverse(), + activeFileWrites, recentFileEvents: recentFileEvents .sort((left, right) => { const leftTime = left.timestamp?.getTime() ?? 0; @@ -1234,6 +1413,7 @@ function deriveHarnessSessionStateFromMessages( ? pendingApprovals : []; const runtimeStatus = extractLatestRuntimeStatus(messages); + const activeFileWrites = extractActiveFileWrites(messages); const toolCalls = collectToolCalls(messages); const activity: HarnessToolActivity = { planning: 0, @@ -1350,6 +1530,7 @@ function deriveHarnessSessionStateFromMessages( latestTodoItems.length > 0 || delegatedTasks.length > 0 || outputSignals.length > 0 || + activeFileWrites.length > 0 || recentFileEvents.length > 0 || Object.values(activity).some((count) => count > 0); @@ -1364,7 +1545,8 @@ function deriveHarnessSessionStateFromMessages( }, activity, delegatedTasks: delegatedTasks.slice(-5).reverse(), - outputSignals: outputSignals.slice(-5).reverse(), + outputSignals: outputSignals.slice(-HARNESS_OUTPUT_SIGNAL_LIMIT).reverse(), + activeFileWrites, recentFileEvents, hasSignals, }; diff --git a/src/components/agent/chat/utils/messageArtifacts.ts b/src/components/agent/chat/utils/messageArtifacts.ts index 572a0ba2e..0b099e7ce 100644 --- a/src/components/agent/chat/utils/messageArtifacts.ts +++ b/src/components/agent/chat/utils/messageArtifacts.ts @@ -4,7 +4,12 @@ import type { ArtifactStatus, ArtifactType, } from "@/lib/artifact/types"; -import type { Message, WriteArtifactContext } from "../types"; +import type { + ArtifactWriteMetadata, + ArtifactWritePhase, + Message, + WriteArtifactContext, +} from "../types"; const MARKDOWN_EXTENSIONS = new Set([ "md", @@ -60,6 +65,7 @@ const ARTIFACT_TYPE_ALIASES: Record = { svg: "svg", text: "document", }; +const ARTIFACT_PREVIEW_MAX_CHARS = 240; function normalizePath(path: string): string { return path.replace(/\\/g, "/").trim(); @@ -88,6 +94,104 @@ function readStringValue( return typeof value === "string" && value.trim() ? value.trim() : undefined; } +function normalizePreviewText( + value: string, + maxChars = ARTIFACT_PREVIEW_MAX_CHARS, +): string { + const normalized = value.trim(); + if (normalized.length <= maxChars) { + return normalized; + } + return `${normalized.slice(0, maxChars).trimEnd()}…`; +} + +export function resolveArtifactWritePhase( + artifact: Pick, +): ArtifactWritePhase | null { + const explicitPhase = + typeof artifact.meta.writePhase === "string" + ? (artifact.meta.writePhase as ArtifactWritePhase) + : null; + if (explicitPhase) { + return explicitPhase; + } + + switch (artifact.status) { + case "pending": + return "preparing"; + case "streaming": + return "streaming"; + case "complete": + return "completed"; + case "error": + return "failed"; + default: + return null; + } +} + +export function formatArtifactWritePhaseLabel( + phase: ArtifactWritePhase | null, +): string { + switch (phase) { + case "preparing": + return "准备写入"; + case "streaming": + return "正在写入"; + case "persisted": + return "已落盘"; + case "completed": + return "已完成"; + case "failed": + return "失败"; + default: + return "待处理"; + } +} + +export function resolveArtifactPreviewText( + artifact: Pick, + maxChars = ARTIFACT_PREVIEW_MAX_CHARS, +): string | undefined { + const previewCandidate = + typeof artifact.meta.previewText === "string" && artifact.meta.previewText.trim() + ? artifact.meta.previewText + : artifact.content; + + if (typeof previewCandidate !== "string" || !previewCandidate.trim()) { + return undefined; + } + + return normalizePreviewText(previewCandidate, maxChars); +} + +export function findMessageArtifact( + message: Pick, + options: { + artifactId?: string; + filePath?: string; + }, +): Artifact | undefined { + const artifacts = message.artifacts || []; + const normalizedPath = options.filePath ? normalizePath(options.filePath) : null; + + return artifacts.find((artifact) => { + if (options.artifactId && artifact.id === options.artifactId) { + return true; + } + + if (!normalizedPath) { + return false; + } + + const artifactPath = + typeof artifact.meta.filePath === "string" + ? normalizePath(artifact.meta.filePath) + : null; + return artifactPath === normalizedPath; + }); +} + export function resolveArtifactTypeFromFile( filePath: string, metadata?: Record, @@ -175,9 +279,22 @@ export function buildArtifactFromWrite({ type === "code" || type === "document" ? resolveArtifactLanguageFromFile(normalizedPath) : undefined; + const previewText = + typeof content === "string" && content.trim() + ? normalizePreviewText(content) + : undefined; + const existingMeta = + context?.artifact?.meta && typeof context.artifact.meta === "object" + ? (context.artifact.meta as ArtifactWriteMetadata) + : undefined; const baseMeta: ArtifactMeta = { + ...(existingMeta || {}), ...(metadata || {}), ...(language ? { language } : {}), + ...(previewText && + !(typeof metadata?.previewText === "string" && metadata.previewText.trim()) + ? { previewText } + : {}), filePath: normalizedPath, filename: title, source: context?.source, @@ -194,7 +311,7 @@ export function buildArtifactFromWrite({ content, status: context?.status || "complete", meta: baseMeta, - position: context?.artifact?.position || { start: 0, end: 0 }, + position: context?.artifact?.position || { start: 0, end: content.length }, createdAt: context?.artifact?.createdAt || now, updatedAt: now, error: context?.artifact?.error, @@ -242,6 +359,15 @@ export function updateMessageArtifactsStatus( ? { ...artifact, status: nextStatus, + meta: { + ...artifact.meta, + writePhase: + nextStatus === "complete" + ? "completed" + : nextStatus === "error" + ? "failed" + : artifact.meta.writePhase, + }, updatedAt: Date.now(), } : artifact, diff --git a/src/components/agent/chat/utils/searchQueryGrouping.test.ts b/src/components/agent/chat/utils/searchQueryGrouping.test.ts new file mode 100644 index 000000000..cc80b040b --- /dev/null +++ b/src/components/agent/chat/utils/searchQueryGrouping.test.ts @@ -0,0 +1,37 @@ +import { describe, expect, it } from "vitest"; +import { + classifySearchQuerySemantic, + summarizeSearchQuerySemantics, +} from "./searchQueryGrouping"; + +describe("searchQueryGrouping", () => { + it("应识别中文日期、英文日期与头条检索语义", () => { + expect(classifySearchQuerySemantic("2026年3月13日 国际新闻").label).toBe( + "中文日期检索", + ); + expect( + classifySearchQuerySemantic("March 13 2026 international news").label, + ).toBe("英文日期检索"); + expect( + classifySearchQuerySemantic("March 13 2026 world headlines").label, + ).toBe("头条检索"); + }); + + it("应汇总搜索语义标签数量", () => { + const summary = summarizeSearchQuerySemantics([ + "2026年3月13日 国际新闻", + "March 13 2026 international news", + "March 13 2026 world headlines", + "今日国际新闻", + ]); + + expect(summary).toEqual( + expect.arrayContaining([ + expect.objectContaining({ label: "中文日期检索", count: 1 }), + expect.objectContaining({ label: "英文日期检索", count: 1 }), + expect.objectContaining({ label: "头条检索", count: 1 }), + expect.objectContaining({ label: "中文检索", count: 1 }), + ]), + ); + }); +}); diff --git a/src/components/agent/chat/utils/searchQueryGrouping.ts b/src/components/agent/chat/utils/searchQueryGrouping.ts new file mode 100644 index 000000000..d3496fb96 --- /dev/null +++ b/src/components/agent/chat/utils/searchQueryGrouping.ts @@ -0,0 +1,56 @@ +export interface SearchQuerySemantic { + key: string; + label: string; +} + +export interface SearchQuerySemanticSummary extends SearchQuerySemantic { + count: number; +} + +const CJK_RE = /[\u4e00-\u9fff]/; +const ZH_DATE_RE = /\d{4}年\d{1,2}月\d{1,2}日|\d{1,2}月\d{1,2}日/; +const EN_DATE_RE = + /\b(?:jan|january|feb|february|mar|march|apr|april|may|jun|june|jul|july|aug|august|sep|sept|september|oct|october|nov|november|dec|december)\b|\b\d{4}-\d{1,2}-\d{1,2}\b/i; +const HEADLINE_RE = /头条|要闻|快讯|headlines?|roundup|briefing|digest|brief/i; + +export function classifySearchQuerySemantic(query?: string | null): SearchQuerySemantic { + const normalized = query?.trim() || ""; + const hasCjk = CJK_RE.test(normalized); + const hasZhDate = ZH_DATE_RE.test(normalized); + const hasEnDate = EN_DATE_RE.test(normalized); + const hasHeadlineHint = HEADLINE_RE.test(normalized); + + if (hasHeadlineHint) { + return { key: "headlines", label: "头条检索" }; + } + if (hasCjk && hasZhDate) { + return { key: "zh_date", label: "中文日期检索" }; + } + if (!hasCjk && hasEnDate) { + return { key: "en_date", label: "英文日期检索" }; + } + if (hasCjk) { + return { key: "zh_general", label: "中文检索" }; + } + return { key: "en_general", label: "英文检索" }; +} + +export function summarizeSearchQuerySemantics( + queries: Array, +): SearchQuerySemanticSummary[] { + const counts = new Map(); + + for (const query of queries) { + const semantic = classifySearchQuerySemantic(query); + const existing = counts.get(semantic.key); + if (existing) { + existing.count += 1; + continue; + } + counts.set(semantic.key, { ...semantic, count: 1 }); + } + + return Array.from(counts.values()).sort((left, right) => + left.key.localeCompare(right.key), + ); +} diff --git a/src/components/agent/chat/utils/searchResultPreview.ts b/src/components/agent/chat/utils/searchResultPreview.ts new file mode 100644 index 000000000..131ec08d1 --- /dev/null +++ b/src/components/agent/chat/utils/searchResultPreview.ts @@ -0,0 +1,277 @@ +export interface SearchResultPreviewItem { + id: string; + title: string; + url: string; + hostname: string; + snippet?: string; +} + +const URL_PATTERN_SOURCE = String.raw`\bhttps?:\/\/[^\s<>"'\`]+`; +const URL_TRAILING_PUNCTUATION = /[),.;!?]+$/; +const SEARCH_MARKDOWN_LINK_RE = /\[([^\]]+)\]\((https?:\/\/[^\s)]+)\)/g; + +export const SEARCH_RESULT_LIST_LIMIT = 10; + +function createUrlPattern(): RegExp { + return new RegExp(URL_PATTERN_SOURCE, "gi"); +} + +function normalizeUrlCandidate(rawUrl: string): { + url: string; + trailing: string; +} { + const normalized = rawUrl.replace(URL_TRAILING_PUNCTUATION, ""); + return { + url: normalized || rawUrl, + trailing: rawUrl.slice((normalized || rawUrl).length), + }; +} + +function findFirstUrl(...values: Array): string | undefined { + for (const value of values) { + if (!value) { + continue; + } + const match = value.match(createUrlPattern()); + if (!match || match.length === 0) { + continue; + } + return normalizeUrlCandidate(match[0]).url; + } + return undefined; +} + +function normalizeSearchText(value: string): string { + return value + .trim() + .replace(/^[\s>*•·\-–—\d().::\]]+/, "") + .replace(/\s+/g, " ") + .trim(); +} + +export function getHostnameFromUrl(url: string): string { + try { + return new URL(url).hostname.replace(/^www\./, ""); + } catch { + return url; + } +} + +function extractSearchResultFromRecord( + record: Record, + index: number, +): SearchResultPreviewItem | null { + const url = + (typeof record.url === "string" && record.url.trim()) || + (typeof record.link === "string" && record.link.trim()) || + (typeof record.href === "string" && record.href.trim()) || + (typeof record.sourceUrl === "string" && record.sourceUrl.trim()) || + (typeof record.source_url === "string" && record.source_url.trim()) || + ""; + if (!url) { + return null; + } + + const title = + (typeof record.title === "string" && normalizeSearchText(record.title)) || + (typeof record.name === "string" && normalizeSearchText(record.name)) || + (typeof record.headline === "string" && + normalizeSearchText(record.headline)) || + getHostnameFromUrl(url); + const snippet = + (typeof record.summary === "string" && normalizeSearchText(record.summary)) || + (typeof record.snippet === "string" && normalizeSearchText(record.snippet)) || + (typeof record.description === "string" && + normalizeSearchText(record.description)) || + (typeof record.content === "string" && normalizeSearchText(record.content)) || + (typeof record.preview === "string" && normalizeSearchText(record.preview)) || + (typeof record.text === "string" && normalizeSearchText(record.text)) || + undefined; + + return { + id: `search-record-${index}-${url}`, + title, + url, + hostname: getHostnameFromUrl(url), + snippet: snippet || undefined, + }; +} + +function parseSearchResultRecords(rawText: string): SearchResultPreviewItem[] { + const trimmed = rawText.trim(); + if (!trimmed) { + return []; + } + + const candidates = [trimmed]; + const fencedMatch = trimmed.match(/```(?:json)?\s*([\s\S]*?)\s*```/i); + if (fencedMatch?.[1]) { + candidates.unshift(fencedMatch[1]); + } + + const seenUrls = new Set(); + for (const candidate of candidates) { + try { + const parsed = JSON.parse(candidate) as unknown; + const queue: unknown[] = [parsed]; + const entries: SearchResultPreviewItem[] = []; + + while (queue.length > 0) { + const current = queue.shift(); + if (!current) { + continue; + } + + if (Array.isArray(current)) { + queue.push(...current); + continue; + } + + if (typeof current !== "object") { + continue; + } + + const record = current as Record; + const extracted = extractSearchResultFromRecord(record, entries.length); + if (extracted && !seenUrls.has(extracted.url)) { + seenUrls.add(extracted.url); + entries.push(extracted); + if (entries.length >= SEARCH_RESULT_LIST_LIMIT) { + return entries; + } + } + + for (const key of ["results", "items", "sources", "citations", "data"]) { + const nested = record[key]; + if (nested) { + queue.push(nested); + } + } + } + + if (entries.length > 0) { + return entries; + } + } catch { + continue; + } + } + + return []; +} + +function parseSearchResultText(rawText: string): SearchResultPreviewItem[] { + const normalizedText = rawText.trim(); + if (!normalizedText) { + return []; + } + + const entries: SearchResultPreviewItem[] = []; + const seenUrls = new Set(); + + for (const match of normalizedText.matchAll(SEARCH_MARKDOWN_LINK_RE)) { + const url = normalizeUrlCandidate(match[2] || "").url; + if (!url || seenUrls.has(url)) { + continue; + } + seenUrls.add(url); + entries.push({ + id: `search-markdown-${entries.length}-${url}`, + title: normalizeSearchText(match[1] || "") || getHostnameFromUrl(url), + url, + hostname: getHostnameFromUrl(url), + }); + if (entries.length >= SEARCH_RESULT_LIST_LIMIT) { + return entries; + } + } + + const lines = normalizedText + .split(/\r?\n/) + .map((line) => line.trim()) + .filter(Boolean); + + for (let index = 0; index < lines.length; index += 1) { + const currentLine = lines[index]; + const url = findFirstUrl(currentLine); + if (!url || seenUrls.has(url)) { + continue; + } + + let title = normalizeSearchText(currentLine.replace(url, "")); + if (!title && index > 0) { + const previousLine = normalizeSearchText(lines[index - 1] || ""); + if (previousLine && !findFirstUrl(previousLine)) { + title = previousLine; + } + } + + const snippetLines: string[] = []; + for (let nextIndex = index + 1; nextIndex < lines.length; nextIndex += 1) { + const nextLine = normalizeSearchText(lines[nextIndex] || ""); + if (!nextLine || findFirstUrl(nextLine)) { + break; + } + snippetLines.push(nextLine); + if (snippetLines.length >= 2 || snippetLines.join(" ").length >= 180) { + break; + } + } + + seenUrls.add(url); + entries.push({ + id: `search-text-${entries.length}-${url}`, + title: title || getHostnameFromUrl(url), + url, + hostname: getHostnameFromUrl(url), + snippet: snippetLines.join(" ").trim() || undefined, + }); + + if (entries.length >= SEARCH_RESULT_LIST_LIMIT) { + break; + } + } + + if (entries.length > 0) { + return entries; + } + + for (const match of normalizedText.matchAll(createUrlPattern())) { + const url = normalizeUrlCandidate(match[0] || "").url; + if (!url || seenUrls.has(url)) { + continue; + } + seenUrls.add(url); + entries.push({ + id: `search-url-${entries.length}-${url}`, + title: getHostnameFromUrl(url), + url, + hostname: getHostnameFromUrl(url), + }); + if (entries.length >= SEARCH_RESULT_LIST_LIMIT) { + break; + } + } + + return entries; +} + +export function resolveSearchResultPreviewItemsFromText( + rawText?: string | null, +): SearchResultPreviewItem[] { + const normalizedText = rawText?.trim(); + if (!normalizedText) { + return []; + } + + const structuredEntries = parseSearchResultRecords(normalizedText); + if (structuredEntries.length > 0) { + return structuredEntries; + } + + return parseSearchResultText(normalizedText); +} + +export function isUnifiedWebSearchToolName(toolName: string): boolean { + return toolName.replace(/[\s_-]+/g, "").trim().toLowerCase() === "websearch"; +} diff --git a/src/components/artifact/ArtifactCanvasOverlay.tsx b/src/components/artifact/ArtifactCanvasOverlay.tsx new file mode 100644 index 000000000..91c153519 --- /dev/null +++ b/src/components/artifact/ArtifactCanvasOverlay.tsx @@ -0,0 +1,76 @@ +import { AlertCircle, Loader2, Sparkles } from "lucide-react"; +import { Badge } from "@/components/ui/badge"; +import { cn } from "@/lib/utils"; +import type { ArtifactDisplayOverlayState } from "@/components/agent/chat/hooks/useArtifactDisplayState"; + +export interface ArtifactCanvasOverlayProps { + overlay: ArtifactDisplayOverlayState; + className?: string; +} + +export function ArtifactCanvasOverlay({ + overlay, + className, +}: ArtifactCanvasOverlayProps) { + const isFailed = overlay.phase === "failed"; + const isCompleted = overlay.phase === "finalized_empty"; + + return ( +
+
+
+
+ {isFailed ? ( + + ) : overlay.showProgress ? ( + + ) : ( + + )} +
+
+
+ + {overlay.phaseLabel} + + + {overlay.displayName} + +
+
+ {overlay.title} +
+
+ {overlay.detail} +
+ {overlay.filePath !== overlay.displayName ? ( +
+ {overlay.filePath} +
+ ) : null} +
+
+ {overlay.showProgress ? ( +
+
+
+ ) : null} +
+
+ ); +} + +export default ArtifactCanvasOverlay; diff --git a/src/components/artifact/ArtifactRenderer.tsx b/src/components/artifact/ArtifactRenderer.tsx index 72c5dcab8..fc7e59065 100644 --- a/src/components/artifact/ArtifactRenderer.tsx +++ b/src/components/artifact/ArtifactRenderer.tsx @@ -13,10 +13,22 @@ import React, { useEffect, useRef, } from "react"; -import { Loader2, AlertTriangle, RefreshCw } from "lucide-react"; +import { + AlertCircle, + AlertTriangle, + FileCode2, + FileText, + LayoutTemplate, + Loader2, + RefreshCw, +} from "lucide-react"; import { cn } from "@/lib/utils"; import { artifactRegistry } from "@/lib/artifact/registry"; import { useDebouncedValue } from "@/lib/artifact/hooks"; +import { + formatArtifactWritePhaseLabel, + resolveArtifactWritePhase, +} from "@/components/agent/chat/utils/messageArtifacts"; import { ErrorFallbackRenderer } from "./ErrorFallbackRenderer"; import { CanvasAdapter } from "./CanvasAdapter"; import type { Artifact, ArtifactRendererProps } from "@/lib/artifact/types"; @@ -95,14 +107,326 @@ const RendererSkeleton: React.FC<{ tone?: "dark" | "light" }> = memo( tone === "light" ? "text-muted-foreground" : "text-gray-400", )} > - - 加载渲染器... + + 加载渲染器...
), ); RendererSkeleton.displayName = "RendererSkeleton"; +type EmptyArtifactSurfaceMode = "writing" | "finished" | "failed"; +type EmptyArtifactSkeletonKind = "document" | "code" | "preview" | "generic"; + +interface EmptyArtifactSurfaceState { + mode: EmptyArtifactSurfaceMode; + title: string; + detail: string; + skeletonKind: EmptyArtifactSkeletonKind; +} + +function resolveArtifactPath(artifact: Pick): string { + if (typeof artifact.meta.filePath === "string" && artifact.meta.filePath.trim()) { + return artifact.meta.filePath.trim(); + } + if (typeof artifact.meta.filename === "string" && artifact.meta.filename.trim()) { + return artifact.meta.filename.trim(); + } + return artifact.title; +} + +function resolveEmptyArtifactSkeletonKind( + artifact: Pick, +): EmptyArtifactSkeletonKind { + if (artifact.type === "document" || artifact.type === "canvas:document") { + return "document"; + } + if (artifact.type === "code") { + return "code"; + } + if ( + artifact.type === "html" || + artifact.type === "svg" || + artifact.type === "react" || + artifact.type === "mermaid" + ) { + return "preview"; + } + return "generic"; +} + +function resolveEmptyArtifactSurfaceState( + artifact: Pick, +): EmptyArtifactSurfaceState | null { + if (artifact.content.trim()) { + return null; + } + + const writePhase = resolveArtifactWritePhase(artifact); + const skeletonKind = resolveEmptyArtifactSkeletonKind(artifact); + + if ( + artifact.status === "error" || + writePhase === "failed" + ) { + return { + mode: "failed", + title: "写入未完成", + detail: + artifact.error?.trim() || + "文件写入过程中出现异常,暂时没有可渲染的内容。", + skeletonKind, + }; + } + + if ( + artifact.status === "complete" || + writePhase === "completed" || + writePhase === "persisted" + ) { + return { + mode: "finished", + title: "写入已结束", + detail: "文件已经创建完成,但当前还没有可直接预览的内容。", + skeletonKind, + }; + } + + if ( + artifact.status === "pending" || + artifact.status === "streaming" || + writePhase === "preparing" || + writePhase === "streaming" + ) { + const phaseLabel = formatArtifactWritePhaseLabel(writePhase); + return { + mode: "writing", + title: phaseLabel, + detail: + writePhase === "preparing" + ? "文件已创建,正在生成首段内容。" + : "内容正在持续写入,首段到达后会立即替换这块骨架。", + skeletonKind, + }; + } + + return null; +} + +const SkeletonBar: React.FC<{ + className?: string; + tone?: "dark" | "light"; +}> = memo(({ className, tone = "dark" }) => ( +
+)); +SkeletonBar.displayName = "SkeletonBar"; + +const DocumentSkeleton: React.FC<{ tone?: "dark" | "light" }> = memo( + ({ tone = "dark" }) => ( +
+ + + + +
+ +
+ + + +
+ ), +); +DocumentSkeleton.displayName = "DocumentSkeleton"; + +const CodeSkeleton: React.FC<{ tone?: "dark" | "light" }> = memo( + ({ tone = "dark" }) => ( +
+ {Array.from({ length: 8 }).map((_, index) => ( +
+
+ {index + 1} +
+ +
+ ))} +
+ ), +); +CodeSkeleton.displayName = "CodeSkeleton"; + +const PreviewSkeleton: React.FC<{ tone?: "dark" | "light" }> = memo( + ({ tone = "dark" }) => ( +
+
+ + +
+ ), +); +PreviewSkeleton.displayName = "PreviewSkeleton"; + +const GenericSkeleton: React.FC<{ tone?: "dark" | "light" }> = memo( + ({ tone = "dark" }) => ( +
+
+ +
+ + + +
+ ), +); +GenericSkeleton.displayName = "GenericSkeleton"; + +const EmptyArtifactSurface: React.FC<{ + artifact: Artifact; + state: EmptyArtifactSurfaceState; + tone?: "dark" | "light"; +}> = memo(({ artifact, state, tone = "dark" }) => { + const filePath = resolveArtifactPath(artifact); + const isWriting = state.mode === "writing"; + const isFailed = state.mode === "failed"; + const Icon = isFailed + ? AlertCircle + : state.skeletonKind === "code" + ? FileCode2 + : state.skeletonKind === "preview" + ? LayoutTemplate + : FileText; + + return ( +
+
+
+
+ +
+
+
+ {state.title} +
+
+ {state.detail} +
+
+ {filePath} +
+
+
+ {isWriting ? ( +
+
+
+ ) : null} +
+
+ {state.skeletonKind === "document" ? ( + + ) : state.skeletonKind === "code" ? ( + + ) : state.skeletonKind === "preview" ? ( + + ) : ( + + )} +
+
+ ); +}); +EmptyArtifactSurface.displayName = "EmptyArtifactSurface"; + /** * 流式状态指示器 Props * @requirements 11.1 @@ -315,6 +639,22 @@ export const ArtifactRenderer: React.FC = memo( // 获取渲染器注册项 const entry = artifactRegistry.get(artifact.type); + const emptySurfaceState = resolveEmptyArtifactSurfaceState(artifact); + + if (emptySurfaceState) { + return ( +
+ + {(isStreaming || isCompleting) && ( + + )} +
+ ); + } // 如果有错误且选择显示源码,直接显示源码 if (renderError && showSourceOnError) { diff --git a/src/components/artifact/ArtifactRenderer.ui.test.tsx b/src/components/artifact/ArtifactRenderer.ui.test.tsx new file mode 100644 index 000000000..b48830a42 --- /dev/null +++ b/src/components/artifact/ArtifactRenderer.ui.test.tsx @@ -0,0 +1,113 @@ +import { act } from "react"; +import { createRoot, type Root } from "react-dom/client"; +import { afterEach, beforeEach, describe, expect, it } from "vitest"; +import { ArtifactRenderer } from "./ArtifactRenderer"; +import type { Artifact } from "@/lib/artifact/types"; + +interface MountedRenderer { + container: HTMLDivElement; + root: Root; +} + +const mountedRenderers: MountedRenderer[] = []; + +function createArtifact(overrides: Partial = {}): Artifact { + const content = overrides.content ?? ""; + return { + id: overrides.id ?? "artifact-1", + type: overrides.type ?? "document", + title: overrides.title ?? "demo.md", + content, + status: overrides.status ?? "pending", + meta: { + filePath: overrides.meta?.filePath ?? "workspace/demo.md", + filename: overrides.meta?.filename ?? "demo.md", + ...overrides.meta, + }, + position: overrides.position ?? { start: 0, end: content.length }, + createdAt: overrides.createdAt ?? 1, + updatedAt: overrides.updatedAt ?? 1, + error: overrides.error, + }; +} + +function renderArtifact(artifact: Artifact) { + const container = document.createElement("div"); + document.body.appendChild(container); + const root = createRoot(container); + + act(() => { + root.render(); + }); + + mountedRenderers.push({ container, root }); + return container; +} + +beforeEach(() => { + ( + globalThis as typeof globalThis & { + IS_REACT_ACT_ENVIRONMENT?: boolean; + } + ).IS_REACT_ACT_ENVIRONMENT = true; +}); + +afterEach(() => { + while (mountedRenderers.length > 0) { + const mounted = mountedRenderers.pop(); + if (!mounted) { + break; + } + act(() => { + mounted.root.unmount(); + }); + mounted.container.remove(); + } +}); + +describe("ArtifactRenderer 空内容态", () => { + it("流式写入但暂无内容时应展示类型化骨架", () => { + const container = renderArtifact( + createArtifact({ + type: "code", + title: "index.ts", + status: "streaming", + meta: { + filePath: "workspace/index.ts", + writePhase: "streaming", + language: "typescript", + }, + }), + ); + + const surface = container.querySelector( + "[data-testid=\"artifact-empty-surface\"]", + ); + + expect(surface).not.toBeNull(); + expect(surface?.getAttribute("data-empty-mode")).toBe("writing"); + expect(container.textContent).toContain("正在写入"); + expect(container.textContent).toContain("workspace/index.ts"); + }); + + it("失败且没有内容时应展示错误解释态", () => { + const container = renderArtifact( + createArtifact({ + status: "error", + error: "保存失败", + meta: { + filePath: "workspace/broken.md", + writePhase: "failed", + }, + }), + ); + + const surface = container.querySelector( + "[data-testid=\"artifact-empty-surface\"]", + ); + + expect(surface?.getAttribute("data-empty-mode")).toBe("failed"); + expect(container.textContent).toContain("写入未完成"); + expect(container.textContent).toContain("保存失败"); + }); +}); diff --git a/src/components/artifact/ArtifactToolbar.tsx b/src/components/artifact/ArtifactToolbar.tsx index a3afa68cf..16a7cc0e4 100644 --- a/src/components/artifact/ArtifactToolbar.tsx +++ b/src/components/artifact/ArtifactToolbar.tsx @@ -19,8 +19,13 @@ import { Monitor, } from "lucide-react"; import { cn } from "@/lib/utils"; +import { Badge } from "@/components/ui/badge"; import { artifactRegistry } from "@/lib/artifact/registry"; import type { Artifact } from "@/lib/artifact/types"; +import { + formatArtifactWritePhaseLabel, + resolveArtifactWritePhase, +} from "@/components/agent/chat/utils/messageArtifacts"; /** 视图模式类型 */ type ViewMode = "source" | "preview"; @@ -178,6 +183,8 @@ export interface ArtifactToolbarProps { onPreviewSizeChange?: (size: PreviewSize) => void; /** 工具栏色调 */ tone?: "dark" | "light"; + /** 额外的展示状态标签 */ + displayBadgeLabel?: string; } /** @@ -328,6 +335,7 @@ export const ArtifactToolbar: React.FC = memo( previewSize = "desktop", onPreviewSizeChange, tone = "dark", + displayBadgeLabel, }) => { const [copied, setCopied] = useState(false); @@ -340,6 +348,7 @@ export const ArtifactToolbar: React.FC = memo( const language = artifact.meta.language?.toLowerCase() || ""; const canPreview = isCode && PREVIEWABLE_LANGUAGES.includes(language); const supportsSharedViewMode = isDocument || canPreview; + const writePhase = resolveArtifactWritePhase(artifact); /** * 复制内容到剪贴板 @@ -487,6 +496,32 @@ export const ArtifactToolbar: React.FC = memo( > {artifact.title} + {writePhase ? ( + + {formatArtifactWritePhaseLabel(writePhase)} + + ) : null} + {displayBadgeLabel ? ( + + {displayBadgeLabel} + + ) : null}
{/* 操作按钮区域 */} diff --git a/src/components/artifact/index.ts b/src/components/artifact/index.ts index b325b6966..2c6fbb692 100644 --- a/src/components/artifact/index.ts +++ b/src/components/artifact/index.ts @@ -30,6 +30,13 @@ export type { ArtifactPanelProps } from "./ArtifactPanel"; export { ArtifactToolbar } from "./ArtifactToolbar"; export type { ArtifactToolbarProps } from "./ArtifactToolbar"; +/** + * Artifact 画布过渡遮罩组件 + * 在文件写入开始但首段内容尚未到达前展示稳定反馈 + */ +export { ArtifactCanvasOverlay } from "./ArtifactCanvasOverlay"; +export type { ArtifactCanvasOverlayProps } from "./ArtifactCanvasOverlay"; + /** * Artifact 列表组件 * 显示当前消息中的所有 artifacts diff --git a/src/components/input-kit/BaseComposer.tsx b/src/components/input-kit/BaseComposer.tsx index 54d09c548..a884ceba9 100644 --- a/src/components/input-kit/BaseComposer.tsx +++ b/src/components/input-kit/BaseComposer.tsx @@ -28,6 +28,7 @@ export interface BaseComposerProps { hasAdditionalContent?: boolean; rows?: number; autoFocus?: boolean; + allowSendWhileLoading?: boolean; children: (context: BaseComposerRenderContext) => React.ReactNode; } @@ -50,6 +51,7 @@ export const BaseComposer: React.FC = ({ hasAdditionalContent = false, rows = 1, autoFocus = false, + allowSendWhileLoading = false, children, }) => { const internalTextareaRef = useRef(null); @@ -59,8 +61,10 @@ export const BaseComposer: React.FC = ({ return text.trim().length > 0 || hasAdditionalContent; }, [hasAdditionalContent, text]); - const canSend = hasContent && !disabled && !isLoading; - const isPrimaryDisabled = !isLoading && !canSend; + const canSend = + hasContent && !disabled && (!isLoading || allowSendWhileLoading); + const isPrimaryDisabled = + isLoading && !allowSendWhileLoading ? false : !canSend; useEffect(() => { const textarea = textareaRef.current; @@ -101,7 +105,7 @@ export const BaseComposer: React.FC = ({ ); const onPrimaryAction = useCallback(() => { - if (isLoading) { + if (isLoading && !allowSendWhileLoading) { onStop?.(); return; } @@ -111,7 +115,7 @@ export const BaseComposer: React.FC = ({ } onSend(); - }, [canSend, isLoading, onSend, onStop]); + }, [allowSendWhileLoading, canSend, isLoading, onSend, onStop]); const handleKeyDown = useCallback( (event: React.KeyboardEvent) => { diff --git a/src/components/provider-pool/ProviderPoolPage.tsx b/src/components/provider-pool/ProviderPoolPage.tsx index 91fb6cf0c..6861b08f3 100644 --- a/src/components/provider-pool/ProviderPoolPage.tsx +++ b/src/components/provider-pool/ProviderPoolPage.tsx @@ -35,7 +35,6 @@ import { ProviderIcon } from "@/icons/providers"; import { ApiKeyProviderSection, AddCustomProviderModal } from "./api-key"; import type { ApiKeyProviderSectionRef } from "./api-key"; import { RelayProvidersSection } from "./RelayProvidersSection"; -import { ModelRegistryTab } from "./ModelRegistryTab"; import { AsrProviderSection } from "@/components/voice"; import type { AddCustomProviderRequest } from "@/lib/api/apiKeyProvider"; import { @@ -85,7 +84,7 @@ const isConfigTab = (tab: TabType): tab is ConfigTabType => { }; // 分类类型 -type CategoryType = "oauth" | "apikey" | "connect" | "models" | "voice"; +type CategoryType = "oauth" | "apikey" | "connect" | "voice"; export const ProviderPoolPage = forwardRef< ProviderPoolPageRef, @@ -398,19 +397,6 @@ export const ProviderPoolPage = forwardRef< > Connect -
)} - {/* 模型库分类 */} - {activeCategory === "models" && } - {/* 语音服务分类 */} {activeCategory === "voice" && (
diff --git a/src/components/provider-pool/README.md b/src/components/provider-pool/README.md index 7a21064e7..eb8898eba 100644 --- a/src/components/provider-pool/README.md +++ b/src/components/provider-pool/README.md @@ -4,35 +4,35 @@ ## 组件列表 -| 文件 | 描述 | -|------|------| -| `ProviderPoolPage.tsx` | 凭证池管理主页面,支持 OAuth 凭证卡片布局和 API Key 左右分栏布局 | -| `CredentialCard.tsx` | OAuth 凭证卡片组件,显示健康状态、使用统计和操作按钮 | -| `CredentialCardContextMenu.tsx` | 凭证卡片右键菜单组件 | -| `AddCredentialModal.tsx` | 添加凭证模态框组件 | -| `EditCredentialModal.tsx` | 编辑凭证模态框组件 | -| `ErrorDisplay.tsx` | 错误显示组件 | -| `UsageDisplay.tsx` | 用量显示组件 | -| `RelayProvidersSection.tsx` | Connect 中转商列表组件,展示已验证的中转服务商 | -| `VertexAISection.tsx` | Vertex AI 配置区域组件 | -| `AmpConfigSection.tsx` | Amp CLI 配置区域组件 | -| `GeminiApiKeySection.tsx` | Gemini API Key 配置区域组件 | -| `CodexSection.tsx` | Codex 配置区域组件 | -| `IFlowSection.tsx` | iFlow 配置区域组件 | -| `OAuthPluginTab.tsx` | OAuth 插件标签页组件 | -| `index.ts` | 组件导出入口 | +| 文件 | 描述 | +| ------------------------------- | ---------------------------------------------------------------- | +| `ProviderPoolPage.tsx` | 凭证池管理主页面,支持 OAuth 凭证卡片布局和 API Key 左右分栏布局 | +| `CredentialCard.tsx` | OAuth 凭证卡片组件,显示健康状态、使用统计和操作按钮 | +| `CredentialCardContextMenu.tsx` | 凭证卡片右键菜单组件 | +| `AddCredentialModal.tsx` | 添加凭证模态框组件 | +| `EditCredentialModal.tsx` | 编辑凭证模态框组件 | +| `ErrorDisplay.tsx` | 错误显示组件 | +| `UsageDisplay.tsx` | 用量显示组件 | +| `RelayProvidersSection.tsx` | Connect 中转商列表组件,展示已验证的中转服务商 | +| `VertexAISection.tsx` | Vertex AI 配置区域组件 | +| `AmpConfigSection.tsx` | Amp CLI 配置区域组件 | +| `GeminiApiKeySection.tsx` | Gemini API Key 配置区域组件 | +| `CodexSection.tsx` | Codex 配置区域组件 | +| `IFlowSection.tsx` | iFlow 配置区域组件 | +| `OAuthPluginTab.tsx` | OAuth 插件标签页组件 | +| `index.ts` | 组件导出入口 | ## 子目录 -| 目录 | 描述 | -|------|------| -| `api-key/` | API Key Provider 管理组件(左右分栏布局) | -| `credential-forms/` | 各类凭证表单组件 | +| 目录 | 描述 | +| ------------------- | ----------------------------------------- | +| `api-key/` | API Key Provider 管理组件(左右分栏布局) | +| `credential-forms/` | 各类凭证表单组件 | ## 测试文件 -| 文件 | 描述 | -|------|------| +| 文件 | 描述 | +| ------------------------ | --------------------------------------------- | | `CredentialCard.test.ts` | Property 3 属性测试:OAuth 凭证卡片信息完整性 | ## 使用示例 @@ -41,9 +41,7 @@ import { ProviderPoolPage } from "@/components/provider-pool"; function App() { - return ( - - ); + return ; } ``` @@ -57,7 +55,8 @@ function App() { ## 架构说明 ProviderPoolPage 支持四种分类: + 1. **OAuth 凭证** - 使用卡片式布局显示 OAuth 类型凭证 2. **API Key** - 使用左右分栏布局(ApiKeyProviderSection) -3. **OAuth 插件** - 第三方 OAuth 插件管理 -4. **Connect** - 中转商列表,支持浏览和一键获取 API Key +3. **Connect** - 中转商列表,支持浏览和一键获取 API Key +4. **语音服务** - 语音 Provider 管理入口 diff --git a/src/components/provider-pool/api-key/AddCustomProviderModal.test.ts b/src/components/provider-pool/api-key/AddCustomProviderModal.test.ts index a327c58d3..157740abd 100644 --- a/src/components/provider-pool/api-key/AddCustomProviderModal.test.ts +++ b/src/components/provider-pool/api-key/AddCustomProviderModal.test.ts @@ -16,6 +16,7 @@ import { isFormValid, hasRequiredFields, } from "./AddCustomProviderModal"; +import { PROVIDER_TYPE_VALUES } from "./ProviderConfigForm.utils"; import type { ProviderType } from "@/lib/types/provider"; // ============================================================================ @@ -23,19 +24,7 @@ import type { ProviderType } from "@/lib/types/provider"; // ============================================================================ /** 所有有效的 Provider 类型 */ -const VALID_PROVIDER_TYPES: ProviderType[] = [ - "openai", - "openai-response", - "anthropic", - "gemini", - "azure-openai", - "vertexai", - "aws-bedrock", - "ollama", - "fal", - "new-api", - "gateway", -]; +const VALID_PROVIDER_TYPES: ProviderType[] = PROVIDER_TYPE_VALUES; /** * 生成有效的 Provider 类型 diff --git a/src/components/provider-pool/api-key/AddCustomProviderModal.tsx b/src/components/provider-pool/api-key/AddCustomProviderModal.tsx index ca832a7cd..4bf03307a 100644 --- a/src/components/provider-pool/api-key/AddCustomProviderModal.tsx +++ b/src/components/provider-pool/api-key/AddCustomProviderModal.tsx @@ -34,44 +34,16 @@ import { type AddCustomProviderRequest, type SystemProviderCatalogItem, } from "@/lib/api/apiKeyProvider"; +import { + isSupportedProviderType, + PROVIDER_TYPE_FIELDS, + PROVIDER_TYPE_OPTIONS, +} from "./ProviderConfigForm.utils"; // ============================================================================ // 常量 // ============================================================================ -/** 支持的 Provider 类型列表 */ -const PROVIDER_TYPES: { value: ProviderType; label: string }[] = [ - { value: "openai", label: "OpenAI 兼容" }, - { value: "openai-response", label: "OpenAI Responses API" }, - { value: "anthropic", label: "Anthropic" }, - { value: "anthropic-compatible", label: "Anthropic 兼容" }, - { value: "gemini", label: "Gemini" }, - { value: "azure-openai", label: "Azure OpenAI" }, - { value: "vertexai", label: "VertexAI" }, - { value: "aws-bedrock", label: "AWS Bedrock" }, - { value: "ollama", label: "Ollama" }, - { value: "fal", label: "Fal" }, - { value: "new-api", label: "New API" }, - { value: "gateway", label: "Vercel AI Gateway" }, -]; - -/** Provider 类型对应的额外字段 */ -const PROVIDER_TYPE_EXTRA_FIELDS: Record = { - openai: [], - "openai-response": [], - codex: [], - anthropic: [], - "anthropic-compatible": [], // Anthropic 兼容格式,无需额外字段 - gemini: [], - "azure-openai": ["apiVersion"], - vertexai: ["project", "location"], - "aws-bedrock": ["region"], - ollama: [], - fal: [], - "new-api": [], - gateway: [], -}; - /** 已知厂商配置 */ interface KnownProvider { id: string; @@ -189,24 +161,7 @@ const FALLBACK_KNOWN_PROVIDERS: KnownProvider[] = [ /** 将 Catalog 返回的 provider type 收敛到前端 ProviderType */ function normalizeCatalogProviderType(providerType: string): ProviderType { - switch (providerType) { - case "openai": - case "openai-response": - case "codex": - case "anthropic": - case "anthropic-compatible": - case "gemini": - case "azure-openai": - case "vertexai": - case "aws-bedrock": - case "ollama": - case "fal": - case "new-api": - case "gateway": - return providerType; - default: - return "openai"; - } + return isSupportedProviderType(providerType) ? providerType : "openai"; } function buildKnownProvidersFromCatalog( @@ -524,7 +479,7 @@ export const AddCustomProviderModal: React.FC = ({ // 获取当前类型需要的额外字段 const extraFields = useMemo( - () => PROVIDER_TYPE_EXTRA_FIELDS[formState.type] || [], + () => PROVIDER_TYPE_FIELDS[formState.type] || [], [formState.type], ); @@ -740,7 +695,7 @@ export const AddCustomProviderModal: React.FC = ({ - {PROVIDER_TYPES.map((type) => ( + {PROVIDER_TYPE_OPTIONS.map((type) => ( {type.label} diff --git a/src/components/provider-pool/api-key/ApiKeyProviderSection.tsx b/src/components/provider-pool/api-key/ApiKeyProviderSection.tsx index 16e2b8b40..d5b1823d1 100644 --- a/src/components/provider-pool/api-key/ApiKeyProviderSection.tsx +++ b/src/components/provider-pool/api-key/ApiKeyProviderSection.tsx @@ -246,9 +246,15 @@ export const ApiKeyProviderSection = forwardRef< return (
+
+
+ {/* 左侧:Provider 列表 */} setShowImportExportDialog(true)} - className="flex-shrink-0" + className="flex-shrink-0 bg-card" /> {/* 右侧:Provider 设置面板 */} -
+
+
= fc.constantFrom( /** * Provider 类型与其额外字段的映射 */ -const EXPECTED_EXTRA_FIELDS: Record = { - openai: [], - "openai-response": [], - codex: [], - anthropic: [], - "anthropic-compatible": [], - gemini: [], - "azure-openai": ["apiVersion"], - vertexai: ["project", "location"], - "aws-bedrock": ["region"], - ollama: [], - fal: [], - "new-api": [], - gateway: [], -}; +const EXPECTED_EXTRA_FIELDS: Record = + PROVIDER_TYPE_FIELDS; + +function createModel( + overrides: Partial & + Pick, +): EnhancedModelMetadata { + return { + id: overrides.id, + display_name: overrides.display_name, + provider_id: overrides.provider_id ?? "openai", + provider_name: overrides.provider_name ?? "OpenAI", + family: overrides.family ?? null, + tier: overrides.tier ?? "pro", + capabilities: overrides.capabilities ?? { + vision: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: true, + }, + pricing: overrides.pricing ?? null, + limits: overrides.limits ?? { + context_length: null, + max_output_tokens: null, + requests_per_minute: null, + tokens_per_minute: null, + }, + status: overrides.status ?? "active", + release_date: overrides.release_date ?? null, + is_latest: overrides.is_latest ?? false, + description: overrides.description ?? null, + source: overrides.source ?? "local", + created_at: overrides.created_at ?? 0, + updated_at: overrides.updated_at ?? 0, + }; +} // ============================================================================ // Property 7: Provider 类型处理正确性 @@ -234,3 +252,66 @@ describe("Property 7: Provider 类型处理正确性", () => { }); }); }); + +describe("模型辅助函数", () => { + test("parseCustomModelsValue 应去重并保留输入顺序", () => { + expect( + parseCustomModelsValue( + "gpt-5.3-codex, babbage-002, GPT-5.3-codex, , gpt-5.2", + ), + ).toEqual(["gpt-5.3-codex", "babbage-002", "gpt-5.2"]); + }); + + test("serializeCustomModels 应输出稳定的逗号分隔字符串", () => { + expect( + serializeCustomModels(["gpt-5.3-codex", "gpt-5.2", "GPT-5.3-codex"]), + ).toBe("gpt-5.3-codex, gpt-5.2"); + }); + + test("sortSelectableModels 应优先最新和带发布日期的模型", () => { + const models = [ + createModel({ + id: "babbage-002", + display_name: "babbage-002", + }), + createModel({ + id: "gpt-5.2", + display_name: "GPT-5.2", + release_date: "2025-12-11", + }), + createModel({ + id: "gpt-5.3-codex", + display_name: "GPT-5.3 Codex", + release_date: "2026-02-05", + is_latest: true, + }), + ]; + + expect(sortSelectableModels(models).map((model) => model.id)).toEqual([ + "gpt-5.3-codex", + "gpt-5.2", + "babbage-002", + ]); + }); + + test("getLatestSelectableModel 不应把按字母序靠前的旧模型当成最新", () => { + const latestModel = getLatestSelectableModel([ + createModel({ + id: "babbage-002", + display_name: "babbage-002", + }), + createModel({ + id: "gpt-5.3-codex", + display_name: "GPT-5.3 Codex", + release_date: "2026-02-05", + is_latest: true, + }), + ]); + + expect(latestModel?.id).toBe("gpt-5.3-codex"); + }); + + test("getLatestSelectableModel 在空列表时应返回 null", () => { + expect(getLatestSelectableModel([])).toBeNull(); + }); +}); diff --git a/src/components/provider-pool/api-key/ProviderConfigForm.tsx b/src/components/provider-pool/api-key/ProviderConfigForm.tsx index bd6195fa3..f711bbf2c 100644 --- a/src/components/provider-pool/api-key/ProviderConfigForm.tsx +++ b/src/components/provider-pool/api-key/ProviderConfigForm.tsx @@ -7,10 +7,20 @@ * **Validates: Requirements 4.1, 4.2, 5.3-5.5** */ -import React, { useState, useEffect, useCallback, useRef } from "react"; +import React, { + forwardRef, + useState, + useEffect, + useCallback, + useRef, + useMemo, + useImperativeHandle, +} from "react"; import { cn } from "@/lib/utils"; import { Input } from "@/components/ui/input"; import { Label } from "@/components/ui/label"; +import { Button } from "@/components/ui/button"; +import { Badge } from "@/components/ui/badge"; import { Select, SelectContent, @@ -23,6 +33,19 @@ import type { UpdateProviderRequest, } from "@/lib/api/apiKeyProvider"; import type { ProviderType } from "@/lib/types/provider"; +import type { EnhancedModelMetadata } from "@/lib/types/modelRegistry"; +import type { ConfiguredProvider } from "@/hooks/useConfiguredProviders"; +import { useProviderModels } from "@/hooks/useProviderModels"; +import { resolveRegistryProviderId } from "./providerTypeMapping"; +import { + dedupeModelIds, + getLatestSelectableModel, + parseCustomModelsValue, + PROVIDER_TYPE_FIELDS, + PROVIDER_TYPE_OPTIONS, + serializeCustomModels, +} from "./ProviderConfigForm.utils"; +import { Plus, Star, X } from "lucide-react"; // ============================================================================ // 常量 @@ -31,40 +54,6 @@ import type { ProviderType } from "@/lib/types/provider"; /** 防抖延迟时间(毫秒) */ const DEBOUNCE_DELAY = 500; -/** 支持的 Provider 类型列表 */ -const PROVIDER_TYPES: { value: ProviderType; label: string }[] = [ - { value: "openai", label: "OpenAI 兼容" }, - { value: "openai-response", label: "OpenAI Responses API" }, - { value: "codex", label: "Codex CLI" }, - { value: "anthropic", label: "Anthropic" }, - { value: "anthropic-compatible", label: "Anthropic 兼容" }, - { value: "gemini", label: "Gemini" }, - { value: "azure-openai", label: "Azure OpenAI" }, - { value: "vertexai", label: "VertexAI" }, - { value: "aws-bedrock", label: "AWS Bedrock" }, - { value: "ollama", label: "Ollama" }, - { value: "fal", label: "Fal" }, - { value: "new-api", label: "New API" }, - { value: "gateway", label: "Vercel AI Gateway" }, -]; - -/** Provider 类型对应的额外字段配置 */ -const PROVIDER_TYPE_FIELDS: Record = { - openai: [], - "openai-response": [], - codex: [], - anthropic: [], - "anthropic-compatible": [], // Anthropic 兼容格式,无需额外字段 - gemini: [], - "azure-openai": ["apiVersion"], - vertexai: ["project", "location"], - "aws-bedrock": ["region"], - ollama: [], - fal: [], - "new-api": [], - gateway: [], -}; - /** 字段标签映射 */ const FIELD_LABELS: Record = { apiHost: "API Host", @@ -101,12 +90,23 @@ export interface ProviderConfigFormProps { provider: ProviderWithKeysDisplay; /** 更新回调 */ onUpdate?: (id: string, request: UpdateProviderRequest) => Promise; + /** 当前模型列表变化回调 */ + onModelsChange?: (models: string[]) => void; + /** 推荐最新模型变化回调 */ + onRecommendedLatestModelChange?: (modelId: string | null) => void; /** 是否正在加载 */ loading?: boolean; /** 额外的 CSS 类名 */ className?: string; } +export interface ProviderConfigFormRef { + /** 将模型设为默认模型(置顶) */ + setDefaultModel: (modelId: string) => void; + /** 追加模型到列表 */ + addModels: (modelIds: string[]) => void; +} + interface FormState { providerType: ProviderType; apiHost: string; @@ -117,6 +117,10 @@ interface FormState { customModels: string; } +function hasRegistryBackedMetadata(model: EnhancedModelMetadata): boolean { + return Boolean(model.is_latest || model.release_date); +} + // ============================================================================ // 组件实现 // ============================================================================ @@ -141,34 +145,23 @@ interface FormState { * /> * ``` */ -export const ProviderConfigForm: React.FC = ({ - provider, - onUpdate, - loading = false, - className, -}) => { - // 表单状态 - const [formState, setFormState] = useState({ - providerType: (provider.type as ProviderType) || "openai", - apiHost: provider.api_host || "", - apiVersion: provider.api_version || "", - project: provider.project || "", - location: provider.location || "", - region: provider.region || "", - customModels: (provider.custom_models || []).join(", "), - }); - - // 保存状态 - const [isSaving, setIsSaving] = useState(false); - const [saveError, setSaveError] = useState(null); - const [lastSaved, setLastSaved] = useState(null); - - // 防抖定时器 - const debounceTimerRef = useRef | null>(null); - - // 当 provider 变化时,重置表单状态 - useEffect(() => { - setFormState({ +export const ProviderConfigForm = forwardRef< + ProviderConfigFormRef, + ProviderConfigFormProps +>( + ( + { + provider, + onUpdate, + onModelsChange, + onRecommendedLatestModelChange, + loading = false, + className, + }, + ref, + ) => { + // 表单状态 + const [formState, setFormState] = useState({ providerType: (provider.type as ProviderType) || "openai", apiHost: provider.api_host || "", apiVersion: provider.api_version || "", @@ -177,302 +170,548 @@ export const ProviderConfigForm: React.FC = ({ region: provider.region || "", customModels: (provider.custom_models || []).join(", "), }); - setSaveError(null); - }, [ - provider.id, - provider.type, - provider.api_host, - provider.api_version, - provider.project, - provider.location, - provider.region, - provider.custom_models, - ]); - // 保存配置 - const saveConfig = useCallback( - async (state: FormState) => { - if (!onUpdate) return; + // 保存状态 + const [isSaving, setIsSaving] = useState(false); + const [saveError, setSaveError] = useState(null); + const [lastSaved, setLastSaved] = useState(null); + const [modelDraft, setModelDraft] = useState(""); - setIsSaving(true); + // 防抖定时器 + const debounceTimerRef = useRef | null>(null); + + const selectedModels = useMemo( + () => parseCustomModelsValue(formState.customModels), + [formState.customModels], + ); + + const configuredProvider = useMemo( + () => ({ + key: provider.id, + label: provider.name, + registryId: provider.id, + fallbackRegistryId: resolveRegistryProviderId(provider.id, { + providerType: formState.providerType, + }), + type: formState.providerType, + providerId: provider.id, + customModels: selectedModels, + }), + [formState.providerType, provider.id, provider.name, selectedModels], + ); + + const { + models: localCandidateModels, + loading: localModelsLoading, + error: localModelsError, + } = useProviderModels(configuredProvider, { + returnFullMetadata: true, + }); + + const latestLocalModel = useMemo(() => { + const localModelsWithMetadata = localCandidateModels.filter( + hasRegistryBackedMetadata, + ); + return getLatestSelectableModel(localModelsWithMetadata); + }, [localCandidateModels]); + + const recommendedLatestModel = useMemo(() => { + if (latestLocalModel) { + return latestLocalModel; + } + + if (localModelsLoading) { + return null; + } + + return getLatestSelectableModel(localCandidateModels); + }, [latestLocalModel, localCandidateModels, localModelsLoading]); + + // 当 provider 变化时,重置表单状态 + useEffect(() => { + setFormState({ + providerType: (provider.type as ProviderType) || "openai", + apiHost: provider.api_host || "", + apiVersion: provider.api_version || "", + project: provider.project || "", + location: provider.location || "", + region: provider.region || "", + customModels: (provider.custom_models || []).join(", "), + }); setSaveError(null); + setModelDraft(""); + }, [ + provider.id, + provider.type, + provider.api_host, + provider.api_version, + provider.project, + provider.location, + provider.region, + provider.custom_models, + ]); - try { - // 解析自定义模型列表(逗号分隔) - const customModels = state.customModels + // 保存配置 + const saveConfig = useCallback( + async (state: FormState) => { + if (!onUpdate) return; + + setIsSaving(true); + setSaveError(null); + + try { + // 解析自定义模型列表(逗号分隔) + const customModels = state.customModels + .split(",") + .map((m) => m.trim()) + .filter((m) => m.length > 0); + + const request: UpdateProviderRequest = { + // 只有自定义 Provider 才发送 type 字段 + type: !provider.is_system ? state.providerType : undefined, + api_host: state.apiHost || undefined, + api_version: state.apiVersion || undefined, + project: state.project || undefined, + location: state.location || undefined, + region: state.region || undefined, + custom_models: customModels.length > 0 ? customModels : undefined, + }; + + await onUpdate(provider.id, request); + setLastSaved(new Date()); + } catch (e) { + setSaveError(e instanceof Error ? e.message : "保存失败"); + } finally { + setIsSaving(false); + } + }, + [provider.id, provider.is_system, onUpdate], + ); + + // 防抖保存 + const debouncedSave = useCallback( + (state: FormState) => { + if (debounceTimerRef.current) { + clearTimeout(debounceTimerRef.current); + } + + debounceTimerRef.current = setTimeout(() => { + saveConfig(state); + }, DEBOUNCE_DELAY); + }, + [saveConfig], + ); + + // 清理定时器 + useEffect(() => { + return () => { + if (debounceTimerRef.current) { + clearTimeout(debounceTimerRef.current); + } + }; + }, []); + + // 处理字段变化 + const handleFieldChange = useCallback( + (field: keyof FormState, value: string) => { + setFormState((previousState) => { + const newState = { ...previousState, [field]: value }; + debouncedSave(newState); + return newState; + }); + }, + [debouncedSave], + ); + + const applyCustomModels = useCallback( + (models: string[]) => { + handleFieldChange("customModels", serializeCustomModels(models)); + }, + [handleFieldChange], + ); + + const setDefaultModel = useCallback( + (modelId: string) => { + const nextModels = selectedModels.filter( + (currentModel) => + currentModel.toLowerCase() !== modelId.toLowerCase(), + ); + applyCustomModels([modelId, ...nextModels]); + }, + [applyCustomModels, selectedModels], + ); + + const addModels = useCallback( + (modelIds: string[]) => { + const normalizedModels = dedupeModelIds(modelIds); + if (normalizedModels.length === 0) { + return; + } + + applyCustomModels([...selectedModels, ...normalizedModels]); + }, + [applyCustomModels, selectedModels], + ); + + useImperativeHandle( + ref, + () => ({ + setDefaultModel, + addModels, + }), + [addModels, setDefaultModel], + ); + + const handleAddModelDraft = useCallback(() => { + const draftModels = dedupeModelIds( + modelDraft .split(",") - .map((m) => m.trim()) - .filter((m) => m.length > 0); + .map((item) => item.trim()) + .filter((item) => item.length > 0), + ); - const request: UpdateProviderRequest = { - // 只有自定义 Provider 才发送 type 字段 - type: !provider.is_system ? state.providerType : undefined, - api_host: state.apiHost || undefined, - api_version: state.apiVersion || undefined, - project: state.project || undefined, - location: state.location || undefined, - region: state.region || undefined, - custom_models: customModels.length > 0 ? customModels : undefined, - }; - - await onUpdate(provider.id, request); - setLastSaved(new Date()); - } catch (e) { - setSaveError(e instanceof Error ? e.message : "保存失败"); - } finally { - setIsSaving(false); - } - }, - [provider.id, provider.is_system, onUpdate], - ); - - // 防抖保存 - const debouncedSave = useCallback( - (state: FormState) => { - if (debounceTimerRef.current) { - clearTimeout(debounceTimerRef.current); + if (draftModels.length === 0) { + return; } - debounceTimerRef.current = setTimeout(() => { - saveConfig(state); - }, DEBOUNCE_DELAY); - }, - [saveConfig], - ); + addModels(draftModels); + setModelDraft(""); + }, [addModels, modelDraft]); - // 清理定时器 - useEffect(() => { - return () => { - if (debounceTimerRef.current) { - clearTimeout(debounceTimerRef.current); + const handleRemoveModel = useCallback( + (modelId: string) => { + applyCustomModels( + selectedModels.filter( + (currentModel) => + currentModel.toLowerCase() !== modelId.toLowerCase(), + ), + ); + }, + [applyCustomModels, selectedModels], + ); + + useEffect(() => { + if (selectedModels.length > 0 || !recommendedLatestModel) { + return; } + + applyCustomModels([recommendedLatestModel.id]); + }, [applyCustomModels, recommendedLatestModel, selectedModels.length]); + + useEffect(() => { + onModelsChange?.(selectedModels); + }, [onModelsChange, selectedModels]); + + useEffect(() => { + onRecommendedLatestModelChange?.(recommendedLatestModel?.id ?? null); + }, [onRecommendedLatestModelChange, recommendedLatestModel]); + + // 获取当前 Provider 类型需要显示的额外字段 + // 使用 formState 中的 providerType,这样修改类型后会立即更新显示的字段 + const extraFields = PROVIDER_TYPE_FIELDS[formState.providerType] || []; + + // 格式化最后保存时间 + const formatLastSaved = (date: Date | null): string => { + if (!date) return ""; + return `已保存于 ${date.toLocaleTimeString("zh-CN")}`; }; - }, []); - // 处理字段变化 - const handleFieldChange = (field: keyof FormState, value: string) => { - const newState = { ...formState, [field]: value }; - setFormState(newState); - debouncedSave(newState); - }; - - // 获取当前 Provider 类型需要显示的额外字段 - // 使用 formState 中的 providerType,这样修改类型后会立即更新显示的字段 - const extraFields = PROVIDER_TYPE_FIELDS[formState.providerType] || []; - - // 格式化最后保存时间 - const formatLastSaved = (date: Date | null): string => { - if (!date) return ""; - return `已保存于 ${date.toLocaleTimeString("zh-CN")}`; - }; - - return ( -
- {/* Provider 类型选择器(仅自定义 Provider 显示) */} - {!provider.is_system && ( -
- - -

- 选择 API 协议类型,不同类型使用不同的请求格式 -

-
- )} - - {/* API Host 字段(所有 Provider 都有) */} -
- - handleFieldChange("apiHost", e.target.value)} - placeholder={FIELD_PLACEHOLDERS.apiHost} - disabled={loading || isSaving} - data-testid="api-host-input" - /> -

- {FIELD_HELP_TEXT.apiHost} -

-
- - {/* Azure OpenAI: API Version */} - {extraFields.includes("apiVersion") && ( -
- - handleFieldChange("apiVersion", e.target.value)} - placeholder={FIELD_PLACEHOLDERS.apiVersion} - disabled={loading || isSaving} - data-testid="api-version-input" - /> -

- {FIELD_HELP_TEXT.apiVersion} -

-
- )} - - {/* VertexAI: Project */} - {extraFields.includes("project") && ( -
- - handleFieldChange("project", e.target.value)} - placeholder={FIELD_PLACEHOLDERS.project} - disabled={loading || isSaving} - data-testid="project-input" - /> -

- {FIELD_HELP_TEXT.project} -

-
- )} - - {/* VertexAI: Location */} - {extraFields.includes("location") && ( -
- - handleFieldChange("location", e.target.value)} - placeholder={FIELD_PLACEHOLDERS.location} - disabled={loading || isSaving} - data-testid="location-input" - /> -

- {FIELD_HELP_TEXT.location} -

-
- )} - - {/* AWS Bedrock: Region */} - {extraFields.includes("region") && ( -
- - handleFieldChange("region", e.target.value)} - placeholder={FIELD_PLACEHOLDERS.region} - disabled={loading || isSaving} - data-testid="region-input" - /> -

- {FIELD_HELP_TEXT.region} -

-
- )} - - {/* 自定义模型列表 */} -
- - handleFieldChange("customModels", e.target.value)} - placeholder="glm-4, glm-4-flash, glm-4.7" - disabled={loading || isSaving} - data-testid="custom-models-input" - /> -

- 该 Provider 支持的模型列表,用逗号分隔。用于不支持 /models 接口的 - Provider(如智谱) -

-
- - {/* 保存状态指示 */} -
- {isSaving ? ( - - 保存中... - - ) : saveError ? ( - - {saveError} - - ) : lastSaved ? ( - - {formatLastSaved(lastSaved)} - - ) : ( - + return ( +
+ {/* Provider 类型选择器(仅自定义 Provider 显示) */} + {!provider.is_system && ( +
+ + +

+ 选择 API 协议类型,不同类型使用不同的请求格式 +

+
)} + + {/* API Host 字段(所有 Provider 都有) */} +
+ + handleFieldChange("apiHost", e.target.value)} + placeholder={FIELD_PLACEHOLDERS.apiHost} + disabled={loading || isSaving} + data-testid="api-host-input" + /> +

+ {FIELD_HELP_TEXT.apiHost} +

+
+ + {/* Azure OpenAI: API Version */} + {extraFields.includes("apiVersion") && ( +
+ + handleFieldChange("apiVersion", e.target.value)} + placeholder={FIELD_PLACEHOLDERS.apiVersion} + disabled={loading || isSaving} + data-testid="api-version-input" + /> +

+ {FIELD_HELP_TEXT.apiVersion} +

+
+ )} + + {/* VertexAI: Project */} + {extraFields.includes("project") && ( +
+ + handleFieldChange("project", e.target.value)} + placeholder={FIELD_PLACEHOLDERS.project} + disabled={loading || isSaving} + data-testid="project-input" + /> +

+ {FIELD_HELP_TEXT.project} +

+
+ )} + + {/* VertexAI: Location */} + {extraFields.includes("location") && ( +
+ + handleFieldChange("location", e.target.value)} + placeholder={FIELD_PLACEHOLDERS.location} + disabled={loading || isSaving} + data-testid="location-input" + /> +

+ {FIELD_HELP_TEXT.location} +

+
+ )} + + {/* AWS Bedrock: Region */} + {extraFields.includes("region") && ( +
+ + handleFieldChange("region", e.target.value)} + placeholder={FIELD_PLACEHOLDERS.region} + disabled={loading || isSaving} + data-testid="region-input" + /> +

+ {FIELD_HELP_TEXT.region} +

+
+ )} + + {/* 自定义模型列表 */} +
+
+
+
+ +

+ 手动添加模型后会保留在这里;默认模型请在右侧“模型能力”列表中点击选择。 +

+
+
+ + + +
+
+ {selectedModels.length > 0 ? ( + selectedModels.map((modelId, index) => { + const isLatest = recommendedLatestModel?.id === modelId; + return ( +
+ + {modelId} + + {index === 0 ? ( + 默认 + ) : null} + {isLatest ? ( + 最新 + ) : null} + {index > 0 ? ( + + ) : null} + +
+ ); + }) + ) : ( +

+ 尚未选择模型。检测到可用模型后,系统会默认填入最新模型。 +

+ )} +
+
+ +
+ setModelDraft(e.target.value)} + onKeyDown={(e) => { + if (e.key === "Enter" || e.key === ",") { + e.preventDefault(); + handleAddModelDraft(); + } + }} + placeholder="手动输入模型 ID,按 Enter 添加" + autoCapitalize="none" + autoCorrect="off" + spellCheck={false} + disabled={loading || isSaving} + data-testid="custom-models-input" + /> + +
+ +
+ + 第一个模型会作为默认模型,用于测试与默认请求;若未显式选择,则自动使用最新模型。 + + {recommendedLatestModel ? ( + + 当前推荐最新模型: + + {recommendedLatestModel.id} + + + ) : null} +
+ {localModelsError ? ( +

{localModelsError}

+ ) : null} + {localModelsLoading ? ( +

+ 正在加载模型列表... +

+ ) : null} +
+
+ + {/* 保存状态指示 */} +
+ {isSaving ? ( + + 保存中... + + ) : saveError ? ( + + {saveError} + + ) : lastSaved ? ( + + {formatLastSaved(lastSaved)} + + ) : ( + + )} +
-
- ); -}; + ); + }, +); -// ============================================================================ -// 辅助函数(用于测试) -// ============================================================================ - -/** - * 获取指定 Provider 类型需要显示的字段列表 - * 用于属性测试验证 Provider 类型处理正确性 - */ -export function getFieldsForProviderType(type: ProviderType): string[] { - const baseFields = ["apiHost"]; - const extraFields = PROVIDER_TYPE_FIELDS[type] || []; - return [...baseFields, ...extraFields]; -} - -/** - * 验证 Provider 类型是否需要特定字段 - */ -export function providerTypeRequiresField( - type: ProviderType, - field: string, -): boolean { - if (field === "apiHost") return true; - const extraFields = PROVIDER_TYPE_FIELDS[type] || []; - return extraFields.includes(field); -} +ProviderConfigForm.displayName = "ProviderConfigForm"; export default ProviderConfigForm; diff --git a/src/components/provider-pool/api-key/ProviderConfigForm.utils.ts b/src/components/provider-pool/api-key/ProviderConfigForm.utils.ts new file mode 100644 index 000000000..e2c2a689f --- /dev/null +++ b/src/components/provider-pool/api-key/ProviderConfigForm.utils.ts @@ -0,0 +1,140 @@ +/** + * @file ProviderConfigForm 工具函数 + * @description Provider 配置表单的模型与字段辅助逻辑 + * @module components/provider-pool/api-key/ProviderConfigForm.utils + */ + +import type { ProviderType } from "@/lib/types/provider"; +import type { EnhancedModelMetadata } from "@/lib/types/modelRegistry"; + +/** 支持的 Provider 类型列表 */ +export const PROVIDER_TYPE_OPTIONS: { value: ProviderType; label: string }[] = [ + { value: "openai", label: "OpenAI 兼容" }, + { value: "openai-response", label: "OpenAI Responses API" }, + { value: "codex", label: "Codex CLI" }, + { value: "anthropic", label: "Anthropic" }, + { value: "anthropic-compatible", label: "Anthropic 兼容" }, + { value: "gemini", label: "Gemini" }, + { value: "azure-openai", label: "Azure OpenAI" }, + { value: "vertexai", label: "VertexAI" }, + { value: "aws-bedrock", label: "AWS Bedrock" }, + { value: "ollama", label: "Ollama" }, + { value: "fal", label: "Fal" }, + { value: "new-api", label: "New API" }, + { value: "gateway", label: "Vercel AI Gateway" }, +]; + +/** 支持的 Provider 类型值列表 */ +export const PROVIDER_TYPE_VALUES: ProviderType[] = PROVIDER_TYPE_OPTIONS.map( + (option) => option.value, +); + +/** Provider 类型对应的额外字段配置 */ +export const PROVIDER_TYPE_FIELDS: Record = { + openai: [], + "openai-response": [], + codex: [], + anthropic: [], + "anthropic-compatible": [], + gemini: [], + "azure-openai": ["apiVersion"], + vertexai: ["project", "location"], + "aws-bedrock": ["region"], + ollama: [], + fal: [], + "new-api": [], + gateway: [], +}; + +export function isSupportedProviderType( + providerType: string, +): providerType is ProviderType { + return PROVIDER_TYPE_VALUES.includes(providerType as ProviderType); +} + +export function dedupeModelIds(modelIds: string[]): string[] { + const seen = new Set(); + const result: string[] = []; + + for (const modelId of modelIds) { + const trimmed = modelId.trim(); + if (!trimmed) { + continue; + } + + const key = trimmed.toLowerCase(); + if (seen.has(key)) { + continue; + } + + seen.add(key); + result.push(trimmed); + } + + return result; +} + +export function parseCustomModelsValue(value: string): string[] { + return dedupeModelIds( + value + .split(",") + .map((item) => item.trim()) + .filter((item) => item.length > 0), + ); +} + +export function serializeCustomModels(models: string[]): string { + return dedupeModelIds(models).join(", "); +} + +export function sortSelectableModels( + models: EnhancedModelMetadata[], +): EnhancedModelMetadata[] { + return [...models].sort((a, b) => { + if (a.is_latest && !b.is_latest) return -1; + if (!a.is_latest && b.is_latest) return 1; + + if (a.release_date && b.release_date && a.release_date !== b.release_date) { + return b.release_date.localeCompare(a.release_date); + } + if (a.release_date && !b.release_date) return -1; + if (!a.release_date && b.release_date) return 1; + + const tierWeight: Record = { max: 3, pro: 2, mini: 1 }; + const aTierWeight = tierWeight[a.tier] ?? 0; + const bTierWeight = tierWeight[b.tier] ?? 0; + if (aTierWeight !== bTierWeight) { + return bTierWeight - aTierWeight; + } + + return a.display_name.localeCompare(b.display_name); + }); +} + +export function getLatestSelectableModel( + models: EnhancedModelMetadata[], +): EnhancedModelMetadata | null { + return sortSelectableModels(models)[0] ?? null; +} + +/** + * 获取指定 Provider 类型需要显示的字段列表 + * 用于属性测试验证 Provider 类型处理正确性 + */ +export function getFieldsForProviderType(type: ProviderType): string[] { + const baseFields = ["apiHost"]; + const extraFields = PROVIDER_TYPE_FIELDS[type] || []; + return [...baseFields, ...extraFields]; +} + +/** + * 验证 Provider 类型是否需要特定字段 + */ +export function providerTypeRequiresField( + type: ProviderType, + field: string, +): boolean { + if (field === "apiHost") return true; + const extraFields = PROVIDER_TYPE_FIELDS[type] || []; + return extraFields.includes(field); +} diff --git a/src/components/provider-pool/api-key/ProviderModelList.tsx b/src/components/provider-pool/api-key/ProviderModelList.tsx index da9315dac..9e36c5504 100644 --- a/src/components/provider-pool/api-key/ProviderModelList.tsx +++ b/src/components/provider-pool/api-key/ProviderModelList.tsx @@ -12,12 +12,14 @@ import { Wrench, Brain, Sparkles, + Check, Loader2, RefreshCw, Cloud, HardDrive, } from "lucide-react"; import { Button } from "@/components/ui/button"; +import { Badge } from "@/components/ui/badge"; import { Tooltip, TooltipContent, @@ -31,6 +33,7 @@ import { buildCatalogAliasMap, resolveRegistryProviderId, } from "./providerTypeMapping"; +import { getLatestSelectableModel } from "./ProviderConfigForm.utils"; import { invoke } from "@tauri-apps/api/core"; // ============================================================================ @@ -42,6 +45,14 @@ export interface ProviderModelListProps { providerId: string; /** Provider 类型(API 协议),如 "anthropic", "openai", "gemini" */ providerType: string; + /** 当前默认模型 ID */ + selectedModelId?: string | null; + /** 推荐最新模型 ID */ + latestModelId?: string | null; + /** 点击模型时设为默认模型 */ + onSelectModel?: (modelId: string) => void; + /** 当前列表解析出的最新模型变化回调 */ + onLatestModelResolved?: (modelId: string | null) => void; /** 是否有可用的 API Key(用于显示刷新按钮) */ hasApiKey?: boolean; /** 额外的 CSS 类名 */ @@ -71,6 +82,17 @@ interface FetchModelsResult { should_prompt_error?: boolean; } +interface CachedProviderModels { + models: EnhancedModelMetadata[]; + source: "Api" | "LocalFallback" | null; + error: string | null; + requestUrl: string | null; + diagnosticHint: string | null; + shouldPromptError: boolean; +} + +const providerModelsCache = new Map(); + function buildApiDiagnosticLines(result: { error: string | null; request_url?: string | null; @@ -99,15 +121,42 @@ function buildApiDiagnosticLines(result: { interface ModelItemProps { model: EnhancedModelMetadata; + isDefault: boolean; + isLatest: boolean; + onSelect?: (modelId: string) => void; } /** * 单个模型项 */ -const ModelItem: React.FC = ({ model }) => { +const ModelItem: React.FC = ({ + model, + isDefault, + isLatest, + onSelect, +}) => { return (
onSelect(model.id) : undefined} + role={onSelect ? "button" : undefined} + tabIndex={onSelect ? 0 : undefined} + onKeyDown={ + onSelect + ? (event) => { + if (event.key === "Enter" || event.key === " ") { + event.preventDefault(); + onSelect(model.id); + } + } + : undefined + } data-testid={`model-item-${model.id}`} >
@@ -115,12 +164,15 @@ const ModelItem: React.FC = ({ model }) => { {model.display_name} + {isLatest ? 最新 : null} + {isDefault ? 默认 : null}
{model.id}
{/* 能力标签 */}
+ {isDefault && } {model.capabilities.vision && ( = ({ model }) => { export const ProviderModelList: React.FC = ({ providerId, providerType, + selectedModelId = null, + latestModelId = null, + onSelectModel, + onLatestModelResolved, hasApiKey = false, className, maxItems, @@ -262,6 +318,27 @@ export const ProviderModelList: React.FC = ({ null, ); const [apiShouldPromptError, setApiShouldPromptError] = useState(false); + const cacheKey = `${providerId}:${providerType}`; + + useEffect(() => { + const cached = providerModelsCache.get(cacheKey); + if (!cached) { + setApiModels(null); + setApiSource(null); + setApiError(null); + setApiRequestUrl(null); + setApiDiagnosticHint(null); + setApiShouldPromptError(false); + return; + } + + setApiModels(cached.models); + setApiSource(cached.source); + setApiError(cached.error); + setApiRequestUrl(cached.requestUrl); + setApiDiagnosticHint(cached.diagnosticHint); + setApiShouldPromptError(cached.shouldPromptError); + }, [cacheKey]); // 从 API 获取模型列表(自动获取 API Key) const handleRefreshFromApi = useCallback(async () => { @@ -287,7 +364,18 @@ export const ProviderModelList: React.FC = ({ setApiShouldPromptError(Boolean(result.should_prompt_error)); if (result.error) { setApiError(result.error); + } else { + setApiError(null); } + + providerModelsCache.set(cacheKey, { + models: result.models, + source: result.source ?? null, + error: result.error ?? null, + requestUrl: result.request_url ?? null, + diagnosticHint: result.diagnostic_hint ?? null, + shouldPromptError: Boolean(result.should_prompt_error), + }); } else { setApiError("返回结果格式错误"); } @@ -296,7 +384,7 @@ export const ProviderModelList: React.FC = ({ } finally { setRefreshing(false); } - }, [providerId]); + }, [cacheKey, providerId]); // 使用 API 模型或本地模型 const displayModelsSource = apiModels ?? models; @@ -315,6 +403,16 @@ export const ProviderModelList: React.FC = ({ }, [displayModelsSource, maxItems]); const hasMore = maxItems && displayModelsSource.length > maxItems; + const resolvedLatestModelId = + latestModelId ?? getLatestSelectableModel(displayModelsSource)?.id ?? null; + const effectiveDefaultModelId = selectedModelId ?? resolvedLatestModelId; + const resolvedLatestModelKey = resolvedLatestModelId?.toLowerCase() ?? null; + const effectiveDefaultModelKey = + effectiveDefaultModelId?.toLowerCase() ?? null; + + useEffect(() => { + onLatestModelResolved?.(resolvedLatestModelId); + }, [onLatestModelResolved, resolvedLatestModelId]); // 加载状态 if (loading && !apiModels) { @@ -349,10 +447,17 @@ export const ProviderModelList: React.FC = ({ return (
-

- - 支持的模型 -

+
+

+ + 支持的模型 +

+ {onSelectModel ? ( +

+ 点击模型即可设为默认模型;未显式选择时,自动使用最新模型。 +

+ ) : null} +
{hasApiKey && ( @@ -503,7 +608,13 @@ export const ProviderModelList: React.FC = ({ {/* 模型列表 */}
{displayModels.map((model) => ( - + ))}
diff --git a/src/components/provider-pool/api-key/ProviderSetting.tsx b/src/components/provider-pool/api-key/ProviderSetting.tsx index 03cd6591c..376bfad06 100644 --- a/src/components/provider-pool/api-key/ProviderSetting.tsx +++ b/src/components/provider-pool/api-key/ProviderSetting.tsx @@ -7,10 +7,11 @@ * **Validates: Requirements 4.1, 6.3, 6.4** */ -import React, { useState } from "react"; +import React, { useCallback, useEffect, useRef, useState } from "react"; import { cn } from "@/lib/utils"; import { Switch } from "@/components/ui/switch"; import { Button } from "@/components/ui/button"; +import { Badge } from "@/components/ui/badge"; import { Dialog, DialogContent, @@ -23,7 +24,10 @@ import { Textarea } from "@/components/ui/textarea"; import { Trash2 } from "lucide-react"; import { ProviderIcon } from "@/icons/providers"; import { ApiKeyList } from "./ApiKeyList"; -import { ProviderConfigForm } from "./ProviderConfigForm"; +import { + ProviderConfigForm, + type ProviderConfigFormRef, +} from "./ProviderConfigForm"; import { ConnectionTestButton, ConnectionTestResult, @@ -103,10 +107,48 @@ export const ProviderSetting: React.FC = ({ loading = false, className, }) => { + const providerConfigFormRef = useRef(null); const [chatDialogOpen, setChatDialogOpen] = useState(false); const [chatPrompt, setChatPrompt] = useState("hello"); const [chatTesting, setChatTesting] = useState(false); const [chatResult, setChatResult] = useState(null); + const [draftCustomModels, setDraftCustomModels] = useState( + provider?.custom_models ?? [], + ); + const [recommendedLatestModelId, setRecommendedLatestModelId] = useState< + string | null + >(null); + const enabledApiKeyCount = + provider?.api_keys?.filter((apiKey) => apiKey.enabled).length ?? 0; + const defaultModel = draftCustomModels[0] ?? recommendedLatestModelId ?? null; + + useEffect(() => { + setDraftCustomModels(provider?.custom_models ?? []); + setRecommendedLatestModelId(null); + }, [provider?.id, provider?.custom_models]); + + const handleModelsChange = useCallback((models: string[]) => { + setDraftCustomModels(models); + }, []); + + const handleRecommendedLatestModelChange = useCallback( + (modelId: string | null) => { + setRecommendedLatestModelId(modelId); + }, + [], + ); + + const handleSelectDefaultModel = useCallback((modelId: string) => { + providerConfigFormRef.current?.setDefaultModel(modelId); + }, []); + + useEffect(() => { + if (draftCustomModels.length > 0 || !recommendedLatestModelId) { + return; + } + + providerConfigFormRef.current?.setDefaultModel(recommendedLatestModelId); + }, [draftCustomModels.length, recommendedLatestModelId]); const handleChatTest = async () => { if (!onTestChat || chatTesting || !provider) return; @@ -136,19 +178,61 @@ export const ProviderSetting: React.FC = ({ return (
-
-

请从左侧列表选择一个 Provider

-

选择后可在此处配置 API Key 和其他设置

+
+

+ 请从左侧列表选择一个 Provider +

+

+ 选择后可在此处集中管理 API Key、模型、连接测试与支持模型信息。 +

); } + const providerHostLabel = (() => { + try { + const url = new URL(provider.api_host); + return `${url.host}${url.pathname === "/" ? "" : url.pathname}`; + } catch { + return provider.api_host; + } + })(); + + const summaryItems = [ + { + label: "可用密钥", + value: `${enabledApiKeyCount}`, + hint: `共 ${provider.api_keys?.length ?? 0} 个 API Key`, + compact: false, + }, + { + label: "默认模型", + value: defaultModel ?? "未设置", + hint: defaultModel + ? "第一个模型用于默认请求与测试" + : "可在下方配置中指定", + compact: true, + }, + { + label: "协议类型", + value: provider.type, + hint: provider.is_system ? "系统预设 Provider" : "自定义 Provider", + compact: true, + }, + { + label: "接口地址", + value: providerHostLabel, + hint: provider.api_host, + compact: true, + }, + ]; + // 处理启用/禁用切换 const handleToggleEnabled = async (enabled: boolean) => { if (onUpdate) { @@ -158,216 +242,305 @@ export const ProviderSetting: React.FC = ({ return (
- {/* Provider 头部 */} -
- {/* 图标 */} - - - {/* 名称和类型 */} -
-

+
+ {/* Provider 头部 */} +
- {provider.name} -

-

- 类型: {provider.type} - {provider.is_system && ( - - 系统预设 - - )} -

-
- - {/* 启用开关 */} -
- - {provider.enabled ? "已启用" : "已禁用"} - - -
- - {/* 删除按钮(仅自定义 Provider) */} - {!provider.is_system && onDeleteProvider && ( - - )} -
- - {/* 内容区域 */} -
- {/* API Key 列表 */} -
- -
- - {/* 分隔线 */} -
- - {/* Provider 配置表单 */} -
-

配置

- -
- - {/* 分隔线 */} -
- - {/* 连接测试 */} -
-

连接测试

-
- - -
- {(provider.api_keys?.length ?? 0) === 0 && ( -

- 请先添加 API Key 后再进行连接测试 -

- )} - - - - - 对话测试 - - 发送一条最小对话请求,直接查看返回内容或原始错误,便于排查模型/权限/路由问题。 - - - -
-