diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 712aa7f47..fe9f270c0 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -41,9 +41,6 @@ jobs: - platform: windows-2022 target: x86_64-pc-windows-msvc name: Windows-x64 - - platform: ubuntu-22.04 - target: x86_64-unknown-linux-gnu - name: Linux-x64 runs-on: ${{ matrix.platform }} @@ -63,12 +60,6 @@ jobs: with: version: 9 - - name: Install Linux dependencies - if: matrix.platform == 'ubuntu-22.04' - run: | - sudo apt-get update - sudo apt-get install -y libwebkit2gtk-4.1-dev libappindicator3-dev librsvg2-dev patchelf libssl-dev pkg-config libasound2-dev libxdo-dev - - name: Setup Rust uses: dtolnay/rust-toolchain@stable with: @@ -160,30 +151,6 @@ jobs: security set-key-partition-list -S apple-tool:,apple: -k "$KEYCHAIN_PASSWORD" $KEYCHAIN_PATH security list-keychain -d user -s $KEYCHAIN_PATH - - name: Build Tauri app (Linux) - if: matrix.platform == 'ubuntu-22.04' - uses: tauri-apps/tauri-action@v0 - env: - GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} - # 禁用 LTO 加速编译(正式发布可改为 thin) - CARGO_PROFILE_RELEASE_LTO: "off" - # 增加并行编译单元 - CARGO_PROFILE_RELEASE_CODEGEN_UNITS: 32 - CARGO_INCREMENTAL: 0 - SCCACHE_GHA_ENABLED: "true" - RUSTC_WRAPPER: sccache - OPENSSL_NO_VENDOR: "1" - with: - tauriScript: npx tauri - projectPath: src-tauri - tagName: ${{ github.event.inputs.version || github.ref_name }} - releaseName: "Lime ${{ github.event.inputs.version || github.ref_name }}" - releaseBody: ${{ env.RELEASE_BODY }} - releaseDraft: false - prerelease: false - # 默认不启用 voice feature(包含 whisper-rs,编译很慢) - args: --target ${{ matrix.target }} - - name: Build Tauri app (Windows offline) if: matrix.platform == 'windows-2022' uses: tauri-apps/tauri-action@v0 diff --git a/.gitignore b/.gitignore index 509c9acc2..6cd66731d 100644 --- a/.gitignore +++ b/.gitignore @@ -65,9 +65,11 @@ openspec # Build artifacts (alternate cargo target dirs) src-tauri/target*/ +src-tauri/.codex* logo-lime.png lime-claw.png lime.db +.codex-* \ No newline at end of file diff --git a/README.md b/README.md index adc71c1fa..a10d83c34 100644 --- a/README.md +++ b/README.md @@ -1,25 +1,54 @@
+Lime Logo + # Lime -**以创作为中心的本地优先 AI Agent 交互工作台** +### 青柠一下,灵感即来 -一句话:用 Skills 组织经验与流程,用 MCP 接入标准能力,用 Claw 渠道把 Agent 带到飞书、Telegram 等入口,让创作、研究、执行与交付在同一个工作环境里闭合。 +**从一句想法,到成稿、成图、成片、成事** + +本地优先的 AI Agent 创作工作台
--- +## 界面预览 + +Lime 把创作工作台、任务会话与 Provider 管理整合在同一个本地优先桌面应用里,下面是三个核心界面的快速预览: + +### 工作台首页 + +Lime 工作台首页预览 + +在同一个 Workspace 中组织任务、技能、自动化、浏览器协助与对话输入,让一句想法继续推进为可执行结果。 + +### 任务与会话协作 + +Lime 任务与会话协作预览 + +围绕单个任务持续补充上下文、追踪执行轨迹,并在同一界面中完成对话推进、模型切换与结果沉淀。 + +### Provider 与凭证管理 + +Lime Provider 与凭证管理预览 + +统一管理 API Key、Connect、语音服务与 OAuth 凭证,并在同一处完成启用状态、模型配置与连接测试。 + +--- + ## 这是什么 -Lime 是一个基于 Tauri 的桌面应用,面向创作者、内容团队与轻知识工作者。它把 Workspace、Agent、Skills、MCP、Claw 渠道和 Artifact 交付整合到同一个桌面环境里,让工作从输入需求直接走向可沉淀、可复用、可继续执行的结果。 +Lime 是一个基于 Tauri 的桌面应用,面向创作者、内容团队与轻知识工作者。它把 Workspace、Agent、Skills、MCP、Claw 渠道和 Artifact 交付整合到同一个桌面环境里,让一句想法可以继续推进为成稿、成图、成片,并最终走向可执行、可协作、可沉淀的结果。 -你可以在一个地方完成: +你可以在一个地方完成从想法到交付的整条链路: -- 在项目里与 Agent 协作 -- 生成和编辑文档、脚本、图文方案等产物 -- 使用浏览器、终端、MCP 和插件扩展执行空间 -- 让结果沉淀为可复用的记忆、风格和版本资产 +- 成稿:生成和编辑文档、脚本、提纲、长文等内容产物 +- 成图:产出图文方案、海报草稿与视觉素材 +- 成片:围绕视频脚本、分镜、素材与创作流程持续协作 +- 成事:通过 Workspace、MCP、Claw 和插件把结果继续执行、协作与交付 +- 沉淀:让对话、版本、记忆、风格和项目资产可复用、可继续 --- @@ -126,6 +155,7 @@ brew install --cask lime 从 [Releases](https://github.com/aiclientproxy/lime/releases) 下载对应平台安装包。 +- 当前仅提供 macOS 与 Windows 发布包,Linux 桌面端已暂停支持 - Windows 用户优先下载 `Lime_*_x64-offline-setup.exe`(NSIS 离线安装器,内置 WebView2,安装更完整) - 如果只想下载更小的安装器,且当前网络可稳定访问微软下载源,再选择 `Lime_*_x64-online-setup.exe` - 如被 SmartScreen 拦截,属于未签名或签名信誉不足的 Windows 常见提示,不代表安装包必然损坏 @@ -161,3 +191,15 @@ npm run tauri build 本项目仅供学习研究使用,用户需自行承担使用风险。 本项目不直接提供 AI 模型服务,模型能力由第三方提供商提供。 + +--- + +
+ +### 微信交流 + +Lime 微信交流群二维码 + +扫码加微信,备注 `Lime`,拉你进群讨论。 + +
diff --git a/RELEASE_NOTES.md b/RELEASE_NOTES.md index cb8c5acb7..af962eb44 100644 --- a/RELEASE_NOTES.md +++ b/RELEASE_NOTES.md @@ -1,26 +1,32 @@ -## Lime v0.90.0 +## Lime v0.91.0 ### ✨ 主要更新 -- **Aster Agent 聊天链路完成收口**:前端统一走 `useAgentChatUnified -> useAsterAgentChat`,移除旧 `useAgentChat` / `agentStore` compat 路径;会话恢复、流式消息、工具状态与 topic snapshot 改为模块化协作 -- **Agent 会话与时间线持久化增强**:Aster session store、agent timeline DAO、数据库 schema 与迁移链路继续补强,聊天历史、统计与项目上下文构建更稳定 -- **托盘模型快捷切换上线**:新增全局托盘模型同步、主题感知的快捷模型组与模型选择联动,聊天页、侧边栏和模型选择器体验同步更新 -- **调试与性能诊断能力补齐**:新增前端 debug 上报 API、Tauri profiling 启动脚本、Perfetto / Tokio Console 支持与配套文档,运行时排障更直接 -- **OpenClaw 与运行时集成继续完善**:Browser Runtime、OpenClaw 页面与后端服务、Dev Bridge 查询和技能服务继续收口,提升桌面侧运行时协同 -- **品牌图标与托盘资产刷新**:应用图标、托盘状态图与启动页、侧边栏视觉资源同步更新 +- **Aster 运行时队列正式接入 Lime**:桌面端补齐 runtime queue service、Aster state support 与 session store 协作,Agent 会话恢复、排队执行和 runtime item 映射进一步收口 +- **Agent 聊天输入链路继续统一**:输入栏、空状态、图片附件、模型选择与会话 hooks 继续围绕 Aster 聊天主链路整理,减少旧 compat 路径分叉 +- **模型能力与视觉提示增强**:新增模型能力徽章、视觉能力提示与 provider model list 整理,模型选择和多模态提示更直接 +- **Skills / Social Post 执行链路补齐**:新增技能执行运行时与社交内容技能集成,Aster Skills 在 Lime 内的发现、执行与同步更完整 +- **数据库与治理清理继续推进**:移除旧 `unified_chat` / `tool_hooks` / `three stage workflow` 相关残留,统一到现役 Aster Agent、Memory 与 Workspace 路径 ### ⚠️ 兼容性说明 -- Agent 聊天现役事实源已统一到 Aster 后端;旧 `useAgentChat` / `agentStore` 已移除,后续新功能不再沿 compat 路径扩展 -- Profiling 诊断能力仅在显式开发启动流程下启用;release / 生产构建默认忽略这些调试开关 +- Aster 相关聊天、会话与时间线事实源进一步集中到新的 runtime / session store 路径,旧 compat API 不再建议继续扩展 +- 模型可见性与能力展示依赖新的 provider model 推断逻辑,历史仅按名称匹配的前端分支需要逐步淘汰 + +### 🔗 依赖同步 + +- `src-tauri/Cargo.toml` 中的 `aster-rust` 依赖固定到 `v0.19.0` ### 🧪 测试 -- 发布前执行:`cargo test`、`cargo fmt --all`、`cargo clippy`、`npm run lint` +- 发布前执行:`cd src-tauri && cargo test` +- 发布前执行:`cd src-tauri && cargo fmt --all --check` +- 发布前执行:`cd src-tauri && cargo clippy` +- 发布前执行:`npm run lint` ### 📝 文档 -- 更新 Agent / Aster 集成、治理与 profiling 相关文档,补充前端诊断与发布说明 +- 更新 Aster 集成、治理、Skills 与发布相关文档,补充当前现役架构与发布说明 ### 📦 Windows 下载说明 @@ -30,4 +36,4 @@ --- -**完整变更**: v0.89.1...v0.90.0 +**完整变更**: v0.90.0...v0.91.0 diff --git a/docs/aiprompts/aster-integration.md b/docs/aiprompts/aster-integration.md index 36cf3e260..154aca4c1 100644 --- a/docs/aiprompts/aster-integration.md +++ b/docs/aiprompts/aster-integration.md @@ -7,6 +7,8 @@ Lime 已完整集成 aster-rust 框架,包括凭证池桥接。 ## 当前事实源 - `Aster thread / turn / item runtime` 是运行态事实源。 +- Aster shared runtime store 与全局 session store 必须在 Lime bootstrap 启动期显式初始化;启动恢复时由 runtime support 统一完成 legacy queue 迁移与 queued session 枚举。 +- 运行时命令与 service 只允许读取已准备好的 shared store / shared queue service,不再在热路径偷偷 fallback 到默认路径。 - Lime 只负责事件映射、数据库投影和 UI 派生,不再伪造核心 runtime item。 - 会话删除统一收口到存储边界;命令层和 Dev Bridge 不应直接调用 `AgentDao::delete_session`。 - 需要恢复运行态时,优先从 Aster runtime 恢复,再映射到 Lime timeline。 diff --git a/docs/aiprompts/commands.md b/docs/aiprompts/commands.md index 747f70128..3a231fdc7 100644 --- a/docs/aiprompts/commands.md +++ b/docs/aiprompts/commands.md @@ -23,9 +23,12 @@ Tauri 命令是前端与 Rust 后端通信的边界,但前端业务代码**不 ## 当前事实源 -- 聊天主命令:`chat_*` +- Agent / Codex 主命令:`agent_runtime_*` +- 运行态摘要主链:Aster `runtime_status` item -> timeline `turn_summary` +- `chat_*` 已停止注册,且不再纳入 `commands::mod` 编译图;旧 General / Creator / 历史桥接如仍需恢复,必须显式走新的 compat 评审 - 旧 `general_chat_*` 前端 compat 网关与 Rust 命令已删除 -- 当前剩余治理重点:统计、记忆等旁路仍在读取 `general_chat_*` 历史表 +- `Tauri runtime_status` 事件只保留前端瞬时状态用途,不再作为 timeline 事实源 +- 当前剩余治理重点:统计、记忆等旁路继续按 `runtime context` 与 `durable knowledge` 分层收口 ## 治理案例:记忆系统 @@ -45,7 +48,8 @@ Tauri 命令是前端与 Rust 后端通信的边界,但前端业务代码**不 4. 如果存在旧命令又无任何调用,就直接删掉命令注册、桥接和 mock,不要继续保留空兼容壳 同理,对话系统也不应该重新引回已经删除的 `general_chat_*` 命令; -后续如需扩展聊天能力,应继续收敛到 `chat_*` 与对应网关。 +后续如需扩展 Agent / Codex 工作流,应继续收敛到 `agent_runtime_*` 与对应网关; +`chat_*` 只允许作为 dead-candidate 参考,不应重新回到 `commands::mod` 或 `generate_handler!`。 ## 目录结构 diff --git a/docs/aiprompts/content-creator.md b/docs/aiprompts/content-creator.md index 14fbe2daf..567d615cc 100644 --- a/docs/aiprompts/content-creator.md +++ b/docs/aiprompts/content-creator.md @@ -11,7 +11,7 @@ ↓ 用户发送消息 → useAgentChatUnified.sendMessage() ↓ -统一收口到 useAsterAgentChat / agent_runtime_* +统一收口到 useAgentChatUnified / useAsterAgentChat / agent_runtime_* ↓ 第一条消息时注入 systemPrompt → 发送到 Aster Agent ↓ @@ -40,7 +40,7 @@ src/components/ │ └── parser.ts # A2UI 和 write_file 解析器 ├── agent/chat/ │ ├── hooks/ -│ │ ├── index.ts # useAgentChatUnified 统一入口 +│ │ ├── index.ts # Agent chat hooks 导出入口 │ │ └── useAsterAgentChat.ts # Agent 聊天主 Hook │ ├── components/ │ │ ├── StreamingRenderer.tsx # 流式渲染(解析 write_file) @@ -119,7 +119,7 @@ export function parseAIResponse( - `write_file` - 完整的文件写入 - `pending_write_file` - 流式传输中的文件写入 -### 3. useAgentChatUnified / useAsterAgentChat - systemPrompt 注入 +### 3. useAsterAgentChat - systemPrompt 注入 在发送第一条消息时注入 systemPrompt,并通过现役 runtime adapter 提交到 `agent_runtime_*`。 diff --git a/docs/aiprompts/governance.md b/docs/aiprompts/governance.md index 4fc5a637a..c92d99bfc 100644 --- a/docs/aiprompts/governance.md +++ b/docs/aiprompts/governance.md @@ -177,7 +177,7 @@ npm run governance:legacy-report 例如: -> 聊天能力后续统一收敛到 `useUnifiedChat + chat_* + ChatDao`。 +> 聊天能力后续统一收敛到 `useAgentChatUnified -> useAsterAgentChat -> agent_runtime_* + lime_core::database::agent_session_repository`。 ### 第三步:优先做减法 @@ -207,8 +207,8 @@ npm run governance:legacy-report 以聊天系统为例,遇到新旧并存时,必须同时问这几个问题: -- 前端唯一入口是不是 `useAgentChatUnified` / `useAsterAgentChat`,还是 `useChat` / `useAgentChat` 还在继续长逻辑? -- Rust 唯一入口是不是 `chat_*`,还是 `general_chat_*` / `agent_*` / `aster_agent_*` 还在平行演进? +- 前端唯一入口是不是 `useAgentChatUnified -> useAsterAgentChat`,还是 `useChat` / `useAgentChat` / `useUnifiedChat` 还在继续长逻辑? +- Rust 唯一入口是不是 `agent_runtime_*`,还是 `chat_*` / `general_chat_*` / `agent_*` / `aster_agent_*` 还在平行演进? - 数据事实源是不是同一组表 / 同一套 Repository,还是还在同时写 `agent_*` 与 `general_chat_*`? - 统计、记忆等旁路是不是已经切到新路径,还是还在读旧表? diff --git a/docs/aiprompts/hooks.md b/docs/aiprompts/hooks.md index b86ef5c17..02786951b 100644 --- a/docs/aiprompts/hooks.md +++ b/docs/aiprompts/hooks.md @@ -9,7 +9,6 @@ ``` src/hooks/ ├── index.ts # 导出入口 -├── useUnifiedChat.ts # 统一对话 Hook(新) ├── useProviderPool.ts # 凭证池管理 ├── useOAuthCredentials.ts # OAuth 凭证 ├── useFlowEvents.ts # 流量事件 @@ -22,69 +21,24 @@ src/hooks/ - 新的前端能力优先落在 `src/lib/api/*`,再由 Hook 或组件消费。 - 历史 `useTauri.ts` 兼容聚合层已删除,不要重新引入新的“大一统 API Hook”。 -- 旧聊天链路优先迁移到 `@/hooks/useUnifiedChat`,不要继续扩散 `useChat` / compat Hook。 +- Agent 工作台统一走 `src/components/agent/chat/hooks/index.ts` 暴露的 `useAgentChatUnified`,底层实现委托 `useAsterAgentChat`。 +- 历史 `@/hooks/useUnifiedChat` 与 `src/lib/api/unified-chat.ts` 已删除,不要重建 compat Hook / API。 ## 核心 Hooks -### useUnifiedChat(统一对话) +### useAgentChatUnified / useAsterAgentChat(现役 Agent 对话) -统一的对话 Hook,支持三种模式:Agent、General、Creator。 +现役 Agent / Codex 工作台事实源: -```typescript -import { useUnifiedChat } from "@/hooks/useUnifiedChat"; - -// Agent 模式 - 支持工具调用 -const { messages, sendMessage, stopGeneration } = useUnifiedChat({ - mode: "agent", - providerType: "claude", - model: "claude-sonnet-4-20250514", -}); - -// Creator 模式 - 支持画布输出 -const creatorChat = useUnifiedChat({ - mode: "creator", - systemPrompt: "你是内容创作助手...", - harnessConfig: { - theme: "social-media", - artifactMode: "version-chain", - }, - onHarnessEvent: (event) => { - /* 接收阶段推进、产物创建等语义事件 */ - }, - onArtifactUpdate: (artifact) => { - /* 接收产物快照 */ - }, - onCanvasUpdate: (path, content) => { - /* 更新画布 */ - }, - onWriteFile: (content, fileName) => { - /* 文件写入 */ - }, -}); - -// General 模式 - 纯文本对话 -const generalChat = useUnifiedChat({ mode: "general" }); -``` - -**返回值**: - -- `session` - 当前会话 -- `messages` - 消息列表 -- `isLoading` / `isSending` - 状态 -- `createSession()` / `loadSession()` / `deleteSession()` - 会话管理 -- `sendMessage()` / `stopGeneration()` - 消息操作 -- `configureProvider()` - Provider 配置 - -**补充说明**: - -- Creator 模式现在支持 `harnessConfig`、`onHarnessEvent`、`onArtifactUpdate` -- 社媒内容推荐把 `` 结果投影为“版本链产物”,而不是只按文件名覆盖 +- `useAgentChatUnified -> useAsterAgentChat -> useAgentContext / useAgentSession / useAgentTools / useAgentStream` +- 命令主链:`agent_runtime_submit_turn -> runtime items(plan / runtime_status / artifact / tool / action) -> action_required -> respond_action` +- 适用场景:Agent 工作台、任务执行、工具审批、timeline 渲染 **相关文件**: -- 类型定义:`src/types/chat.ts` -- API 封装:`src/lib/api/unified-chat.ts` -- 架构文档:`docs/prd/chat-architecture-redesign.md` +- 统一入口:`src/components/agent/chat/hooks/index.ts` +- Hook 实现:`src/components/agent/chat/hooks/useAsterAgentChat.ts` +- API 封装:`src/lib/api/agentRuntime.ts` ### useProviderPool diff --git a/docs/aiprompts/services.md b/docs/aiprompts/services.md index 22fcdac6b..b2d33045f 100644 --- a/docs/aiprompts/services.md +++ b/docs/aiprompts/services.md @@ -22,7 +22,7 @@ src-tauri/src/services/ ## 核心服务 > 注意:`general_chat/` 兼容壳已删除。 -> 新功能与新治理都应直接落到 unified chat / `chat_*` 体系,不要重新引回旧入口。 +> 新功能与新治理都应直接落到 `agent_runtime_*` 与现役 `agent/chat` 体系,不要重新引回旧入口。 > `ProviderPoolService::select_credential_with_fallback_legacy` 也已删除,凭证选择统一走现役 `select_credential_with_fallback`。 ### ProviderPoolService diff --git a/docs/content/01.introduction/2.installation.md b/docs/content/01.introduction/2.installation.md index 9c1a5b8c5..f64ab0c8f 100644 --- a/docs/content/01.introduction/2.installation.md +++ b/docs/content/01.introduction/2.installation.md @@ -20,6 +20,10 @@ navigation: [下载 Lime](https://github.com/aiclientproxy/lime/releases) +::alert{type="warning"} +Lime 当前仅提供 macOS 与 Windows 桌面端安装包,Linux 版本已暂停支持,不再发布 `.deb` 或 `AppImage`。 +:: + ### 安装包 | 平台 | 文件名 | 说明 | diff --git a/docs/content/06.development/3.building.md b/docs/content/06.development/3.building.md index 2642f0e62..4039f8201 100644 --- a/docs/content/06.development/3.building.md +++ b/docs/content/06.development/3.building.md @@ -11,6 +11,10 @@ navigation: ## 本地开发 +::alert{type="warning"} +Lime 桌面端当前仅支持 macOS 与 Windows,本页不再提供 Linux 打包与发布说明。 +:: + ### 环境准备 1. **安装 Node.js** @@ -32,7 +36,7 @@ npm install -g pnpm 3. **安装 Rust** ```bash -# macOS/Linux +# macOS curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh # Windows @@ -52,12 +56,6 @@ xcode-select --install - 安装 Visual Studio Build Tools - 安装 WebView2(开发模式必需;对外分发时默认推荐在线小包,离线或受限网络环境再提供离线大包) -**Linux:** - -```bash -sudo apt install libwebkit2gtk-4.1-dev build-essential curl wget file libssl-dev libayatana-appindicator3-dev librsvg2-dev -``` - ### 启动开发 ```bash @@ -121,7 +119,6 @@ pnpm tauri build --debug | ------- | --------------------------------------- | | macOS | `src-tauri/target/release/bundle/dmg/` | | Windows | `src-tauri/target/release/bundle/nsis/` | -| Linux | `src-tauri/target/release/bundle/deb/` | ### 跨平台构建 @@ -150,13 +147,6 @@ pnpm tauri build --target x86_64-pc-windows-msvc --config src-tauri/tauri.window > 建议默认对外分发在线小包;只有内网、离线或受限网络环境,再提供离线大包。 -#### Linux 构建 - -```bash -# 构建 deb 包 -pnpm tauri build --target x86_64-unknown-linux-gnu -``` - ## 版本管理 ### 更新版本号 @@ -208,7 +198,6 @@ git push origin v1.0.1 | macOS | arm64 | macos-latest | | macOS | x64 | macos-13 | | Windows | x64 | windows-2022 | -| Linux | x64 | ubuntu-latest | ## 调试 @@ -229,7 +218,6 @@ RUST_LOG=debug pnpm tauri dev | ------- | -------------------------------- | | macOS | `~/Library/Logs/Lime/` | | Windows | `%APPDATA%\Lime\logs\` | -| Linux | `~/.local/share/lime/logs/` | ## 常见问题 diff --git a/docs/images/067c7d64-e116-4a30-b533-748873166f37.png b/docs/images/067c7d64-e116-4a30-b533-748873166f37.png deleted file mode 100644 index 15e989ccd..000000000 Binary files a/docs/images/067c7d64-e116-4a30-b533-748873166f37.png and /dev/null differ diff --git a/docs/images/151b4355-821c-4bda-a731-c4367b6b8716.png b/docs/images/151b4355-821c-4bda-a731-c4367b6b8716.png deleted file mode 100644 index d38507ddc..000000000 Binary files a/docs/images/151b4355-821c-4bda-a731-c4367b6b8716.png and /dev/null differ diff --git a/docs/images/25eb018a-5be2-4f82-ba22-e68f39160cac.png b/docs/images/25eb018a-5be2-4f82-ba22-e68f39160cac.png deleted file mode 100644 index 8fe5b9470..000000000 Binary files a/docs/images/25eb018a-5be2-4f82-ba22-e68f39160cac.png and /dev/null differ diff --git a/docs/images/943663ed-b17c-4b32-a74c-c0243ffb3dea.png b/docs/images/943663ed-b17c-4b32-a74c-c0243ffb3dea.png deleted file mode 100644 index 71b933c8b..000000000 Binary files a/docs/images/943663ed-b17c-4b32-a74c-c0243ffb3dea.png and /dev/null differ diff --git a/docs/images/aee62eb5-3aeb-4454-b14d-24b1d5f9a0fe.png b/docs/images/aee62eb5-3aeb-4454-b14d-24b1d5f9a0fe.png deleted file mode 100644 index 98fb49fd3..000000000 Binary files a/docs/images/aee62eb5-3aeb-4454-b14d-24b1d5f9a0fe.png and /dev/null differ diff --git a/docs/images/c7d8236b-ea6c-4496-ada5-288cd0a01738.png b/docs/images/c7d8236b-ea6c-4496-ada5-288cd0a01738.png deleted file mode 100644 index 20f6d97e5..000000000 Binary files a/docs/images/c7d8236b-ea6c-4496-ada5-288cd0a01738.png and /dev/null differ diff --git a/docs/images/coso.jpg b/docs/images/coso.jpg new file mode 100644 index 000000000..da158adc3 Binary files /dev/null and b/docs/images/coso.jpg differ diff --git a/docs/images/ffc70018-aa5f-4738-883d-045614488608.png b/docs/images/ffc70018-aa5f-4738-883d-045614488608.png deleted file mode 100644 index 2b4113451..000000000 Binary files a/docs/images/ffc70018-aa5f-4738-883d-045614488608.png and /dev/null differ diff --git a/docs/images/lime-claw-readme.png b/docs/images/lime-claw-readme.png new file mode 100644 index 000000000..8104cf430 Binary files /dev/null and b/docs/images/lime-claw-readme.png differ diff --git a/docs/images/screenshot-20260319-114622.png b/docs/images/screenshot-20260319-114622.png new file mode 100644 index 000000000..dda929a22 Binary files /dev/null and b/docs/images/screenshot-20260319-114622.png differ diff --git a/docs/images/screenshot-20260319-114807.png b/docs/images/screenshot-20260319-114807.png new file mode 100644 index 000000000..55361c94f Binary files /dev/null and b/docs/images/screenshot-20260319-114807.png differ diff --git a/docs/images/screenshot-20260319-115011.png b/docs/images/screenshot-20260319-115011.png new file mode 100644 index 000000000..933f06454 Binary files /dev/null and b/docs/images/screenshot-20260319-115011.png differ diff --git a/eslint.config.js b/eslint.config.js index f0489b8ed..367856be6 100644 --- a/eslint.config.js +++ b/eslint.config.js @@ -47,7 +47,7 @@ const generalChatRestrictedPaths = [ name: "@/components/general-chat", importNames: ["useChat"], message: - "general-chat 的 useChat 属于旧路径,请优先使用 @/hooks/useUnifiedChat 或当前现役聊天入口。", + "general-chat 的 useChat 属于旧路径,请优先使用 useAgentChatUnified / useAsterAgentChat 或当前现役聊天入口。", }, { name: "@/components/general-chat", @@ -65,18 +65,18 @@ const generalChatRestrictedPaths = [ name: "@/components/general-chat/hooks", importNames: ["useChat"], message: - "general-chat/hooks/useChat 属于旧路径,请优先使用 @/hooks/useUnifiedChat 或当前现役聊天入口。", + "general-chat/hooks/useChat 属于旧路径,请优先使用 useAgentChatUnified / useAsterAgentChat 或当前现役聊天入口。", }, { name: "@/components/general-chat/hooks", importNames: ["useSession", "useStreaming"], message: - "general-chat/hooks 下的 useSession/useStreaming 属于兼容实现,请优先接入统一对话链路。", + "general-chat/hooks 下的 useSession/useStreaming 属于兼容实现,请优先接入 useAgentChatUnified / useAsterAgentChat / agent_runtime_*。", }, { name: "@/components/general-chat/hooks/useChat", message: - "general-chat/hooks/useChat 属于旧路径,请优先使用 @/hooks/useUnifiedChat 或当前现役聊天入口。", + "general-chat/hooks/useChat 属于旧路径,请优先使用 useAgentChatUnified / useAsterAgentChat 或当前现役聊天入口。", }, { name: "@/components/general-chat/GeneralChatPage", @@ -400,6 +400,11 @@ const generalChatRestrictedPaths = [ message: "agentRuntime 中这些旧命名 helper 只允许留在兼容层或兼容测试;业务层请改用 agent_runtime_* 对应 API 或 useAsterAgentChat。", }, + { + name: "@/hooks/useUnifiedChat", + message: + "useUnifiedChat 已删除;Agent 工作台请改用 useAgentChatUnified / useAsterAgentChat,旧 General/Creator 路径如需恢复,必须基于 agent_runtime_* 重新设计。", + }, { name: "@/lib/api/agent", importNames: [ @@ -895,7 +900,7 @@ const unifiedChatCommandSelectors = [ ].map((command) => ({ selector: `CallExpression[callee.name='safeInvoke'][arguments.0.value='${command}'], CallExpression[callee.name='invoke'][arguments.0.value='${command}']`, message: - "统一对话命令请统一通过 `src/lib/api/unified-chat.ts` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", + "`chat_*` 兼容命令已停用;不要继续直接拼接旧命令名,现役链路请统一收敛到 agent_runtime_*。", })); const unifiedMemoryCommandSelectors = [ @@ -1179,7 +1184,6 @@ export default [ "src/lib/api/memoryFeedback.ts", "src/lib/api/toolHooks.ts", "src/lib/api/terminal.ts", - "src/lib/api/unified-chat.ts", "src/lib/api/unifiedMemory.ts", "src/lib/api/serverRuntime.ts", "src/lib/api/logs.ts", diff --git a/extensions/lime-chrome/README.md b/extensions/lime-chrome/README.md index 1c1d1f497..d62bc3cf7 100644 --- a/extensions/lime-chrome/README.md +++ b/extensions/lime-chrome/README.md @@ -67,5 +67,5 @@ npm run bridge:e2e -- --server ws://127.0.0.1:8787 --key proxy_cast --profile de ## 兼容说明 - 扩展只负责浏览器侧采集与动作执行。 -- Agent 侧通过 `aster_agent_cmd` 与 `unified_chat_cmd` 注册的浏览器 MCP 兼容工具访问。 +- Agent 侧通过 `aster_agent_cmd` 暴露的现役浏览器 MCP 工具访问。 - 若你同时使用独立 Chrome Profile(Tauri `open_chrome_profile_window`),请在对应 Profile 内安装该扩展,并使用不同 `profileKey` 做隔离。 diff --git a/package.json b/package.json index 58fa63995..b26f15a65 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "lime", "private": true, - "version": "0.90.0", + "version": "0.91.0", "type": "module", "engines": { "node": ">=22.0.0" diff --git a/scripts/report-legacy-surfaces.mjs b/scripts/report-legacy-surfaces.mjs index 9179ec5be..58a2f0e52 100644 --- a/scripts/report-legacy-surfaces.mjs +++ b/scripts/report-legacy-surfaces.mjs @@ -41,17 +41,17 @@ const agentLegacyHelperSurfaceMonitors = ( const importSurfaceMonitors = [ { id: "general-chat-root-entry", - classification: "deprecated", - description: "旧 general-chat 根导出入口", + classification: "dead-candidate", + description: "已删除的旧 general-chat 根导出入口", targets: ["src/components/general-chat/index.ts"], allowedPaths: [], }, { id: "general-chat-page-entry", - classification: "deprecated", - description: "旧 general-chat 页面实现入口", + classification: "dead-candidate", + description: "已删除的旧 general-chat 页面实现入口", targets: ["src/components/general-chat/GeneralChatPage.tsx"], - allowedPaths: ["src/components/general-chat/index.ts"], + allowedPaths: [], }, { id: "general-chat-legacy-session-hook", @@ -62,22 +62,22 @@ const importSurfaceMonitors = [ }, { id: "general-chat-legacy-streaming-hook", - classification: "compat", - description: "旧 general-chat 流式兼容 Hook", + classification: "dead-candidate", + description: "已删除的旧 general-chat 流式兼容 Hook", targets: ["src/components/general-chat/hooks/useStreaming.ts"], - allowedPaths: ["src/components/general-chat/GeneralChatPage.tsx"], + allowedPaths: [], }, { id: "general-chat-compat-gateway", classification: "dead-candidate", description: "general-chat compat API 网关", targets: ["src/lib/api/generalChatCompat.ts"], - allowedPaths: ["src/components/general-chat/store/useGeneralChatStore.ts"], + allowedPaths: [], }, { id: "agent-compat-gateway", - classification: "deprecated", - description: "Agent / Aster compat API 网关", + classification: "dead-candidate", + description: "已删除的 Agent / Aster compat API 网关", targets: ["src/lib/api/agentCompat.ts"], allowedPaths: [], }, @@ -97,22 +97,22 @@ const importSurfaceMonitors = [ }, { id: "heartbeat-api-gateway", - classification: "deprecated", - description: "旧 heartbeat 前端 API 入口", + classification: "dead-candidate", + description: "已删除的旧 heartbeat 前端 API 入口", targets: ["src/lib/api/heartbeat.ts"], allowedPaths: [], }, { id: "heartbeat-settings-page-entry", - classification: "deprecated", - description: "旧 heartbeat 设置页入口", + classification: "dead-candidate", + description: "已删除的旧 heartbeat 设置页入口", targets: ["src/components/settings-v2/system/heartbeat/index.tsx"], allowedPaths: [], }, { id: "assistant-settings-page-entry", - classification: "deprecated", - description: "旧助理服务设置页入口", + classification: "dead-candidate", + description: "已删除的旧助理服务设置页入口", targets: ["src/components/settings-v2/agent/assistant/index.tsx"], allowedPaths: [], }, @@ -130,14 +130,56 @@ const importSurfaceMonitors = [ targets: ["src/stores/agentStore.ts"], allowedPaths: [], }, + { + id: "use-unified-chat-compat-hook", + classification: "dead-candidate", + description: "useUnifiedChat compat Hook 入口", + targets: ["src/hooks/useUnifiedChat.ts"], + allowedPaths: [], + }, + { + id: "unified-chat-compat-gateway", + classification: "dead-candidate", + description: "unified-chat compat API 网关", + targets: ["src/lib/api/unified-chat.ts"], + allowedPaths: [], + }, + { + id: "three-stage-workflow-hook-entry", + classification: "dead-candidate", + description: "旧 three-stage workflow React Hook 入口", + targets: ["src/hooks/useThreeStageWorkflow.ts"], + allowedPaths: [], + }, + { + id: "three-stage-workflow-manager-entry", + classification: "dead-candidate", + description: "旧 three-stage workflow 管理器入口", + targets: ["src/lib/workflow/threeStageWorkflow.ts"], + allowedPaths: [], + }, + { + id: "tool-hooks-api-gateway", + classification: "dead-candidate", + description: "旧 tool hooks 前端 API 网关", + targets: ["src/lib/api/toolHooks.ts"], + allowedPaths: [], + }, + { + id: "context-memory-legacy-api-gateway", + classification: "dead-candidate", + description: "旧 context memory 前端 API 网关", + targets: ["src/lib/api/contextMemory.ts"], + allowedPaths: [], + }, ]; const commandSurfaceMonitors = [ ...agentLegacyCommandSurfaceMonitors, { id: "general-chat-compat-commands", - classification: "compat", - description: "general_chat compat 命令前端边界", + classification: "dead-candidate", + description: "已零引用的 general_chat compat 命令前端边界", commands: [ "general_chat_get_session", "general_chat_list_sessions", @@ -146,31 +188,64 @@ const commandSurfaceMonitors = [ "general_chat_rename_session", "general_chat_get_messages", ], - allowedPaths: ["src/lib/api/generalChatCompat.ts"], + allowedPaths: [], }, { id: "conversation-memory-legacy-commands", - classification: "compat", - description: "旧 conversation memory 命令前端边界", + classification: "dead-candidate", + description: "已零引用的旧 conversation memory 命令前端边界", commands: [ "get_conversation_memory_overview", "get_conversation_memory_stats", "request_conversation_memory_analysis", "cleanup_conversation_memory", ], - allowedPaths: ["src/lib/api/memoryRuntime.ts"], + allowedPaths: [], + }, + { + id: "context-memory-legacy-commands", + classification: "dead-candidate", + description: "旧 context memory 命令前端边界", + commands: [ + "save_memory_entry", + "get_session_memories", + "get_memory_context", + "record_error", + "should_avoid_operation", + "mark_error_resolved", + "get_memory_stats", + "cleanup_expired_memories", + ], + allowedPaths: [], + }, + { + id: "chat-compat-commands", + classification: "dead-candidate", + description: "chat_* compat 命令前端边界", + commands: [ + "chat_create_session", + "chat_list_sessions", + "chat_get_session", + "chat_delete_session", + "chat_rename_session", + "chat_get_messages", + "chat_send_message", + "chat_stop_generation", + "chat_configure_provider", + ], + allowedPaths: [], }, { id: "prompt-switch-legacy-command", - classification: "deprecated", - description: "旧 prompt 切换命令前端边界", + classification: "dead-candidate", + description: "已零引用的旧 prompt 切换命令前端边界", commands: ["switch_prompt"], allowedPaths: [], }, { id: "api-key-legacy-migration-commands", - classification: "deprecated", - description: "旧 API Key 迁移命令前端边界", + classification: "dead-candidate", + description: "已零引用的旧 API Key 迁移命令前端边界", commands: [ "get_legacy_api_key_credentials", "migrate_legacy_api_key_credentials", @@ -180,8 +255,8 @@ const commandSurfaceMonitors = [ }, { id: "heartbeat-legacy-commands", - classification: "deprecated", - description: "旧 heartbeat 命令前端边界", + classification: "dead-candidate", + description: "已零引用的旧 heartbeat 命令前端边界", commands: [ "get_heartbeat_config", "update_heartbeat_config", @@ -203,14 +278,29 @@ const commandSurfaceMonitors = [ ], allowedPaths: [], }, + { + id: "tool-hooks-legacy-commands", + classification: "dead-candidate", + description: "旧 tool hooks 命令前端边界", + commands: [ + "execute_hooks", + "add_hook_rule", + "remove_hook_rule", + "toggle_hook_rule", + "get_hook_rules", + "get_hook_execution_stats", + "clear_hook_execution_stats", + ], + allowedPaths: [], + }, ]; const frontendTextSurfaceMonitors = [ ...agentLegacyHelperSurfaceMonitors, { id: "frontend-assistant-settings-surfaces", - classification: "deprecated", - description: "前端助理服务设置页与配置面回流", + classification: "dead-candidate", + description: "已零引用的前端助理服务设置页与配置面回流", patterns: [ "SettingsTabs.Assistant", "settings.tab.assistant", @@ -225,8 +315,8 @@ const frontendTextSurfaceMonitors = [ }, { id: "stores-root-barrel-imports", - classification: "deprecated", - description: "从 @/stores 根 barrel 回流遗留 store 导入", + classification: "dead-candidate", + description: "已零引用的 @/stores 根 barrel 回流", patterns: ['from "@/stores"', "from '@/stores'"], allowedPaths: [], }, @@ -235,8 +325,8 @@ const frontendTextSurfaceMonitors = [ const rustTextSurfaceMonitors = [ { id: "rust-general-chat-dao", - classification: "deprecated", - description: "Rust 业务层 direct GeneralChatDao 依赖", + classification: "dead-candidate", + description: "已零引用的 Rust 业务层 direct GeneralChatDao 依赖", patterns: ["GeneralChatDao", "database::dao::general_chat"], allowedPaths: [], }, @@ -248,15 +338,14 @@ const rustTextSurfaceMonitors = [ allowedPaths: [ "src-tauri/crates/core/src/app_paths.rs", "src-tauri/crates/core/src/database/migration/general_chat_migration.rs", - "src-tauri/crates/core/src/database/pending_general_chat.rs", "src-tauri/crates/core/src/database/migration.rs", "src-tauri/crates/core/src/database/schema.rs", ], }, { id: "rust-legacy-general-helper-usage", - classification: "compat", - description: "Rust runtime pending general raw helper 扩散", + classification: "dead-candidate", + description: "Rust runtime pending general raw helper 回流", patterns: [ "load_pending_general_session_messages_raw", "load_pending_general_messages_raw", @@ -269,15 +358,27 @@ const rustTextSurfaceMonitors = [ "count_unmigrated_legacy_general_messages", "sum_unmigrated_legacy_general_message_chars", ], - allowedPaths: [ - "src-tauri/crates/core/src/database/pending_general_chat.rs", - "src-tauri/crates/core/src/database/mod.rs", + allowedPaths: [], + }, + { + id: "rust-pending-general-wrapper-usage", + classification: "dead-candidate", + description: "Rust 业务层 pending general 兼容 wrapper 回流", + patterns: [ + "load_pending_general_messages(", + "load_pending_general_session_messages(", + "count_pending_general_sessions(", + "count_pending_general_messages(", + "sum_pending_general_message_chars(", + "summarize_pending_general(", ], + includePathPrefixes: ["src-tauri/src", "src-tauri/crates"], + allowedPaths: [], }, { id: "rust-legacy-general-module-imports", - classification: "deprecated", - description: "Rust 外部模块直接引用 pending/legacy general 子模块", + classification: "dead-candidate", + description: "已零引用的 Rust 外部模块 direct pending/legacy general 子模块", patterns: [ "crate::database::legacy_general_chat::", "lime_core::database::legacy_general_chat::", @@ -296,14 +397,12 @@ const rustTextSurfaceMonitors = [ ], allowedPaths: [ "src-tauri/crates/core/src/database/migration/general_chat_migration.rs", - "src-tauri/crates/core/src/database/migration.rs", - "src-tauri/crates/core/src/database/mod.rs", ], }, { id: "rust-services-crate-general-chat-compat", - classification: "deprecated", - description: "services crate 内部继续依赖 general_chat 兼容壳", + classification: "dead-candidate", + description: "已零引用的 services crate general_chat 兼容壳回流", patterns: [ "use crate::general_chat::", "use crate::general_chat::{", @@ -314,22 +413,22 @@ const rustTextSurfaceMonitors = [ }, { id: "rust-cross-crate-general-chat-compat", - classification: "deprecated", - description: "跨 crate 引回 lime_services::general_chat 兼容壳", + classification: "dead-candidate", + description: "已零引用的跨 crate lime_services::general_chat 兼容壳回流", patterns: ["lime_services::general_chat::"], allowedPaths: [], }, { id: "rust-provider-pool-legacy-selector", - classification: "deprecated", - description: "provider pool legacy 凭证选择兼容方法", + classification: "dead-candidate", + description: "已零引用的 provider pool legacy 凭证选择兼容方法", patterns: ["select_credential_with_fallback_legacy"], allowedPaths: [], }, { id: "rust-memory-legacy-command-shells", - classification: "deprecated", - description: "旧 conversation memory Rust 命令壳回流", + classification: "dead-candidate", + description: "已零引用的旧 conversation memory Rust 命令壳回流", patterns: [ "get_conversation_memory_stats", "get_conversation_memory_overview", @@ -340,9 +439,138 @@ const rustTextSurfaceMonitors = [ allowedPaths: [], }, { - id: "rust-request-tool-policy-compat-service", + id: "rust-memory-profile-prompt-helper-leak", classification: "deprecated", - description: "request_tool_policy 旧服务壳回流", + description: "低层 build_memory_profile_prompt helper 泄漏到统一装配边界之外", + patterns: ["build_memory_profile_prompt("], + includePathPrefixes: ["src-tauri/src"], + allowedPaths: ["src-tauri/src/services/memory_profile_prompt_service.rs"], + }, + { + id: "rust-memory-sources-prompt-helper-leak", + classification: "deprecated", + description: "低层 build_memory_sources_prompt helper 泄漏到统一装配边界之外", + patterns: ["build_memory_sources_prompt("], + includePathPrefixes: ["src-tauri/src"], + allowedPaths: [ + "src-tauri/src/services/memory_profile_prompt_service.rs", + "src-tauri/src/services/memory_source_resolver_service.rs", + ], + }, + { + id: "rust-project-session-config-helper-leak", + classification: "dead-candidate", + description: "旧 create_session_config_with_project helper 回流", + patterns: ["create_session_config_with_project("], + allowedPaths: [], + }, + { + id: "rust-unified-chat-command-module-leak", + classification: "dead-candidate", + description: "Rust unified_chat compat 命令模块重新回到 commands 编译图", + patterns: ["pub mod unified_chat_cmd;"], + includePathPrefixes: ["src-tauri/src/commands"], + allowedPaths: [], + }, + { + id: "rust-tool-hooks-command-module-leak", + classification: "dead-candidate", + description: "Rust tool_hooks 旧命令模块重新回到 commands 编译图", + patterns: ["pub mod tool_hooks;"], + includePathPrefixes: ["src-tauri/src/commands"], + allowedPaths: [], + }, + { + id: "rust-tool-hooks-service-module-leak", + classification: "dead-candidate", + description: "services crate 旧 tool_hooks_service 模块重新回到编译图", + patterns: ["pub mod tool_hooks_service;"], + includePathPrefixes: ["src-tauri/crates/services/src"], + allowedPaths: [], + }, + { + id: "rust-three-stage-workflow-tool-leak", + classification: "dead-candidate", + description: "legacy three_stage_workflow 工具名重新回到 Lime Rust 编译图", + patterns: ['"three_stage_workflow"', "three_stage_workflow"], + includePathPrefixes: ["src-tauri/src", "src-tauri/crates"], + allowedPaths: [], + }, + { + id: "rust-skill-workflow-tauri-orchestration-leak", + classification: "dead-candidate", + description: "已零引用的 skill workflow 执行主链回流到 Tauri skills 模块", + patterns: [ + "SessionConfigBuilder::new(", + "convert_agent_event(", + "WriteArtifactEventEmitter::new(", + ".reply(", + ], + includePathPrefixes: ["src-tauri/src/skills"], + allowedPaths: [], + }, + { + id: "rust-skill-prompt-command-orchestration-leak", + classification: "dead-candidate", + description: "已零引用的 skill prompt 执行主链回流到 skill_exec_cmd 命令层", + patterns: [ + "SessionConfigBuilder::new(", + "convert_agent_event(", + "WriteArtifactEventEmitter::new(", + ".reply(", + ], + includePathPrefixes: ["src-tauri/src/commands/skill_exec_cmd.rs"], + allowedPaths: [], + }, + { + id: "rust-skill-runtime-command-bootstrap-leak", + classification: "dead-candidate", + description: "已零引用的 skill runtime 准备与 provider fallback 回流到 skill_exec_cmd 命令层", + patterns: [ + "ensure_browser_mcp_tools_registered(", + "ensure_social_image_tool_registered(", + "ensure_creation_task_tools_registered(", + "build_memory_profile_prompt(", + ".configure_provider_from_pool(", + "TauriExecutionCallback::new(", + ], + includePathPrefixes: ["src-tauri/src/commands/skill_exec_cmd.rs"], + allowedPaths: [], + }, + { + id: "rust-skill-catalog-command-leak", + classification: "dead-candidate", + description: "已零引用的 skill catalog 枚举与详情装配回流到 skill_exec_cmd 命令层", + patterns: [ + "get_skill_roots(", + "load_skills_from_directory(", + "find_skill_by_name(", + "invalid_skill_message(", + "load_skill_from_file(", + "parse_skill_frontmatter(", + "parse_allowed_tools(", + "parse_boolean(", + ], + includePathPrefixes: ["src-tauri/src/commands/skill_exec_cmd.rs"], + allowedPaths: [], + }, + { + id: "rust-skill-mode-branch-command-leak", + classification: "dead-candidate", + description: "已零引用的 skill execution_mode 分支回流到 skill_exec_cmd 命令层", + patterns: [ + "skill.execution_mode == \"workflow\"", + "!skill.workflow_steps.is_empty()", + "execute_skill_workflow(", + "execute_skill_prompt(", + ], + includePathPrefixes: ["src-tauri/src/commands/skill_exec_cmd.rs"], + allowedPaths: [], + }, + { + id: "rust-request-tool-policy-compat-service", + classification: "dead-candidate", + description: "已零引用的 request_tool_policy 旧服务壳回流", patterns: [ "crate::services::request_tool_policy_prompt_service::", "lime_lib::services::request_tool_policy_prompt_service::", @@ -352,8 +580,8 @@ const rustTextSurfaceMonitors = [ }, { id: "rust-service-agent-table-query-leak", - classification: "deprecated", - description: "Tauri service 层 direct agent_sessions/agent_messages 查询回流", + classification: "dead-candidate", + description: "已零引用的 Tauri service 层 direct agent_sessions/agent_messages 查询回流", patterns: [ "FROM agent_sessions s", "FROM agent_messages m", @@ -364,33 +592,90 @@ const rustTextSurfaceMonitors = [ }, { id: "rust-service-model-usage-table-query-leak", - classification: "deprecated", - description: "Tauri service 层 direct model_usage_stats 查询回流", + classification: "dead-candidate", + description: "已零引用的 Tauri service 层 direct model_usage_stats 查询回流", patterns: ["FROM model_usage_stats", "SELECT COUNT(*) FROM model_usage_stats"], includePathPrefixes: ["src-tauri/src/services"], allowedPaths: [], }, + { + id: "rust-dev-bridge-unified-memory-sql-leak", + classification: "dead-candidate", + description: "已零引用的 DevBridge unified_memory SQL 绕路回流", + patterns: [ + "FROM unified_memory", + "INSERT INTO unified_memory", + "UPDATE unified_memory", + "DELETE FROM unified_memory", + ], + includePathPrefixes: ["src-tauri/src/dev_bridge/dispatcher/memory.rs"], + allowedPaths: [], + }, + { + id: "rust-agent-runtime-legacy-queue-table-leak", + classification: "deprecated", + description: "legacy runtime queue 表名从数据库迁移边界向外扩散", + patterns: ["agent_runtime_queued_turns"], + includePathPrefixes: ["src-tauri/src", "src-tauri/crates"], + allowedPaths: ["src-tauri/crates/core/src/database/agent_runtime_queue_repository.rs"], + }, + { + id: "rust-agent-runtime-legacy-queue-migration-leak", + classification: "dead-candidate", + description: "已零引用的 legacy runtime queue 启动迁移 helper 回流到其他模块", + patterns: ["migrate_legacy_runtime_queue_to_aster_store("], + includePathPrefixes: ["src-tauri/src", "src-tauri/crates"], + allowedPaths: [], + }, + { + id: "rust-agent-session-legacy-todo-state-leak", + classification: "dead-candidate", + description: "已零引用的 Lime 业务层 direct legacy TodoState 读取回流", + patterns: ["TodoState::from_extension_data(", "TodoState::new("], + includePathPrefixes: ["src-tauri/src", "src-tauri/crates"], + allowedPaths: [], + }, + { + id: "rust-agent-session-structured-todo-helper-bypass", + classification: "dead-candidate", + description: "已零引用的 Lime 业务层绕过 unified todo helper 直接读取 TodoListState", + patterns: ["TodoListState::from_extension_data(", "TodoListState::from_markdown("], + includePathPrefixes: ["src-tauri/src", "src-tauri/crates"], + allowedPaths: [], + }, + { + id: "rust-services-default-workspace-query-leak", + classification: "dead-candidate", + description: "已零引用的 services crate direct 默认 workspace root 查询回流", + patterns: ["SELECT root_path FROM workspaces WHERE is_default = 1 LIMIT 1"], + includePathPrefixes: ["src-tauri/crates/services/src"], + allowedPaths: [], + }, { id: "rust-agent-session-direct-record-access", classification: "deprecated", description: "Rust 上层模块 direct Agent session 记录与消息读写回流", patterns: [ "AgentDao::get_session(", + "AgentDao::get_session_with_messages(", "AgentDao::list_sessions(", + "AgentDao::list_session_overviews(", "AgentDao::get_message_count(", "AgentDao::get_messages(", + "AgentDao::get_session_overview(", "AgentDao::session_exists(", "AgentDao::update_title(", "AgentDao::update_session_time(", + "AgentDao::rename_session(", "AgentDao::update_working_dir(", "AgentDao::update_execution_strategy(", ], - allowedPaths: ["src-tauri/crates/agent/src/session_store.rs"], + allowedPaths: ["src-tauri/crates/core/src/database/agent_session_repository.rs"], }, { id: "rust-agent-session-direct-delete", - classification: "deprecated", - description: "Rust 业务层 direct AgentDao::delete_session 回流", + classification: "dead-candidate", + description: "已零引用的 Rust 业务层 direct AgentDao::delete_session 回流", patterns: ["AgentDao::delete_session("], allowedPaths: [], }, @@ -399,12 +684,191 @@ const rustTextSurfaceMonitors = [ classification: "deprecated", description: "Rust 业务层 direct AgentDao::create_session 回流", patterns: ["AgentDao::create_session("], - allowedPaths: ["src-tauri/crates/agent/src/session_store.rs"], + allowedPaths: ["src-tauri/crates/core/src/database/agent_session_repository.rs"], + }, + { + id: "rust-agent-dao-row-type-leak", + classification: "dead-candidate", + description: "agent crate 重新泄漏 core DAO row 类型", + patterns: ["AgentSessionOverviewRow"], + includePathPrefixes: ["src-tauri/crates/agent/src"], + allowedPaths: [], + }, + { + id: "rust-agent-workspace-query-leak", + classification: "dead-candidate", + description: "agent runtime direct workspaces 绑定查询回流", + patterns: ["SELECT id FROM workspaces WHERE root_path = ? LIMIT 1"], + includePathPrefixes: ["src-tauri/crates/agent/src"], + allowedPaths: [], + }, + { + id: "rust-agent-session-record-repository-module-leak", + classification: "dead-candidate", + description: "lime-agent 本地 session_record_repository 壳重新回到编译图", + patterns: ["mod session_record_repository;"], + includePathPrefixes: ["src-tauri/crates/agent/src/lib.rs"], + allowedPaths: [], + }, + { + id: "rust-agent-session-compat-surface", + classification: "dead-candidate", + description: "agent crate 旧 compat session API 回流", + patterns: [ + "CompatSessionInfo", + "list_compat_sessions_sync(", + "get_compat_session_sync(", + ], + includePathPrefixes: ["src-tauri/crates/agent/src"], + allowedPaths: [], + }, + { + id: "rust-agent-session-store-public-module-leak", + classification: "dead-candidate", + description: "lime-agent 重新对 crate 外暴露 session_store 模块", + patterns: ["pub mod session_store;"], + includePathPrefixes: ["src-tauri/crates/agent/src/lib.rs"], + allowedPaths: [], + }, + { + id: "rust-agent-session-store-direct-module-usage", + classification: "dead-candidate", + description: "上层模块重新 direct 依赖 lime_agent::session_store 模块路径", + patterns: ["lime_agent::session_store::"], + includePathPrefixes: ["src-tauri/src"], + allowedPaths: [], + }, + { + id: "rust-agent-session-record-create-api-leak", + classification: "dead-candidate", + description: "lime-agent crate 根重新暴露内部 session record 创建 API", + patterns: ["create_session_record_sync,", "CreateSessionRecordInput,"], + includePathPrefixes: ["src-tauri/crates/agent/src/lib.rs"], + allowedPaths: [], + }, + { + id: "rust-agent-integration-public-module-leak", + classification: "dead-candidate", + description: "agent 集成模块重新对外暴露 integration 模块路径", + patterns: ["pub mod integration;"], + includePathPrefixes: ["src-tauri/src/agent/mod.rs"], + allowedPaths: [], + }, + { + id: "rust-agent-integration-module-compiled-leak", + classification: "dead-candidate", + description: "已退出编译图的 agent integration 模块重新回到 agent 根模块", + patterns: ["mod integration;"], + includePathPrefixes: ["src-tauri/src/agent/mod.rs"], + allowedPaths: [], + }, + { + id: "rust-agent-integration-direct-module-usage", + classification: "dead-candidate", + description: "应用层重新 direct 依赖 crate::agent::integration 模块路径", + patterns: ["crate::agent::integration::"], + includePathPrefixes: ["src-tauri/src"], + allowedPaths: [], + }, + { + id: "rust-agent-subagent-direct-module-usage", + classification: "dead-candidate", + description: "应用层重新 direct 依赖 crate::agent::subagent_scheduler 模块路径", + patterns: ["crate::agent::subagent_scheduler::"], + includePathPrefixes: ["src-tauri/src"], + allowedPaths: [], + }, + { + id: "rust-aster-runtime-snapshot-helper-leak", + classification: "dead-candidate", + description: "已零引用的 Lime 业务层 direct Aster runtime snapshot helper 回流", + patterns: ["load_session_runtime_snapshot("], + includePathPrefixes: ["src-tauri/src", "src-tauri/crates"], + allowedPaths: [], + }, + { + id: "rust-aster-runtime-store-leak", + classification: "dead-candidate", + description: "已零引用的 Lime 业务层 direct Aster shared runtime store 回流", + patterns: [ + "shared_thread_runtime_store(", + "initialize_shared_thread_runtime_store(", + "initialize_shared_sqlite_thread_runtime_store(", + ], + includePathPrefixes: ["src-tauri/src", "src-tauri/crates"], + allowedPaths: [], + }, + { + id: "rust-aster-runtime-store-public-require-api-leak", + classification: "dead-candidate", + description: "Aster runtime support 重新对 crate 外暴露 require_aster_thread_runtime_store", + patterns: ["pub fn require_aster_thread_runtime_store("], + includePathPrefixes: ["src-tauri/crates/agent/src/aster_runtime_support.rs"], + allowedPaths: [], + }, + { + id: "rust-aster-runtime-init-return-store-api-leak", + classification: "dead-candidate", + description: "Aster runtime 启动初始化 API 重新向 crate 外返回 runtime store", + patterns: [ + "pub fn initialize_aster_thread_runtime_store() -> Result, String>", + ], + includePathPrefixes: ["src-tauri/crates/agent/src/aster_runtime_support.rs"], + allowedPaths: [], + }, + { + id: "rust-aster-runtime-public-legacy-init-helper-leak", + classification: "dead-candidate", + description: "Aster runtime support 重新对 crate 外暴露旧 initialize_aster_thread_runtime_store helper", + patterns: ["pub fn initialize_aster_thread_runtime_store("], + includePathPrefixes: ["src-tauri/crates/agent/src/aster_runtime_support.rs"], + allowedPaths: [], + }, + { + id: "rust-aster-runtime-snapshot-root-export-leak", + classification: "dead-candidate", + description: "lime-agent crate 根重新暴露 load_aster_runtime_snapshot helper", + patterns: ["load_aster_runtime_snapshot"], + includePathPrefixes: ["src-tauri/crates/agent/src/lib.rs"], + allowedPaths: [], + }, + { + id: "rust-aster-runtime-queue-service-leak", + classification: "dead-candidate", + description: "已零引用的 Lime 业务层 direct Aster shared runtime queue service 回流", + patterns: [], + regexPatterns: [String.raw`(? - monitor.patterns.some((pattern) => sourceCode.includes(pattern)); + monitor.patterns.some((pattern) => sourceCode.includes(pattern)) || + (monitor.regexPatterns ?? []).some((pattern) => + new RegExp(pattern, "m").test(sourceCode), + ); const references = filteredRuntimeSources .filter((file) => matchesPattern(file.sourceCode)) @@ -1019,15 +1504,43 @@ function evaluateTextCountMonitor(monitor, runtimeSources, testSources) { }; } +function getImportStatus(result) { + return result.violations.length > 0 + ? "违规" + : result.references.length === 0 && result.existingTargets.length === 0 + ? "已删除" + : result.references.length === 0 + ? "零引用" + : "受控"; +} + +function getCommandStatus(result) { + const flattenedReferences = [...result.referencesByCommand.values()].flat(); + const uniqueReferences = [...new Set(flattenedReferences)].sort(); + return result.violations.length > 0 + ? "违规" + : uniqueReferences.length === 0 + ? "零引用" + : "受控"; +} + +function getTextStatus(result) { + return result.violations.length > 0 + ? "违规" + : result.references.length === 0 + ? "零引用" + : "受控"; +} + +function isStatusClassificationDrift(status, classification) { + return ( + (status === "已删除" || status === "零引用") && + classification !== "dead-candidate" + ); +} + function printImportReport(result) { - const status = - result.violations.length > 0 - ? "违规" - : result.references.length === 0 && result.existingTargets.length === 0 - ? "已删除" - : result.references.length === 0 - ? "零引用" - : "受控"; + const status = getImportStatus(result); console.log( `- [${status}] ${result.id} (${result.classification}):${result.description}`, @@ -1046,14 +1559,7 @@ function printImportReport(result) { } function printCommandReport(result) { - const flattenedReferences = [...result.referencesByCommand.values()].flat(); - const uniqueReferences = [...new Set(flattenedReferences)].sort(); - const status = - result.violations.length > 0 - ? "违规" - : uniqueReferences.length === 0 - ? "零引用" - : "受控"; + const status = getCommandStatus(result); console.log( `- [${status}] ${result.id} (${result.classification}):${result.description}`, @@ -1074,17 +1580,16 @@ function printCommandReport(result) { } function printTextReport(result) { - const status = - result.violations.length > 0 - ? "违规" - : result.references.length === 0 - ? "零引用" - : "受控"; + const status = getTextStatus(result); console.log( `- [${status}] ${result.id} (${result.classification}):${result.description}`, ); - console.log(` 关键字:${result.patterns.join(", ")}`); + const keywords = [ + ...result.patterns, + ...(result.regexPatterns ?? []).map((pattern) => `regex:${pattern}`), + ]; + console.log(` 关键字:${keywords.join(", ")}`); console.log(` 允许引用:${result.allowedPaths.join(", ") || "无"}`); console.log(` 实际引用:\n${formatPaths(result.references)}`); console.log(` 测试引用:\n${formatPaths(result.testReferences)}`); @@ -1169,6 +1674,65 @@ const zeroReferenceCandidates = importResults result.references.length === 0 && result.existingTargets.length > 0, ) .map((result) => `${result.id} (${result.description})`); +const classificationDriftCandidates = [ + ...importResults + .filter((result) => + isStatusClassificationDrift( + getImportStatus(result), + result.classification, + ), + ) + .map( + (result) => + `${result.id} -> ${result.classification} / ${getImportStatus(result)}`, + ), + ...commandResults + .filter((result) => + isStatusClassificationDrift( + getCommandStatus(result), + result.classification, + ), + ) + .map( + (result) => + `${result.id} -> ${result.classification} / ${getCommandStatus(result)}`, + ), + ...frontendTextResults + .filter((result) => + isStatusClassificationDrift( + getTextStatus(result), + result.classification, + ), + ) + .map( + (result) => + `${result.id} -> ${result.classification} / ${getTextStatus(result)}`, + ), + ...rustTextResults + .filter((result) => + isStatusClassificationDrift( + getTextStatus(result), + result.classification, + ), + ) + .map( + (result) => + `${result.id} -> ${result.classification} / ${getTextStatus(result)}`, + ), + ...rustTextCountResults + .filter((result) => + isStatusClassificationDrift( + result.runtimeMatches.length === 0 ? "零引用" : "受控", + result.classification, + ), + ) + .map( + (result) => + `${result.id} -> ${result.classification} / ${ + result.runtimeMatches.length === 0 ? "零引用" : "受控" + }`, + ), +]; const violations = [ ...importResults.flatMap((result) => result.violations.map((item) => `${result.id} -> ${item}`), @@ -1225,6 +1789,10 @@ console.log(`- 零引用候选:${zeroReferenceCandidates.length}`); for (const candidate of zeroReferenceCandidates) { console.log(` - ${candidate}`); } +console.log(`- 分类漂移候选:${classificationDriftCandidates.length}`); +for (const candidate of classificationDriftCandidates) { + console.log(` - ${candidate}`); +} console.log(`- 边界违规:${violations.length}`); for (const violation of violations) { console.log(` - ${violation}`); diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 398633789..670910617 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -369,7 +369,7 @@ checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" [[package]] name = "aster-core" -version = "0.18.0" +version = "0.19.0" dependencies = [ "ahash", "anyhow", @@ -461,7 +461,7 @@ dependencies = [ [[package]] name = "aster-models" -version = "0.18.0" +version = "0.19.0" dependencies = [ "serde", "serde_json", @@ -1313,7 +1313,7 @@ dependencies = [ "bitflags 2.11.0", "cexpr", "clang-sys", - "itertools 0.13.0", + "itertools 0.12.1", "proc-macro2", "quote", "regex", @@ -2399,7 +2399,7 @@ dependencies = [ "dtoa-short", "itoa", "matches", - "phf 0.10.1", + "phf 0.8.0", "proc-macro2", "quote", "smallvec", @@ -2415,7 +2415,7 @@ dependencies = [ "cssparser-macros", "dtoa-short", "itoa", - "phf 0.11.3", + "phf 0.8.0", "smallvec", ] @@ -4336,7 +4336,7 @@ dependencies = [ "js-sys", "log", "wasm-bindgen", - "windows-core 0.57.0", + "windows-core 0.56.0", ] [[package]] @@ -4716,15 +4716,6 @@ dependencies = [ "either", ] -[[package]] -name = "itertools" -version = "0.13.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "413ee7dfc52ee1a4949ceeb7dbc8a33f2d6c088194d9f922fb8318faf1f01186" -dependencies = [ - "either", -] - [[package]] name = "itertools" version = "0.14.0" @@ -5071,7 +5062,7 @@ dependencies = [ [[package]] name = "lime" -version = "0.90.0" +version = "0.91.0" dependencies = [ "anyhow", "arboard", @@ -5174,7 +5165,7 @@ dependencies = [ [[package]] name = "lime-agent" -version = "0.90.0" +version = "0.91.0" dependencies = [ "aster-core", "async-trait", @@ -5185,6 +5176,7 @@ dependencies = [ "lime-mcp", "lime-providers", "lime-services", + "lime-skills", "regex", "rmcp", "rusqlite", @@ -5200,7 +5192,7 @@ dependencies = [ [[package]] name = "lime-browser-runtime" -version = "0.90.0" +version = "0.91.0" dependencies = [ "chrono", "futures", @@ -5217,7 +5209,7 @@ dependencies = [ [[package]] name = "lime-config" -version = "0.90.0" +version = "0.91.0" dependencies = [ "async-trait", "lime-core", @@ -5233,7 +5225,7 @@ dependencies = [ [[package]] name = "lime-core" -version = "0.90.0" +version = "0.91.0" dependencies = [ "aster-models", "async-trait", @@ -5273,7 +5265,7 @@ dependencies = [ [[package]] name = "lime-credential" -version = "0.90.0" +version = "0.91.0" dependencies = [ "axum 0.7.9", "base64 0.22.1", @@ -5308,7 +5300,7 @@ dependencies = [ [[package]] name = "lime-gateway" -version = "0.90.0" +version = "0.91.0" dependencies = [ "axum 0.7.9", "chrono", @@ -5329,7 +5321,7 @@ dependencies = [ [[package]] name = "lime-infra" -version = "0.90.0" +version = "0.91.0" dependencies = [ "chrono", "dashmap 5.5.3", @@ -5349,7 +5341,7 @@ dependencies = [ [[package]] name = "lime-mcp" -version = "0.90.0" +version = "0.91.0" dependencies = [ "async-trait", "dirs 5.0.1", @@ -5381,7 +5373,7 @@ dependencies = [ [[package]] name = "lime-processor" -version = "0.90.0" +version = "0.91.0" dependencies = [ "async-trait", "lime-core", @@ -5400,7 +5392,7 @@ dependencies = [ [[package]] name = "lime-providers" -version = "0.90.0" +version = "0.91.0" dependencies = [ "anyhow", "async-stream", @@ -5454,7 +5446,7 @@ dependencies = [ [[package]] name = "lime-server" -version = "0.90.0" +version = "0.91.0" dependencies = [ "aster-core", "async-stream", @@ -5499,7 +5491,7 @@ dependencies = [ [[package]] name = "lime-server-utils" -version = "0.90.0" +version = "0.91.0" dependencies = [ "axum 0.7.9", "futures", @@ -5514,7 +5506,7 @@ dependencies = [ [[package]] name = "lime-services" -version = "0.90.0" +version = "0.91.0" dependencies = [ "anyhow", "aster-core", @@ -5556,7 +5548,7 @@ dependencies = [ [[package]] name = "lime-skills" -version = "0.90.0" +version = "0.91.0" dependencies = [ "async-trait", "dirs 5.0.1", @@ -5574,7 +5566,7 @@ dependencies = [ [[package]] name = "lime-terminal" -version = "0.90.0" +version = "0.91.0" dependencies = [ "async-trait", "base64 0.22.1", @@ -5601,7 +5593,7 @@ dependencies = [ [[package]] name = "lime-websocket" -version = "0.90.0" +version = "0.91.0" dependencies = [ "axum 0.7.9", "chrono", @@ -6274,7 +6266,7 @@ version = "0.7.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ff32365de1b6743cb203b710788263c44a03de03802daf96092f2da4fe6ba4d7" dependencies = [ - "proc-macro-crate 2.0.2", + "proc-macro-crate 1.3.1", "proc-macro2", "quote", "syn 2.0.117", @@ -7001,7 +6993,9 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3dfb61232e34fcb633f43d12c58f83c1df82962dcdfa565a4e866ffc17dafe12" dependencies = [ + "phf_macros 0.8.0", "phf_shared 0.8.0", + "proc-macro-hack", ] [[package]] @@ -7010,9 +7004,7 @@ version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fabbf1ead8a5bcbc20f5f8b939ee3f5b0f6f281b6ad3468b84656b658b455259" dependencies = [ - "phf_macros 0.10.0", "phf_shared 0.10.0", - "proc-macro-hack", ] [[package]] @@ -7116,12 +7108,12 @@ dependencies = [ [[package]] name = "phf_macros" -version = "0.10.0" +version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "58fdf3184dd560f160dd73922bea2d5cd6e8f064bf4b13110abd81b03697b4e0" +checksum = "7f6fde18ff429ffc8fe78e2bf7f8b7a5a5a6e2a8b58bc5a9ac69198bbda9189c" dependencies = [ - "phf_generator 0.10.0", - "phf_shared 0.10.0", + "phf_generator 0.8.0", + "phf_shared 0.8.0", "proc-macro-hack", "proc-macro2", "quote", @@ -7533,7 +7525,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8a56d757972c98b346a9b766e3f02746cde6dd1cd1d1d563472929fdd74bec4d" dependencies = [ "anyhow", - "itertools 0.13.0", + "itertools 0.12.1", "proc-macro2", "quote", "syn 2.0.117", @@ -9000,7 +8992,7 @@ version = "3.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8b1fdf65dd6331831494dd616b30351c38e96e45921a27745cf98490458b90bb" dependencies = [ - "dirs 6.0.0", + "dirs 4.0.0", ] [[package]] diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 41853deea..44a8b6844 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -3,7 +3,7 @@ members = ["crates/*"] resolver = "2" [workspace.package] -version = "0.90.0" +version = "0.91.0" edition = "2021" authors = ["coso"] repository = "https://github.com/aiclientproxy/lime" @@ -127,8 +127,8 @@ enigo = "0.3" # 如需联调本地 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.18.0" } -aster-models = { git = "https://github.com/astercloud/aster-rust", tag = "v0.18.0" } +aster = { package = "aster-core", git = "https://github.com/astercloud/aster-rust", tag = "v0.19.0" } +aster-models = { git = "https://github.com/astercloud/aster-rust", tag = "v0.19.0" } # MCP (Model Context Protocol) rmcp = { version = "0.12.0", features = ["client", "transport-io", "transport-child-process"] } @@ -191,7 +191,7 @@ version = "2.4" [package] name = "lime" -version = "0.90.0" +version = "0.91.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" diff --git a/src-tauri/crates/agent/Cargo.toml b/src-tauri/crates/agent/Cargo.toml index 1e3238e39..c339f37ea 100644 --- a/src-tauri/crates/agent/Cargo.toml +++ b/src-tauri/crates/agent/Cargo.toml @@ -10,6 +10,7 @@ lime-core.workspace = true lime-mcp.workspace = true lime-services.workspace = true lime-providers.workspace = true +lime-skills.workspace = true aster.workspace = true rmcp.workspace = true serde.workspace = true diff --git a/src-tauri/crates/agent/src/ask_bridge.rs b/src-tauri/crates/agent/src/ask_bridge.rs index 264e0b62c..b2100ef59 100644 --- a/src-tauri/crates/agent/src/ask_bridge.rs +++ b/src-tauri/crates/agent/src/ask_bridge.rs @@ -4,6 +4,8 @@ //! 通过 elicitation 事件把问题发送到前端并等待用户输入。 use aster::action_required_manager::ActionRequiredManager; +use aster::conversation::message::ActionRequiredScope; +use aster::session_context::{current_action_scope, current_session_id}; use aster::tools::AskCallback; use serde_json::{json, Value}; use std::time::Duration; @@ -15,9 +17,11 @@ pub fn create_ask_callback() -> AskCallback { std::sync::Arc::new(|question: String, options: Option>| { Box::pin(async move { let requested_schema = build_requested_schema(&question, options.as_deref()); + let scope = resolve_action_scope(); match ActionRequiredManager::global() - .request_and_wait( + .request_and_wait_scoped( + scope, question.clone(), requested_schema, Duration::from_secs(DEFAULT_ASK_TIMEOUT_SECS), @@ -38,6 +42,17 @@ pub fn create_ask_callback() -> AskCallback { }) } +fn resolve_action_scope() -> ActionRequiredScope { + current_action_scope().unwrap_or_else(|| { + let session_id = current_session_id(); + ActionRequiredScope { + session_id: session_id.clone(), + thread_id: session_id, + turn_id: None, + } + }) +} + /// 构建 elicitation 的请求 schema fn build_requested_schema(question: &str, options: Option<&[String]>) -> Value { if let Some(options) = options { @@ -107,3 +122,39 @@ pub fn extract_response(user_data: &Value) -> Option { .filter(|s| !s.is_empty()), } } + +#[cfg(test)] +mod tests { + use super::*; + use aster::session_context::{with_action_scope, with_session_id}; + + #[tokio::test] + async fn resolve_action_scope_prefers_runtime_scope() { + let scope = ActionRequiredScope { + session_id: Some("session-1".to_string()), + thread_id: Some("thread-1".to_string()), + turn_id: Some("turn-1".to_string()), + }; + + let resolved = with_action_scope(scope.clone(), async { resolve_action_scope() }).await; + + assert_eq!(resolved, scope); + } + + #[tokio::test] + async fn resolve_action_scope_falls_back_to_session_id() { + let resolved = with_session_id(Some("session-2".to_string()), async { + resolve_action_scope() + }) + .await; + + assert_eq!( + resolved, + ActionRequiredScope { + session_id: Some("session-2".to_string()), + thread_id: Some("session-2".to_string()), + turn_id: None, + } + ); + } +} diff --git a/src-tauri/crates/agent/src/aster_runtime_support.rs b/src-tauri/crates/agent/src/aster_runtime_support.rs new file mode 100644 index 000000000..291150cf2 --- /dev/null +++ b/src-tauri/crates/agent/src/aster_runtime_support.rs @@ -0,0 +1,326 @@ +//! Aster runtime 支持模块 +//! +//! 收口 Lime 对 Aster thread runtime store 的访问边界, +//! 避免业务层散落依赖上游 free function。 + +use crate::aster_state::QueuedTurnTask; +use crate::queued_turn::QueuedTurnSnapshot; +use aster::session::{ + initialize_shared_session_runtime_with_root, load_shared_session_runtime_snapshot, + require_shared_session_runtime_store, QueuedTurnRuntime, SessionRuntimeSnapshot, + ThreadRuntimeStore, +}; +use lime_core::app_paths; +use lime_core::database::agent_runtime_queue_repository::{self, LegacyRuntimeQueuedTurn}; +use lime_core::database::{lock_db, DbConnection}; +use lime_services::aster_session_store::LimeSessionStore; +use serde_json::Value; +use std::collections::{HashMap, HashSet}; +use std::future::Future; +use std::path::PathBuf; +use std::sync::{Arc, OnceLock}; + +const QUEUED_TURN_EVENT_NAME_METADATA_KEY: &str = "event_name"; +const DEFAULT_QUEUE_EVENT_NAME: &str = "agent_stream"; +static ASTER_RUNTIME_ROOT: OnceLock> = OnceLock::new(); + +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +struct LegacyRuntimeQueueMigrationReport { + pub session_count: usize, + pub migrated_turn_count: usize, + pub skipped_existing_turn_count: usize, + pub invalid_turn_count: usize, +} + +pub(crate) fn ensure_aster_runtime_dirs() -> Result { + ASTER_RUNTIME_ROOT + .get_or_init(initialize_aster_runtime_dirs) + .clone() +} + +pub(crate) fn require_aster_runtime_dirs() -> Result { + match ASTER_RUNTIME_ROOT.get() { + Some(result) => result.clone(), + None => Err( + "Aster 运行时尚未初始化;应在应用启动期先调用 ensure_aster_runtime_dirs()".to_string(), + ), + } +} + +#[cfg(test)] +pub(crate) fn ensure_aster_runtime_dirs_with_root(root: PathBuf) -> Result { + ASTER_RUNTIME_ROOT + .get_or_init(|| initialize_aster_runtime_dirs_with_root(root)) + .clone() +} + +/// 启动期显式初始化 Aster runtime 目录、共享 runtime store 与全局 session store。 +pub fn initialize_aster_runtime(db: DbConnection) -> Result<(), String> { + let runtime_root = ensure_aster_runtime_dirs()?; + let session_store = Arc::new(LimeSessionStore::new(db.clone())); + let migration_db = db.clone(); + + block_on_aster_runtime_init(async move { + initialize_shared_session_runtime_with_root(runtime_root, Some(session_store)) + .await + .map_err(|error| format!("初始化 Aster runtime 失败: {error}"))?; + migrate_legacy_runtime_queue_to_aster_store(&migration_db).await?; + Ok(()) + }) +} + +fn block_on_aster_runtime_init(future: F) -> Result<(), String> +where + F: Future>, +{ + if let Ok(handle) = tokio::runtime::Handle::try_current() { + return handle.block_on(future); + } + + #[cfg(target_os = "windows")] + tracing::info!("[AsterRuntime] Windows 平台 - 创建 Tokio Runtime (IOCP)"); + + #[cfg(target_os = "macos")] + tracing::info!("[AsterRuntime] macOS 平台 - 创建 Tokio Runtime (kqueue)"); + + #[cfg(target_os = "linux")] + tracing::info!("[AsterRuntime] Linux 平台 - 创建 Tokio Runtime (epoll)"); + + let runtime = tokio::runtime::Builder::new_multi_thread() + .worker_threads(2) + .thread_name("lime-runtime") + .enable_io() + .enable_time() + .build() + .map_err(|error| format!("创建 Tokio Runtime 失败: {error}"))?; + runtime.block_on(future) +} + +fn initialize_aster_runtime_dirs() -> Result { + initialize_aster_runtime_dirs_with_root(app_paths::resolve_aster_dir()?) +} + +fn initialize_aster_runtime_dirs_with_root(root: PathBuf) -> Result { + let runtime_root = root.clone(); + block_on_aster_runtime_init(async move { + initialize_shared_session_runtime_with_root(runtime_root, None) + .await + .map_err(|error| format!("初始化 Aster runtime 失败: {error}")) + })?; + Ok(root) +} + +/// 获取 Lime 当前统一使用的 Aster runtime store。 +pub(crate) fn require_aster_runtime_store() -> Result, String> { + ensure_aster_runtime_dirs()?; + require_shared_session_runtime_store().map_err(|error| error.to_string()) +} + +/// 读取会话 runtime snapshot。 +pub(crate) async fn load_aster_runtime_snapshot( + session_id: &str, +) -> Result { + ensure_aster_runtime_dirs()?; + load_shared_session_runtime_snapshot(session_id) + .await + .map_err(|error| format!("读取 runtime snapshot 失败: {error}")) +} + +pub(crate) async fn list_aster_runtime_queued_turns( + session_id: &str, +) -> Result, String> { + let store = require_aster_runtime_store()?; + store + .list_queued_turns(session_id) + .await + .map_err(|error| format!("读取 queued runtime turns 失败: {error}")) +} + +async fn list_aster_runtime_queued_turn_session_ids() -> Result, String> { + let store = require_aster_runtime_store()?; + store + .list_queued_turn_session_ids() + .await + .map_err(|error| format!("读取 queued runtime session ids 失败: {error}")) +} + +/// 启动恢复统一入口:只在这里完成当前 queued session 枚举。 +pub(crate) async fn prepare_aster_runtime_queue_resumption() -> Result, String> { + ensure_aster_runtime_dirs()?; + list_aster_runtime_queued_turn_session_ids().await +} + +async fn enqueue_aster_runtime_turn( + queued_turn: QueuedTurnRuntime, +) -> Result { + let store = require_aster_runtime_store()?; + store + .enqueue_turn(queued_turn) + .await + .map_err(|error| format!("写入 queued runtime turn 失败: {error}")) +} + +pub(crate) async fn remove_aster_runtime_queued_turn( + queued_turn_id: &str, +) -> Result, String> { + let store = require_aster_runtime_store()?; + store + .remove_queued_turn(queued_turn_id) + .await + .map_err(|error| format!("删除 queued runtime turn 失败: {error}")) +} + +pub(crate) async fn clear_aster_runtime_queued_turns( + session_id: &str, +) -> Result, String> { + let store = require_aster_runtime_store()?; + store + .clear_queued_turns(session_id) + .await + .map_err(|error| format!("清空 queued runtime turns 失败: {error}")) +} + +pub(crate) fn queued_turn_runtime_from_task(task: &QueuedTurnTask) -> QueuedTurnRuntime { + build_queued_turn_runtime(QueuedTurnRuntimeInput { + queued_turn_id: &task.queued_turn_id, + session_id: &task.session_id, + event_name: &task.event_name, + message_preview: &task.message_preview, + message_text: &task.message_text, + created_at: task.created_at, + image_count: task.image_count, + payload: task.payload.clone(), + }) +} + +pub(crate) fn queued_turn_event_name_from_runtime(queued_turn: &QueuedTurnRuntime) -> String { + queued_turn + .metadata + .get(QUEUED_TURN_EVENT_NAME_METADATA_KEY) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or(DEFAULT_QUEUE_EVENT_NAME) + .to_string() +} + +pub(crate) fn queued_turn_snapshot_from_runtime( + queued_turn: &QueuedTurnRuntime, + position: usize, +) -> QueuedTurnSnapshot { + QueuedTurnSnapshot { + queued_turn_id: queued_turn.queued_turn_id.clone(), + message_preview: queued_turn.message_preview.clone(), + message_text: queued_turn.message_text.clone(), + created_at: queued_turn.created_at, + image_count: queued_turn.image_count, + position, + } +} + +async fn migrate_legacy_runtime_queue_to_aster_store( + db: &DbConnection, +) -> Result { + ensure_aster_runtime_dirs()?; + let snapshot = { + let conn = lock_db(db)?; + agent_runtime_queue_repository::load_legacy_runtime_queue_snapshot(&conn)? + }; + let Some(snapshot) = snapshot else { + return Ok(LegacyRuntimeQueueMigrationReport::default()); + }; + + let mut report = LegacyRuntimeQueueMigrationReport { + session_count: snapshot.sessions.len(), + invalid_turn_count: snapshot.invalid_turn_count, + ..LegacyRuntimeQueueMigrationReport::default() + }; + + for session in snapshot.sessions { + let mut existing_ids = list_aster_runtime_queued_turns(&session.session_id) + .await? + .into_iter() + .map(|queued_turn| queued_turn.queued_turn_id) + .collect::>(); + + for queued_turn in session.turns { + if !existing_ids.insert(queued_turn.queued_turn_id.clone()) { + report.skipped_existing_turn_count += 1; + continue; + } + + enqueue_aster_runtime_turn(queued_turn_runtime_from_legacy_turn(&queued_turn)).await?; + report.migrated_turn_count += 1; + } + } + + { + let conn = lock_db(db)?; + agent_runtime_queue_repository::drop_legacy_runtime_queue_table(&conn)?; + } + tracing::info!( + "[AsterAgent][Queue] legacy 排队队列已导入到 Aster store: sessions={}, migrated={}, skipped_existing={}, invalid={}", + report.session_count, + report.migrated_turn_count, + report.skipped_existing_turn_count, + report.invalid_turn_count + ); + Ok(report) +} + +fn queued_turn_runtime_from_legacy_turn( + queued_turn: &LegacyRuntimeQueuedTurn, +) -> QueuedTurnRuntime { + build_queued_turn_runtime(QueuedTurnRuntimeInput { + queued_turn_id: &queued_turn.queued_turn_id, + session_id: &queued_turn.session_id, + event_name: &queued_turn.event_name, + message_preview: &queued_turn.message_preview, + message_text: &queued_turn.message_text, + created_at: queued_turn.created_at, + image_count: queued_turn.image_count, + payload: queued_turn.payload.clone(), + }) +} + +struct QueuedTurnRuntimeInput<'a> { + queued_turn_id: &'a str, + session_id: &'a str, + event_name: &'a str, + message_preview: &'a str, + message_text: &'a str, + created_at: i64, + image_count: usize, + payload: Value, +} + +fn build_queued_turn_runtime(input: QueuedTurnRuntimeInput<'_>) -> QueuedTurnRuntime { + let QueuedTurnRuntimeInput { + queued_turn_id, + session_id, + event_name, + message_preview, + message_text, + created_at, + image_count, + payload, + } = input; + let mut metadata = HashMap::new(); + if !event_name.trim().is_empty() { + metadata.insert( + QUEUED_TURN_EVENT_NAME_METADATA_KEY.to_string(), + Value::String(event_name.to_string()), + ); + } + + QueuedTurnRuntime { + queued_turn_id: queued_turn_id.to_string(), + session_id: session_id.to_string(), + message_preview: message_preview.to_string(), + message_text: message_text.to_string(), + created_at, + image_count, + payload, + metadata, + } +} diff --git a/src-tauri/crates/agent/src/aster_state.rs b/src-tauri/crates/agent/src/aster_state.rs index e7f278b73..22a6e216a 100644 --- a/src-tauri/crates/agent/src/aster_state.rs +++ b/src-tauri/crates/agent/src/aster_state.rs @@ -7,7 +7,7 @@ //! ## 重要:SessionStore 注入 //! //! 为了让 Aster Agent 的消息存储到 Lime 数据库,必须在创建 Agent 时 -//! 注入 `LimeSessionStore`。使用 `init_agent_with_db()` 方法而不是 `init_agent()`。 +//! 注入 `LimeSessionStore`,并统一通过 `init_agent_with_db()` 初始化。 //! //! ## Agent 身份配置 //! @@ -17,30 +17,27 @@ //! //! ## Skills 集成 //! -//! Agent 初始化时会自动加载 `~/.lime/skills/` 目录下的 Skills 到 +//! Agent 初始化时会自动加载 Lime 当前应用数据目录中的 Skills 到 //! aster-rust 的 global_registry,使 AI 能够自动发现和调用这些 Skills。 //! //! 参考文档:`docs/prd/chat-architecture-redesign.md` -use aster::agents::{Agent, SessionConfig}; +use aster::agents::Agent; 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; use tokio::sync::RwLock; use tokio_util::sync::CancellationToken; use crate::credential_bridge::{create_aster_provider, AsterProviderConfig, CredentialBridge}; +#[cfg(test)] use crate::queued_turn::QueuedTurnSnapshot; use lime_core::database::DbConnection; use lime_services::aster_session_store::LimeSessionStore; -use std::collections::{HashMap, VecDeque}; -use std::sync::Mutex; - async fn configure_lime_native_file_tools(agent: &Agent) { let shared_history = create_shared_history(); let registry_arc = agent.tool_registry().clone(); @@ -68,34 +65,8 @@ pub struct QueuedTurnTask { 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(), - } - } - + #[cfg(test)] fn snapshot(&self, position: usize) -> QueuedTurnSnapshot { QueuedTurnSnapshot { queued_turn_id: self.queued_turn_id.clone(), @@ -108,178 +79,6 @@ impl QueuedTurnTask { } } -#[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 { @@ -311,8 +110,6 @@ pub struct AsterAgentState { initialized_cache: Arc, /// Provider 配置状态缓存(避免每次都获取锁) provider_configured_cache: Arc, - /// 会话级 turn 队列 - turn_queue: SessionTurnQueueManager, } impl Clone for AsterAgentState { @@ -324,7 +121,6 @@ impl Clone for AsterAgentState { credential_bridge: CredentialBridge::new(), initialized_cache: self.initialized_cache.clone(), provider_configured_cache: self.provider_configured_cache.clone(), - turn_queue: self.turn_queue.clone(), } } } @@ -345,66 +141,21 @@ 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(), } } - fn resolve_aster_path_root() -> Result { - if let Ok(raw) = std::env::var("ASTER_PATH_ROOT") { - let trimmed = raw.trim(); - if !trimmed.is_empty() { - return Ok(PathBuf::from(trimmed)); - } - } - - let home_dir = dirs::home_dir().ok_or_else(|| "无法获取用户目录".to_string())?; - Ok(home_dir.join(".lime").join("aster")) - } - - fn ensure_aster_runtime_dirs() -> Result { - let root = Self::resolve_aster_path_root()?; - let root_string = root.to_string_lossy().to_string(); - - if std::env::var("ASTER_PATH_ROOT") - .ok() - .map(|value| value.trim().is_empty()) - .unwrap_or(true) - { - std::env::set_var("ASTER_PATH_ROOT", &root_string); - } - - let dirs = [ - root.join("config"), - root.join("data"), - root.join("state"), - root.join("state").join("logs"), - ]; - - for dir in dirs { - std::fs::create_dir_all(&dir).map_err(|e| { - format!( - "初始化 Aster 运行目录失败: {} ({})", - dir.to_string_lossy(), - e - ) - })?; - } - - Ok(root) - } - /// 初始化 Agent(带数据库连接) /// /// 创建 Agent 并注入 LimeSessionStore,确保消息存储到 Lime 数据库。 /// 同时设置 Lime 专属的 Agent 身份(名称、语言、描述)。 - /// 自动加载 `~/.lime/skills/` 目录下的 Skills 到 aster-rust 的 global_registry。 + /// 自动加载 Lime 当前应用数据目录中的 Skills 到 aster-rust 的 global_registry。 /// - /// **推荐使用此方法**而不是 `init_agent()`。 + /// 这是 Lime 当前唯一支持的 Agent 初始化入口。 /// /// # 参数 /// - `db`: 数据库连接,用于创建 SessionStore pub async fn init_agent_with_db(&self, db: &DbConnection) -> Result<(), String> { - let runtime_root = Self::ensure_aster_runtime_dirs()?; + let runtime_root = crate::aster_runtime_support::require_aster_runtime_dirs()?; tracing::info!( "[AsterAgent] Aster 运行目录已准备: {}", runtime_root.to_string_lossy() @@ -423,9 +174,10 @@ impl AsterAgentState { // 创建 Agent(启用 Ask/LSP 回调)并注入 SessionStore let tool_config = crate::create_lime_tool_config(); + let runtime_store = crate::aster_runtime_support::require_aster_runtime_store()?; let agent = Agent::with_tool_config(tool_config) .with_session_store(session_store) - .with_thread_runtime_store(aster::session::shared_thread_runtime_store()); + .with_thread_runtime_store(runtime_store); // 验证 session_store 是否被正确设置 let has_store = agent.session_store().is_some(); @@ -465,29 +217,6 @@ impl AsterAgentState { crate::reload_lime_skills(); } - /// 初始化 Agent(无数据库版本) - /// - /// **警告**:此方法创建的 Agent 不会将消息存储到 Lime 数据库, - /// 消息会存储到 Aster 默认的 `~/.aster/sessions.db`。 - /// - /// 建议使用 `init_agent_with_db()` 代替。 - #[deprecated( - since = "0.1.0", - note = "请使用 init_agent_with_db() 以确保消息存储到 Lime 数据库" - )] - pub async fn init_agent(&self) -> Result<(), String> { - let mut agent_guard = self.agent.write().await; - if agent_guard.is_none() { - let agent = Agent::new() - .with_thread_runtime_store(aster::session::shared_thread_runtime_store()); - *agent_guard = Some(agent); - tracing::warn!( - "[AsterAgent] Agent 初始化(无 SessionStore),消息将存储到 Aster 默认数据库" - ); - } - Ok(()) - } - /// 配置 Provider /// /// 根据配置创建并设置 Provider @@ -757,11 +486,6 @@ 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(); @@ -809,25 +533,6 @@ impl AsterAgentState { crate::build_project_system_prompt(db, project_id) } - /// 创建带项目上下文的会话配置 - /// - /// 自动加载项目配置并构建 SessionConfig。 - /// - /// # 参数 - /// - `db`: 数据库连接 - /// - `session_id`: 会话 ID - /// - `project_id`: 项目 ID(可选,如果为 None 则不注入项目上下文) - /// - /// # 返回 - /// - 构建好的 SessionConfig - pub fn create_session_config_with_project( - db: &DbConnection, - session_id: &str, - project_id: Option<&str>, - ) -> SessionConfig { - crate::create_session_config_with_project(db, session_id, project_id) - } - /// 注册 MCP 桥接客户端 /// /// 将 Lime 托管的 MCP 客户端注册到 Aster Agent 的 ExtensionManager, @@ -893,6 +598,7 @@ pub use crate::aster_state_support::{message_helpers, SessionConfigBuilder}; mod tests { use super::*; use std::fs; + use std::sync::{Arc, Mutex}; use tempfile::TempDir; #[tokio::test] @@ -900,8 +606,20 @@ mod tests { let state = AsterAgentState::new(); assert!(!state.is_initialized().await); - #[allow(deprecated)] - state.init_agent().await.unwrap(); + let runtime_dir = TempDir::new().unwrap(); + crate::aster_runtime_support::ensure_aster_runtime_dirs_with_root( + runtime_dir.path().to_path_buf(), + ) + .unwrap(); + + let db: DbConnection = + Arc::new(Mutex::new(rusqlite::Connection::open_in_memory().unwrap())); + { + let conn = db.lock().unwrap(); + lime_core::database::schema::create_tables(&conn).unwrap(); + } + + state.init_agent_with_db(&db).await.unwrap(); assert!(state.is_initialized().await); } @@ -921,105 +639,33 @@ mod tests { } #[test] - fn test_session_turn_queue_manager() { - let manager = SessionTurnQueueManager::default(); + fn test_session_turn_queue_manager_execution_gate() { + let gate = SessionTurnExecutionGate::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") - ); + assert!(gate.try_start("session-queue")); + assert!(gate.is_active("session-queue")); + assert!(!gate.try_start("session-queue")); + assert!(gate.finish("session-queue")); + assert!(!gate.is_active("session-queue")); } #[test] - fn test_session_turn_queue_manager_restore_pending() { - let manager = SessionTurnQueueManager::default(); + fn test_session_turn_queue_manager_snapshot() { + let task = 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" }), + }; - 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") - ); + let snapshot = task.snapshot(2); + assert_eq!(snapshot.queued_turn_id, "turn-restore-1"); + assert_eq!(snapshot.position, 2); + assert_eq!(snapshot.message_text, "restore body 1"); } // ========================================================================= @@ -1112,7 +758,7 @@ description: {} #[test] fn test_reload_lime_skills_no_panic() { // 这个测试确保 reload_lime_skills 在各种情况下都不会 panic - // 即使 ~/.lime/skills/ 目录不存在 + // 即使当前 Skills 目录不存在 AsterAgentState::reload_lime_skills(); // 如果没有 panic,测试通过 } diff --git a/src-tauri/crates/agent/src/aster_state_support.rs b/src-tauri/crates/agent/src/aster_state_support.rs index 91aa4c8e1..4d42928e5 100644 --- a/src-tauri/crates/agent/src/aster_state_support.rs +++ b/src-tauri/crates/agent/src/aster_state_support.rs @@ -7,6 +7,7 @@ use aster::agents::{AgentIdentity, SessionConfig}; use aster::session::TurnContextOverride; use aster::skills::{global_registry, load_skills_from_directory, SkillSource}; use aster::tools::ToolRegistrationConfig; +use lime_core::app_paths; use lime_core::database::{lock_db, DbConnection}; use lime_services::project_context_builder::ProjectContextBuilder; @@ -34,23 +35,17 @@ pub fn create_lime_tool_config() -> ToolRegistrationConfig { /// 加载 Lime Skills 到 aster-rust 的 global_registry fn load_lime_skills() { - let home = match dirs::home_dir() { - Some(home_dir) => home_dir, - None => { - tracing::warn!("[AsterAgent] 无法获取 home 目录,跳过 Skills 加载"); + let skills_dir = match app_paths::resolve_skills_dir() { + Ok(path) => path, + Err(error) => { + tracing::warn!( + "[AsterAgent] 解析 Lime Skills 目录失败,跳过加载: {}", + error + ); return; } }; - let skills_dir = home.join(".lime").join("skills"); - if !skills_dir.exists() { - tracing::info!( - "[AsterAgent] Lime Skills 目录不存在: {:?},跳过加载", - skills_dir - ); - return; - } - let skills = load_skills_from_directory(&skills_dir, SkillSource::User); let skill_count = skills.len(); @@ -82,19 +77,6 @@ pub fn build_project_system_prompt(db: &DbConnection, project_id: &str) -> Resul .map_err(|e| format!("构建项目上下文失败: {e}")) } -/// 创建带项目上下文的会话配置 -pub fn create_session_config_with_project( - db: &DbConnection, - session_id: &str, - project_id: Option<&str>, -) -> SessionConfig { - let system_prompt = project_id.and_then(|pid| build_project_system_prompt(db, pid).ok()); - - SessionConfigBuilder::new(session_id) - .system_prompt(system_prompt.unwrap_or_default()) - .build() -} - /// 会话配置构建器 pub struct SessionConfigBuilder { id: String, diff --git a/src-tauri/crates/agent/src/event_converter.rs b/src-tauri/crates/agent/src/event_converter.rs index c6de9696e..7d1bc3a4e 100644 --- a/src-tauri/crates/agent/src/event_converter.rs +++ b/src-tauri/crates/agent/src/event_converter.rs @@ -871,6 +871,29 @@ fn convert_item_status( } } +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 convert_item_payload(payload: ItemRuntimePayload) -> AgentThreadItemPayload { match payload { ItemRuntimePayload::UserMessage { content } => { @@ -879,6 +902,26 @@ fn convert_item_payload(payload: ItemRuntimePayload) -> AgentThreadItemPayload { ItemRuntimePayload::AgentMessage { text } => { AgentThreadItemPayload::AgentMessage { text, phase: None } } + ItemRuntimePayload::Plan { text } => AgentThreadItemPayload::Plan { text }, + ItemRuntimePayload::RuntimeStatus { + phase: _, + title, + detail, + checkpoints, + } => AgentThreadItemPayload::TurnSummary { + text: format_runtime_status_text(&title, &detail, &checkpoints), + }, + ItemRuntimePayload::FileArtifact { + path, + source, + content, + metadata, + } => AgentThreadItemPayload::FileArtifact { + path, + source, + content, + metadata, + }, ItemRuntimePayload::Reasoning { text } => AgentThreadItemPayload::Reasoning { text, summary: None, @@ -1396,6 +1439,118 @@ mod tests { } } + #[test] + fn test_convert_item_started_plan_runtime_item() { + let now = chrono::Utc::now(); + let item = ItemRuntime { + id: "plan:turn-1".to_string(), + thread_id: "thread-1".to_string(), + turn_id: "turn-1".to_string(), + sequence: 2, + status: ItemStatus::InProgress, + started_at: now, + completed_at: None, + updated_at: now, + payload: ItemRuntimePayload::Plan { + text: "- 调研\n- 实现".to_string(), + }, + }; + + let events = convert_agent_event(AgentEvent::ItemStarted { item }); + assert_eq!(events.len(), 1); + match &events[0] { + TauriAgentEvent::ItemStarted { item } => match &item.payload { + AgentThreadItemPayload::Plan { text } => { + assert_eq!(text, "- 调研\n- 实现"); + } + other => panic!("Unexpected payload: {other:?}"), + }, + other => panic!("Expected ItemStarted event, got {other:?}"), + } + } + + #[test] + fn test_convert_item_started_file_artifact_runtime_item() { + let now = chrono::Utc::now(); + let item = ItemRuntime { + id: "artifact-1".to_string(), + thread_id: "thread-1".to_string(), + turn_id: "turn-1".to_string(), + sequence: 3, + status: ItemStatus::Completed, + started_at: now, + completed_at: Some(now), + updated_at: now, + payload: ItemRuntimePayload::FileArtifact { + path: "/tmp/result.md".to_string(), + source: "tool_result".to_string(), + content: None, + metadata: Some(serde_json::json!({ + "output_file": "/tmp/result.md", + "artifact_id": "artifact-1" + })), + }, + }; + + let events = convert_agent_event(AgentEvent::ItemStarted { item }); + assert_eq!(events.len(), 1); + match &events[0] { + TauriAgentEvent::ItemStarted { item } => match &item.payload { + AgentThreadItemPayload::FileArtifact { + path, + source, + content, + metadata, + } => { + assert_eq!(path, "/tmp/result.md"); + assert_eq!(source, "tool_result"); + assert_eq!(content, &None); + assert_eq!( + metadata.as_ref().and_then(|value| value.get("artifact_id")), + Some(&serde_json::json!("artifact-1")) + ); + } + other => panic!("Unexpected payload: {other:?}"), + }, + other => panic!("Expected ItemStarted event, got {other:?}"), + } + } + + #[test] + fn test_convert_item_updated_runtime_status_runtime_item() { + let now = chrono::Utc::now(); + let item = ItemRuntime { + id: "turn_summary:turn-1".to_string(), + thread_id: "thread-1".to_string(), + turn_id: "turn-1".to_string(), + sequence: 4, + status: ItemStatus::InProgress, + started_at: now, + completed_at: None, + updated_at: now, + payload: ItemRuntimePayload::RuntimeStatus { + phase: "routing".to_string(), + title: "已决定:先规划再输出".to_string(), + detail: "当前请求更像计划拆解,会先输出结构化行动路径。".to_string(), + checkpoints: vec!["检测到计划需求".to_string(), "优先整理关键步骤".to_string()], + }, + }; + + let events = convert_agent_event(AgentEvent::ItemUpdated { item }); + assert_eq!(events.len(), 1); + match &events[0] { + TauriAgentEvent::ItemUpdated { item } => match &item.payload { + AgentThreadItemPayload::TurnSummary { text } => { + assert!(text.contains("已决定:先规划再输出")); + assert!(text.contains("当前请求更像计划拆解")); + assert!(text.contains("• 检测到计划需求")); + } + other => panic!("Unexpected payload: {other:?}"), + }, + other => panic!("Expected ItemUpdated event, got {other:?}"), + } + } + #[test] fn test_convert_item_started_request_user_input() { let now = chrono::Utc::now(); diff --git a/src-tauri/crates/agent/src/lib.rs b/src-tauri/crates/agent/src/lib.rs index 649f05223..998da2564 100644 --- a/src-tauri/crates/agent/src/lib.rs +++ b/src-tauri/crates/agent/src/lib.rs @@ -11,6 +11,7 @@ #![allow(clippy::borrowed_box)] pub mod ask_bridge; +pub mod aster_runtime_support; pub mod aster_state; pub mod aster_state_support; pub mod credential_bridge; @@ -22,8 +23,10 @@ pub mod mcp_bridge; pub mod prompt; pub mod queued_turn; pub mod request_tool_policy; -pub mod session_store; +pub mod runtime_queue; +mod session_store; pub mod shell_security; +pub mod skill_execution; pub mod subagent_scheduler; pub mod tool_io_offload; pub mod tool_permissions; @@ -31,11 +34,11 @@ 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_runtime_support::initialize_aster_runtime; +pub use aster_state::{AsterAgentState, ProviderConfig, QueuedTurnTask}; pub use aster_state_support::{ - build_project_system_prompt, create_lime_identity, create_lime_tool_config, - create_session_config_with_project, message_helpers, reload_lime_skills, SessionConfigBuilder, + build_project_system_prompt, create_lime_identity, create_lime_tool_config, message_helpers, + reload_lime_skills, SessionConfigBuilder, }; pub use credential_bridge::{ create_aster_provider, AsterProviderConfig, CredentialBridge, CredentialBridgeError, @@ -59,14 +62,23 @@ pub use request_tool_policy::{ ReplyAttemptError, RequestToolPolicy, RequestToolPolicyMode, StreamReplyExecution, WebSearchExecutionTracker, REQUEST_TOOL_POLICY_MARKER, }; +pub use runtime_queue::{ + clear_runtime_queue, list_runtime_queue_snapshots, remove_runtime_queued_turn, + resume_persisted_runtime_queues_on_startup, resume_runtime_queue_if_needed, + submit_runtime_turn, RuntimeQueueEventEmitter, RuntimeQueueExecutor, +}; pub use session_store::{ - create_session_record_sync, create_session_sync, get_compat_session_sync, - get_persisted_session_metadata_sync, get_session_sync, list_compat_sessions_sync, - list_sessions_sync, list_title_preview_messages_sync, update_session_execution_strategy_sync, - update_session_working_dir_sync, CompatSessionInfo, CreateSessionRecordInput, - PersistedSessionMetadata, SessionDetail, SessionInfo, SessionTitlePreviewMessage, + create_session_sync, delete_session, get_persisted_session_metadata_sync, + get_runtime_session_detail, get_session_sync, list_sessions_sync, + list_title_preview_messages_sync, rename_session_sync, update_session_execution_strategy_sync, + update_session_working_dir_sync, PersistedSessionMetadata, SessionDetail, SessionInfo, + SessionTitlePreviewMessage, SessionTodoItem, }; pub use shell_security::ShellSecurityChecker; +pub use skill_execution::{ + execute_skill_prompt, execute_skill_workflow, SkillEventEmitter, SkillExecutionError, + SkillExecutionResult, SkillWorkflowExecution, StepResult, +}; pub use subagent_scheduler::{ LimeScheduler, LimeSubAgentExecutor, SchedulerEventEmitter, SubAgentProgressEvent, SubAgentRole, }; diff --git a/src-tauri/crates/agent/src/prompt/instruction_discovery.rs b/src-tauri/crates/agent/src/prompt/instruction_discovery.rs index 84f64f8f2..36114382c 100644 --- a/src-tauri/crates/agent/src/prompt/instruction_discovery.rs +++ b/src-tauri/crates/agent/src/prompt/instruction_discovery.rs @@ -3,19 +3,14 @@ //! 从文件系统发现并加载多层级的 AGENT.md 指令文件, //! 按优先级从低到高:全局 -> 项目根 -> 当前目录 +use lime_core::app_paths; use std::collections::HashSet; use std::path::{Path, PathBuf}; use std::sync::RwLock; use std::time::{Duration, Instant}; /// 支持的指令文件名列表(按优先级排序) -const INSTRUCTION_FILENAMES: &[&str] = &[ - "AGENT.md", - ".agent.md", - "agent.md", - ".lime/AGENT.md", - ".lime/instructions.md", -]; +const INSTRUCTION_FILENAMES: &[&str] = &["AGENT.md", ".agent.md", "agent.md"]; // 保留旧常量供测试使用(第一优先级文件名) #[cfg(test)] @@ -24,7 +19,7 @@ const INSTRUCTION_FILENAME: &str = "AGENT.md"; /// 指令来源,按优先级从低到高 #[derive(Debug, Clone, PartialEq)] pub enum InstructionSource { - /// ~/.lime/AGENT.md + /// 应用级用户指令文件(统一收口到 app_paths::resolve_user_memory_path) Global, /// 项目根目录/AGENT.md Project, @@ -51,16 +46,19 @@ fn find_instruction_file(dir: &Path) -> Option { None } +fn find_global_instruction_file() -> Option { + app_paths::resolve_user_memory_path() + .ok() + .filter(|path| path.is_file()) +} + /// 从文件系统发现并加载层级化指令 /// 返回按优先级排序的指令列表(低优先级在前) pub fn discover_instructions(working_dir: &Path) -> Vec { let mut layers = Vec::new(); - // 1. 全局: ~/.lime/ 下查找指令文件 - if let Some(home) = dirs::home_dir() { - let global_dir = home.join(".lime"); - // 全局层只查找 AGENT.md(不递归子目录模式) - let global_path = global_dir.join("AGENT.md"); + // 1. 全局: 统一用户指令文件 + if let Some(global_path) = find_global_instruction_file() { if let Some(layer) = load_layer(&global_path, InstructionSource::Global) { layers.push(layer); } @@ -384,7 +382,7 @@ mod tests { InstructionLayer { source: InstructionSource::Global, content: "global rule".to_string(), - path: PathBuf::from("/home/.lime/AGENT.md"), + path: PathBuf::from("/appdata/lime/AGENTS.md"), }, InstructionLayer { source: InstructionSource::Project, @@ -460,6 +458,22 @@ mod tests { assert!(project[0].content.contains("primary agent")); } + #[test] + fn test_project_discovery_ignores_legacy_lime_subdir_files() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + let legacy_dir = tmp.path().join(".lime"); + fs::create_dir(&legacy_dir).unwrap(); + fs::write(legacy_dir.join("AGENT.md"), "legacy nested agent").unwrap(); + + let layers = discover_instructions(tmp.path()); + let project: Vec<_> = layers + .iter() + .filter(|l| l.source == InstructionSource::Project) + .collect(); + assert!(project.is_empty()); + } + #[test] fn test_include_directive() { let tmp = TempDir::new().unwrap(); diff --git a/src-tauri/crates/agent/src/request_tool_policy.rs b/src-tauri/crates/agent/src/request_tool_policy.rs index c68ea6935..d11bf9bc4 100644 --- a/src-tauri/crates/agent/src/request_tool_policy.rs +++ b/src-tauri/crates/agent/src/request_tool_policy.rs @@ -1014,6 +1014,41 @@ fn duplicate_session_config(config: &aster::agents::SessionConfig) -> aster::age } } +async fn emit_runtime_status_with_projection( + agent: &Agent, + session_config: &aster::agents::SessionConfig, + status: TauriRuntimeStatus, + on_event: &mut F, +) where + F: FnMut(&TauriAgentEvent), +{ + match agent + .upsert_runtime_status_item( + session_config, + status.phase.clone(), + status.title.clone(), + status.detail.clone(), + status.checkpoints.clone(), + ) + .await + { + Ok(agent_event) => { + for event in convert_agent_event(agent_event) { + on_event(&event); + } + } + Err(error) => { + tracing::warn!( + "[AsterAgent][RuntimeStatus] 写入 runtime item 失败,降级仅发 transient 事件: {}", + error + ); + } + } + + let event = TauriAgentEvent::RuntimeStatus { status }; + on_event(&event); +} + fn should_retry_after_empty_reply( preflight_execution: &PreflightToolExecution, current_text_output: &str, @@ -1320,6 +1355,30 @@ pub async fn stream_reply_with_policy( agent: &Agent, message_text: &str, working_directory: Option<&Path>, + session_config: aster::agents::SessionConfig, + cancel_token: Option, + request_tool_policy: &RequestToolPolicy, + on_event: F, +) -> Result +where + F: FnMut(&TauriAgentEvent), +{ + stream_message_reply_with_policy( + agent, + Message::user().with_text(message_text), + working_directory, + session_config, + cancel_token, + request_tool_policy, + on_event, + ) + .await +} + +pub async fn stream_message_reply_with_policy( + agent: &Agent, + user_message: Message, + working_directory: Option<&Path>, mut session_config: aster::agents::SessionConfig, cancel_token: Option, request_tool_policy: &RequestToolPolicy, @@ -1328,11 +1387,12 @@ pub async fn stream_reply_with_policy( where F: FnMut(&TauriAgentEvent), { + let message_text = user_message.as_concat_text(); let mut web_search_tracker = WebSearchExecutionTracker::default(); let preflight = execute_web_search_preflight_if_needed( agent, &session_config.id, - message_text, + &message_text, working_directory, cancel_token.clone(), request_tool_policy, @@ -1368,7 +1428,7 @@ where let mut diagnostics = StreamEventDiagnostics::default(); stream_agent_reply_once( agent, - Message::user().with_text(message_text), + user_message, duplicate_session_config(&session_config), cancel_token.clone(), request_tool_policy, @@ -1393,12 +1453,15 @@ where session_config.id, web_search_tracker.format_attempts() ); - let status = TauriAgentEvent::RuntimeStatus { - status: build_web_search_synthesis_runtime_status( + emit_runtime_status_with_projection( + agent, + &session_config, + build_web_search_synthesis_runtime_status( preflight_execution.coverage_summary.as_deref(), ), - }; - on_event(&status); + &mut on_event, + ) + .await; session_config.system_prompt = merge_system_prompt_with_web_search_synthesis_instruction( session_config.system_prompt.take(), ); diff --git a/src-tauri/crates/agent/src/runtime_queue.rs b/src-tauri/crates/agent/src/runtime_queue.rs new file mode 100644 index 000000000..36dea5bd6 --- /dev/null +++ b/src-tauri/crates/agent/src/runtime_queue.rs @@ -0,0 +1,277 @@ +use crate::aster_runtime_support::{ + clear_aster_runtime_queued_turns, list_aster_runtime_queued_turns, + prepare_aster_runtime_queue_resumption, queued_turn_event_name_from_runtime, + queued_turn_runtime_from_task, queued_turn_snapshot_from_runtime, + remove_aster_runtime_queued_turn, +}; +use crate::{QueuedTurnSnapshot, QueuedTurnTask, TauriAgentEvent}; +use aster::session::{ + require_shared_session_runtime_queue_service, QueuedTurnRuntime, RuntimeQueueSubmitResult, +}; +use futures::future::BoxFuture; +use serde_json::Value; +use std::sync::Arc; + +pub type RuntimeQueueExecutor = + Arc BoxFuture<'static, Result<(), String>> + Send + Sync>; + +pub type RuntimeQueueEventEmitter = Arc; + +fn emit_runtime_queue_event( + emitter: &RuntimeQueueEventEmitter, + event_name: &str, + event: TauriAgentEvent, +) { + emitter(event_name.to_string(), event); +} + +async fn continue_runtime_queue_after_turn( + session_id: String, + context: C, + executor: RuntimeQueueExecutor, + emitter: RuntimeQueueEventEmitter, +) -> Result +where + C: Clone + Send + Sync + 'static, +{ + start_next_runtime_queue_turn(session_id, false, context, executor, emitter).await +} + +async fn start_next_runtime_queue_turn( + session_id: String, + acquire_gate: bool, + context: C, + executor: RuntimeQueueExecutor, + emitter: RuntimeQueueEventEmitter, +) -> Result +where + C: Clone + Send + Sync + 'static, +{ + let runtime_queue_service = require_shared_session_runtime_queue_service() + .map_err(|error| format!("读取 runtime queue service 失败: {error}"))?; + let next_queued_turn = match if acquire_gate { + runtime_queue_service.resume_if_idle(&session_id).await + } else { + runtime_queue_service + .finish_turn_and_take_next(&session_id) + .await + } { + Ok(next_queued_turn) => next_queued_turn, + Err(error) => { + return Err(format!("读取下一条 runtime queue turn 失败: {}", error)); + } + }; + let Some(next_queued_turn) = next_queued_turn else { + return Ok(false); + }; + + let event_name = queued_turn_event_name_from_runtime(&next_queued_turn); + emit_runtime_queue_event( + &emitter, + &event_name, + TauriAgentEvent::QueueStarted { + session_id: session_id.clone(), + queued_turn_id: next_queued_turn.queued_turn_id.clone(), + }, + ); + + let runtime_handle = tokio::runtime::Handle::current(); + let executor_for_task = executor.clone(); + let emitter_for_task = emitter.clone(); + let context_for_task = context.clone(); + let session_id_for_task = session_id.clone(); + let payload = next_queued_turn.payload; + tokio::task::spawn_blocking(move || { + runtime_handle.block_on(async move { + let result = executor_for_task(context_for_task.clone(), payload).await; + if let Err(error) = continue_runtime_queue_after_turn( + session_id_for_task, + context_for_task.clone(), + executor_for_task.clone(), + emitter_for_task.clone(), + ) + .await + { + tracing::warn!("[AsterAgent][Queue] 调度下一条排队 turn 失败: {}", error); + } + if let Err(error) = result { + tracing::warn!("[AsterAgent][Queue] 队列任务执行失败: {}", error); + } + }); + }); + Ok(true) +} + +pub async fn resume_runtime_queue_if_needed( + session_id: String, + context: C, + executor: RuntimeQueueExecutor, + emitter: RuntimeQueueEventEmitter, +) -> Result +where + C: Clone + Send + Sync + 'static, +{ + if list_aster_runtime_queued_turns(&session_id) + .await? + .is_empty() + { + return Ok(false); + } + + start_next_runtime_queue_turn(session_id, true, context, executor, emitter).await +} + +pub async fn submit_runtime_turn( + queued_task: QueuedTurnTask, + queue_if_busy: bool, + context: C, + executor: RuntimeQueueExecutor, + emitter: RuntimeQueueEventEmitter, +) -> Result<(), String> +where + C: Clone + Send + Sync + 'static, +{ + let runtime_queue_service = require_shared_session_runtime_queue_service() + .map_err(|error| format!("读取 runtime queue service 失败: {error}"))?; + let session_id = queued_task.session_id.clone(); + let _ = resume_runtime_queue_if_needed( + session_id.clone(), + context.clone(), + executor.clone(), + emitter.clone(), + ) + .await?; + + match runtime_queue_service + .submit_turn(queued_turn_runtime_from_task(&queued_task), queue_if_busy) + .await + .map_err(|error| format!("提交 runtime queue turn 失败: {error}"))? + { + RuntimeQueueSubmitResult::StartNow => { + let result = executor(context.clone(), queued_task.payload).await; + if let Err(error) = + continue_runtime_queue_after_turn(session_id, context, executor.clone(), emitter) + .await + { + tracing::warn!("[AsterAgent][Queue] 调度下一条排队 turn 失败: {}", error); + } + result + } + RuntimeQueueSubmitResult::Busy => Err("当前会话仍在生成,无法立即开始执行".to_string()), + RuntimeQueueSubmitResult::Enqueued { + queued_turn, + position, + } => { + emit_runtime_queue_event( + &emitter, + &queued_turn_event_name_from_runtime(&queued_turn), + TauriAgentEvent::QueueAdded { + session_id, + queued_turn: queued_turn_snapshot_from_runtime(&queued_turn, position), + }, + ); + Ok(()) + } + } +} + +pub async fn clear_runtime_queue( + session_id: &str, + emitter: RuntimeQueueEventEmitter, +) -> Result, String> { + let cleared = clear_aster_runtime_queued_turns(session_id).await?; + if cleared.is_empty() { + return Ok(cleared); + } + + let queued_turn_ids = cleared + .iter() + .map(|queued_turn| queued_turn.queued_turn_id.clone()) + .collect::>(); + for queued_turn in &cleared { + emit_runtime_queue_event( + &emitter, + &queued_turn_event_name_from_runtime(queued_turn), + TauriAgentEvent::QueueCleared { + session_id: session_id.to_string(), + queued_turn_ids: queued_turn_ids.clone(), + }, + ); + } + + Ok(cleared) +} + +pub async fn list_runtime_queue_snapshots( + session_id: &str, +) -> Result, String> { + Ok(list_aster_runtime_queued_turns(session_id) + .await? + .iter() + .enumerate() + .map(|(index, queued_turn)| queued_turn_snapshot_from_runtime(queued_turn, index + 1)) + .collect()) +} + +pub async fn remove_runtime_queued_turn( + session_id: &str, + queued_turn_id: &str, + emitter: RuntimeQueueEventEmitter, +) -> Result { + let queued_turns = list_aster_runtime_queued_turns(session_id).await?; + let Some(existing) = queued_turns + .into_iter() + .find(|queued_turn| queued_turn.queued_turn_id == queued_turn_id) + else { + return Ok(false); + }; + + let removed = remove_aster_runtime_queued_turn(queued_turn_id).await?; + let Some(queued_turn) = removed else { + return Ok(false); + }; + + emit_runtime_queue_event( + &emitter, + &queued_turn_event_name_from_runtime(&existing), + TauriAgentEvent::QueueRemoved { + session_id: session_id.to_string(), + queued_turn_id: queued_turn.queued_turn_id, + }, + ); + Ok(true) +} + +pub async fn resume_persisted_runtime_queues_on_startup( + context: C, + executor: RuntimeQueueExecutor, + emitter: RuntimeQueueEventEmitter, +) -> Result +where + C: Clone + Send + Sync + 'static, +{ + let session_ids = prepare_aster_runtime_queue_resumption().await?; + if session_ids.is_empty() { + return Ok(0); + } + + let mut resumed = 0usize; + for session_id in session_ids { + if resume_runtime_queue_if_needed( + session_id.clone(), + context.clone(), + executor.clone(), + emitter.clone(), + ) + .await? + { + resumed += 1; + tracing::info!( + "[AsterAgent][Queue] 启动阶段已恢复会话排队执行: session_id={}", + session_id + ); + } + } + + Ok(resumed) +} diff --git a/src-tauri/crates/agent/src/session_store.rs b/src-tauri/crates/agent/src/session_store.rs index 88b9b1786..64e38c445 100644 --- a/src-tauri/crates/agent/src/session_store.rs +++ b/src-tauri/crates/agent/src/session_store.rs @@ -1,26 +1,36 @@ //! Agent 会话存储服务 //! //! 提供会话创建、列表查询、详情查询能力。 -//! 数据来源为 Lime 数据库(AgentDao)。 +//! 数据事实源收敛到 lime_core::database::agent_session_repository + Lime 数据库。 +use aster::session::extension_data::{resolve_todo_list_state, TodoListItem, TodoListItemStatus}; +use aster::session::SessionRuntimeSnapshot; use chrono::Utc; use lime_core::agent::types::{AgentMessage, AgentSession, ContentPart, MessageContent}; -use lime_core::database::dao::agent::AgentDao; +use lime_core::database::agent_session_repository::{ + self, SessionRecordDetail, SessionRecordMetadata, SessionRecordOverview, + SessionRecordPreviewMessage, +}; use lime_core::database::dao::agent_timeline::{ AgentThreadItem, AgentThreadTurn, AgentTimelineDao, }; use lime_core::database::DbConnection; use lime_core::workspace::WorkspaceManager; use lime_services::aster_session_store::LimeSessionStore; -use rusqlite::{Connection, OptionalExtension}; +use std::collections::HashMap; use uuid::Uuid; -use crate::event_converter::{TauriMessage, TauriMessageContent}; +use crate::aster_runtime_support::load_aster_runtime_snapshot; +use crate::event_converter::{ + convert_item_runtime, convert_turn_runtime, TauriMessage, TauriMessageContent, +}; use crate::tool_io_offload::{ build_history_tool_io_eviction_plan_for_model, force_offload_plain_tool_output_for_history, force_offload_tool_arguments_for_history, maybe_offload_plain_tool_output, maybe_offload_tool_arguments, }; +#[cfg(test)] +use lime_core::database::dao::agent::AgentDao; /// 会话信息(简化版) #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] @@ -51,25 +61,28 @@ pub struct SessionDetail { pub execution_strategy: Option, pub turns: Vec, pub items: Vec, + #[serde(default)] + pub todo_items: Vec, } -/// 兼容旧 Agent API 的会话摘要 -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct CompatSessionInfo { - pub session_id: String, - pub provider_type: String, - pub model: Option, - pub title: Option, - pub created_at: String, - pub last_activity: String, - pub messages_count: usize, - pub workspace_id: Option, - pub working_dir: Option, - pub execution_strategy: Option, +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum SessionTodoStatus { + Pending, + InProgress, + Completed, +} + +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)] +pub struct SessionTodoItem { + pub content: String, + pub status: SessionTodoStatus, + #[serde(skip_serializing_if = "Option::is_none")] + pub active_form: Option, } #[derive(Debug, Clone, Default)] -pub struct CreateSessionRecordInput { +pub(crate) struct CreateSessionRecordInput { pub session_id: Option, pub title: Option, pub model: Option, @@ -110,96 +123,147 @@ fn normalize_optional_nonempty_body(value: Option) -> Option { } } -fn load_agent_session_record( - conn: &Connection, - session_id: &str, -) -> Result, String> { - AgentDao::get_session(conn, session_id).map_err(|e| format!("获取会话失败: {e}")) +fn map_session_todo_status(status: TodoListItemStatus) -> SessionTodoStatus { + match status { + TodoListItemStatus::Pending => SessionTodoStatus::Pending, + TodoListItemStatus::InProgress => SessionTodoStatus::InProgress, + TodoListItemStatus::Completed => SessionTodoStatus::Completed, + } } -fn load_agent_session_messages( - conn: &Connection, - session_id: &str, -) -> Result, String> { - AgentDao::get_messages(conn, session_id).map_err(|e| format!("获取消息失败: {e}")) -} - -fn resolve_workspace_id_by_working_dir( - conn: &Connection, - working_dir: Option<&str>, -) -> Option { - let resolved_working_dir = working_dir?.trim(); - if resolved_working_dir.is_empty() { +fn map_session_todo_item(item: TodoListItem) -> Option { + let content = item.content.trim().to_string(); + if content.is_empty() { return None; } - match conn - .query_row( - "SELECT id FROM workspaces WHERE root_path = ? LIMIT 1", - [resolved_working_dir], - |row| row.get::<_, String>(0), - ) - .optional() - { - Ok(workspace_id) => workspace_id, + let active_form = normalize_optional_nonempty_body(Some(item.active_form)); + Some(SessionTodoItem { + content, + status: map_session_todo_status(item.status), + active_form, + }) +} + +fn load_session_todo_items_from_conn( + conn: &rusqlite::Connection, + session_id: &str, +) -> Vec { + let extension_data = match LimeSessionStore::load_extension_data_from_conn(conn, session_id) { + Ok(extension_data) => extension_data, Err(error) => { tracing::warn!( - "[SessionStore] 解析 workspace_id 失败,已降级忽略: working_dir={}, error={}", - resolved_working_dir, + "[SessionStore] 读取 session todo 状态失败: session_id={}, error={}", + session_id, error ); - None + return Vec::new(); + } + }; + + resolve_todo_list_state(&extension_data) + .map(|todo_list| { + todo_list + .items + .into_iter() + .filter_map(map_session_todo_item) + .collect() + }) + .unwrap_or_default() +} + +fn sort_runtime_turns(turns: &mut [AgentThreadTurn]) { + turns.sort_by(|left, right| { + left.started_at + .cmp(&right.started_at) + .then(left.created_at.cmp(&right.created_at)) + .then(left.id.cmp(&right.id)) + }); +} + +fn sort_runtime_items(items: &mut [AgentThreadItem], turn_started_at: &HashMap) { + items.sort_by(|left, right| { + let left_turn_started = turn_started_at + .get(&left.turn_id) + .map(String::as_str) + .unwrap_or(left.started_at.as_str()); + let right_turn_started = turn_started_at + .get(&right.turn_id) + .map(String::as_str) + .unwrap_or(right.started_at.as_str()); + + left_turn_started + .cmp(right_turn_started) + .then(left.sequence.cmp(&right.sequence)) + .then(left.turn_id.cmp(&right.turn_id)) + .then(left.started_at.cmp(&right.started_at)) + .then(left.id.cmp(&right.id)) + }); +} + +fn apply_aster_runtime_snapshot(detail: &mut SessionDetail, snapshot: &SessionRuntimeSnapshot) { + if let Some(thread) = snapshot.threads.first() { + detail.thread_id = thread.thread.id.clone(); + } + + if snapshot.threads.is_empty() { + return; + } + + let mut turns_by_id = detail + .turns + .drain(..) + .map(|turn| (turn.id.clone(), turn)) + .collect::>(); + for thread in &snapshot.threads { + for turn in &thread.turns { + turns_by_id.insert(turn.id.clone(), convert_turn_runtime(turn.clone())); } } + detail.turns = turns_by_id.into_values().collect(); + sort_runtime_turns(&mut detail.turns); + + let turn_started_at = detail + .turns + .iter() + .map(|turn| (turn.id.clone(), turn.started_at.clone())) + .collect::>(); + + let mut items_by_id = detail + .items + .drain(..) + .map(|item| (item.id.clone(), item)) + .collect::>(); + for thread in &snapshot.threads { + for item in &thread.items { + items_by_id.insert(item.id.clone(), convert_item_runtime(item.clone())); + } + } + detail.items = items_by_id.into_values().collect(); + sort_runtime_items(&mut detail.items, &turn_started_at); } -fn build_runtime_session_info( - conn: &Connection, - session: AgentSession, - messages_count: usize, -) -> SessionInfo { - let working_dir = session.working_dir.clone(); - let workspace_id = resolve_workspace_id_by_working_dir(conn, working_dir.as_deref()); +fn build_runtime_session_info(overview: SessionRecordOverview) -> SessionInfo { + let working_dir = overview.working_dir; + let workspace_id = overview.workspace_id; SessionInfo { - id: session.id, - name: session.title.unwrap_or_else(|| "未命名".to_string()), - created_at: chrono::DateTime::parse_from_rfc3339(&session.created_at) + id: overview.id, + name: overview.title.unwrap_or_else(|| "未命名".to_string()), + created_at: chrono::DateTime::parse_from_rfc3339(&overview.created_at) .map(|dt| dt.timestamp()) .unwrap_or(0), - updated_at: chrono::DateTime::parse_from_rfc3339(&session.updated_at) + updated_at: chrono::DateTime::parse_from_rfc3339(&overview.updated_at) .map(|dt| dt.timestamp()) .unwrap_or(0), - messages_count, - execution_strategy: session.execution_strategy, - model: Some(session.model), + messages_count: overview.messages_count, + execution_strategy: overview.execution_strategy, + model: Some(overview.model), working_dir, workspace_id, } } -fn build_compat_session_info( - conn: &Connection, - session: AgentSession, - messages_count: usize, -) -> CompatSessionInfo { - let working_dir = session.working_dir.clone(); - let workspace_id = resolve_workspace_id_by_working_dir(conn, working_dir.as_deref()); - - CompatSessionInfo { - session_id: session.id, - provider_type: "aster".to_string(), - model: Some(session.model), - title: session.title, - created_at: session.created_at, - last_activity: session.updated_at, - messages_count, - workspace_id, - working_dir, - execution_strategy: session.execution_strategy, - } -} - /// 解析会话 working_dir(优先入参,其次 workspace_id) fn resolve_session_working_dir( db: &DbConnection, @@ -251,7 +315,7 @@ fn resolve_optional_session_working_dir( } /// 创建并持久化会话记录 -pub fn create_session_record_sync( +pub(crate) fn create_session_record_sync( db: &DbConnection, input: CreateSessionRecordInput, ) -> Result { @@ -273,7 +337,7 @@ pub fn create_session_record_sync( }; let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - AgentDao::create_session(&conn, &session).map_err(|e| format!("创建会话失败: {e}"))?; + agent_session_repository::create_session(&conn, &session)?; Ok(session) } @@ -303,29 +367,11 @@ pub fn create_session_sync( /// 列出所有会话 pub fn list_sessions_sync(db: &DbConnection) -> Result, String> { let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let sessions = AgentDao::list_sessions(&conn).map_err(|e| format!("获取会话列表失败: {e}"))?; + let sessions = agent_session_repository::list_session_overviews(&conn)?; Ok(sessions .into_iter() - .map(|session| { - let messages_count = AgentDao::get_message_count(&conn, &session.id).unwrap_or(0); - build_runtime_session_info(&conn, session, messages_count) - }) - .collect()) -} - -/// 列出兼容旧 Agent API 的会话摘要 -pub fn list_compat_sessions_sync(db: &DbConnection) -> Result, String> { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - let sessions = AgentDao::list_sessions(&conn).map_err(|e| format!("获取会话列表失败: {e}"))?; - - Ok(sessions - .into_iter() - .map(|session| { - let messages_count = AgentDao::get_message_count(&conn, &session.id).unwrap_or(0); - build_compat_session_info(&conn, session, messages_count) - }) + .map(build_runtime_session_info) .collect()) } @@ -334,13 +380,15 @@ pub fn get_persisted_session_metadata_sync( session_id: &str, ) -> Result, String> { let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - let session = load_agent_session_record(&conn, session_id)?; + let session = agent_session_repository::get_persisted_session_metadata(&conn, session_id)?; - Ok(session.map(|session| PersistedSessionMetadata { - system_prompt: session.system_prompt, - working_dir: session.working_dir, - execution_strategy: session.execution_strategy, - })) + Ok( + session.map(|metadata: SessionRecordMetadata| PersistedSessionMetadata { + system_prompt: metadata.system_prompt, + working_dir: metadata.working_dir, + execution_strategy: metadata.execution_strategy, + }), + ) } pub fn list_title_preview_messages_sync( @@ -348,21 +396,17 @@ pub fn list_title_preview_messages_sync( session_id: &str, limit: usize, ) -> Result, String> { - if limit == 0 { - return Ok(Vec::new()); - } - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - let messages = load_agent_session_messages(&conn, session_id)?; + let messages = agent_session_repository::list_title_preview_messages(&conn, session_id, limit)?; Ok(messages .into_iter() - .filter(|msg| msg.role == "user" || msg.role == "assistant") - .take(limit) - .map(|msg| SessionTitlePreviewMessage { - role: msg.role, - content: msg.content.as_text(), - }) + .map( + |msg: SessionRecordPreviewMessage| SessionTitlePreviewMessage { + role: msg.role, + content: msg.content, + }, + ) .collect()) } @@ -370,18 +414,20 @@ pub fn list_title_preview_messages_sync( pub fn get_session_sync(db: &DbConnection, session_id: &str) -> Result { let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - let session = load_agent_session_record(&conn, session_id)? + let SessionRecordDetail { + session, + workspace_id, + } = agent_session_repository::get_session_with_messages(&conn, session_id)? .ok_or_else(|| format!("会话不存在: {session_id}"))?; - let messages = load_agent_session_messages(&conn, session_id)?; let turns = AgentTimelineDao::list_turns_by_thread(&conn, session_id) .map_err(|e| format!("获取 turn 历史失败: {e}"))?; let items = AgentTimelineDao::list_items_by_thread(&conn, session_id) .map_err(|e| format!("获取 item 历史失败: {e}"))?; let working_dir = session.working_dir.clone(); - let workspace_id = resolve_workspace_id_by_working_dir(&conn, working_dir.as_deref()); + let todo_items = load_session_todo_items_from_conn(&conn, session_id); - let tauri_messages = convert_agent_messages(&messages, Some(session.model.as_str())); + let tauri_messages = convert_agent_messages(&session.messages, Some(session.model.as_str())); tracing::debug!( "[SessionStore] 会话消息转换完成: session_id={}, messages_count={}", @@ -406,21 +452,28 @@ pub fn get_session_sync(db: &DbConnection, session_id: &str) -> Result Result { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - let session = load_agent_session_record(&conn, session_id)? - .ok_or_else(|| format!("会话不存在: {session_id}"))?; +) -> Result { + let mut detail = get_session_sync(db, session_id)?; - let messages_count = AgentDao::get_message_count(&conn, session_id).unwrap_or(0); + match load_aster_runtime_snapshot(session_id).await { + Ok(snapshot) => apply_aster_runtime_snapshot(&mut detail, &snapshot), + Err(error) => { + tracing::warn!( + "[SessionStore] 读取 Aster runtime snapshot 失败: session_id={}, error={}", + session_id, + error + ); + } + } - Ok(build_compat_session_info(&conn, session, messages_count)) + Ok(detail) } /// 重命名会话 @@ -431,12 +484,8 @@ pub fn rename_session_sync(db: &DbConnection, session_id: &str, name: &str) -> R } let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - AgentDao::update_title(&conn, session_id, trimmed_name) - .map_err(|e| format!("更新会话标题失败: {e}"))?; - let now = Utc::now().to_rfc3339(); - AgentDao::update_session_time(&conn, session_id, &now) - .map_err(|e| format!("更新会话时间失败: {e}"))?; + agent_session_repository::rename_session(&conn, session_id, trimmed_name, &now)?; Ok(()) } @@ -452,8 +501,7 @@ pub fn update_session_working_dir_sync( } let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - AgentDao::update_working_dir(&conn, session_id, trimmed_working_dir) - .map_err(|e| format!("更新 session working_dir 失败: {e}"))?; + agent_session_repository::update_session_working_dir(&conn, session_id, trimmed_working_dir)?; Ok(()) } @@ -464,8 +512,11 @@ pub fn update_session_execution_strategy_sync( execution_strategy: &str, ) -> Result<(), String> { let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - AgentDao::update_execution_strategy(&conn, session_id, execution_strategy) - .map_err(|e| format!("更新会话执行策略失败: {e}"))?; + agent_session_repository::update_session_execution_strategy( + &conn, + session_id, + execution_strategy, + )?; Ok(()) } @@ -992,6 +1043,22 @@ mod tests { assert_eq!(detail.workspace_id.as_deref(), Some("workspace-4")); } + #[test] + fn rename_session_sync_should_update_session_title() { + let db = create_test_db(); + insert_test_session_with_message( + &db, + "session-rename", + "/tmp/lime-workspace-5", + "原始消息", + ); + + rename_session_sync(&db, "session-rename", "新的会话标题").expect("rename session"); + + let session = get_session_sync(&db, "session-rename").expect("get session"); + assert_eq!(session.name, "新的会话标题"); + } + #[test] fn list_title_preview_messages_sync_should_only_keep_chat_roles() { let db = create_test_db(); diff --git a/src-tauri/crates/agent/src/skill_execution.rs b/src-tauri/crates/agent/src/skill_execution.rs new file mode 100644 index 000000000..1e2018358 --- /dev/null +++ b/src-tauri/crates/agent/src/skill_execution.rs @@ -0,0 +1,336 @@ +use crate::{ + convert_agent_event, AsterAgentState, SessionConfigBuilder, TauriAgentEvent, + WriteArtifactEventEmitter, +}; +use aster::agents::SessionConfig; +use aster::conversation::message::Message; +use futures::StreamExt; +use lime_skills::{ExecutionCallback, LoadedSkillDefinition}; +use serde::{Deserialize, Serialize}; +use std::sync::Arc; + +pub type SkillEventEmitter = Arc; + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct StepResult { + pub step_id: String, + pub step_name: String, + pub success: bool, + pub output: Option, + pub error: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SkillExecutionResult { + pub success: bool, + pub output: Option, + pub error: Option, + pub steps_completed: Vec, +} + +pub struct SkillWorkflowExecution<'a> { + pub aster_state: &'a AsterAgentState, + pub skill: &'a LoadedSkillDefinition, + pub user_input: &'a str, + pub execution_id: &'a str, + pub session_id: &'a str, + pub callback: &'a dyn ExecutionCallback, + pub memory_prompt: Option<&'a str>, + pub emitter: SkillEventEmitter, +} + +#[derive(Debug, Clone)] +pub enum SkillExecutionError { + SessionInitFailed(String), +} + +impl std::fmt::Display for SkillExecutionError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::SessionInitFailed(message) => write!(f, "{message}"), + } + } +} + +impl std::error::Error for SkillExecutionError {} + +struct StreamedSkillReply { + output: String, + error: Option, +} + +fn emit_skill_event(emitter: &SkillEventEmitter, event_name: &str, event: TauriAgentEvent) { + emitter(event_name.to_string(), event); +} + +fn build_step_system_prompt( + skill_markdown: &str, + step_name: &str, + step_number: usize, + total_steps: usize, + step_prompt: &str, + memory_prompt: Option<&str>, +) -> String { + let base_prompt = format!( + "{skill_markdown}\n\n---\n\n## 当前步骤: {step_name} ({step_number}/{total_steps})\n\n{step_prompt}" + ); + if let Some(memory_prompt) = memory_prompt { + format!("{base_prompt}\n\n{memory_prompt}") + } else { + base_prompt + } +} + +fn build_step_input(user_input: &str, accumulated_context: &str, is_first_step: bool) -> String { + if is_first_step { + accumulated_context.to_string() + } else { + format!("原始需求:{user_input}\n\n前序步骤输出:\n{accumulated_context}") + } +} + +fn build_prompt_system_prompt(skill_markdown: &str, memory_prompt: Option<&str>) -> String { + if let Some(memory_prompt) = memory_prompt { + format!("{skill_markdown}\n\n{memory_prompt}") + } else { + skill_markdown.to_string() + } +} + +async fn stream_skill_session( + aster_state: &AsterAgentState, + session_id: &str, + event_name: &str, + session_config: SessionConfig, + user_message: Message, + emitter: &SkillEventEmitter, +) -> Result { + let agent_arc = aster_state.get_agent_arc(); + let guard = agent_arc.read().await; + let agent = guard.as_ref().ok_or_else(|| { + SkillExecutionError::SessionInitFailed("Agent not initialized".to_string()) + })?; + + let cancel_token = aster_state.create_cancel_token(session_id).await; + let stream_result = agent + .reply(user_message, session_config, Some(cancel_token.clone())) + .await; + + let mut output = String::new(); + let mut error: Option = None; + let mut write_artifact_emitter = WriteArtifactEventEmitter::new(session_id.to_string()); + + match stream_result { + Ok(mut stream) => { + while let Some(event_result) = stream.next().await { + match event_result { + Ok(agent_event) => { + let tauri_events = convert_agent_event(agent_event); + for mut tauri_event in tauri_events { + let extra_events = + write_artifact_emitter.process_event(&mut tauri_event); + for extra_event in extra_events { + emit_skill_event(emitter, event_name, extra_event); + } + if let TauriAgentEvent::TextDelta { ref text } = tauri_event { + output.push_str(text); + } + emit_skill_event(emitter, event_name, tauri_event); + } + } + Err(stream_error) => { + error = Some(format!("Stream error: {stream_error}")); + break; + } + } + } + } + Err(agent_error) => { + error = Some(format!("Agent error: {agent_error}")); + } + } + + aster_state.remove_cancel_token(session_id).await; + + Ok(StreamedSkillReply { output, error }) +} + +pub async fn execute_skill_workflow( + request: SkillWorkflowExecution<'_>, +) -> Result { + let SkillWorkflowExecution { + aster_state, + skill, + user_input, + execution_id, + session_id, + callback, + memory_prompt, + emitter, + } = request; + let steps = &skill.workflow_steps; + let total_steps = steps.len(); + let event_name = format!("skill-exec-{execution_id}"); + let mut steps_completed = Vec::new(); + let mut accumulated_context = user_input.to_string(); + let mut final_output = String::new(); + + tracing::info!( + "[execute_skill_workflow] 开始 workflow 执行: steps={}, skill={}", + total_steps, + skill.skill_name + ); + + for (idx, step) in steps.iter().enumerate() { + let step_num = idx + 1; + callback.on_step_start(&step.id, &step.name, step_num, total_steps); + + tracing::info!( + "[execute_skill_workflow] 执行步骤 {}/{}: id={}, name={}", + step_num, + total_steps, + step.id, + step.name + ); + + let step_system_prompt = build_step_system_prompt( + &skill.markdown_content, + &step.name, + step_num, + total_steps, + &step.prompt, + memory_prompt, + ); + let step_session_id = format!("{session_id}-step-{}", step.id); + let session_config = SessionConfigBuilder::new(&step_session_id) + .system_prompt(step_system_prompt) + .include_context_trace(true) + .build(); + let step_input = build_step_input(user_input, &accumulated_context, idx == 0); + let user_message = Message::user().with_text(&step_input); + + let reply = stream_skill_session( + aster_state, + &step_session_id, + &event_name, + session_config, + user_message, + &emitter, + ) + .await?; + + if let Some(error) = &reply.error { + callback.on_step_error(&step.id, error, false); + steps_completed.push(StepResult { + step_id: step.id.clone(), + step_name: step.name.clone(), + success: false, + output: None, + error: Some(error.clone()), + }); + + let final_error = format!("步骤 '{}' 执行失败: {}", step.name, error); + callback.on_complete(false, None, Some(&final_error)); + emit_skill_event( + &emitter, + &event_name, + TauriAgentEvent::FinalDone { usage: None }, + ); + + return Ok(SkillExecutionResult { + success: false, + output: None, + error: Some(final_error), + steps_completed, + }); + } + + callback.on_step_complete(&step.id, &reply.output); + steps_completed.push(StepResult { + step_id: step.id.clone(), + step_name: step.name.clone(), + success: true, + output: Some(reply.output.clone()), + error: None, + }); + accumulated_context = reply.output.clone(); + final_output = reply.output; + } + + callback.on_complete(true, Some(&final_output), None); + emit_skill_event( + &emitter, + &event_name, + TauriAgentEvent::FinalDone { usage: None }, + ); + + tracing::info!( + "[execute_skill_workflow] Workflow 执行完成: skill={}, steps_completed={}", + skill.skill_name, + steps_completed.len() + ); + + Ok(SkillExecutionResult { + success: true, + output: Some(final_output), + error: None, + steps_completed, + }) +} + +pub async fn execute_skill_prompt( + aster_state: &AsterAgentState, + skill: &LoadedSkillDefinition, + user_input: &str, + execution_id: &str, + session_id: &str, + memory_prompt: Option<&str>, + emitter: SkillEventEmitter, +) -> Result { + let event_name = format!("skill-exec-{execution_id}"); + let session_config = SessionConfigBuilder::new(session_id) + .system_prompt(build_prompt_system_prompt( + &skill.markdown_content, + memory_prompt, + )) + .include_context_trace(true) + .build(); + let user_message = Message::user().with_text(user_input); + let reply = stream_skill_session( + aster_state, + session_id, + &event_name, + session_config, + user_message, + &emitter, + ) + .await?; + + if let Some(error) = reply.error { + return Ok(SkillExecutionResult { + success: false, + output: None, + error: Some(error.clone()), + steps_completed: vec![StepResult { + step_id: "main".to_string(), + step_name: skill.display_name.clone(), + success: false, + output: None, + error: Some(error), + }], + }); + } + + Ok(SkillExecutionResult { + success: true, + output: Some(reply.output.clone()), + error: None, + steps_completed: vec![StepResult { + step_id: "main".to_string(), + step_name: skill.display_name.clone(), + success: true, + output: Some(reply.output), + error: None, + }], + }) +} diff --git a/src-tauri/crates/core/src/app_paths.rs b/src-tauri/crates/core/src/app_paths.rs index 2ebede5db..cdf8eeb8d 100644 --- a/src-tauri/crates/core/src/app_paths.rs +++ b/src-tauri/crates/core/src/app_paths.rs @@ -10,6 +10,7 @@ const COMPAT_HOME_DIR_NAME: &str = ".lime"; const DATABASE_FILE_NAME: &str = "lime.db"; const LEGACY_DATABASE_FILE_NAME: &str = "proxycast.db"; const MIGRATION_MARKER_FILE: &str = ".migration_completed"; +const LEGACY_USER_MEMORY_FILE_NAMES: &[&str] = &["AGENTS.md", "AGENT.md", "instructions.md"]; const USER_SIGNAL_TABLES: &[&str] = &[ "contents", "agent_sessions", @@ -77,6 +78,10 @@ pub fn resolve_skills_dir() -> Result { resolve_runtime_subdir("skills") } +pub fn resolve_aster_dir() -> Result { + resolve_runtime_subdir("aster") +} + pub fn resolve_project_skills_dir() -> Option { std::env::current_dir() .ok() @@ -209,6 +214,18 @@ fn resolve_default_project_dir_from_roots( resolve_default_project_dir_from_source_roots(preferred_root, &[legacy_root.to_path_buf()]) } +#[cfg(test)] +fn resolve_aster_dir_from_roots( + preferred_root: &Path, + legacy_root: &Path, +) -> Result { + resolve_subdir_with_legacy_copy_from_source_roots( + preferred_root, + &[legacy_root.to_path_buf()], + "aster", + ) +} + fn resolve_user_memory_path_from_source_roots( preferred_root: &Path, legacy_roots: &[PathBuf], @@ -218,9 +235,9 @@ fn resolve_user_memory_path_from_source_roots( return Ok(preferred_path); } - let legacy_path = legacy_roots + let legacy_path = LEGACY_USER_MEMORY_FILE_NAMES .iter() - .map(|root| root.join("AGENTS.md")) + .flat_map(|file_name| legacy_roots.iter().map(move |root| root.join(file_name))) .find(|path| path.exists()); let Some(legacy_path) = legacy_path else { return Ok(preferred_path); @@ -737,6 +754,24 @@ mod tests { ); } + #[test] + fn resolve_aster_dir_copies_legacy_runtime_directories() { + let temp = tempdir().unwrap(); + let preferred_root = temp.path().join("appdata").join("lime"); + let legacy_root = temp.path().join("home").join(".lime"); + let legacy_aster_dir = legacy_root.join("aster").join("state").join("logs"); + fs::create_dir_all(&legacy_aster_dir).unwrap(); + fs::write(legacy_aster_dir.join("runtime.log"), "legacy runtime").unwrap(); + + let resolved = resolve_aster_dir_from_roots(&preferred_root, &legacy_root).unwrap(); + + assert_eq!(resolved, preferred_root.join("aster")); + assert_eq!( + fs::read_to_string(resolved.join("state").join("logs").join("runtime.log")).unwrap(), + "legacy runtime" + ); + } + #[test] fn resolve_project_skills_dir_from_cwd_builds_agents_skills_path() { let cwd = Path::new("/tmp/workspace"); @@ -759,6 +794,21 @@ mod tests { assert_eq!(fs::read_to_string(expected).unwrap(), "legacy agents"); } + #[test] + fn resolve_user_memory_path_copies_legacy_agent_file() { + let temp = tempdir().unwrap(); + let preferred_root = temp.path().join("appdata").join("lime"); + let legacy_root = temp.path().join("home").join(".lime"); + fs::create_dir_all(&legacy_root).unwrap(); + fs::write(legacy_root.join("AGENT.md"), "legacy agent").unwrap(); + + let resolved = resolve_user_memory_path_from_roots(&preferred_root, &legacy_root).unwrap(); + + let expected = preferred_root.join("AGENTS.md"); + assert_eq!(resolved, expected); + assert_eq!(fs::read_to_string(expected).unwrap(), "legacy agent"); + } + #[test] fn resolve_default_project_dir_creates_default_subdirectory() { let temp = tempdir().unwrap(); diff --git a/src-tauri/crates/core/src/config/types.rs b/src-tauri/crates/core/src/config/types.rs index cd1a834ab..85243815e 100644 --- a/src-tauri/crates/core/src/config/types.rs +++ b/src-tauri/crates/core/src/config/types.rs @@ -2404,7 +2404,7 @@ impl Default for MemorySourcesConfig { managed_policy_path: None, project_memory_paths: vec!["AGENTS.md".to_string(), ".agents/AGENTS.md".to_string()], project_rule_dirs: vec![".agents/rules".to_string()], - user_memory_path: Some("~/.lime/AGENTS.md".to_string()), + user_memory_path: None, project_local_memory_path: Some("AGENTS.local.md".to_string()), } } @@ -2422,7 +2422,7 @@ pub struct MemoryAutoConfig { /// 启动时加载 MEMORY 入口的最大行数 #[serde(default = "default_memory_auto_max_loaded_lines")] pub max_loaded_lines: u32, - /// 自动记忆根目录(可选,默认 ~/.lime/projects//memory) + /// 自动记忆根目录(可选,默认位于应用数据目录下的 projects//memory) #[serde(default, skip_serializing_if = "Option::is_none")] pub root_dir: Option, } diff --git a/src-tauri/crates/core/src/database/README.md b/src-tauri/crates/core/src/database/README.md index 62fb1db41..9a502d818 100644 --- a/src-tauri/crates/core/src/database/README.md +++ b/src-tauri/crates/core/src/database/README.md @@ -23,10 +23,9 @@ - `providers` - Provider 配置 - `settings` - 应用设置 -### Legacy 通用对话表 +### Legacy 通用对话迁移面 -- `general_chat_sessions` - 历史通用对话会话(legacy 保留) -- `general_chat_messages` - 历史通用对话消息(legacy 保留) +- `general_chat_sessions` / `general_chat_messages` - 只用于启动期识别并迁移历史安装里的旧表,不再作为新库默认 schema 的一部分 ### 功能表 diff --git a/src-tauri/crates/core/src/database/agent_runtime_queue_repository.rs b/src-tauri/crates/core/src/database/agent_runtime_queue_repository.rs new file mode 100644 index 000000000..11c2667c4 --- /dev/null +++ b/src-tauri/crates/core/src/database/agent_runtime_queue_repository.rs @@ -0,0 +1,258 @@ +//! Agent runtime queue legacy 数据访问边界。 +//! +//! 统一收口旧 `agent_runtime_queued_turns` 表的读取、校验与删除, +//! 避免上层业务边界继续散落 legacy SQL 和表名。 + +use rusqlite::Connection; +use serde_json::Value; + +const LEGACY_RUNTIME_QUEUE_TABLE: &str = "agent_runtime_queued_turns"; +const LEGACY_RUNTIME_QUEUE_SESSION_INDEX: &str = "idx_agent_runtime_queued_turns_session"; + +#[derive(Debug, Clone, PartialEq)] +pub struct LegacyRuntimeQueuedTurn { + pub queued_turn_id: String, + pub session_id: String, + pub event_name: String, + pub message_preview: String, + pub message_text: String, + pub payload: Value, + pub image_count: usize, + pub created_at: i64, +} + +#[derive(Debug, Clone, Default, PartialEq)] +pub struct LegacyRuntimeQueueSession { + pub session_id: String, + pub turns: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq)] +pub struct LegacyRuntimeQueueSnapshot { + pub sessions: Vec, + pub invalid_turn_count: usize, +} + +#[derive(Debug)] +struct LegacyRuntimeQueuedTurnRecord { + queued_turn_id: String, + session_id: String, + event_name: String, + message_preview: String, + message_text: String, + payload_json: String, + image_count: usize, + created_at: i64, +} + +pub fn load_legacy_runtime_queue_snapshot( + conn: &Connection, +) -> Result, String> { + if !legacy_runtime_queue_table_exists(conn)? { + return Ok(None); + } + + let mut stmt = conn + .prepare( + "SELECT + queued_turn_id, + session_id, + event_name, + message_preview, + message_text, + payload_json, + image_count, + created_at + FROM agent_runtime_queued_turns + ORDER BY session_id ASC, id ASC", + ) + .map_err(|error| format!("读取 legacy 排队 turn 失败: {error}"))?; + let rows = stmt + .query_map([], |row| { + Ok(LegacyRuntimeQueuedTurnRecord { + queued_turn_id: row.get(0)?, + session_id: row.get(1)?, + event_name: row.get(2)?, + message_preview: row.get(3)?, + message_text: row.get(4)?, + payload_json: row.get(5)?, + image_count: row.get::<_, i64>(6)? as usize, + created_at: row.get(7)?, + }) + }) + .map_err(|error| format!("读取 legacy 排队 turn 失败: {error}"))?; + + let mut sessions = Vec::new(); + let mut invalid_turn_count = 0usize; + + for row in rows { + let record = row.map_err(|error| format!("读取 legacy 排队 turn 失败: {error}"))?; + match serde_json::from_str::(&record.payload_json) { + Ok(payload) => { + if sessions + .last() + .map(|session: &LegacyRuntimeQueueSession| session.session_id.as_str()) + != Some(record.session_id.as_str()) + { + sessions.push(LegacyRuntimeQueueSession { + session_id: record.session_id.clone(), + turns: Vec::new(), + }); + } + + if let Some(session) = sessions.last_mut() { + session.turns.push(LegacyRuntimeQueuedTurn { + 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, + payload, + image_count: record.image_count, + created_at: record.created_at, + }); + } + } + Err(error) => { + invalid_turn_count += 1; + tracing::warn!( + "[AgentRuntimeQueueRepository] 跳过损坏的 legacy 排队 turn: session_id={}, queued_turn_id={}, error={}", + record.session_id, + record.queued_turn_id, + error + ); + } + } + } + + Ok(Some(LegacyRuntimeQueueSnapshot { + sessions, + invalid_turn_count, + })) +} + +pub fn drop_legacy_runtime_queue_table(conn: &Connection) -> Result<(), String> { + if !legacy_runtime_queue_table_exists(conn)? { + return Ok(()); + } + + conn.execute_batch( + "DROP INDEX IF EXISTS idx_agent_runtime_queued_turns_session; + DROP TABLE IF EXISTS agent_runtime_queued_turns;", + ) + .map_err(|error| format!("删除 legacy 排队表失败: {error}"))?; + + tracing::info!( + "[AgentRuntimeQueueRepository] 已删除 legacy 排队表: table={}, index={}", + LEGACY_RUNTIME_QUEUE_TABLE, + LEGACY_RUNTIME_QUEUE_SESSION_INDEX + ); + Ok(()) +} + +fn legacy_runtime_queue_table_exists(conn: &Connection) -> Result { + match conn.query_row( + "SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = ?1 LIMIT 1", + [LEGACY_RUNTIME_QUEUE_TABLE], + |_| Ok(()), + ) { + Ok(()) => Ok(true), + Err(rusqlite::Error::QueryReturnedNoRows) => Ok(false), + Err(error) => Err(format!("检测 legacy 排队表失败: {error}")), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn setup_legacy_runtime_queue_db() -> Connection { + let conn = Connection::open_in_memory().unwrap(); + conn.execute_batch( + " + CREATE TABLE agent_runtime_queued_turns ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + queued_turn_id TEXT NOT NULL, + 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 + ); + CREATE INDEX idx_agent_runtime_queued_turns_session + ON agent_runtime_queued_turns(session_id); + ", + ) + .unwrap(); + conn + } + + #[test] + fn load_legacy_runtime_queue_snapshot_groups_sessions_and_skips_invalid_payloads() { + let conn = setup_legacy_runtime_queue_db(); + 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)", + ( + "turn-1", + "session-a", + "agent_stream", + "preview-1", + "message-1", + r#"{"message":"ok"}"#, + 1i64, + 100i64, + ), + ) + .unwrap(); + 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)", + ( + "turn-2", + "session-b", + "agent_stream", + "preview-2", + "message-2", + "not-json", + 0i64, + 200i64, + ), + ) + .unwrap(); + + let snapshot = load_legacy_runtime_queue_snapshot(&conn) + .unwrap() + .expect("snapshot should exist"); + + assert_eq!(snapshot.invalid_turn_count, 1); + assert_eq!(snapshot.sessions.len(), 1); + assert_eq!(snapshot.sessions[0].session_id, "session-a"); + assert_eq!(snapshot.sessions[0].turns.len(), 1); + assert_eq!(snapshot.sessions[0].turns[0].queued_turn_id, "turn-1"); + assert_eq!( + snapshot.sessions[0].turns[0].payload, + serde_json::json!({ "message": "ok" }) + ); + } + + #[test] + fn drop_legacy_runtime_queue_table_removes_table_and_index() { + let conn = setup_legacy_runtime_queue_db(); + drop_legacy_runtime_queue_table(&conn).unwrap(); + + let table_exists = conn + .query_row( + "SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = ?1 LIMIT 1", + [LEGACY_RUNTIME_QUEUE_TABLE], + |_| Ok(()), + ) + .is_ok(); + + assert!(!table_exists); + } +} diff --git a/src-tauri/crates/core/src/database/agent_session_repository.rs b/src-tauri/crates/core/src/database/agent_session_repository.rs new file mode 100644 index 000000000..ab65e8ce2 --- /dev/null +++ b/src-tauri/crates/core/src/database/agent_session_repository.rs @@ -0,0 +1,196 @@ +//! Agent 会话持久化访问边界。 +//! +//! 统一收口 Agent session 的数据库读写与 workspace 绑定解析, +//! 避免上层 crate 继续散落 direct AgentDao 调用或手写 workspace SQL。 + +use crate::agent::types::AgentSession; +use crate::database::dao::agent::{AgentDao, AgentSessionOverviewRow}; +use rusqlite::{Connection, OptionalExtension}; + +#[derive(Debug, Clone)] +pub struct SessionRecordOverview { + pub id: String, + pub model: String, + pub system_prompt: Option, + pub title: Option, + pub created_at: String, + pub updated_at: String, + pub working_dir: Option, + pub workspace_id: Option, + pub execution_strategy: Option, + pub messages_count: usize, +} + +#[derive(Debug, Clone)] +pub struct SessionRecordDetail { + pub session: AgentSession, + pub workspace_id: Option, +} + +#[derive(Debug, Clone, Default)] +pub struct SessionRecordMetadata { + pub system_prompt: Option, + pub working_dir: Option, + pub execution_strategy: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SessionRecordPreviewMessage { + pub role: String, + pub content: String, +} + +fn resolve_workspace_id_by_working_dir( + conn: &Connection, + working_dir: Option<&str>, +) -> Option { + let resolved_working_dir = working_dir?.trim(); + if resolved_working_dir.is_empty() { + return None; + } + + match conn + .query_row( + "SELECT id FROM workspaces WHERE root_path = ? LIMIT 1", + [resolved_working_dir], + |row| row.get::<_, String>(0), + ) + .optional() + { + Ok(workspace_id) => workspace_id, + Err(error) => { + tracing::warn!( + "[AgentSessionRepository] 解析 workspace_id 失败,已降级忽略: working_dir={}, error={}", + resolved_working_dir, + error + ); + None + } + } +} + +fn map_session_overview( + conn: &Connection, + overview: AgentSessionOverviewRow, +) -> SessionRecordOverview { + let working_dir = overview.session.working_dir; + let workspace_id = resolve_workspace_id_by_working_dir(conn, working_dir.as_deref()); + + SessionRecordOverview { + id: overview.session.id, + model: overview.session.model, + system_prompt: overview.session.system_prompt, + title: overview.session.title, + created_at: overview.session.created_at, + updated_at: overview.session.updated_at, + working_dir, + workspace_id, + execution_strategy: overview.session.execution_strategy, + messages_count: overview.messages_count, + } +} + +pub fn create_session(conn: &Connection, session: &AgentSession) -> Result<(), String> { + AgentDao::create_session(conn, session).map_err(|error| format!("创建会话失败: {error}")) +} + +pub fn list_session_overviews(conn: &Connection) -> Result, String> { + AgentDao::list_session_overviews(conn) + .map(|rows| { + rows.into_iter() + .map(|row| map_session_overview(conn, row)) + .collect() + }) + .map_err(|error| format!("获取会话列表失败: {error}")) +} + +pub fn get_session_overview( + conn: &Connection, + session_id: &str, +) -> Result, String> { + AgentDao::get_session_overview(conn, session_id) + .map(|row| row.map(|row| map_session_overview(conn, row))) + .map_err(|error| format!("获取会话失败: {error}")) +} + +pub fn get_session_with_messages( + conn: &Connection, + session_id: &str, +) -> Result, String> { + AgentDao::get_session_with_messages(conn, session_id) + .map(|session| { + session.map(|session| SessionRecordDetail { + workspace_id: resolve_workspace_id_by_working_dir( + conn, + session.working_dir.as_deref(), + ), + session, + }) + }) + .map_err(|error| format!("获取会话详情失败: {error}")) +} + +pub fn get_persisted_session_metadata( + conn: &Connection, + session_id: &str, +) -> Result, String> { + get_session_overview(conn, session_id).map(|overview| { + overview.map(|overview| SessionRecordMetadata { + system_prompt: overview.system_prompt, + working_dir: overview.working_dir, + execution_strategy: overview.execution_strategy, + }) + }) +} + +pub fn list_title_preview_messages( + conn: &Connection, + session_id: &str, + limit: usize, +) -> Result, String> { + if limit == 0 { + return Ok(Vec::new()); + } + + AgentDao::get_messages(conn, session_id) + .map(|messages| { + messages + .into_iter() + .filter(|msg| msg.role == "user" || msg.role == "assistant") + .take(limit) + .map(|msg| SessionRecordPreviewMessage { + role: msg.role, + content: msg.content.as_text(), + }) + .collect() + }) + .map_err(|error| format!("获取标题预览消息失败: {error}")) +} + +pub fn rename_session( + conn: &Connection, + session_id: &str, + title: &str, + updated_at: &str, +) -> Result<(), String> { + AgentDao::rename_session(conn, session_id, title, updated_at) + .map_err(|error| format!("重命名会话失败: {error}")) +} + +pub fn update_session_working_dir( + conn: &Connection, + session_id: &str, + working_dir: &str, +) -> Result<(), String> { + AgentDao::update_working_dir(conn, session_id, working_dir) + .map_err(|error| format!("更新 session working_dir 失败: {error}")) +} + +pub fn update_session_execution_strategy( + conn: &Connection, + session_id: &str, + execution_strategy: &str, +) -> Result<(), String> { + AgentDao::update_execution_strategy(conn, session_id, execution_strategy) + .map_err(|error| format!("更新会话执行策略失败: {error}")) +} diff --git a/src-tauri/crates/core/src/database/dao/agent.rs b/src-tauri/crates/core/src/database/dao/agent.rs index 7b14b07e4..2d242b63d 100644 --- a/src-tauri/crates/core/src/database/dao/agent.rs +++ b/src-tauri/crates/core/src/database/dao/agent.rs @@ -463,6 +463,12 @@ pub struct AgentModelUsageRow { pub content_chars: u64, } +#[derive(Debug, Clone)] +pub struct AgentSessionOverviewRow { + pub session: AgentSession, + pub messages_count: usize, +} + #[derive(Debug, Clone, PartialEq, Eq)] pub struct AgentMessageTextRow { pub session_id: String, @@ -496,6 +502,20 @@ fn parse_datetime_or_timestamp_to_millis(value: &str) -> Option { }) } +fn map_agent_session_row(row: &rusqlite::Row) -> Result { + Ok(AgentSession { + id: row.get(0)?, + model: row.get(1)?, + messages: Vec::new(), + system_prompt: row.get(2)?, + title: row.get(3)?, + created_at: row.get(4)?, + updated_at: row.get(5)?, + working_dir: row.get(6)?, + execution_strategy: row.get(7)?, + }) +} + impl AgentDao { /// 创建新会话 pub fn create_session( @@ -532,17 +552,7 @@ impl AgentDao { let mut rows = stmt.query([session_id])?; if let Some(row) = rows.next()? { - Ok(Some(AgentSession { - id: row.get(0)?, - model: row.get(1)?, - messages: Vec::new(), // 消息需要单独加载 - system_prompt: row.get(2)?, - title: row.get(3)?, - created_at: row.get(4)?, - updated_at: row.get(5)?, - working_dir: row.get(6)?, - execution_strategy: row.get(7)?, - })) + Ok(Some(map_agent_session_row(row)?)) } else { Ok(None) } @@ -569,21 +579,59 @@ impl AgentDao { FROM agent_sessions ORDER BY updated_at DESC", )?; - let sessions = stmt.query_map([], |row| { - Ok(AgentSession { - id: row.get(0)?, - model: row.get(1)?, - messages: Vec::new(), - system_prompt: row.get(2)?, - title: row.get(3)?, - created_at: row.get(4)?, - updated_at: row.get(5)?, - working_dir: row.get(6)?, - execution_strategy: row.get(7)?, + let sessions = stmt.query_map([], map_agent_session_row)?; + + sessions.collect() + } + + pub fn list_session_overviews( + conn: &Connection, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT s.id, s.model, s.system_prompt, s.title, s.created_at, s.updated_at, + s.working_dir, s.execution_strategy, COUNT(m.id) AS messages_count + FROM agent_sessions s + LEFT JOIN agent_messages m ON m.session_id = s.id + GROUP BY s.id, s.model, s.system_prompt, s.title, s.created_at, s.updated_at, + s.working_dir, s.execution_strategy + ORDER BY s.updated_at DESC", + )?; + + let rows = stmt.query_map([], |row| { + let messages_count: i64 = row.get(8)?; + Ok(AgentSessionOverviewRow { + session: map_agent_session_row(row)?, + messages_count: messages_count.max(0) as usize, }) })?; - sessions.collect() + rows.collect() + } + + pub fn get_session_overview( + conn: &Connection, + session_id: &str, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT s.id, s.model, s.system_prompt, s.title, s.created_at, s.updated_at, + s.working_dir, s.execution_strategy, COUNT(m.id) AS messages_count + FROM agent_sessions s + LEFT JOIN agent_messages m ON m.session_id = s.id + WHERE s.id = ?1 + GROUP BY s.id, s.model, s.system_prompt, s.title, s.created_at, s.updated_at, + s.working_dir, s.execution_strategy", + )?; + let mut rows = stmt.query([session_id])?; + + if let Some(row) = rows.next()? { + let messages_count: i64 = row.get(8)?; + Ok(Some(AgentSessionOverviewRow { + session: map_agent_session_row(row)?, + messages_count: messages_count.max(0) as usize, + })) + } else { + Ok(None) + } } /// 获取会话的消息数量 @@ -960,6 +1008,19 @@ impl AgentDao { Ok(()) } + pub fn rename_session( + conn: &Connection, + session_id: &str, + title: &str, + updated_at: &str, + ) -> Result<(), rusqlite::Error> { + conn.execute( + "UPDATE agent_sessions SET title = ?1, updated_at = ?2 WHERE id = ?3", + params![title, updated_at, session_id], + )?; + Ok(()) + } + /// 更新会话工作目录 pub fn update_working_dir( conn: &Connection, @@ -1277,4 +1338,73 @@ mod tests { assert_eq!(general_rows[0].content, "第二条 general 消息"); assert!(general_rows[0].timestamp_ms > 0); } + + #[test] + fn session_overview_queries_and_rename_should_work() { + let conn = setup_pattern_test_db(); + + conn.execute( + "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at, working_dir, execution_strategy) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", + params![ + "session-a", + "claude-sonnet-4", + "system-a", + "标题 A", + "2026-03-10T10:00:00+08:00", + "2026-03-10T10:00:00+08:00", + "/tmp/a", + "react" + ], + ) + .unwrap(); + conn.execute( + "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at, working_dir, execution_strategy) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", + params![ + "session-b", + "gpt-4.1", + "system-b", + "标题 B", + "2026-03-11T10:00:00+08:00", + "2026-03-11T10:00:00+08:00", + "/tmp/b", + "auto" + ], + ) + .unwrap(); + conn.execute( + "INSERT INTO agent_messages (session_id, role, content_json, timestamp) VALUES (?1, ?2, ?3, ?4)", + params!["session-a", "user", r#"[{"type":"text","text":"消息一"}]"#, "2026-03-10T10:01:00+08:00"], + ) + .unwrap(); + conn.execute( + "INSERT INTO agent_messages (session_id, role, content_json, timestamp) VALUES (?1, ?2, ?3, ?4)", + params!["session-a", "assistant", r#"[{"type":"text","text":"消息二"}]"#, "2026-03-10T10:02:00+08:00"], + ) + .unwrap(); + + let overviews = AgentDao::list_session_overviews(&conn).unwrap(); + assert_eq!(overviews.len(), 2); + assert_eq!(overviews[0].session.id, "session-b"); + assert_eq!(overviews[0].messages_count, 0); + assert_eq!(overviews[1].session.id, "session-a"); + assert_eq!(overviews[1].messages_count, 2); + + let overview = AgentDao::get_session_overview(&conn, "session-a") + .unwrap() + .expect("session-a overview"); + assert_eq!(overview.session.title.as_deref(), Some("标题 A")); + assert_eq!(overview.messages_count, 2); + assert_eq!(overview.session.working_dir.as_deref(), Some("/tmp/a")); + + AgentDao::rename_session(&conn, "session-a", "新的标题", "2026-03-12T09:00:00+08:00") + .unwrap(); + + let renamed = AgentDao::get_session_overview(&conn, "session-a") + .unwrap() + .expect("renamed overview"); + assert_eq!(renamed.session.title.as_deref(), Some("新的标题")); + assert_eq!(renamed.session.updated_at, "2026-03-12T09:00:00+08:00"); + } } 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 deleted file mode 100644 index e15e61ec5..000000000 --- a/src-tauri/crates/core/src/database/dao/agent_runtime_queue.rs +++ /dev/null @@ -1,183 +0,0 @@ -//! 统一运行时排队 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 759e2c22d..bf6ceabff 100644 --- a/src-tauri/crates/core/src/database/dao/mod.rs +++ b/src-tauri/crates/core/src/database/dao/mod.rs @@ -1,7 +1,6 @@ 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 automation_job; diff --git a/src-tauri/crates/core/src/database/migration/general_chat_migration.rs b/src-tauri/crates/core/src/database/migration/general_chat_migration.rs index a79b59563..6a2ed77fc 100644 --- a/src-tauri/crates/core/src/database/migration/general_chat_migration.rs +++ b/src-tauri/crates/core/src/database/migration/general_chat_migration.rs @@ -3,6 +3,8 @@ use rusqlite::{params, Connection}; use super::{is_true_setting, mark_true_setting}; pub const GENERAL_CHAT_MIGRATION_COMPLETED_KEY: &str = "migrated_general_chat_to_unified"; +const LEGACY_GENERAL_CHAT_SESSIONS_TABLE: &str = "general_chat_sessions"; +const LEGACY_GENERAL_CHAT_MESSAGES_TABLE: &str = "general_chat_messages"; pub fn is_general_chat_migration_completed(conn: &Connection) -> bool { is_true_setting(conn, GENERAL_CHAT_MIGRATION_COMPLETED_KEY) @@ -19,28 +21,38 @@ pub fn migrate_general_chat_to_unified(conn: &Connection) -> Result 0 { + let migrated_sessions = migrate_general_sessions(conn)?; + tracing::info!("[迁移] 迁移了 {} 个会话", migrated_sessions); + migrated_sessions + } else { + 0 + }; - let migrated_messages = migrate_general_messages(conn)?; - tracing::info!("[迁移] 迁移了 {} 条消息", migrated_messages); + let migrated_messages = if general_message_count > 0 { + let migrated_messages = migrate_general_messages(conn)?; + tracing::info!("[迁移] 迁移了 {} 条消息", migrated_messages); + migrated_messages + } else { + 0 + }; mark_true_setting(conn, GENERAL_CHAT_MIGRATION_COMPLETED_KEY)?; @@ -48,6 +60,34 @@ pub fn migrate_general_chat_to_unified(conn: &Connection) -> Result Result, String> { + if !legacy_table_exists(conn, table_name)? { + return Ok(None); + } + + conn.query_row(&format!("SELECT COUNT(*) FROM {table_name}"), [], |row| { + row.get(0) + }) + .map(Some) + .map_err(|error| format!("查询 legacy 表 {table_name} 行数失败: {error}")) +} + +fn legacy_table_exists(conn: &Connection, table_name: &str) -> Result { + conn.query_row( + "SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = ?1 LIMIT 1", + [table_name], + |_| Ok(()), + ) + .map(|_| true) + .or_else(|error| match error { + rusqlite::Error::QueryReturnedNoRows => Ok(false), + _ => Err(format!("检测 legacy 表 {table_name} 是否存在失败: {error}")), + }) +} + fn migrate_general_sessions(conn: &Connection) -> Result { let mut stmt = conn .prepare( @@ -274,16 +314,14 @@ fn convert_general_content_to_json(content: &str, blocks: &Option) -> St pub fn check_general_chat_migration_status(conn: &Connection) -> GeneralChatMigrationStatus { let migrated = is_general_chat_migration_completed(conn); - let general_sessions: i64 = conn - .query_row("SELECT COUNT(*) FROM general_chat_sessions", [], |row| { - row.get(0) - }) + let general_sessions = legacy_general_chat_row_count(conn, LEGACY_GENERAL_CHAT_SESSIONS_TABLE) + .ok() + .flatten() .unwrap_or(0); - let general_messages: i64 = conn - .query_row("SELECT COUNT(*) FROM general_chat_messages", [], |row| { - row.get(0) - }) + let general_messages = legacy_general_chat_row_count(conn, LEGACY_GENERAL_CHAT_MESSAGES_TABLE) + .ok() + .flatten() .unwrap_or(0); let unified_general_sessions: i64 = conn @@ -365,6 +403,39 @@ mod tests { conn } + fn setup_unified_only_db() -> Connection { + let conn = Connection::open_in_memory().unwrap(); + conn.execute_batch( + " + CREATE TABLE settings ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL + ); + CREATE TABLE agent_sessions ( + id TEXT PRIMARY KEY, + model TEXT NOT NULL, + system_prompt TEXT, + title TEXT, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + working_dir TEXT, + execution_strategy TEXT NOT NULL DEFAULT 'react' + ); + CREATE TABLE agent_messages ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + session_id TEXT NOT NULL, + role TEXT NOT NULL, + content_json TEXT NOT NULL, + timestamp TEXT NOT NULL, + tool_calls_json TEXT, + tool_call_id TEXT + ); + ", + ) + .unwrap(); + conn + } + #[test] fn migrate_general_chat_to_unified_is_safe_to_rerun() { let conn = setup_general_chat_migration_db(); @@ -434,4 +505,19 @@ mod tests { let completed = check_general_chat_migration_status(&conn); assert!(!completed.needs_migration); } + + #[test] + fn migrate_general_chat_without_legacy_tables_marks_completed() { + let conn = setup_unified_only_db(); + + let migrated = migrate_general_chat_to_unified(&conn).unwrap(); + assert_eq!(migrated, 0); + assert!(is_general_chat_migration_completed(&conn)); + + let status = check_general_chat_migration_status(&conn); + assert_eq!(status.general_sessions_count, 0); + assert_eq!(status.general_messages_count, 0); + assert_eq!(status.migrated_sessions_count, 0); + assert!(!status.needs_migration); + } } diff --git a/src-tauri/crates/core/src/database/migration_v2.rs b/src-tauri/crates/core/src/database/migration_v2.rs index 909a05856..e82ee97d7 100644 --- a/src-tauri/crates/core/src/database/migration_v2.rs +++ b/src-tauri/crates/core/src/database/migration_v2.rs @@ -11,6 +11,7 @@ use rusqlite::{params, Connection}; use uuid::Uuid; use crate::app_paths; +use crate::workspace::WorkspaceManager; use super::migration_support::{ is_migration_completed, mark_migration_completed, run_in_transaction, @@ -113,18 +114,9 @@ fn get_or_create_default_project( where F: Fn() -> Result, { - // 检查是否已存在默认项目 - let existing_id: Option = conn - .query_row( - "SELECT id FROM workspaces WHERE is_default = 1", - [], - |row| row.get(0), - ) - .ok(); - - if let Some(id) = existing_id { - tracing::info!("[迁移] 找到现有默认项目: {}", id); - return Ok(id); + if let Some(workspace) = WorkspaceManager::get_default_from_conn(conn)? { + tracing::info!("[迁移] 找到现有默认项目: {}", workspace.id); + return Ok(workspace.id); } // 创建新的默认项目 @@ -233,14 +225,7 @@ fn verify_migration(conn: &Connection) -> Result<(), String> { )); } - // 验证默认项目存在 - let default_exists: bool = conn - .query_row( - "SELECT EXISTS(SELECT 1 FROM workspaces WHERE is_default = 1)", - [], - |row| row.get(0), - ) - .unwrap_or(false); + let default_exists = WorkspaceManager::get_default_from_conn(conn)?.is_some(); if !default_exists { return Err("迁移验证失败: 默认项目不存在".to_string()); @@ -298,12 +283,10 @@ impl MigrationResult { /// /// 如果默认项目不存在,返回 None pub fn get_default_project_id(conn: &Connection) -> Option { - conn.query_row( - "SELECT id FROM workspaces WHERE is_default = 1", - [], - |row| row.get(0), - ) - .ok() + WorkspaceManager::get_default_from_conn(conn) + .ok() + .flatten() + .map(|workspace| workspace.id) } /// 确保默认项目存在 @@ -320,6 +303,7 @@ pub fn ensure_default_project(conn: &Connection) -> Result { #[cfg(test)] mod tests { use super::*; + use crate::workspace::WorkspaceManager; use rusqlite::Connection; /// 创建测试数据库 @@ -395,13 +379,9 @@ mod tests { assert!(result.stats.is_some()); // 验证默认项目存在 - let default_exists: bool = conn - .query_row( - "SELECT EXISTS(SELECT 1 FROM workspaces WHERE is_default = 1)", - [], - |row| row.get(0), - ) - .unwrap(); + let default_exists = WorkspaceManager::get_default_from_conn(&conn) + .unwrap() + .is_some(); assert!(default_exists); } diff --git a/src-tauri/crates/core/src/database/migration_v4.rs b/src-tauri/crates/core/src/database/migration_v4.rs index 6c7052165..028d428c7 100644 --- a/src-tauri/crates/core/src/database/migration_v4.rs +++ b/src-tauri/crates/core/src/database/migration_v4.rs @@ -10,6 +10,7 @@ use rusqlite::{params, Connection}; use crate::app_paths; +use crate::workspace::WorkspaceManager; use super::migration_support::{ is_migration_completed, mark_migration_completed, run_in_transaction, @@ -176,17 +177,8 @@ fn execute_migration( /// 获取默认 workspace 的 root_path fn get_default_workspace_path(conn: &Connection) -> Result, String> { - let result = conn.query_row( - "SELECT root_path FROM workspaces WHERE is_default = 1 LIMIT 1", - [], - |row| row.get::<_, String>(0), - ); - - match result { - Ok(path) => Ok(Some(path)), - Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None), - Err(e) => Err(format!("查询默认 workspace 失败: {e}")), - } + WorkspaceManager::get_default_root_path_from_conn(conn) + .map(|path| path.map(|path| path.to_string_lossy().to_string())) } fn count_corrupted_workspaces(conn: &Connection) -> i64 { @@ -206,3 +198,39 @@ fn count_corrupted_sessions(conn: &Connection) -> i64 { ) .unwrap_or(0) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::database::schema; + + fn setup_test_db() -> Connection { + let conn = Connection::open_in_memory().expect("创建内存数据库失败"); + schema::create_tables(&conn).expect("初始化表结构失败"); + conn + } + + #[test] + fn get_default_workspace_path_should_return_none_without_default_workspace() { + let conn = setup_test_db(); + + let path = get_default_workspace_path(&conn).expect("查询默认 workspace 失败"); + + assert_eq!(path, None); + } + + #[test] + fn get_default_workspace_path_should_return_default_workspace_root() { + let conn = setup_test_db(); + conn.execute( + "INSERT INTO workspaces (id, name, workspace_type, root_path, is_default, settings_json, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, 1, '{}', 0, 0)", + params!["workspace-default", "默认项目", "general", "/tmp/lime-default"], + ) + .expect("插入默认 workspace 失败"); + + let path = get_default_workspace_path(&conn).expect("查询默认 workspace 失败"); + + assert_eq!(path.as_deref(), Some("/tmp/lime-default")); + } +} diff --git a/src-tauri/crates/core/src/database/mod.rs b/src-tauri/crates/core/src/database/mod.rs index 15a269f13..b36c65ca3 100644 --- a/src-tauri/crates/core/src/database/mod.rs +++ b/src-tauri/crates/core/src/database/mod.rs @@ -1,10 +1,11 @@ +pub mod agent_runtime_queue_repository; +pub mod agent_session_repository; pub mod dao; pub mod migration; mod migration_support; pub mod migration_v2; pub mod migration_v3; pub mod migration_v4; -mod pending_general_chat; pub mod schema; mod startup_migrations; pub mod system_providers; @@ -30,133 +31,6 @@ pub struct ConversationWindowSummary { pub content_chars: i64, } -impl ConversationWindowSummary { - pub fn merge(self, other: Self) -> Self { - Self { - session_count: self.session_count + other.session_count, - message_count: self.message_count + other.message_count, - content_chars: self.content_chars + other.content_chars, - } - } -} - -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct PendingGeneralMessage { - pub id: String, - pub session_id: String, - pub role: String, - pub content: String, - pub created_at: i64, -} - -impl From for PendingGeneralMessage { - fn from(message: pending_general_chat::PendingGeneralMessageRow) -> Self { - Self { - id: message.id, - session_id: message.session_id, - role: message.role, - content: message.content, - created_at: message.created_at, - } - } -} - -fn run_pending_general_query( - conn: &Connection, - empty_value: T, - query: F, -) -> Result -where - F: FnOnce(&Connection) -> Result, -{ - if migration::is_general_chat_migration_completed(conn) { - return Ok(empty_value); - } - - query(conn) -} - -pub fn load_pending_general_session_messages( - conn: &Connection, - session_id: &str, -) -> Result, rusqlite::Error> { - run_pending_general_query(conn, Vec::new(), |tx| { - pending_general_chat::load_pending_general_session_messages_raw(tx, session_id) - .map(|messages| messages.into_iter().map(Into::into).collect()) - }) -} - -pub fn load_pending_general_messages( - conn: &Connection, - from_timestamp_ms: Option, - to_timestamp_ms: Option, - limit: usize, -) -> Result, rusqlite::Error> { - run_pending_general_query(conn, Vec::new(), |tx| { - pending_general_chat::load_pending_general_messages_raw( - tx, - from_timestamp_ms, - to_timestamp_ms, - limit, - ) - .map(|messages| messages.into_iter().map(Into::into).collect()) - }) -} - -pub fn count_pending_general_sessions( - conn: &Connection, - from_timestamp_ms: Option, - to_timestamp_ms: Option, -) -> Result { - run_pending_general_query(conn, 0, |tx| { - pending_general_chat::count_pending_general_sessions_raw( - tx, - from_timestamp_ms, - to_timestamp_ms, - ) - }) -} - -pub fn count_pending_general_messages( - conn: &Connection, - from_timestamp_ms: Option, - to_timestamp_ms: Option, -) -> Result { - run_pending_general_query(conn, 0, |tx| { - pending_general_chat::count_pending_general_messages_raw( - tx, - from_timestamp_ms, - to_timestamp_ms, - ) - }) -} - -pub fn sum_pending_general_message_chars( - conn: &Connection, - from_timestamp_ms: Option, - to_timestamp_ms: Option, -) -> Result { - run_pending_general_query(conn, 0, |tx| { - pending_general_chat::sum_pending_general_message_chars_raw( - tx, - from_timestamp_ms, - to_timestamp_ms, - ) - }) -} - -pub fn summarize_pending_general( - conn: &Connection, - from_timestamp_ms: Option, - to_timestamp_ms: Option, -) -> Result { - Ok(ConversationWindowSummary { - session_count: count_pending_general_sessions(conn, from_timestamp_ms, to_timestamp_ms)?, - message_count: count_pending_general_messages(conn, from_timestamp_ms, to_timestamp_ms)?, - content_chars: sum_pending_general_message_chars(conn, from_timestamp_ms, to_timestamp_ms)?, - }) -} - /// 获取数据库连接锁(自动处理 poisoned lock) pub fn lock_db(db: &DbConnection) -> Result, String> { match db.lock() { @@ -201,42 +75,3 @@ pub fn init_database() -> Result { Ok(Arc::new(Mutex::new(conn))) } - -#[cfg(test)] -mod tests { - use super::*; - - fn setup_completed_general_migration_db() -> Connection { - let conn = Connection::open_in_memory().unwrap(); - conn.execute( - "CREATE TABLE settings (key TEXT PRIMARY KEY, value TEXT NOT NULL)", - [], - ) - .unwrap(); - conn.execute( - "INSERT INTO settings (key, value) VALUES (?1, 'true')", - [migration::GENERAL_CHAT_MIGRATION_COMPLETED_KEY], - ) - .unwrap(); - conn - } - - #[test] - fn pending_general_queries_short_circuit_after_migration_completed() { - let conn = setup_completed_general_migration_db(); - - let messages = load_pending_general_messages(&conn, None, None, 10).unwrap(); - let session_messages = load_pending_general_session_messages(&conn, "session-1").unwrap(); - let session_count = count_pending_general_sessions(&conn, None, None).unwrap(); - let message_count = count_pending_general_messages(&conn, None, None).unwrap(); - let char_count = sum_pending_general_message_chars(&conn, None, None).unwrap(); - let summary = summarize_pending_general(&conn, None, None).unwrap(); - - assert!(messages.is_empty()); - assert!(session_messages.is_empty()); - assert_eq!(session_count, 0); - assert_eq!(message_count, 0); - assert_eq!(char_count, 0); - assert_eq!(summary, ConversationWindowSummary::default()); - } -} diff --git a/src-tauri/crates/core/src/database/pending_general_chat.rs b/src-tauri/crates/core/src/database/pending_general_chat.rs deleted file mode 100644 index f9e29f7b0..000000000 --- a/src-tauri/crates/core/src/database/pending_general_chat.rs +++ /dev/null @@ -1,537 +0,0 @@ -use crate::general_chat::{ChatMessage, ContentBlock, MessageRole}; -use rusqlite::{params, Connection}; - -const GENERAL_MODE_PATTERN: &str = "general:%"; - -#[derive(Debug, Clone, PartialEq, Eq)] -pub(super) struct PendingGeneralMessageRow { - pub id: String, - pub session_id: String, - pub role: String, - pub content: String, - pub created_at: i64, -} - -fn get_pending_general_messages( - conn: &Connection, - session_id: &str, - limit: Option, - before_id: Option<&str>, -) -> Result, rusqlite::Error> { - if !has_pending_general_messages_table(conn)? { - return Ok(Vec::new()); - } - - let before_filter = r#" - AND ( - NOT EXISTS ( - SELECT 1 - FROM general_chat_messages before_message - WHERE before_message.session_id = ?1 - AND before_message.id = ?2 - ) - OR created_at < ( - SELECT before_message.created_at - FROM general_chat_messages before_message - WHERE before_message.session_id = ?1 - AND before_message.id = ?2 - ) - OR ( - created_at = ( - SELECT before_message.created_at - FROM general_chat_messages before_message - WHERE before_message.session_id = ?1 - AND before_message.id = ?2 - ) - AND id < ?2 - ) - ) - "#; - - let query = match (limit, before_id) { - (Some(lim), Some(_)) => { - format!( - "SELECT id, session_id, role, content, blocks, status, created_at, metadata - FROM general_chat_messages - WHERE session_id = ?1 - {before_filter} - ORDER BY created_at DESC, id DESC - LIMIT {lim}" - ) - } - (Some(lim), None) => { - format!( - "SELECT id, session_id, role, content, blocks, status, created_at, metadata - FROM general_chat_messages - WHERE session_id = ?1 - ORDER BY created_at DESC, id DESC - LIMIT {lim}" - ) - } - (None, Some(_)) => { - format!( - "SELECT id, session_id, role, content, blocks, status, created_at, metadata - FROM general_chat_messages - WHERE session_id = ?1 - {before_filter} - ORDER BY created_at ASC, id ASC" - ) - } - (None, None) => "SELECT id, session_id, role, content, blocks, status, created_at, metadata - FROM general_chat_messages - WHERE session_id = ?1 - ORDER BY created_at ASC, id ASC" - .to_string(), - }; - - let mut stmt = conn.prepare(&query)?; - let rows = if before_id.is_some() { - stmt.query_map( - params![session_id, before_id], - map_pending_general_chat_message_row, - )? - } else { - stmt.query_map(params![session_id], map_pending_general_chat_message_row)? - }; - - let mut messages = rows.collect::, _>>()?; - if limit.is_some() { - messages.reverse(); - } - - Ok(messages) -} - -fn has_pending_general_messages_table(conn: &Connection) -> Result { - table_exists(conn, "general_chat_messages") -} - -fn has_pending_general_sessions_table(conn: &Connection) -> Result { - table_exists(conn, "general_chat_sessions") -} - -pub(super) fn load_pending_general_session_messages_raw( - conn: &Connection, - session_id: &str, -) -> Result, rusqlite::Error> { - get_pending_general_messages(conn, session_id, None, None).map(|messages| { - messages - .into_iter() - .map(|message| PendingGeneralMessageRow { - id: message.id, - session_id: message.session_id, - role: stringify_message_role(&message.role).to_string(), - content: message.content, - created_at: message.created_at, - }) - .collect() - }) -} - -pub(super) fn load_pending_general_messages_raw( - conn: &Connection, - from_timestamp_ms: Option, - to_timestamp_ms: Option, - limit: usize, -) -> Result, rusqlite::Error> { - if !has_pending_general_messages_table(conn)? { - return Ok(Vec::new()); - } - - let has_agent_sessions = table_exists(conn, "agent_sessions")?; - let mut stmt = if has_agent_sessions { - conn.prepare( - "SELECT m.id, m.session_id, m.role, m.content, m.created_at - FROM general_chat_messages m - WHERE NOT EXISTS ( - SELECT 1 - FROM agent_sessions s - WHERE s.id = m.session_id - AND s.model LIKE ?1 - ) - AND (?2 IS NULL OR m.created_at >= ?2) - AND (?3 IS NULL OR m.created_at <= ?3) - ORDER BY m.created_at DESC - LIMIT ?4", - )? - } else { - conn.prepare( - "SELECT m.id, m.session_id, m.role, m.content, m.created_at - FROM general_chat_messages m - WHERE (?1 IS NULL OR m.created_at >= ?1) - AND (?2 IS NULL OR m.created_at <= ?2) - ORDER BY m.created_at DESC - LIMIT ?3", - )? - }; - - let rows = if has_agent_sessions { - stmt.query_map( - params![ - GENERAL_MODE_PATTERN, - from_timestamp_ms, - to_timestamp_ms, - limit as i64 - ], - map_pending_general_message_row, - )? - } else { - stmt.query_map( - params![from_timestamp_ms, to_timestamp_ms, limit as i64], - map_pending_general_message_row, - )? - }; - - rows.collect() -} - -pub(super) fn count_pending_general_sessions_raw( - conn: &Connection, - from_timestamp_ms: Option, - to_timestamp_ms: Option, -) -> Result { - if !has_pending_general_sessions_table(conn)? { - return Ok(0); - } - - if table_exists(conn, "agent_sessions")? { - return conn.query_row( - "SELECT COUNT(*) - FROM general_chat_sessions s - WHERE NOT EXISTS ( - SELECT 1 - FROM agent_sessions unified - WHERE unified.id = s.id - AND unified.model LIKE ?1 - ) - AND (?2 IS NULL OR s.created_at >= ?2) - AND (?3 IS NULL OR s.created_at < ?3)", - params![GENERAL_MODE_PATTERN, from_timestamp_ms, to_timestamp_ms], - |row| row.get(0), - ); - } - - conn.query_row( - "SELECT COUNT(*) - FROM general_chat_sessions s - WHERE (?1 IS NULL OR s.created_at >= ?1) - AND (?2 IS NULL OR s.created_at < ?2)", - params![from_timestamp_ms, to_timestamp_ms], - |row| row.get(0), - ) -} - -pub(super) fn count_pending_general_messages_raw( - conn: &Connection, - from_timestamp_ms: Option, - to_timestamp_ms: Option, -) -> Result { - if !has_pending_general_messages_table(conn)? { - return Ok(0); - } - - if table_exists(conn, "agent_sessions")? { - return conn.query_row( - "SELECT COUNT(*) - FROM general_chat_messages m - WHERE NOT EXISTS ( - SELECT 1 - FROM agent_sessions unified - WHERE unified.id = m.session_id - AND unified.model LIKE ?1 - ) - AND (?2 IS NULL OR m.created_at >= ?2) - AND (?3 IS NULL OR m.created_at < ?3)", - params![GENERAL_MODE_PATTERN, from_timestamp_ms, to_timestamp_ms], - |row| row.get(0), - ); - } - - conn.query_row( - "SELECT COUNT(*) - FROM general_chat_messages m - WHERE (?1 IS NULL OR m.created_at >= ?1) - AND (?2 IS NULL OR m.created_at < ?2)", - params![from_timestamp_ms, to_timestamp_ms], - |row| row.get(0), - ) -} - -pub(super) fn sum_pending_general_message_chars_raw( - conn: &Connection, - from_timestamp_ms: Option, - to_timestamp_ms: Option, -) -> Result { - if !has_pending_general_messages_table(conn)? { - return Ok(0); - } - - if table_exists(conn, "agent_sessions")? { - return conn.query_row( - "SELECT COALESCE(SUM(LENGTH(m.content)), 0) - FROM general_chat_messages m - WHERE NOT EXISTS ( - SELECT 1 - FROM agent_sessions unified - WHERE unified.id = m.session_id - AND unified.model LIKE ?1 - ) - AND (?2 IS NULL OR m.created_at >= ?2) - AND (?3 IS NULL OR m.created_at < ?3)", - params![GENERAL_MODE_PATTERN, from_timestamp_ms, to_timestamp_ms], - |row| row.get(0), - ); - } - - conn.query_row( - "SELECT COALESCE(SUM(LENGTH(m.content)), 0) - FROM general_chat_messages m - WHERE (?1 IS NULL OR m.created_at >= ?1) - AND (?2 IS NULL OR m.created_at < ?2)", - params![from_timestamp_ms, to_timestamp_ms], - |row| row.get(0), - ) -} - -fn map_pending_general_message_row( - row: &rusqlite::Row, -) -> Result { - Ok(PendingGeneralMessageRow { - id: row.get(0)?, - session_id: row.get(1)?, - role: row.get(2)?, - content: row.get(3)?, - created_at: row.get(4)?, - }) -} - -fn map_pending_general_chat_message_row( - row: &rusqlite::Row, -) -> Result { - let role_str: String = row.get(2)?; - let blocks_json: Option = row.get(4)?; - let blocks: Option> = blocks_json - .map(|json| serde_json::from_str(&json)) - .transpose() - .map_err(|e| { - rusqlite::Error::FromSqlConversionFailure(4, rusqlite::types::Type::Text, Box::new(e)) - })?; - - let metadata_json: Option = row.get(7)?; - let metadata = metadata_json - .map(|json| serde_json::from_str(&json)) - .transpose() - .map_err(|e| { - rusqlite::Error::FromSqlConversionFailure(7, rusqlite::types::Type::Text, Box::new(e)) - })?; - - Ok(ChatMessage { - id: row.get(0)?, - session_id: row.get(1)?, - role: parse_message_role(&role_str), - content: row.get(3)?, - blocks, - status: row.get(5)?, - created_at: row.get(6)?, - metadata, - }) -} - -fn parse_message_role(role: &str) -> MessageRole { - match role { - "assistant" => MessageRole::Assistant, - "system" => MessageRole::System, - _ => MessageRole::User, - } -} - -fn stringify_message_role(role: &MessageRole) -> &'static str { - match role { - MessageRole::User => "user", - MessageRole::Assistant => "assistant", - MessageRole::System => "system", - } -} - -fn table_exists(conn: &Connection, table_name: &str) -> Result { - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = ?1", - [table_name], - |row| row.get(0), - )?; - - Ok(count > 0) -} - -#[cfg(test)] -mod tests { - use super::*; - - fn create_test_schema(conn: &Connection) { - conn.execute_batch( - " - CREATE TABLE agent_sessions ( - id TEXT PRIMARY KEY, - model TEXT NOT NULL, - created_at TEXT NOT NULL, - updated_at TEXT NOT NULL - ); - CREATE TABLE general_chat_sessions ( - id TEXT PRIMARY KEY, - name TEXT NOT NULL, - created_at INTEGER NOT NULL, - updated_at INTEGER NOT NULL, - metadata TEXT - ); - CREATE TABLE general_chat_messages ( - id TEXT PRIMARY KEY, - session_id TEXT NOT NULL, - role TEXT NOT NULL, - content TEXT NOT NULL, - blocks TEXT, - status TEXT NOT NULL DEFAULT 'complete', - created_at INTEGER NOT NULL, - metadata TEXT - ); - ", - ) - .unwrap(); - } - - #[test] - fn pending_messages_support_limit_blocks_and_pagination() { - let conn = Connection::open_in_memory().unwrap(); - create_test_schema(&conn); - - let now = chrono::Utc::now().timestamp_millis(); - conn.execute( - "INSERT INTO general_chat_sessions (id, name, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", - params!["session-1", "测试会话", now, now], - ) - .unwrap(); - - conn.execute( - "INSERT INTO general_chat_messages (id, session_id, role, content, blocks, status, created_at, metadata) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", - params![ - "msg-0", - "session-1", - "assistant", - "这是一段代码:", - r#"[{"type":"code","content":"fn main() {}","language":"rust"}]"#, - "complete", - now, - Option::::None, - ], - ) - .unwrap(); - - for index in 1..=5 { - conn.execute( - "INSERT INTO general_chat_messages (id, session_id, role, content, blocks, status, created_at, metadata) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", - params![ - format!("msg-{index}"), - "session-1", - "user", - format!("消息 {index}"), - Option::::None, - "complete", - now + index as i64, - Option::::None, - ], - ) - .unwrap(); - } - - let limited = get_pending_general_messages(&conn, "session-1", Some(3), None).unwrap(); - assert_eq!(limited.len(), 3); - assert_eq!(limited[0].id, "msg-3"); - assert_eq!(limited[2].id, "msg-5"); - - let before = - get_pending_general_messages(&conn, "session-1", Some(10), Some("msg-3")).unwrap(); - let before_ids = before - .iter() - .map(|item| item.id.as_str()) - .collect::>(); - assert_eq!(before_ids, vec!["msg-0", "msg-1", "msg-2"]); - - let with_blocks = get_pending_general_messages(&conn, "session-1", None, None).unwrap(); - assert_eq!( - with_blocks[0] - .blocks - .as_ref() - .and_then(|blocks| blocks.first()) - .and_then(|block| block.language.as_deref()), - Some("rust") - ); - } - - #[test] - fn pending_messages_exclude_migrated_sessions() { - let conn = Connection::open_in_memory().unwrap(); - create_test_schema(&conn); - - conn.execute( - "INSERT INTO general_chat_sessions (id, name, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", - params!["legacy-only", "legacy", 1000i64, 1000i64], - ) - .unwrap(); - conn.execute( - "INSERT INTO general_chat_sessions (id, name, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", - params!["migrated", "migrated", 2000i64, 2000i64], - ) - .unwrap(); - conn.execute( - "INSERT INTO general_chat_messages (id, session_id, role, content, created_at) VALUES (?1, ?2, ?3, ?4, ?5)", - params!["gm-1", "legacy-only", "user", "legacy message", 1000i64], - ) - .unwrap(); - conn.execute( - "INSERT INTO general_chat_messages (id, session_id, role, content, created_at) VALUES (?1, ?2, ?3, ?4, ?5)", - params!["gm-2", "migrated", "assistant", "migrated message", 2000i64], - ) - .unwrap(); - conn.execute( - "INSERT INTO agent_sessions (id, model, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", - params![ - "migrated", - "general:default", - "2026-03-12T10:00:00+08:00", - "2026-03-12T10:00:00+08:00" - ], - ) - .unwrap(); - - let messages = load_pending_general_messages_raw(&conn, None, None, 20).unwrap(); - assert_eq!(messages.len(), 1); - assert_eq!(messages[0].session_id, "legacy-only"); - } - - #[test] - fn counters_return_zero_when_legacy_tables_are_missing() { - let conn = Connection::open_in_memory().unwrap(); - conn.execute( - "CREATE TABLE agent_sessions (id TEXT PRIMARY KEY, model TEXT NOT NULL, created_at TEXT NOT NULL, updated_at TEXT NOT NULL)", - [], - ) - .unwrap(); - - assert_eq!( - count_pending_general_sessions_raw(&conn, None, None).unwrap(), - 0 - ); - assert_eq!( - count_pending_general_messages_raw(&conn, None, None).unwrap(), - 0 - ); - assert_eq!( - sum_pending_general_message_chars_raw(&conn, None, None).unwrap(), - 0 - ); - assert!(load_pending_general_session_messages_raw(&conn, "missing") - .unwrap() - .is_empty()); - } -} diff --git a/src-tauri/crates/core/src/database/schema.rs b/src-tauri/crates/core/src/database/schema.rs index 2d7f20a33..8aea5380c 100644 --- a/src-tauri/crates/core/src/database/schema.rs +++ b/src-tauri/crates/core/src/database/schema.rs @@ -473,7 +473,21 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> { created_at TEXT NOT NULL, updated_at TEXT NOT NULL, working_dir TEXT, - execution_strategy TEXT NOT NULL DEFAULT 'react' + execution_strategy TEXT NOT NULL DEFAULT 'react', + session_type TEXT NOT NULL DEFAULT 'user', + user_set_name INTEGER NOT NULL DEFAULT 0, + extension_data_json TEXT NOT NULL DEFAULT '{}', + total_tokens INTEGER, + input_tokens INTEGER, + output_tokens INTEGER, + accumulated_total_tokens INTEGER, + accumulated_input_tokens INTEGER, + accumulated_output_tokens INTEGER, + schedule_id TEXT, + recipe_json TEXT, + user_recipe_values_json TEXT, + provider_name TEXT, + model_config_json TEXT )", [], )?; @@ -489,6 +503,56 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> { "ALTER TABLE agent_sessions ADD COLUMN execution_strategy TEXT NOT NULL DEFAULT 'react'", [], ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN session_type TEXT NOT NULL DEFAULT 'user'", + [], + ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN user_set_name INTEGER NOT NULL DEFAULT 0", + [], + ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN extension_data_json TEXT NOT NULL DEFAULT '{}'", + [], + ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN total_tokens INTEGER", + [], + ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN input_tokens INTEGER", + [], + ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN output_tokens INTEGER", + [], + ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN accumulated_total_tokens INTEGER", + [], + ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN accumulated_input_tokens INTEGER", + [], + ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN accumulated_output_tokens INTEGER", + [], + ); + let _ = conn.execute("ALTER TABLE agent_sessions ADD COLUMN schedule_id TEXT", []); + let _ = conn.execute("ALTER TABLE agent_sessions ADD COLUMN recipe_json TEXT", []); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN user_recipe_values_json TEXT", + [], + ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN provider_name TEXT", + [], + ); + let _ = conn.execute( + "ALTER TABLE agent_sessions ADD COLUMN model_config_json TEXT", + [], + ); // Agent 消息表 // 存储每个会话的消息历史 @@ -562,82 +626,6 @@ 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 相关表 - // ============================================================================ - - // 通用对话会话表 - // 存储通用对话的会话元数据 - // _Requirements: 1.6, 1.7_ - conn.execute( - "CREATE TABLE IF NOT EXISTS general_chat_sessions ( - id TEXT PRIMARY KEY, - name TEXT NOT NULL, - created_at INTEGER NOT NULL, - updated_at INTEGER NOT NULL, - metadata TEXT - )", - [], - )?; - - // 通用对话消息表 - // 存储每个会话的消息历史 - // _Requirements: 1.6, 1.7_ - conn.execute( - "CREATE TABLE IF NOT EXISTS general_chat_messages ( - id TEXT PRIMARY KEY, - session_id TEXT NOT NULL, - role TEXT NOT NULL CHECK (role IN ('user', 'assistant', 'system')), - content TEXT NOT NULL, - blocks TEXT, - status TEXT NOT NULL DEFAULT 'complete', - created_at INTEGER NOT NULL, - metadata TEXT, - FOREIGN KEY (session_id) REFERENCES general_chat_sessions(id) ON DELETE CASCADE - )", - [], - )?; - - // 创建 general_chat_messages 索引 - conn.execute( - "CREATE INDEX IF NOT EXISTS idx_messages_session_id ON general_chat_messages(session_id)", - [], - )?; - conn.execute( - "CREATE INDEX IF NOT EXISTS idx_messages_created_at ON general_chat_messages(created_at)", - [], - )?; - - // 创建 general_chat_sessions 索引 - conn.execute( - "CREATE INDEX IF NOT EXISTS idx_sessions_updated_at ON general_chat_sessions(updated_at)", - [], - )?; - // ============================================================================ // Workspace 相关表 // ============================================================================ diff --git a/src-tauri/crates/core/src/tool_calling.rs b/src-tauri/crates/core/src/tool_calling.rs index 8af82423f..9d51b0f29 100644 --- a/src-tauri/crates/core/src/tool_calling.rs +++ b/src-tauri/crates/core/src/tool_calling.rs @@ -205,7 +205,7 @@ pub fn resolve_tool_input_examples(tool_name: &str, schema: &Value) -> Vec Result, String> { let conn = self.db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; + Self::get_default_from_conn(&conn) + } + + /// 基于现有数据库连接获取默认 workspace,避免重复加锁。 + pub fn get_default_from_conn(conn: &rusqlite::Connection) -> Result, String> { let result = conn.query_row( "SELECT id, name, workspace_type, root_path, is_default, settings_json, created_at, updated_at, icon, color, is_favorite, is_archived, tags_json FROM workspaces WHERE is_default = 1", [], - |row| { - Self::row_to_workspace(row) - }, + Self::row_to_workspace, ); match result { @@ -446,6 +449,12 @@ impl WorkspaceManager { } } + pub fn get_default_root_path_from_conn( + conn: &rusqlite::Connection, + ) -> Result, String> { + Ok(Self::get_default_from_conn(conn)?.map(|workspace| workspace.root_path)) + } + /// 从数据库行解析 Workspace fn row_to_workspace(row: &rusqlite::Row) -> Result { let id: String = row.get(0)?; diff --git a/src-tauri/crates/services/src/aster_session_store.rs b/src-tauri/crates/services/src/aster_session_store.rs index 7e884d398..9f4827136 100644 --- a/src-tauri/crates/services/src/aster_session_store.rs +++ b/src-tauri/crates/services/src/aster_session_store.rs @@ -18,10 +18,12 @@ use aster::session::{ }; use async_trait::async_trait; use chrono::Utc; +use lime_core::app_paths; use lime_core::database::DbConnection; +use lime_core::workspace::WorkspaceManager; +use serde::de::DeserializeOwned; use std::collections::HashMap; -use std::fs; -use std::path::PathBuf; +use std::path::{Path, PathBuf}; /// Lime 的 SessionStore 实现 /// @@ -36,6 +38,107 @@ impl LimeSessionStore { Self { db } } + pub fn load_extension_data_from_conn( + conn: &rusqlite::Connection, + session_id: &str, + ) -> Result { + let extension_data_json: String = conn + .query_row( + "SELECT extension_data_json FROM agent_sessions WHERE id = ?1", + rusqlite::params![session_id], + |row| row.get(0), + ) + .map_err(|e| anyhow!("读取 extension_data 失败: {e}"))?; + + Ok(serde_json::from_str(&extension_data_json).unwrap_or_default()) + } + + pub fn load_extension_data_sync(db: &DbConnection, session_id: &str) -> Result { + let conn = db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; + Self::load_extension_data_from_conn(&conn, session_id) + } + + fn normalize_optional_text(value: Option) -> Option { + let value = value?; + let trimmed = value.trim(); + if trimmed.is_empty() { + None + } else { + Some(trimmed.to_string()) + } + } + + fn parse_optional_json(raw: Option) -> Option { + raw.and_then(|text| serde_json::from_str(&text).ok()) + } + + fn resolve_session_type(raw: Option, model: &str) -> SessionType { + let parsed_model = model.parse::().ok(); + match raw + .as_deref() + .and_then(|value| value.parse::().ok()) + { + Some(SessionType::User) if matches!(parsed_model, Some(parsed) if parsed != SessionType::User) => { + parsed_model.unwrap_or(SessionType::User) + } + Some(session_type) => session_type, + None => parsed_model.unwrap_or(SessionType::User), + } + } + + fn default_model_name() -> String { + "agent:default".to_string() + } + + fn insert_session_row( + conn: &rusqlite::Connection, + id: &str, + title: &str, + working_dir: &Path, + session_type: SessionType, + ) -> Result<()> { + let now = Utc::now().to_rfc3339(); + conn.execute( + "INSERT INTO agent_sessions ( + id, model, system_prompt, title, created_at, updated_at, working_dir, + execution_strategy, session_type, user_set_name, extension_data_json + ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11)", + rusqlite::params![ + id, + Self::default_model_name(), + None::, + title, + now, + now, + working_dir.to_string_lossy().to_string(), + "react", + session_type.to_string(), + false, + serde_json::to_string(&ExtensionData::default()) + .map_err(|e| anyhow!("序列化 extension_data 失败: {e}"))?, + ], + ) + .map_err(|e| anyhow!("创建会话失败: {e}"))?; + Ok(()) + } + + fn ensure_session_row(conn: &rusqlite::Connection, session_id: &str) -> Result<()> { + let session_exists: bool = conn + .query_row( + "SELECT 1 FROM agent_sessions WHERE id = ?", + [session_id], + |_| Ok(true), + ) + .unwrap_or(false); + + if session_exists { + return Ok(()); + } + + let working_dir = Self::resolve_session_working_dir(conn); + Self::insert_session_row(conn, session_id, "新对话", &working_dir, SessionType::User) + } + /// 将 Message 的 role 转换为字符串 /// 通过检查 Message::user() 和 Message::assistant() 的 role 来判断 fn message_role_to_string(message: &Message) -> String { @@ -50,38 +153,23 @@ impl LimeSessionStore { /// 解析会话 working_dir(优先默认 workspace,其次应用默认项目目录) fn resolve_session_working_dir(conn: &rusqlite::Connection) -> PathBuf { - // 1) 优先使用默认 workspace(is_default = 1) - let default_workspace_path: Option = conn - .query_row( - "SELECT root_path FROM workspaces WHERE is_default = 1 LIMIT 1", - [], - |row| row.get(0), - ) - .ok(); - - if let Some(path) = default_workspace_path { - if !path.trim().is_empty() { - let pb = PathBuf::from(path); - return if pb.is_absolute() { - pb - } else { - std::env::current_dir() - .unwrap_or_else(|_| PathBuf::from(".")) - .join(pb) - }; + if let Some(path) = WorkspaceManager::get_default_root_path_from_conn(conn) + .ok() + .flatten() + { + let normalized = Self::normalize_working_dir(path); + if !normalized.as_os_str().is_empty() { + return normalized; } } - // 2) 回退到 ~/.lime/projects/default - if let Some(home) = dirs::home_dir() { - let fallback = home.join(".lime").join("projects").join("default"); - if !fallback.exists() { - let _ = fs::create_dir_all(&fallback); - } - return fallback; + if let Ok(default_project_dir) = app_paths::resolve_default_project_dir() { + return default_project_dir; } - // 3) 最终回退到进程当前目录 + tracing::warn!( + "[SessionStore] 解析默认 working_dir 失败,已回退当前目录;建议检查 app_paths 配置" + ); std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")) } @@ -120,26 +208,8 @@ impl SessionStore for LimeSessionStore { ) -> Result { let id = uuid::Uuid::new_v4().to_string(); let now = Utc::now(); - let now_str = now.to_rfc3339(); - let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; - - let type_str = session_type.to_string(); - - conn.execute( - "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at, working_dir) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)", - rusqlite::params![ - id, - type_str, - None::, - name, - now_str, - now_str, - working_dir.to_string_lossy().to_string() - ], - ) - .map_err(|e| anyhow!("创建会话失败: {e}"))?; + Self::insert_session_row(&conn, &id, &name, &working_dir, session_type)?; Ok(Session { id, @@ -174,40 +244,17 @@ impl SessionStore for LimeSessionStore { ); let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; - - // 检查 session 是否存在 - let session_exists: bool = conn - .query_row("SELECT 1 FROM agent_sessions WHERE id = ?", [id], |_| { - Ok(true) - }) - .unwrap_or(false); - - tracing::info!("[SessionStore] session_exists={}", session_exists); - - // 如果不存在,自动创建 - if !session_exists { - let now = Utc::now().to_rfc3339(); - let working_dir = Self::resolve_session_working_dir(&conn); - conn.execute( - "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at, working_dir) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)", - rusqlite::params![ - id, - "agent:default", - None::, - "新对话", - now, - now, - working_dir.to_string_lossy().to_string() - ], - ) - .map_err(|e| anyhow!("自动创建会话失败: {e}"))?; - tracing::info!("[SessionStore] get_session 自动创建会话: {}", id); - } + Self::ensure_session_row(&conn, id)?; + tracing::info!("[SessionStore] get_session 已确保会话存在: {}", id); let mut stmt = conn .prepare( - "SELECT id, model, system_prompt, title, created_at, updated_at, working_dir + "SELECT id, model, system_prompt, title, created_at, updated_at, working_dir, + session_type, user_set_name, extension_data_json, + total_tokens, input_tokens, output_tokens, + accumulated_total_tokens, accumulated_input_tokens, accumulated_output_tokens, + schedule_id, recipe_json, user_recipe_values_json, + provider_name, model_config_json FROM agent_sessions WHERE id = ?", ) .map_err(|e| anyhow!("准备查询失败: {e}"))?; @@ -222,12 +269,47 @@ impl SessionStore for LimeSessionStore { row.get::<_, String>(4)?, row.get::<_, String>(5)?, row.get::<_, Option>(6)?, + row.get::<_, Option>(7)?, + row.get::<_, bool>(8)?, + row.get::<_, String>(9)?, + row.get::<_, Option>(10)?, + row.get::<_, Option>(11)?, + row.get::<_, Option>(12)?, + row.get::<_, Option>(13)?, + row.get::<_, Option>(14)?, + row.get::<_, Option>(15)?, + row.get::<_, Option>(16)?, + row.get::<_, Option>(17)?, + row.get::<_, Option>(18)?, + row.get::<_, Option>(19)?, + row.get::<_, Option>(20)?, )) }) .map_err(|e| anyhow!("会话不存在: {e}"))?; - let (id, model, _system_prompt, title, created_at, updated_at, db_working_dir) = - session_row; + let ( + id, + model, + _system_prompt, + title, + created_at, + updated_at, + db_working_dir, + session_type_raw, + user_set_name, + extension_data_json, + total_tokens, + input_tokens, + output_tokens, + accumulated_total_tokens, + accumulated_input_tokens, + accumulated_output_tokens, + schedule_id, + recipe_json, + user_recipe_values_json, + provider_name, + model_config_json, + ) = session_row; let created_at = chrono::DateTime::parse_from_rfc3339(&created_at) .map(|dt| dt.with_timezone(&Utc)) @@ -236,7 +318,7 @@ impl SessionStore for LimeSessionStore { .map(|dt| dt.with_timezone(&Utc)) .unwrap_or_else(|_| Utc::now()); - let session_type = model.parse().unwrap_or(SessionType::User); + let session_type = Self::resolve_session_type(session_type_raw, &model); let working_dir = Self::parse_session_working_dir(&conn, db_working_dir); let conversation = if include_messages { @@ -251,27 +333,29 @@ impl SessionStore for LimeSessionStore { id: id.to_string(), working_dir, name: title.unwrap_or_else(|| "未命名会话".to_string()), - user_set_name: false, + user_set_name, session_type, created_at, updated_at, - extension_data: ExtensionData::default(), - total_tokens: None, - input_tokens: None, - output_tokens: None, - accumulated_total_tokens: None, - accumulated_input_tokens: None, - accumulated_output_tokens: None, - schedule_id: None, - recipe: None, - user_recipe_values: None, + extension_data: serde_json::from_str(&extension_data_json).unwrap_or_default(), + total_tokens, + input_tokens, + output_tokens, + accumulated_total_tokens, + accumulated_input_tokens, + accumulated_output_tokens, + schedule_id, + recipe: Self::parse_optional_json(recipe_json), + user_recipe_values: Self::parse_optional_json(user_recipe_values_json), conversation, message_count, - provider_name: None, - model_config: match model.trim() { - "" | "agent:default" => None, - normalized => ModelConfig::new(normalized).ok(), - }, + provider_name, + model_config: Self::parse_optional_json(model_config_json).or_else(|| { + match model.trim() { + "" | "agent:default" => None, + normalized => ModelConfig::new(normalized).ok(), + } + }), }) } @@ -282,35 +366,7 @@ impl SessionStore for LimeSessionStore { ); let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; - - // 检查会话是否存在,如果不存在则自动创建 - let session_exists: bool = conn - .query_row( - "SELECT 1 FROM agent_sessions WHERE id = ?", - [session_id], - |_| Ok(true), - ) - .unwrap_or(false); - - if !session_exists { - let now = Utc::now().to_rfc3339(); - let working_dir = Self::resolve_session_working_dir(&conn); - conn.execute( - "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at, working_dir) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)", - rusqlite::params![ - session_id, - "agent:default", - None::, - "新对话", - now, - now, - working_dir.to_string_lossy().to_string() - ], - ) - .map_err(|e| anyhow!("自动创建会话失败: {e}"))?; - tracing::info!("[SessionStore] 自动创建会话: {}", session_id); - } + Self::ensure_session_row(&conn, session_id)?; let role = Self::message_role_to_string(message); let content_json = serde_json::to_string(&message.content) @@ -425,7 +481,12 @@ impl SessionStore for LimeSessionStore { let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; let mut stmt = conn.prepare( - "SELECT id, model, system_prompt, title, created_at, updated_at, working_dir + "SELECT id, model, system_prompt, title, created_at, updated_at, working_dir, + session_type, user_set_name, extension_data_json, + total_tokens, input_tokens, output_tokens, + accumulated_total_tokens, accumulated_input_tokens, accumulated_output_tokens, + schedule_id, recipe_json, user_recipe_values_json, + provider_name, model_config_json FROM agent_sessions ORDER BY updated_at DESC", )?; @@ -437,43 +498,106 @@ impl SessionStore for LimeSessionStore { let created_at: String = row.get(4)?; let updated_at: String = row.get(5)?; let working_dir: Option = row.get(6)?; + let session_type: Option = row.get(7)?; + let user_set_name: bool = row.get(8)?; + let extension_data_json: String = row.get(9)?; + let total_tokens: Option = row.get(10)?; + let input_tokens: Option = row.get(11)?; + let output_tokens: Option = row.get(12)?; + let accumulated_total_tokens: Option = row.get(13)?; + let accumulated_input_tokens: Option = row.get(14)?; + let accumulated_output_tokens: Option = row.get(15)?; + let schedule_id: Option = row.get(16)?; + let recipe_json: Option = row.get(17)?; + let user_recipe_values_json: Option = row.get(18)?; + let provider_name: Option = row.get(19)?; + let model_config_json: Option = row.get(20)?; - Ok((id, model, title, created_at, updated_at, working_dir)) + Ok(( + id, + model, + title, + created_at, + updated_at, + working_dir, + session_type, + user_set_name, + extension_data_json, + total_tokens, + input_tokens, + output_tokens, + accumulated_total_tokens, + accumulated_input_tokens, + accumulated_output_tokens, + schedule_id, + recipe_json, + user_recipe_values_json, + provider_name, + model_config_json, + )) })? .filter_map(|r| r.ok()) .map( - |(id, model, title, created_at, updated_at, db_working_dir)| { + |( + id, + model, + title, + created_at, + updated_at, + db_working_dir, + session_type_raw, + user_set_name, + extension_data_json, + total_tokens, + input_tokens, + output_tokens, + accumulated_total_tokens, + accumulated_input_tokens, + accumulated_output_tokens, + schedule_id, + recipe_json, + user_recipe_values_json, + provider_name, + model_config_json, + )| { let created_at = chrono::DateTime::parse_from_rfc3339(&created_at) .map(|dt| dt.with_timezone(&Utc)) .unwrap_or_else(|_| Utc::now()); let updated_at = chrono::DateTime::parse_from_rfc3339(&updated_at) .map(|dt| dt.with_timezone(&Utc)) .unwrap_or_else(|_| Utc::now()); - let session_type = model.parse().unwrap_or(SessionType::User); + let session_type = Self::resolve_session_type(session_type_raw, &model); let working_dir = Self::parse_session_working_dir(&conn, db_working_dir); + let message_count = self.count_messages(&conn, &id).unwrap_or(0); Session { id, working_dir, name: title.unwrap_or_else(|| "未命名会话".to_string()), - user_set_name: false, + user_set_name, session_type, created_at, updated_at, - extension_data: ExtensionData::default(), - total_tokens: None, - input_tokens: None, - output_tokens: None, - accumulated_total_tokens: None, - accumulated_input_tokens: None, - accumulated_output_tokens: None, - schedule_id: None, - recipe: None, - user_recipe_values: None, + extension_data: serde_json::from_str(&extension_data_json) + .unwrap_or_default(), + total_tokens, + input_tokens, + output_tokens, + accumulated_total_tokens, + accumulated_input_tokens, + accumulated_output_tokens, + schedule_id, + recipe: Self::parse_optional_json(recipe_json), + user_recipe_values: Self::parse_optional_json(user_recipe_values_json), conversation: None, - message_count: 0, - provider_name: None, - model_config: None, + message_count, + provider_name, + model_config: Self::parse_optional_json(model_config_json).or_else(|| { + match model.trim() { + "" | "agent:default" => None, + normalized => ModelConfig::new(normalized).ok(), + } + }), } }, ) @@ -504,10 +628,16 @@ impl SessionStore for LimeSessionStore { let total_sessions: i64 = conn.query_row("SELECT COUNT(*) FROM agent_sessions", [], |row| row.get(0))?; + let total_tokens: i64 = conn.query_row( + "SELECT COALESCE(SUM(COALESCE(accumulated_total_tokens, total_tokens, 0)), 0) + FROM agent_sessions", + [], + |row| row.get(0), + )?; Ok(SessionInsights { total_sessions: total_sessions as usize, - total_tokens: 0, + total_tokens, }) } @@ -528,6 +658,36 @@ impl SessionStore for LimeSessionStore { ) .await?; + self.update_session_name(&new_session.id, session.name.clone(), session.user_set_name) + .await?; + self.update_extension_data(&new_session.id, session.extension_data.clone()) + .await?; + self.update_token_stats( + &new_session.id, + TokenStatsUpdate { + schedule_id: session.schedule_id.clone(), + total_tokens: session.total_tokens, + input_tokens: session.input_tokens, + output_tokens: session.output_tokens, + accumulated_total: session.accumulated_total_tokens, + accumulated_input: session.accumulated_input_tokens, + accumulated_output: session.accumulated_output_tokens, + }, + ) + .await?; + self.update_provider_config( + &new_session.id, + session.provider_name.clone(), + session.model_config.clone(), + ) + .await?; + self.update_recipe( + &new_session.id, + session.recipe.clone(), + session.user_recipe_values.clone(), + ) + .await?; + if let Some(conversation) = &session.conversation { self.replace_conversation(&new_session.id, conversation) .await?; @@ -538,15 +698,47 @@ impl SessionStore for LimeSessionStore { async fn copy_session(&self, session_id: &str, new_name: String) -> Result { let original = self.get_session(session_id, true).await?; + let created_session_name = new_name.clone(); + let persisted_session_name = new_name.clone(); let new_session = self .create_session( original.working_dir.clone(), - new_name, + created_session_name, original.session_type, ) .await?; + self.update_session_name(&new_session.id, persisted_session_name, true) + .await?; + self.update_extension_data(&new_session.id, original.extension_data.clone()) + .await?; + self.update_token_stats( + &new_session.id, + TokenStatsUpdate { + schedule_id: original.schedule_id.clone(), + total_tokens: original.total_tokens, + input_tokens: original.input_tokens, + output_tokens: original.output_tokens, + accumulated_total: original.accumulated_total_tokens, + accumulated_input: original.accumulated_input_tokens, + accumulated_output: original.accumulated_output_tokens, + }, + ) + .await?; + self.update_provider_config( + &new_session.id, + original.provider_name.clone(), + original.model_config.clone(), + ) + .await?; + self.update_recipe( + &new_session.id, + original.recipe.clone(), + original.user_recipe_values.clone(), + ) + .await?; + if let Some(conversation) = &original.conversation { self.replace_conversation(&new_session.id, conversation) .await?; @@ -573,25 +765,59 @@ impl SessionStore for LimeSessionStore { &self, session_id: &str, name: String, - _user_set: bool, + user_set: bool, ) -> Result<()> { let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; + let now = Utc::now().to_rfc3339(); conn.execute( - "UPDATE agent_sessions SET title = ? WHERE id = ?", - rusqlite::params![name, session_id], + "UPDATE agent_sessions SET title = ?1, user_set_name = ?2, updated_at = ?3 WHERE id = ?4", + rusqlite::params![name, user_set, now, session_id], )?; Ok(()) } async fn update_extension_data( &self, - _session_id: &str, - _extension_data: ExtensionData, + session_id: &str, + extension_data: ExtensionData, ) -> Result<()> { + let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; + let now = Utc::now().to_rfc3339(); + let extension_data_json = serde_json::to_string(&extension_data) + .map_err(|e| anyhow!("序列化 extension_data 失败: {e}"))?; + conn.execute( + "UPDATE agent_sessions SET extension_data_json = ?1, updated_at = ?2 WHERE id = ?3", + rusqlite::params![extension_data_json, now, session_id], + )?; Ok(()) } - async fn update_token_stats(&self, _session_id: &str, _stats: TokenStatsUpdate) -> Result<()> { + async fn update_token_stats(&self, session_id: &str, stats: TokenStatsUpdate) -> Result<()> { + let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; + let now = Utc::now().to_rfc3339(); + conn.execute( + "UPDATE agent_sessions SET + total_tokens = COALESCE(?1, total_tokens), + input_tokens = COALESCE(?2, input_tokens), + output_tokens = COALESCE(?3, output_tokens), + accumulated_total_tokens = COALESCE(?4, accumulated_total_tokens), + accumulated_input_tokens = COALESCE(?5, accumulated_input_tokens), + accumulated_output_tokens = COALESCE(?6, accumulated_output_tokens), + schedule_id = COALESCE(?7, schedule_id), + updated_at = ?8 + WHERE id = ?9", + rusqlite::params![ + stats.total_tokens, + stats.input_tokens, + stats.output_tokens, + stats.accumulated_total, + stats.accumulated_input, + stats.accumulated_output, + Self::normalize_optional_text(stats.schedule_id), + now, + session_id, + ], + )?; Ok(()) } @@ -601,27 +827,67 @@ impl SessionStore for LimeSessionStore { provider_name: Option, model_config: Option, ) -> Result<()> { - if let Some(model_name) = model_config + let normalized_provider_name = Self::normalize_optional_text(provider_name); + let normalized_model_name = model_config .as_ref() .map(|config| config.model_name.trim().to_string()) - .filter(|value| !value.is_empty()) - .or(provider_name.filter(|value| !value.trim().is_empty())) - { - let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; - conn.execute( - "UPDATE agent_sessions SET model = ? WHERE id = ?", - rusqlite::params![model_name, session_id], - )?; + .filter(|value| !value.is_empty()); + let model_config_json = model_config + .as_ref() + .map(serde_json::to_string) + .transpose() + .map_err(|e| anyhow!("序列化 model_config 失败: {e}"))?; + + if normalized_provider_name.is_none() && normalized_model_name.is_none() { + return Ok(()); } + + let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; + let now = Utc::now().to_rfc3339(); + conn.execute( + "UPDATE agent_sessions SET + provider_name = COALESCE(?1, provider_name), + model = COALESCE(?2, model), + model_config_json = CASE WHEN ?3 IS NULL THEN model_config_json ELSE ?3 END, + updated_at = ?4 + WHERE id = ?5", + rusqlite::params![ + normalized_provider_name, + normalized_model_name, + model_config_json, + now, + session_id, + ], + )?; Ok(()) } async fn update_recipe( &self, - _session_id: &str, - _recipe: Option, - _user_recipe_values: Option>, + session_id: &str, + recipe: Option, + user_recipe_values: Option>, ) -> Result<()> { + let conn = self.db.lock().map_err(|e| anyhow!("数据库锁定失败: {e}"))?; + let now = Utc::now().to_rfc3339(); + let recipe_json = recipe + .as_ref() + .map(serde_json::to_string) + .transpose() + .map_err(|e| anyhow!("序列化 recipe 失败: {e}"))?; + let user_recipe_values_json = user_recipe_values + .as_ref() + .map(serde_json::to_string) + .transpose() + .map_err(|e| anyhow!("序列化 user_recipe_values 失败: {e}"))?; + conn.execute( + "UPDATE agent_sessions SET + recipe_json = ?1, + user_recipe_values_json = ?2, + updated_at = ?3 + WHERE id = ?4", + rusqlite::params![recipe_json, user_recipe_values_json, now, session_id], + )?; Ok(()) } @@ -788,7 +1054,41 @@ mod tests { use aster::session::{SessionStore, SessionType}; use lime_core::database::schema::create_tables; use rusqlite::Connection; + use std::ffi::OsString; use std::sync::{Arc, Mutex}; + use tempfile::tempdir; + + fn env_lock() -> &'static Mutex<()> { + static LOCK: std::sync::OnceLock> = std::sync::OnceLock::new(); + LOCK.get_or_init(|| Mutex::new(())) + } + + struct EnvGuard { + values: Vec<(&'static str, Option)>, + } + + impl EnvGuard { + fn set(entries: &[(&'static str, OsString)]) -> Self { + let mut values = Vec::new(); + for (key, value) in entries { + values.push((*key, std::env::var_os(key))); + std::env::set_var(key, value); + } + Self { values } + } + } + + impl Drop for EnvGuard { + fn drop(&mut self) { + for (key, previous) in self.values.drain(..) { + if let Some(value) = previous { + std::env::set_var(key, value); + } else { + std::env::remove_var(key); + } + } + } + } fn setup_test_store() -> LimeSessionStore { let conn = Connection::open_in_memory().expect("创建内存数据库失败"); @@ -828,4 +1128,175 @@ mod tests { assert_eq!(persisted_model, "gpt-4.1"); } + + #[tokio::test] + async fn get_session_should_prefer_default_workspace_root_when_missing_row() { + let store = setup_test_store(); + let workspace_root = std::env::temp_dir().join("lime-aster-default-workspace"); + let conn = store.db.lock().expect("锁数据库"); + conn.execute( + "INSERT INTO workspaces (id, name, workspace_type, root_path, is_default, settings_json, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, 1, '{}', 0, 0)", + rusqlite::params![ + "workspace-default", + "默认工作区", + "general", + workspace_root.to_string_lossy().to_string(), + ], + ) + .expect("插入默认 workspace 失败"); + drop(conn); + + let session = store + .get_session("missing-default-workspace-session", false) + .await + .expect("读取缺失会话失败"); + + assert_eq!(session.working_dir, workspace_root); + } + + #[tokio::test] + async fn get_session_should_fallback_to_app_paths_default_project_dir() { + let _env_guard = env_lock().lock().expect("锁环境变量"); + let temp = tempdir().expect("创建临时目录失败"); + let home = temp.path().join("home"); + let app_data = temp.path().join("appdata"); + std::fs::create_dir_all(&home).expect("创建 home 目录失败"); + std::fs::create_dir_all(&app_data).expect("创建 appdata 目录失败"); + let _guard = EnvGuard::set(&[ + ("HOME", home.as_os_str().to_os_string()), + ("XDG_DATA_HOME", app_data.as_os_str().to_os_string()), + ("APPDATA", app_data.as_os_str().to_os_string()), + ("LOCALAPPDATA", app_data.as_os_str().to_os_string()), + ]); + + let store = setup_test_store(); + let session = store + .get_session("missing-fallback-session", false) + .await + .expect("读取缺失会话失败"); + let expected = app_paths::resolve_default_project_dir().expect("解析默认项目目录失败"); + + assert_eq!(session.working_dir, expected); + assert!(session.working_dir.is_absolute()); + assert!(session + .working_dir + .ends_with(PathBuf::from("projects").join("default"))); + assert!(!session + .working_dir + .to_string_lossy() + .contains(".lime/projects/default")); + } + + #[tokio::test] + async fn update_session_metadata_should_roundtrip() { + let store = setup_test_store(); + let session = store + .create_session( + PathBuf::from("."), + "元数据测试".to_string(), + SessionType::SubAgent, + ) + .await + .expect("创建会话失败"); + + let mut extension_data = ExtensionData::new(); + extension_data.set_extension_state("todo", "v0", serde_json::json!({"items":["a"]})); + + store + .update_session_name(&session.id, "已命名会话".to_string(), true) + .await + .expect("更新名称失败"); + store + .update_extension_data(&session.id, extension_data.clone()) + .await + .expect("更新 extension_data 失败"); + store + .update_token_stats( + &session.id, + TokenStatsUpdate { + schedule_id: Some("job-1".to_string()), + total_tokens: Some(100), + input_tokens: Some(60), + output_tokens: Some(40), + accumulated_total: Some(300), + accumulated_input: Some(180), + accumulated_output: Some(120), + }, + ) + .await + .expect("更新 token 统计失败"); + store + .update_provider_config( + &session.id, + Some("openai".to_string()), + Some(ModelConfig::new("gpt-4.1").expect("model config")), + ) + .await + .expect("更新 provider 配置失败"); + store + .update_recipe( + &session.id, + Some(Recipe { + version: "1.0.0".to_string(), + title: "demo".to_string(), + description: "demo recipe".to_string(), + instructions: None, + prompt: None, + extensions: None, + settings: None, + activities: None, + author: None, + parameters: None, + response: None, + sub_recipes: None, + retry: None, + }), + Some(HashMap::from([( + "temperature".to_string(), + "0.2".to_string(), + )])), + ) + .await + .expect("更新 recipe 失败"); + + let loaded = store + .get_session(&session.id, false) + .await + .expect("读取会话失败"); + + assert_eq!(loaded.name, "已命名会话"); + assert!(loaded.user_set_name); + assert_eq!(loaded.session_type, SessionType::SubAgent); + assert_eq!(loaded.total_tokens, Some(100)); + assert_eq!(loaded.accumulated_total_tokens, Some(300)); + assert_eq!(loaded.schedule_id.as_deref(), Some("job-1")); + assert_eq!(loaded.provider_name.as_deref(), Some("openai")); + assert_eq!( + loaded + .model_config + .as_ref() + .map(|config| config.model_name.as_str()), + Some("gpt-4.1") + ); + assert_eq!( + loaded + .extension_data + .get_extension_state("todo", "v0") + .cloned(), + extension_data.get_extension_state("todo", "v0").cloned() + ); + assert_eq!( + loaded.recipe.as_ref().map(|recipe| recipe.title.as_str()), + Some("demo") + ); + assert_eq!( + loaded + .user_recipe_values + .as_ref() + .and_then(|values| values.get("temperature")) + .map(String::as_str), + Some("0.2") + ); + } } diff --git a/src-tauri/crates/services/src/lib.rs b/src-tauri/crates/services/src/lib.rs index 76c0c21ce..175daa9ea 100644 --- a/src-tauri/crates/services/src/lib.rs +++ b/src-tauri/crates/services/src/lib.rs @@ -43,7 +43,6 @@ //! - `session_context_service` - 会话上下文服务 //! - `ai_summary_service` - AI 摘要服务 //! - `project_context_builder` - 项目上下文构建器 -//! - `tool_hooks_service` - 工具钩子服务 //! - `kiro_event_service` - Kiro 事件服务 //! - `api_key_provider_service` - API Key Provider 服务 //! - `provider_pool_service` - Provider 池服务 @@ -90,7 +89,6 @@ pub mod content_creator; pub mod ai_summary_service; pub mod project_context_builder; pub mod session_context_service; -pub mod tool_hooks_service; // 事件服务 pub mod kiro_event_service; diff --git a/src-tauri/crates/services/src/session_context_service.rs b/src-tauri/crates/services/src/session_context_service.rs index 810d61412..5fa355957 100644 --- a/src-tauri/crates/services/src/session_context_service.rs +++ b/src-tauri/crates/services/src/session_context_service.rs @@ -4,7 +4,6 @@ use crate::ai_summary_service::AISummaryService; use lime_core::database::dao::chat::{ChatDao, ChatMessage as UnifiedChatMessage, ChatMode}; -use lime_core::database::load_pending_general_session_messages; use lime_core::general_chat::{ChatMessage, MessageRole}; use rusqlite::Connection; use serde::{Deserialize, Serialize}; @@ -184,52 +183,27 @@ impl SessionContextService { let session = ChatDao::get_session(conn, session_id).map_err(|e| format!("获取统一会话失败: {e}"))?; - if let Some(session) = session { - if session.mode != ChatMode::General { - debug!( - "会话 {} 不是 general 模式,跳过上下文加载: {:?}", - session_id, session.mode - ); - return Ok(vec![]); - } + let Some(session) = session else { + debug!("会话 {} 未命中 unified chat,跳过上下文加载", session_id); + return Ok(vec![]); + }; - let unified_messages = ChatDao::get_messages(conn, session_id, None) - .map_err(|e| format!("获取统一消息失败: {e}"))?; - - return Ok(unified_messages - .into_iter() - .map(Self::convert_unified_message) - .collect()); + if session.mode != ChatMode::General { + debug!( + "会话 {} 不是 general 模式,跳过上下文加载: {:?}", + session_id, session.mode + ); + return Ok(vec![]); } - debug!( - "会话 {} 未命中 unified chat,回退待迁移 general 历史表", - session_id - ); - - Self::load_pending_general_messages(conn, session_id) - } - - fn load_pending_general_messages( - conn: &Connection, - session_id: &str, - ) -> Result, String> { - let messages = load_pending_general_session_messages(conn, session_id) - .map_err(|e| format!("查询待迁移 general 消息失败: {e}"))?; - - Ok(messages - .into_iter() - .map(|message| ChatMessage { - id: message.id, - session_id: message.session_id, - role: Self::convert_legacy_role(&message.role), - content: message.content, - blocks: None, - status: "complete".to_string(), - created_at: message.created_at, - metadata: None, + ChatDao::get_messages(conn, session_id, None) + .map_err(|e| format!("获取统一消息失败: {e}")) + .map(|messages| { + messages + .into_iter() + .map(Self::convert_unified_message) + .collect() }) - .collect()) } fn convert_unified_message(message: UnifiedChatMessage) -> ChatMessage { @@ -253,14 +227,6 @@ impl SessionContextService { } } - fn convert_legacy_role(role: &str) -> MessageRole { - match role { - "assistant" => MessageRole::Assistant, - "system" => MessageRole::System, - _ => MessageRole::User, - } - } - fn extract_unified_text_content(content: &serde_json::Value) -> String { match content { serde_json::Value::Null => String::new(), @@ -604,7 +570,6 @@ mod tests { use lime_core::database::dao::chat::{ ChatDao, ChatMessage as UnifiedChatMessage, ChatMode, ChatSession as UnifiedChatSession, }; - use lime_core::general_chat::ChatSession; use rusqlite::Connection; fn setup_test_db() -> Connection { @@ -649,36 +614,6 @@ mod tests { ) .unwrap(); - // 创建会话表 - conn.execute( - "CREATE TABLE general_chat_sessions ( - id TEXT PRIMARY KEY, - name TEXT NOT NULL, - created_at INTEGER NOT NULL, - updated_at INTEGER NOT NULL, - metadata TEXT - )", - [], - ) - .unwrap(); - - // 创建消息表 - conn.execute( - "CREATE TABLE general_chat_messages ( - id TEXT PRIMARY KEY, - session_id TEXT NOT NULL, - role TEXT NOT NULL CHECK (role IN ('user', 'assistant', 'system')), - content TEXT NOT NULL, - blocks TEXT, - status TEXT NOT NULL DEFAULT 'complete', - created_at INTEGER NOT NULL, - metadata TEXT, - FOREIGN KEY (session_id) REFERENCES general_chat_sessions(id) ON DELETE CASCADE - )", - [], - ) - .unwrap(); - conn.execute("PRAGMA foreign_keys = ON", []).unwrap(); conn } @@ -749,45 +684,6 @@ mod tests { .collect() } - fn insert_legacy_session(conn: &Connection, session: &ChatSession) { - conn.execute( - "INSERT INTO general_chat_sessions (id, name, created_at, updated_at, metadata) - VALUES (?1, ?2, ?3, ?4, ?5)", - rusqlite::params![ - session.id, - session.name, - session.created_at, - session.updated_at, - Option::::None, - ], - ) - .unwrap(); - } - - fn insert_legacy_message(conn: &Connection, message: &ChatMessage) { - let role = match message.role { - MessageRole::User => "user", - MessageRole::Assistant => "assistant", - MessageRole::System => "system", - }; - - conn.execute( - "INSERT INTO general_chat_messages (id, session_id, role, content, blocks, status, created_at, metadata) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)", - rusqlite::params![ - message.id, - message.session_id, - role, - message.content, - Option::::None, - message.status, - message.created_at, - Option::::None, - ], - ) - .unwrap(); - } - #[tokio::test] async fn test_context_service_creation() { let conn = Arc::new(Mutex::new(setup_test_db())); @@ -880,47 +776,22 @@ mod tests { } #[tokio::test] - async fn test_get_effective_context_falls_back_to_legacy_messages() { + async fn test_get_effective_context_returns_empty_when_unified_session_missing() { let conn = Arc::new(Mutex::new(setup_test_db())); let service = SessionContextService::new(conn.clone(), ContextWindowConfig::default()); - let legacy_session = ChatSession { - id: "legacy-session".to_string(), - name: "旧会话".to_string(), - created_at: chrono::Utc::now().timestamp_millis(), - updated_at: chrono::Utc::now().timestamp_millis(), - metadata: None, - }; - - { - let conn_guard = conn.lock().unwrap(); - insert_legacy_session(&conn_guard, &legacy_session); - for message in create_test_messages("legacy-session", 3) { - insert_legacy_message(&conn_guard, &message); - } - } - let context = service .get_effective_context("legacy-session") .await .unwrap(); - assert_eq!(context.len(), 3); - assert_eq!(context[0].content, "这是第 1 条消息,包含一些测试内容"); + assert!(context.is_empty()); } #[tokio::test] - async fn test_get_effective_context_skips_legacy_fallback_after_general_migration_completed() { + async fn test_get_effective_context_keeps_empty_when_general_migration_completed() { let conn = Arc::new(Mutex::new(setup_test_db())); let service = SessionContextService::new(conn.clone(), ContextWindowConfig::default()); - let legacy_session = ChatSession { - id: "legacy-session".to_string(), - name: "旧会话".to_string(), - created_at: chrono::Utc::now().timestamp_millis(), - updated_at: chrono::Utc::now().timestamp_millis(), - metadata: None, - }; - { let conn_guard = conn.lock().unwrap(); conn_guard @@ -929,10 +800,6 @@ mod tests { rusqlite::params!["migrated_general_chat_to_unified", "true"], ) .unwrap(); - insert_legacy_session(&conn_guard, &legacy_session); - for message in create_test_messages("legacy-session", 3) { - insert_legacy_message(&conn_guard, &message); - } } let context = service diff --git a/src-tauri/crates/services/src/tool_hooks_service.rs b/src-tauri/crates/services/src/tool_hooks_service.rs deleted file mode 100644 index 50db3870c..000000000 --- a/src-tauri/crates/services/src/tool_hooks_service.rs +++ /dev/null @@ -1,715 +0,0 @@ -//! 工具钩子管理服务 -//! -//! 提供工具执行前后的钩子机制,用于自动化上下文记忆管理 - -use crate::context_memory_service::{ContextMemoryService, MemoryEntry, MemoryFileType}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::sync::{Arc, Mutex}; -use tracing::{debug, error, info}; - -/// 钩子触发时机 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -pub enum HookTrigger { - /// 会话开始时 - SessionStart, - /// 工具使用前 - PreToolUse, - /// 工具使用后 - PostToolUse, - /// 会话停止时 - Stop, -} - -/// 钩子动作类型 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -pub enum HookAction { - /// 保存发现 - SaveFinding { - title: String, - content: String, - tags: Vec, - priority: u8, - }, - /// 更新任务计划 - UpdateTaskPlan { - title: String, - content: String, - priority: u8, - }, - /// 记录进度 - LogProgress { title: String, content: String }, - /// 记录错误 - RecordError { - error_description: String, - attempted_solution: String, - }, - /// 自定义动作 - Custom { - action_type: String, - parameters: HashMap, - }, -} - -/// 钩子规则 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HookRule { - /// 规则 ID - pub id: String, - /// 规则名称 - pub name: String, - /// 描述 - pub description: String, - /// 触发时机 - pub trigger: HookTrigger, - /// 触发条件 - pub conditions: Vec, - /// 执行动作 - pub actions: Vec, - /// 是否启用 - pub enabled: bool, - /// 优先级 (数字越小优先级越高) - pub priority: u32, - /// 创建时间 - pub created_at: i64, -} - -/// 钩子条件 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -pub enum HookCondition { - /// 工具名称匹配 - ToolNameEquals(String), - /// 工具名称包含 - ToolNameContains(String), - /// 消息内容包含 - MessageContains(String), - /// 会话消息数量大于 - MessageCountGreaterThan(usize), - /// 错误次数大于 - ErrorCountGreaterThan(u32), - /// 自定义条件 - Custom { - condition_type: String, - parameters: HashMap, - }, -} - -/// 钩子执行上下文 -#[derive(Debug, Clone)] -pub struct HookContext { - /// 会话 ID - pub session_id: String, - /// 工具名称(如果适用) - pub tool_name: Option, - /// 工具参数(如果适用) - pub tool_parameters: Option>, - /// 工具结果(如果适用) - pub tool_result: Option, - /// 消息内容 - pub message_content: Option, - /// 会话消息数量 - pub message_count: usize, - /// 错误信息(如果适用) - pub error_info: Option, - /// 额外元数据 - pub metadata: HashMap, -} - -/// 工具钩子管理器 -pub struct ToolHooksService { - /// 钩子规则 - rules: Arc>>, - /// 上下文记忆服务 - memory_service: Arc, - /// 执行统计 - execution_stats: Arc>>, -} - -/// 钩子执行统计 -#[derive(Debug, Clone, Default, Serialize, Deserialize)] -pub struct HookExecutionStats { - /// 执行次数 - pub execution_count: u64, - /// 成功次数 - pub success_count: u64, - /// 失败次数 - pub failure_count: u64, - /// 最后执行时间 - pub last_execution_at: i64, - /// 平均执行时间(毫秒) - pub average_execution_time_ms: f64, -} - -impl ToolHooksService { - /// 创建新的工具钩子服务 - pub fn new(memory_service: Arc) -> Self { - let service = Self { - rules: Arc::new(Mutex::new(Vec::new())), - memory_service, - execution_stats: Arc::new(Mutex::new(HashMap::new())), - }; - - // 注册默认钩子规则 - service.register_default_hooks(); - - service - } - - /// 注册默认钩子规则 - fn register_default_hooks(&self) { - let default_rules = vec![ - // 会话开始时创建任务计划 - HookRule { - id: "session-start-task-plan".to_string(), - name: "会话开始任务计划".to_string(), - description: "会话开始时自动创建任务计划记录".to_string(), - trigger: HookTrigger::SessionStart, - conditions: vec![], - actions: vec![HookAction::UpdateTaskPlan { - title: "会话任务计划".to_string(), - content: "新会话开始,等待用户输入任务目标...".to_string(), - priority: 3, - }], - enabled: true, - priority: 1, - created_at: chrono::Utc::now().timestamp_millis(), - }, - // 工具使用后记录进度(2-Action 规则) - HookRule { - id: "post-tool-progress-log".to_string(), - name: "工具使用进度记录".to_string(), - description: "工具使用后自动记录进度(2-Action 规则)".to_string(), - trigger: HookTrigger::PostToolUse, - conditions: vec![], - actions: vec![HookAction::LogProgress { - title: "工具执行记录".to_string(), - content: "工具执行完成,结果已记录".to_string(), - }], - enabled: true, - priority: 2, - created_at: chrono::Utc::now().timestamp_millis(), - }, - // 重要发现自动保存 - HookRule { - id: "important-finding-save".to_string(), - name: "重要发现保存".to_string(), - description: "检测到重要信息时自动保存".to_string(), - trigger: HookTrigger::PostToolUse, - conditions: vec![ - HookCondition::MessageContains("重要".to_string()), - HookCondition::MessageContains("发现".to_string()), - ], - actions: vec![HookAction::SaveFinding { - title: "重要发现".to_string(), - content: "检测到重要信息,已自动保存".to_string(), - tags: vec!["重要".to_string(), "自动保存".to_string()], - priority: 4, - }], - enabled: true, - priority: 3, - created_at: chrono::Utc::now().timestamp_millis(), - }, - // 错误自动记录 - HookRule { - id: "error-auto-record".to_string(), - name: "错误自动记录".to_string(), - description: "检测到错误时自动记录".to_string(), - trigger: HookTrigger::PostToolUse, - conditions: vec![HookCondition::MessageContains("错误".to_string())], - actions: vec![HookAction::RecordError { - error_description: "检测到错误".to_string(), - attempted_solution: "正在尝试解决".to_string(), - }], - enabled: true, - priority: 1, - created_at: chrono::Utc::now().timestamp_millis(), - }, - // 会话停止时保存摘要 - HookRule { - id: "session-stop-summary".to_string(), - name: "会话停止摘要".to_string(), - description: "会话停止时保存会话摘要".to_string(), - trigger: HookTrigger::Stop, - conditions: vec![HookCondition::MessageCountGreaterThan(5)], - actions: vec![HookAction::SaveFinding { - title: "会话摘要".to_string(), - content: "会话已结束,主要成果和发现已记录".to_string(), - tags: vec!["摘要".to_string(), "会话结束".to_string()], - priority: 3, - }], - enabled: true, - priority: 2, - created_at: chrono::Utc::now().timestamp_millis(), - }, - ]; - - let mut rules = self.rules.lock().unwrap(); - rules.extend(default_rules); - info!("已注册 {} 个默认钩子规则", rules.len()); - } - - /// 执行钩子 - pub fn execute_hooks(&self, trigger: HookTrigger, context: &HookContext) -> Result<(), String> { - let rules = self.rules.lock().map_err(|e| e.to_string())?; - - // 获取匹配的规则并按优先级排序 - let mut matching_rules: Vec<_> = rules - .iter() - .filter(|rule| rule.enabled && rule.trigger == trigger) - .filter(|rule| self.evaluate_conditions(rule, context)) - .collect(); - - matching_rules.sort_by_key(|rule| rule.priority); - - debug!( - "触发钩子 {:?},匹配到 {} 个规则 (会话: {})", - trigger, - matching_rules.len(), - context.session_id - ); - - // 执行匹配的规则 - for rule in matching_rules { - if let Err(e) = self.execute_rule(rule, context) { - error!("执行钩子规则失败 {}: {}", rule.name, e); - self.update_execution_stats(&rule.id, false, 0.0); - } else { - self.update_execution_stats(&rule.id, true, 0.0); - } - } - - Ok(()) - } - - /// 评估钩子条件 - fn evaluate_conditions(&self, rule: &HookRule, context: &HookContext) -> bool { - if rule.conditions.is_empty() { - return true; - } - - for condition in &rule.conditions { - if !self.evaluate_single_condition(condition, context) { - return false; - } - } - - true - } - - /// 评估单个条件 - fn evaluate_single_condition(&self, condition: &HookCondition, context: &HookContext) -> bool { - match condition { - HookCondition::ToolNameEquals(name) => context.tool_name.as_ref() == Some(name), - HookCondition::ToolNameContains(substring) => context - .tool_name - .as_ref() - .is_some_and(|tn| tn.contains(substring)), - HookCondition::MessageContains(substring) => { - context - .message_content - .as_ref() - .is_some_and(|mc| mc.contains(substring)) - || context - .tool_result - .as_ref() - .is_some_and(|tr| tr.contains(substring)) - } - HookCondition::MessageCountGreaterThan(count) => context.message_count > *count, - HookCondition::ErrorCountGreaterThan(_count) => { - // 这里可以从 memory_service 获取错误计数 - context.error_info.is_some() - } - HookCondition::Custom { - condition_type: _, - parameters: _, - } => { - // 自定义条件的实现 - true - } - } - } - - /// 执行钩子规则 - fn execute_rule(&self, rule: &HookRule, context: &HookContext) -> Result<(), String> { - let start_time = std::time::Instant::now(); - - for action in &rule.actions { - self.execute_action(action, context)?; - } - - let execution_time = start_time.elapsed().as_millis() as f64; - self.update_execution_stats(&rule.id, true, execution_time); - - debug!( - "执行钩子规则成功: {} (耗时: {:.2}ms)", - rule.name, execution_time - ); - Ok(()) - } - - /// 执行钩子动作 - fn execute_action(&self, action: &HookAction, context: &HookContext) -> Result<(), String> { - match action { - HookAction::SaveFinding { - title, - content, - tags, - priority, - } => { - let entry = MemoryEntry { - id: uuid::Uuid::new_v4().to_string(), - session_id: context.session_id.clone(), - file_type: MemoryFileType::Findings, - title: self.interpolate_template(title, context), - content: self.interpolate_template(content, context), - tags: tags.clone(), - priority: *priority, - created_at: chrono::Utc::now().timestamp_millis(), - updated_at: chrono::Utc::now().timestamp_millis(), - archived: false, - }; - self.memory_service.save_memory_entry(&entry)?; - } - - HookAction::UpdateTaskPlan { - title, - content, - priority, - } => { - let entry = MemoryEntry { - id: uuid::Uuid::new_v4().to_string(), - session_id: context.session_id.clone(), - file_type: MemoryFileType::TaskPlan, - title: self.interpolate_template(title, context), - content: self.interpolate_template(content, context), - tags: vec!["任务计划".to_string()], - priority: *priority, - created_at: chrono::Utc::now().timestamp_millis(), - updated_at: chrono::Utc::now().timestamp_millis(), - archived: false, - }; - self.memory_service.save_memory_entry(&entry)?; - } - - HookAction::LogProgress { title, content } => { - let entry = MemoryEntry { - id: uuid::Uuid::new_v4().to_string(), - session_id: context.session_id.clone(), - file_type: MemoryFileType::Progress, - title: self.interpolate_template(title, context), - content: self.interpolate_template(content, context), - tags: vec!["进度".to_string()], - priority: 2, - created_at: chrono::Utc::now().timestamp_millis(), - updated_at: chrono::Utc::now().timestamp_millis(), - archived: false, - }; - self.memory_service.save_memory_entry(&entry)?; - } - - HookAction::RecordError { - error_description, - attempted_solution, - } => { - let error_desc = self.interpolate_template(error_description, context); - let solution = self.interpolate_template(attempted_solution, context); - self.memory_service - .record_error(&context.session_id, &error_desc, &solution)?; - } - - HookAction::Custom { - action_type: _, - parameters: _, - } => { - // 自定义动作的实现 - debug!("执行自定义钩子动作"); - } - } - - Ok(()) - } - - /// 模板插值 - fn interpolate_template(&self, template: &str, context: &HookContext) -> String { - let mut result = template.to_string(); - - // 替换常见的模板变量 - result = result.replace("{session_id}", &context.session_id); - - if let Some(tool_name) = &context.tool_name { - result = result.replace("{tool_name}", tool_name); - } - - if let Some(message_content) = &context.message_content { - let preview = if message_content.len() > 100 { - format!("{}...", &message_content[..100]) - } else { - message_content.clone() - }; - result = result.replace("{message_preview}", &preview); - } - - if let Some(tool_result) = &context.tool_result { - let preview = if tool_result.len() > 200 { - format!("{}...", &tool_result[..200]) - } else { - tool_result.clone() - }; - result = result.replace("{tool_result_preview}", &preview); - } - - result = result.replace("{message_count}", &context.message_count.to_string()); - result = result.replace( - "{timestamp}", - &chrono::Utc::now().format("%Y-%m-%d %H:%M:%S").to_string(), - ); - - // 替换元数据变量 - for (key, value) in &context.metadata { - result = result.replace(&format!("{{{key}}}"), value); - } - - result - } - - /// 更新执行统计 - fn update_execution_stats(&self, rule_id: &str, success: bool, execution_time_ms: f64) { - let mut stats = self.execution_stats.lock().unwrap(); - let entry = stats.entry(rule_id.to_string()).or_default(); - - entry.execution_count += 1; - if success { - entry.success_count += 1; - } else { - entry.failure_count += 1; - } - entry.last_execution_at = chrono::Utc::now().timestamp_millis(); - - // 更新平均执行时间 - if execution_time_ms > 0.0 { - let total_time = entry.average_execution_time_ms * (entry.execution_count - 1) as f64; - entry.average_execution_time_ms = - (total_time + execution_time_ms) / entry.execution_count as f64; - } - } - - /// 添加钩子规则 - pub fn add_hook_rule(&self, rule: HookRule) -> Result<(), String> { - let mut rules = self.rules.lock().map_err(|e| e.to_string())?; - - // 检查是否已存在相同 ID 的规则 - if rules.iter().any(|r| r.id == rule.id) { - return Err(format!("钩子规则 ID 已存在: {}", rule.id)); - } - - rules.push(rule.clone()); - info!("已添加钩子规则: {}", rule.name); - Ok(()) - } - - /// 移除钩子规则 - pub fn remove_hook_rule(&self, rule_id: &str) -> Result<(), String> { - let mut rules = self.rules.lock().map_err(|e| e.to_string())?; - - let initial_len = rules.len(); - rules.retain(|r| r.id != rule_id); - - if rules.len() == initial_len { - return Err(format!("未找到钩子规则: {rule_id}")); - } - - info!("已移除钩子规则: {}", rule_id); - Ok(()) - } - - /// 启用/禁用钩子规则 - pub fn toggle_hook_rule(&self, rule_id: &str, enabled: bool) -> Result<(), String> { - let mut rules = self.rules.lock().map_err(|e| e.to_string())?; - - if let Some(rule) = rules.iter_mut().find(|r| r.id == rule_id) { - rule.enabled = enabled; - info!( - "钩子规则 {} 已{}", - rule.name, - if enabled { "启用" } else { "禁用" } - ); - Ok(()) - } else { - Err(format!("未找到钩子规则: {rule_id}")) - } - } - - /// 获取所有钩子规则 - pub fn get_hook_rules(&self) -> Result, String> { - let rules = self.rules.lock().map_err(|e| e.to_string())?; - Ok(rules.clone()) - } - - /// 获取执行统计 - pub fn get_execution_stats(&self) -> Result, String> { - let stats = self.execution_stats.lock().map_err(|e| e.to_string())?; - Ok(stats.clone()) - } - - /// 清理执行统计 - pub fn clear_execution_stats(&self) -> Result<(), String> { - let mut stats = self.execution_stats.lock().map_err(|e| e.to_string())?; - stats.clear(); - info!("已清理钩子执行统计"); - Ok(()) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::context_memory_service::ContextMemoryConfig; - use tempfile::TempDir; - - fn create_test_services() -> (Arc, ToolHooksService, TempDir) { - let temp_dir = TempDir::new().unwrap(); - let memory_config = ContextMemoryConfig { - memory_dir: temp_dir.path().to_path_buf(), - ..Default::default() - }; - - let memory_service = Arc::new(ContextMemoryService::new(memory_config).unwrap()); - let hooks_service = ToolHooksService::new(memory_service.clone()); - - (memory_service, hooks_service, temp_dir) - } - - #[test] - fn test_hooks_service_creation() { - let (_memory_service, hooks_service, _temp_dir) = create_test_services(); - - let rules = hooks_service.get_hook_rules().unwrap(); - assert!(!rules.is_empty()); // 应该有默认规则 - } - - #[test] - fn test_session_start_hook() { - let (_memory_service, hooks_service, _temp_dir) = create_test_services(); - - let context = HookContext { - session_id: "test-session".to_string(), - tool_name: None, - tool_parameters: None, - tool_result: None, - message_content: None, - message_count: 0, - error_info: None, - metadata: HashMap::new(), - }; - - hooks_service - .execute_hooks(HookTrigger::SessionStart, &context) - .unwrap(); - - // 验证任务计划是否被创建 - let memories = _memory_service - .get_session_memories("test-session", Some(MemoryFileType::TaskPlan)) - .unwrap(); - assert!(!memories.is_empty()); - } - - #[test] - fn test_error_recording_hook() { - let (_memory_service, hooks_service, _temp_dir) = create_test_services(); - - let context = HookContext { - session_id: "test-session".to_string(), - tool_name: Some("test_tool".to_string()), - tool_parameters: None, - tool_result: Some("发生了一个错误".to_string()), - message_content: Some("这里有一个错误需要处理".to_string()), - message_count: 5, - error_info: Some("测试错误".to_string()), - metadata: HashMap::new(), - }; - - hooks_service - .execute_hooks(HookTrigger::PostToolUse, &context) - .unwrap(); - - // 验证错误是否被记录 - let stats = _memory_service.get_memory_stats("test-session").unwrap(); - assert!(stats.unresolved_errors > 0); - } - - #[test] - fn test_custom_hook_rule() { - let (_memory_service, hooks_service, _temp_dir) = create_test_services(); - - let custom_rule = HookRule { - id: "custom-test-rule".to_string(), - name: "自定义测试规则".to_string(), - description: "测试自定义钩子规则".to_string(), - trigger: HookTrigger::PostToolUse, - conditions: vec![HookCondition::ToolNameEquals("custom_tool".to_string())], - actions: vec![HookAction::SaveFinding { - title: "自定义发现".to_string(), - content: "这是一个自定义钩子触发的发现".to_string(), - tags: vec!["自定义".to_string()], - priority: 3, - }], - enabled: true, - priority: 1, - created_at: chrono::Utc::now().timestamp_millis(), - }; - - hooks_service.add_hook_rule(custom_rule).unwrap(); - - let context = HookContext { - session_id: "test-session".to_string(), - tool_name: Some("custom_tool".to_string()), - tool_parameters: None, - tool_result: None, - message_content: None, - message_count: 0, - error_info: None, - metadata: HashMap::new(), - }; - - hooks_service - .execute_hooks(HookTrigger::PostToolUse, &context) - .unwrap(); - - // 验证自定义发现是否被保存 - let memories = _memory_service - .get_session_memories("test-session", Some(MemoryFileType::Findings)) - .unwrap(); - assert!(memories.iter().any(|m| m.title == "自定义发现")); - } - - #[test] - fn test_template_interpolation() { - let (_memory_service, hooks_service, _temp_dir) = create_test_services(); - - let context = HookContext { - session_id: "test-session-123".to_string(), - tool_name: Some("test_tool".to_string()), - tool_parameters: None, - tool_result: None, - message_content: Some("这是测试消息".to_string()), - message_count: 42, - error_info: None, - metadata: { - let mut map = HashMap::new(); - map.insert("custom_var".to_string(), "custom_value".to_string()); - map - }, - }; - - let template = "会话 {session_id} 使用工具 {tool_name},消息数量: {message_count},自定义变量: {custom_var}"; - let result = hooks_service.interpolate_template(template, &context); - - assert!(result.contains("test-session-123")); - assert!(result.contains("test_tool")); - assert!(result.contains("42")); - assert!(result.contains("custom_value")); - } -} diff --git a/src-tauri/src/agent/aster_agent.rs b/src-tauri/src/agent/aster_agent.rs index 8d4f40249..3fdd22abc 100644 --- a/src-tauri/src/agent/aster_agent.rs +++ b/src-tauri/src/agent/aster_agent.rs @@ -10,7 +10,7 @@ use futures::StreamExt; use lime_agent::{convert_agent_event, TauriAgentEvent, WriteArtifactEventEmitter}; use tauri::{AppHandle, Emitter}; -pub use lime_agent::session_store::{ +pub use lime_agent::{ PersistedSessionMetadata, SessionDetail, SessionInfo, SessionTitlePreviewMessage, }; @@ -121,30 +121,31 @@ impl AsterAgentWrapper { workspace_id: String, execution_strategy: Option, ) -> Result { - lime_agent::session_store::create_session_sync( - db, - name, - working_dir, - workspace_id, - execution_strategy, - ) + lime_agent::create_session_sync(db, name, working_dir, workspace_id, execution_strategy) } /// 列出所有会话 pub fn list_sessions_sync(db: &DbConnection) -> Result, String> { - lime_agent::session_store::list_sessions_sync(db) + lime_agent::list_sessions_sync(db) } /// 获取会话详情 pub fn get_session_sync(db: &DbConnection, session_id: &str) -> Result { - lime_agent::session_store::get_session_sync(db, session_id) + lime_agent::get_session_sync(db, session_id) + } + + pub async fn get_runtime_session_detail( + db: &DbConnection, + session_id: &str, + ) -> Result { + lime_agent::get_runtime_session_detail(db, session_id).await } pub fn get_persisted_session_metadata_sync( db: &DbConnection, session_id: &str, ) -> Result, String> { - lime_agent::session_store::get_persisted_session_metadata_sync(db, session_id) + lime_agent::get_persisted_session_metadata_sync(db, session_id) } pub fn list_title_preview_messages_sync( @@ -152,7 +153,7 @@ impl AsterAgentWrapper { session_id: &str, limit: usize, ) -> Result, String> { - lime_agent::session_store::list_title_preview_messages_sync(db, session_id, limit) + lime_agent::list_title_preview_messages_sync(db, session_id, limit) } /// 重命名会话 @@ -161,7 +162,7 @@ impl AsterAgentWrapper { session_id: &str, name: &str, ) -> Result<(), String> { - lime_agent::session_store::rename_session_sync(db, session_id, name) + lime_agent::rename_session_sync(db, session_id, name) } pub fn update_session_working_dir_sync( @@ -169,7 +170,7 @@ impl AsterAgentWrapper { session_id: &str, working_dir: &str, ) -> Result<(), String> { - lime_agent::session_store::update_session_working_dir_sync(db, session_id, working_dir) + lime_agent::update_session_working_dir_sync(db, session_id, working_dir) } pub fn update_session_execution_strategy_sync( @@ -177,16 +178,12 @@ impl AsterAgentWrapper { session_id: &str, execution_strategy: &str, ) -> Result<(), String> { - lime_agent::session_store::update_session_execution_strategy_sync( - db, - session_id, - execution_strategy, - ) + lime_agent::update_session_execution_strategy_sync(db, session_id, execution_strategy) } /// 删除会话 pub async fn delete_session(db: &DbConnection, session_id: &str) -> Result<(), String> { - lime_agent::session_store::delete_session(db, session_id).await + lime_agent::delete_session(db, session_id).await } } diff --git a/src-tauri/src/agent/integration.rs b/src-tauri/src/agent/integration.rs new file mode 100644 index 000000000..359c226a7 --- /dev/null +++ b/src-tauri/src/agent/integration.rs @@ -0,0 +1,4 @@ +//! 旧 agent integration 壳层已退出编译图。 +//! +//! Aster runtime 的启动初始化与全局 session store 注册 +//! 已统一收口到 `lime_agent::initialize_aster_runtime`。 diff --git a/src-tauri/src/agent/mod.rs b/src-tauri/src/agent/mod.rs index 27469a81d..c978f1307 100644 --- a/src-tauri/src/agent/mod.rs +++ b/src-tauri/src/agent/mod.rs @@ -3,10 +3,11 @@ //! 纯逻辑部分已迁移到 lime-agent crate, //! 本模块保留深耦合部分(Aster 状态与 Tauri 桥接)。 -pub mod aster_agent; +mod aster_agent; pub mod aster_state; -pub mod credential_bridge; -pub mod subagent_scheduler; +mod credential_bridge; +pub mod runtime_queue_service; +mod subagent_scheduler; // 从 lime-agent crate re-export pub use lime_agent::event_converter; @@ -23,8 +24,8 @@ pub use credential_bridge::{ create_aster_provider, AsterProviderConfig, CredentialBridge, CredentialBridgeError, }; pub use lime_agent::{ - convert_agent_event, convert_to_tauri_message, QueueInsertResult, QueuedTurnSnapshot, - QueuedTurnTask, SessionTurnQueueManager, TauriAgentEvent, + convert_agent_event, convert_to_tauri_message, initialize_aster_runtime, QueuedTurnSnapshot, + QueuedTurnTask, TauriAgentEvent, }; pub use subagent_scheduler::{ LimeScheduler, LimeSubAgentExecutor, SubAgentProgressEvent, SubAgentRole, diff --git a/src-tauri/src/agent/runtime_queue_service.rs b/src-tauri/src/agent/runtime_queue_service.rs new file mode 100644 index 000000000..72ad464ac --- /dev/null +++ b/src-tauri/src/agent/runtime_queue_service.rs @@ -0,0 +1,209 @@ +//! Agent runtime queue 共享服务边界。 +//! +//! 命令层只保留 Tauri 状态装配; +//! queue 的纯调度与数据事实源统一委托给 `lime-agent`。 + +use super::aster_state::AsterAgentState; +use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState; +use crate::config::GlobalConfigManagerState; +use crate::database::DbConnection; +use crate::mcp::McpManagerState; +use crate::services::automation_service::AutomationServiceState; +use crate::LogState; +use aster::session::QueuedTurnRuntime; +use lime_agent::{ + clear_runtime_queue as clear_runtime_queue_impl, + list_runtime_queue_snapshots as list_runtime_queue_snapshots_impl, + remove_runtime_queued_turn as remove_runtime_queued_turn_impl, + resume_persisted_runtime_queues_on_startup as resume_persisted_runtime_queues_on_startup_impl, + resume_runtime_queue_if_needed as resume_runtime_queue_if_needed_impl, + submit_runtime_turn as submit_runtime_turn_impl, QueuedTurnSnapshot, QueuedTurnTask, + RuntimeQueueEventEmitter, RuntimeQueueExecutor as SharedRuntimeQueueExecutor, TauriAgentEvent, +}; +use serde_json::Value; +use tauri::{AppHandle, Emitter}; + +pub(crate) type RuntimeQueueExecutor = SharedRuntimeQueueExecutor; + +pub(crate) struct AgentRuntimeQueueContext { + pub(crate) app: AppHandle, + pub(crate) state: AsterAgentState, + pub(crate) db: DbConnection, + pub(crate) api_key_provider_service: ApiKeyProviderServiceState, + pub(crate) logs: LogState, + pub(crate) config_manager: GlobalConfigManagerState, + pub(crate) mcp_manager: McpManagerState, + pub(crate) automation_state: AutomationServiceState, +} + +impl Clone for AgentRuntimeQueueContext { + 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(), + automation_state: self.automation_state.clone(), + } + } +} + +fn build_runtime_queue_context( + app: AppHandle, + state: &AsterAgentState, + db: &DbConnection, + api_key_provider_service: &ApiKeyProviderServiceState, + logs: &LogState, + config_manager: &GlobalConfigManagerState, + mcp_manager: &McpManagerState, + automation_state: &AutomationServiceState, +) -> AgentRuntimeQueueContext { + AgentRuntimeQueueContext { + 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(), + automation_state: automation_state.clone(), + } +} + +fn build_runtime_queue_event_emitter(app: &AppHandle) -> RuntimeQueueEventEmitter { + let app = app.clone(); + std::sync::Arc::new(move |event_name: String, event: TauriAgentEvent| { + if let Err(error) = app.emit(&event_name, &event) { + tracing::warn!( + "[AsterAgent][Queue] 发送队列事件失败: event_name={}, error={}", + event_name, + error + ); + } + }) +} + +pub(crate) async fn resume_runtime_queue_if_needed( + app: AppHandle, + state: &AsterAgentState, + db: &DbConnection, + api_key_provider_service: &ApiKeyProviderServiceState, + logs: &LogState, + config_manager: &GlobalConfigManagerState, + mcp_manager: &McpManagerState, + automation_state: &AutomationServiceState, + session_id: String, + executor: RuntimeQueueExecutor, +) -> Result { + let context = build_runtime_queue_context( + app, + state, + db, + api_key_provider_service, + logs, + config_manager, + mcp_manager, + automation_state, + ); + resume_runtime_queue_if_needed_impl( + session_id, + context.clone(), + executor, + build_runtime_queue_event_emitter(&context.app), + ) + .await +} + +pub(crate) async fn submit_runtime_turn( + app: AppHandle, + state: &AsterAgentState, + db: &DbConnection, + api_key_provider_service: &ApiKeyProviderServiceState, + logs: &LogState, + config_manager: &GlobalConfigManagerState, + mcp_manager: &McpManagerState, + automation_state: &AutomationServiceState, + queued_task: QueuedTurnTask, + queue_if_busy: bool, + executor: RuntimeQueueExecutor, +) -> Result<(), String> { + let context = build_runtime_queue_context( + app, + state, + db, + api_key_provider_service, + logs, + config_manager, + mcp_manager, + automation_state, + ); + submit_runtime_turn_impl( + queued_task, + queue_if_busy, + context.clone(), + executor, + build_runtime_queue_event_emitter(&context.app), + ) + .await +} + +pub(crate) async fn clear_runtime_queue( + app: &AppHandle, + session_id: &str, +) -> Result, String> { + clear_runtime_queue_impl(session_id, build_runtime_queue_event_emitter(app)).await +} + +pub(crate) async fn list_runtime_queue_snapshots( + session_id: &str, +) -> Result, String> { + list_runtime_queue_snapshots_impl(session_id).await +} + +pub(crate) async fn remove_runtime_queued_turn( + app: &AppHandle, + session_id: &str, + queued_turn_id: &str, +) -> Result { + remove_runtime_queued_turn_impl( + session_id, + queued_turn_id, + build_runtime_queue_event_emitter(app), + ) + .await +} + +pub(crate) async 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, + automation_state: &AutomationServiceState, + executor: RuntimeQueueExecutor, +) -> Result { + let context = build_runtime_queue_context( + app.clone(), + state, + db, + api_key_provider_service, + logs, + config_manager, + mcp_manager, + automation_state, + ); + + resume_persisted_runtime_queues_on_startup_impl( + context, + executor, + build_runtime_queue_event_emitter(&app), + ) + .await +} diff --git a/src-tauri/src/app/bootstrap.rs b/src-tauri/src/app/bootstrap.rs index 549bc76b0..9d890e6c1 100644 --- a/src-tauri/src/app/bootstrap.rs +++ b/src-tauri/src/app/bootstrap.rs @@ -5,7 +5,7 @@ use std::sync::Arc; use tokio::sync::RwLock; -use crate::agent::AsterAgentState; +use crate::agent::{initialize_aster_runtime, AsterAgentState}; use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState; use crate::commands::connect_cmd::ConnectStateWrapper; use crate::commands::context_memory::ContextMemoryServiceState; @@ -19,7 +19,6 @@ use crate::commands::resilience_cmd::ResilienceConfigState; use crate::commands::session_files_cmd::SessionFilesState; use crate::commands::skill_cmd::SkillServiceState; use crate::commands::terminal_cmd::TerminalManagerState; -use crate::commands::tool_hooks::ToolHooksServiceState; use crate::commands::webview_cmd::{ ChromeProfileManagerWrapper, WebviewManagerState, WebviewManagerWrapper, }; @@ -36,12 +35,10 @@ use lime_core::config::{Config, ConfigManager}; use lime_scheduler::AgentScheduler; use lime_server as server; use lime_services::api_key_provider_service::ApiKeyProviderService; -use lime_services::aster_session_store::LimeSessionStore; use lime_services::context_memory_service::{ContextMemoryConfig, ContextMemoryService}; use lime_services::provider_pool_service::ProviderPoolService; use lime_services::skill_service::SkillService; use lime_services::token_cache_service::TokenCacheService; -use lime_services::tool_hooks_service::ToolHooksService; use lime_services::update_check_service::UpdateCheckServiceState; use super::types::{AppState, LogState, TokenCacheServiceState}; @@ -75,7 +72,6 @@ pub struct AppStates { pub update_check_service: UpdateCheckServiceState, pub session_files: SessionFilesState, pub context_memory_service: ContextMemoryServiceState, - pub tool_hooks_service: ToolHooksServiceState, pub recording_service: RecordingServiceState, pub mcp_manager: McpManagerState, pub automation_service: AutomationServiceState, @@ -131,6 +127,8 @@ pub fn init_states(config: &Config) -> Result { } } + initialize_aster_runtime(db.clone()).map_err(|e| format!("Aster 运行时初始化失败: {e}"))?; + // 服务状态 let skill_service = SkillService::new().map_err(|e| format!("SkillService 初始化失败: {e}"))?; let skill_service_state = SkillServiceState(Arc::new(skill_service)); @@ -181,44 +179,6 @@ pub fn init_states(config: &Config) -> Result { let (telemetry_state, shared_stats, shared_tokens, shared_logger) = init_telemetry(config)?; // 其他状态 - // 设置 Aster 全局 session store(使用 Lime 数据库) - let session_store = Arc::new(LimeSessionStore::new(db.clone())); - - // 使用 tokio runtime 来设置全局 store - // 使用 Builder 模式以获得更好的跨平台兼容性 - let rt = tokio::runtime::Handle::try_current().unwrap_or_else(|_| { - // Windows: IOCP, macOS: kqueue, Linux: epoll/io-uring - #[cfg(target_os = "windows")] - tracing::info!("[Bootstrap] Windows 平台 - 创建 Tokio Runtime (IOCP)"); - - #[cfg(target_os = "macos")] - tracing::info!("[Bootstrap] macOS 平台 - 创建 Tokio Runtime (kqueue)"); - - #[cfg(target_os = "linux")] - tracing::info!("[Bootstrap] Linux 平台 - 创建 Tokio Runtime (epoll)"); - - // 使用 Builder 模式获得更多控制,提高 Windows 兼容性 - tokio::runtime::Builder::new_multi_thread() - .worker_threads(2) // 限制线程数,避免 Windows 资源问题 - .thread_name("lime-runtime") - .enable_io() - .enable_time() - .build() - .expect("Failed to create tokio runtime: 系统资源不足或配置错误") - .handle() - .clone() - }); - rt.block_on(async { - if let Err(e) = aster::session::set_global_session_store(session_store).await { - tracing::warn!( - "[Bootstrap] 设置全局 session store 失败(可能已设置): {}", - e - ); - } else { - tracing::info!("[Bootstrap] 已设置 Aster 全局 session store"); - } - }); - let aster_agent_state = AsterAgentState::new(); let orchestrator_state = OrchestratorState::new(); @@ -282,13 +242,7 @@ pub fn init_states(config: &Config) -> Result { let context_memory_config = build_context_memory_config(config); let context_memory_service = ContextMemoryService::new(context_memory_config) .map_err(|e| format!("ContextMemoryService 初始化失败: {e}"))?; - let context_memory_service_arc = Arc::new(context_memory_service); - let context_memory_service_state = - ContextMemoryServiceState(context_memory_service_arc.clone()); - - // 初始化工具钩子服务 - let tool_hooks_service = ToolHooksService::new(context_memory_service_arc.clone()); - let tool_hooks_service_state = ToolHooksServiceState(Arc::new(tool_hooks_service)); + let context_memory_service_state = ContextMemoryServiceState(Arc::new(context_memory_service)); // 录音服务(使用独立线程 + channel 通信解决 cpal::Stream 不是 Send 的问题) let recording_service_state = create_recording_service_state(); @@ -339,7 +293,6 @@ pub fn init_states(config: &Config) -> Result { update_check_service: update_check_service_state, session_files: session_files_state, context_memory_service: context_memory_service_state, - tool_hooks_service: tool_hooks_service_state, recording_service: recording_service_state, mcp_manager: mcp_manager_state, automation_service: automation_service_state, diff --git a/src-tauri/src/app/mod.rs b/src-tauri/src/app/mod.rs index c718879d0..519bbb748 100644 --- a/src-tauri/src/app/mod.rs +++ b/src-tauri/src/app/mod.rs @@ -5,24 +5,21 @@ //! ## 模块结构 //! - `types` - 核心类型定义(ProviderType 等) //! - `state` - 状态类型和初始化 -//! - `setup` - Tauri setup hook //! - `commands` - 内置 Tauri 命令 //! - `utils` - 辅助函数 //! - `bootstrap` - 应用启动引导(配置验证、状态初始化) -//! - `runner` - 应用运行器(Tauri Builder 配置和命令注册) +//! - `runner` - 应用运行器(Tauri Builder 配置、setup 和命令注册) pub mod bootstrap; pub mod commands; pub mod runner; pub mod scheduler_service; -mod setup; mod state; mod types; mod utils; pub use runner::run; pub use scheduler_service::{SchedulerService, SchedulerServiceConfig}; -pub use setup::setup_app; pub use state::*; pub use types::*; pub use utils::*; diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index 613063740..d856d0639 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -21,6 +21,19 @@ fn should_minimize_to_tray(window_label: &str, minimize_to_tray: bool) -> bool { minimize_to_tray && window_label == MAIN_WINDOW_LABEL } +fn reveal_main_window(window: &tauri::WebviewWindow) { + let run_action = |action: &str, operation: &dyn Fn() -> tauri::Result<()>| { + if let Err(error) = operation() { + tracing::warn!("[启动] 主窗口{}失败: {}", action, error); + } + }; + + run_action("取消最小化", &|| window.unminimize()); + run_action("最大化", &|| window.maximize()); + run_action("显示", &|| window.show()); + run_action("聚焦", &|| window.set_focus()); +} + /// 运行 Tauri 应用 /// /// 这是应用的主入口点,负责: @@ -87,7 +100,6 @@ pub fn run() { update_check_service: update_check_service_state, session_files: session_files_state, context_memory_service, - tool_hooks_service, recording_service, mcp_manager: mcp_manager_state, automation_service: automation_service_state, @@ -144,10 +156,7 @@ pub fn run() { // 将窗口带到前台 if let Some(window) = app.get_webview_window("main") { - let _ = window.unminimize(); - let _ = window.maximize(); - let _ = window.show(); - let _ = window.set_focus(); + reveal_main_window(&window); } })); @@ -177,7 +186,6 @@ pub fn run() { .manage(update_check_service_state) .manage(session_files_state) .manage(context_memory_service) - .manage(tool_hooks_service) .manage(recording_service) .manage(mcp_manager_state) .manage(automation_service_state) @@ -216,20 +224,17 @@ pub fn run() { } }) .setup(move |app| { - // 启动时先最大化再显示,避免用户看到“先小窗后展开”的过程 + // 启动时先最大化再显示,避免用户看到“先小窗后展开”的过程。 if let Some(main_window) = app.get_webview_window("main") { - if let Err(e) = main_window.maximize() { - tracing::warn!("[启动] 主窗口最大化失败: {}", e); - } - if let Err(e) = main_window.show() { - tracing::warn!("[启动] 主窗口显示失败: {}", e); - } + reveal_main_window(&main_window); #[cfg(debug_assertions)] if crate::profiling::should_open_webview_devtools() { main_window.open_devtools(); tracing::info!("[Profiling] 已自动打开主窗口 WebView DevTools"); } + } else { + tracing::warn!("[启动] 未找到主窗口,无法执行启动展示流程"); } #[cfg(target_os = "windows")] @@ -360,26 +365,30 @@ pub fn run() { automation_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, - &automation_state, - ) { - Ok(resumed) if resumed > 0 => { - tracing::info!("[启动] 已恢复 {} 个会话的排队执行", resumed); + tauri::async_runtime::spawn(async move { + 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, + &automation_state, + ) + .await + { + Ok(resumed) if resumed > 0 => { + tracing::info!("[启动] 已恢复 {} 个会话的排队执行", resumed); + } + Ok(_) => { + tracing::debug!("[启动] 无需恢复持久化排队执行"); + } + Err(error) => { + tracing::warn!("[启动] 恢复持久化排队执行失败: {}", error); + } } - Ok(_) => { - tracing::debug!("[启动] 无需恢复持久化排队执行"); - } - Err(error) => { - tracing::warn!("[启动] 恢复持久化排队执行失败: {}", error); - } - } + }); } #[cfg(debug_assertions)] @@ -1585,16 +1594,6 @@ pub fn run() { commands::document_import_cmd::import_document, commands::document_import_cmd::import_document_to_session, commands::document_import_cmd::save_exported_document, - // Unified Chat commands(统一对话 API,后续治理收口入口) - commands::unified_chat_cmd::chat_create_session, - commands::unified_chat_cmd::chat_list_sessions, - commands::unified_chat_cmd::chat_get_session, - commands::unified_chat_cmd::chat_delete_session, - commands::unified_chat_cmd::chat_rename_session, - commands::unified_chat_cmd::chat_get_messages, - commands::unified_chat_cmd::chat_send_message, - commands::unified_chat_cmd::chat_stop_generation, - commands::unified_chat_cmd::chat_configure_provider, // Workspace commands commands::workspace_cmd::workspace_create, commands::workspace_cmd::workspace_list, @@ -1716,15 +1715,6 @@ pub fn run() { commands::memory_cmd::outline_node_update, commands::memory_cmd::outline_node_delete, commands::memory_cmd::project_memory_get, - // Context Memory commands - commands::context_memory::save_memory_entry, - commands::context_memory::get_session_memories, - commands::context_memory::get_memory_context, - commands::context_memory::record_error, - commands::context_memory::should_avoid_operation, - commands::context_memory::mark_error_resolved, - commands::context_memory::get_memory_stats, - commands::context_memory::cleanup_expired_memories, // Usage Stats commands commands::usage_stats_cmd::get_usage_stats, commands::usage_stats_cmd::get_model_usage_ranking, @@ -1757,14 +1747,6 @@ pub fn run() { // File Upload commands commands::file_upload_cmd::upload_avatar, commands::file_upload_cmd::delete_avatar, - // Tool Hooks commands - commands::tool_hooks::execute_hooks, - commands::tool_hooks::add_hook_rule, - commands::tool_hooks::remove_hook_rule, - commands::tool_hooks::toggle_hook_rule, - commands::tool_hooks::get_hook_rules, - commands::tool_hooks::get_hook_execution_stats, - commands::tool_hooks::clear_hook_execution_stats, // ASR commands commands::asr_cmd::get_asr_credentials, commands::asr_cmd::add_asr_credential, diff --git a/src-tauri/src/app/setup.rs b/src-tauri/src/app/setup.rs deleted file mode 100644 index 019a0eced..000000000 --- a/src-tauri/src/app/setup.rs +++ /dev/null @@ -1,289 +0,0 @@ -//! Tauri Setup Hook -//! -//! 包含应用启动时的初始化逻辑。 - -use std::sync::Arc; -use tauri::{App, Manager}; - -// use crate::agent::tools::{set_term_scrollback_tool_app_handle, set_terminal_tool_app_handle}; -use crate::agent::AsterAgentState; -use crate::database; -use crate::skills::ensure_default_local_skills; -use crate::telemetry; -use crate::tray::{TrayIconStatus, TrayManager, TrayStateSnapshot}; -use lime_scheduler::AgentScheduler; -use lime_services::aster_session_store::LimeSessionStore; -use lime_services::provider_pool_service::ProviderPoolService; -use lime_services::token_cache_service::TokenCacheService; - -use super::scheduler_service::{SchedulerService, SchedulerServiceConfig}; -use super::types::{AppState, LogState, TrayManagerState}; - -/// Tauri setup hook -/// -/// 在应用启动时执行初始化逻辑 -#[allow(clippy::too_many_arguments)] -pub fn setup_app( - app: &mut App, - state: AppState, - logs: LogState, - db: database::DbConnection, - pool_service: Arc, - token_cache: Arc, - shared_stats: Arc>, - shared_tokens: Arc>, - shared_logger: Arc, -) -> Result<(), Box> { - // 注册全局 SessionStore(作为后备方案) - // 注意:主要的 SessionStore 注入在 AsterAgentState::init_agent_with_db() 中完成 - // 这里的全局注册是为了兼容可能直接使用 SessionManager 静态方法的代码 - let session_store = Arc::new(LimeSessionStore::new(db.clone())); - tauri::async_runtime::block_on(async { - if let Err(e) = aster::session::set_global_session_store(session_store).await { - tracing::warn!("[启动] 注册全局 SessionStore 失败(可能已注册): {}", e); - } else { - tracing::info!("[启动] 全局 LimeSessionStore 已注册(后备方案)"); - } - }); - - // 初始化托盘管理器 - match TrayManager::new(app.handle()) { - Ok(tray_manager) => { - tracing::info!("[启动] 托盘管理器初始化成功"); - let tray_state: TrayManagerState = - TrayManagerState(Arc::new(tokio::sync::RwLock::new(Some(tray_manager)))); - app.manage(tray_state); - } - Err(e) => { - tracing::error!("[启动] 托盘管理器初始化失败: {}", e); - let tray_state: TrayManagerState = - TrayManagerState(Arc::new(tokio::sync::RwLock::new(None))); - app.manage(tray_state); - } - } - - // 初始化 AsterAgentState - let aster_agent_state = AsterAgentState::new(); - app.manage(aster_agent_state); - - // TODO: 重新实现 TerminalTool 和 TermScrollbackTool 的 AppHandle 设置 - // 当前暂时注释掉,等待适配 aster-rust 工具系统 - // set_terminal_tool_app_handle(app.handle().clone()); - // tracing::info!("[启动] TerminalTool AppHandle 已设置"); - - // set_term_scrollback_tool_app_handle(app.handle().clone()); - // tracing::info!("[启动] TermScrollbackTool AppHandle 已设置"); - - // 初始化默认 skill repos - { - let conn = lime_core::database::lock_db(&db)?; - database::dao::skills::SkillDao::init_default_skill_repos(&conn) - .expect("Failed to initialize default skill repos"); - } - match ensure_default_local_skills() { - Ok(installed) if installed.is_empty() => { - tracing::info!("[启动] 默认本地 Skills 已存在,跳过写入"); - } - Ok(installed) => { - tracing::info!("[启动] 默认本地 Skills 安装完成: {}", installed.join(", ")); - } - Err(error) => { - tracing::warn!("[启动] 安装默认本地 Skills 失败: {}", error); - } - } - - // 初始化调度器数据库表 - if let Err(e) = AgentScheduler::init_tables(&db) { - tracing::error!("[启动] 调度器表初始化失败: {}", e); - } else { - tracing::info!("[启动] 调度器表初始化成功"); - } - - // 启动调度器服务 - let scheduler_config = SchedulerServiceConfig::default(); - let scheduler_service = SchedulerService::new(db.clone(), scheduler_config); - scheduler_service.start(db.clone()); - tracing::info!("[启动] 调度器服务已启动"); - - // 将调度器服务注册为 Tauri 状态,以便后续访问 - app.manage(Arc::new(scheduler_service)); - - // 自动启动服务器 - let app_handle = app.handle().clone(); - tauri::async_runtime::spawn(async move { - start_server_async( - state, - logs, - db, - pool_service, - token_cache, - shared_stats, - shared_tokens, - shared_logger, - app_handle, - ) - .await; - }); - - Ok(()) -} - -/// 异步启动服务器 -#[allow(clippy::too_many_arguments)] -async fn start_server_async( - state: AppState, - logs: LogState, - db: database::DbConnection, - pool_service: Arc, - token_cache: Arc, - shared_stats: Arc>, - shared_tokens: Arc>, - shared_logger: Arc, - app_handle: tauri::AppHandle, -) { - let mut available_credentials = 0usize; - let mut total_credentials = 0usize; - - // 先加载凭证池中的凭证 - { - logs.write().await.add("info", "[启动] 正在加载凭证池..."); - - match pool_service.get_overview(&db) { - Ok(overview) => { - let mut loaded_types = Vec::new(); - - for provider_overview in overview { - let enabled_credentials: Vec<_> = provider_overview - .credentials - .iter() - .filter(|credential| !credential.is_disabled) - .collect(); - let count = enabled_credentials.len(); - if count > 0 { - total_credentials += count; - available_credentials += enabled_credentials - .iter() - .filter(|credential| credential.is_healthy) - .count(); - let provider_name = match provider_overview.provider_type.as_str() { - "kiro" => "Kiro", - "gemini" => "Gemini", - "qwen" => "通义千问", - "antigravity" => "Antigravity", - "openai" => "OpenAI", - "claude" => "Claude", - "codex" => "Codex", - "claude_oauth" => "Claude OAuth", - "iflow" => "iFlow", - _ => &provider_overview.provider_type, - }; - loaded_types.push(format!("{provider_name} ({count} 个)")); - } - } - - if loaded_types.is_empty() { - logs.write().await.add("warn", "[启动] 未找到任何可用凭证"); - } else { - let message = format!( - "[启动] 凭证已加载: {} (共 {} 个)", - loaded_types.join(", "), - total_credentials - ); - logs.write().await.add("info", &message); - } - } - Err(e) => { - logs.write() - .await - .add("warn", &format!("[启动] 获取凭证池信息失败: {e}")); - } - } - - // 兼容性:仍然尝试加载旧的 Kiro 凭证(如果存在) - let mut s = state.write().await; - if let Err(e) = s.kiro_provider.load_credentials().await { - logs.write() - .await - .add("debug", &format!("[启动] 旧版 Kiro 凭证加载失败: {e}")); - } - } - - // 启动服务器 - let server_started; - let server_address; - { - let mut s = state.write().await; - logs.write() - .await - .add("info", "[启动] 正在自动启动服务器..."); - match s - .start_with_telemetry( - logs.clone(), - pool_service, - token_cache, - Some(db), - Some(shared_stats), - Some(shared_tokens), - Some(shared_logger), - ) - .await - { - Ok(_) => { - // 获取服务器实际使用的地址(可能已经自动切换到有效的 IP) - let status = s.status(); - let host = status.host; - let port = status.port; - logs.write() - .await - .add("info", &format!("[启动] 服务器已启动: {host}:{port}")); - server_started = true; - server_address = format!("{host}:{port}"); - } - Err(e) => { - logs.write() - .await - .add("error", &format!("[启动] 服务器启动失败: {e}")); - server_started = false; - server_address = String::new(); - } - } - } - - // 更新托盘状态 - if let Some(tray_state) = app_handle.try_state::>() { - let tray_guard = tray_state.0.read().await; - if let Some(tray_manager) = tray_guard.as_ref() { - let current_state = tray_manager.get_state().await; - let icon_status = if !server_started { - TrayIconStatus::Stopped - } else if total_credentials > 0 && available_credentials == 0 { - TrayIconStatus::Error - } else if available_credentials < total_credentials { - TrayIconStatus::Warning - } else { - TrayIconStatus::Running - }; - - let snapshot = TrayStateSnapshot { - icon_status, - server_running: server_started, - server_address, - available_credentials, - total_credentials, - today_requests: current_state.today_requests, - auto_start_enabled: current_state.auto_start_enabled, - current_model_provider_type: current_state.current_model_provider_type, - current_model_provider_label: current_state.current_model_provider_label, - current_model: current_state.current_model, - current_theme_label: current_state.current_theme_label, - quick_model_groups: current_state.quick_model_groups, - }; - - if let Err(e) = tray_manager.update_state(snapshot).await { - tracing::error!("[启动] 更新托盘状态失败: {}", e); - } else { - tracing::info!("[启动] 托盘状态已更新"); - } - } - } -} diff --git a/src-tauri/src/app/state.rs b/src-tauri/src/app/state.rs index afe27c66f..8c6774a9a 100644 --- a/src-tauri/src/app/state.rs +++ b/src-tauri/src/app/state.rs @@ -14,7 +14,6 @@ use crate::commands::plugin_install_cmd::PluginInstallerState; use crate::commands::provider_pool_cmd::{CredentialSyncServiceState, ProviderPoolServiceState}; use crate::commands::resilience_cmd::ResilienceConfigState; use crate::commands::skill_cmd::SkillServiceState; -use crate::commands::tool_hooks::ToolHooksServiceState; use crate::config::{GlobalConfigManager, GlobalConfigManagerState}; use crate::database; use crate::plugin; @@ -26,7 +25,6 @@ use lime_services::context_memory_service::{ContextMemoryConfig, ContextMemorySe use lime_services::provider_pool_service::ProviderPoolService; use lime_services::skill_service::SkillService; use lime_services::token_cache_service::TokenCacheService; -use lime_services::tool_hooks_service::ToolHooksService; use super::types::{AppState, LogState, TokenCacheServiceState}; use crate::logger; @@ -61,7 +59,6 @@ pub struct ServiceStates { pub plugin_installer: PluginInstallerState, pub orchestrator: OrchestratorState, pub context_memory_service: ContextMemoryServiceState, - pub tool_hooks_service: ToolHooksServiceState, pub workflow_service: Arc>, pub progress_store: Arc>, } @@ -113,10 +110,6 @@ pub fn init_service_states() -> ServiceStates { .expect("Failed to initialize ContextMemoryService"); let context_memory_service_state = ContextMemoryServiceState(Arc::new(context_memory_service)); - // Initialize ToolHooksService - let tool_hooks_service = ToolHooksService::new(context_memory_service_state.0.clone()); - let tool_hooks_service_state = ToolHooksServiceState(Arc::new(tool_hooks_service)); - // Initialize WorkflowService let workflow_service = WorkflowService::new(); let workflow_service_state = Arc::new(RwLock::new(workflow_service)); @@ -138,7 +131,6 @@ pub fn init_service_states() -> ServiceStates { plugin_installer: plugin_installer_state, orchestrator: orchestrator_state, context_memory_service: context_memory_service_state, - tool_hooks_service: tool_hooks_service_state, workflow_service: workflow_service_state, progress_store: progress_store_state, } diff --git a/src-tauri/src/commands/aster_agent_cmd.rs b/src-tauri/src/commands/aster_agent_cmd.rs index 9db775cfd..89c98cab0 100644 --- a/src-tauri/src/commands/aster_agent_cmd.rs +++ b/src-tauri/src/commands/aster_agent_cmd.rs @@ -5,9 +5,17 @@ //! 支持从 Lime 凭证池自动选择凭证 use crate::agent::aster_state::{ProviderConfig, SessionConfigBuilder}; +use crate::agent::runtime_queue_service::{ + clear_runtime_queue as clear_runtime_queue_service, + list_runtime_queue_snapshots as list_runtime_queue_snapshots_service, + remove_runtime_queued_turn as remove_runtime_queued_turn_service, + resume_persisted_runtime_queues_on_startup as resume_persisted_runtime_queues_on_startup_service, + resume_runtime_queue_if_needed as resume_runtime_queue_if_needed_service, + submit_runtime_turn as submit_runtime_turn_service, RuntimeQueueExecutor, +}; use crate::agent::{ - AsterAgentState, AsterAgentWrapper, LimeScheduler, QueueInsertResult, QueuedTurnSnapshot, - QueuedTurnTask, SessionDetail, SessionInfo, SubAgentRole, TauriAgentEvent, + AsterAgentState, AsterAgentWrapper, LimeScheduler, QueuedTurnSnapshot, QueuedTurnTask, + SessionDetail, SessionInfo, SubAgentRole, TauriAgentEvent, }; use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState; use crate::commands::webview_cmd::{ @@ -15,18 +23,13 @@ use crate::commands::webview_cmd::{ BrowserBackendType, }; use crate::config::{GlobalConfigManager, GlobalConfigManagerState}; -use crate::database::dao::agent_runtime_queue::{ - AgentRuntimeQueuedTurnDao, NewAgentRuntimeQueuedTurnRecord, -}; use crate::database::DbConnection; use crate::mcp::{McpManagerState, McpServerConfig}; -use crate::services::agent_timeline_service::{ - build_action_response_value, complete_action_item, AgentTimelineRecorder, -}; +use crate::services::agent_timeline_service::AgentTimelineRecorder; use crate::services::automation_service::AutomationServiceState; use crate::services::execution_tracker_service::{ExecutionTracker, RunFinishDecision, RunSource}; use crate::services::memory_profile_prompt_service::{ - merge_system_prompt_with_memory_profile, merge_system_prompt_with_memory_sources, + merge_system_prompt_with_memory_context, MemoryPromptContext, }; 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; @@ -46,7 +49,7 @@ use aster::permission::{Permission, PermissionConfirmation, PrincipalType}; use aster::sandbox::{ detect_best_sandbox, execute_in_sandbox, ResourceLimits, SandboxConfig as ProcessSandboxConfig, }; -use aster::session::{SessionRuntimeSnapshot, TurnContextOverride}; +use aster::session::TurnContextOverride; use aster::tools::task_output_tool::TaskOutputInput; use aster::tools::{ BashTool, KillShellTool, PermissionBehavior, PermissionCheckResult, TaskManager, @@ -54,19 +57,18 @@ use aster::tools::{ MAX_OUTPUT_LENGTH, }; use async_trait::async_trait; -use futures::StreamExt; +use futures::{FutureExt, StreamExt}; #[cfg(test)] use lime_agent::request_tool_policy::REQUEST_TOOL_POLICY_MARKER; use lime_agent::request_tool_policy::{ merge_system_prompt_with_request_tool_policy, resolve_request_tool_policy_with_mode, - stream_reply_with_policy, ReplyAttemptError, RequestToolPolicy, RequestToolPolicyMode, + stream_message_reply_with_policy, ReplyAttemptError, RequestToolPolicy, RequestToolPolicyMode, }; use lime_agent::{ - convert_item_runtime, convert_turn_runtime, 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, + 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 lime_core::database::dao::agent_timeline::{AgentThreadItem, AgentThreadTurn}; use lime_services::api_key_provider_service::ApiKeyProviderService; use lime_services::mcp_service::McpService; use lime_services::video_generation_service::{ @@ -545,6 +547,8 @@ pub struct AgentRuntimeSessionDetail { pub turns: Vec, pub items: Vec, #[serde(default)] + pub todo_items: Vec, + #[serde(default)] pub queued_turns: Vec, } @@ -560,82 +564,12 @@ impl AgentRuntimeSessionDetail { execution_strategy: detail.execution_strategy, turns: detail.turns, items: detail.items, + todo_items: detail.todo_items, queued_turns, } } } -fn sort_runtime_turns(turns: &mut [AgentThreadTurn]) { - turns.sort_by(|left, right| { - left.started_at - .cmp(&right.started_at) - .then(left.created_at.cmp(&right.created_at)) - .then(left.id.cmp(&right.id)) - }); -} - -fn sort_runtime_items(items: &mut [AgentThreadItem], turn_started_at: &HashMap) { - items.sort_by(|left, right| { - let left_turn_started = turn_started_at - .get(&left.turn_id) - .map(String::as_str) - .unwrap_or(left.started_at.as_str()); - let right_turn_started = turn_started_at - .get(&right.turn_id) - .map(String::as_str) - .unwrap_or(right.started_at.as_str()); - - left_turn_started - .cmp(right_turn_started) - .then(left.sequence.cmp(&right.sequence)) - .then(left.turn_id.cmp(&right.turn_id)) - .then(left.started_at.cmp(&right.started_at)) - .then(left.id.cmp(&right.id)) - }); -} - -fn apply_aster_runtime_snapshot(detail: &mut SessionDetail, snapshot: &SessionRuntimeSnapshot) { - if let Some(thread) = snapshot.threads.first() { - detail.thread_id = thread.thread.id.clone(); - } - - if snapshot.threads.is_empty() { - return; - } - - let mut turns_by_id = detail - .turns - .drain(..) - .map(|turn| (turn.id.clone(), turn)) - .collect::>(); - for thread in &snapshot.threads { - for turn in &thread.turns { - turns_by_id.insert(turn.id.clone(), convert_turn_runtime(turn.clone())); - } - } - detail.turns = turns_by_id.into_values().collect(); - sort_runtime_turns(&mut detail.turns); - - let turn_started_at = detail - .turns - .iter() - .map(|turn| (turn.id.clone(), turn.started_at.clone())) - .collect::>(); - - let mut items_by_id = detail - .items - .drain(..) - .map(|item| (item.id.clone(), item)) - .collect::>(); - for thread in &snapshot.threads { - for item in &thread.items { - items_by_id.insert(item.id.clone(), convert_item_runtime(item.clone())); - } - } - detail.items = items_by_id.into_values().collect(); - sort_runtime_items(&mut detail.items, &turn_started_at); -} - #[derive(Debug, Clone, Copy, Deserialize, PartialEq, Eq)] #[serde(rename_all = "snake_case")] pub enum AgentRuntimeActionType { @@ -1633,6 +1567,100 @@ fn build_turn_runtime_statuses( ) } +fn emit_projected_runtime_item_event( + app: &AppHandle, + event_name: &str, + timeline_recorder: &Arc>, + workspace_root: &str, + event: TauriAgentEvent, +) { + if let Err(error) = app.emit(event_name, &event) { + tracing::warn!("[AsterAgent] 发送 runtime item 投影事件失败: {}", error); + } + + let mut recorder = match timeline_recorder.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + if let Err(error) = recorder.record_runtime_event(app, event_name, &event, workspace_root) { + tracing::warn!( + "[AsterAgent] 记录 runtime item 投影事件失败(已降级继续): {}", + error + ); + } +} + +async fn emit_runtime_status_with_projection( + agent: &Agent, + app: &AppHandle, + event_name: &str, + timeline_recorder: &Arc>, + workspace_root: &str, + session_config: &aster::agents::SessionConfig, + status: TauriRuntimeStatus, +) { + match agent + .upsert_runtime_status_item( + session_config, + status.phase.clone(), + status.title.clone(), + status.detail.clone(), + status.checkpoints.clone(), + ) + .await + { + Ok(agent_event) => { + for event in lime_agent::convert_agent_event(agent_event) { + emit_projected_runtime_item_event( + app, + event_name, + timeline_recorder, + workspace_root, + event, + ); + } + } + Err(error) => { + tracing::warn!( + "[AsterAgent] 写入 runtime_status item 失败,降级仅发送 transient 事件: {}", + error + ); + } + } + + let runtime_event = TauriAgentEvent::RuntimeStatus { status }; + if let Err(error) = app.emit(event_name, &runtime_event) { + tracing::warn!("[AsterAgent] 发送 runtime_status 失败: {}", error); + } +} + +async fn complete_runtime_status_projection( + agent: &Agent, + app: &AppHandle, + event_name: &str, + timeline_recorder: &Arc>, + workspace_root: &str, + session_config: &aster::agents::SessionConfig, +) { + match agent.complete_runtime_status_item(session_config).await { + Ok(Some(agent_event)) => { + for event in lime_agent::convert_agent_event(agent_event) { + emit_projected_runtime_item_event( + app, + event_name, + timeline_recorder, + workspace_root, + event, + ); + } + } + Ok(None) => {} + Err(error) => { + tracing::warn!("[AsterAgent] 完成 runtime_status item 失败: {}", error); + } + } +} + fn extend_map_with_harness_fields( target: &mut serde_json::Map, request_metadata: Option<&serde_json::Value>, @@ -2139,7 +2167,7 @@ async fn stream_reply_once( agent: &Agent, app: &AppHandle, event_name: &str, - message_text: &str, + user_message: Message, working_directory: Option<&Path>, session_config: aster::agents::SessionConfig, cancel_token: CancellationToken, @@ -2149,9 +2177,9 @@ async fn stream_reply_once( where F: FnMut(&TauriAgentEvent), { - stream_reply_with_policy( + stream_message_reply_with_policy( agent, - message_text, + user_message, working_directory, session_config, Some(cancel_token), @@ -2167,6 +2195,29 @@ where .map(|_| ()) } +fn build_runtime_user_message(message_text: &str, images: Option<&[ImageInput]>) -> Message { + let mut message = Message::user(); + + if !message_text.is_empty() { + message = message.with_text(message_text); + } + + if let Some(images) = images { + for image in images { + if image.data.trim().is_empty() || image.media_type.trim().is_empty() { + continue; + } + message = message.with_image(image.data.clone(), image.media_type.clone()); + } + } + + if message.content.is_empty() { + return Message::user().with_text(message_text); + } + + message +} + /// 基于 aster::sandbox 的本地 bash 强隔离工具 #[derive(Debug)] struct WorkspaceSandboxedBashTool { @@ -5323,7 +5374,6 @@ async fn apply_workspace_sandbox_permissions( "WebSearch", "ask", "tool_search", - "three_stage_workflow", SOCIAL_IMAGE_TOOL_NAME, LIME_CREATE_VIDEO_TASK_TOOL_NAME, LIME_CREATE_BROADCAST_TASK_TOOL_NAME, @@ -5687,11 +5737,10 @@ async fn execute_aster_chat_request( } }; - let prompt_with_memory = merge_system_prompt_with_memory_sources( - merge_system_prompt_with_memory_profile(resolved_prompt, &runtime_config), + let prompt_with_memory = merge_system_prompt_with_memory_context( + resolved_prompt, &runtime_config, - Path::new(&workspace_root), - None, + MemoryPromptContext::with_working_dir(Path::new(&workspace_root)), ); let merged_prompt = merge_system_prompt_with_auto_continue( merge_system_prompt_with_elicitation_context( @@ -5874,6 +5923,33 @@ async fn execute_aster_chat_request( resolved_turn_id.clone(), request.message.clone(), )?)); + let include_context_trace = runtime_config.memory.enabled; + let turn_context = build_turn_context_override(request.metadata.as_ref()); + let runtime_status_session_config = { + let mut session_config_builder = SessionConfigBuilder::new(session_id) + .thread_id(resolved_thread_id.clone()) + .turn_id(resolved_turn_id.clone()); + if let Some(turn_context) = turn_context.clone() { + session_config_builder = session_config_builder.turn_context(turn_context); + } + session_config_builder.build() + }; + + // 获取 Agent Arc 并保持 guard 在整个流处理期间存活 + let guard = agent_arc.read().await; + let agent = guard.as_ref().ok_or("Agent not initialized")?; + if let Err(error) = agent + .ensure_runtime_turn_initialized( + &runtime_status_session_config, + Some(request.message.clone()), + ) + .await + { + tracing::warn!( + "[AsterAgent] 初始化 runtime turn 失败,后续降级继续: {}", + error + ); + } let (initial_runtime_status, decided_runtime_status) = build_turn_runtime_statuses( &request, @@ -5885,27 +5961,17 @@ async fn execute_aster_chat_request( .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_runtime_event(app, &request.event_name, &event, workspace_root.as_str()) - { - tracing::warn!("[AsterAgent] 记录 runtime_status 失败: {}", error); - } + emit_runtime_status_with_projection( + agent, + app, + &request.event_name, + &timeline_recorder, + workspace_root.as_str(), + &runtime_status_session_config, + status, + ) + .await; } - - // 获取 Agent Arc 并保持 guard 在整个流处理期间存活 - let guard = agent_arc.read().await; - let agent = guard.as_ref().ok_or("Agent not initialized")?; - - let include_context_trace = runtime_config.memory.enabled; - let turn_context = build_turn_context_override(request.metadata.as_ref()); let resolved_thread_id_for_session = resolved_thread_id.clone(); let resolved_turn_id_for_session = resolved_turn_id.clone(); @@ -5941,7 +6007,7 @@ async fn execute_aster_chat_request( agent, app, &request.event_name, - &request.message, + build_runtime_user_message(&request.message, request.images.as_deref()), Some(Path::new(&workspace_root)), build_session_config(), cancel_token.clone(), @@ -6013,7 +6079,10 @@ async fn execute_aster_chat_request( agent, &app, &request.event_name, - &request.message, + build_runtime_user_message( + &request.message, + request.images.as_deref(), + ), Some(Path::new(&workspace_root)), build_session_config(), cancel_token.clone(), @@ -6109,6 +6178,15 @@ async fn execute_aster_chat_request( match final_result { Ok(()) => { + complete_runtime_status_projection( + agent, + app, + &request.event_name, + &timeline_recorder, + workspace_root.as_str(), + &runtime_status_session_config, + ) + .await; { let mut recorder = match timeline_recorder.lock() { Ok(guard) => guard, @@ -6124,6 +6202,15 @@ async fn execute_aster_chat_request( } } Err(e) => { + complete_runtime_status_projection( + agent, + app, + &request.event_name, + &timeline_recorder, + workspace_root.as_str(), + &runtime_status_session_config, + ) + .await; { let mut recorder = match timeline_recorder.lock() { Ok(guard) => guard, @@ -6151,60 +6238,6 @@ async fn execute_aster_chat_request( Ok(()) } -struct AgentRuntimeExecutionContext { - app: AppHandle, - state: AsterAgentState, - db: DbConnection, - api_key_provider_service: ApiKeyProviderServiceState, - logs: LogState, - config_manager: GlobalConfigManagerState, - mcp_manager: McpManagerState, - automation_state: AutomationServiceState, -} - -impl AgentRuntimeExecutionContext { - fn from_states( - app: AppHandle, - state: &AsterAgentState, - db: &DbConnection, - api_key_provider_service: &ApiKeyProviderServiceState, - logs: &LogState, - config_manager: &GlobalConfigManagerState, - mcp_manager: &McpManagerState, - automation_state: &AutomationServiceState, - ) -> 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(), - automation_state: automation_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(), - automation_state: self.automation_state.clone(), - } - } -} - fn build_queued_turn_preview(message: &str) -> String { let compact = message.split_whitespace().collect::>().join(" "); if compact.is_empty() { @@ -6252,254 +6285,28 @@ fn deserialize_queued_turn_request(payload: serde_json::Value) -> Result, -) -> 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); - } +fn build_runtime_queue_executor() -> RuntimeQueueExecutor { + Arc::new(|context, payload| { + async move { + let request = deserialize_queued_turn_request(payload)?; + execute_aster_chat_request( + &context.app, + &context.state, + &context.db, + &context.api_key_provider_service, + &context.logs, + &context.config_manager, + &context.mcp_manager, + &context.automation_state, + request, + ) + .await } - } - - 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) + .boxed() + }) } -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.automation_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( +pub async fn resume_persisted_runtime_queues_on_startup( app: AppHandle, state: &AsterAgentState, db: &DbConnection, @@ -6509,33 +6316,18 @@ pub fn resume_persisted_runtime_queues_on_startup( mcp_manager: &McpManagerState, automation_state: &AutomationServiceState, ) -> 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, - automation_state, - ); - if resume_runtime_queue_if_needed(context, session_id.clone())? { - resumed += 1; - tracing::info!( - "[AsterAgent][Queue] 启动阶段已恢复会话排队执行: session_id={}", - session_id - ); - } - } - - Ok(resumed) + resume_persisted_runtime_queues_on_startup_service( + app, + state, + db, + api_key_provider_service, + logs, + config_manager, + mcp_manager, + automation_state, + build_runtime_queue_executor(), + ) + .await } /// 统一运行时:提交一个 turn。 @@ -6554,9 +6346,7 @@ pub async fn agent_runtime_submit_turn( 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( + submit_runtime_turn_service( app, state.inner(), db.inner(), @@ -6565,39 +6355,11 @@ pub async fn agent_runtime_submit_turn( config_manager.inner(), mcp_manager.inner(), automation_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(()) - } - } + queued_task, + queue_if_busy, + build_runtime_queue_executor(), + ) + .await } /// 统一运行时:中断当前 turn。 @@ -6605,13 +6367,11 @@ pub async fn agent_runtime_submit_turn( pub async fn agent_runtime_interrupt_turn( app: AppHandle, state: State<'_, AsterAgentState>, - db: State<'_, DbConnection>, request: AgentRuntimeInterruptTurnRequest, ) -> Result { 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); + let cleared = clear_runtime_queue_service(&app, &session_id).await?; Ok(cancelled || !cleared.is_empty()) } @@ -6735,14 +6495,6 @@ fn list_runtime_sessions_internal(db: &DbConnection) -> Result, AsterAgentWrapper::list_sessions_sync(db) } -fn get_runtime_session_detail_internal( - db: &DbConnection, - session_id: &str, -) -> Result { - tracing::info!("[AsterAgent] 获取会话: {}", session_id); - AsterAgentWrapper::get_session_sync(db, session_id) -} - /// 统一运行时:获取会话详情。 #[tauri::command] pub async fn agent_runtime_get_session( @@ -6756,44 +6508,31 @@ pub async fn agent_runtime_get_session( automation_state: State<'_, AutomationServiceState>, session_id: String, ) -> Result { - ensure_runtime_queue_loaded(state.inner(), db.inner(), &session_id)?; - let mut detail = get_runtime_session_detail_internal(db.inner(), &session_id)?; + tracing::info!("[AsterAgent] 获取运行时会话: {}", session_id); + let detail = AsterAgentWrapper::get_runtime_session_detail(db.inner(), &session_id).await?; - let agent_arc = state.get_agent_arc(); - let guard = agent_arc.read().await; - if let Some(agent) = guard.as_ref() { - match agent.runtime_snapshot(&session_id).await { - Ok(snapshot) => apply_aster_runtime_snapshot(&mut detail, &snapshot), - Err(error) => { - tracing::warn!( - "[AsterAgent] 读取 Aster runtime snapshot 失败: session_id={}, error={}", - session_id, - error - ); - } - } - } - - 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(), - automation_state.inner(), + if let Err(error) = resume_runtime_queue_if_needed_service( + app, + state.inner(), + db.inner(), + api_key_provider_service.inner(), + logs.inner(), + config_manager.inner(), + mcp_manager.inner(), + automation_state.inner(), + session_id.clone(), + build_runtime_queue_executor(), + ) + .await + { + tracing::warn!( + "[AsterAgent][Queue] 获取会话后恢复排队执行失败: session_id={}, error={}", + session_id, + error ); - if let Err(error) = resume_runtime_queue_if_needed(context, session_id.clone()) { - tracing::warn!( - "[AsterAgent][Queue] 获取会话后恢复排队执行失败: session_id={}, error={}", - session_id, - error - ); - } } + + let queued_turns = list_runtime_queue_snapshots_service(&session_id).await?; Ok(AgentRuntimeSessionDetail::from_session_detail( detail, queued_turns, @@ -6804,8 +6543,6 @@ pub async fn agent_runtime_get_session( #[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(); @@ -6814,25 +6551,7 @@ pub async fn agent_runtime_remove_queued_turn( 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) + remove_runtime_queued_turn_service(&app, &session_id, &queued_turn_id).await } fn rename_runtime_session_internal( @@ -6892,7 +6611,7 @@ pub async fn agent_runtime_delete_session( ) -> Result<(), String> { 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); + let _ = clear_runtime_queue_service(&app, &trimmed_session_id).await; delete_runtime_session_internal(db.inner(), &trimmed_session_id).await } @@ -6980,16 +6699,9 @@ fn build_runtime_action_user_data(request: &AgentRuntimeRespondActionRequest) -> #[tauri::command] pub async fn agent_runtime_respond_action( state: State<'_, AsterAgentState>, - db: State<'_, DbConnection>, request: AgentRuntimeRespondActionRequest, ) -> Result<(), String> { - let response_value = build_action_response_value( - request.confirmed, - request.response.as_deref(), - request.user_data.as_ref(), - ); - - let result = match request.action_type { + match request.action_type { AgentRuntimeActionType::ToolConfirmation => { confirm_runtime_action_internal( state.inner(), @@ -7014,13 +6726,7 @@ pub async fn agent_runtime_respond_action( ) .await } - }; - - if result.is_ok() { - complete_action_item(db.inner(), &request.request_id, response_value)?; } - - result } async fn submit_runtime_elicitation_response_internal( @@ -7626,6 +7332,27 @@ mod tests { ); } + #[test] + fn test_build_runtime_user_message_includes_images() { + let message = build_runtime_user_message( + "这个是什么", + Some(&[ImageInput { + data: "aGVsbG8=".to_string(), + media_type: "image/png".to_string(), + }]), + ); + + assert_eq!(message.as_concat_text(), "这个是什么"); + assert_eq!(message.content.len(), 2); + + if let MessageContent::Image(image) = &message.content[1] { + assert_eq!(image.data, "aGVsbG8="); + assert_eq!(image.mime_type, "image/png"); + } else { + panic!("expected image content in runtime user message"); + } + } + #[test] fn test_build_runtime_action_user_data_prefers_structured_payload() { let request = AgentRuntimeRespondActionRequest { diff --git a/src-tauri/src/commands/context_memory.rs b/src-tauri/src/commands/context_memory.rs index c142d925f..118917285 100644 --- a/src-tauri/src/commands/context_memory.rs +++ b/src-tauri/src/commands/context_memory.rs @@ -1,179 +1,6 @@ -//! 上下文记忆管理相关的 Tauri 命令 +//! 上下文记忆运行时共享状态。 -use crate::config::GlobalConfigManagerState; -use lime_services::context_memory_service::{ - ContextMemoryService, MemoryEntry, MemoryFileType, MemoryStats, -}; -use serde::{Deserialize, Serialize}; +use lime_services::context_memory_service::ContextMemoryService; use std::sync::Arc; -use tauri::State; -use tracing::{debug, info}; pub struct ContextMemoryServiceState(pub Arc); - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SaveMemoryRequest { - pub session_id: String, - pub file_type: MemoryFileType, - pub title: String, - pub content: String, - pub tags: Vec, - pub priority: u8, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RecordErrorRequest { - pub session_id: String, - pub error_description: String, - pub attempted_solution: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ResolveErrorRequest { - pub session_id: String, - pub error_description: String, - pub resolution: String, -} - -#[tauri::command] -pub async fn save_memory_entry( - memory_service: State<'_, ContextMemoryServiceState>, - request: SaveMemoryRequest, -) -> Result<(), String> { - debug!( - "保存记忆条目: {} (会话: {})", - request.title, request.session_id - ); - - let entry = MemoryEntry { - id: uuid::Uuid::new_v4().to_string(), - session_id: request.session_id.clone(), - file_type: request.file_type, - title: request.title, - content: request.content, - tags: request.tags, - priority: request.priority, - created_at: chrono::Utc::now().timestamp_millis(), - updated_at: chrono::Utc::now().timestamp_millis(), - archived: false, - }; - - memory_service.0.save_memory_entry(&entry)?; - info!("记忆条目保存成功: {}", entry.title); - Ok(()) -} - -#[tauri::command] -pub async fn get_session_memories( - memory_service: State<'_, ContextMemoryServiceState>, - session_id: String, - file_type: Option, -) -> Result, String> { - debug!("获取会话记忆: {} (类型: {:?})", session_id, file_type); - let memories = memory_service - .0 - .get_session_memories(&session_id, file_type)?; - info!("获取到 {} 个记忆条目", memories.len()); - Ok(memories) -} - -#[tauri::command] -pub async fn get_memory_context( - memory_service: State<'_, ContextMemoryServiceState>, - session_id: String, -) -> Result { - debug!("获取记忆上下文: {}", session_id); - let context = memory_service.0.get_memory_context(&session_id)?; - info!("记忆上下文长度: {} 字符", context.len()); - Ok(context) -} - -#[tauri::command] -pub async fn record_error( - memory_service: State<'_, ContextMemoryServiceState>, - request: RecordErrorRequest, -) -> Result<(), String> { - debug!( - "记录错误: {} (会话: {})", - request.error_description, request.session_id - ); - memory_service.0.record_error( - &request.session_id, - &request.error_description, - &request.attempted_solution, - )?; - info!("错误记录成功"); - Ok(()) -} - -#[tauri::command] -pub async fn should_avoid_operation( - memory_service: State<'_, ContextMemoryServiceState>, - session_id: String, - operation_description: String, -) -> Result { - debug!( - "检查是否避免操作: {} (会话: {})", - operation_description, session_id - ); - let should_avoid = memory_service - .0 - .should_avoid_operation(&session_id, &operation_description); - if should_avoid { - info!("建议避免操作: {}", operation_description); - } - Ok(should_avoid) -} - -#[tauri::command] -pub async fn mark_error_resolved( - memory_service: State<'_, ContextMemoryServiceState>, - request: ResolveErrorRequest, -) -> Result<(), String> { - debug!( - "标记错误已解决: {} (会话: {})", - request.error_description, request.session_id - ); - memory_service.0.mark_error_resolved( - &request.session_id, - &request.error_description, - &request.resolution, - )?; - info!("错误已标记为解决"); - Ok(()) -} - -#[tauri::command] -pub async fn get_memory_stats( - memory_service: State<'_, ContextMemoryServiceState>, - session_id: String, -) -> Result { - debug!("获取记忆统计: {}", session_id); - let stats = memory_service.0.get_memory_stats(&session_id)?; - info!( - "记忆统计: {} 个活跃记忆, {} 个未解决错误", - stats.active_memories, stats.unresolved_errors - ); - Ok(stats) -} - -#[tauri::command] -pub async fn cleanup_expired_memories( - memory_service: State<'_, ContextMemoryServiceState>, - global_config: State<'_, GlobalConfigManagerState>, -) -> Result<(), String> { - debug!("清理过期记忆"); - let memory_config = global_config.config().memory; - if matches!(memory_config.auto_cleanup, Some(false)) { - info!("自动清理已关闭,跳过过期记忆清理"); - return Ok(()); - } - - let retention_days = memory_config.retention_days.unwrap_or(30).clamp(1, 3650); - - memory_service - .0 - .cleanup_expired_memories_with_retention_days(retention_days)?; - info!("过期记忆清理完成"); - Ok(()) -} diff --git a/src-tauri/src/commands/ecommerce_review_reply_cmd.rs b/src-tauri/src/commands/ecommerce_review_reply_cmd.rs index e773a818d..0409830c0 100644 --- a/src-tauri/src/commands/ecommerce_review_reply_cmd.rs +++ b/src-tauri/src/commands/ecommerce_review_reply_cmd.rs @@ -7,9 +7,9 @@ use tauri::State; use crate::agent::AsterAgentState; use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState; -use crate::commands::skill_exec_cmd::{execute_skill, SkillExecutionResult}; use crate::config::GlobalConfigManagerState; use crate::database::DbConnection; +use crate::skills::{execute_named_skill, SkillExecutionRequest, SkillExecutionResult}; /// 电商差评回复请求 #[derive(Debug, Clone, Serialize, Deserialize)] @@ -72,19 +72,20 @@ pub async fn execute_ecommerce_review_reply( .unwrap_or_default() ); - // 调用通用的 execute_skill - execute_skill( - app_handle, - db, - api_key_provider_service, - config_manager, - aster_state, - "ecommerce-review-reply".to_string(), - user_input, - Some("anthropic".to_string()), // 优先使用 Anthropic - request.model, - request.execution_id, - None, // session_id + execute_named_skill( + &app_handle, + db.inner(), + api_key_provider_service.inner(), + config_manager.inner(), + aster_state.inner(), + SkillExecutionRequest { + skill_name: "ecommerce-review-reply".to_string(), + user_input, + provider_override: Some("anthropic".to_string()), + model_override: request.model, + execution_id: request.execution_id, + session_id: None, + }, ) .await } diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 04e7c4af2..2a17d4702 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -64,9 +64,7 @@ pub mod telemetry_cmd; pub mod template_cmd; pub mod terminal_cmd; pub mod theme_context_cmd; -pub mod tool_hooks; pub mod tray_cmd; -pub mod unified_chat_cmd; pub mod unified_memory_cmd; pub mod update_cmd; pub mod usage_cmd; diff --git a/src-tauri/src/commands/persona_cmd.rs b/src-tauri/src/commands/persona_cmd.rs index 9c28d907e..c6332fc59 100644 --- a/src-tauri/src/commands/persona_cmd.rs +++ b/src-tauri/src/commands/persona_cmd.rs @@ -25,6 +25,7 @@ use crate::models::project_model::{ BrandPersona, BrandPersonaExtension, BrandPersonaTemplate, CreateBrandExtensionRequest, CreatePersonaRequest, Persona, PersonaTemplate, PersonaUpdate, UpdateBrandExtensionRequest, }; +use crate::services::memory_profile_prompt_service::{build_memory_prompt, MemoryPromptContext}; use lime_services::persona_service::PersonaService; // ============================================================================ @@ -361,9 +362,7 @@ pub async fn generate_persona( crate::agent::aster_state::SessionConfigBuilder::new(&session_id) .include_context_trace(true); if let Some(memory_prompt) = - crate::services::memory_profile_prompt_service::build_memory_profile_prompt( - &config_manager.config(), - ) + build_memory_prompt(&config_manager.config(), MemoryPromptContext::default()) { session_config_builder = session_config_builder.system_prompt(memory_prompt); } diff --git a/src-tauri/src/commands/skill_error.rs b/src-tauri/src/commands/skill_error.rs index c570ffc08..333ea609f 100644 --- a/src-tauri/src/commands/skill_error.rs +++ b/src-tauri/src/commands/skill_error.rs @@ -7,7 +7,6 @@ pub const SKILL_ERR_CATALOG_UNAVAILABLE: &str = "skill_catalog_unavailable"; pub const SKILL_ERR_NOT_FOUND: &str = "skill_not_found"; pub const SKILL_ERR_SESSION_INIT_FAILED: &str = "skill_session_init_failed"; pub const SKILL_ERR_PROVIDER_UNAVAILABLE: &str = "skill_provider_unavailable"; -pub const SKILL_ERR_STREAM_FAILED: &str = "skill_stream_failed"; pub const SKILL_ERR_EXECUTE_FAILED: &str = "skill_execute_failed"; pub fn format_skill_error(code: &str, message: impl AsRef) -> String { diff --git a/src-tauri/src/commands/skill_exec_cmd.rs b/src-tauri/src/commands/skill_exec_cmd.rs index 20ffb0a93..a1701a9c3 100644 --- a/src-tauri/src/commands/skill_exec_cmd.rs +++ b/src-tauri/src/commands/skill_exec_cmd.rs @@ -6,572 +6,30 @@ //! - `get_skill_detail`: 获取 Skill 详情 //! //! ## 依赖 -//! - `AsterAgentState`: Aster Agent 状态管理,提供完整的工具集支持 -//! - `TauriExecutionCallback`: 执行进度回调 -//! - `ProviderPoolService`: 凭证池服务 +//! - `AsterAgentState`: Aster Agent 状态管理,提供底层 Agent 执行能力 +//! - `skills/runtime`: skill 执行前置准备、provider fallback 与 run metadata 边界 +//! - `ExecutionTracker`: 统一执行记录写入 //! //! ## Requirements //! - 3.1: execute_skill 命令接受 skill_name 和 user_input 参数 //! - 4.1: list_executable_skills 返回所有可执行的 skills //! - 5.1: get_skill_detail 接受 skill_name 参数 -use futures::StreamExt; -use serde::{Deserialize, Serialize}; -use tauri::{Emitter, State}; -use uuid::Uuid; +use tauri::State; -use aster::conversation::message::Message; -use chrono::Utc; - -use crate::agent::aster_state::SessionConfigBuilder; -use crate::agent::{AsterAgentState, TauriAgentEvent}; +use crate::agent::AsterAgentState; use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState; -use crate::commands::aster_agent_cmd::{ - ensure_browser_mcp_tools_registered, ensure_creation_task_tools_registered, - ensure_social_image_tool_registered, -}; -use crate::commands::skill_error::{ - format_skill_error, map_find_skill_error, SKILL_ERR_CATALOG_UNAVAILABLE, - SKILL_ERR_EXECUTE_FAILED, SKILL_ERR_PROVIDER_UNAVAILABLE, SKILL_ERR_SESSION_INIT_FAILED, - SKILL_ERR_STREAM_FAILED, -}; use crate::config::GlobalConfigManagerState; 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 lime_agent::event_converter::{convert_agent_event, TauriArtifactSnapshot, TauriToolResult}; -use lime_agent::WriteArtifactEventEmitter; -use lime_skills::{ - find_skill_by_name, get_skill_roots, load_skills_from_directory, ExecutionCallback, - LoadedSkillDefinition, -}; -#[cfg(test)] -use lime_skills::{ - load_skill_from_file, parse_allowed_tools, parse_boolean, parse_skill_frontmatter, +use crate::skills::{ + execute_named_skill, get_skill_detail_info, list_executable_skill_catalog, ExecutableSkillInfo, + SkillDetailInfo, SkillExecutionRequest, SkillExecutionResult, }; // ============================================================================ // 公开类型定义 // ============================================================================ -/// 可执行 Skill 信息 -/// -/// 用于 list_executable_skills 命令的返回类型 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ExecutableSkillInfo { - /// Skill 名称(唯一标识) - pub name: String, - /// 显示名称 - pub display_name: String, - /// Skill 描述 - pub description: String, - /// 执行模式:prompt, workflow, agent - pub execution_mode: String, - /// 是否有 workflow 定义 - pub has_workflow: bool, - /// 指定的 Provider(可选) - pub provider: Option, - /// 指定的 Model(可选) - pub model: Option, - /// 参数提示(可选) - pub argument_hint: Option, -} - -/// Workflow 步骤信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct WorkflowStepInfo { - /// 步骤 ID - pub id: String, - /// 步骤名称 - pub name: String, - /// 依赖的步骤 ID 列表 - pub dependencies: Vec, -} - -/// Skill 详情信息 -/// -/// 用于 get_skill_detail 命令的返回类型 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SkillDetailInfo { - /// 基本信息 - #[serde(flatten)] - pub basic: ExecutableSkillInfo, - /// Markdown 内容 - pub markdown_content: String, - /// Workflow 步骤(如果有) - pub workflow_steps: Option>, - /// 允许的工具列表(可选) - pub allowed_tools: Option>, - /// 使用场景说明(可选) - pub when_to_use: Option, -} - -/// 步骤执行结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct StepResult { - /// 步骤 ID - pub step_id: String, - /// 步骤名称 - pub step_name: String, - /// 是否成功 - pub success: bool, - /// 输出内容 - pub output: Option, - /// 错误信息 - pub error: Option, -} - -/// Skill 执行结果 -/// -/// 用于 execute_skill 命令的返回类型 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SkillExecutionResult { - /// 是否成功 - pub success: bool, - /// 最终输出 - pub output: Option, - /// 错误信息 - pub error: Option, - /// 已完成的步骤结果 - pub steps_completed: Vec, -} - -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"; -const SOCIAL_POST_EMPTY_FALLBACK_CONTENT: &str = "# 社媒文案\n\n(生成结果为空,请重试。)"; -const SOCIAL_POST_FALLBACK_COVER_URL: &str = "cover-generation-failed"; -const SOCIAL_POST_FALLBACK_COVER_NOTE: &str = "封面图生成失败,可稍后仅重试配图。"; -const SOCIAL_POST_DEFAULT_IMAGE_SIZE: &str = "1024x1024"; - -#[derive(Debug, Clone)] -struct SocialSkillOutputEnvelope { - final_output: String, - file_path: String, - file_content: String, -} - -fn infer_theme_workbench_gate_key(skill_name: &str, user_input: &str) -> &'static str { - let probe = format!("{} {}", skill_name, user_input).to_lowercase(); - if probe.contains("publish") - || probe.contains("adapt") - || probe.contains("distribution") - || probe.contains("release") - || probe.contains("发布") - || probe.contains("分发") - || probe.contains("平台适配") - { - return "publish_confirm"; - } - if probe.contains("topic") - || probe.contains("research") - || probe.contains("trend") - || probe.contains("idea") - || probe.contains("选题") - || probe.contains("方向") - || probe.contains("调研") - || probe.contains("洞察") - { - return "topic_select"; - } - "write_mode" -} - -fn normalize_social_post_output( - skill_name: &str, - user_input: &str, - execution_id: &str, - raw_output: &str, -) -> Option { - if skill_name != SOCIAL_POST_WITH_COVER_SKILL_NAME { - return None; - } - - let generated_path = build_social_post_file_path(user_input, execution_id); - if let Some((range, existing_path, content)) = extract_first_write_file_block(raw_output) { - let normalized_content = normalize_social_markdown_contract(&content); - let has_existing_path = existing_path.is_some(); - let path = existing_path.unwrap_or_else(|| generated_path.clone()); - - if has_existing_path { - if normalized_content != content { - let normalized_block = build_write_file_block(&path, &normalized_content); - let mut rebuilt = String::new(); - rebuilt.push_str(&raw_output[..range.start]); - rebuilt.push_str(&normalized_block); - rebuilt.push_str(&raw_output[range.end..]); - return Some(SocialSkillOutputEnvelope { - final_output: rebuilt, - file_path: path, - file_content: normalized_content, - }); - } - return Some(SocialSkillOutputEnvelope { - final_output: raw_output.to_string(), - file_path: path, - file_content: normalized_content, - }); - } - - let normalized_block = build_write_file_block(&path, &normalized_content); - let mut rebuilt = String::new(); - rebuilt.push_str(&raw_output[..range.start]); - rebuilt.push_str(&normalized_block); - rebuilt.push_str(&raw_output[range.end..]); - - return Some(SocialSkillOutputEnvelope { - final_output: rebuilt, - file_path: path, - file_content: normalized_content, - }); - } - - let normalized_content = normalize_social_markdown_contract(raw_output); - Some(SocialSkillOutputEnvelope { - final_output: build_write_file_block(&generated_path, &normalized_content), - file_path: generated_path, - file_content: normalized_content, - }) -} - -fn extract_first_write_file_block( - raw_output: &str, -) -> Option<(std::ops::Range, Option, String)> { - let open_start = raw_output.find("')?; - let open_end = open_start + open_end_offset; - let open_tag = &raw_output[open_start..=open_end]; - - let content_start = open_end + 1; - let close_tag = ""; - let close_offset = raw_output[content_start..].find(close_tag)?; - let close_start = content_start + close_offset; - let block_end = close_start + close_tag.len(); - - let content = raw_output[content_start..close_start].trim().to_string(); - let path = extract_write_file_path(open_tag); - Some((open_start..block_end, path, content)) -} - -fn extract_write_file_path(open_tag: &str) -> Option { - let path_idx = open_tag.find("path")?; - let after_path = &open_tag[path_idx + "path".len()..]; - let equal_idx = after_path.find('=')?; - let value = after_path[equal_idx + 1..].trim_start(); - let quote = value.chars().next()?; - if quote != '"' && quote != '\'' { - return None; - } - - let rest = &value[quote.len_utf8()..]; - let end_idx = rest.find(quote)?; - let path = rest[..end_idx].trim(); - if path.is_empty() { - None - } else { - Some(path.to_string()) - } -} - -fn normalize_social_output_content(content: &str) -> String { - let trimmed = content.trim(); - if trimmed.is_empty() { - SOCIAL_POST_EMPTY_FALLBACK_CONTENT.to_string() - } else { - trimmed.to_string() - } -} - -fn normalize_social_markdown_contract(content: &str) -> String { - let mut normalized = normalize_social_output_content(content); - if !normalized.contains("![封面图](") { - normalized = format!("{normalized}\n\n![封面图]({SOCIAL_POST_FALLBACK_COVER_URL})"); - } - normalized -} - -fn extract_cover_url_from_markdown(content: &str) -> Option { - for line in content.lines() { - let trimmed = line.trim(); - if !trimmed.starts_with("![") { - continue; - } - let open = trimmed.find("](")?; - let close = trimmed.rfind(')')?; - if close <= open + 2 { - continue; - } - let url = trimmed[(open + 2)..close].trim(); - if !url.is_empty() { - return Some(url.to_string()); - } - } - None -} - -fn extract_detail_value(content: &str, label: &str) -> Option { - let probe = format!("- {label}:"); - for line in content.lines() { - let trimmed = line.trim(); - if let Some(value) = trimmed.strip_prefix(&probe) { - let value = value.trim(); - if !value.is_empty() { - return Some(value.to_string()); - } - } - } - None -} - -fn derive_social_auxiliary_paths(article_path: &str) -> (String, String) { - let base = article_path.strip_suffix(".md").unwrap_or(article_path); - ( - format!("{base}.cover.json"), - format!("{base}.publish-pack.json"), - ) -} - -fn collect_social_artifact_paths_from_output(output: Option<&str>) -> Vec { - let Some(raw_output) = output else { - return Vec::new(); - }; - let Some((_, maybe_path, _)) = extract_first_write_file_block(raw_output) else { - return Vec::new(); - }; - let Some(article_path) = maybe_path else { - return Vec::new(); - }; - let (cover_meta_path, publish_pack_path) = derive_social_auxiliary_paths(&article_path); - vec![article_path, cover_meta_path, publish_pack_path] -} - -fn summarize_social_content(content: &str) -> String { - let compact = content - .lines() - .filter(|line| !line.trim().starts_with('#')) - .collect::>() - .join(" "); - let compact = compact.split_whitespace().collect::>().join(" "); - compact.chars().take(180).collect() -} - -fn build_social_auxiliary_file_payloads( - execution_id: &str, - user_input: &str, - article_path: &str, - article_content: &str, -) -> Vec<(String, String)> { - let (cover_meta_path, publish_pack_path) = derive_social_auxiliary_paths(article_path); - let cover_url = extract_cover_url_from_markdown(article_content) - .unwrap_or_else(|| SOCIAL_POST_FALLBACK_COVER_URL.to_string()); - let cover_prompt = - extract_detail_value(article_content, "提示词").unwrap_or_else(|| "未提供".to_string()); - let cover_size = extract_detail_value(article_content, "尺寸") - .unwrap_or_else(|| SOCIAL_POST_DEFAULT_IMAGE_SIZE.to_string()); - let cover_status = extract_detail_value(article_content, "状态").unwrap_or_else(|| { - if cover_url == SOCIAL_POST_FALLBACK_COVER_URL { - "失败".to_string() - } else { - "成功".to_string() - } - }); - let cover_remark = extract_detail_value(article_content, "备注").unwrap_or_else(|| { - if cover_status == "失败" { - SOCIAL_POST_FALLBACK_COVER_NOTE.to_string() - } else { - "".to_string() - } - }); - - let cover_meta = serde_json::json!({ - "execution_id": execution_id, - "article_path": article_path, - "cover_url": cover_url, - "prompt": cover_prompt, - "size": cover_size, - "status": cover_status, - "remark": cover_remark, - "generated_at": Utc::now().to_rfc3339(), - }); - - let publish_pack = serde_json::json!({ - "execution_id": execution_id, - "pipeline": ["topic_select", "write_mode", "publish_confirm"], - "article_path": article_path, - "cover_meta_path": cover_meta_path, - "source_input": user_input, - "recommended_channels": ["xiaohongshu", "wechat"], - "summary": summarize_social_content(article_content), - "generated_at": Utc::now().to_rfc3339(), - }); - - vec![ - ( - cover_meta_path, - serde_json::to_string_pretty(&cover_meta).unwrap_or_else(|_| cover_meta.to_string()), - ), - ( - publish_pack_path, - serde_json::to_string_pretty(&publish_pack) - .unwrap_or_else(|_| publish_pack.to_string()), - ), - ] -} - -fn build_write_file_block(file_path: &str, file_content: &str) -> String { - format!("\n{file_content}\n") -} - -fn build_social_post_file_path(user_input: &str, execution_id: &str) -> String { - let timestamp = Utc::now().format("%Y%m%d-%H%M%S"); - let slug = build_social_post_slug(user_input); - let suffix = build_execution_suffix(execution_id); - format!("{SOCIAL_POST_OUTPUT_DIR}/{timestamp}-{slug}-{suffix}.md") -} - -fn build_social_post_slug(user_input: &str) -> String { - let mut normalized = String::new(); - let mut last_was_dash = false; - - for ch in user_input.chars() { - if ch.is_ascii_alphanumeric() { - normalized.push(ch.to_ascii_lowercase()); - last_was_dash = false; - continue; - } - - if !last_was_dash { - normalized.push('-'); - last_was_dash = true; - } - } - - let trimmed = normalized.trim_matches('-'); - let truncated: String = trimmed.chars().take(24).collect(); - if truncated.is_empty() { - "post".to_string() - } else { - truncated - } -} - -fn build_execution_suffix(execution_id: &str) -> String { - let normalized: String = execution_id - .chars() - .filter(|ch| ch.is_ascii_alphanumeric()) - .take(6) - .collect(); - if normalized.is_empty() { - "run".to_string() - } else { - normalized.to_ascii_lowercase() - } -} - -fn build_social_tool_event_id(execution_id: &str, file_path: &str) -> String { - let mut hash: u32 = 0x811c9dc5; - for byte in file_path.as_bytes() { - hash ^= u32::from(*byte); - hash = hash.wrapping_mul(0x01000193); - } - format!("social-write-{execution_id}-{hash:08x}") -} - -fn emit_social_write_file_events( - app_handle: &tauri::AppHandle, - execution_id: &str, - file_path: &str, - file_content: &str, -) { - 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(), - tool_id: tool_id.clone(), - arguments: Some(arguments), - }; - if let Err(err) = app_handle.emit(&event_name, &tool_start) { - 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 { - success: true, - output: format!("写入社媒文稿: {file_path}"), - error: None, - images: None, - metadata: Some(tool_end_metadata), - }, - }; - if let Err(err) = app_handle.emit(&event_name, &tool_end) { - tracing::warn!("[execute_skill] 发送社媒写入工具完成事件失败: {}", err); - } -} - /// 执行 Skill /// /// 加载并执行指定的 Skill,使用 Aster Agent 系统提供完整的工具集支持。 @@ -609,621 +67,22 @@ pub async fn execute_skill( execution_id: Option, session_id: Option, ) -> Result { - // 生成执行 ID,并优先复用前端会话 ID(提升 /skill 与主会话上下文一致性) - let execution_id = execution_id.unwrap_or_else(|| Uuid::new_v4().to_string()); - let session_id = session_id.unwrap_or_else(|| format!("skill-exec-{}", Uuid::new_v4())); - let inferred_gate_key = - infer_theme_workbench_gate_key(skill_name.as_str(), user_input.as_str()); - let memory_profile_prompt = build_memory_profile_prompt(&config_manager.config()); - let tracker = ExecutionTracker::new(db.inner().clone()); - - tracker - .with_run_custom( - RunSource::Skill, - Some(skill_name.clone()), - Some(session_id.clone()), - Some(serde_json::json!({ - "execution_id": execution_id.clone(), - "skill_name": skill_name.clone(), - "gate_key": inferred_gate_key, - "provider_override": provider_override.clone(), - "model_override": model_override.clone(), - })), - async { - tracing::info!( - "[execute_skill] 开始执行 Skill: name={}, execution_id={}, session_id={}, provider_override={:?}, model_override={:?}", + execute_named_skill( + &app_handle, + db.inner(), + api_key_provider_service.inner(), + config_manager.inner(), + aster_state.inner(), + SkillExecutionRequest { skill_name, + user_input, + provider_override, + model_override, execution_id, session_id, - provider_override, - model_override - ); - - // 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( - SKILL_ERR_EXECUTE_FAILED, - format!("Skill '{skill_name}' 已禁用模型调用,无法执行"), - )); - } - - // 2. 创建 TauriExecutionCallback - let callback = TauriExecutionCallback::new(app_handle.clone(), execution_id.clone()); - - // 3. 初始化 Agent(如果未初始化) - if !aster_state.is_initialized().await { - tracing::info!("[execute_skill] Agent 未初始化,开始初始化..."); - aster_state.init_agent_with_db(&db).await.map_err(|e| { - format_skill_error( - SKILL_ERR_SESSION_INIT_FAILED, - format!("初始化 Agent 失败: {e}"), - ) - })?; - tracing::info!("[execute_skill] Agent 初始化完成"); - } - ensure_browser_mcp_tools_registered(aster_state.inner()) - .await - .map_err(|e| { - format_skill_error( - SKILL_ERR_SESSION_INIT_FAILED, - format!("注册浏览器工具失败: {e}"), - ) - })?; - ensure_social_image_tool_registered(aster_state.inner(), config_manager.inner()) - .await - .map_err(|e| { - format_skill_error( - SKILL_ERR_SESSION_INIT_FAILED, - format!("注册社媒生图工具失败: {e}"), - ) - })?; - ensure_creation_task_tools_registered( - aster_state.inner(), - db.inner(), - api_key_provider_service.inner(), - &app_handle, - ) - .await - .map_err(|e| { - format_skill_error( - SKILL_ERR_SESSION_INIT_FAILED, - format!("注册创作任务工具失败: {e}"), - ) - })?; - - // 4. 配置 Provider(从凭证池选择,支持 fallback) - let preferred_provider = provider_override - .clone() - .or_else(|| skill.provider.clone()) - .unwrap_or_else(|| "anthropic".to_string()); - - let preferred_model = model_override - .clone() - .or_else(|| skill.model.clone()) - .unwrap_or_else(|| "claude-sonnet-4-20250514".to_string()); - - // 支持工具调用的 Provider fallback 列表 - // 注意:provider 名称需要与 ProviderType::FromStr 匹配 - let fallback_providers: Vec<(&str, &str)> = vec![ - ("anthropic", "claude-sonnet-4-20250514"), - ("openai", "gpt-4o"), - ("gemini", "gemini-2.0-flash"), - ]; - - let mut configure_result = aster_state - .configure_provider_from_pool(&db, &preferred_provider, &preferred_model, &session_id) - .await; - - if configure_result.is_err() { - tracing::warn!( - "[execute_skill] 首选 Provider {} 配置失败: {:?},尝试 fallback", - preferred_provider, - configure_result.as_ref().err() - ); - - for (fb_provider, fb_model) in &fallback_providers { - if *fb_provider == preferred_provider { - continue; - } - match aster_state - .configure_provider_from_pool(&db, fb_provider, fb_model, &session_id) - .await - { - Ok(config) => { - tracing::info!( - "[execute_skill] Fallback 到 {} / {} 成功", - fb_provider, - fb_model - ); - configure_result = Ok(config); - break; - } - Err(e) => { - tracing::warn!("[execute_skill] Fallback {} 也失败: {}", fb_provider, e); - } - } - } - } - - let configured_provider = configure_result.map_err(|e| { - format_skill_error( - SKILL_ERR_PROVIDER_UNAVAILABLE, - format!( - "无法配置任何可用的 Provider(需要支持工具调用的 Provider,如 Anthropic、OpenAI 或 Google): {e}" - ), - ) - })?; - - let resolved_provider = configured_provider.provider_name.clone(); - let resolved_model = configured_provider.model_name.clone(); - - tracing::info!( - "[execute_skill] Provider 配置成功: requested={} / {}, resolved={} / {}", - preferred_provider, - preferred_model, - resolved_provider, - resolved_model - ); - - // 5. 根据 execution_mode 分支执行 - if skill.execution_mode == "workflow" && !skill.workflow_steps.is_empty() { - // ========== Workflow 模式:按步骤顺序执行 ========== - execute_skill_workflow( - &app_handle, - &aster_state, - &skill, - &user_input, - &execution_id, - &session_id, - &callback, - memory_profile_prompt.as_deref(), - ) - .await - } else { - // ========== Prompt 模式:单次执行 ========== - execute_skill_prompt( - &app_handle, - &aster_state, - &skill, - &user_input, - &execution_id, - &session_id, - &callback, - memory_profile_prompt.as_deref(), - ) - .await - } - }, - |result| match result { - Ok(exec_result) if exec_result.success => { - let artifact_paths = if skill_name == SOCIAL_POST_WITH_COVER_SKILL_NAME { - collect_social_artifact_paths_from_output(exec_result.output.as_deref()) - } else { - Vec::new() - }; - let metadata = if skill_name == SOCIAL_POST_WITH_COVER_SKILL_NAME { - serde_json::json!({ - "skill_name": skill_name, - "execution_id": execution_id, - "workflow": "social_content_pipeline_v1", - "version_id": execution_id, - "stages": ["topic_select", "write_mode", "publish_confirm"], - "artifact_paths": artifact_paths, - "provider_override": provider_override, - "model_override": model_override, - "requested_provider": provider_override, - "requested_model": model_override, - }) - } else { - serde_json::json!({ - "skill_name": skill_name, - "execution_id": execution_id, - "provider_override": provider_override, - "model_override": model_override, - "requested_provider": provider_override, - "requested_model": model_override, - }) - }; - RunFinishDecision { - status: crate::database::dao::agent_run::AgentRunStatus::Success, - error_code: None, - error_message: None, - metadata: Some(metadata), - } - } - Ok(exec_result) => RunFinishDecision { - status: crate::database::dao::agent_run::AgentRunStatus::Error, - error_code: Some("skill_execute_failed".to_string()), - error_message: exec_result.error.clone(), - metadata: Some(serde_json::json!({ - "skill_name": skill_name, - "execution_id": execution_id, - "success": false, - "provider_override": provider_override, - "model_override": model_override, - "requested_provider": provider_override, - "requested_model": model_override, - })), - }, - Err(err) => RunFinishDecision { - status: crate::database::dao::agent_run::AgentRunStatus::Error, - error_code: Some("skill_execute_failed".to_string()), - error_message: Some(err.clone()), - metadata: Some(serde_json::json!({ - "skill_name": skill_name, - "execution_id": execution_id, - "provider_override": provider_override, - "model_override": model_override, - "requested_provider": provider_override, - "requested_model": model_override, - })), - }, - }, - ) - .await -} - -/// Prompt 模式执行(单步) -async fn execute_skill_prompt( - app_handle: &tauri::AppHandle, - aster_state: &AsterAgentState, - skill: &lime_skills::LoadedSkillDefinition, - user_input: &str, - execution_id: &str, - session_id: &str, - callback: &TauriExecutionCallback, - memory_profile_prompt: Option<&str>, -) -> Result { - // 发送步骤开始事件 - callback.on_step_start("main", &skill.display_name, 1, 1); - - // 构建 SessionConfig - let mut combined_prompt = skill.markdown_content.clone(); - if let Some(memory_prompt) = memory_profile_prompt { - combined_prompt = format!("{combined_prompt}\n\n{memory_prompt}"); - } - - let session_config = SessionConfigBuilder::new(session_id) - .system_prompt(combined_prompt) - .include_context_trace(true) - .build(); - - let user_message = Message::user().with_text(user_input); - - // 获取 Agent 并执行 - let agent_arc = aster_state.get_agent_arc(); - let guard = agent_arc.read().await; - let agent = guard.as_ref().ok_or_else(|| { - format_skill_error(SKILL_ERR_SESSION_INIT_FAILED, "Agent not initialized") - })?; - - let cancel_token = aster_state.create_cancel_token(session_id).await; - let stream_result = agent - .reply(user_message, session_config, Some(cancel_token.clone())) - .await; - - let mut final_output = String::new(); - 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) => { - while let Some(event_result) = stream.next().await { - match event_result { - Ok(agent_event) => { - let tauri_events = convert_agent_event(agent_event); - 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); - } - if let Err(e) = app_handle.emit(&event_name, &tauri_event) { - tracing::error!("[execute_skill] 发送事件失败: {}", e); - } - } - } - Err(e) => { - has_error = true; - error_message = Some(format_skill_error( - SKILL_ERR_STREAM_FAILED, - format!("Stream error: {e}"), - )); - tracing::error!("[execute_skill] 流处理错误: {}", e); - } - } - } - - let done_event = TauriAgentEvent::FinalDone { usage: None }; - if let Err(e) = app_handle.emit(&event_name, &done_event) { - tracing::error!("[execute_skill] 发送完成事件失败: {}", e); - } - } - Err(e) => { - has_error = true; - error_message = Some(format_skill_error( - SKILL_ERR_STREAM_FAILED, - format!("Agent error: {e}"), - )); - tracing::error!("[execute_skill] Agent 错误: {}", e); - } - } - - aster_state.remove_cancel_token(session_id).await; - - if has_error { - let err_msg = error_message - .unwrap_or_else(|| format_skill_error(SKILL_ERR_EXECUTE_FAILED, "Unknown error")); - callback.on_step_error("main", &err_msg, false); - callback.on_complete(false, None, Some(&err_msg)); - - Ok(SkillExecutionResult { - success: false, - output: None, - error: Some(err_msg.clone()), - steps_completed: vec![StepResult { - step_id: "main".to_string(), - step_name: skill.display_name.clone(), - success: false, - output: None, - error: Some(err_msg), - }], - }) - } else { - let normalized_output = normalize_social_post_output( - &skill.skill_name, - user_input, - execution_id, - &final_output, - ); - let output_for_return = if let Some(ref social_output) = normalized_output { - emit_social_write_file_events( - app_handle, - execution_id, - &social_output.file_path, - &social_output.file_content, - ); - for (artifact_path, artifact_content) in build_social_auxiliary_file_payloads( - execution_id, - user_input, - &social_output.file_path, - &social_output.file_content, - ) { - emit_social_write_file_events( - app_handle, - execution_id, - &artifact_path, - &artifact_content, - ); - } - social_output.final_output.clone() - } else { - final_output.clone() - }; - - callback.on_step_complete("main", &output_for_return); - callback.on_complete(true, Some(&output_for_return), None); - - Ok(SkillExecutionResult { - success: true, - output: Some(output_for_return.clone()), - error: None, - steps_completed: vec![StepResult { - step_id: "main".to_string(), - step_name: skill.display_name.clone(), - success: true, - output: Some(output_for_return), - error: None, - }], - }) - } -} - -/// Workflow 模式执行(多步骤顺序执行) -async fn execute_skill_workflow( - app_handle: &tauri::AppHandle, - aster_state: &AsterAgentState, - skill: &lime_skills::LoadedSkillDefinition, - user_input: &str, - execution_id: &str, - session_id: &str, - callback: &TauriExecutionCallback, - memory_profile_prompt: Option<&str>, -) -> Result { - let steps = &skill.workflow_steps; - let total_steps = steps.len(); - let event_name = format!("skill-exec-{execution_id}"); - let mut steps_completed = Vec::new(); - let mut accumulated_context = user_input.to_string(); - let mut final_output = String::new(); - - tracing::info!( - "[execute_skill_workflow] 开始 workflow 执行: steps={}, skill={}", - total_steps, - skill.skill_name - ); - - for (idx, step) in steps.iter().enumerate() { - let step_num = idx + 1; - - // 发送步骤开始事件 - callback.on_step_start(&step.id, &step.name, step_num, total_steps); - - tracing::info!( - "[execute_skill_workflow] 执行步骤 {}/{}: id={}, name={}", - step_num, - total_steps, - step.id, - step.name - ); - - // 构建该步骤的 system_prompt:基础 skill prompt + 步骤 prompt - let step_system_prompt = format!( - "{}\n\n---\n\n## 当前步骤: {} ({}/{})\n\n{}", - skill.markdown_content, step.name, step_num, total_steps, step.prompt - ); - - let step_prompt_with_memory = if let Some(memory_prompt) = memory_profile_prompt { - format!("{step_system_prompt}\n\n{memory_prompt}") - } else { - step_system_prompt - }; - - let step_session_id = format!("{}-step-{}", session_id, step.id); - let session_config = SessionConfigBuilder::new(&step_session_id) - .system_prompt(step_prompt_with_memory) - .include_context_trace(true) - .build(); - - // 用户消息 = 原始输入 + 前序步骤的累积上下文 - let step_input = if idx == 0 { - accumulated_context.clone() - } else { - format!("原始需求:{user_input}\n\n前序步骤输出:\n{accumulated_context}") - }; - - let user_message = Message::user().with_text(&step_input); - - // 获取 Agent 并执行 - let agent_arc = aster_state.get_agent_arc(); - let guard = agent_arc.read().await; - let agent = guard.as_ref().ok_or_else(|| { - format_skill_error(SKILL_ERR_SESSION_INIT_FAILED, "Agent not initialized") - })?; - - let cancel_token = aster_state.create_cancel_token(&step_session_id).await; - let stream_result = agent - .reply(user_message, session_config, Some(cancel_token.clone())) - .await; - - 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) => { - while let Some(event_result) = stream.next().await { - match event_result { - Ok(agent_event) => { - let tauri_events = convert_agent_event(agent_event); - 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); - } - if let Err(e) = app_handle.emit(&event_name, &tauri_event) { - tracing::error!("[execute_skill_workflow] 发送事件失败: {}", e); - } - } - } - Err(e) => { - step_error = Some(format!("Stream error: {e}")); - tracing::error!( - "[execute_skill_workflow] 步骤 {} 流处理错误: {}", - step.id, - e - ); - break; - } - } - } - } - Err(e) => { - step_error = Some(format!("Agent error: {e}")); - tracing::error!( - "[execute_skill_workflow] 步骤 {} Agent 错误: {}", - step.id, - e - ); - } - } - - aster_state.remove_cancel_token(&step_session_id).await; - - if let Some(err) = &step_error { - callback.on_step_error(&step.id, err, false); - steps_completed.push(StepResult { - step_id: step.id.clone(), - step_name: step.name.clone(), - success: false, - output: None, - error: Some(err.clone()), - }); - - // 步骤失败,终止 workflow - let err_msg = format_skill_error( - SKILL_ERR_EXECUTE_FAILED, - format!("步骤 '{}' 执行失败: {}", step.name, err), - ); - callback.on_complete(false, None, Some(&err_msg)); - - let done_event = TauriAgentEvent::FinalDone { usage: None }; - let _ = app_handle.emit(&event_name, &done_event); - - return Ok(SkillExecutionResult { - success: false, - output: None, - error: Some(err_msg), - steps_completed, - }); - } - - // 步骤成功 - callback.on_step_complete(&step.id, &step_output); - steps_completed.push(StepResult { - step_id: step.id.clone(), - step_name: step.name.clone(), - success: true, - output: Some(step_output.clone()), - error: None, - }); - - // 累积上下文供下一步使用 - accumulated_context = step_output.clone(); - final_output = step_output; - } - - // 所有步骤完成 - callback.on_complete(true, Some(&final_output), None); - - let done_event = TauriAgentEvent::FinalDone { usage: None }; - let _ = app_handle.emit(&event_name, &done_event); - - tracing::info!( - "[execute_skill_workflow] Workflow 执行完成: skill={}, steps_completed={}", - skill.skill_name, - steps_completed.len() - ); - - Ok(SkillExecutionResult { - success: true, - output: Some(final_output), - error: None, - steps_completed, - }) + }, + ) + .await } /// 列出可执行的 Skills @@ -1243,46 +102,7 @@ async fn execute_skill_workflow( /// - 4.5: 过滤未通过标准校验的 skills #[tauri::command] pub async fn list_executable_skills() -> Result, String> { - let skill_roots = get_skill_roots(); - if skill_roots.is_empty() { - return Err(format_skill_error( - SKILL_ERR_CATALOG_UNAVAILABLE, - "无法获取 Skills 目录", - )); - } - - 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 - .into_iter() - .filter(|s| !s.disable_model_invocation) - .map(|s| ExecutableSkillInfo { - name: s.skill_name, - display_name: s.display_name, - description: s.description, - execution_mode: s.execution_mode.clone(), - has_workflow: s.execution_mode == "workflow", - provider: s.provider, - model: s.model, - argument_hint: s.argument_hint, - }) - .collect(); - - tracing::info!( - "[list_executable_skills] 返回 {} 个可执行 Skills", - executable_skills.len() - ); - - Ok(executable_skills) + list_executable_skill_catalog() } /// 获取 Skill 详情 @@ -1303,497 +123,5 @@ pub async fn list_executable_skills() -> Result, String /// - 5.4: skill 不存在时返回错误 #[tauri::command] pub async fn get_skill_detail(skill_name: String) -> Result { - // 查找 skill(Requirements 5.1, 5.4) - let skill = find_skill_by_name(&skill_name).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 { - basic: ExecutableSkillInfo { - name: skill.skill_name, - display_name: skill.display_name, - description: skill.description, - execution_mode: skill.execution_mode.clone(), - has_workflow: skill.execution_mode == "workflow", - provider: skill.provider, - model: skill.model, - argument_hint: skill.argument_hint, - }, - markdown_content: skill.markdown_content, - workflow_steps: if skill.workflow_steps.is_empty() { - None - } else { - Some( - skill - .workflow_steps - .iter() - .map(|s| WorkflowStepInfo { - id: s.id.clone(), - name: s.name.clone(), - dependencies: Vec::new(), - }) - .collect(), - ) - }, - allowed_tools: skill.allowed_tools, - when_to_use: skill.when_to_use, - }; - - tracing::info!("[get_skill_detail] 返回 Skill 详情: name={}", skill_name); - - Ok(detail) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_executable_skill_info_serialization() { - let info = ExecutableSkillInfo { - name: "test-skill".to_string(), - display_name: "Test Skill".to_string(), - description: "A test skill".to_string(), - execution_mode: "prompt".to_string(), - has_workflow: false, - provider: None, - model: None, - argument_hint: Some("Enter your query".to_string()), - }; - - let json = serde_json::to_string(&info).unwrap(); - assert!(json.contains("test-skill")); - assert!(json.contains("Test Skill")); - } - - #[test] - fn test_skill_execution_result_serialization() { - let result = SkillExecutionResult { - success: true, - output: Some("Hello, world!".to_string()), - error: None, - steps_completed: vec![StepResult { - step_id: "step-1".to_string(), - step_name: "Process".to_string(), - success: true, - output: Some("Done".to_string()), - error: None, - }], - }; - - let json = serde_json::to_string(&result).unwrap(); - assert!(json.contains("\"success\":true")); - assert!(json.contains("Hello, world!")); - assert!(json.contains("step-1")); - } - - #[test] - fn test_skill_detail_info_serialization() { - let detail = SkillDetailInfo { - basic: ExecutableSkillInfo { - name: "workflow-skill".to_string(), - display_name: "Workflow Skill".to_string(), - description: "A workflow skill".to_string(), - execution_mode: "workflow".to_string(), - has_workflow: true, - provider: Some("claude".to_string()), - model: Some("claude-sonnet-4-5-20250514".to_string()), - argument_hint: None, - }, - markdown_content: "# Workflow Skill\n\nThis is a workflow skill.".to_string(), - workflow_steps: Some(vec![ - WorkflowStepInfo { - id: "step-1".to_string(), - name: "Initialize".to_string(), - dependencies: vec![], - }, - WorkflowStepInfo { - id: "step-2".to_string(), - name: "Process".to_string(), - dependencies: vec!["step-1".to_string()], - }, - ]), - allowed_tools: Some(vec!["read_file".to_string(), "write_file".to_string()]), - when_to_use: Some("Use this skill for complex workflows".to_string()), - }; - - let json = serde_json::to_string(&detail).unwrap(); - assert!(json.contains("workflow-skill")); - assert!(json.contains("workflow_steps")); - assert!(json.contains("step-1")); - assert!(json.contains("step-2")); - } - - #[test] - fn test_parse_skill_frontmatter_basic() { - let content = r#"--- -name: test-skill -description: A test skill -metadata: - lime_model_preference: claude-sonnet-4-5-20250514 - lime_provider_preference: claude ---- - -# Test Skill - -This is the body content. -"#; - let (fm, body) = parse_skill_frontmatter(content); - assert_eq!(fm.name, Some("test-skill".to_string())); - assert_eq!(fm.description, Some("A test skill".to_string())); - assert_eq!(fm.model, Some("claude-sonnet-4-5-20250514".to_string())); - assert_eq!(fm.provider, Some("claude".to_string())); - assert!(body.contains("# Test Skill")); - assert!(body.contains("This is the body content.")); - } - - #[test] - fn test_parse_skill_frontmatter_no_frontmatter() { - let content = "# Just content\nNo frontmatter here."; - let (fm, body) = parse_skill_frontmatter(content); - assert!(fm.name.is_none()); - assert_eq!(body, content); - } - - #[test] - fn test_parse_skill_frontmatter_with_quotes() { - let content = r#"--- -name: "quoted-name" -description: 'single quoted' ---- -Body -"#; - let (fm, _) = parse_skill_frontmatter(content); - assert_eq!(fm.name, Some("quoted-name".to_string())); - assert_eq!(fm.description, Some("single quoted".to_string())); - } - - #[test] - fn test_normalize_social_post_output_wraps_plain_markdown() { - let normalized = normalize_social_post_output( - SOCIAL_POST_WITH_COVER_SKILL_NAME, - "春季上新", - "exec123456", - "# 标题\n\n正文内容", - ) - .expect("should normalize"); - - assert!(normalized - .final_output - .contains("\n# 标题\n\n正文\n"; - let normalized = normalize_social_post_output( - SOCIAL_POST_WITH_COVER_SKILL_NAME, - "春季上新", - "exec123456", - raw_output, - ) - .expect("should normalize"); - - assert_eq!(normalized.file_path, "social-posts/custom-post.md"); - assert!(normalized - .final_output - .contains("social-posts/custom-post.md")); - assert!(normalized.file_content.contains("# 标题")); - assert!(normalized.file_content.contains("![封面图](")); - } - - #[test] - fn test_normalize_social_post_output_injects_missing_path() { - let raw_output = "前置说明\n\n# 标题\n\n正文\n\n后置说明"; - let normalized = normalize_social_post_output( - SOCIAL_POST_WITH_COVER_SKILL_NAME, - "spring launch", - "exec123456", - raw_output, - ) - .expect("should normalize"); - - assert!(normalized.final_output.contains("前置说明")); - assert!(normalized.final_output.contains("后置说明")); - assert!(normalized - .final_output - .contains("\n# 标题\n\n正文\n"; - let paths = collect_social_artifact_paths_from_output(Some(output)); - assert_eq!(paths.len(), 3); - assert_eq!(paths[0], "social-posts/demo.md"); - assert!(paths[1].ends_with(".cover.json")); - assert!(paths[2].ends_with(".publish-pack.json")); - } - - #[test] - fn test_build_social_post_slug_fallback_to_post() { - assert_eq!(build_social_post_slug(""), "post"); - assert_eq!(build_social_post_slug("!!!"), "post"); - assert_eq!( - build_social_post_slug("Spring Launch 2026"), - "spring-launch-2026" - ); - } - - #[test] - fn test_parse_allowed_tools() { - assert_eq!(parse_allowed_tools(None), None); - assert_eq!(parse_allowed_tools(Some("")), None); - assert_eq!( - parse_allowed_tools(Some("tool1")), - Some(vec!["tool1".to_string()]) - ); - assert_eq!( - parse_allowed_tools(Some("tool1, tool2, tool3")), - Some(vec![ - "tool1".to_string(), - "tool2".to_string(), - "tool3".to_string() - ]) - ); - } - - #[test] - fn test_parse_boolean() { - assert!(!parse_boolean(None, false)); - assert!(parse_boolean(None, true)); - assert!(parse_boolean(Some("true"), false)); - assert!(parse_boolean(Some("TRUE"), false)); - assert!(parse_boolean(Some("1"), false)); - assert!(parse_boolean(Some("yes"), false)); - assert!(!parse_boolean(Some("false"), true)); - assert!(!parse_boolean(Some("no"), true)); - } - - #[test] - fn test_load_skill_from_file() { - use tempfile::TempDir; - - let temp_dir = TempDir::new().unwrap(); - let skill_dir = temp_dir.path().join("my-skill"); - std::fs::create_dir(&skill_dir).unwrap(); - - let skill_file = skill_dir.join("SKILL.md"); - std::fs::write( - &skill_file, - r#"--- -name: my-skill -description: Test skill description -allowed-tools: tool1, tool2 -metadata: - lime_model_preference: gpt-4 - lime_provider_preference: openai ---- - -# My Skill - -Instructions here. -"#, - ) - .unwrap(); - - let skill = load_skill_from_file("my-skill", &skill_file).unwrap(); - - assert_eq!(skill.skill_name, "my-skill"); - assert_eq!(skill.display_name, "my-skill"); - assert_eq!(skill.description, "Test skill description"); - assert_eq!( - skill.allowed_tools, - Some(vec!["tool1".to_string(), "tool2".to_string()]) - ); - assert_eq!(skill.model, Some("gpt-4".to_string())); - assert_eq!(skill.provider, Some("openai".to_string())); - assert!(!skill.disable_model_invocation); - assert_eq!(skill.execution_mode, "prompt"); - 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: - lime_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.lime_workflow_ref"))); - assert!(skill.workflow_steps.is_empty()); - } - - #[test] - fn test_load_skills_from_directory() { - use tempfile::TempDir; - - let temp_dir = TempDir::new().unwrap(); - let skills_dir = temp_dir.path(); - - // 创建 skill 1 - let skill1_dir = skills_dir.join("skill-one"); - std::fs::create_dir(&skill1_dir).unwrap(); - std::fs::write( - skill1_dir.join("SKILL.md"), - r#"--- -name: skill-one -description: First skill ---- -Content 1 -"#, - ) - .unwrap(); - - // 创建 skill 2 - let skill2_dir = skills_dir.join("skill-two"); - std::fs::create_dir(&skill2_dir).unwrap(); - std::fs::write( - skill2_dir.join("SKILL.md"), - r#"--- -name: skill-two -description: Second skill -disable-model-invocation: true ---- -Content 2 -"#, - ) - .unwrap(); - - let skills = load_skills_from_directory(skills_dir); - - assert_eq!(skills.len(), 2); - let names: Vec<_> = skills.iter().map(|s| s.skill_name.as_str()).collect(); - assert!(names.contains(&"skill-one")); - assert!(names.contains(&"skill-two")); - - // 验证 disable_model_invocation 被正确解析 - let skill_two = skills.iter().find(|s| s.skill_name == "skill-two").unwrap(); - assert!(skill_two.disable_model_invocation); - } - - #[test] - fn test_load_skills_from_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: - lime_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")); - assert!(skills.is_empty()); - } - - #[test] - fn test_bundled_social_post_with_cover_skill_contract() { - let skill_file = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR")) - .join("resources/default-skills/social_post_with_cover/SKILL.md"); - - assert!(skill_file.exists()); - let content = std::fs::read_to_string(&skill_file).unwrap(); - let skill = load_skill_from_file("social_post_with_cover", &skill_file).unwrap(); - - 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![ - "social_generate_cover_image".to_string(), - "search_query".to_string(), - ]) - ); - assert!(content.contains("); - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ExecuteHooksRequest { - pub trigger: HookTrigger, - pub context: HookContextData, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HookContextData { - pub session_id: String, - pub tool_name: Option, - pub tool_parameters: Option>, - pub tool_result: Option, - pub message_content: Option, - pub message_count: usize, - pub error_info: Option, - pub metadata: HashMap, -} - -impl From for HookContext { - fn from(data: HookContextData) -> Self { - Self { - session_id: data.session_id, - tool_name: data.tool_name, - tool_parameters: data.tool_parameters, - tool_result: data.tool_result, - message_content: data.message_content, - message_count: data.message_count, - error_info: data.error_info, - metadata: data.metadata, - } - } -} - -#[tauri::command] -pub async fn execute_hooks( - hooks_service: State<'_, ToolHooksServiceState>, - request: ExecuteHooksRequest, -) -> Result<(), String> { - debug!( - "执行钩子: {:?} (会话: {})", - request.trigger, request.context.session_id - ); - let context: HookContext = request.context.into(); - hooks_service.0.execute_hooks(request.trigger, &context)?; - info!("钩子执行完成"); - Ok(()) -} - -#[tauri::command] -pub async fn add_hook_rule( - hooks_service: State<'_, ToolHooksServiceState>, - rule: HookRule, -) -> Result<(), String> { - debug!("添加钩子规则: {}", rule.name); - hooks_service.0.add_hook_rule(rule.clone())?; - info!("钩子规则添加成功: {}", rule.name); - Ok(()) -} - -#[tauri::command] -pub async fn remove_hook_rule( - hooks_service: State<'_, ToolHooksServiceState>, - rule_id: String, -) -> Result<(), String> { - debug!("移除钩子规则: {}", rule_id); - hooks_service.0.remove_hook_rule(&rule_id)?; - info!("钩子规则移除成功: {}", rule_id); - Ok(()) -} - -#[tauri::command] -pub async fn toggle_hook_rule( - hooks_service: State<'_, ToolHooksServiceState>, - rule_id: String, - enabled: bool, -) -> Result<(), String> { - debug!("切换钩子规则状态: {} -> {}", rule_id, enabled); - hooks_service.0.toggle_hook_rule(&rule_id, enabled)?; - info!("钩子规则状态切换成功: {} -> {}", rule_id, enabled); - Ok(()) -} - -#[tauri::command] -pub async fn get_hook_rules( - hooks_service: State<'_, ToolHooksServiceState>, -) -> Result, String> { - debug!("获取所有钩子规则"); - let rules = hooks_service.0.get_hook_rules()?; - info!("获取到 {} 个钩子规则", rules.len()); - Ok(rules) -} - -#[tauri::command] -pub async fn get_hook_execution_stats( - hooks_service: State<'_, ToolHooksServiceState>, -) -> Result, String> { - debug!("获取钩子执行统计"); - let stats = hooks_service.0.get_execution_stats()?; - info!("获取到 {} 个规则的执行统计", stats.len()); - Ok(stats) -} - -#[tauri::command] -pub async fn clear_hook_execution_stats( - hooks_service: State<'_, ToolHooksServiceState>, -) -> Result<(), String> { - debug!("清理钩子执行统计"); - hooks_service.0.clear_execution_stats()?; - info!("钩子执行统计清理完成"); - Ok(()) -} diff --git a/src-tauri/src/commands/unified_chat_cmd.rs b/src-tauri/src/commands/unified_chat_cmd.rs deleted file mode 100644 index 80aae9145..000000000 --- a/src-tauri/src/commands/unified_chat_cmd.rs +++ /dev/null @@ -1,847 +0,0 @@ -//! 统一对话命令模块 -//! -//! 提供统一的对话 API,支持多种对话模式: -//! - Agent: AI Agent 模式,支持工具调用 -//! - General: 通用对话模式,纯文本 -//! - Creator: 内容创作模式,支持画布输出 -//! -//! ## 设计原则 -//! - 单一入口:所有对话场景使用同一套 API -//! - 模式化设计:通过 ChatMode 区分不同场景 -//! - Aster 引擎:底层使用 Aster Agent 处理对话 -//! -//! ## 参考文档 -//! - `docs/prd/chat-architecture-redesign.md` - -use crate::agent::aster_state::SessionConfigBuilder; -use crate::agent::{AsterAgentState, TauriAgentEvent}; -use crate::commands::aster_agent_cmd::ensure_browser_mcp_tools_registered; -use crate::config::GlobalConfigManagerState; -use crate::database::dao::chat::{ChatDao, ChatMessage, ChatMode, ChatSession}; -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::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 lime_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}; -use tracing::Instrument; - -const CODE_EXECUTION_EXTENSION_NAME: &str = "code_execution"; - -// ============================================================================ -// 请求/响应结构 -// ============================================================================ - -/// 创建会话请求 -#[derive(Debug, Deserialize)] -pub struct CreateSessionRequest { - /// 对话模式 - pub mode: ChatMode, - /// 会话标题(可选) - pub title: Option, - /// 系统提示词(可选) - pub system_prompt: Option, - /// Provider 类型(可选) - pub provider_type: Option, - /// 模型名称(可选) - pub model: Option, - /// 扩展元数据(可选) - pub metadata: Option, -} - -/// 发送消息请求 -#[derive(Debug, Deserialize)] -pub struct SendMessageRequest { - /// 会话 ID - #[serde(alias = "sessionId")] - pub session_id: String, - /// 消息内容 - pub message: String, - /// 事件名称(用于前端监听) - #[serde(alias = "eventName")] - pub event_name: String, - /// 图片输入(可选,用于多模态对话) - /// TODO: 实现图片处理逻辑,将图片转换为 Aster Message 的 ImageContent - pub images: Option>, - /// 请求级联网搜索开关 - #[serde(default, alias = "webSearch")] - pub web_search: Option, - /// 联网搜索模式(disabled / allowed / required) - #[serde(default, alias = "searchMode")] - pub search_mode: Option, -} - -/// 图片输入 -#[derive(Debug, Deserialize)] -pub struct ImageInput { - /// Base64 编码的图片数据 - pub data: String, - /// 图片 MIME 类型,如 "image/png", "image/jpeg" - pub media_type: String, -} - -/// 会话信息响应 -#[derive(Debug, Serialize)] -pub struct SessionResponse { - pub id: String, - pub mode: ChatMode, - pub title: Option, - pub model: Option, - pub created_at: String, - pub updated_at: String, - pub message_count: usize, -} - -impl From for SessionResponse { - fn from(session: ChatSession) -> Self { - Self { - id: session.id, - mode: session.mode, - title: session.title, - model: session.model, - created_at: session.created_at, - updated_at: session.updated_at, - message_count: 0, - } - } -} - -// ============================================================================ -// 会话管理命令 -// ============================================================================ - -/// 创建新会话 -/// -/// 统一的会话创建入口,支持所有对话模式 -#[tauri::command] -pub async fn chat_create_session( - db: State<'_, DbConnection>, - agent_state: State<'_, AsterAgentState>, - config_manager: State<'_, GlobalConfigManagerState>, - request: CreateSessionRequest, -) -> Result { - let now = chrono::Utc::now().to_rfc3339(); - let session_id = uuid::Uuid::new_v4().to_string(); - - let config = config_manager.config(); - let working_dir = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from(".")); - let merged_system_prompt = merge_system_prompt_with_web_search( - merge_system_prompt_with_memory_sources( - merge_system_prompt_with_memory_profile(request.system_prompt.clone(), &config), - &config, - &working_dir, - None, - ), - &config, - ); - - // 创建会话 - let session = ChatSession { - id: session_id.clone(), - mode: request.mode, - title: request.title, - system_prompt: merged_system_prompt, - model: request.model.clone(), - provider_type: request.provider_type.clone(), - credential_uuid: None, - metadata: request.metadata, - created_at: now.clone(), - updated_at: now, - }; - - // 保存到数据库(异步化) - { - let db = db.inner().clone(); - let session_clone = session.clone(); - tokio::task::spawn_blocking(move || { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - ChatDao::create_session(&conn, &session_clone).map_err(|e| format!("创建会话失败: {e}")) - }) - .await - .map_err(|e| format!("任务执行失败: {e}"))??; - } - - // 初始化 Aster Agent(如果是 Agent 或 Creator 模式) - if matches!(request.mode, ChatMode::Agent | ChatMode::Creator) { - agent_state.init_agent_with_db(&db).await?; - - // 如果指定了 Provider,配置它 - if let (Some(provider_type), Some(model)) = (&request.provider_type, &request.model) { - agent_state - .configure_provider_from_pool(&db, provider_type, model, &session_id) - .await?; - } - } - - tracing::info!( - "[UnifiedChat] 创建会话: id={}, mode={:?}", - session_id, - request.mode - ); - - Ok(SessionResponse::from(session)) -} - -/// 获取会话列表 -/// -/// 可选按模式过滤 -#[tauri::command] -pub async fn chat_list_sessions( - db: State<'_, DbConnection>, - mode: Option, -) -> Result, String> { - let db = db.inner().clone(); - tokio::task::spawn_blocking(move || { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let sessions = - ChatDao::list_sessions(&conn, mode).map_err(|e| format!("获取会话列表失败: {e}"))?; - - let mut result: Vec = Vec::new(); - for session in sessions { - let message_count = ChatDao::get_message_count(&conn, &session.id).unwrap_or(0); - let mut resp = SessionResponse::from(session); - resp.message_count = message_count; - result.push(resp); - } - - Ok(result) - }) - .await - .map_err(|e| format!("任务执行失败: {e}"))? -} - -/// 获取会话详情 -#[tauri::command] -pub async fn chat_get_session( - db: State<'_, DbConnection>, - session_id: String, -) -> Result { - let db = db.inner().clone(); - tokio::task::spawn_blocking(move || { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let session = ChatDao::get_session(&conn, &session_id) - .map_err(|e| format!("获取会话失败: {e}"))? - .ok_or_else(|| "会话不存在".to_string())?; - - let message_count = ChatDao::get_message_count(&conn, &session_id).unwrap_or(0); - let mut resp = SessionResponse::from(session); - resp.message_count = message_count; - - Ok(resp) - }) - .await - .map_err(|e| format!("任务执行失败: {e}"))? -} - -/// 删除会话 -#[tauri::command] -pub async fn chat_delete_session( - db: State<'_, DbConnection>, - session_id: String, -) -> Result { - let db = db.inner().clone(); - tokio::task::spawn_blocking(move || { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let deleted = ChatDao::delete_session(&conn, &session_id) - .map_err(|e| format!("删除会话失败: {e}"))?; - - if deleted { - tracing::info!("[UnifiedChat] 删除会话: id={}", session_id); - } - - Ok(deleted) - }) - .await - .map_err(|e| format!("任务执行失败: {e}"))? -} - -/// 重命名会话 -#[tauri::command] -pub async fn chat_rename_session( - db: State<'_, DbConnection>, - session_id: String, - title: String, -) -> Result<(), String> { - let db = db.inner().clone(); - tokio::task::spawn_blocking(move || { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - ChatDao::update_title(&conn, &session_id, &title) - .map_err(|e| format!("重命名会话失败: {e}"))?; - - tracing::info!( - "[UnifiedChat] 重命名会话: id={}, title={}", - session_id, - title - ); - - Ok(()) - }) - .await - .map_err(|e| format!("任务执行失败: {e}"))? -} - -// ============================================================================ -// 消息管理命令 -// ============================================================================ - -/// 获取会话消息列表 -#[tauri::command] -pub async fn chat_get_messages( - db: State<'_, DbConnection>, - session_id: String, - limit: Option, -) -> Result, String> { - let db = db.inner().clone(); - tokio::task::spawn_blocking(move || { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let messages = ChatDao::get_messages(&conn, &session_id, limit) - .map_err(|e| format!("获取消息失败: {e}"))?; - - Ok(messages) - }) - .await - .map_err(|e| format!("任务执行失败: {e}"))? -} - -/// 发送消息并获取流式响应 -/// -/// 统一的消息发送入口,根据会话模式选择处理方式 -#[tracing::instrument( - name = "chat_send_message", - skip(app, db, agent_state, config_manager, request), - fields( - session_id = %request.session_id, - event_name = %request.event_name, - image_count = request.images.as_ref().map(|items| items.len()).unwrap_or(0) - ) -)] -#[tauri::command] -pub async fn chat_send_message( - app: AppHandle, - db: State<'_, DbConnection>, - agent_state: State<'_, AsterAgentState>, - config_manager: State<'_, GlobalConfigManagerState>, - request: SendMessageRequest, -) -> Result<(), String> { - let start_time = std::time::Instant::now(); - - let image_count = request.images.as_ref().map(|v| v.len()).unwrap_or(0); - tracing::info!( - "[UnifiedChat] 发送消息: session={}, event={}, images={}", - request.session_id, - request.event_name, - image_count - ); - - // TODO: 实现图片处理逻辑,将图片转换为 Aster Message 的 ImageContent - if let Some(images) = &request.images { - for (i, img) in images.iter().enumerate() { - tracing::debug!( - "[UnifiedChat] 图片 {}: media_type={}, data_len={}", - i, - img.media_type, - img.data.len() - ); - } - if !images.is_empty() { - tracing::warn!( - "[UnifiedChat] 图片输入暂未实现,忽略 {} 张图片", - images.len() - ); - } - } - - // 获取会话信息(异步化数据库操作) - let db_start = std::time::Instant::now(); - let session = { - let db = db.inner().clone(); - let session_id = request.session_id.clone(); - tokio::task::spawn_blocking(move || { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - ChatDao::get_session(&conn, &session_id) - .map_err(|e| format!("获取会话失败: {e}"))? - .ok_or_else(|| "会话不存在".to_string()) - }) - .instrument(tracing::info_span!("chat_send_message.load_session")) - .await - .map_err(|e| format!("任务执行失败: {e}"))?? - }; - let db_elapsed = db_start.elapsed(); - tracing::debug!("[UnifiedChat] 数据库查询耗时: {:?}", db_elapsed); - - // 根据模式处理 - let config = config_manager.config(); - apply_web_search_runtime_env(&config); - let working_dir = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from(".")); - let merged_system_prompt = tracing::debug_span!("chat_send_message.prepare_system_prompt") - .in_scope(|| { - merge_system_prompt_with_web_search( - merge_system_prompt_with_memory_sources( - merge_system_prompt_with_memory_profile(session.system_prompt.clone(), &config), - &config, - &working_dir, - None, - ), - &config, - ) - }); - - 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={:?}, 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.search_mode.as_str() - ); - - let result = send_message_with_aster( - &app, - &db, - &agent_state, - &request.session_id, - &request.message, - &request.event_name, - merged_system_prompt.as_deref(), - config.memory.enabled, - &request_tool_policy, - ) - .instrument(tracing::info_span!( - "chat_send_message.dispatch_agent", - effective_web_search = request_tool_policy.effective_web_search, - search_mode = %request_tool_policy.search_mode.as_str() - )) - .await; - - let total_elapsed = start_time.elapsed(); - tracing::info!( - "[UnifiedChat] 消息发送完成: session={}, 总耗时={:?}", - request.session_id, - total_elapsed - ); - - result -} - -/// 使用 Aster Agent 发送消息 -#[tracing::instrument( - name = "send_message_with_aster", - skip(app, db, agent_state, message, system_prompt, request_tool_policy), - fields( - session_id = %session_id, - event_name = %event_name, - message_len = message.len(), - include_context_trace = include_context_trace, - effective_web_search = request_tool_policy.effective_web_search - ) -)] -async fn send_message_with_aster( - app: &AppHandle, - db: &DbConnection, - agent_state: &AsterAgentState, - session_id: &str, - message: &str, - event_name: &str, - system_prompt: Option<&str>, - include_context_trace: bool, - request_tool_policy: &RequestToolPolicy, -) -> Result<(), String> { - let start_time = std::time::Instant::now(); - tracing::info!( - "[UnifiedChat][WebSearchGuard] session={}, effective_web_search={}", - session_id, - request_tool_policy.effective_web_search - ); - - // 确保 Agent 已初始化 - let init_start = std::time::Instant::now(); - async { - if !agent_state.is_initialized().await { - agent_state.init_agent_with_db(db).await?; - } - ensure_browser_mcp_tools_registered(agent_state).await?; - Ok::<(), String>(()) - } - .instrument(tracing::info_span!( - "send_message_with_aster.ensure_agent_ready" - )) - .await?; - let init_elapsed = init_start.elapsed(); - tracing::debug!("[UnifiedChat] Agent 初始化检查耗时: {:?}", init_elapsed); - - // 检查 Provider 是否已配置 - let provider_check_start = std::time::Instant::now(); - let is_provider_configured = - async { Ok::(agent_state.is_provider_configured().await) } - .instrument(tracing::debug_span!( - "send_message_with_aster.check_provider_config" - )) - .await?; - if !is_provider_configured { - return Err("Provider 未配置,请先配置凭证".to_string()); - } - let provider_check_elapsed = provider_check_start.elapsed(); - tracing::debug!( - "[UnifiedChat] Provider 配置检查耗时: {:?}", - provider_check_elapsed - ); - - // 创建取消令牌 - let cancel_token = agent_state.create_cancel_token(session_id).await; - - let effective_system_prompt = merge_system_prompt_with_request_tool_policy( - system_prompt.map(|prompt| prompt.to_string()), - request_tool_policy, - ); - - let user_message = Message::user().with_text(message); - let mut session_config_builder = SessionConfigBuilder::new(session_id); - if let Some(prompt) = effective_system_prompt { - session_config_builder = session_config_builder.system_prompt(prompt); - } - let mut session_config = session_config_builder - .include_context_trace(include_context_trace) - .build(); - - // 获取 Agent 引用 - let agent_arc = agent_state.get_agent_arc(); - let guard = agent_arc.read().await; - let agent = guard.as_ref().ok_or("Agent 未初始化")?; - - let mut removed_extension: Option = None; - if request_tool_policy.requires_web_search() { - let extension_configs = agent.get_extension_configs().await; - if let Some(extension) = extension_configs - .into_iter() - .find(|extension| extension.name() == CODE_EXECUTION_EXTENSION_NAME) - { - match agent.remove_extension(CODE_EXECUTION_EXTENSION_NAME).await { - Ok(_) => { - removed_extension = Some(extension); - tracing::info!( - "[UnifiedChat] 当前会话优先联网搜索,临时关闭 {} 扩展", - CODE_EXECUTION_EXTENSION_NAME - ); - } - Err(error) => { - tracing::warn!( - "[UnifiedChat] 移除 {} 扩展失败: {}", - CODE_EXECUTION_EXTENSION_NAME, - error - ); - } - } - } else { - tracing::info!( - "[UnifiedChat][WebSearchGuard] session={}, 未检测到 {} 扩展,无需移除", - session_id, - CODE_EXECUTION_EXTENSION_NAME - ); - } - } - - // 调用 Agent - let reply_start = std::time::Instant::now(); - let mut web_search_tracker = WebSearchExecutionTracker::default(); - let preflight = execute_web_search_preflight_if_needed( - agent, - session_id, - message, - None, - Some(cancel_token.clone()), - request_tool_policy, - &mut web_search_tracker, - ) - .instrument(tracing::info_span!( - "send_message_with_aster.web_search_preflight" - )) - .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); - } - } - } - Err(error) => { - let error_event = TauriAgentEvent::Error { - message: format!( - "{error}\n尝试记录: {}", - web_search_tracker.format_attempts() - ), - }; - let _ = app.emit(event_name, &error_event); - agent_state.remove_cancel_token(session_id).await; - if let Some(extension) = removed_extension { - if let Err(restore_error) = agent.add_extension(extension).await { - tracing::warn!( - "[UnifiedChat] 预调用失败后恢复 {} 扩展失败: {}", - CODE_EXECUTION_EXTENSION_NAME, - restore_error - ); - } - } - return Err(error); - } - } - - let stream_result = agent - .reply(user_message, session_config, Some(cancel_token.clone())) - .instrument(tracing::info_span!("send_message_with_aster.reply")) - .await; - - 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) => { - while let Some(event_result) = stream - .next() - .instrument(tracing::trace_span!( - "send_message_with_aster.next_stream_event" - )) - .await - { - match event_result { - Ok(agent_event) => { - // 记录首个 chunk 时间(TTFB) - if first_chunk_time.is_none() { - first_chunk_time = Some(std::time::Instant::now()); - let ttfb = first_chunk_time.unwrap() - reply_start; - tracing::info!("[UnifiedChat] TTFB (首字节时间): {:?}", ttfb); - } - chunk_count += 1; - - let tauri_events = convert_agent_event(agent_event); - 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( - request_tool_policy, - tool_id, - tool_name, - ), - TauriAgentEvent::ToolEnd { tool_id, result } => { - web_search_tracker.record_tool_end( - request_tool_policy, - tool_id, - result.success, - result.error.as_deref(), - ); - } - _ => {} - } - if let Err(e) = app.emit(event_name, &tauri_event) { - tracing::error!("[UnifiedChat] 发送事件失败: {}", e); - } - } - } - Err(e) => { - let message = format!("流错误: {e}"); - let error_event = TauriAgentEvent::Error { - message: message.clone(), - }; - let _ = app.emit(event_name, &error_event); - stream_error = Some(message); - } - } - } - - if stream_error.is_none() { - if let Err(validation_error) = - web_search_tracker.validate_web_search_requirement(request_tool_policy) - { - let error_event = TauriAgentEvent::Error { - message: validation_error.clone(), - }; - let _ = app.emit(event_name, &error_event); - stream_error = Some(validation_error); - } - } - - 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 }; - let _ = app.emit(event_name, &done_event); - } - - let stream_elapsed = start_time.elapsed(); - tracing::info!( - "[UnifiedChat] 流式传输完成: session={}, chunks={}, 总耗时={:?}", - session_id, - chunk_count, - stream_elapsed - ); - } - Err(e) => { - let message = format!("Agent 错误: {e}"); - let error_event = TauriAgentEvent::Error { - message: message.clone(), - }; - let _ = app.emit(event_name, &error_event); - stream_error = Some(message); - } - } - - if let Some(extension) = removed_extension { - if let Err(error) = agent.add_extension(extension).await { - tracing::warn!( - "[UnifiedChat] 恢复 {} 扩展失败: {}", - CODE_EXECUTION_EXTENSION_NAME, - error - ); - } - } - - // 清理取消令牌 - agent_state.remove_cancel_token(session_id).await; - - if let Some(error) = stream_error { - return Err(error); - } - - Ok(()) -} - -/// 停止生成 -#[tauri::command] -pub async fn chat_stop_generation( - agent_state: State<'_, AsterAgentState>, - session_id: String, -) -> Result { - tracing::info!("[UnifiedChat] 停止生成: session={}", session_id); - Ok(agent_state.cancel_session(&session_id).await) -} - -/// 配置会话的 Provider -#[tauri::command] -pub async fn chat_configure_provider( - db: State<'_, DbConnection>, - agent_state: State<'_, AsterAgentState>, - session_id: String, - provider_type: String, - model: String, -) -> Result<(), String> { - tracing::info!( - "[UnifiedChat] 配置 Provider: session={}, provider={}, model={}", - session_id, - provider_type, - model - ); - - // 确保 Agent 已初始化 - agent_state.init_agent_with_db(&db).await?; - - // 配置 Provider - agent_state - .configure_provider_from_pool(&db, &provider_type, &model, &session_id) - .await?; - - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - use lime_agent::resolve_request_tool_policy; - - #[test] - fn test_send_message_request_deserialize_web_search_camel_case() { - let payload = serde_json::json!({ - "sessionId": "session-1", - "message": "hello", - "eventName": "event-1", - "webSearch": true - }); - let request: SendMessageRequest = - serde_json::from_value(payload).expect("deserialize request"); - assert_eq!(request.web_search, Some(true)); - assert_eq!(request.session_id, "session-1"); - assert_eq!(request.event_name, "event-1"); - } - - #[test] - fn test_send_message_request_deserialize_web_search_snake_case() { - let payload = serde_json::json!({ - "session_id": "session-1", - "message": "hello", - "event_name": "event-1", - "web_search": false - }); - let request: SendMessageRequest = - serde_json::from_value(payload).expect("deserialize request"); - assert_eq!(request.web_search, Some(false)); - } - - #[test] - fn test_unified_effective_web_search_uses_request_override() { - let mode_default = true; - let policy = resolve_request_tool_policy(Some(false), mode_default); - assert!(!policy.effective_web_search); - } -} diff --git a/src-tauri/src/commands/unified_memory_cmd.rs b/src-tauri/src/commands/unified_memory_cmd.rs index 14cb57293..b9748bf6c 100644 --- a/src-tauri/src/commands/unified_memory_cmd.rs +++ b/src-tauri/src/commands/unified_memory_cmd.rs @@ -104,7 +104,13 @@ pub async fn unified_memory_list( info!("[Unified Memory] List memories: {:?}", filters); let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; + list_unified_memories(&conn, filters) +} +pub(crate) fn list_unified_memories( + conn: &rusqlite::Connection, + filters: ListFilters, +) -> Result, String> { let archived = filters.archived.unwrap_or(false); let sort_by = normalize_sort_by(filters.sort_by.as_deref()); let order = normalize_sort_order(filters.order.as_deref()); @@ -374,7 +380,12 @@ pub async fn unified_memory_stats( info!("[Unified Memory] Stats"); let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; + collect_unified_memory_stats(&conn) +} +pub(crate) fn collect_unified_memory_stats( + conn: &rusqlite::Connection, +) -> Result { let (total_entries, memory_count, storage_used): (i64, i64, i64) = conn .query_row( "SELECT COUNT(*), COUNT(DISTINCT session_id), COALESCE(SUM(length(title) + length(content) + length(summary) + length(tags)), 0) FROM unified_memory WHERE archived = 0", diff --git a/src-tauri/src/commands/workspace_cmd.rs b/src-tauri/src/commands/workspace_cmd.rs index 8ed53dabb..ce8da113d 100644 --- a/src-tauri/src/commands/workspace_cmd.rs +++ b/src-tauri/src/commands/workspace_cmd.rs @@ -19,7 +19,10 @@ use crate::services::workspace_health_service::{ use crate::workspace::{ Workspace, WorkspaceManager, WorkspaceSettings, WorkspaceType, WorkspaceUpdate, }; -use lime_core::app_paths; +use crate::workspace_support::{ + get_or_create_default_project as load_or_create_default_project, + get_workspace_projects_root_dir, sanitize_project_dir_name, +}; use lime_core::database::lock_db; use lime_services::project_context_builder::ProjectContextBuilder; use serde::{Deserialize, Serialize}; @@ -28,31 +31,6 @@ use std::sync::Arc; use tauri::State; use tokio::sync::RwLock; -/// 获取统一的项目根目录 -fn get_workspace_projects_root_dir() -> Result { - app_paths::resolve_projects_dir() -} - -/// 规范化项目目录名,避免非法路径字符 -fn sanitize_project_dir_name(name: &str) -> String { - let sanitized: String = name - .trim() - .chars() - .map(|ch| match ch { - '/' | '\\' | ':' | '*' | '?' | '"' | '<' | '>' | '|' => '_', - _ if ch.is_control() => '_', - _ => ch, - }) - .collect(); - - let trimmed = sanitized.trim().trim_matches('.').to_string(); - if trimmed.is_empty() { - "未命名项目".to_string() - } else { - trimmed - } -} - /// Workspace 管理器状态 #[allow(dead_code)] pub struct WorkspaceManagerState(pub Arc>>); @@ -386,30 +364,7 @@ pub async fn get_or_create_default_project( db: State<'_, DbConnection>, ) -> Result { let manager = WorkspaceManager::new(db.inner().clone()); - - // 先尝试获取默认项目 - if let Some(workspace) = manager.get_default()? { - return Ok(workspace.into()); - } - - // 不存在则创建默认项目 - let default_project_path = get_workspace_projects_root_dir()?.join("default"); - - std::fs::create_dir_all(&default_project_path) - .map_err(|e| format!("创建默认项目目录失败: {e}"))?; - - let workspace = manager.create_with_type( - "默认项目".to_string(), - default_project_path, - WorkspaceType::Persistent, - )?; - - // 设置为默认 - manager.set_default(&workspace.id)?; - - // 重新获取以确保 is_default 标志正确 - let workspace = manager.get(&workspace.id)?.ok_or("创建默认项目失败")?; - Ok(workspace.into()) + Ok(load_or_create_default_project(&manager)?.into()) } /// 获取项目上下文 diff --git a/src-tauri/src/dev_bridge/dispatcher/memory.rs b/src-tauri/src/dev_bridge/dispatcher/memory.rs index 527e58916..f367de7bc 100644 --- a/src-tauri/src/dev_bridge/dispatcher/memory.rs +++ b/src-tauri/src/dev_bridge/dispatcher/memory.rs @@ -1,112 +1,12 @@ use super::{args_or_default, parse_optional_nested_arg}; +use crate::commands::unified_memory_cmd::{ + collect_unified_memory_stats, list_unified_memories, ListFilters, +}; use crate::dev_bridge::DevBridgeState; -use lime_memory::{MemoryCategory, MemoryMetadata, MemorySource, MemoryType, UnifiedMemory}; -use rusqlite::{params_from_iter, types::Value}; use serde_json::Value as JsonValue; type DynError = Box; -fn parse_unified_memory_row(row: &rusqlite::Row) -> Result { - let id: String = row.get(0)?; - let session_id: String = row.get(1)?; - let memory_type_json: String = row.get(2)?; - let category_json: String = row.get(3)?; - let title: String = row.get(4)?; - let content: String = row.get(5)?; - let summary: String = row.get(6)?; - let tags_json: String = row.get(7)?; - let confidence: f32 = row.get(8)?; - let importance: i64 = row.get(9)?; - let access_count: i64 = row.get(10)?; - let last_accessed_at: Option = row.get(11)?; - let source_json: String = row.get(12)?; - let created_at: i64 = row.get(13)?; - let updated_at: i64 = row.get(14)?; - let archived: i64 = row.get(15)?; - - let memory_type: MemoryType = serde_json::from_str(&memory_type_json) - .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?; - let category: MemoryCategory = serde_json::from_str(&category_json) - .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?; - let tags: Vec = serde_json::from_str(&tags_json) - .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?; - let source: MemorySource = serde_json::from_str(&source_json) - .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?; - - Ok(UnifiedMemory { - id, - session_id, - memory_type, - category, - title, - content, - summary, - tags, - metadata: MemoryMetadata { - confidence, - importance: importance.clamp(0, 10) as u8, - access_count: access_count.max(0) as u32, - last_accessed_at, - source, - embedding: None, - }, - created_at, - updated_at, - archived: archived != 0, - }) -} - -fn unified_memory_category_to_key(category: &MemoryCategory) -> &'static str { - match category { - MemoryCategory::Identity => "identity", - MemoryCategory::Context => "context", - MemoryCategory::Preference => "preference", - MemoryCategory::Experience => "experience", - MemoryCategory::Activity => "activity", - } -} - -fn ordered_unified_categories() -> [&'static str; 5] { - [ - "identity", - "context", - "preference", - "experience", - "activity", - ] -} - -fn normalize_unified_category_value(value: &str) -> Option<&'static str> { - if let Ok(category) = serde_json::from_str::(value) { - return Some(unified_memory_category_to_key(&category)); - } - - match value.trim_matches('"').to_lowercase().as_str() { - "identity" | "身份" => Some("identity"), - "context" | "情境" | "上下文" => Some("context"), - "preference" | "偏好" => Some("preference"), - "experience" | "经验" => Some("experience"), - "activity" | "活动" => Some("activity"), - _ => None, - } -} - -fn normalize_unified_sort_by(sort_by: Option<&str>) -> &'static str { - match sort_by.unwrap_or("updated_at") { - "created_at" => "created_at", - "importance" => "importance", - "access_count" => "access_count", - _ => "updated_at", - } -} - -fn normalize_unified_sort_order(order: Option<&str>) -> &'static str { - match order.unwrap_or("desc").to_lowercase().as_str() { - "asc" => "ASC", - _ => "DESC", - } -} - pub(super) fn try_handle( state: &DevBridgeState, cmd: &str, @@ -124,116 +24,18 @@ pub(super) fn try_handle( }; let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let (total_entries, memory_count, storage_used): (i64, i64, i64) = conn - .query_row( - "SELECT COUNT(*), COUNT(DISTINCT session_id), COALESCE(SUM(length(title) + length(content) + length(summary) + length(tags)), 0) FROM unified_memory WHERE archived = 0", - [], - |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?)), - ) - .map_err(|e| format!("统计记忆失败: {e}"))?; - - let mut category_counts: std::collections::HashMap = - std::collections::HashMap::new(); - let mut stmt = conn - .prepare( - "SELECT category, COUNT(*) FROM unified_memory WHERE archived = 0 GROUP BY category", - ) - .map_err(|e| format!("构建分类统计查询失败: {e}"))?; - - let rows = stmt - .query_map([], |row| { - let category_raw: String = row.get(0)?; - let count: i64 = row.get(1)?; - Ok((category_raw, count)) - }) - .map_err(|e| format!("分类统计查询失败: {e}"))?; - - for row in rows.flatten() { - if let Some(category) = normalize_unified_category_value(&row.0) { - category_counts.insert(category.to_string(), row.1.max(0) as u32); - } - } - - let categories = ordered_unified_categories() - .iter() - .map( - |category| crate::commands::unified_memory_cmd::MemoryCategoryStat { - category: (*category).to_string(), - count: *category_counts.get(*category).unwrap_or(&0), - }, - ) - .collect(); - - let response = crate::commands::unified_memory_cmd::MemoryStatsResponse { - total_entries: total_entries.max(0) as u32, - storage_used: storage_used.max(0) as u64, - memory_count: memory_count.max(0) as u32, - categories, - }; - - serde_json::to_value(response)? + serde_json::to_value(collect_unified_memory_stats(&conn)?)? } "unified_memory_list" => { let args = args_or_default(args); - let filters: Option = - parse_optional_nested_arg(&args, "filters")?; - let filters = filters.unwrap_or_default(); + let filters: Option = parse_optional_nested_arg(&args, "filters")?; let Some(db) = &state.db else { return Ok(Some(serde_json::json!([]))); }; let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - let archived = filters.archived.unwrap_or(false); - let sort_by = normalize_unified_sort_by(filters.sort_by.as_deref()); - let order = normalize_unified_sort_order(filters.order.as_deref()); - let limit = filters.limit.unwrap_or(120).clamp(1, 1000) as i64; - let offset = filters.offset.unwrap_or(0) as i64; - - let mut where_parts = vec!["archived = ?".to_string()]; - let mut values: Vec = vec![Value::from(if archived { 1 } else { 0 })]; - - if let Some(session_id) = filters.session_id.filter(|value| !value.trim().is_empty()) { - where_parts.push("session_id = ?".to_string()); - values.push(Value::from(session_id)); - } - - if let Some(memory_type) = filters.memory_type { - let encoded = serde_json::to_string(&memory_type) - .map_err(|e| format!("序列化 memory_type 失败: {e}"))?; - where_parts.push("memory_type = ?".to_string()); - values.push(Value::from(encoded)); - } - - if let Some(category) = filters.category { - let encoded = serde_json::to_string(&category) - .map_err(|e| format!("序列化 category 失败: {e}"))?; - where_parts.push("category = ?".to_string()); - values.push(Value::from(encoded)); - } - - let sql = format!( - "SELECT id, session_id, memory_type, category, title, content, summary, tags, confidence, importance, access_count, last_accessed_at, source, created_at, updated_at, archived FROM unified_memory WHERE {} ORDER BY {} {} LIMIT ? OFFSET ?", - where_parts.join(" AND "), - sort_by, - order, - ); - - values.push(Value::from(limit)); - values.push(Value::from(offset)); - - let mut stmt = conn - .prepare(&sql) - .map_err(|e| format!("构建查询失败: {e}"))?; - - let memories = stmt - .query_map(params_from_iter(values), parse_unified_memory_row) - .map_err(|e| format!("查询记忆失败: {e}"))? - .collect::, rusqlite::Error>>() - .map_err(|e| format!("解析记忆失败: {e}"))?; - - serde_json::to_value(memories)? + serde_json::to_value(list_unified_memories(&conn, filters.unwrap_or_default())?)? } _ => return Ok(None), }; diff --git a/src-tauri/src/dev_bridge/dispatcher/workspace.rs b/src-tauri/src/dev_bridge/dispatcher/workspace.rs index 0bc56863f..f9c1b031d 100644 --- a/src-tauri/src/dev_bridge/dispatcher/workspace.rs +++ b/src-tauri/src/dev_bridge/dispatcher/workspace.rs @@ -26,29 +26,6 @@ fn get_optional_bool_arg(args: &JsonValue, primary: &str, secondary: &str) -> Op .and_then(|value| value.as_bool()) } -fn get_workspace_projects_root_dir() -> Result { - lime_core::app_paths::resolve_projects_dir() -} - -fn sanitize_project_dir_name(name: &str) -> String { - let sanitized: String = name - .trim() - .chars() - .map(|ch| match ch { - '/' | '\\' | ':' | '*' | '?' | '"' | '<' | '>' | '|' => '_', - _ if ch.is_control() => '_', - _ => ch, - }) - .collect(); - - let trimmed = sanitized.trim().trim_matches('.').to_string(); - if trimmed.is_empty() { - "未命名项目".to_string() - } else { - trimmed - } -} - fn to_workspace_list_item_json(workspace: T) -> Result where WorkspaceListItem: From, @@ -75,25 +52,6 @@ fn build_ensure_result( } } -fn create_default_project_if_missing(manager: &WorkspaceManager) -> Result { - if let Some(workspace) = manager.get_default()? { - return to_workspace_list_item_json(workspace); - } - - let default_project_path = get_workspace_projects_root_dir()?.join("default"); - std::fs::create_dir_all(&default_project_path) - .map_err(|e| format!("创建默认项目目录失败: {e}"))?; - - let workspace = manager.create_with_type( - "默认项目".to_string(), - default_project_path, - WorkspaceType::Persistent, - )?; - manager.set_default(&workspace.id)?; - let workspace = manager.get(&workspace.id)?.ok_or("创建默认项目失败")?; - to_workspace_list_item_json(workspace) -} - fn remove_workspace_directory_if_requested( manager: &WorkspaceManager, workspace_id: &str, diff --git a/src-tauri/src/dev_bridge/dispatcher/workspace/management.rs b/src-tauri/src/dev_bridge/dispatcher/workspace/management.rs index 35ad5c23b..effc14bb9 100644 --- a/src-tauri/src/dev_bridge/dispatcher/workspace/management.rs +++ b/src-tauri/src/dev_bridge/dispatcher/workspace/management.rs @@ -1,11 +1,11 @@ use super::{ - args_or_default, create_default_project_if_missing, ensure_update_root_path, - ensure_valid_workspace_root, get_optional_bool_arg, get_string_arg, parse_nested_arg, - remove_workspace_directory_if_requested, to_workspace_list_item_json, workspace_manager, - CreateWorkspaceRequest, DynError, PathBuf, UpdateWorkspaceRequest, WorkspaceType, - WorkspaceUpdate, + args_or_default, ensure_update_root_path, ensure_valid_workspace_root, get_optional_bool_arg, + get_string_arg, parse_nested_arg, remove_workspace_directory_if_requested, + to_workspace_list_item_json, workspace_manager, CreateWorkspaceRequest, DynError, PathBuf, + UpdateWorkspaceRequest, WorkspaceType, WorkspaceUpdate, }; use crate::dev_bridge::DevBridgeState; +use crate::workspace_support::get_or_create_default_project; use serde_json::Value as JsonValue; pub(super) fn try_handle( @@ -75,7 +75,7 @@ pub(super) fn try_handle( } "get_or_create_default_project" => { let manager = workspace_manager(state)?; - create_default_project_if_missing(&manager)? + to_workspace_list_item_json(get_or_create_default_project(&manager)?)? } _ => return Ok(None), }; diff --git a/src-tauri/src/dev_bridge/dispatcher/workspace/queries.rs b/src-tauri/src/dev_bridge/dispatcher/workspace/queries.rs index ed661888c..a2be0676a 100644 --- a/src-tauri/src/dev_bridge/dispatcher/workspace/queries.rs +++ b/src-tauri/src/dev_bridge/dispatcher/workspace/queries.rs @@ -1,8 +1,8 @@ use super::{ - args_or_default, get_string_arg, get_workspace_projects_root_dir, sanitize_project_dir_name, - workspace_manager, DynError, PathBuf, WorkspaceListItem, + args_or_default, get_string_arg, workspace_manager, DynError, PathBuf, WorkspaceListItem, }; use crate::dev_bridge::DevBridgeState; +use crate::workspace_support::{get_workspace_projects_root_dir, sanitize_project_dir_name}; use serde_json::Value as JsonValue; pub(super) fn try_handle( diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 0fb85469f..9e62b9a7f 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -16,6 +16,9 @@ // 该警告来自 cocoa/objc 依赖的 msg_send! 宏,是已知的 issue #![allow(unexpected_cfgs)] +#[cfg(target_os = "linux")] +compile_error!("Lime 已暂停 Linux 桌面端支持,请使用 macOS 或 Windows 构建。"); + // 从 providers crate 重新导出(保持 crate::xxx 路径兼容) pub use lime_providers::providers; @@ -56,6 +59,7 @@ mod dev_bridge; mod logger; mod profiling; mod theme; +mod workspace_support; use lime_core::models; // 测试模块 diff --git a/src-tauri/src/services/agent_timeline_service.rs b/src-tauri/src/services/agent_timeline_service.rs index 121864b64..45b96a39f 100644 --- a/src-tauri/src/services/agent_timeline_service.rs +++ b/src-tauri/src/services/agent_timeline_service.rs @@ -5,139 +5,16 @@ use lime_core::database::dao::agent_timeline::{ AgentThreadTurnStatus, AgentTimelineDao, }; use lime_core::database::{lock_db, DbConnection}; -use serde_json::{json, Value}; +use serde_json::Value; use std::collections::HashMap; use tauri::{AppHandle, Emitter}; -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); } } -fn as_object(value: &Value) -> Option<&serde_json::Map> { - value.as_object() -} - -#[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 { - continue; - }; - for key in [ - "path", - "file_path", - "filePath", - "output_file", - "output_path", - "outputPath", - "artifact_path", - "artifact_paths", - "absolute_path", - "absolutePath", - ] { - let Some(value) = object.get(key) else { - continue; - }; - 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")) @@ -165,18 +42,6 @@ fn resolve_artifact_item_source(metadata: Option<&Value>) -> String { .unwrap_or_else(|| "artifact_snapshot".to_string()) } -fn extract_proposed_plan_block(text: &str) -> Option { - let start = text.find(PROPOSED_PLAN_OPEN)?; - let remainder = &text[start + PROPOSED_PLAN_OPEN.len()..]; - let end = remainder.find(PROPOSED_PLAN_CLOSE)?; - let content = remainder[..end].trim(); - if content.is_empty() { - None - } else { - Some(content.to_string()) - } -} - #[derive(Debug)] pub struct AgentTimelineRecorder { db: DbConnection, @@ -187,7 +52,6 @@ pub struct AgentTimelineRecorder { item_sequences: HashMap, item_statuses: HashMap, plan_text: Option, - turn_summary_text: Option, } impl AgentTimelineRecorder { @@ -228,7 +92,6 @@ impl AgentTimelineRecorder { item_sequences: HashMap::new(), item_statuses: HashMap::new(), plan_text: None, - turn_summary_text: None, }) } @@ -265,7 +128,6 @@ impl AgentTimelineRecorder { item.clone(), TauriAgentEvent::ItemStarted { item: item.clone() }, )?; - self.maybe_project_plan_item(app, event_name, item)?; } TauriAgentEvent::ItemUpdated { item } => { self.persist_runtime_item( @@ -274,7 +136,6 @@ impl AgentTimelineRecorder { item.clone(), TauriAgentEvent::ItemUpdated { item: item.clone() }, )?; - self.maybe_project_plan_item(app, event_name, item)?; } TauriAgentEvent::ItemCompleted { item } => { self.persist_runtime_item( @@ -283,52 +144,9 @@ impl AgentTimelineRecorder { item.clone(), TauriAgentEvent::ItemCompleted { item: item.clone() }, )?; - self.maybe_project_plan_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::ToolEnd { tool_id, result } => { - let metadata_value = result - .metadata - .as_ref() - .and_then(|metadata| serde_json::to_value(metadata).ok()); - - for artifact in extract_file_artifacts(None, metadata_value.as_ref()) { - let artifact_path = artifact.path.clone(); - let status = resolve_artifact_item_status(metadata_value.as_ref()); - let file_item = self.build_item( - 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: artifact_path, - source: "tool_result".to_string(), - content: None, - metadata: metadata_value.clone(), - }, - ); - self.persist_and_emit_item(app, event_name, file_item)?; - } } + TauriAgentEvent::RuntimeStatus { .. } => {} + TauriAgentEvent::ToolEnd { .. } => {} TauriAgentEvent::ArtifactSnapshot { artifact } => { let metadata_value = artifact .metadata @@ -480,18 +298,6 @@ impl AgentTimelineRecorder { 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(()) } @@ -592,98 +398,8 @@ impl AgentTimelineRecorder { self.item_statuses .insert(item.id.clone(), item.status.clone()); - if let AgentThreadItemPayload::AgentMessage { text, .. } = &item.payload { - self.plan_text = extract_proposed_plan_block(text); + if let AgentThreadItemPayload::Plan { text } = &item.payload { + self.plan_text = Some(text.clone()); } } - - fn maybe_project_plan_item( - &mut self, - app: &AppHandle, - event_name: &str, - item: &AgentThreadItem, - ) -> Result<(), String> { - let AgentThreadItemPayload::AgentMessage { text, .. } = &item.payload else { - return Ok(()); - }; - let Some(plan_text) = extract_proposed_plan_block(text) else { - return Ok(()); - }; - self.plan_text = Some(plan_text.clone()); - let plan_item = self.build_item( - format!("plan:{}", self.turn_id), - item.status.clone(), - item.completed_at.clone(), - AgentThreadItemPayload::Plan { text: plan_text }, - ); - self.persist_and_emit_item(app, event_name, plan_item)?; - Ok(()) - } -} - -pub fn complete_action_item( - db: &DbConnection, - request_id: &str, - response: Option, -) -> Result<(), String> { - let conn = lock_db(db)?; - let Some(mut item) = AgentTimelineDao::get_item(&conn, request_id) - .map_err(|e| format!("读取 action item 失败: {e}"))? - else { - return Ok(()); - }; - - let payload = match item.payload { - AgentThreadItemPayload::ApprovalRequest { - request_id, - action_type, - prompt, - tool_name, - arguments, - .. - } => AgentThreadItemPayload::ApprovalRequest { - request_id, - action_type, - prompt, - tool_name, - arguments, - response, - }, - AgentThreadItemPayload::RequestUserInput { - request_id, - action_type, - prompt, - questions, - .. - } => AgentThreadItemPayload::RequestUserInput { - request_id, - action_type, - prompt, - questions, - response, - }, - other => other, - }; - - let now = Utc::now().to_rfc3339(); - item.status = AgentThreadItemStatus::Completed; - item.completed_at = Some(now.clone()); - item.updated_at = now; - item.payload = payload; - - AgentTimelineDao::upsert_item(&conn, &item).map_err(|e| format!("更新 action item 失败: {e}")) -} - -pub fn build_action_response_value( - confirmed: bool, - response: Option<&str>, - user_data: Option<&Value>, -) -> Option { - if let Some(value) = user_data { - return Some(value.clone()); - } - if !confirmed { - return Some(json!({ "confirmed": false })); - } - response.map(|value| Value::String(value.to_string())) } diff --git a/src-tauri/src/services/chat_history_service.rs b/src-tauri/src/services/chat_history_service.rs index 7da3616cf..250dfa0bb 100644 --- a/src-tauri/src/services/chat_history_service.rs +++ b/src-tauri/src/services/chat_history_service.rs @@ -1,5 +1,4 @@ use crate::database::dao::agent::{AgentDao, AgentModelPatternMatch}; -use crate::database::load_pending_general_messages; use chrono::{Local, TimeZone}; use rusqlite::Connection; use std::collections::HashSet; @@ -24,15 +23,6 @@ pub fn load_memory_source_candidates( let mut candidates = Vec::new(); let mut seen = HashSet::new(); - load_pending_general_candidates( - conn, - from_timestamp, - to_timestamp, - limit, - min_message_length, - &mut candidates, - &mut seen, - )?; load_unified_general_candidates( conn, from_timestamp, @@ -58,33 +48,6 @@ pub fn load_memory_source_candidates( Ok(candidates) } -fn load_pending_general_candidates( - conn: &Connection, - from_timestamp: Option, - to_timestamp: Option, - limit: usize, - min_message_length: usize, - candidates: &mut Vec, - seen: &mut HashSet, -) -> Result<(), String> { - let rows = load_pending_general_messages(conn, from_timestamp, to_timestamp, limit) - .map_err(|e| format!("读取待迁移 general 消息失败: {e}"))?; - - for row in rows { - push_candidate( - candidates, - seen, - row.session_id, - row.role, - row.content, - normalize_timestamp(row.created_at), - min_message_length, - ); - } - - Ok(()) -} - fn load_unified_general_candidates( conn: &Connection, from_timestamp: Option, @@ -274,7 +237,7 @@ mod tests { } #[test] - fn load_memory_source_candidates_merges_unified_and_legacy_without_duplicates() { + fn load_memory_source_candidates_only_reads_unified_general_and_agent_messages() { let conn = Connection::open_in_memory().expect("open in memory db"); create_test_schema(&conn); @@ -299,28 +262,6 @@ mod tests { ) .unwrap(); - conn.execute( - "INSERT INTO general_chat_sessions (id, name, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", - params!["general-migrated", "旧会话", 1_741_744_000_000i64, 1_741_744_000_000i64], - ) - .unwrap(); - conn.execute( - "INSERT INTO general_chat_sessions (id, name, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", - params!["legacy-only", "旧会话2", 1_741_744_100_000i64, 1_741_744_100_000i64], - ) - .unwrap(); - - conn.execute( - "INSERT INTO general_chat_messages (id, session_id, role, content, created_at) VALUES (?1, ?2, ?3, ?4, ?5)", - params!["g1", "general-migrated", "user", "这条消息已经迁移", 1_741_744_000_000i64], - ) - .unwrap(); - conn.execute( - "INSERT INTO general_chat_messages (id, session_id, role, content, created_at) VALUES (?1, ?2, ?3, ?4, ?5)", - params!["g2", "legacy-only", "assistant", "这条消息仍在旧表中", 1_741_744_100_000i64], - ) - .unwrap(); - conn.execute( "INSERT INTO agent_messages (session_id, role, content_json, timestamp) VALUES (?1, ?2, ?3, ?4)", params![ @@ -349,23 +290,16 @@ mod tests { .iter() .map(|item| item.session_id.as_str()) .collect::>(); - assert_eq!(candidates.len(), 3); + assert_eq!(candidates.len(), 2); assert!(session_ids.contains(&"general-migrated")); - assert!(session_ids.contains(&"legacy-only")); assert!(session_ids.contains(&"agent-1")); } #[test] - fn load_memory_source_candidates_skips_legacy_general_after_migration_completed() { + fn load_memory_source_candidates_ignores_pending_general_tables() { let conn = Connection::open_in_memory().expect("open in memory db"); create_test_schema(&conn); - conn.execute( - "INSERT INTO settings (key, value) VALUES (?1, ?2)", - params!["migrated_general_chat_to_unified", "true"], - ) - .unwrap(); - conn.execute( "INSERT INTO agent_sessions (id, model, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", params![ diff --git a/src-tauri/src/services/conversation_statistics_service.rs b/src-tauri/src/services/conversation_statistics_service.rs index 3256e0df9..858f67e61 100644 --- a/src-tauri/src/services/conversation_statistics_service.rs +++ b/src-tauri/src/services/conversation_statistics_service.rs @@ -4,7 +4,7 @@ use crate::database::dao::agent::{AgentDao, AgentModelPatternMatch}; use crate::database::dao::orchestrator::OrchestratorDao; -use crate::database::{summarize_pending_general, ConversationWindowSummary}; +use crate::database::ConversationWindowSummary; use chrono::{DateTime, Datelike, Duration, Local, TimeZone, Timelike}; use rusqlite::Connection; use serde::{Deserialize, Serialize}; @@ -242,16 +242,12 @@ fn summarize_general_window( from_timestamp_ms: Option, to_timestamp_ms: Option, ) -> Result { - let unified = summarize_unified_window( + summarize_unified_window( conn, AgentModelPatternMatch::Like, from_timestamp_ms, to_timestamp_ms, - )?; - let pending = summarize_pending_general(conn, from_timestamp_ms, to_timestamp_ms) - .map_err(|e| format!("查询待迁移 general 摘要失败: {e}"))?; - - Ok(unified.merge(pending)) + ) } fn summarize_agent_window( @@ -298,6 +294,7 @@ fn build_conversation_stats(windows: ConversationWindowTriplet) -> ConversationS } } +#[cfg(test)] fn query_general_chat_stats( conn: &Connection, today_start: &DateTime, @@ -307,6 +304,7 @@ fn query_general_chat_stats( .map(build_conversation_stats) } +#[cfg(test)] fn query_agent_chat_stats( conn: &Connection, today_start: &DateTime, @@ -641,7 +639,7 @@ mod tests { } #[test] - fn stats_ignore_legacy_general_after_migration_completed() { + fn stats_ignore_legacy_general_tables_during_runtime() { let conn = Connection::open_in_memory().expect("open in memory db"); create_test_schema(&conn); @@ -651,12 +649,6 @@ mod tests { .expect("build datetime"); let now_ms = now.timestamp_millis(); - conn.execute( - "INSERT INTO settings (key, value) VALUES (?1, ?2)", - params!["migrated_general_chat_to_unified", "true"], - ) - .unwrap(); - conn.execute( "INSERT INTO general_chat_sessions (id, name, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", params!["legacy-only", "旧通用会话", now_ms, now_ms], diff --git a/src-tauri/src/services/memory_profile_prompt_service.rs b/src-tauri/src/services/memory_profile_prompt_service.rs index 2a2cb51aa..3bac7133b 100644 --- a/src-tauri/src/services/memory_profile_prompt_service.rs +++ b/src-tauri/src/services/memory_profile_prompt_service.rs @@ -1,7 +1,7 @@ -//! 记忆画像提示词服务 +//! 记忆提示词装配服务 //! -//! 将设置页中的记忆画像(学习状态、擅长领域、解释偏好、难题偏好) -//! 转换为可注入到系统提示词中的统一指令片段。 +//! 将设置页中的记忆画像与配置化记忆来源统一装配为可注入到 system prompt +//! 的单一记忆指令片段,避免调用方继续各自决定拼装顺序。 use lime_core::config::Config; use std::path::Path; @@ -11,6 +11,26 @@ use crate::services::memory_source_resolver_service::build_memory_sources_prompt const MEMORY_PROFILE_PROMPT_MARKER: &str = "【用户记忆画像偏好】"; const MEMORY_SOURCE_PROMPT_MARKER: &str = "【记忆来源补充指令】"; +#[derive(Debug, Clone, Copy, Default)] +pub struct MemoryPromptContext<'a> { + pub working_dir: Option<&'a Path>, + pub active_relative_path: Option<&'a str>, +} + +impl<'a> MemoryPromptContext<'a> { + pub fn with_working_dir(working_dir: &'a Path) -> Self { + Self { + working_dir: Some(working_dir), + active_relative_path: None, + } + } + + pub fn with_active_relative_path(mut self, active_relative_path: Option<&'a str>) -> Self { + self.active_relative_path = active_relative_path; + self + } +} + fn normalize_text(input: &str) -> Option { let trimmed = input.trim(); if trimmed.is_empty() { @@ -32,7 +52,7 @@ fn normalize_list(items: &[String]) -> Vec { /// 仅在以下条件满足时返回: /// - 记忆功能已启用 /// - 至少有一项画像字段有值 -pub fn build_memory_profile_prompt(config: &Config) -> Option { +fn build_memory_profile_prompt(config: &Config) -> Option { let memory = &config.memory; if !memory.enabled { return None; @@ -83,59 +103,73 @@ pub fn build_memory_profile_prompt(config: &Config) -> Option { Some(lines.join("\n")) } -/// 合并基础系统提示词与记忆画像提示词 -/// -/// - 已包含画像标记时不会重复追加 -/// - 任一方为空时返回另一方 -pub fn merge_system_prompt_with_memory_profile( - base_prompt: Option, +fn build_memory_sources_prompt_for_context( config: &Config, + context: MemoryPromptContext<'_>, ) -> Option { - let memory_prompt = build_memory_profile_prompt(config); + let working_dir = context.working_dir?; + if !config.memory.enabled { + return None; + } - match (base_prompt, memory_prompt) { - (Some(base), Some(memory)) => { - if base.contains(MEMORY_PROFILE_PROMPT_MARKER) { + build_memory_sources_prompt(config, working_dir, context.active_relative_path, 4000) +} + +fn merge_prompt_section( + base_prompt: Option, + section_prompt: Option, + marker: &str, +) -> Option { + match (base_prompt, section_prompt) { + (Some(base), Some(section)) => { + if base.contains(marker) { Some(base) } else if base.trim().is_empty() { - Some(memory) + Some(section) } else { - Some(format!("{base}\n\n{memory}")) + Some(format!("{base}\n\n{section}")) } } (Some(base), None) => Some(base), - (None, Some(memory)) => Some(memory), + (None, Some(section)) => Some(section), (None, None) => None, } } -pub fn merge_system_prompt_with_memory_sources( +pub fn build_memory_prompt(config: &Config, context: MemoryPromptContext<'_>) -> Option { + let with_profile = merge_prompt_section( + None, + build_memory_profile_prompt(config), + MEMORY_PROFILE_PROMPT_MARKER, + ); + + merge_prompt_section( + with_profile, + build_memory_sources_prompt_for_context(config, context), + MEMORY_SOURCE_PROMPT_MARKER, + ) +} + +/// 合并基础系统提示词与统一记忆提示词。 +/// +/// - 画像与来源统一在同一边界内拼装 +/// - 已包含对应 marker 时不会重复追加 +pub fn merge_system_prompt_with_memory_context( base_prompt: Option, config: &Config, - working_dir: &Path, - active_relative_path: Option<&str>, + context: MemoryPromptContext<'_>, ) -> Option { - if !config.memory.enabled { - return base_prompt; - } + let with_profile = merge_prompt_section( + base_prompt, + build_memory_profile_prompt(config), + MEMORY_PROFILE_PROMPT_MARKER, + ); - let memory_sources_prompt = - build_memory_sources_prompt(config, working_dir, active_relative_path, 4000); - - match (base_prompt, memory_sources_prompt) { - (Some(base), Some(source_prompt)) => { - if base.contains(MEMORY_SOURCE_PROMPT_MARKER) { - Some(base) - } else if base.trim().is_empty() { - Some(source_prompt) - } else { - Some(format!("{base}\n\n{source_prompt}")) - } - } - (Some(base), None) => Some(base), - (None, Some(source_prompt)) => Some(source_prompt), - (None, None) => None, - } + merge_prompt_section( + with_profile, + build_memory_sources_prompt_for_context(config, context), + MEMORY_SOURCE_PROMPT_MARKER, + ) } #[cfg(test)] @@ -192,7 +226,11 @@ mod tests { config.memory.profile = Some(profile); let base = Some("前置内容\n\n【用户记忆画像偏好】\n已有内容".to_string()); - let merged = merge_system_prompt_with_memory_profile(base.clone(), &config); + let merged = merge_system_prompt_with_memory_context( + base.clone(), + &config, + MemoryPromptContext::default(), + ); assert_eq!(merged, base); } @@ -210,10 +248,40 @@ mod tests { config.memory.sources.project_memory_paths = vec!["AGENTS.md".to_string()]; config.memory.sources.project_rule_dirs = Vec::new(); - let merged = merge_system_prompt_with_memory_sources(None, &config, tmp.path(), None) - .expect("should build sources prompt"); + let merged = merge_system_prompt_with_memory_context( + None, + &config, + MemoryPromptContext::with_working_dir(tmp.path()), + ) + .expect("should build sources prompt"); assert!(merged.contains("【记忆来源补充指令】")); assert!(merged.contains("偏好简洁输出")); } + + #[test] + fn should_build_combined_memory_prompt() { + let tmp = TempDir::new().expect("create temp dir"); + fs::write(tmp.path().join("AGENTS.md"), "# 项目记忆\n- 保持简洁") + .expect("write memory file"); + + let mut config = Config::default(); + config.memory.enabled = true; + let mut profile = config.memory.profile.clone().unwrap_or_default(); + profile.current_status = Some("高级开发者".to_string()); + config.memory.profile = Some(profile); + config.memory.sources.project_memory_paths = vec!["AGENTS.md".to_string()]; + config.memory.sources.project_rule_dirs = Vec::new(); + config.memory.sources.managed_policy_path = Some("missing-managed.md".to_string()); + config.memory.sources.user_memory_path = Some("missing-user.md".to_string()); + + let prompt = + build_memory_prompt(&config, MemoryPromptContext::with_working_dir(tmp.path())) + .expect("should build combined prompt"); + + assert!(prompt.contains("【用户记忆画像偏好】")); + assert!(prompt.contains("高级开发者")); + assert!(prompt.contains("【记忆来源补充指令】")); + assert!(prompt.contains("保持简洁")); + } } diff --git a/src-tauri/src/services/openclaw_service.rs b/src-tauri/src/services/openclaw_service.rs index c30a81cd6..430632994 100644 --- a/src-tauri/src/services/openclaw_service.rs +++ b/src-tauri/src/services/openclaw_service.rs @@ -125,6 +125,12 @@ pub struct EnvironmentDiagnostics { pub supplemental_search_dirs: Vec, #[serde(default)] pub supplemental_command_candidates: Vec, + #[serde(default)] + pub git_where_candidates: Vec, + #[serde(default)] + pub git_supplemental_search_dirs: Vec, + #[serde(default)] + pub git_supplemental_command_candidates: Vec, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -4089,23 +4095,30 @@ fn find_command_in_bin_dir(command_name: &str, bin_dir: &Path) -> Option Result, String> { +fn collect_existing_unique_dirs(candidates: I) -> Vec +where + I: IntoIterator, +{ let mut dirs = Vec::new(); let mut seen = HashSet::new(); - let mut push_dir = |dir: PathBuf| { + for dir in candidates { if dir.as_os_str().is_empty() || !dir.exists() { - return; + continue; } if seen.insert(dir.clone()) { dirs.push(dir); } - }; + } - push_dir(preferred_bin_dir.to_path_buf()); + dirs +} + +async fn collect_preferred_runtime_command_dirs( + command_name: &str, + preferred_bin_dir: &Path, +) -> Result, String> { + let mut candidate_dirs = vec![preferred_bin_dir.to_path_buf()]; if command_name == "openclaw" { if let Some(npm_path) = find_command_in_bin_dir("npm", preferred_bin_dir) @@ -4113,13 +4126,13 @@ async fn collect_preferred_runtime_command_dirs( { if let Some(prefix) = detect_npm_global_prefix(&npm_path).await { for dir in npm_global_command_dirs(&prefix) { - push_dir(dir); + candidate_dirs.push(dir); } } } } - Ok(dirs) + Ok(collect_existing_unique_dirs(candidate_dirs)) } async fn collect_preferred_runtime_command_candidates( @@ -4252,9 +4265,17 @@ async fn select_command_candidate( return select_node_runtime_candidate(candidates).await; } + if command_name == "git" { + return Ok(select_best_git_candidate(candidates)); + } + Ok(candidates.into_iter().next()) } +fn select_best_git_candidate(candidates: Vec) -> Option { + select_preferred_path_candidate(candidates.clone()).or_else(|| candidates.into_iter().next()) +} + async fn select_preferred_runtime_candidate( command_name: &str, candidates: &[PathBuf], @@ -4287,68 +4308,61 @@ async fn select_preferred_runtime_candidate( } fn find_all_commands_in_known_locations(command_name: &str) -> Vec { - let search_dirs = collect_known_command_search_dirs(); + let search_dirs = collect_known_command_search_dirs(command_name); find_all_commands_in_paths(command_name, &search_dirs) } -fn collect_known_command_search_dirs() -> Vec { +fn collect_known_command_search_dirs(_command_name: &str) -> Vec { let mut search_dirs = Vec::new(); - let mut seen = HashSet::new(); - - let mut push_dir = |dir: PathBuf| { - if dir.as_os_str().is_empty() || !dir.exists() { - return; - } - if seen.insert(dir.clone()) { - search_dirs.push(dir); - } - }; if let Some(path_var) = std::env::var_os("PATH") { - for dir in std::env::split_paths(&path_var) { - push_dir(dir); - } + search_dirs.extend(std::env::split_paths(&path_var)); } if let Some(home) = home_dir() { - push_dir(home.join(".npm-global/bin")); - push_dir(home.join(".local/bin")); - push_dir(home.join(".bun/bin")); - push_dir(home.join(".volta/bin")); - push_dir(home.join(".asdf/shims")); - push_dir(home.join(".local/share/mise/shims")); - push_dir(home.join("Library/PhpWebStudy/env/node/bin")); + search_dirs.extend([ + home.join(".npm-global/bin"), + home.join(".local/bin"), + home.join(".bun/bin"), + home.join(".volta/bin"), + home.join(".asdf/shims"), + home.join(".local/share/mise/shims"), + home.join("Library/PhpWebStudy/env/node/bin"), + ]); let nvm_versions = home.join(".nvm/versions/node"); if let Ok(entries) = std::fs::read_dir(nvm_versions) { for entry in entries.flatten() { - push_dir(entry.path().join("bin")); + search_dirs.push(entry.path().join("bin")); } } let fnm_versions = home.join(".fnm/node-versions"); if let Ok(entries) = std::fs::read_dir(fnm_versions) { for entry in entries.flatten() { - push_dir(entry.path().join("installation/bin")); + search_dirs.push(entry.path().join("installation/bin")); } } } #[cfg(target_os = "windows")] { - for dir in windows_known_command_dirs_from_env() { - push_dir(dir); + search_dirs.extend(windows_known_command_dirs_from_env()); + if _command_name == "git" { + search_dirs.extend(windows_known_git_command_dirs_from_env()); } } if cfg!(target_os = "macos") { - push_dir(PathBuf::from("/opt/homebrew/bin")); - push_dir(PathBuf::from("/usr/local/bin")); - push_dir(PathBuf::from("/usr/bin")); - push_dir(PathBuf::from("/bin")); + search_dirs.extend([ + PathBuf::from("/opt/homebrew/bin"), + PathBuf::from("/usr/local/bin"), + PathBuf::from("/usr/bin"), + PathBuf::from("/bin"), + ]); } - search_dirs + collect_existing_unique_dirs(search_dirs) } #[cfg(target_os = "windows")] @@ -4386,17 +4400,64 @@ fn windows_known_command_dirs_from_env() -> Vec { dirs } -fn find_all_commands_in_paths(command_name: &str, search_dirs: &[PathBuf]) -> Vec { - #[cfg(target_os = "windows")] - let candidates = [ - format!("{command_name}.exe"), - format!("{command_name}.cmd"), - format!("{command_name}.bat"), - command_name.to_string(), - ]; +#[cfg(any(target_os = "windows", test))] +fn windows_git_install_dir_variants(root: PathBuf) -> Vec { + vec![ + root.join("cmd"), + root.join("bin"), + root.join("mingw64").join("bin"), + ] +} - #[cfg(not(target_os = "windows"))] - let candidates = [command_name.to_string()]; +#[cfg(target_os = "windows")] +fn windows_known_git_command_dirs_from_env() -> Vec { + let mut dirs = Vec::new(); + + if let Some(program_files) = std::env::var_os("ProgramFiles") { + dirs.extend(windows_git_install_dir_variants( + PathBuf::from(program_files).join("Git"), + )); + } + + if let Some(program_files_x86) = std::env::var_os("ProgramFiles(x86)") { + dirs.extend(windows_git_install_dir_variants( + PathBuf::from(program_files_x86).join("Git"), + )); + } + + if let Some(localappdata) = std::env::var_os("LOCALAPPDATA") { + dirs.extend(windows_git_install_dir_variants( + PathBuf::from(localappdata).join("Programs").join("Git"), + )); + } + + if let Some(home) = home_dir() { + dirs.extend(windows_git_install_dir_variants( + home.join("scoop").join("apps").join("git").join("current"), + )); + } + + dirs +} + +fn find_all_commands_in_paths(command_name: &str, search_dirs: &[PathBuf]) -> Vec { + find_all_commands_in_paths_for(current_shell_platform(), command_name, search_dirs) +} + +fn find_all_commands_in_paths_for( + platform: ShellPlatform, + command_name: &str, + search_dirs: &[PathBuf], +) -> Vec { + let candidates = match platform { + ShellPlatform::Windows => vec![ + format!("{command_name}.exe"), + format!("{command_name}.cmd"), + format!("{command_name}.bat"), + command_name.to_string(), + ], + ShellPlatform::Unix => vec![command_name.to_string()], + }; let mut matches = Vec::new(); let mut seen = HashSet::new(); @@ -4844,6 +4905,24 @@ async fn collect_environment_diagnostics() -> EnvironmentDiagnostics { .and_then(find_installed_openclaw_package_details) .map(|package| package.path.display().to_string()); + #[cfg(target_os = "windows")] + let git_where_candidates = find_commands_via_where("git") + .await + .unwrap_or_default() + .into_iter() + .map(|path| path.display().to_string()) + .collect(); + + #[cfg(not(target_os = "windows"))] + let git_where_candidates = Vec::new(); + + let git_supplemental_search_dirs = collect_supplemental_git_search_dirs(); + let git_supplemental_command_candidates = + find_all_commands_in_paths("git", &git_supplemental_search_dirs) + .into_iter() + .map(|path| path.display().to_string()) + .collect(); + EnvironmentDiagnostics { npm_path, npm_global_prefix, @@ -4854,36 +4933,40 @@ async fn collect_environment_diagnostics() -> EnvironmentDiagnostics { .map(|path| path.display().to_string()) .collect(), supplemental_command_candidates, + git_where_candidates, + git_supplemental_search_dirs: git_supplemental_search_dirs + .into_iter() + .map(|path| path.display().to_string()) + .collect(), + git_supplemental_command_candidates, } } fn collect_supplemental_openclaw_search_dirs(npm_global_prefix: Option<&str>) -> Vec { let mut dirs = Vec::new(); - let mut seen = HashSet::new(); - - let mut push_dir = |dir: PathBuf| { - if dir.as_os_str().is_empty() || !dir.exists() { - return; - } - if seen.insert(dir.clone()) { - dirs.push(dir); - } - }; #[cfg(target_os = "windows")] { - for dir in windows_known_command_dirs_from_env() { - push_dir(dir); - } + dirs.extend(windows_known_command_dirs_from_env()); } if let Some(prefix) = npm_global_prefix { - for dir in npm_global_command_dirs(prefix) { - push_dir(dir); - } + dirs.extend(npm_global_command_dirs(prefix)); } - dirs + collect_existing_unique_dirs(dirs) +} + +fn collect_supplemental_git_search_dirs() -> Vec { + #[cfg(target_os = "windows")] + { + return collect_existing_unique_dirs(windows_known_git_command_dirs_from_env()); + } + + #[cfg(not(target_os = "windows"))] + { + Vec::new() + } } async fn select_best_node_candidate(candidates: Vec) -> Result, String> { @@ -5263,15 +5346,15 @@ mod tests { npm_global_node_modules_dirs_for, package_registry_for_package_spec, parse_semver_from_text, resolve_openclaw_cli_entry_from_package_manifest, resolve_openclaw_command_from_runtime_candidate, resolve_windows_dependency_install_plan, - runtime_candidate_matches_install_root, sanitize_runtime_config, + runtime_candidate_matches_install_root, sanitize_runtime_config, select_best_git_candidate, select_best_semver_candidate, select_gateway_start_failure_detail, select_openclaw_update_failure_detail, select_preferred_path_candidate, shell_command_escape_for, shell_npm_prefix_assignment_for, shell_path_assignment_for, trim_trailing_slash, windows_dependency_action_result, windows_dependency_setup_message, - windows_install_block_result, windows_manual_install_message, DependencyKind, - DependencyStatus, EnvironmentDiagnostics, OpenClawRuntimeCandidate, - ResolvedOpenClawCommand, ShellPlatform, WindowsDependencyInstallPlan, NPM_MIRROR_CN, - OPENCLAW_CN_PACKAGE, OPENCLAW_DEFAULT_PACKAGE, + windows_git_install_dir_variants, windows_install_block_result, + windows_manual_install_message, DependencyKind, DependencyStatus, EnvironmentDiagnostics, + OpenClawRuntimeCandidate, ResolvedOpenClawCommand, ShellPlatform, + WindowsDependencyInstallPlan, NPM_MIRROR_CN, OPENCLAW_CN_PACKAGE, OPENCLAW_DEFAULT_PACKAGE, }; use crate::database::dao::api_key_provider::{ApiKeyProvider, ApiProviderType, ProviderGroup}; use chrono::Utc; @@ -5924,6 +6007,43 @@ mod tests { ); } + #[test] + fn git_candidate_selection_prefers_executable_extension() { + let preferred = select_best_git_candidate(vec![ + PathBuf::from(r"C:\Program Files\Git\cmd\git.cmd"), + PathBuf::from(r"C:\Program Files\Git\cmd\git.exe"), + ]); + + assert_eq!( + preferred, + Some(PathBuf::from(r"C:\Program Files\Git\cmd\git.exe")) + ); + } + + #[test] + fn windows_git_install_dir_variants_cover_common_layouts() { + let git_root = build_unique_temp_dir("git-layout-root"); + let cmd_dir = git_root.join("cmd"); + let bin_dir = git_root.join("bin"); + fs::create_dir_all(&cmd_dir).unwrap(); + fs::create_dir_all(&bin_dir).unwrap(); + fs::write(cmd_dir.join("git.exe"), "").unwrap(); + fs::write(bin_dir.join("git.cmd"), "").unwrap(); + + let matches = super::find_all_commands_in_paths_for( + ShellPlatform::Windows, + "git", + &windows_git_install_dir_variants(git_root.clone()), + ); + + let _ = fs::remove_dir_all(&git_root); + + assert_eq!( + matches, + vec![cmd_dir.join("git.exe"), bin_dir.join("git.cmd")] + ); + } + #[test] fn windows_npm_global_command_dirs_use_prefix_root() { assert_eq!( diff --git a/src-tauri/src/skills/README.md b/src-tauri/src/skills/README.md index 56a73a78b..b19cdd6f2 100644 --- a/src-tauri/src/skills/README.md +++ b/src-tauri/src/skills/README.md @@ -7,8 +7,12 @@ | 文件 | 说明 | |------|------| | `mod.rs` | 模块导出 | +| `catalog.rs` | skill catalog 枚举、详情 DTO 与标准合规校验边界 | +| `execution.rs` | Skill prompt/workflow 的 Tauri emitter、execution_mode 路由、按 skill_name 执行主链与错误码适配壳(纯执行主链位于 `crates/agent/src/skill_execution.rs`) | | `llm_provider.rs` | 桥接层(纯逻辑已迁移到 `crates/skills/src/lime_llm_provider.rs`) | | `execution_callback.rs` | TauriExecutionCallback 实现(保留在主 crate) | +| `runtime.rs` | skill 执行前置准备、provider fallback、run metadata 投影边界 | +| `social_post.rs` | `social_post_with_cover` 的后处理与产物事件投影 | ## Skills 集成架构 @@ -78,6 +82,23 @@ agent/aster_state.rs └── reload_lime_skills() skills/ +├── catalog.rs +│ ├── 可执行 skill 列表枚举 +│ ├── skill 详情 DTO 投影 +│ └── 标准合规校验错误归一 +├── execution.rs (Tauri 适配壳) +│ ├── execute_named_skill 统一技能执行入口 +│ └── crates/agent/src/skill_execution.rs +│ ├── prompt/workflow 执行主链 +│ ├── Aster reply 流桥接 +│ └── SkillExecutionResult / StepResult +├── runtime.rs +│ ├── agent/tool 初始化准备 +│ ├── provider fallback 与统一 memory prompt +│ └── skill run start/finish metadata 投影 +├── social_post.rs +│ ├── social_post_with_cover 结果标准化 +│ └── 社媒产物 Tool/Artifact 事件补投影 ├── llm_provider.rs (桥接) │ └── crates/skills/src/lime_llm_provider.rs │ ├── ProviderPoolService (凭证池管理) @@ -103,6 +124,12 @@ commands/skill_cmd.rs - Agent Skills 是唯一标准格式 - Lime 私有能力统一写入 `metadata.lime_*` - Workflow 不再推荐使用 `steps-json` 内联,优先通过 `metadata.lime_workflow_ref` 指向 `references/` 下文件 +- skill catalog 的枚举、详情 DTO 与标准合规校验统一收口到 `skills/catalog.rs` +- execution_mode -> prompt/workflow 的路由统一收口到 `skills/execution.rs` +- 按 `skill_name` 执行、execution tracker 包装与场景命令复用统一收口到 `skills/execution.rs` +- prompt/workflow 的执行、session 构建、流事件桥接统一收口到 `lime-agent::skill_execution` +- skill 执行前的 agent/tool 初始化、provider fallback 与 execution tracker metadata 统一收口到 `skills/runtime.rs` +- `social_post_with_cover` 的结果标准化与补充 Artifact 事件统一收口到 `skills/social_post.rs` - 服务层和执行层共用 `SkillService::inspect_*` inspection 结果作为标准合规事实源,并向前端暴露标准合规状态与资源摘要 - 无效 Skill 仍可在管理页中看到检查结果,但不会进入运行时自动加载和可执行列表 - 管理链路支持创建最小标准 Skill 脚手架,新建结果会立即经过统一 inspection 校验 diff --git a/src-tauri/src/skills/catalog.rs b/src-tauri/src/skills/catalog.rs new file mode 100644 index 000000000..f8e63fb52 --- /dev/null +++ b/src-tauri/src/skills/catalog.rs @@ -0,0 +1,532 @@ +use serde::{Deserialize, Serialize}; + +use crate::commands::skill_error::{ + format_skill_error, map_find_skill_error, SKILL_ERR_CATALOG_UNAVAILABLE, + SKILL_ERR_EXECUTE_FAILED, +}; +use lime_skills::{ + find_skill_by_name, get_skill_roots, load_skills_from_directory, LoadedSkillDefinition, +}; + +/// 可执行 Skill 信息 +/// +/// 用于 list_executable_skills 命令的返回类型 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ExecutableSkillInfo { + /// Skill 名称(唯一标识) + pub name: String, + /// 显示名称 + pub display_name: String, + /// Skill 描述 + pub description: String, + /// 执行模式:prompt, workflow, agent + pub execution_mode: String, + /// 是否有 workflow 定义 + pub has_workflow: bool, + /// 指定的 Provider(可选) + pub provider: Option, + /// 指定的 Model(可选) + pub model: Option, + /// 参数提示(可选) + pub argument_hint: Option, +} + +/// Workflow 步骤信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct WorkflowStepInfo { + /// 步骤 ID + pub id: String, + /// 步骤名称 + pub name: String, + /// 依赖的步骤 ID 列表 + pub dependencies: Vec, +} + +/// Skill 详情信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SkillDetailInfo { + /// 基本信息 + #[serde(flatten)] + pub basic: ExecutableSkillInfo, + /// Markdown 内容 + pub markdown_content: String, + /// Workflow 步骤(如果有) + pub workflow_steps: Option>, + /// 允许的工具列表(可选) + pub allowed_tools: Option>, + /// 使用场景说明(可选) + pub when_to_use: Option, +} + +pub 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("; ") + )) +} + +pub fn load_executable_skill_definition(skill_name: &str) -> Result { + 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( + SKILL_ERR_EXECUTE_FAILED, + format!("Skill '{skill_name}' 已禁用模型调用,无法执行"), + )); + } + + Ok(skill) +} + +fn to_executable_skill_info(skill: LoadedSkillDefinition) -> ExecutableSkillInfo { + ExecutableSkillInfo { + name: skill.skill_name, + display_name: skill.display_name, + description: skill.description, + execution_mode: skill.execution_mode.clone(), + has_workflow: skill.execution_mode == "workflow", + provider: skill.provider, + model: skill.model, + argument_hint: skill.argument_hint, + } +} + +pub fn list_executable_skill_catalog() -> Result, String> { + let skill_roots = get_skill_roots(); + if skill_roots.is_empty() { + return Err(format_skill_error( + SKILL_ERR_CATALOG_UNAVAILABLE, + "无法获取 Skills 目录", + )); + } + + 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); + } + } + } + + let executable_skills: Vec = all_skills + .into_iter() + .filter(|skill| !skill.disable_model_invocation) + .map(to_executable_skill_info) + .collect(); + + tracing::info!( + "[list_executable_skills] 返回 {} 个可执行 Skills", + executable_skills.len() + ); + + Ok(executable_skills) +} + +pub fn get_skill_detail_info(skill_name: &str) -> Result { + let skill = load_executable_skill_definition(skill_name)?; + + let detail = SkillDetailInfo { + basic: ExecutableSkillInfo { + name: skill.skill_name, + display_name: skill.display_name, + description: skill.description, + execution_mode: skill.execution_mode.clone(), + has_workflow: skill.execution_mode == "workflow", + provider: skill.provider, + model: skill.model, + argument_hint: skill.argument_hint, + }, + markdown_content: skill.markdown_content, + workflow_steps: if skill.workflow_steps.is_empty() { + None + } else { + Some( + skill + .workflow_steps + .iter() + .map(|step| WorkflowStepInfo { + id: step.id.clone(), + name: step.name.clone(), + dependencies: Vec::new(), + }) + .collect(), + ) + }, + allowed_tools: skill.allowed_tools, + when_to_use: skill.when_to_use, + }; + + tracing::info!("[get_skill_detail] 返回 Skill 详情: name={}", skill_name); + + Ok(detail) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::skills::{SkillExecutionResult, StepResult}; + use lime_skills::{ + load_skill_from_file, parse_allowed_tools, parse_boolean, parse_skill_frontmatter, + }; + + #[test] + fn test_executable_skill_info_serialization() { + let info = ExecutableSkillInfo { + name: "test-skill".to_string(), + display_name: "Test Skill".to_string(), + description: "A test skill".to_string(), + execution_mode: "prompt".to_string(), + has_workflow: false, + provider: None, + model: None, + argument_hint: Some("Enter your query".to_string()), + }; + + let json = serde_json::to_string(&info).unwrap(); + assert!(json.contains("test-skill")); + assert!(json.contains("Test Skill")); + } + + #[test] + fn test_skill_execution_result_serialization() { + let result = SkillExecutionResult { + success: true, + output: Some("Hello, world!".to_string()), + error: None, + steps_completed: vec![StepResult { + step_id: "step-1".to_string(), + step_name: "Process".to_string(), + success: true, + output: Some("Done".to_string()), + error: None, + }], + }; + + let json = serde_json::to_string(&result).unwrap(); + assert!(json.contains("\"success\":true")); + assert!(json.contains("Hello, world!")); + assert!(json.contains("step-1")); + } + + #[test] + fn test_skill_detail_info_serialization() { + let detail = SkillDetailInfo { + basic: ExecutableSkillInfo { + name: "workflow-skill".to_string(), + display_name: "Workflow Skill".to_string(), + description: "A workflow skill".to_string(), + execution_mode: "workflow".to_string(), + has_workflow: true, + provider: Some("claude".to_string()), + model: Some("claude-sonnet-4-5-20250514".to_string()), + argument_hint: None, + }, + markdown_content: "# Workflow Skill\n\nThis is a workflow skill.".to_string(), + workflow_steps: Some(vec![ + WorkflowStepInfo { + id: "step-1".to_string(), + name: "Initialize".to_string(), + dependencies: vec![], + }, + WorkflowStepInfo { + id: "step-2".to_string(), + name: "Process".to_string(), + dependencies: vec!["step-1".to_string()], + }, + ]), + allowed_tools: Some(vec!["read_file".to_string(), "write_file".to_string()]), + when_to_use: Some("Use this skill for complex workflows".to_string()), + }; + + let json = serde_json::to_string(&detail).unwrap(); + assert!(json.contains("workflow-skill")); + assert!(json.contains("workflow_steps")); + assert!(json.contains("step-1")); + assert!(json.contains("step-2")); + } + + #[test] + fn test_parse_skill_frontmatter_basic() { + let content = r#"--- +name: test-skill +description: A test skill +metadata: + lime_model_preference: claude-sonnet-4-5-20250514 + lime_provider_preference: claude +--- + +# Test Skill + +This is the body content. +"#; + let (fm, body) = parse_skill_frontmatter(content); + assert_eq!(fm.name, Some("test-skill".to_string())); + assert_eq!(fm.description, Some("A test skill".to_string())); + assert_eq!(fm.model, Some("claude-sonnet-4-5-20250514".to_string())); + assert_eq!(fm.provider, Some("claude".to_string())); + assert!(body.contains("# Test Skill")); + assert!(body.contains("This is the body content.")); + } + + #[test] + fn test_parse_skill_frontmatter_no_frontmatter() { + let content = "# Just content\nNo frontmatter here."; + let (fm, body) = parse_skill_frontmatter(content); + assert!(fm.name.is_none()); + assert_eq!(body, content); + } + + #[test] + fn test_parse_skill_frontmatter_with_quotes() { + let content = r#"--- +name: "quoted-name" +description: 'single quoted' +--- +Body +"#; + let (fm, _) = parse_skill_frontmatter(content); + assert_eq!(fm.name, Some("quoted-name".to_string())); + assert_eq!(fm.description, Some("single quoted".to_string())); + } + + #[test] + fn test_parse_allowed_tools() { + assert_eq!(parse_allowed_tools(None), None); + assert_eq!(parse_allowed_tools(Some("")), None); + assert_eq!( + parse_allowed_tools(Some("tool1")), + Some(vec!["tool1".to_string()]) + ); + assert_eq!( + parse_allowed_tools(Some("tool1, tool2, tool3")), + Some(vec![ + "tool1".to_string(), + "tool2".to_string(), + "tool3".to_string() + ]) + ); + } + + #[test] + fn test_parse_boolean() { + assert!(!parse_boolean(None, false)); + assert!(parse_boolean(None, true)); + assert!(parse_boolean(Some("true"), false)); + assert!(parse_boolean(Some("TRUE"), false)); + assert!(parse_boolean(Some("1"), false)); + assert!(parse_boolean(Some("yes"), false)); + assert!(!parse_boolean(Some("false"), true)); + assert!(!parse_boolean(Some("no"), true)); + } + + #[test] + fn test_load_skill_from_file() { + use tempfile::TempDir; + + let temp_dir = TempDir::new().unwrap(); + let skill_dir = temp_dir.path().join("my-skill"); + std::fs::create_dir(&skill_dir).unwrap(); + + let skill_file = skill_dir.join("SKILL.md"); + std::fs::write( + &skill_file, + r#"--- +name: my-skill +description: Test skill description +allowed-tools: tool1, tool2 +metadata: + lime_model_preference: gpt-4 + lime_provider_preference: openai +--- + +# My Skill + +Instructions here. +"#, + ) + .unwrap(); + + let skill = load_skill_from_file("my-skill", &skill_file).unwrap(); + + assert_eq!(skill.skill_name, "my-skill"); + assert_eq!(skill.display_name, "my-skill"); + assert_eq!(skill.description, "Test skill description"); + assert_eq!( + skill.allowed_tools, + Some(vec!["tool1".to_string(), "tool2".to_string()]) + ); + assert_eq!(skill.model, Some("gpt-4".to_string())); + assert_eq!(skill.provider, Some("openai".to_string())); + assert!(!skill.disable_model_invocation); + assert_eq!(skill.execution_mode, "prompt"); + 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: + lime_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.lime_workflow_ref"))); + assert!(skill.workflow_steps.is_empty()); + } + + #[test] + fn test_load_skills_from_directory() { + use tempfile::TempDir; + + let temp_dir = TempDir::new().unwrap(); + let skills_dir = temp_dir.path(); + + let skill1_dir = skills_dir.join("skill-one"); + std::fs::create_dir(&skill1_dir).unwrap(); + std::fs::write( + skill1_dir.join("SKILL.md"), + r#"--- +name: skill-one +description: First skill +--- +Content 1 +"#, + ) + .unwrap(); + + let skill2_dir = skills_dir.join("skill-two"); + std::fs::create_dir(&skill2_dir).unwrap(); + std::fs::write( + skill2_dir.join("SKILL.md"), + r#"--- +name: skill-two +description: Second skill +disable-model-invocation: true +--- +Content 2 +"#, + ) + .unwrap(); + + let skills = load_skills_from_directory(skills_dir); + + assert_eq!(skills.len(), 2); + let names: Vec<_> = skills + .iter() + .map(|skill| skill.skill_name.as_str()) + .collect(); + assert!(names.contains(&"skill-one")); + assert!(names.contains(&"skill-two")); + + let skill_two = skills + .iter() + .find(|skill| skill.skill_name == "skill-two") + .unwrap(); + 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: + lime_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")); + assert!(skills.is_empty()); + } + + #[test] + fn test_bundled_social_post_with_cover_skill_contract() { + let skill_file = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .join("resources/default-skills/social_post_with_cover/SKILL.md"); + + assert!(skill_file.exists()); + let content = std::fs::read_to_string(&skill_file).unwrap(); + let skill = load_skill_from_file("social_post_with_cover", &skill_file).unwrap(); + + 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![ + "social_generate_cover_image".to_string(), + "search_query".to_string(), + ]) + ); + assert!(content.contains(", + pub model_override: Option, + pub execution_id: Option, + pub session_id: Option, +} + +fn ensure_skill_error_code(code: &str, message: &str) -> String { + if message.contains('|') { + message.to_string() + } else { + format_skill_error(code, message) + } +} + +struct TauriExecutionCallbackAdapter<'a> { + inner: &'a TauriExecutionCallback, +} + +impl<'a> TauriExecutionCallbackAdapter<'a> { + fn new(inner: &'a TauriExecutionCallback) -> Self { + Self { inner } + } +} + +impl ExecutionCallback for TauriExecutionCallbackAdapter<'_> { + fn on_step_start( + &self, + step_id: &str, + step_name: &str, + current_step: usize, + total_steps: usize, + ) { + self.inner + .on_step_start(step_id, step_name, current_step, total_steps); + } + + fn on_step_complete(&self, step_id: &str, output: &str) { + self.inner.on_step_complete(step_id, output); + } + + fn on_step_error(&self, step_id: &str, error: &str, will_retry: bool) { + self.inner.on_step_error(step_id, error, will_retry); + } + + fn on_complete(&self, success: bool, final_output: Option<&str>, error: Option<&str>) { + let mapped_error = if success { + error.map(|value| value.to_string()) + } else { + error.map(|value| ensure_skill_error_code(SKILL_ERR_EXECUTE_FAILED, value)) + }; + self.inner + .on_complete(success, final_output, mapped_error.as_deref()); + } +} + +fn create_skill_event_emitter(app_handle: &AppHandle) -> SkillEventEmitter { + let app_handle = app_handle.clone(); + Arc::new(move |event_name: String, event: TauriAgentEvent| { + if let Err(error) = app_handle.emit(&event_name, &event) { + tracing::error!("[execute_skill_workflow] 发送事件失败: {}", error); + } + }) +} + +fn emit_skill_final_done(app_handle: &AppHandle, execution_id: &str) { + let event_name = format!("skill-exec-{execution_id}"); + if let Err(error) = app_handle.emit(&event_name, TauriAgentEvent::FinalDone { usage: None }) { + tracing::error!("[execute_skill] 发送完成事件失败: {}", error); + } +} + +fn map_execution_error(error: SkillExecutionError) -> String { + match error { + SkillExecutionError::SessionInitFailed(message) => { + format_skill_error(SKILL_ERR_SESSION_INIT_FAILED, message) + } + } +} + +fn map_execution_result(mut result: SkillExecutionResult) -> SkillExecutionResult { + if !result.success { + result.error = result + .error + .take() + .map(|error| ensure_skill_error_code(SKILL_ERR_EXECUTE_FAILED, &error)); + } + result +} + +pub async fn execute_named_skill( + app_handle: &AppHandle, + db: &DbConnection, + api_key_provider_service: &ApiKeyProviderServiceState, + config_manager: &GlobalConfigManagerState, + aster_state: &AsterAgentState, + request: SkillExecutionRequest, +) -> Result { + let SkillExecutionRequest { + skill_name, + user_input, + provider_override, + model_override, + execution_id, + session_id, + } = request; + + let execution_id = execution_id.unwrap_or_else(|| Uuid::new_v4().to_string()); + let session_id = session_id.unwrap_or_else(|| format!("skill-exec-{}", Uuid::new_v4())); + let tracker = ExecutionTracker::new(db.clone()); + let provider_selection = Arc::new(Mutex::new(None)); + let start_metadata = build_skill_run_start_metadata( + skill_name.as_str(), + execution_id.as_str(), + user_input.as_str(), + provider_override.as_deref(), + model_override.as_deref(), + ); + let provider_selection_for_run = Arc::clone(&provider_selection); + let provider_selection_for_finalize = Arc::clone(&provider_selection); + let skill_name_for_run = skill_name.clone(); + let execution_id_for_run = execution_id.clone(); + let session_id_for_run = session_id.clone(); + let user_input_for_run = user_input.clone(); + let provider_override_for_run = provider_override.clone(); + let model_override_for_run = model_override.clone(); + let skill_name_for_finalize = skill_name.clone(); + let execution_id_for_finalize = execution_id.clone(); + let provider_override_for_finalize = provider_override.clone(); + let model_override_for_finalize = model_override.clone(); + let app_handle = app_handle.clone(); + let db = db.clone(); + let api_key_provider_service = ApiKeyProviderServiceState(api_key_provider_service.0.clone()); + let config_manager = GlobalConfigManagerState(config_manager.0.clone()); + let aster_state = aster_state.clone(); + + tracker + .with_run_custom( + RunSource::Skill, + Some(skill_name.clone()), + Some(session_id.clone()), + Some(start_metadata), + async move { + tracing::info!( + "[execute_skill] 开始执行 Skill: name={}, execution_id={}, session_id={}, provider_override={:?}, model_override={:?}", + skill_name_for_run, + execution_id_for_run, + session_id_for_run, + provider_override_for_run, + model_override_for_run + ); + + let skill = load_executable_skill_definition(&skill_name_for_run)?; + let prepared = prepare_skill_execution( + &app_handle, + &db, + &api_key_provider_service, + &config_manager, + &aster_state, + &skill, + &execution_id_for_run, + &session_id_for_run, + provider_override_for_run.as_deref(), + model_override_for_run.as_deref(), + ) + .await?; + + if let Ok(mut slot) = provider_selection_for_run.lock() { + *slot = Some(prepared.provider_selection.clone()); + } else { + tracing::warn!( + "[execute_skill] provider 选择状态锁定失败,运行记录将缺少 resolved provider 元数据" + ); + } + + execute_skill_definition( + &app_handle, + &aster_state, + &skill, + &user_input_for_run, + &execution_id_for_run, + &session_id_for_run, + &prepared.callback, + prepared.memory_prompt.as_deref(), + ) + .await + }, + move |result| { + let provider_selection = provider_selection_for_finalize + .lock() + .ok() + .and_then(|slot| slot.as_ref().cloned()); + build_skill_run_finish_decision( + &skill_name_for_finalize, + &execution_id_for_finalize, + provider_override_for_finalize.as_deref(), + model_override_for_finalize.as_deref(), + provider_selection.as_ref(), + result, + ) + }, + ) + .await +} + +pub async fn execute_skill_prompt( + app_handle: &AppHandle, + aster_state: &AsterAgentState, + skill: &LoadedSkillDefinition, + user_input: &str, + execution_id: &str, + session_id: &str, + callback: &TauriExecutionCallback, + memory_prompt: Option<&str>, +) -> Result { + let callback_adapter = TauriExecutionCallbackAdapter::new(callback); + callback_adapter.on_step_start("main", &skill.display_name, 1, 1); + + let mut result = map_execution_result( + execute_agent_skill_prompt( + aster_state, + skill, + user_input, + execution_id, + session_id, + memory_prompt, + create_skill_event_emitter(app_handle), + ) + .await + .map_err(map_execution_error)?, + ); + + if !result.success { + let error_message = result + .error + .clone() + .unwrap_or_else(|| format_skill_error(SKILL_ERR_EXECUTE_FAILED, "Unknown error")); + callback_adapter.on_step_error("main", &error_message, false); + callback_adapter.on_complete(false, None, Some(&error_message)); + emit_skill_final_done(app_handle, execution_id); + return Ok(result); + } + + let final_output = finalize_skill_output( + app_handle, + &skill.skill_name, + user_input, + execution_id, + result.output.as_deref().unwrap_or(""), + ); + result.output = Some(final_output.clone()); + if let Some(step_result) = result.steps_completed.get_mut(0) { + step_result.output = Some(final_output.clone()); + } + + callback_adapter.on_step_complete("main", &final_output); + callback_adapter.on_complete(true, Some(&final_output), None); + emit_skill_final_done(app_handle, execution_id); + Ok(result) +} + +pub async fn execute_skill_workflow( + app_handle: &AppHandle, + aster_state: &AsterAgentState, + skill: &LoadedSkillDefinition, + user_input: &str, + execution_id: &str, + session_id: &str, + callback: &TauriExecutionCallback, + memory_prompt: Option<&str>, +) -> Result { + let callback_adapter = TauriExecutionCallbackAdapter::new(callback); + execute_agent_skill_workflow(SkillWorkflowExecution { + aster_state, + skill, + user_input, + execution_id, + session_id, + callback: &callback_adapter, + memory_prompt, + emitter: create_skill_event_emitter(app_handle), + }) + .await + .map(map_execution_result) + .map_err(map_execution_error) +} + +pub async fn execute_skill_definition( + app_handle: &AppHandle, + aster_state: &AsterAgentState, + skill: &LoadedSkillDefinition, + user_input: &str, + execution_id: &str, + session_id: &str, + callback: &TauriExecutionCallback, + memory_prompt: Option<&str>, +) -> Result { + if skill.execution_mode == "workflow" && !skill.workflow_steps.is_empty() { + execute_skill_workflow( + app_handle, + aster_state, + skill, + user_input, + execution_id, + session_id, + callback, + memory_prompt, + ) + .await + } else { + execute_skill_prompt( + app_handle, + aster_state, + skill, + user_input, + execution_id, + session_id, + callback, + memory_prompt, + ) + .await + } +} + +pub use lime_agent::{SkillExecutionResult, StepResult}; diff --git a/src-tauri/src/skills/mod.rs b/src-tauri/src/skills/mod.rs index c2905578f..7e039edc7 100644 --- a/src-tauri/src/skills/mod.rs +++ b/src-tauri/src/skills/mod.rs @@ -1,12 +1,30 @@ //! Skills 集成模块 //! //! 纯逻辑已迁移到 `lime-skills` crate, -//! 本模块保留 Tauri 相关实现和兼容导出层。 +//! workflow 执行主链已继续下沉到 `lime-agent`, +//! 本模块只保留 Tauri 适配与兼容导出层。 +mod catalog; mod default_skills; +mod execution; mod execution_callback; mod llm_provider; +mod runtime; +mod social_post; +pub use catalog::{ + get_skill_detail_info, list_executable_skill_catalog, load_executable_skill_definition, + ExecutableSkillInfo, SkillDetailInfo, WorkflowStepInfo, +}; +pub use execution::{ + execute_named_skill, execute_skill_definition, execute_skill_prompt, execute_skill_workflow, + SkillExecutionRequest, SkillExecutionResult, StepResult, +}; +pub use runtime::{ + build_skill_run_finish_decision, build_skill_run_start_metadata, prepare_skill_execution, + PreparedSkillExecution, SkillProviderSelection, +}; +pub use social_post::{collect_social_artifact_paths_from_output, infer_theme_workbench_gate_key}; // Tauri 实现(留在主 crate) pub use default_skills::ensure_default_local_skills; pub use execution_callback::TauriExecutionCallback; diff --git a/src-tauri/src/skills/runtime.rs b/src-tauri/src/skills/runtime.rs new file mode 100644 index 000000000..743091f5e --- /dev/null +++ b/src-tauri/src/skills/runtime.rs @@ -0,0 +1,368 @@ +use crate::agent::AsterAgentState; +use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState; +use crate::commands::aster_agent_cmd::{ + ensure_browser_mcp_tools_registered, ensure_creation_task_tools_registered, + ensure_social_image_tool_registered, +}; +use crate::commands::skill_error::{ + format_skill_error, SKILL_ERR_PROVIDER_UNAVAILABLE, SKILL_ERR_SESSION_INIT_FAILED, +}; +use crate::config::GlobalConfigManagerState; +use crate::database::dao::agent_run::AgentRunStatus; +use crate::database::DbConnection; +use crate::services::execution_tracker_service::RunFinishDecision; +use crate::services::memory_profile_prompt_service::{build_memory_prompt, MemoryPromptContext}; +use lime_skills::LoadedSkillDefinition; +use std::path::Path; + +use super::execution::SkillExecutionResult; +use super::execution_callback::TauriExecutionCallback; +use super::social_post::{ + collect_social_artifact_paths_from_output, infer_theme_workbench_gate_key, +}; + +#[derive(Debug, Clone)] +pub struct SkillProviderSelection { + pub requested_provider: String, + pub requested_model: String, + pub resolved_provider: String, + pub resolved_model: String, +} + +pub struct PreparedSkillExecution { + pub callback: TauriExecutionCallback, + pub memory_prompt: Option, + pub provider_selection: SkillProviderSelection, +} + +const DEFAULT_SKILL_PROVIDER: &str = "anthropic"; +const DEFAULT_SKILL_MODEL: &str = "claude-sonnet-4-20250514"; +const FALLBACK_TOOL_CAPABLE_PROVIDERS: &[(&str, &str)] = &[ + ("anthropic", "claude-sonnet-4-20250514"), + ("openai", "gpt-4o"), + ("gemini", "gemini-2.0-flash"), +]; +const SOCIAL_POST_WITH_COVER_SKILL_NAME: &str = "social_post_with_cover"; + +fn build_skill_memory_prompt( + db: &DbConnection, + config_manager: &GlobalConfigManagerState, + session_id: &str, +) -> Option { + let config = config_manager.config(); + let session_working_dir = lime_agent::get_session_sync(db, session_id) + .ok() + .and_then(|session| session.working_dir) + .filter(|path| !path.trim().is_empty()); + let context = MemoryPromptContext { + working_dir: session_working_dir.as_deref().map(Path::new), + active_relative_path: None, + }; + + build_memory_prompt(&config, context) +} + +fn resolve_requested_provider( + skill: &LoadedSkillDefinition, + provider_override: Option<&str>, + model_override: Option<&str>, +) -> (String, String) { + let requested_provider = provider_override + .map(|value| value.to_string()) + .or_else(|| skill.provider.clone()) + .unwrap_or_else(|| DEFAULT_SKILL_PROVIDER.to_string()); + let requested_model = model_override + .map(|value| value.to_string()) + .or_else(|| skill.model.clone()) + .unwrap_or_else(|| DEFAULT_SKILL_MODEL.to_string()); + (requested_provider, requested_model) +} + +async fn ensure_skill_agent_ready( + app_handle: &tauri::AppHandle, + db: &DbConnection, + api_key_provider_service: &ApiKeyProviderServiceState, + config_manager: &GlobalConfigManagerState, + aster_state: &AsterAgentState, +) -> Result<(), String> { + if !aster_state.is_initialized().await { + tracing::info!("[execute_skill] Agent 未初始化,开始初始化..."); + aster_state.init_agent_with_db(db).await.map_err(|error| { + format_skill_error( + SKILL_ERR_SESSION_INIT_FAILED, + format!("初始化 Agent 失败: {error}"), + ) + })?; + tracing::info!("[execute_skill] Agent 初始化完成"); + } + + ensure_browser_mcp_tools_registered(aster_state) + .await + .map_err(|error| { + format_skill_error( + SKILL_ERR_SESSION_INIT_FAILED, + format!("注册浏览器工具失败: {error}"), + ) + })?; + ensure_social_image_tool_registered(aster_state, config_manager) + .await + .map_err(|error| { + format_skill_error( + SKILL_ERR_SESSION_INIT_FAILED, + format!("注册社媒生图工具失败: {error}"), + ) + })?; + ensure_creation_task_tools_registered(aster_state, db, api_key_provider_service, app_handle) + .await + .map_err(|error| { + format_skill_error( + SKILL_ERR_SESSION_INIT_FAILED, + format!("注册创作任务工具失败: {error}"), + ) + })?; + + Ok(()) +} + +async fn configure_skill_provider_with_fallback( + aster_state: &AsterAgentState, + db: &DbConnection, + session_id: &str, + requested_provider: &str, + requested_model: &str, +) -> Result { + let mut configure_result = aster_state + .configure_provider_from_pool(db, requested_provider, requested_model, session_id) + .await; + + if configure_result.is_err() { + tracing::warn!( + "[execute_skill] 首选 Provider {} 配置失败: {:?},尝试 fallback", + requested_provider, + configure_result.as_ref().err() + ); + + for (fallback_provider, fallback_model) in FALLBACK_TOOL_CAPABLE_PROVIDERS { + if *fallback_provider == requested_provider { + continue; + } + match aster_state + .configure_provider_from_pool(db, fallback_provider, fallback_model, session_id) + .await + { + Ok(config) => { + tracing::info!( + "[execute_skill] Fallback 到 {} / {} 成功", + fallback_provider, + fallback_model + ); + configure_result = Ok(config); + break; + } + Err(error) => { + tracing::warn!( + "[execute_skill] Fallback {} 也失败: {}", + fallback_provider, + error + ); + } + } + } + } + + let configured_provider = configure_result.map_err(|error| { + format_skill_error( + SKILL_ERR_PROVIDER_UNAVAILABLE, + format!( + "无法配置任何可用的 Provider(需要支持工具调用的 Provider,如 Anthropic、OpenAI 或 Google): {error}" + ), + ) + })?; + + Ok(SkillProviderSelection { + requested_provider: requested_provider.to_string(), + requested_model: requested_model.to_string(), + resolved_provider: configured_provider.provider_name, + resolved_model: configured_provider.model_name, + }) +} + +pub async fn prepare_skill_execution( + app_handle: &tauri::AppHandle, + db: &DbConnection, + api_key_provider_service: &ApiKeyProviderServiceState, + config_manager: &GlobalConfigManagerState, + aster_state: &AsterAgentState, + skill: &LoadedSkillDefinition, + execution_id: &str, + session_id: &str, + provider_override: Option<&str>, + model_override: Option<&str>, +) -> Result { + ensure_skill_agent_ready( + app_handle, + db, + api_key_provider_service, + config_manager, + aster_state, + ) + .await?; + + let (requested_provider, requested_model) = + resolve_requested_provider(skill, provider_override, model_override); + let provider_selection = configure_skill_provider_with_fallback( + aster_state, + db, + session_id, + &requested_provider, + &requested_model, + ) + .await?; + + tracing::info!( + "[execute_skill] Provider 配置成功: requested={} / {}, resolved={} / {}", + provider_selection.requested_provider, + provider_selection.requested_model, + provider_selection.resolved_provider, + provider_selection.resolved_model + ); + + Ok(PreparedSkillExecution { + callback: TauriExecutionCallback::new(app_handle.clone(), execution_id.to_string()), + memory_prompt: build_skill_memory_prompt(db, config_manager, session_id), + provider_selection, + }) +} + +pub fn build_skill_run_start_metadata( + skill_name: &str, + execution_id: &str, + user_input: &str, + provider_override: Option<&str>, + model_override: Option<&str>, +) -> serde_json::Value { + serde_json::json!({ + "execution_id": execution_id, + "skill_name": skill_name, + "gate_key": infer_theme_workbench_gate_key(skill_name, user_input), + "provider_override": provider_override, + "model_override": model_override, + }) +} + +fn build_success_metadata( + skill_name: &str, + execution_id: &str, + provider_override: Option<&str>, + model_override: Option<&str>, + provider_selection: Option<&SkillProviderSelection>, + artifact_paths: Vec, +) -> serde_json::Value { + let mut metadata = serde_json::json!({ + "skill_name": skill_name, + "execution_id": execution_id, + "provider_override": provider_override, + "model_override": model_override, + }); + + if let Some(selection) = provider_selection { + metadata["requested_provider"] = serde_json::json!(selection.requested_provider); + metadata["requested_model"] = serde_json::json!(selection.requested_model); + metadata["resolved_provider"] = serde_json::json!(selection.resolved_provider); + metadata["resolved_model"] = serde_json::json!(selection.resolved_model); + } else { + metadata["requested_provider"] = serde_json::json!(provider_override); + metadata["requested_model"] = serde_json::json!(model_override); + } + + if skill_name == SOCIAL_POST_WITH_COVER_SKILL_NAME { + metadata["workflow"] = serde_json::json!("social_content_pipeline_v1"); + metadata["version_id"] = serde_json::json!(execution_id); + metadata["stages"] = serde_json::json!(["topic_select", "write_mode", "publish_confirm"]); + metadata["artifact_paths"] = serde_json::json!(artifact_paths); + } + + metadata +} + +fn build_error_metadata( + skill_name: &str, + execution_id: &str, + provider_override: Option<&str>, + model_override: Option<&str>, + provider_selection: Option<&SkillProviderSelection>, + success: Option, +) -> serde_json::Value { + let mut metadata = serde_json::json!({ + "skill_name": skill_name, + "execution_id": execution_id, + "provider_override": provider_override, + "model_override": model_override, + }); + + if let Some(value) = success { + metadata["success"] = serde_json::json!(value); + } + if let Some(selection) = provider_selection { + metadata["requested_provider"] = serde_json::json!(selection.requested_provider); + metadata["requested_model"] = serde_json::json!(selection.requested_model); + metadata["resolved_provider"] = serde_json::json!(selection.resolved_provider); + metadata["resolved_model"] = serde_json::json!(selection.resolved_model); + } else { + metadata["requested_provider"] = serde_json::json!(provider_override); + metadata["requested_model"] = serde_json::json!(model_override); + } + + metadata +} + +pub fn build_skill_run_finish_decision( + skill_name: &str, + execution_id: &str, + provider_override: Option<&str>, + model_override: Option<&str>, + provider_selection: Option<&SkillProviderSelection>, + result: &Result, +) -> RunFinishDecision { + match result { + Ok(execution) if execution.success => RunFinishDecision { + status: AgentRunStatus::Success, + error_code: None, + error_message: None, + metadata: Some(build_success_metadata( + skill_name, + execution_id, + provider_override, + model_override, + provider_selection, + collect_social_artifact_paths_from_output(execution.output.as_deref()), + )), + }, + Ok(execution) => RunFinishDecision { + status: AgentRunStatus::Error, + error_code: Some("skill_execute_failed".to_string()), + error_message: execution.error.clone(), + metadata: Some(build_error_metadata( + skill_name, + execution_id, + provider_override, + model_override, + provider_selection, + Some(false), + )), + }, + Err(error) => RunFinishDecision { + status: AgentRunStatus::Error, + error_code: Some("skill_execute_failed".to_string()), + error_message: Some(error.clone()), + metadata: Some(build_error_metadata( + skill_name, + execution_id, + provider_override, + model_override, + provider_selection, + None, + )), + }, + } +} diff --git a/src-tauri/src/skills/social_post.rs b/src-tauri/src/skills/social_post.rs new file mode 100644 index 000000000..a8d83766c --- /dev/null +++ b/src-tauri/src/skills/social_post.rs @@ -0,0 +1,561 @@ +use crate::agent::TauriAgentEvent; +use chrono::Utc; +use lime_agent::event_converter::{TauriArtifactSnapshot, TauriToolResult}; +use tauri::{AppHandle, Emitter}; + +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"; +const SOCIAL_POST_EMPTY_FALLBACK_CONTENT: &str = "# 社媒文案\n\n(生成结果为空,请重试。)"; +const SOCIAL_POST_FALLBACK_COVER_URL: &str = "cover-generation-failed"; +const SOCIAL_POST_FALLBACK_COVER_NOTE: &str = "封面图生成失败,可稍后仅重试配图。"; +const SOCIAL_POST_DEFAULT_IMAGE_SIZE: &str = "1024x1024"; + +#[derive(Debug, Clone)] +struct SocialSkillOutputEnvelope { + final_output: String, + file_path: String, + file_content: String, +} + +pub fn infer_theme_workbench_gate_key(skill_name: &str, user_input: &str) -> &'static str { + let probe = format!("{} {}", skill_name, user_input).to_lowercase(); + if probe.contains("publish") + || probe.contains("adapt") + || probe.contains("distribution") + || probe.contains("release") + || probe.contains("发布") + || probe.contains("分发") + || probe.contains("平台适配") + { + return "publish_confirm"; + } + if probe.contains("topic") + || probe.contains("research") + || probe.contains("trend") + || probe.contains("idea") + || probe.contains("选题") + || probe.contains("方向") + || probe.contains("调研") + || probe.contains("洞察") + { + return "topic_select"; + } + "write_mode" +} + +pub fn finalize_skill_output( + app_handle: &AppHandle, + skill_name: &str, + user_input: &str, + execution_id: &str, + raw_output: &str, +) -> String { + let Some(social_output) = + normalize_social_post_output(skill_name, user_input, execution_id, raw_output) + else { + return raw_output.to_string(); + }; + + emit_social_write_file_events( + app_handle, + execution_id, + &social_output.file_path, + &social_output.file_content, + ); + for (artifact_path, artifact_content) in build_social_auxiliary_file_payloads( + execution_id, + user_input, + &social_output.file_path, + &social_output.file_content, + ) { + emit_social_write_file_events(app_handle, execution_id, &artifact_path, &artifact_content); + } + + social_output.final_output +} + +pub fn collect_social_artifact_paths_from_output(output: Option<&str>) -> Vec { + let Some(raw_output) = output else { + return Vec::new(); + }; + let Some((_, maybe_path, _)) = extract_first_write_file_block(raw_output) else { + return Vec::new(); + }; + let Some(article_path) = maybe_path else { + return Vec::new(); + }; + let (cover_meta_path, publish_pack_path) = derive_social_auxiliary_paths(&article_path); + vec![article_path, cover_meta_path, publish_pack_path] +} + +fn normalize_social_post_output( + skill_name: &str, + user_input: &str, + execution_id: &str, + raw_output: &str, +) -> Option { + if skill_name != SOCIAL_POST_WITH_COVER_SKILL_NAME { + return None; + } + + let generated_path = build_social_post_file_path(user_input, execution_id); + if let Some((range, existing_path, content)) = extract_first_write_file_block(raw_output) { + let normalized_content = normalize_social_markdown_contract(&content); + let has_existing_path = existing_path.is_some(); + let path = existing_path.unwrap_or_else(|| generated_path.clone()); + + if has_existing_path { + if normalized_content != content { + let normalized_block = build_write_file_block(&path, &normalized_content); + let mut rebuilt = String::new(); + rebuilt.push_str(&raw_output[..range.start]); + rebuilt.push_str(&normalized_block); + rebuilt.push_str(&raw_output[range.end..]); + return Some(SocialSkillOutputEnvelope { + final_output: rebuilt, + file_path: path, + file_content: normalized_content, + }); + } + return Some(SocialSkillOutputEnvelope { + final_output: raw_output.to_string(), + file_path: path, + file_content: normalized_content, + }); + } + + let normalized_block = build_write_file_block(&path, &normalized_content); + let mut rebuilt = String::new(); + rebuilt.push_str(&raw_output[..range.start]); + rebuilt.push_str(&normalized_block); + rebuilt.push_str(&raw_output[range.end..]); + + return Some(SocialSkillOutputEnvelope { + final_output: rebuilt, + file_path: path, + file_content: normalized_content, + }); + } + + let normalized_content = normalize_social_markdown_contract(raw_output); + Some(SocialSkillOutputEnvelope { + final_output: build_write_file_block(&generated_path, &normalized_content), + file_path: generated_path, + file_content: normalized_content, + }) +} + +fn extract_first_write_file_block( + raw_output: &str, +) -> Option<(std::ops::Range, Option, String)> { + let open_start = raw_output.find("')?; + let open_end = open_start + open_end_offset; + let open_tag = &raw_output[open_start..=open_end]; + + let content_start = open_end + 1; + let close_tag = ""; + let close_offset = raw_output[content_start..].find(close_tag)?; + let close_start = content_start + close_offset; + let block_end = close_start + close_tag.len(); + + let content = raw_output[content_start..close_start].trim().to_string(); + let path = extract_write_file_path(open_tag); + Some((open_start..block_end, path, content)) +} + +fn extract_write_file_path(open_tag: &str) -> Option { + let path_idx = open_tag.find("path")?; + let after_path = &open_tag[path_idx + "path".len()..]; + let equal_idx = after_path.find('=')?; + let value = after_path[equal_idx + 1..].trim_start(); + let quote = value.chars().next()?; + if quote != '"' && quote != '\'' { + return None; + } + + let rest = &value[quote.len_utf8()..]; + let end_idx = rest.find(quote)?; + let path = rest[..end_idx].trim(); + if path.is_empty() { + None + } else { + Some(path.to_string()) + } +} + +fn normalize_social_output_content(content: &str) -> String { + let trimmed = content.trim(); + if trimmed.is_empty() { + SOCIAL_POST_EMPTY_FALLBACK_CONTENT.to_string() + } else { + trimmed.to_string() + } +} + +fn normalize_social_markdown_contract(content: &str) -> String { + let mut normalized = normalize_social_output_content(content); + if !normalized.contains("![封面图](") { + normalized = format!("{normalized}\n\n![封面图]({SOCIAL_POST_FALLBACK_COVER_URL})"); + } + normalized +} + +fn extract_cover_url_from_markdown(content: &str) -> Option { + for line in content.lines() { + let trimmed = line.trim(); + if !trimmed.starts_with("![") { + continue; + } + let open = trimmed.find("](")?; + let close = trimmed.rfind(')')?; + if close <= open + 2 { + continue; + } + let url = trimmed[(open + 2)..close].trim(); + if !url.is_empty() { + return Some(url.to_string()); + } + } + None +} + +fn extract_detail_value(content: &str, label: &str) -> Option { + let probe = format!("- {label}:"); + for line in content.lines() { + let trimmed = line.trim(); + if let Some(value) = trimmed.strip_prefix(&probe) { + let value = value.trim(); + if !value.is_empty() { + return Some(value.to_string()); + } + } + } + None +} + +fn derive_social_auxiliary_paths(article_path: &str) -> (String, String) { + let base = article_path.strip_suffix(".md").unwrap_or(article_path); + ( + format!("{base}.cover.json"), + format!("{base}.publish-pack.json"), + ) +} + +fn summarize_social_content(content: &str) -> String { + let compact = content + .lines() + .filter(|line| !line.trim().starts_with('#')) + .collect::>() + .join(" "); + let compact = compact.split_whitespace().collect::>().join(" "); + compact.chars().take(180).collect() +} + +fn build_social_auxiliary_file_payloads( + execution_id: &str, + user_input: &str, + article_path: &str, + article_content: &str, +) -> Vec<(String, String)> { + let (cover_meta_path, publish_pack_path) = derive_social_auxiliary_paths(article_path); + let cover_url = extract_cover_url_from_markdown(article_content) + .unwrap_or_else(|| SOCIAL_POST_FALLBACK_COVER_URL.to_string()); + let cover_prompt = + extract_detail_value(article_content, "提示词").unwrap_or_else(|| "未提供".to_string()); + let cover_size = extract_detail_value(article_content, "尺寸") + .unwrap_or_else(|| SOCIAL_POST_DEFAULT_IMAGE_SIZE.to_string()); + let cover_status = extract_detail_value(article_content, "状态").unwrap_or_else(|| { + if cover_url == SOCIAL_POST_FALLBACK_COVER_URL { + "失败".to_string() + } else { + "成功".to_string() + } + }); + let cover_remark = extract_detail_value(article_content, "备注").unwrap_or_else(|| { + if cover_status == "失败" { + SOCIAL_POST_FALLBACK_COVER_NOTE.to_string() + } else { + "".to_string() + } + }); + + let cover_meta = serde_json::json!({ + "execution_id": execution_id, + "article_path": article_path, + "cover_url": cover_url, + "prompt": cover_prompt, + "size": cover_size, + "status": cover_status, + "remark": cover_remark, + "generated_at": Utc::now().to_rfc3339(), + }); + + let publish_pack = serde_json::json!({ + "execution_id": execution_id, + "pipeline": ["topic_select", "write_mode", "publish_confirm"], + "article_path": article_path, + "cover_meta_path": cover_meta_path, + "source_input": user_input, + "recommended_channels": ["xiaohongshu", "wechat"], + "summary": summarize_social_content(article_content), + "generated_at": Utc::now().to_rfc3339(), + }); + + vec![ + ( + cover_meta_path, + serde_json::to_string_pretty(&cover_meta).unwrap_or_else(|_| cover_meta.to_string()), + ), + ( + publish_pack_path, + serde_json::to_string_pretty(&publish_pack) + .unwrap_or_else(|_| publish_pack.to_string()), + ), + ] +} + +fn build_write_file_block(file_path: &str, file_content: &str) -> String { + format!("\n{file_content}\n") +} + +fn build_social_post_file_path(user_input: &str, execution_id: &str) -> String { + let timestamp = Utc::now().format("%Y%m%d-%H%M%S"); + let slug = build_social_post_slug(user_input); + let suffix = build_execution_suffix(execution_id); + format!("{SOCIAL_POST_OUTPUT_DIR}/{timestamp}-{slug}-{suffix}.md") +} + +fn build_social_post_slug(user_input: &str) -> String { + let mut normalized = String::new(); + let mut last_was_dash = false; + + for ch in user_input.chars() { + if ch.is_ascii_alphanumeric() { + normalized.push(ch.to_ascii_lowercase()); + last_was_dash = false; + continue; + } + + if !last_was_dash { + normalized.push('-'); + last_was_dash = true; + } + } + + let trimmed = normalized.trim_matches('-'); + let truncated: String = trimmed.chars().take(24).collect(); + if truncated.is_empty() { + "post".to_string() + } else { + truncated + } +} + +fn build_execution_suffix(execution_id: &str) -> String { + let normalized: String = execution_id + .chars() + .filter(|ch| ch.is_ascii_alphanumeric()) + .take(6) + .collect(); + if normalized.is_empty() { + "run".to_string() + } else { + normalized.to_ascii_lowercase() + } +} + +fn build_social_tool_event_id(execution_id: &str, file_path: &str) -> String { + let mut hash: u32 = 0x811c9dc5; + for byte in file_path.as_bytes() { + hash ^= u32::from(*byte); + hash = hash.wrapping_mul(0x01000193); + } + format!("social-write-{execution_id}-{hash:08x}") +} + +fn emit_social_write_file_events( + app_handle: &AppHandle, + execution_id: &str, + file_path: &str, + file_content: &str, +) { + 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(), + tool_id: tool_id.clone(), + arguments: Some(arguments), + }; + if let Err(err) = app_handle.emit(&event_name, &tool_start) { + 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 { + success: true, + output: format!("写入社媒文稿: {file_path}"), + error: None, + images: None, + metadata: Some(tool_end_metadata), + }, + }; + if let Err(err) = app_handle.emit(&event_name, &tool_end) { + tracing::warn!("[execute_skill] 发送社媒写入工具完成事件失败: {}", err); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_normalize_social_post_output_wraps_plain_markdown() { + let normalized = normalize_social_post_output( + SOCIAL_POST_WITH_COVER_SKILL_NAME, + "春季上新", + "exec123456", + "# 标题\n\n正文内容", + ) + .expect("should normalize"); + + assert!(normalized + .final_output + .contains("\n# 标题\n\n正文\n"; + let normalized = normalize_social_post_output( + SOCIAL_POST_WITH_COVER_SKILL_NAME, + "春季上新", + "exec123456", + raw_output, + ) + .expect("should normalize"); + + assert_eq!(normalized.file_path, "social-posts/custom-post.md"); + assert!(normalized + .final_output + .contains("social-posts/custom-post.md")); + assert!(normalized.file_content.contains("# 标题")); + assert!(normalized.file_content.contains("![封面图](")); + } + + #[test] + fn test_normalize_social_post_output_injects_missing_path() { + let raw_output = "前置说明\n\n# 标题\n\n正文\n\n后置说明"; + let normalized = normalize_social_post_output( + SOCIAL_POST_WITH_COVER_SKILL_NAME, + "spring launch", + "exec123456", + raw_output, + ) + .expect("should normalize"); + + assert!(normalized.final_output.contains("前置说明")); + assert!(normalized.final_output.contains("后置说明")); + assert!(normalized + .final_output + .contains("\n# 标题\n\n正文\n"; + let paths = collect_social_artifact_paths_from_output(Some(output)); + assert_eq!(paths.len(), 3); + assert_eq!(paths[0], "social-posts/demo.md"); + assert!(paths[1].ends_with(".cover.json")); + assert!(paths[2].ends_with(".publish-pack.json")); + } + + #[test] + fn test_build_social_post_slug_fallback_to_post() { + assert_eq!(build_social_post_slug(""), "post"); + assert_eq!(build_social_post_slug("!!!"), "post"); + assert_eq!( + build_social_post_slug("Spring Launch 2026"), + "spring-launch-2026" + ); + } +} diff --git a/src-tauri/src/workspace_support.rs b/src-tauri/src/workspace_support.rs new file mode 100644 index 000000000..a4401cdbd --- /dev/null +++ b/src-tauri/src/workspace_support.rs @@ -0,0 +1,68 @@ +use crate::workspace::{Workspace, WorkspaceManager, WorkspaceType}; +use lime_core::app_paths; +use std::path::PathBuf; + +pub(crate) fn get_workspace_projects_root_dir() -> Result { + app_paths::resolve_projects_dir() +} + +pub(crate) fn resolve_default_project_path() -> Result { + app_paths::resolve_default_project_dir() +} + +pub(crate) fn sanitize_project_dir_name(name: &str) -> String { + let sanitized: String = name + .trim() + .chars() + .map(|ch| match ch { + '/' | '\\' | ':' | '*' | '?' | '"' | '<' | '>' | '|' => '_', + _ if ch.is_control() => '_', + _ => ch, + }) + .collect(); + + let trimmed = sanitized.trim().trim_matches('.').to_string(); + if trimmed.is_empty() { + "未命名项目".to_string() + } else { + trimmed + } +} + +pub(crate) fn get_or_create_default_project( + manager: &WorkspaceManager, +) -> Result { + if let Some(workspace) = manager.get_default()? { + return Ok(workspace); + } + + let default_project_path = resolve_default_project_path()?; + std::fs::create_dir_all(&default_project_path) + .map_err(|e| format!("创建默认项目目录失败: {e}"))?; + + let workspace = manager.create_with_type( + "默认项目".to_string(), + default_project_path, + WorkspaceType::Persistent, + )?; + manager.set_default(&workspace.id)?; + + manager + .get(&workspace.id)? + .ok_or_else(|| "创建默认项目失败".to_string()) +} + +#[cfg(test)] +mod tests { + use super::sanitize_project_dir_name; + + #[test] + fn sanitize_project_dir_name_should_replace_invalid_chars() { + assert_eq!(sanitize_project_dir_name(" a/b:c*?d "), "a_b_c__d"); + } + + #[test] + fn sanitize_project_dir_name_should_fallback_when_empty() { + assert_eq!(sanitize_project_dir_name(" .. "), "未命名项目"); + } +} diff --git a/src-tauri/tauri.conf.headless.json b/src-tauri/tauri.conf.headless.json index f03e07182..51b201e89 100644 --- a/src-tauri/tauri.conf.headless.json +++ b/src-tauri/tauri.conf.headless.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "Lime", - "version": "0.90.0", + "version": "0.91.0", "identifier": "com.lime.app", "build": { "beforeDevCommand": "npm run dev:web-bridge", diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 5e2e82480..790a4e0a9 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": "Lime", - "version": "0.90.0", + "version": "0.91.0", "identifier": "com.lime.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src/App.tsx b/src/App.tsx index 921f96aa0..82f389da0 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -8,26 +8,12 @@ * _需求: 2.2, 3.2, 5.2_ */ -import React, { useState, useEffect, useCallback } from "react"; +import React, { Suspense, lazy, useState, useEffect, useCallback } from "react"; import styled from "styled-components"; import { getWindowsStartupDiagnostics } from "@/lib/api/serverRuntime"; import { withI18nPatch } from "./i18n/withI18nPatch"; import { SplashScreen } from "./components/SplashScreen"; import { AppSidebar } from "./components/AppSidebar"; -import { SettingsPageV2 } from "./components/settings-v2"; -import { ToolsPage } from "./components/tools/ToolsPage"; -import { ResourcesPage } from "./components/resources"; -import { MemoryPage } from "./components/memory"; -import { StylePage } from "./components/style"; -import { AgentChatPage } from "./components/agent"; -import { PluginsPage } from "./components/plugins/PluginsPage"; -import { ImageGenPage } from "./components/image-gen"; -import { AutomationPage } from "./components/automation"; -import { OpenClawPage } from "./components/openclaw"; -import { RecentImageInsertFloating } from "./components/image-gen/RecentImageInsertFloating"; -import { CreateProjectDialog } from "./components/projects/CreateProjectDialog"; -import { WorkbenchPage } from "./components/workspace"; -import { BrowserRuntimeWorkspace } from "@/features/browser-runtime"; import { ProjectType, createProject, @@ -35,14 +21,7 @@ import { isUserProjectType, resolveProjectRootPath, } from "./lib/api/project"; -import { - TerminalWorkspace, - SysinfoView, - FileBrowserView, - WebView, -} from "./components/terminal"; -import { OnboardingWizard, useOnboardingState } from "./components/onboarding"; -import { ConnectConfirmDialog } from "./components/connect"; +import { useOnboardingState } from "./components/onboarding"; import { showRegistryLoadError } from "./lib/utils/connectError"; import { useDeepLink } from "./hooks/useDeepLink"; import { useRelayRegistry } from "./hooks/useRelayRegistry"; @@ -116,6 +95,123 @@ const THEME_WORKSPACE_PAGES: ThemeWorkspacePage[] = [ "workspace-novel", ]; +const SettingsPageV2 = lazy(() => + import("./components/settings-v2").then((module) => ({ + default: module.SettingsPageV2, + })), +); +const ToolsPage = lazy(() => + import("./components/tools/ToolsPage").then((module) => ({ + default: module.ToolsPage, + })), +); +const ResourcesPage = lazy(() => + import("./components/resources").then((module) => ({ + default: module.ResourcesPage, + })), +); +const MemoryPage = lazy(() => + import("./components/memory").then((module) => ({ + default: module.MemoryPage, + })), +); +const StylePage = lazy(() => + import("./components/style").then((module) => ({ + default: module.StylePage, + })), +); +const PluginsPage = lazy(() => + import("./components/plugins/PluginsPage").then((module) => ({ + default: module.PluginsPage, + })), +); +const ImageGenPage = lazy(() => + import("./components/image-gen").then((module) => ({ + default: module.ImageGenPage, + })), +); +const AutomationPage = lazy(() => + import("./components/automation").then((module) => ({ + default: module.AutomationPage, + })), +); +const OpenClawPage = lazy(() => + import("./components/openclaw").then((module) => ({ + default: module.OpenClawPage, + })), +); +const RecentImageInsertFloating = lazy(() => + import("./components/image-gen/RecentImageInsertFloating").then((module) => ({ + default: module.RecentImageInsertFloating, + })), +); +const CreateProjectDialog = lazy(() => + import("./components/projects/CreateProjectDialog").then((module) => ({ + default: module.CreateProjectDialog, + })), +); +const WorkbenchPage = lazy(() => + import("./components/workspace").then((module) => ({ + default: module.WorkbenchPage, + })), +); +const BrowserRuntimeWorkspace = lazy(() => + import("@/features/browser-runtime").then((module) => ({ + default: module.BrowserRuntimeWorkspace, + })), +); +const TerminalWorkspace = lazy(() => + import("./components/terminal").then((module) => ({ + default: module.TerminalWorkspace, + })), +); +const SysinfoView = lazy(() => + import("./components/terminal").then((module) => ({ + default: module.SysinfoView, + })), +); +const FileBrowserView = lazy(() => + import("./components/terminal").then((module) => ({ + default: module.FileBrowserView, + })), +); +const WebView = lazy(() => + import("./components/terminal").then((module) => ({ + default: module.WebView, + })), +); +const OnboardingWizard = lazy(() => + import("./components/onboarding").then((module) => ({ + default: module.OnboardingWizard, + })), +); +const ConnectConfirmDialog = lazy(() => + import("./components/connect").then((module) => ({ + default: module.ConnectConfirmDialog, + })), +); +const AgentChatPage = lazy(() => + import("./components/agent/chat").then((module) => ({ + default: module.AgentChatPage, + })), +); + +const pageLoadingFallback = ( +
+ 页面加载中... +
+); + function isTauriDesktopEnvironment(): boolean { return hasTauriInvokeCapability(); } @@ -691,7 +787,11 @@ function AppContent() { } if (needsOnboarding) { - return ; + return ( + + + + ); } const currentAgentParams = pageParams as AgentPageParams; @@ -726,35 +826,43 @@ function AppContent() { /> )} - {renderCurrentPage()} + + {renderCurrentPage()} + - + + + - + + + - { - setProjectDialogOpen(open); - if (!open) { - setPendingRecommendation(null); - } - }} - onSubmit={handleCreateProjectFromRecommendation} - defaultType={pendingRecommendation?.projectType} - defaultName={pendingRecommendation?.projectName} - /> + + { + setProjectDialogOpen(open); + if (!open) { + setPendingRecommendation(null); + } + }} + onSubmit={handleCreateProjectFromRecommendation} + defaultType={pendingRecommendation?.projectType} + defaultName={pendingRecommendation?.projectName} + /> + diff --git a/src/components/SplashScreen.tsx b/src/components/SplashScreen.tsx index 2fe21d29f..fdffeac79 100644 --- a/src/components/SplashScreen.tsx +++ b/src/components/SplashScreen.tsx @@ -46,8 +46,16 @@ const Container = styled.div<{ $isExiting: boolean }>` justify-content: center; overflow: hidden; background: - radial-gradient(circle at 20% 18%, rgba(132, 204, 22, 0.18), transparent 30%), - radial-gradient(circle at 78% 12%, rgba(250, 204, 21, 0.14), transparent 28%), + radial-gradient( + circle at 20% 18%, + rgba(132, 204, 22, 0.18), + transparent 30% + ), + radial-gradient( + circle at 78% 12%, + rgba(250, 204, 21, 0.14), + transparent 28% + ), radial-gradient(circle at 50% 84%, rgba(34, 197, 94, 0.1), transparent 28%), linear-gradient( 180deg, @@ -108,8 +116,12 @@ const LogoGlow = styled.div` position: absolute; inset: 12% 12% 16%; border-radius: 999px; - background: - radial-gradient(circle, rgba(163, 230, 53, 0.34) 0%, rgba(163, 230, 53, 0.12) 44%, transparent 72%); + background: radial-gradient( + circle, + rgba(163, 230, 53, 0.34) 0%, + rgba(163, 230, 53, 0.12) 44%, + transparent 72% + ); filter: blur(24px); animation: ${glowPulse} 2.8s ease-in-out infinite; `; @@ -170,12 +182,11 @@ const ProgressTrack = styled.div` width: min(320px, 72vw); height: 8px; border-radius: 999px; - background: - linear-gradient( - 90deg, - hsl(var(--muted) / 0.82) 0%, - hsl(var(--muted) / 0.96) 100% - ); + background: linear-gradient( + 90deg, + hsl(var(--muted) / 0.82) 0%, + hsl(var(--muted) / 0.96) 100% + ); box-shadow: inset 0 1px 0 rgba(255, 255, 255, 0.32), 0 12px 28px rgba(15, 23, 42, 0.08); @@ -186,12 +197,11 @@ const ProgressBar = styled.div` inset: 0 auto 0 0; width: 44%; border-radius: inherit; - background: - linear-gradient( - 90deg, - rgba(132, 204, 22, 0.96) 0%, - rgba(250, 204, 21, 0.9) 100% - ); + background: linear-gradient( + 90deg, + rgba(132, 204, 22, 0.96) 0%, + rgba(250, 204, 21, 0.9) 100% + ); box-shadow: 0 0 24px rgba(163, 230, 53, 0.35); animation: ${progressShift} 1.6s ease-in-out infinite; `; @@ -199,11 +209,13 @@ const ProgressBar = styled.div` interface SplashScreenProps { onComplete: () => void; duration?: number; + exitDuration?: number; } export function SplashScreen({ onComplete, - duration = 1500, + duration = 220, + exitDuration = 180, }: SplashScreenProps) { const [isExiting, setIsExiting] = useState(false); @@ -214,13 +226,13 @@ export function SplashScreen({ const completeTimer = setTimeout(() => { onComplete(); - }, duration + 500); + }, duration + exitDuration); return () => { clearTimeout(exitTimer); clearTimeout(completeTimer); }; - }, [duration, onComplete]); + }, [duration, exitDuration, onComplete]); return ( diff --git a/src/components/agent/chat/components/EmptyState.tsx b/src/components/agent/chat/components/EmptyState.tsx index afc5d9c64..cc1a328cc 100644 --- a/src/components/agent/chat/components/EmptyState.tsx +++ b/src/components/agent/chat/components/EmptyState.tsx @@ -50,6 +50,10 @@ import type { Character } from "@/lib/api/memory"; import type { Skill } from "@/lib/api/skills"; import type { MessageImage } from "../types"; import { isGeneralResearchTheme } from "../utils/generalAgentPrompt"; +import { + getClipboardImageCandidates, + readImageAttachment, +} from "../utils/imageAttachments"; // Import Assets import capabilitySkillsPlaceholder from "@/assets/claw-home/capability-skills-placeholder.svg"; @@ -528,28 +532,45 @@ export const EmptyState: React.FC = ({ if (!files || files.length === 0) return; Array.from(files).forEach((file) => { - if (!file.type.startsWith("image/")) { - return; - } - - const reader = new FileReader(); - reader.onload = (event) => { - const base64 = event.target?.result as string; - const base64Data = base64.split(",")[1]; - setPendingImages((prev) => [ - ...prev, - { - data: base64Data, - mediaType: file.type, - }, - ]); - }; - reader.readAsDataURL(file); + void readImageAttachment(file) + .then((image) => { + setPendingImages((prev) => [...prev, image]); + }) + .catch(() => { + toast.error(`图片读取失败: ${file.name || "未命名图片"}`); + }); }); e.target.value = ""; }; + const handlePaste = (event: React.ClipboardEvent) => { + const imageFiles = getClipboardImageCandidates(event.clipboardData); + if (imageFiles.length === 0) { + return; + } + + event.preventDefault(); + imageFiles.forEach(({ file, mediaType }, index) => { + void readImageAttachment(file, mediaType) + .then((image) => { + setPendingImages((prev) => [...prev, image]); + if (index === 0) { + toast.success("已粘贴图片"); + } + }) + .catch(() => { + toast.error(`图片读取失败: ${file.name || "未命名图片"}`); + }); + }); + }; + + const handleRemoveImage = (index: number) => { + setPendingImages((prev) => + prev.filter((_, currentIndex) => currentIndex !== index), + ); + }; + const handleSend = () => { if (!input.trim() && !isEntryTheme && pendingImages.length === 0) return; const imagesToSend = pendingImages.length > 0 ? pendingImages : undefined; @@ -1054,8 +1075,10 @@ export const EmptyState: React.FC = ({ onSubagentEnabledChange={onSubagentEnabledChange} webSearchEnabled={webSearchEnabled} onWebSearchEnabledChange={onWebSearchEnabledChange} - pendingImagesCount={pendingImages.length} + pendingImages={pendingImages} onFileSelect={handleFileSelect} + onPaste={handlePaste} + onRemoveImage={handleRemoveImage} /> ); diff --git a/src/components/agent/chat/components/EmptyStateComposerPanel.test.tsx b/src/components/agent/chat/components/EmptyStateComposerPanel.test.tsx new file mode 100644 index 000000000..1fd1adc6c --- /dev/null +++ b/src/components/agent/chat/components/EmptyStateComposerPanel.test.tsx @@ -0,0 +1,169 @@ +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 { EmptyStateComposerPanel } from "./EmptyStateComposerPanel"; + +vi.mock("./ChatModelSelector", () => ({ + ChatModelSelector: () =>
, +})); + +vi.mock("./Inputbar/components/CharacterMention", () => ({ + CharacterMention: () =>
, +})); + +vi.mock("./Inputbar/components/SkillBadge", () => ({ + SkillBadge: () =>
, +})); + +vi.mock("./Inputbar/components/SkillSelector", () => ({ + SkillSelector: () =>
, +})); + +const mountedRoots: Array<{ root: Root; container: HTMLDivElement }> = []; + +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(); + } + vi.clearAllMocks(); +}); + +function renderPanel( + props?: Partial>, +) { + const container = document.createElement("div"); + document.body.appendChild(container); + const root = createRoot(container); + + const defaultProps: React.ComponentProps = { + input: "", + setInput: vi.fn(), + placeholder: "输入内容", + onSend: vi.fn(), + activeTheme: "general", + providerType: "openai", + setProviderType: vi.fn(), + model: "gpt-4.1", + setModel: vi.fn(), + executionStrategy: "react", + executionStrategyLabel: "ReAct", + setExecutionStrategy: vi.fn(), + onManageProviders: vi.fn(), + isGeneralTheme: false, + isEntryTheme: false, + entryTaskType: "direct", + entryTaskTypes: [], + getEntryTaskTemplate: vi.fn(), + entryTemplate: { + type: "direct", + label: "直接写作", + description: "直接按需求写作", + pattern: "{input}", + slots: [], + }, + entryPreview: "", + entrySlotValues: {}, + onEntryTaskTypeChange: vi.fn(), + onEntrySlotChange: vi.fn(), + characters: [], + skills: [], + activeSkill: null, + setActiveSkill: vi.fn(), + clearActiveSkill: vi.fn(), + isSkillsLoading: false, + onNavigateToSettings: vi.fn(), + onImportSkill: vi.fn(), + onRefreshSkills: vi.fn(), + showCreationModeSelector: false, + creationMode: "guided", + onCreationModeChange: vi.fn(), + platform: "xiaohongshu", + setPlatform: vi.fn(), + depth: "deep", + setDepth: vi.fn(), + ratio: "3:4", + setRatio: vi.fn(), + style: "minimal", + setStyle: vi.fn(), + ratioPopoverOpen: false, + setRatioPopoverOpen: vi.fn(), + stylePopoverOpen: false, + setStylePopoverOpen: vi.fn(), + thinkingEnabled: false, + onThinkingEnabledChange: vi.fn(), + taskEnabled: false, + onTaskEnabledChange: vi.fn(), + subagentEnabled: false, + onSubagentEnabledChange: vi.fn(), + webSearchEnabled: false, + onWebSearchEnabledChange: vi.fn(), + pendingImages: [], + onFileSelect: vi.fn(), + onPaste: vi.fn(), + onRemoveImage: vi.fn(), + }; + + act(() => { + root.render(); + }); + + mountedRoots.push({ root, container }); + return container; +} + +describe("EmptyStateComposerPanel", () => { + it("应将 onPaste 绑定到输入框", () => { + const onPaste = vi.fn(); + const container = renderPanel({ onPaste }); + const textarea = container.querySelector("textarea"); + + expect(textarea).toBeTruthy(); + + act(() => { + textarea?.dispatchEvent(new Event("paste", { bubbles: true })); + }); + + expect(onPaste).toHaveBeenCalledTimes(1); + }); + + it("有待发送图片时应显示预览并支持删除", () => { + const onRemoveImage = vi.fn(); + const container = renderPanel({ + pendingImages: [ + { + data: "aGVsbG8=", + mediaType: "image/png", + }, + ], + onRemoveImage, + }); + + expect(container.querySelector('img[alt="待发送图片 1"]')).toBeTruthy(); + + const removeButton = container.querySelector( + 'button[aria-label="移除待发送图片 1"]', + ) as HTMLButtonElement | null; + + expect(removeButton).toBeTruthy(); + + act(() => { + removeButton?.click(); + }); + + expect(onRemoveImage).toHaveBeenCalledWith(0); + }); +}); diff --git a/src/components/agent/chat/components/EmptyStateComposerPanel.tsx b/src/components/agent/chat/components/EmptyStateComposerPanel.tsx index 7add633a5..217ab55bb 100644 --- a/src/components/agent/chat/components/EmptyStateComposerPanel.tsx +++ b/src/components/agent/chat/components/EmptyStateComposerPanel.tsx @@ -11,6 +11,7 @@ import { Paperclip, Search, Workflow, + X, } from "lucide-react"; import { Button } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; @@ -41,6 +42,7 @@ import type { } from "./types"; import type { Character } from "@/lib/api/memory"; import type { Skill } from "@/lib/api/skills"; +import type { MessageImage } from "../types"; import iconXhs from "@/assets/platforms/xhs.png"; import iconGzh from "@/assets/platforms/gzh.png"; @@ -227,6 +229,52 @@ const Toolbar = styled.div` border-bottom-right-radius: 24px; `; +const PendingImagesRow = styled.div` + display: flex; + flex-wrap: wrap; + gap: 10px; + padding: 0 18px 12px; +`; + +const PendingImageItem = styled.div` + position: relative; + width: 64px; + height: 64px; + overflow: hidden; + border-radius: 16px; + border: 1px solid rgba(226, 232, 240, 0.92); + background: rgba(248, 250, 252, 0.96); + box-shadow: 0 10px 20px -18px rgba(15, 23, 42, 0.3); +`; + +const PendingImagePreview = styled.img` + width: 100%; + height: 100%; + object-fit: cover; +`; + +const PendingImageRemoveButton = styled.button` + position: absolute; + top: 6px; + right: 6px; + display: inline-flex; + align-items: center; + justify-content: center; + width: 20px; + height: 20px; + border: none; + border-radius: 999px; + color: #fff; + background: rgba(15, 23, 42, 0.72); + box-shadow: 0 6px 14px -10px rgba(15, 23, 42, 0.5); + cursor: pointer; + transition: background-color 0.16s ease; + + &:hover { + background: rgba(220, 38, 38, 0.92); + } +`; + const ToolLoginLeft = styled.div` display: flex; align-items: center; @@ -473,8 +521,10 @@ interface EmptyStateComposerPanelProps { onSubagentEnabledChange?: (enabled: boolean) => void; webSearchEnabled: boolean; onWebSearchEnabledChange?: (enabled: boolean) => void; - pendingImagesCount: number; + pendingImages: MessageImage[]; onFileSelect: (event: React.ChangeEvent) => void; + onPaste?: (event: React.ClipboardEvent) => void; + onRemoveImage?: (index: number) => void; } export function EmptyStateComposerPanel({ @@ -533,8 +583,10 @@ export function EmptyStateComposerPanel({ onSubagentEnabledChange, webSearchEnabled, onWebSearchEnabledChange, - pendingImagesCount, + pendingImages, onFileSelect, + onPaste, + onRemoveImage, }: EmptyStateComposerPanelProps) { const textareaRef = useRef(null); const imageInputRef = useRef(null); @@ -614,6 +666,7 @@ export function EmptyStateComposerPanel({ value={input} onChange={(event) => setInput(event.target.value)} onKeyDown={handleKeyDown} + onPaste={onPaste} placeholder={placeholder} /> @@ -636,10 +689,29 @@ export function EmptyStateComposerPanel({ onChange={onFileSelect} /> - {pendingImagesCount > 0 ? ( -
- 已添加图片 {pendingImagesCount} 张 -
+ {pendingImages.length > 0 ? ( + <> +
+ 已添加图片 {pendingImages.length} 张 +
+ + {pendingImages.map((image, index) => ( + + + onRemoveImage?.(index)} + > + + + + ))} + + ) : null} @@ -801,7 +873,7 @@ export function EmptyStateComposerPanel({ @@ -851,7 +923,7 @@ export function EmptyStateComposerPanel({ @@ -995,9 +1067,7 @@ export function EmptyStateComposerPanel({ 开始生成 diff --git a/src/components/agent/chat/components/HarnessStatusPanel.test.tsx b/src/components/agent/chat/components/HarnessStatusPanel.test.tsx index 1a24a15a0..4d8f19387 100644 --- a/src/components/agent/chat/components/HarnessStatusPanel.test.tsx +++ b/src/components/agent/chat/components/HarnessStatusPanel.test.tsx @@ -209,6 +209,23 @@ describe("HarnessStatusPanel", () => { expect(document.body.textContent).toContain("等待首个模型事件"); }); + it("仅有计划摘要兜底时也应在工作台显示已就绪计划状态", () => { + renderPanel({ + harnessState: createHarnessState({ + plan: { + phase: "ready", + items: [], + summaryText: "已决定:直接回答优先\n当前请求无需工具介入。", + }, + }), + }); + + expect(document.body.textContent).toContain("计划状态"); + expect(document.body.textContent).toContain("已就绪"); + expect(document.body.textContent).toContain("已决定:直接回答优先"); + expect(document.body.textContent).toContain("规划状态"); + }); + it("存在 activeFileWrites 时应在工作台中展示当前文件写入", () => { renderPanel({ harnessState: createHarnessState({ diff --git a/src/components/agent/chat/components/HarnessStatusPanel.tsx b/src/components/agent/chat/components/HarnessStatusPanel.tsx index 72e075c0f..da2763d34 100644 --- a/src/components/agent/chat/components/HarnessStatusPanel.tsx +++ b/src/components/agent/chat/components/HarnessStatusPanel.tsx @@ -1224,7 +1224,10 @@ export function HarnessStatusPanel({ : harnessState.plan.phase === "ready" ? "已就绪" : "空闲", - hint: harnessState.plan.items[0]?.content || "未检测到显式计划快照", + hint: + harnessState.plan.items[0]?.content || + harnessState.plan.summaryText || + "未检测到显式计划快照", icon: ListChecks, }, { @@ -1247,6 +1250,7 @@ export function HarnessStatusPanel({ harnessState.pendingApprovals.length, harnessState.plan.items, harnessState.plan.phase, + harnessState.plan.summaryText, harnessState.recentFileEvents, harnessState.runtimeStatus, ]); @@ -2136,7 +2140,8 @@ export function HarnessStatusPanel({ )) ) : (
- 已进入规划流程,但暂无可展示的 Todo 快照。 + {harnessState.plan.summaryText || + "已进入规划流程,但暂无可展示的 Todo 快照。"}
)}
diff --git a/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx b/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx index f45efdac0..a02b444ab 100644 --- a/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx +++ b/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx @@ -9,6 +9,7 @@ import { InputbarCore } from "./InputbarCore"; import { SkillSelector } from "./SkillSelector"; import { ThemeWorkbenchStatusPanel } from "./ThemeWorkbenchStatusPanel"; import { InputbarModelExtra } from "./InputbarModelExtra"; +import { InputbarVisionCapabilityNotice } from "./InputbarVisionCapabilityNotice"; import { InputbarExecutionStrategySelect } from "./InputbarExecutionStrategySelect"; import { isGeneralResearchTheme } from "../../../utils/generalAgentPrompt"; import type { @@ -94,6 +95,9 @@ export const InputbarComposerSection: React.FC< }) => { const showSkillSelector = !isThemeWorkbenchVariant && isGeneralResearchTheme(activeTheme); + const currentPendingImages = + (inputAdapter.state.attachments as MessageImage[] | undefined) || + pendingImages; if (renderThemeWorkbenchGeneratingPanel) { return ( @@ -141,8 +145,7 @@ export const InputbarComposerSection: React.FC< executionStrategy={executionStrategy} showExecutionStrategy={false} pendingImages={ - (inputAdapter.state.attachments as MessageImage[] | undefined) || - pendingImages + currentPendingImages } onRemoveImage={onRemoveImage} onPaste={onPaste} @@ -159,7 +162,16 @@ export const InputbarComposerSection: React.FC< showTranslate={!isThemeWorkbenchVariant} showDragHandle={!isThemeWorkbenchVariant} visualVariant={isThemeWorkbenchVariant ? "floating" : "default"} - topExtra={topExtra} + topExtra={ + <> + {topExtra} + 0} + /> + + } activeTheme={activeTheme} queuedTurns={queuedTurns} onRemoveQueuedTurn={onRemoveQueuedTurn} diff --git a/src/components/agent/chat/components/Inputbar/components/InputbarCore.test.tsx b/src/components/agent/chat/components/Inputbar/components/InputbarCore.test.tsx index eeab827c6..fbcfdc2f8 100644 --- a/src/components/agent/chat/components/Inputbar/components/InputbarCore.test.tsx +++ b/src/components/agent/chat/components/Inputbar/components/InputbarCore.test.tsx @@ -179,4 +179,30 @@ describe("InputbarCore", () => { expect(onSend).toHaveBeenCalledTimes(1); expect(onStop).toHaveBeenCalledTimes(1); }); + + it("点击图片删除按钮应触发 onRemoveImage", () => { + const onRemoveImage = vi.fn(); + const container = renderInputbarCore({ + pendingImages: [ + { + data: "aGVsbG8=", + mediaType: "image/png", + }, + ], + onRemoveImage, + }); + + const removeButton = container.querySelector( + 'button[aria-label="移除图片 1"]', + ) as HTMLButtonElement | null; + + expect(removeButton).toBeTruthy(); + + act(() => { + removeButton?.dispatchEvent(new MouseEvent("mousedown", { bubbles: true })); + removeButton?.dispatchEvent(new MouseEvent("click", { bubbles: true })); + }); + + expect(onRemoveImage).toHaveBeenCalledWith(0); + }); }); diff --git a/src/components/agent/chat/components/Inputbar/components/InputbarCore.tsx b/src/components/agent/chat/components/Inputbar/components/InputbarCore.tsx index a9d4ce024..06fbce79b 100644 --- a/src/components/agent/chat/components/Inputbar/components/InputbarCore.tsx +++ b/src/components/agent/chat/components/Inputbar/components/InputbarCore.tsx @@ -182,6 +182,23 @@ export const InputbarCore: React.FC = ({ }); }, [isFloatingVariant, toolMode]); + const handleRemoveImageMouseDown = useCallback( + (event: React.MouseEvent) => { + event.preventDefault(); + event.stopPropagation(); + }, + [], + ); + + const handleRemoveImageClick = useCallback( + (event: React.MouseEvent, index: number) => { + event.preventDefault(); + event.stopPropagation(); + onRemoveImage?.(index); + }, + [onRemoveImage], + ); + return ( = ({ src={`data:${img.mediaType};base64,${img.data}`} alt={`预览 ${index + 1}`} /> - onRemoveImage?.(index)}> + handleRemoveImageClick(event, index)} + > diff --git a/src/components/agent/chat/components/Inputbar/components/InputbarVisionCapabilityNotice.test.tsx b/src/components/agent/chat/components/Inputbar/components/InputbarVisionCapabilityNotice.test.tsx new file mode 100644 index 000000000..6c2360fe7 --- /dev/null +++ b/src/components/agent/chat/components/Inputbar/components/InputbarVisionCapabilityNotice.test.tsx @@ -0,0 +1,120 @@ +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 { InputbarVisionCapabilityNotice } from "./InputbarVisionCapabilityNotice"; + +const mockUseConfiguredProviders = vi.fn(); +const mockUseProviderModels = vi.fn(); +const mockResolveVisionModel = vi.fn(); + +vi.mock("@/hooks/useConfiguredProviders", () => ({ + useConfiguredProviders: (options: unknown) => mockUseConfiguredProviders(options), +})); + +vi.mock("@/hooks/useProviderModels", () => ({ + useProviderModels: (...args: unknown[]) => mockUseProviderModels(...args), +})); + +vi.mock("@/lib/model/visionModelResolver", () => ({ + resolveVisionModel: (...args: unknown[]) => mockResolveVisionModel(...args), +})); + +const mountedRoots: Array<{ root: Root; container: HTMLDivElement }> = []; + +beforeEach(() => { + ( + globalThis as typeof globalThis & { + IS_REACT_ACT_ENVIRONMENT?: boolean; + } + ).IS_REACT_ACT_ENVIRONMENT = true; + + mockUseConfiguredProviders.mockReturnValue({ + providers: [ + { + key: "openai", + label: "OpenAI", + registryId: "openai", + type: "openai", + }, + ], + loading: false, + }); + mockUseProviderModels.mockReturnValue({ + models: [{ id: "gpt-4.1" }, { id: "gpt-4.1-vision" }], + loading: false, + error: null, + modelIds: ["gpt-4.1", "gpt-4.1-vision"], + }); + mockResolveVisionModel.mockReturnValue({ + reason: "already_vision", + targetModelId: "gpt-4.1", + }); +}); + +afterEach(() => { + while (mountedRoots.length > 0) { + const mounted = mountedRoots.pop(); + if (!mounted) break; + act(() => { + mounted.root.unmount(); + }); + mounted.container.remove(); + } + vi.clearAllMocks(); +}); + +function renderNotice( + props?: Partial>, +) { + const container = document.createElement("div"); + document.body.appendChild(container); + const root = createRoot(container); + + act(() => { + root.render( + , + ); + }); + + mountedRoots.push({ root, container }); + return container; +} + +describe("InputbarVisionCapabilityNotice", () => { + it("当前模型支持多模态时不应展示提示", () => { + const container = renderNotice(); + + expect( + container.querySelector('[data-testid="inputbar-vision-warning"]'), + ).toBeNull(); + }); + + it("当前模型不支持多模态时应展示推荐模型提示", () => { + mockResolveVisionModel.mockReturnValue({ + reason: "switched", + targetModelId: "gpt-4.1-vision", + }); + + const container = renderNotice(); + + expect(container.textContent).toContain("gpt-4.1 不支持多模态图片理解"); + expect(container.textContent).toContain("gpt-4.1-vision"); + }); + + it("当前 Provider 没有可用多模态模型时应展示 Provider 级提示", () => { + mockResolveVisionModel.mockReturnValue({ + reason: "no_vision_model", + targetModelId: "", + }); + + const container = renderNotice(); + + expect(container.textContent).toContain("当前 Provider 暂无可用的多模态模型"); + }); +}); diff --git a/src/components/agent/chat/components/Inputbar/components/InputbarVisionCapabilityNotice.tsx b/src/components/agent/chat/components/Inputbar/components/InputbarVisionCapabilityNotice.tsx new file mode 100644 index 000000000..9fca52b68 --- /dev/null +++ b/src/components/agent/chat/components/Inputbar/components/InputbarVisionCapabilityNotice.tsx @@ -0,0 +1,85 @@ +import React, { useMemo } from "react"; +import { AlertCircle } from "lucide-react"; +import { useConfiguredProviders } from "@/hooks/useConfiguredProviders"; +import { useProviderModels } from "@/hooks/useProviderModels"; +import { resolveVisionModel } from "@/lib/model/visionModelResolver"; + +interface InputbarVisionCapabilityNoticeProps { + providerType?: string; + model?: string; + hasPendingImages: boolean; +} + +export const InputbarVisionCapabilityNotice: React.FC< + InputbarVisionCapabilityNoticeProps +> = ({ providerType, model, hasPendingImages }) => { + const shouldInspectCapability = + hasPendingImages && + Boolean(providerType?.trim()) && + Boolean(model?.trim()); + + const { providers, loading: providersLoading } = useConfiguredProviders({ + autoLoad: shouldInspectCapability, + }); + + const selectedProvider = useMemo( + () => providers.find((item) => item.key === providerType), + [providerType, providers], + ); + + const { models, loading: modelsLoading } = useProviderModels( + selectedProvider, + { + returnFullMetadata: true, + autoLoad: shouldInspectCapability && Boolean(selectedProvider), + }, + ); + + const warningMessage = useMemo(() => { + if (!shouldInspectCapability || !model?.trim()) { + return null; + } + if (providersLoading || modelsLoading || !selectedProvider) { + return null; + } + + const visionResult = resolveVisionModel({ + currentModelId: model, + models, + }); + + if (visionResult.reason === "already_vision") { + return null; + } + + if (visionResult.reason === "no_vision_model") { + return "当前 Provider 暂无可用的多模态模型,请切换到支持多模态的 Provider 或模型后再发送图片"; + } + + const suggestedModel = visionResult.targetModelId.trim(); + return suggestedModel + ? `当前模型 ${model} 不支持多模态图片理解,建议切换到 ${suggestedModel} 后再发送图片` + : `当前模型 ${model} 不支持多模态图片理解,请切换到支持多模态的模型后再发送图片`; + }, [ + model, + models, + modelsLoading, + providersLoading, + selectedProvider, + shouldInspectCapability, + ]); + + if (!warningMessage) { + return null; + } + + return ( +
+ + {warningMessage} +
+ ); +}; diff --git a/src/components/agent/chat/components/Inputbar/components/SkillSelector.tsx b/src/components/agent/chat/components/Inputbar/components/SkillSelector.tsx index 968100b40..cf3891789 100644 --- a/src/components/agent/chat/components/Inputbar/components/SkillSelector.tsx +++ b/src/components/agent/chat/components/Inputbar/components/SkillSelector.tsx @@ -195,13 +195,13 @@ export const SkillSelector: React.FC = ({ - -
+ +
技能能力
diff --git a/src/components/agent/chat/components/Inputbar/hooks/useImageAttachments.test.tsx b/src/components/agent/chat/components/Inputbar/hooks/useImageAttachments.test.tsx new file mode 100644 index 000000000..d9b437c0b --- /dev/null +++ b/src/components/agent/chat/components/Inputbar/hooks/useImageAttachments.test.tsx @@ -0,0 +1,203 @@ +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 { useImageAttachments } from "./useImageAttachments"; + +const { toastMock } = vi.hoisted(() => ({ + toastMock: { + info: vi.fn(), + success: vi.fn(), + error: vi.fn(), + }, +})); + +vi.mock("sonner", () => ({ + toast: toastMock, +})); + +interface MountedHarness { + container: HTMLDivElement; + root: Root; +} + +const mountedRoots: MountedHarness[] = []; +const originalFileReader = globalThis.FileReader; + +class MockFileReader { + public onload: ((event: { target: { result: string } }) => void) | null = + null; + public onerror: (() => void) | null = null; + public error: Error | null = null; + + readAsDataURL(file: File) { + this.onload?.({ + target: { + result: `data:${file.type};base64,ZmFrZS1pbWFnZQ==`, + }, + }); + } +} + +function Harness() { + const { pendingImages, handlePaste, handleRemoveImage } = useImageAttachments(); + + return ( +
+ + + +
{pendingImages.length}
+
{pendingImages[0]?.mediaType || ""}
+
+ ); +} + +function renderHarness(): HTMLDivElement { + const container = document.createElement("div"); + document.body.appendChild(container); + const root = createRoot(container); + + act(() => { + root.render(); + }); + + mountedRoots.push({ container, root }); + return container; +} + +beforeEach(() => { + ( + globalThis as typeof globalThis & { + IS_REACT_ACT_ENVIRONMENT?: boolean; + FileReader: typeof FileReader; + } + ).IS_REACT_ACT_ENVIRONMENT = true; + globalThis.FileReader = MockFileReader as never; +}); + +afterEach(() => { + while (mountedRoots.length > 0) { + const mounted = mountedRoots.pop(); + if (!mounted) break; + act(() => { + mounted.root.unmount(); + }); + mounted.container.remove(); + } + globalThis.FileReader = originalFileReader; + vi.clearAllMocks(); +}); + +describe("useImageAttachments", () => { + it("应支持从 clipboardData.files 直接粘贴图片", async () => { + const container = renderHarness(); + const pasteButton = container.querySelector( + '[data-testid="paste-image"]', + ) as HTMLButtonElement | null; + + await act(async () => { + pasteButton?.click(); + await Promise.resolve(); + }); + + expect( + container.querySelector('[data-testid="image-count"]')?.textContent, + ).toBe("1"); + expect( + container.querySelector('[data-testid="image-type"]')?.textContent, + ).toBe("image/png"); + expect(toastMock.success).toHaveBeenCalledWith("已粘贴图片"); + }); + + it("删除图片后应从待发送列表移除", async () => { + const container = renderHarness(); + const pasteButton = container.querySelector( + '[data-testid="paste-image"]', + ) as HTMLButtonElement | null; + const removeButton = container.querySelector( + '[data-testid="remove-image"]', + ) as HTMLButtonElement | null; + + await act(async () => { + pasteButton?.click(); + await Promise.resolve(); + }); + + await act(async () => { + removeButton?.click(); + await Promise.resolve(); + }); + + expect( + container.querySelector('[data-testid="image-count"]')?.textContent, + ).toBe("0"); + }); + + it("应支持从 clipboardData.items 的 type 回退识别粘贴图片", async () => { + const container = renderHarness(); + const pasteButton = container.querySelector( + '[data-testid="paste-image-from-item-type"]', + ) as HTMLButtonElement | null; + + await act(async () => { + pasteButton?.click(); + await Promise.resolve(); + }); + + expect( + container.querySelector('[data-testid="image-count"]')?.textContent, + ).toBe("1"); + expect( + container.querySelector('[data-testid="image-type"]')?.textContent, + ).toBe("image/png"); + expect(toastMock.success).toHaveBeenCalledWith("已粘贴图片"); + }); +}); diff --git a/src/components/agent/chat/components/Inputbar/hooks/useImageAttachments.ts b/src/components/agent/chat/components/Inputbar/hooks/useImageAttachments.ts index 0b2453274..9ca45abae 100644 --- a/src/components/agent/chat/components/Inputbar/hooks/useImageAttachments.ts +++ b/src/components/agent/chat/components/Inputbar/hooks/useImageAttachments.ts @@ -8,53 +8,23 @@ import { } from "react"; import { toast } from "sonner"; import type { MessageImage } from "../../../types"; - -function readImageAsBase64(file: File): Promise { - return new Promise((resolve, reject) => { - const reader = new FileReader(); - - reader.onload = (event) => { - const result = event.target?.result; - if (typeof result !== "string") { - reject(new Error("invalid_result")); - return; - } - - const [, base64Data = ""] = result.split(","); - resolve(base64Data); - }; - - reader.onerror = () => { - reject(reader.error ?? new Error("read_failed")); - }; - - reader.readAsDataURL(file); - }); -} +import { + getClipboardImageCandidates, + readImageAttachment, +} from "../../../utils/imageAttachments"; export function useImageAttachments() { const [pendingImages, setPendingImages] = useState([]); const fileInputRef = useRef(null); const appendImageFile = useCallback( - async (file: File, successMessage?: string) => { - if (!file.type.startsWith("image/")) { - toast.info(`暂不支持该文件类型: ${file.type}`); - return; - } - + async (file: File, successMessage?: string, preferredMediaType?: string) => { try { - const base64Data = await readImageAsBase64(file); - setPendingImages((prev) => [ - ...prev, - { - data: base64Data, - mediaType: file.type, - }, - ]); - toast.success(successMessage ?? `已添加图片: ${file.name}`); + const image = await readImageAttachment(file, preferredMediaType); + setPendingImages((prev) => [...prev, image]); + toast.success(successMessage ?? `已添加图片: ${file.name || "未命名图片"}`); } catch { - toast.error(`图片读取失败: ${file.name}`); + toast.error(`图片读取失败: ${file.name || "未命名图片"}`); } }, [], @@ -84,23 +54,19 @@ export function useImageAttachments() { const handlePaste = useCallback( (event: ClipboardEvent) => { - const items = event.clipboardData?.items; - if (!items) { + const imageFiles = getClipboardImageCandidates(event.clipboardData); + if (imageFiles.length === 0) { return; } - for (const item of items) { - if (!item.type.startsWith("image/")) { - continue; - } - - event.preventDefault(); - const file = item.getAsFile(); - if (file) { - void appendImageFile(file, "已粘贴图片"); - } - break; - } + event.preventDefault(); + imageFiles.forEach(({ file, mediaType }, index) => { + void appendImageFile( + file, + index === 0 ? "已粘贴图片" : undefined, + mediaType, + ); + }); }, [appendImageFile], ); @@ -126,7 +92,9 @@ export function useImageAttachments() { ); const handleRemoveImage = useCallback((index: number) => { - setPendingImages((prev) => prev.filter((_, currentIndex) => currentIndex !== index)); + setPendingImages((prev) => + prev.filter((_, currentIndex) => currentIndex !== index), + ); }, []); const clearPendingImages = useCallback(() => { diff --git a/src/components/agent/chat/components/Inputbar/hooks/useInputbarController.ts b/src/components/agent/chat/components/Inputbar/hooks/useInputbarController.ts index 2afba39b6..bfdbae425 100644 --- a/src/components/agent/chat/components/Inputbar/hooks/useInputbarController.ts +++ b/src/components/agent/chat/components/Inputbar/hooks/useInputbarController.ts @@ -27,7 +27,7 @@ interface UseInputbarControllerParams { thinking?: boolean, textOverride?: string, executionStrategy?: "react" | "code_orchestrated" | "auto", - ) => void; + ) => void | Promise | boolean; onStop?: () => void; isLoading: boolean; disabled?: boolean; diff --git a/src/components/agent/chat/components/Inputbar/hooks/useInputbarSend.ts b/src/components/agent/chat/components/Inputbar/hooks/useInputbarSend.ts index a23466832..d7ebc651b 100644 --- a/src/components/agent/chat/components/Inputbar/hooks/useInputbarSend.ts +++ b/src/components/agent/chat/components/Inputbar/hooks/useInputbarSend.ts @@ -19,7 +19,7 @@ interface UseInputbarSendParams { thinking?: boolean, textOverride?: string, executionStrategy?: "react" | "code_orchestrated" | "auto", - ) => void; + ) => void | Promise | boolean; clearPendingImages: () => void; clearActiveSkill: () => void; } @@ -37,7 +37,7 @@ export function useInputbarSend({ clearPendingImages, clearActiveSkill, }: UseInputbarSendParams) { - return useCallback(() => { + return useCallback(async () => { if (!input.trim() && pendingImages.length === 0) { return; } @@ -63,15 +63,22 @@ export function useInputbarSend({ textOverride = `/${SOCIAL_ARTICLE_SKILL_KEY} ${input}`.trim(); } - onSend( - pendingImages.length > 0 ? pendingImages : undefined, - webSearch, - thinking, - textOverride, - strategy, - ); - clearPendingImages(); - clearActiveSkill(); + try { + const result = await onSend( + pendingImages.length > 0 ? pendingImages : undefined, + webSearch, + thinking, + textOverride, + strategy, + ); + if (result === false) { + return; + } + clearPendingImages(); + clearActiveSkill(); + } catch { + // 发送失败时保留图片与技能,交由上层 toast / 恢复逻辑处理。 + } }, [ activeSkill, activeTheme, diff --git a/src/components/agent/chat/components/Inputbar/index.test.tsx b/src/components/agent/chat/components/Inputbar/index.test.tsx index 9ba657983..00b84900e 100644 --- a/src/components/agent/chat/components/Inputbar/index.test.tsx +++ b/src/components/agent/chat/components/Inputbar/index.test.tsx @@ -276,8 +276,9 @@ describe("Inputbar", () => { ) as HTMLButtonElement | null; expect(sendButton).toBeTruthy(); - act(() => { + await act(async () => { sendButton?.click(); + await Promise.resolve(); }); expect(onSend).toHaveBeenCalledWith( @@ -307,8 +308,9 @@ describe("Inputbar", () => { ) as HTMLButtonElement | null; expect(sendButton).toBeTruthy(); - act(() => { + await act(async () => { sendButton?.click(); + await Promise.resolve(); }); expect(onSend).toHaveBeenCalledWith( diff --git a/src/components/agent/chat/components/Inputbar/index.tsx b/src/components/agent/chat/components/Inputbar/index.tsx index 90619cecd..e414ed90c 100644 --- a/src/components/agent/chat/components/Inputbar/index.tsx +++ b/src/components/agent/chat/components/Inputbar/index.tsx @@ -28,7 +28,7 @@ interface InputbarProps { thinking?: boolean, textOverride?: string, executionStrategy?: "react" | "code_orchestrated" | "auto", - ) => void; + ) => void | Promise | boolean; /** 停止生成回调 */ onStop?: () => void; isLoading: boolean; diff --git a/src/components/agent/chat/components/MarkdownRenderer.test.tsx b/src/components/agent/chat/components/MarkdownRenderer.test.tsx index 46793e08f..edb6e1056 100644 --- a/src/components/agent/chat/components/MarkdownRenderer.test.tsx +++ b/src/components/agent/chat/components/MarkdownRenderer.test.tsx @@ -38,6 +38,12 @@ interface MountedHarness { root: Root; } +interface RenderOptions { + isStreaming?: boolean; + collapseCodeBlocks?: boolean; + shouldCollapseCodeBlock?: (language: string, code: string) => boolean; +} + const mountedRoots: MountedHarness[] = []; beforeEach(() => { @@ -61,14 +67,26 @@ afterEach(() => { vi.clearAllMocks(); }); -function render(content: string, isStreaming = false): HTMLDivElement { +function render( + content: string, + { + isStreaming = false, + collapseCodeBlocks = false, + shouldCollapseCodeBlock, + }: RenderOptions = {}, +): HTMLDivElement { const container = document.createElement("div"); document.body.appendChild(container); const root = createRoot(container); act(() => { root.render( - , + , ); }); @@ -76,23 +94,44 @@ function render(content: string, isStreaming = false): HTMLDivElement { return container; } -function renderHarness(content: string, isStreaming = false) { +function renderHarness( + content: string, + { + isStreaming = false, + collapseCodeBlocks = false, + shouldCollapseCodeBlock, + }: RenderOptions = {}, +) { const container = document.createElement("div"); document.body.appendChild(container); const root = createRoot(container); - const rerender = (nextContent: string, nextIsStreaming = isStreaming) => { + const rerender = ( + nextContent: string, + { + isStreaming: nextIsStreaming = isStreaming, + collapseCodeBlocks: nextCollapseCodeBlocks = collapseCodeBlocks, + shouldCollapseCodeBlock: nextShouldCollapseCodeBlock = + shouldCollapseCodeBlock, + }: RenderOptions = {}, + ) => { act(() => { root.render( , ); }); }; - rerender(content, isStreaming); + rerender(content, { + isStreaming, + collapseCodeBlocks, + shouldCollapseCodeBlock, + }); mountedRoots.push({ container, root }); return { container, rerender }; @@ -108,7 +147,7 @@ describe("MarkdownRenderer", () => { "后置文本", ].join("\n"); - const container = render(content, false); + const container = render(content); expect(container.querySelector(".rendered-html")).not.toBeNull(); expect(container.textContent).toContain("原始 HTML"); @@ -123,7 +162,7 @@ describe("MarkdownRenderer", () => { "结尾文本", ].join("\n"); - const container = render(content, true); + const container = render(content, { isStreaming: true }); expect(container.querySelector(".rendered-html")).toBeNull(); expect(container.textContent).toContain("结尾文本"); @@ -139,10 +178,52 @@ describe("MarkdownRenderer", () => { "结尾文本", ].join("\n"); - const { container, rerender } = renderHarness(content, true); + const { container, rerender } = renderHarness(content, { + isStreaming: true, + }); expect(container.querySelector(".rendered-html")).toBeNull(); - rerender(content, false); + rerender(content, { isStreaming: false }); expect(container.querySelector(".rendered-html")).not.toBeNull(); }); + + it("逐块判定返回 false 时应保持对话内联代码渲染", () => { + const shouldCollapseCodeBlock = vi.fn(() => false); + const content = ["```ts", "const answer = 42;", "```"].join("\n"); + + const container = render(content, { + collapseCodeBlocks: true, + shouldCollapseCodeBlock, + }); + + expect(shouldCollapseCodeBlock).toHaveBeenCalledWith( + "ts", + "const answer = 42;", + ); + expect( + container.querySelector('[data-testid="artifact-placeholder"]'), + ).toBeNull(); + expect( + container.querySelector('[data-testid="syntax-highlighter"]'), + ).not.toBeNull(); + expect(container.textContent).toContain("const answer = 42;"); + }); + + it("逐块判定返回 true 时才应渲染 artifact 占位卡", () => { + const content = ["```tsx", "export default function Demo() {}", "```"].join( + "\n", + ); + + const container = render(content, { + collapseCodeBlocks: true, + shouldCollapseCodeBlock: () => true, + }); + + expect( + container.querySelector('[data-testid="artifact-placeholder"]'), + ).not.toBeNull(); + expect( + container.querySelector('[data-testid="syntax-highlighter"]'), + ).toBeNull(); + }); }); diff --git a/src/components/agent/chat/components/MarkdownRenderer.tsx b/src/components/agent/chat/components/MarkdownRenderer.tsx index 0c3a288b9..4faf85bd6 100644 --- a/src/components/agent/chat/components/MarkdownRenderer.tsx +++ b/src/components/agent/chat/components/MarkdownRenderer.tsx @@ -237,6 +237,8 @@ interface MarkdownRendererProps { renderA2UIInline?: boolean; /** 是否折叠代码块(当画布打开时) */ collapseCodeBlocks?: boolean; + /** 按代码块决定是否折叠 */ + shouldCollapseCodeBlock?: (language: string, code: string) => boolean; /** 代码块点击回调(用于在画布中显示) */ onCodeBlockClick?: (language: string, code: string) => void; /** 是否正在流式生成 */ @@ -249,6 +251,7 @@ export const MarkdownRenderer: React.FC = memo( onA2UISubmit, renderA2UIInline = true, collapseCodeBlocks = false, + shouldCollapseCodeBlock, onCodeBlockClick, isStreaming = false, }) => { @@ -436,7 +439,13 @@ export const MarkdownRenderer: React.FC = memo( } // 如果启用了代码块折叠,显示占位符卡片 - if (collapseCodeBlocks) { + const shouldRenderArtifactPlaceholder = + collapseCodeBlocks && + (shouldCollapseCodeBlock + ? shouldCollapseCodeBlock(language, codeContent) + : true); + + if (shouldRenderArtifactPlaceholder) { const lineCount = codeContent.split("\n").length; return ( void; /** 是否折叠代码块(当画布打开时) */ collapseCodeBlocks?: boolean; + /** 按代码块决定是否折叠 */ + shouldCollapseCodeBlock?: (language: string, code: string) => boolean; /** 代码块点击回调(用于在画布中显示) */ onCodeBlockClick?: (language: string, code: string) => void; /** 是否将待处理问答提升为输入区 A2UI 表单 */ @@ -98,6 +100,7 @@ const MessageListInner: React.FC = ({ onArtifactClick, onPermissionResponse, collapseCodeBlocks, + shouldCollapseCodeBlock, onCodeBlockClick, promoteActionRequestsToA2UI = false, }) => { @@ -259,7 +262,9 @@ const MessageListInner: React.FC = ({ {showIdentity ? ( - {msg.role === "user" ? "用户" : assistantLabel} + + {msg.role === "user" ? "用户" : assistantLabel} + {formatTime(msg.timestamp)} ) : ( @@ -321,8 +326,10 @@ const MessageListInner: React.FC = ({ onFileClick={onFileClick} onPermissionResponse={onPermissionResponse} collapseCodeBlocks={collapseCodeBlocks} + shouldCollapseCodeBlock={shouldCollapseCodeBlock} onCodeBlockClick={onCodeBlockClick} promoteActionRequestsToA2UI={promoteActionRequestsToA2UI} + renderProposedPlanBlocks={!timeline} /> ) : ( = ({ {filePath}
- + {statusLabel} {previewText ? ( diff --git a/src/components/agent/chat/components/SearchResultPreviewList.test.tsx b/src/components/agent/chat/components/SearchResultPreviewList.test.tsx index e67c0e4e1..89121a0fb 100644 --- a/src/components/agent/chat/components/SearchResultPreviewList.test.tsx +++ b/src/components/agent/chat/components/SearchResultPreviewList.test.tsx @@ -90,4 +90,33 @@ describe("SearchResultPreviewList", () => { expect(container.textContent).not.toContain("结果 6"); expect(container.textContent).toContain("展开其余 2 条结果"); }); + + it("悬浮预览在鼠标移出整体区域后应自动关闭", () => { + vi.useFakeTimers(); + const { container } = renderList(); + const trigger = container.querySelector( + 'button[aria-label="预览搜索结果:结果 1"]', + ) as HTMLButtonElement | null; + + act(() => { + trigger?.dispatchEvent( + new MouseEvent("mouseover", { + bubbles: true, + }), + ); + }); + + expect(document.body.textContent).toContain("摘要 1"); + + act(() => { + document.body.dispatchEvent( + new MouseEvent("mousemove", { + bubbles: true, + }), + ); + vi.advanceTimersByTime(140); + }); + + expect(document.body.textContent).not.toContain("摘要 1"); + }); }); diff --git a/src/components/agent/chat/components/SearchResultPreviewList.tsx b/src/components/agent/chat/components/SearchResultPreviewList.tsx index 7f4499ce6..fd4be8ae4 100644 --- a/src/components/agent/chat/components/SearchResultPreviewList.tsx +++ b/src/components/agent/chat/components/SearchResultPreviewList.tsx @@ -23,6 +23,8 @@ function SearchResultHoverCard({ }) { const [open, setOpen] = useState(false); const closeTimerRef = useRef(null); + const triggerRef = useRef(null); + const contentRef = useRef(null); const clearCloseTimer = useCallback(() => { if (closeTimerRef.current !== null && typeof window !== "undefined") { @@ -36,24 +38,74 @@ function SearchResultHoverCard({ setOpen(true); }, [clearCloseTimer]); - const handleScheduleClose = useCallback(() => { + const handleCloseNow = useCallback(() => { clearCloseTimer(); + setOpen(false); + }, [clearCloseTimer]); + + const handleScheduleClose = useCallback(() => { + if (closeTimerRef.current !== null) { + return; + } if (typeof window === "undefined") { - setOpen(false); + handleCloseNow(); return; } closeTimerRef.current = window.setTimeout(() => { setOpen(false); closeTimerRef.current = null; }, 120); - }, [clearCloseTimer]); + }, [handleCloseNow]); + + const isWithinHoverRegion = useCallback((target: EventTarget | null) => { + const node = target instanceof Node ? target : null; + if (!node) { + return false; + } + return Boolean( + triggerRef.current?.contains(node) || contentRef.current?.contains(node), + ); + }, []); useEffect(() => () => clearCloseTimer(), [clearCloseTimer]); + useEffect(() => { + if (!open || typeof document === "undefined") { + return; + } + + const handleDocumentMouseMove = (event: MouseEvent) => { + if (isWithinHoverRegion(event.target)) { + clearCloseTimer(); + return; + } + handleScheduleClose(); + }; + + const handleWindowBlur = () => { + handleCloseNow(); + }; + + document.addEventListener("mousemove", handleDocumentMouseMove, true); + window.addEventListener("blur", handleWindowBlur); + + return () => { + document.removeEventListener("mousemove", handleDocumentMouseMove, true); + window.removeEventListener("blur", handleWindowBlur); + }; + }, [ + clearCloseTimer, + handleCloseNow, + handleScheduleClose, + isWithinHoverRegion, + open, + ]); + return (
- {/* 能力标签 */} -
- {isDefault && } - {model.capabilities.vision && ( - - - - )} - {model.capabilities.tools && ( - - - - )} - {model.capabilities.reasoning && ( - - - - )} +
+ {isDefault ? : null}
); diff --git a/src/components/settings-v2/general/memory/index.test.tsx b/src/components/settings-v2/general/memory/index.test.tsx index 0d6cfcb7b..e3c2ab9e4 100644 --- a/src/components/settings-v2/general/memory/index.test.tsx +++ b/src/components/settings-v2/general/memory/index.test.tsx @@ -30,11 +30,11 @@ vi.mock("@/lib/api/appConfig", () => ({ })); vi.mock("@/lib/api/memoryRuntime", () => ({ - getMemoryOverview: mockGetMemoryOverview, - getMemoryEffectiveSources: mockGetMemoryEffectiveSources, - getMemoryAutoIndex: mockGetMemoryAutoIndex, - toggleMemoryAuto: mockToggleMemoryAuto, - updateMemoryAutoNote: mockUpdateMemoryAutoNote, + getContextMemoryOverview: mockGetMemoryOverview, + getContextMemoryEffectiveSources: mockGetMemoryEffectiveSources, + getContextMemoryAutoIndex: mockGetMemoryAutoIndex, + toggleContextMemoryAuto: mockToggleMemoryAuto, + updateContextMemoryAutoNote: mockUpdateMemoryAutoNote, })); vi.mock("@/lib/api/unifiedMemory", () => ({ @@ -152,7 +152,7 @@ beforeEach(() => { sources: { project_memory_paths: ["AGENTS.md"], project_rule_dirs: [".agents/rules"], - user_memory_path: "~/.lime/AGENTS.md", + user_memory_path: undefined, }, }, }); @@ -230,7 +230,7 @@ describe("MemorySettings", () => { expect(mockGetMemoryAutoIndex).toHaveBeenCalledTimes(1); }); - it("点击立即关闭应调用 toggleMemoryAuto", async () => { + it("点击立即关闭应调用 toggleContextMemoryAuto", async () => { const container = renderComponent(); await flushEffects(); await flushEffects(); diff --git a/src/components/settings-v2/general/memory/index.tsx b/src/components/settings-v2/general/memory/index.tsx index d28aebb7e..6df3dfaa2 100644 --- a/src/components/settings-v2/general/memory/index.tsx +++ b/src/components/settings-v2/general/memory/index.tsx @@ -20,10 +20,11 @@ import { import { Switch } from "@/components/ui/switch"; import { cn } from "@/lib/utils"; import { - getMemoryAutoIndex, - getMemoryEffectiveSources, - toggleMemoryAuto, - updateMemoryAutoNote, + getContextMemoryAutoIndex, + getContextMemoryEffectiveSources, + getContextMemoryOverview, + toggleContextMemoryAuto, + updateContextMemoryAutoNote, type AutoMemoryIndexResponse, type EffectiveMemorySourcesResponse, type MemoryAutoConfig, @@ -31,7 +32,6 @@ import { type MemoryProfileConfig, type MemoryResolveConfig, type MemorySourcesConfig, - getMemoryOverview as getContextMemoryOverview, } from "@/lib/api/memoryRuntime"; import { getConfig, saveConfig, type Config } from "@/lib/api/appConfig"; import { getUnifiedMemoryStats } from "@/lib/api/unifiedMemory"; @@ -101,7 +101,7 @@ function normalizeSources(sources?: MemorySourcesConfig): MemorySourcesConfig { sources.project_rule_dirs.filter((item) => item.trim().length > 0) ? sources.project_rule_dirs : [".agents/rules"], - user_memory_path: sources?.user_memory_path ?? "~/.lime/AGENTS.md", + user_memory_path: sources?.user_memory_path ?? undefined, project_local_memory_path: sources?.project_local_memory_path ?? "AGENTS.local.md", }; @@ -390,8 +390,8 @@ export function MemorySettings() { setLoadingSourceState(true); try { const [sources, index] = await Promise.all([ - getMemoryEffectiveSources().catch(() => null), - getMemoryAutoIndex().catch(() => null), + getContextMemoryEffectiveSources().catch(() => null), + getContextMemoryAutoIndex().catch(() => null), ]); setEffectiveSources(sources); setAutoIndex(index); @@ -499,7 +499,7 @@ export function MemorySettings() { const current = normalizeAuto(draft.auto).enabled ?? true; const next = !current; try { - const result = await toggleMemoryAuto(next); + const result = await toggleContextMemoryAuto(next); setDraft((prev) => ({ ...prev, auto: { @@ -534,7 +534,7 @@ export function MemorySettings() { setSavingAutoNote(true); try { - const index = await updateMemoryAutoNote( + const index = await updateContextMemoryAutoNote( note, autoTopic.trim() || undefined, ); @@ -937,7 +937,7 @@ export function MemorySettings() { })) } className={INPUT_CLASS_NAME} - placeholder="例如 ~/.lime/AGENTS.md" + placeholder="留空时使用应用默认 AGENTS.md 路径" /> diff --git a/src/components/terminal/ai/TerminalAIModeSelector.tsx b/src/components/terminal/ai/TerminalAIModeSelector.tsx index 7e5d329f7..5671818bc 100644 --- a/src/components/terminal/ai/TerminalAIModeSelector.tsx +++ b/src/components/terminal/ai/TerminalAIModeSelector.tsx @@ -293,12 +293,12 @@ export const TerminalAIModeSelector: React.FC = ({
{/* 左侧:Provider 列表 */} -
+
Providers
diff --git a/src/hooks/README.md b/src/hooks/README.md index 856b6e2e1..bcec8dbd6 100644 --- a/src/hooks/README.md +++ b/src/hooks/README.md @@ -6,58 +6,17 @@ | 文件 | 说明 | |------|------| -| `useUnifiedChat.ts` | 统一对话 Hook,支持 Agent/General/Creator 三种模式 | | `useSkillExecution.ts` | Skill 执行 Hook,监听 Tauri 事件并管理执行状态 | -## useUnifiedChat - -统一的对话逻辑 Hook,统一收口 Agent / General / Creator 三类对话入口。 - -### 使用示例 - -```typescript -import { useUnifiedChat } from "@/hooks/useUnifiedChat"; - -// Agent 模式 -const agentChat = useUnifiedChat({ - mode: "agent", - providerType: "claude", - model: "claude-sonnet-4-20250514", -}); - -// 内容创作模式 -const creatorChat = useUnifiedChat({ - mode: "creator", - systemPrompt: "你是一位专业的内容创作助手...", - onCanvasUpdate: (path, content) => { - // 更新画布内容 - }, -}); - -// 通用对话模式 -const generalChat = useUnifiedChat({ - mode: "general", -}); -``` - -### 返回值 - -- `session` - 当前会话 -- `messages` - 消息列表 -- `isLoading` - 加载状态 -- `isSending` - 发送状态 -- `error` - 错误信息 -- `createSession()` - 创建会话 -- `loadSession()` - 加载会话 -- `sendMessage()` - 发送消息 -- `stopGeneration()` - 停止生成 -- `configureProvider()` - 配置 Provider +历史 `useUnifiedChat.ts` compat Hook 已删除。 +新 Agent / Codex 工作台统一走 `src/components/agent/chat/hooks/index.ts` 暴露的 `useAgentChatUnified`; +底层实现为 `src/components/agent/chat/hooks/useAsterAgentChat.ts`; +如果未来要恢复 General / Creator 能力,也应基于 `agent_runtime_*` 重建。 ## 相关文档 - 架构设计:`docs/prd/chat-architecture-redesign.md` - 类型定义:`src/types/chat.ts` -- API 封装:`src/lib/api/unified-chat.ts` ## useSkillExecution diff --git a/src/hooks/useProviderModels.ts b/src/hooks/useProviderModels.ts index 3766fb85e..32c76cce0 100644 --- a/src/hooks/useProviderModels.ts +++ b/src/hooks/useProviderModels.ts @@ -13,6 +13,7 @@ import { getAliasConfigKey, isAliasProvider, } from "@/lib/constants/providerMappings"; +import { inferModelCapabilities } from "@/lib/model/inferModelCapabilities"; import type { ConfiguredProvider } from "./useConfiguredProviders"; import type { EnhancedModelMetadata, @@ -105,14 +106,10 @@ function convertCustomModelsToMetadata( provider_name: providerName, family: null, tier: "pro" as const, - capabilities: { - vision: false, - tools: true, - streaming: true, - json_mode: true, - function_calling: true, - reasoning: modelName.includes("thinking"), - }, + capabilities: inferModelCapabilities({ + modelId: modelName, + providerId, + }), pricing: null, limits: { context_length: null, @@ -149,14 +146,12 @@ function convertAliasModelsToMetadata( provider_name: providerName, family: aliasInfo?.provider || null, tier: "pro" as const, - capabilities: { - vision: false, - tools: true, - streaming: true, - json_mode: true, - function_calling: true, - reasoning: modelName.includes("thinking"), - }, + capabilities: inferModelCapabilities({ + modelId: modelName, + providerId: aliasInfo?.provider || providerId, + family: aliasInfo?.provider || null, + description: aliasInfo?.description || null, + }), pricing: null, limits: { context_length: null, diff --git a/src/hooks/useThreeStageWorkflow.ts b/src/hooks/useThreeStageWorkflow.ts deleted file mode 100644 index 9d75669f8..000000000 --- a/src/hooks/useThreeStageWorkflow.ts +++ /dev/null @@ -1,365 +0,0 @@ -/** - * 三阶段工作流 React Hook - * - * 提供在 React 组件中使用三阶段工作流的便捷接口 - */ - -import { useState, useCallback, useRef, useEffect } from "react"; -import ThreeStageWorkflowManager, { - type WorkflowConfig, - type ActionContext, -} from "../lib/workflow/threeStageWorkflow"; - -export interface UseThreeStageWorkflowOptions { - sessionId: string; - autoInitialize?: boolean; - defaultConfig?: Partial; -} - -export interface WorkflowState { - isInitialized: boolean; - isLoading: boolean; - error: string | null; - currentPhase: number; - visualOperationCount: number; - errorAttempts: Record; -} - -export interface WorkflowActions { - initializeWorkflow: (config: WorkflowConfig) => Promise; - preAction: (context: ActionContext) => Promise; - executeAction: (context: ActionContext, result: string) => Promise; - postAction: ( - context: ActionContext, - result: string, - error?: string, - ) => Promise; - updatePhaseStatus: ( - phaseNumber: number, - status: "pending" | "in_progress" | "complete", - notes?: string, - ) => Promise; - recordFinding: ( - title: string, - content: string, - tags?: string[], - ) => Promise; - recordDecision: (decision: string, rationale: string) => Promise; - checkCompletion: () => Promise<{ isComplete: boolean; summary: string }>; - finalizeWorkflow: () => Promise; - getSessionStats: () => Promise; - reset: () => void; -} - -/** - * 三阶段工作流 Hook - */ -export function useThreeStageWorkflow( - options: UseThreeStageWorkflowOptions, -): [WorkflowState, WorkflowActions] { - const { sessionId, autoInitialize = false, defaultConfig } = options; - - const [state, setState] = useState({ - isInitialized: false, - isLoading: false, - error: null, - currentPhase: 1, - visualOperationCount: 0, - errorAttempts: {}, - }); - - const workflowManagerRef = useRef(null); - - // 初始化工作流管理器 - useEffect(() => { - if (sessionId && !workflowManagerRef.current) { - workflowManagerRef.current = new ThreeStageWorkflowManager(sessionId); - } - }, [sessionId]); - - const updateStats = useCallback(async () => { - if (!workflowManagerRef.current) return; - - try { - const stats = await workflowManagerRef.current.getSessionStats(); - setState((prev) => ({ - ...prev, - visualOperationCount: stats.visualOperationCount, - errorAttempts: stats.errorAttempts, - })); - } catch (error) { - console.warn("更新统计信息失败:", error); - } - }, []); - - const setLoading = useCallback((loading: boolean) => { - setState((prev) => ({ ...prev, isLoading: loading })); - }, []); - - const setError = useCallback((error: string | null) => { - setState((prev) => ({ ...prev, error })); - }, []); - - const initializeWorkflow = useCallback( - async (config: WorkflowConfig) => { - if (!workflowManagerRef.current) return; - - setLoading(true); - setError(null); - - try { - await workflowManagerRef.current.initializeWorkflow(config); - setState((prev) => ({ - ...prev, - isInitialized: true, - currentPhase: 1, - })); - await updateStats(); - } catch (error) { - setError(error instanceof Error ? error.message : "初始化工作流失败"); - } finally { - setLoading(false); - } - }, - [setLoading, setError, updateStats], - ); - - // 自动初始化 - useEffect(() => { - if ( - autoInitialize && - defaultConfig && - !state.isInitialized && - workflowManagerRef.current - ) { - const config: WorkflowConfig = { - sessionId, - projectName: defaultConfig.projectName || "新项目", - goal: defaultConfig.goal || "待定义目标", - phases: defaultConfig.phases || [ - { - number: 1, - name: "需求分析", - status: "in_progress", - tasks: ["理解用户需求", "识别约束条件", "记录发现"], - }, - { - number: 2, - name: "方案设计", - status: "pending", - tasks: ["制定技术方案", "创建项目结构", "记录关键决策"], - }, - { - number: 3, - name: "实施执行", - status: "pending", - tasks: ["按计划执行", "增量测试", "记录进展"], - }, - ], - }; - - initializeWorkflow(config); - } - }, [ - autoInitialize, - defaultConfig, - state.isInitialized, - sessionId, - initializeWorkflow, - ]); - - const preAction = useCallback( - async (context: ActionContext): Promise => { - if (!workflowManagerRef.current) throw new Error("工作流未初始化"); - - setLoading(true); - setError(null); - - try { - const result = await workflowManagerRef.current.preAction(context); - await updateStats(); - return result; - } catch (error) { - const errorMessage = - error instanceof Error ? error.message : "Pre-Action 执行失败"; - setError(errorMessage); - throw error; - } finally { - setLoading(false); - } - }, - [setLoading, setError, updateStats], - ); - - const executeAction = useCallback( - async (context: ActionContext, result: string) => { - if (!workflowManagerRef.current) throw new Error("工作流未初始化"); - - try { - await workflowManagerRef.current.executeAction(context, result); - await updateStats(); - } catch (error) { - setError(error instanceof Error ? error.message : "Action 执行失败"); - throw error; - } - }, - [setError, updateStats], - ); - - const postAction = useCallback( - async ( - context: ActionContext, - result: string, - error?: string, - ): Promise => { - if (!workflowManagerRef.current) throw new Error("工作流未初始化"); - - setLoading(true); - - try { - const message = await workflowManagerRef.current.postAction( - context, - result, - error, - ); - await updateStats(); - return message; - } catch (err) { - const errorMessage = - err instanceof Error ? err.message : "Post-Action 执行失败"; - setError(errorMessage); - throw err; - } finally { - setLoading(false); - } - }, - [setLoading, setError, updateStats], - ); - - const updatePhaseStatus = useCallback( - async ( - phaseNumber: number, - status: "pending" | "in_progress" | "complete", - notes?: string, - ) => { - if (!workflowManagerRef.current) throw new Error("工作流未初始化"); - - try { - await workflowManagerRef.current.updatePhaseStatus( - phaseNumber, - status, - notes, - ); - setState((prev) => ({ ...prev, currentPhase: phaseNumber })); - await updateStats(); - } catch (error) { - setError(error instanceof Error ? error.message : "更新阶段状态失败"); - throw error; - } - }, - [setError, updateStats], - ); - - const recordFinding = useCallback( - async (title: string, content: string, tags: string[] = []) => { - if (!workflowManagerRef.current) throw new Error("工作流未初始化"); - - try { - await workflowManagerRef.current.recordFinding(title, content, tags); - await updateStats(); - } catch (error) { - setError(error instanceof Error ? error.message : "记录发现失败"); - throw error; - } - }, - [setError, updateStats], - ); - - const recordDecision = useCallback( - async (decision: string, rationale: string) => { - if (!workflowManagerRef.current) throw new Error("工作流未初始化"); - - try { - await workflowManagerRef.current.recordDecision(decision, rationale); - await updateStats(); - } catch (error) { - setError(error instanceof Error ? error.message : "记录决策失败"); - throw error; - } - }, - [setError, updateStats], - ); - - const checkCompletion = useCallback(async () => { - if (!workflowManagerRef.current) throw new Error("工作流未初始化"); - - try { - const result = await workflowManagerRef.current.checkCompletion(); - await updateStats(); - return result; - } catch (error) { - setError(error instanceof Error ? error.message : "检查完成状态失败"); - throw error; - } - }, [setError, updateStats]); - - const finalizeWorkflow = useCallback(async (): Promise => { - if (!workflowManagerRef.current) throw new Error("工作流未初始化"); - - setLoading(true); - - try { - const result = await workflowManagerRef.current.finalizeWorkflow(); - setState((prev) => ({ ...prev, isInitialized: false })); - return result; - } catch (error) { - setError(error instanceof Error ? error.message : "结束工作流失败"); - throw error; - } finally { - setLoading(false); - } - }, [setLoading, setError]); - - const getSessionStats = useCallback(async () => { - if (!workflowManagerRef.current) throw new Error("工作流未初始化"); - - try { - return await workflowManagerRef.current.getSessionStats(); - } catch (error) { - setError(error instanceof Error ? error.message : "获取会话统计失败"); - throw error; - } - }, [setError]); - - const reset = useCallback(() => { - setState({ - isInitialized: false, - isLoading: false, - error: null, - currentPhase: 1, - visualOperationCount: 0, - errorAttempts: {}, - }); - workflowManagerRef.current = sessionId - ? new ThreeStageWorkflowManager(sessionId) - : null; - }, [sessionId]); - - const actions: WorkflowActions = { - initializeWorkflow, - preAction, - executeAction, - postAction, - updatePhaseStatus, - recordFinding, - recordDecision, - checkCompletion, - finalizeWorkflow, - getSessionStats, - reset, - }; - - return [state, actions]; -} - -export default useThreeStageWorkflow; diff --git a/src/hooks/useUnifiedChat.ts b/src/hooks/useUnifiedChat.ts deleted file mode 100644 index 1956c7f23..000000000 --- a/src/hooks/useUnifiedChat.ts +++ /dev/null @@ -1,720 +0,0 @@ -/** - * @file useUnifiedChat.ts - * @description 统一对话 Hook - * @module hooks/useUnifiedChat - * - * 提供统一的对话逻辑,支持多种对话模式: - * - Agent: AI Agent 模式,支持工具调用 - * - General: 通用对话模式,纯文本 - * - Creator: 内容创作模式,支持画布输出 - * - * ## 设计原则 - * - 单一入口:所有对话场景使用同一个 Hook - * - 模式化设计:通过 mode 参数区分不同场景 - * - 统一 API:调用后端统一的 unified_chat_cmd - */ - -import { useState, useEffect, useRef, useCallback } from "react"; -import { toast } from "sonner"; -import { safeListen } from "@/lib/dev-bridge"; -import type { UnlistenFn } from "@tauri-apps/api/event"; -import * as chatApi from "@/lib/api/unified-chat"; -import type { - ChatSession, - ChatMessage, - ChatError, - ImageInput, - StreamEvent, - UseUnifiedChatOptions, - UseUnifiedChatReturn, - ToolCall, - CreateSessionRequest, - HarnessArtifactSnapshot, -} from "@/types/chat"; - -// ============================================================================ -// 常量 -// ============================================================================ - -const STORAGE_PREFIX = "unified_chat_"; - -// ============================================================================ -// 辅助函数 -// ============================================================================ - -/** 从 localStorage 加载数据 */ -function loadFromStorage(key: string, defaultValue: T): T { - try { - const stored = localStorage.getItem(`${STORAGE_PREFIX}${key}`); - return stored ? JSON.parse(stored) : defaultValue; - } catch { - return defaultValue; - } -} - -/** 保存数据到 localStorage */ -function saveToStorage(key: string, value: unknown): void { - try { - localStorage.setItem(`${STORAGE_PREFIX}${key}`, JSON.stringify(value)); - } catch (e) { - console.error("[useUnifiedChat] 保存到 localStorage 失败:", e); - } -} - -/** 解析 API 错误 */ -function parseApiError(error: unknown): ChatError { - const message = error instanceof Error ? error.message : String(error); - - // 根据错误消息判断类型 - if (message.includes("network") || message.includes("连接")) { - return { type: "network", message, retryable: true }; - } - if ( - message.includes("auth") || - message.includes("认证") || - message.includes("401") - ) { - return { type: "auth", message, retryable: false }; - } - if ( - message.includes("rate") || - message.includes("限流") || - message.includes("429") - ) { - return { type: "rate_limit", message, retryable: true }; - } - if (message.includes("quota") || message.includes("配额")) { - return { type: "quota", message, retryable: false }; - } - - return { type: "unknown", message, retryable: true }; -} - -// ============================================================================ -// Hook 实现 -// ============================================================================ - -/** - * 统一对话 Hook - */ -export function useUnifiedChat( - options: UseUnifiedChatOptions, -): UseUnifiedChatReturn { - const { - mode, - sessionId: initialSessionId, - systemPrompt, - providerType: initialProviderType, - model: initialModel, - onCanvasUpdate, - onWriteFile, - harnessConfig, - onHarnessEvent, - onArtifactUpdate, - onError, - } = options; - - // ========== 状态 ========== - const [session, setSession] = useState(null); - const [messages, setMessages] = useState([]); - const [isLoading, setIsLoading] = useState(false); - const [isSending, setIsSending] = useState(false); - const [error, setError] = useState(null); - - // Provider 配置 - const [providerType, setProviderType] = useState( - () => initialProviderType || loadFromStorage(`${mode}_provider`, "claude"), - ); - const [model, setModel] = useState( - () => initialModel || loadFromStorage(`${mode}_model`, ""), - ); - - // Refs - const unlistenRef = useRef(null); - const currentMsgIdRef = useRef(null); - const accumulatedContentRef = useRef(""); - const completedWriteFilesRef = useRef>(new Map()); - - // ========== 会话操作 ========== - - /** 创建新会话 */ - const createSession = useCallback( - async (opts?: Partial): Promise => { - try { - setIsLoading(true); - setError(null); - - const response = await chatApi.createSession({ - mode, - title: opts?.title, - systemPrompt: opts?.systemPrompt || systemPrompt, - providerType: opts?.providerType || providerType, - model: opts?.model || model, - metadata: - harnessConfig || opts?.metadata - ? { - ...(opts?.metadata || {}), - ...(harnessConfig ? { harness: harnessConfig } : {}), - } - : undefined, - }); - - const newSession: ChatSession = { - id: response.id, - mode: response.mode, - title: response.title, - model: response.model, - createdAt: response.createdAt, - updatedAt: response.updatedAt, - messageCount: response.messageCount, - }; - - setSession(newSession); - setMessages([]); - completedWriteFilesRef.current.clear(); - saveToStorage(`${mode}_session_id`, response.id); - - console.log("[useUnifiedChat] 创建会话成功:", response.id); - return response.id; - } catch (e) { - const chatError = parseApiError(e); - setError(chatError); - onError?.(chatError); - throw e; - } finally { - setIsLoading(false); - } - }, - [mode, systemPrompt, providerType, model, harnessConfig, onError], - ); - - /** 加载会话 */ - const loadSession = useCallback( - async (sessionId: string): Promise => { - try { - setIsLoading(true); - setError(null); - - // 获取会话详情 - const response = await chatApi.getSession(sessionId); - const loadedSession: ChatSession = { - id: response.id, - mode: response.mode, - title: response.title, - model: response.model, - createdAt: response.createdAt, - updatedAt: response.updatedAt, - messageCount: response.messageCount, - }; - setSession(loadedSession); - - // 获取消息列表 - const loadedMessages = await chatApi.getMessages(sessionId); - setMessages(loadedMessages); - - saveToStorage(`${mode}_session_id`, sessionId); - console.log( - "[useUnifiedChat] 加载会话成功:", - sessionId, - "消息数:", - loadedMessages.length, - ); - } catch (e) { - const chatError = parseApiError(e); - setError(chatError); - onError?.(chatError); - throw e; - } finally { - setIsLoading(false); - } - }, - [mode, onError], - ); - - /** 删除会话 */ - const deleteSession = useCallback( - async (sessionId?: string): Promise => { - const targetId = sessionId || session?.id; - if (!targetId) return; - - try { - await chatApi.deleteSession(targetId); - - if (targetId === session?.id) { - setSession(null); - setMessages([]); - completedWriteFilesRef.current.clear(); - localStorage.removeItem(`${STORAGE_PREFIX}${mode}_session_id`); - } - - toast.success("会话已删除"); - } catch (e) { - const chatError = parseApiError(e); - setError(chatError); - onError?.(chatError); - toast.error("删除会话失败"); - } - }, - [session?.id, mode, onError], - ); - - /** 重命名会话 */ - const renameSession = useCallback( - async (title: string, sessionId?: string): Promise => { - const targetId = sessionId || session?.id; - if (!targetId) return; - - try { - await chatApi.renameSession(targetId, title); - - if (targetId === session?.id) { - setSession((prev) => (prev ? { ...prev, title } : null)); - } - } catch (e) { - const chatError = parseApiError(e); - setError(chatError); - onError?.(chatError); - toast.error("重命名失败"); - } - }, - [session?.id, onError], - ); - - // ========== 消息操作 ========== - - /** 发送消息 */ - const sendMessage = useCallback( - async ( - content: string, - images?: ImageInput[], - webSearch?: boolean, - ): Promise => { - if (!content.trim() && (!images || images.length === 0)) return; - - let activeSessionId = session?.id; - - // 如果没有会话,先创建一个 - if (!activeSessionId) { - try { - activeSessionId = await createSession(); - } catch { - return; - } - } - - // 创建用户消息 - const userMsgId = `user-${Date.now()}`; - const userMessage: ChatMessage = { - id: userMsgId, - sessionId: activeSessionId, - role: "user", - content: content.trim(), - contentBlocks: [{ type: "text", text: content.trim() }], - status: "complete", - createdAt: new Date().toISOString(), - }; - - // 创建助手消息占位符 - const assistantMsgId = `assistant-${Date.now()}`; - const assistantMessage: ChatMessage = { - id: assistantMsgId, - sessionId: activeSessionId, - role: "assistant", - content: "", - contentBlocks: [], - status: "streaming", - createdAt: new Date().toISOString(), - }; - - setMessages((prev) => [...prev, userMessage, assistantMessage]); - setIsSending(true); - currentMsgIdRef.current = assistantMsgId; - accumulatedContentRef.current = ""; - - // 设置事件监听 - const eventName = chatApi.generateEventName(activeSessionId); - - try { - const unlisten = await safeListen(eventName, (event) => { - const data = chatApi.parseStreamEvent(event.payload); - if (!data) return; - - handleStreamEvent(data, assistantMsgId); - }); - - unlistenRef.current = unlisten; - - // 发送消息 - await chatApi.sendMessage({ - sessionId: activeSessionId, - message: content.trim(), - eventName, - webSearch, - images: images?.map((img) => ({ - data: img.data, - media_type: img.mediaType, - })), - }); - } catch (e) { - console.error("[useUnifiedChat] 发送消息失败:", e); - const chatError = parseApiError(e); - - setMessages((prev) => - prev.map((msg) => - msg.id === assistantMsgId - ? { - ...msg, - status: "error", - error: chatError, - content: chatError.message, - } - : msg, - ), - ); - - setIsSending(false); - onError?.(chatError); - - if (unlistenRef.current) { - unlistenRef.current(); - unlistenRef.current = null; - } - } - }, - // eslint-disable-next-line react-hooks/exhaustive-deps - [session?.id, createSession, onError], - ); - - /** 处理流式事件 */ - const handleStreamEvent = useCallback( - (event: StreamEvent, msgId: string): void => { - switch (event.type) { - case "text_delta": - accumulatedContentRef.current += event.text; - setMessages((prev) => - prev.map((msg) => - msg.id === msgId - ? { ...msg, content: accumulatedContentRef.current } - : msg, - ), - ); - - // 检查是否有 write_file 标签(用于画布) - checkWriteFileTag(accumulatedContentRef.current); - break; - - case "thinking_delta": - // 处理思考内容(可选显示) - break; - - case "tool_start": { - const newToolCall: ToolCall = { - id: event.tool_id, - name: event.tool_name, - arguments: event.arguments, - status: "running", - startTime: new Date(), - }; - - setMessages((prev) => - prev.map((msg) => - msg.id === msgId - ? { ...msg, toolCalls: [...(msg.toolCalls || []), newToolCall] } - : msg, - ), - ); - - // 检查是否是文件写入工具 - const toolName = event.tool_name.toLowerCase(); - if (toolName.includes("write") || toolName.includes("create")) { - try { - const args = JSON.parse(event.arguments || "{}"); - const filePath = args.path || args.file_path || args.filePath; - const content = args.content || args.text || ""; - if (filePath && content && onWriteFile) { - onWriteFile(content, filePath); - } - } catch { - // 忽略解析错误 - } - } - break; - } - - case "tool_end": - setMessages((prev) => - prev.map((msg) => { - if (msg.id !== msgId) return msg; - return { - ...msg, - toolCalls: msg.toolCalls?.map((tc) => - tc.id === event.tool_id - ? { - ...tc, - status: event.result.success ? "completed" : "failed", - result: event.result, - endTime: new Date(), - } - : tc, - ), - }; - }), - ); - break; - - case "harness_event": - onHarnessEvent?.(event.event); - break; - - case "artifact_snapshot": - onArtifactUpdate?.(event.artifact); - if (event.artifact.filePath && typeof event.artifact.content === "string") { - onCanvasUpdate?.(event.artifact.filePath, event.artifact.content); - const previousContent = completedWriteFilesRef.current.get( - event.artifact.filePath, - ); - if ( - previousContent !== event.artifact.content && - onWriteFile - ) { - completedWriteFilesRef.current.set( - event.artifact.filePath, - event.artifact.content, - ); - onWriteFile(event.artifact.content, event.artifact.filePath); - } - } - break; - - case "done": - // 单次 API 响应完成,但工具循环可能继续 - break; - - case "final_done": - setMessages((prev) => - prev.map((msg) => - msg.id === msgId - ? { - ...msg, - status: "complete", - content: accumulatedContentRef.current || "(无响应)", - metadata: event.usage - ? { - tokens: { - input: event.usage.input_tokens, - output: event.usage.output_tokens, - }, - } - : undefined, - } - : msg, - ), - ); - setIsSending(false); - cleanup(); - break; - - case "error": { - const chatError: ChatError = { - type: "unknown", - message: event.message, - retryable: true, - }; - - setMessages((prev) => - prev.map((msg) => - msg.id === msgId - ? { - ...msg, - status: "error", - error: chatError, - content: accumulatedContentRef.current || event.message, - } - : msg, - ), - ); - setIsSending(false); - onError?.(chatError); - cleanup(); - break; - } - } - }, - // eslint-disable-next-line react-hooks/exhaustive-deps - [onArtifactUpdate, onCanvasUpdate, onHarnessEvent, onWriteFile, onError], - ); - - /** 检查 write_file 标签 */ - const checkWriteFileTag = useCallback( - (content: string): void => { - // 匹配 content - const regex = /([\s\S]*?)(<\/write_file>)?/g; - let match; - - while ((match = regex.exec(content)) !== null) { - const [, path, fileContent, closeTag] = match; - const isComplete = !!closeTag; - - if (onCanvasUpdate) { - onCanvasUpdate(path, fileContent); - } - - if (onArtifactUpdate) { - const artifact: HarnessArtifactSnapshot = { - artifactId: `write-file:${path}`, - filePath: path, - content: fileContent, - metadata: { - source: "write_file", - complete: isComplete, - }, - }; - onArtifactUpdate(artifact); - } - - if (isComplete && onWriteFile) { - const previousContent = completedWriteFilesRef.current.get(path); - if (previousContent === fileContent) { - continue; - } - completedWriteFilesRef.current.set(path, fileContent); - onWriteFile(fileContent, path); - onHarnessEvent?.({ - kind: previousContent === undefined ? "artifact_created" : "artifact_updated", - summary: `产物已写入 ${path}`, - artifact: { - artifactId: `write-file:${path}`, - filePath: path, - content: fileContent, - metadata: { - source: "write_file", - }, - }, - }); - } - } - }, - [onArtifactUpdate, onCanvasUpdate, onHarnessEvent, onWriteFile], - ); - - /** 清理资源 */ - const cleanup = useCallback((): void => { - if (unlistenRef.current) { - unlistenRef.current(); - unlistenRef.current = null; - } - currentMsgIdRef.current = null; - accumulatedContentRef.current = ""; - }, []); - - /** 停止生成 */ - const stopGeneration = useCallback(async (): Promise => { - if (!session?.id) return; - - try { - await chatApi.stopGeneration(session.id); - - // 更新当前消息状态 - if (currentMsgIdRef.current) { - setMessages((prev) => - prev.map((msg) => - msg.id === currentMsgIdRef.current - ? { - ...msg, - status: "complete", - content: accumulatedContentRef.current || "(已停止)", - } - : msg, - ), - ); - } - - setIsSending(false); - cleanup(); - } catch (e) { - console.error("[useUnifiedChat] 停止生成失败:", e); - } - }, [session?.id, cleanup]); - - /** 清空消息 */ - const clearMessages = useCallback((): void => { - setMessages([]); - setSession(null); - completedWriteFilesRef.current.clear(); - localStorage.removeItem(`${STORAGE_PREFIX}${mode}_session_id`); - toast.success("新对话已创建"); - }, [mode]); - - // ========== Provider 配置 ========== - - /** 配置 Provider */ - const configureProvider = useCallback( - async (newProviderType: string, newModel: string): Promise => { - setProviderType(newProviderType); - setModel(newModel); - saveToStorage(`${mode}_provider`, newProviderType); - saveToStorage(`${mode}_model`, newModel); - - // 如果有活跃会话,更新其 Provider 配置 - if (session?.id) { - try { - await chatApi.configureProvider( - session.id, - newProviderType, - newModel, - ); - } catch (e) { - console.error("[useUnifiedChat] 配置 Provider 失败:", e); - } - } - }, - [mode, session?.id], - ); - - // ========== 初始化 ========== - - useEffect(() => { - // 尝试恢复上次的会话 - const savedSessionId = - initialSessionId || loadFromStorage(`${mode}_session_id`, null); - if (savedSessionId) { - loadSession(savedSessionId).catch(() => { - // 如果加载失败,清除保存的 ID - localStorage.removeItem(`${STORAGE_PREFIX}${mode}_session_id`); - }); - } - - // 清理函数 - return () => { - cleanup(); - }; - // eslint-disable-next-line react-hooks/exhaustive-deps - }, [mode, initialSessionId]); - - // ========== 返回值 ========== - - return { - // 状态 - session, - messages, - isLoading, - isSending, - error, - - // 会话操作 - createSession, - loadSession, - deleteSession, - renameSession, - - // 消息操作 - sendMessage, - stopGeneration, - clearMessages, - - // Provider 配置 - configureProvider, - }; -} - -export default useUnifiedChat; diff --git a/src/lib/api/agentRuntime.ts b/src/lib/api/agentRuntime.ts index 00dce5c40..a3be987fd 100644 --- a/src/lib/api/agentRuntime.ts +++ b/src/lib/api/agentRuntime.ts @@ -94,6 +94,12 @@ export interface AsterSessionInfo { working_dir?: string; } +export interface AsterTodoItem { + content: string; + status: "pending" | "in_progress" | "completed"; + active_form?: string; +} + /** * TauriMessageContent(匹配后端 TauriMessageContent 枚举) */ @@ -135,6 +141,7 @@ export interface AsterSessionDetail { turns?: AgentThreadTurn[]; items?: AgentThreadItem[]; queued_turns?: QueuedTurnSnapshot[]; + todo_items?: AsterTodoItem[]; } export interface AgentTurnConfigSnapshot { diff --git a/src/lib/api/contextMemory.test.ts b/src/lib/api/contextMemory.test.ts deleted file mode 100644 index 87d762e76..000000000 --- a/src/lib/api/contextMemory.test.ts +++ /dev/null @@ -1,105 +0,0 @@ -import { beforeEach, describe, expect, it, vi } from "vitest"; -import { safeInvoke } from "@/lib/dev-bridge"; -import { ContextMemoryAPI } from "./contextMemory"; - -vi.mock("@/lib/dev-bridge", () => ({ - safeInvoke: vi.fn(), -})); - -describe("contextMemory API", () => { - beforeEach(() => { - vi.clearAllMocks(); - }); - - it("应代理基础记忆命令", async () => { - vi.mocked(safeInvoke) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce([{ id: "m1", title: "任务" }]) - .mockResolvedValueOnce("上下文") - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce(true) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce({ session_id: "session-1", active_memories: 2 }) - .mockResolvedValueOnce(undefined); - - await expect( - ContextMemoryAPI.saveMemoryEntry({ - session_id: "session-1", - file_type: "task_plan", - title: "计划", - content: "内容", - tags: ["任务"], - priority: 3, - }), - ).resolves.toBeUndefined(); - await expect( - ContextMemoryAPI.getSessionMemories("session-1"), - ).resolves.toEqual([expect.objectContaining({ id: "m1" })]); - await expect(ContextMemoryAPI.getMemoryContext("session-1")).resolves.toBe( - "上下文", - ); - await expect( - ContextMemoryAPI.recordError({ - session_id: "session-1", - error_description: "错误", - attempted_solution: "方案", - }), - ).resolves.toBeUndefined(); - await expect( - ContextMemoryAPI.shouldAvoidOperation("session-1", "重复操作"), - ).resolves.toBe(true); - await expect( - ContextMemoryAPI.markErrorResolved({ - session_id: "session-1", - error_description: "错误", - resolution: "已修复", - }), - ).resolves.toBeUndefined(); - await expect(ContextMemoryAPI.getMemoryStats("session-1")).resolves.toEqual( - expect.objectContaining({ active_memories: 2 }), - ); - await expect( - ContextMemoryAPI.cleanupExpiredMemories(), - ).resolves.toBeUndefined(); - }); - - it("应通过辅助方法复用基础命令", async () => { - vi.mocked(safeInvoke) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce(false); - - await ContextMemoryAPI.saveTaskPlan("session-2", "计划", "内容"); - await ContextMemoryAPI.saveFinding("session-2", "发现", "内容", ["关键"]); - await ContextMemoryAPI.logProgress("session-2", "进度", "完成一半"); - await ContextMemoryAPI.apply2ActionRule("session-2", "发现了线索"); - await expect( - ContextMemoryAPI.recordErrorWithCheck( - "session-2", - "失败", - "重试", - "再次点击", - ), - ).resolves.toEqual({ shouldAvoid: false }); - - expect(safeInvoke).toHaveBeenCalledWith("save_memory_entry", { - request: expect.objectContaining({ file_type: "task_plan" }), - }); - expect(safeInvoke).toHaveBeenCalledWith("save_memory_entry", { - request: expect.objectContaining({ file_type: "findings" }), - }); - expect(safeInvoke).toHaveBeenCalledWith("save_memory_entry", { - request: expect.objectContaining({ file_type: "progress" }), - }); - expect(safeInvoke).toHaveBeenCalledWith("record_error", { - request: expect.objectContaining({ session_id: "session-2" }), - }); - expect(safeInvoke).toHaveBeenCalledWith("should_avoid_operation", { - sessionId: "session-2", - operationDescription: "再次点击", - }); - }); -}); diff --git a/src/lib/api/contextMemory.ts b/src/lib/api/contextMemory.ts deleted file mode 100644 index 64618b6ee..000000000 --- a/src/lib/api/contextMemory.ts +++ /dev/null @@ -1,230 +0,0 @@ -/** - * 上下文记忆管理 API - * - * 基于文件系统的持久化记忆系统,解决 AI Agent 的上下文丢失、目标漂移、错误重复问题 - */ - -import { safeInvoke } from "@/lib/dev-bridge"; - -export interface MemoryEntry { - id: string; - session_id: string; - file_type: MemoryFileType; - title: string; - content: string; - tags: string[]; - priority: number; - created_at: number; - updated_at: number; - archived: boolean; -} - -export type MemoryFileType = - | "task_plan" - | "findings" - | "progress" - | "error_log"; - -export interface MemoryStats { - session_id: string; - active_memories: number; - archived_memories: number; - unresolved_errors: number; - resolved_errors: number; - memory_by_type: Record; - last_updated: number; -} - -export interface SaveMemoryRequest { - session_id: string; - file_type: MemoryFileType; - title: string; - content: string; - tags: string[]; - priority: number; -} - -export interface RecordErrorRequest { - session_id: string; - error_description: string; - attempted_solution: string; -} - -export interface ResolveErrorRequest { - session_id: string; - error_description: string; - resolution: string; -} - -/** - * 上下文记忆管理 API 类 - */ -export class ContextMemoryAPI { - /** - * 保存记忆条目 - */ - static async saveMemoryEntry(request: SaveMemoryRequest): Promise { - return safeInvoke("save_memory_entry", { request }); - } - - /** - * 获取会话记忆 - */ - static async getSessionMemories( - sessionId: string, - fileType?: MemoryFileType, - ): Promise { - return safeInvoke("get_session_memories", { - sessionId, - fileType: fileType || null, - }); - } - - /** - * 获取记忆上下文(用于 AI 上下文) - */ - static async getMemoryContext(sessionId: string): Promise { - return safeInvoke("get_memory_context", { sessionId }); - } - - /** - * 记录错误 - */ - static async recordError(request: RecordErrorRequest): Promise { - return safeInvoke("record_error", { request }); - } - - /** - * 检查是否应该避免某个操作(3次错误协议) - */ - static async shouldAvoidOperation( - sessionId: string, - operationDescription: string, - ): Promise { - return safeInvoke("should_avoid_operation", { - sessionId, - operationDescription, - }); - } - - /** - * 标记错误已解决 - */ - static async markErrorResolved(request: ResolveErrorRequest): Promise { - return safeInvoke("mark_error_resolved", { request }); - } - - /** - * 获取记忆统计信息 - */ - static async getMemoryStats(sessionId: string): Promise { - return safeInvoke("get_memory_stats", { sessionId }); - } - - /** - * 清理过期记忆 - */ - static async cleanupExpiredMemories(): Promise { - return safeInvoke("cleanup_expired_memories"); - } - - /** - * 保存任务计划记忆 - */ - static async saveTaskPlan( - sessionId: string, - title: string, - content: string, - priority: number = 3, - ): Promise { - return this.saveMemoryEntry({ - session_id: sessionId, - file_type: "task_plan", - title, - content, - tags: ["任务计划"], - priority, - }); - } - - /** - * 保存研究发现 - */ - static async saveFinding( - sessionId: string, - title: string, - content: string, - tags: string[] = [], - priority: number = 4, - ): Promise { - return this.saveMemoryEntry({ - session_id: sessionId, - file_type: "findings", - title, - content, - tags: ["发现", ...tags], - priority, - }); - } - - /** - * 记录进度 - */ - static async logProgress( - sessionId: string, - title: string, - content: string, - ): Promise { - return this.saveMemoryEntry({ - session_id: sessionId, - file_type: "progress", - title, - content, - tags: ["进度"], - priority: 2, - }); - } - - /** - * 应用 2-Action 规则:每2次视觉操作后保存发现 - */ - static async apply2ActionRule( - sessionId: string, - finding: string, - ): Promise { - const timestamp = new Date().toLocaleTimeString(); - return this.saveFinding( - sessionId, - `2-Action 规则发现 (${timestamp})`, - finding, - ["2-Action规则", "自动保存"], - 4, - ); - } - - /** - * 记录错误并检查是否需要避免重复操作 - */ - static async recordErrorWithCheck( - sessionId: string, - errorDescription: string, - attemptedSolution: string, - operationDescription?: string, - ): Promise<{ shouldAvoid: boolean }> { - // 记录错误 - await this.recordError({ - session_id: sessionId, - error_description: errorDescription, - attempted_solution: attemptedSolution, - }); - - // 检查是否应该避免该操作 - const shouldAvoid = operationDescription - ? await this.shouldAvoidOperation(sessionId, operationDescription) - : false; - - return { shouldAvoid }; - } -} - -export default ContextMemoryAPI; diff --git a/src/lib/api/memoryRuntime.test.ts b/src/lib/api/memoryRuntime.test.ts index 157c136b7..7befc1363 100644 --- a/src/lib/api/memoryRuntime.test.ts +++ b/src/lib/api/memoryRuntime.test.ts @@ -1,14 +1,14 @@ import { beforeEach, describe, expect, it, vi } from "vitest"; import { safeInvoke } from "@/lib/dev-bridge"; import { - cleanupMemory, - getMemoryAutoIndex, - getMemoryEffectiveSources, - getMemoryOverview, - getMemoryStats, - requestMemoryAnalysis, - toggleMemoryAuto, - updateMemoryAutoNote, + analyzeContextMemory, + cleanupContextMemory, + getContextMemoryAutoIndex, + getContextMemoryEffectiveSources, + getContextMemoryOverview, + getContextMemoryStats, + toggleContextMemoryAuto, + updateContextMemoryAutoNote, } from "./memoryRuntime"; vi.mock("@/lib/dev-bridge", () => ({ @@ -20,7 +20,7 @@ describe("memoryRuntime API", () => { vi.clearAllMocks(); }); - it("应代理记忆查询命令", async () => { + it("应通过 context memory 命名代理记忆查询命令", async () => { vi.mocked(safeInvoke).mockImplementation(async (command) => { switch (command) { case "memory_runtime_get_stats": @@ -44,22 +44,22 @@ describe("memoryRuntime API", () => { } }); - await expect(getMemoryStats()).resolves.toEqual( + await expect(getContextMemoryStats()).resolves.toEqual( expect.objectContaining({ total_entries: 1 }), ); - await expect(requestMemoryAnalysis()).resolves.toEqual( + await expect(analyzeContextMemory()).resolves.toEqual( expect.objectContaining({ analyzed_sessions: 1 }), ); - await expect(cleanupMemory()).resolves.toEqual( + await expect(cleanupContextMemory()).resolves.toEqual( expect.objectContaining({ cleaned_entries: 1 }), ); - await expect(getMemoryOverview(200)).resolves.toEqual( + await expect(getContextMemoryOverview(200)).resolves.toEqual( expect.objectContaining({ entries: [] }), ); - await expect(getMemoryEffectiveSources()).resolves.toEqual( + await expect(getContextMemoryEffectiveSources()).resolves.toEqual( expect.objectContaining({ sources: [] }), ); - await expect(getMemoryAutoIndex()).resolves.toEqual( + await expect(getContextMemoryAutoIndex()).resolves.toEqual( expect.objectContaining({ items: [] }), ); @@ -81,15 +81,65 @@ describe("memoryRuntime API", () => { }); }); - it("应代理自动记忆开关与写入命令", async () => { + it("应暴露清晰的 context memory 命名", async () => { + vi.mocked(safeInvoke).mockImplementation(async (command) => { + switch (command) { + case "memory_runtime_get_stats": + return { total_entries: 9 }; + case "memory_runtime_request_analysis": + return { analyzed_sessions: 3 }; + case "memory_runtime_cleanup": + return { cleaned_entries: 4 }; + case "memory_runtime_get_overview": + return { stats: {}, categories: [], entries: [] }; + case "memory_get_effective_sources": + return { sources: [] }; + case "memory_get_auto_index": + return { items: [] }; + case "memory_toggle_auto": + return { enabled: true }; + case "memory_update_auto_note": + return { items: [] }; + default: + return null; + } + }); + + await expect(getContextMemoryStats()).resolves.toEqual( + expect.objectContaining({ total_entries: 9 }), + ); + await expect(analyzeContextMemory()).resolves.toEqual( + expect.objectContaining({ analyzed_sessions: 3 }), + ); + await expect(cleanupContextMemory()).resolves.toEqual( + expect.objectContaining({ cleaned_entries: 4 }), + ); + await expect(getContextMemoryOverview()).resolves.toEqual( + expect.objectContaining({ entries: [] }), + ); + await expect(getContextMemoryEffectiveSources()).resolves.toEqual( + expect.objectContaining({ sources: [] }), + ); + await expect(getContextMemoryAutoIndex()).resolves.toEqual( + expect.objectContaining({ items: [] }), + ); + await expect(toggleContextMemoryAuto(true)).resolves.toEqual( + expect.objectContaining({ enabled: true }), + ); + await expect(updateContextMemoryAutoNote("note")).resolves.toEqual( + expect.objectContaining({ items: [] }), + ); + }); + + it("应代理 context memory 自动记忆开关与写入命令", async () => { vi.mocked(safeInvoke) .mockResolvedValueOnce({ enabled: true }) .mockResolvedValueOnce({ items: [] }); - await expect(toggleMemoryAuto(true)).resolves.toEqual( + await expect(toggleContextMemoryAuto(true)).resolves.toEqual( expect.objectContaining({ enabled: true }), ); - await expect(updateMemoryAutoNote("note", "topic")).resolves.toEqual( + await expect(updateContextMemoryAutoNote("note", "topic")).resolves.toEqual( expect.objectContaining({ items: [] }), ); }); diff --git a/src/lib/api/memoryRuntime.ts b/src/lib/api/memoryRuntime.ts index d4e3ccf02..dfa9b20cd 100644 --- a/src/lib/api/memoryRuntime.ts +++ b/src/lib/api/memoryRuntime.ts @@ -28,17 +28,17 @@ export type { MemoryStatsResponse, } from "./memoryRuntimeTypes"; -export async function getMemoryOverview( +export async function getContextMemoryOverview( limit?: number, ): Promise { return safeInvoke("memory_runtime_get_overview", { limit }); } -export async function getMemoryStats(): Promise { +export async function getContextMemoryStats(): Promise { return safeInvoke("memory_runtime_get_stats"); } -export async function requestMemoryAnalysis( +export async function analyzeContextMemory( fromTimestamp?: number, toTimestamp?: number, ): Promise { @@ -48,11 +48,11 @@ export async function requestMemoryAnalysis( }); } -export async function cleanupMemory(): Promise { +export async function cleanupContextMemory(): Promise { return safeInvoke("memory_runtime_cleanup"); } -export async function getMemoryEffectiveSources( +export async function getContextMemoryEffectiveSources( workingDir?: string, activeRelativePath?: string, ): Promise { @@ -62,19 +62,19 @@ export async function getMemoryEffectiveSources( }); } -export async function getMemoryAutoIndex( +export async function getContextMemoryAutoIndex( workingDir?: string, ): Promise { return safeInvoke("memory_get_auto_index", { workingDir }); } -export async function toggleMemoryAuto( +export async function toggleContextMemoryAuto( enabled: boolean, ): Promise { return safeInvoke("memory_toggle_auto", { enabled }); } -export async function updateMemoryAutoNote( +export async function updateContextMemoryAutoNote( note: string, topic?: string, workingDir?: string, diff --git a/src/lib/api/openclaw.ts b/src/lib/api/openclaw.ts index 94efa5b03..cb2955d12 100644 --- a/src/lib/api/openclaw.ts +++ b/src/lib/api/openclaw.ts @@ -47,6 +47,9 @@ export interface OpenClawEnvironmentDiagnostics { whereCandidates?: string[]; supplementalSearchDirs?: string[]; supplementalCommandCandidates?: string[]; + gitWhereCandidates?: string[]; + gitSupplementalSearchDirs?: string[]; + gitSupplementalCommandCandidates?: string[]; } export interface OpenClawActionResult { diff --git a/src/lib/api/toolHooks.test.ts b/src/lib/api/toolHooks.test.ts deleted file mode 100644 index 0eb666d6c..000000000 --- a/src/lib/api/toolHooks.test.ts +++ /dev/null @@ -1,100 +0,0 @@ -import { beforeEach, describe, expect, it, vi } from "vitest"; -import { safeInvoke } from "@/lib/dev-bridge"; -import { ToolHooksAPI } from "./toolHooks"; - -vi.mock("@/lib/dev-bridge", () => ({ - safeInvoke: vi.fn(), -})); - -describe("toolHooks API", () => { - beforeEach(() => { - vi.clearAllMocks(); - }); - - it("应代理钩子规则管理命令", async () => { - vi.mocked(safeInvoke) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce([{ id: "rule-1", name: "规则1" }]) - .mockResolvedValueOnce({ "rule-1": { execution_count: 3 } }) - .mockResolvedValueOnce(undefined); - - const rule = ToolHooksAPI.createCustomRule( - "rule-1", - "规则1", - "说明", - "session_start", - [], - [], - ); - - await expect( - ToolHooksAPI.executeHooks({ - trigger: "session_start", - context: { session_id: "session-1", message_count: 0, metadata: {} }, - }), - ).resolves.toBeUndefined(); - await expect(ToolHooksAPI.addHookRule(rule)).resolves.toBeUndefined(); - await expect( - ToolHooksAPI.removeHookRule("rule-1"), - ).resolves.toBeUndefined(); - await expect( - ToolHooksAPI.toggleHookRule("rule-1", true), - ).resolves.toBeUndefined(); - await expect(ToolHooksAPI.getHookRules()).resolves.toEqual([ - expect.objectContaining({ id: "rule-1" }), - ]); - await expect(ToolHooksAPI.getHookExecutionStats()).resolves.toEqual( - expect.objectContaining({ "rule-1": expect.any(Object) }), - ); - await expect( - ToolHooksAPI.clearHookExecutionStats(), - ).resolves.toBeUndefined(); - }); - - it("应基于上下文辅助方法生成 executeHooks 请求", async () => { - vi.mocked(safeInvoke) - .mockResolvedValue(undefined) - .mockResolvedValue(undefined) - .mockResolvedValue(undefined); - - await ToolHooksAPI.triggerSessionStart("session-1", { source: "test" }); - await ToolHooksAPI.triggerPreToolUse( - "session-1", - "read_file", - { path: "/tmp/a.ts" }, - "读取文件", - 2, - ); - await ToolHooksAPI.triggerStop("session-1", 5, { source: "test" }); - - expect(safeInvoke).toHaveBeenNthCalledWith(1, "execute_hooks", { - request: expect.objectContaining({ - trigger: "session_start", - context: expect.objectContaining({ - session_id: "session-1", - metadata: expect.objectContaining({ source: "test" }), - }), - }), - }); - expect(safeInvoke).toHaveBeenNthCalledWith(2, "execute_hooks", { - request: expect.objectContaining({ - trigger: "pre_tool_use", - context: expect.objectContaining({ - tool_name: "read_file", - message_count: 2, - }), - }), - }); - expect(safeInvoke).toHaveBeenNthCalledWith(3, "execute_hooks", { - request: expect.objectContaining({ - trigger: "stop", - context: expect.objectContaining({ - message_count: 5, - }), - }), - }); - }); -}); diff --git a/src/lib/api/toolHooks.ts b/src/lib/api/toolHooks.ts deleted file mode 100644 index 88c87bc45..000000000 --- a/src/lib/api/toolHooks.ts +++ /dev/null @@ -1,344 +0,0 @@ -/** - * 工具钩子管理 API - * - * 提供工具执行前后的钩子机制,用于自动化上下文记忆管理 - */ - -import { safeInvoke } from "@/lib/dev-bridge"; - -export type HookTrigger = - | "session_start" - | "pre_tool_use" - | "post_tool_use" - | "stop"; - -export interface HookRule { - id: string; - name: string; - description: string; - trigger: HookTrigger; - conditions: HookCondition[]; - actions: HookAction[]; - enabled: boolean; - priority: number; - created_at: number; -} - -export type HookCondition = - | { tool_name_equals: string } - | { tool_name_contains: string } - | { message_contains: string } - | { message_count_greater_than: number } - | { error_count_greater_than: number } - | { custom: { condition_type: string; parameters: Record } }; - -export type HookAction = - | { - save_finding: { - title: string; - content: string; - tags: string[]; - priority: number; - }; - } - | { update_task_plan: { title: string; content: string; priority: number } } - | { log_progress: { title: string; content: string } } - | { record_error: { error_description: string; attempted_solution: string } } - | { custom: { action_type: string; parameters: Record } }; - -export interface HookExecutionStats { - execution_count: number; - success_count: number; - failure_count: number; - last_execution_at: number; - average_execution_time_ms: number; -} - -export interface HookContextData { - session_id: string; - tool_name?: string; - tool_parameters?: Record; - tool_result?: string; - message_content?: string; - message_count: number; - error_info?: string; - metadata: Record; -} - -export interface ExecuteHooksRequest { - trigger: HookTrigger; - context: HookContextData; -} - -/** - * 工具钩子管理 API 类 - */ -export class ToolHooksAPI { - /** - * 执行钩子 - */ - static async executeHooks(request: ExecuteHooksRequest): Promise { - return safeInvoke("execute_hooks", { request }); - } - - /** - * 添加钩子规则 - */ - static async addHookRule(rule: HookRule): Promise { - return safeInvoke("add_hook_rule", { rule }); - } - - /** - * 移除钩子规则 - */ - static async removeHookRule(ruleId: string): Promise { - return safeInvoke("remove_hook_rule", { ruleId }); - } - - /** - * 启用/禁用钩子规则 - */ - static async toggleHookRule(ruleId: string, enabled: boolean): Promise { - return safeInvoke("toggle_hook_rule", { ruleId, enabled }); - } - - /** - * 获取所有钩子规则 - */ - static async getHookRules(): Promise { - return safeInvoke("get_hook_rules"); - } - - /** - * 获取钩子执行统计 - */ - static async getHookExecutionStats(): Promise< - Record - > { - return safeInvoke>( - "get_hook_execution_stats", - ); - } - - /** - * 清理钩子执行统计 - */ - static async clearHookExecutionStats(): Promise { - return safeInvoke("clear_hook_execution_stats"); - } - - /** - * 触发会话开始钩子 - */ - static async triggerSessionStart( - sessionId: string, - metadata: Record = {}, - ): Promise { - return this.executeHooks({ - trigger: "session_start", - context: { - session_id: sessionId, - message_count: 0, - metadata: { - timestamp: new Date().toISOString(), - ...metadata, - }, - }, - }); - } - - /** - * 触发工具使用前钩子 - */ - static async triggerPreToolUse( - sessionId: string, - toolName: string, - toolParameters: Record = {}, - messageContent?: string, - messageCount: number = 0, - ): Promise { - return this.executeHooks({ - trigger: "pre_tool_use", - context: { - session_id: sessionId, - tool_name: toolName, - tool_parameters: toolParameters, - message_content: messageContent, - message_count: messageCount, - metadata: { - timestamp: new Date().toISOString(), - tool_name: toolName, - }, - }, - }); - } - - /** - * 触发工具使用后钩子 - */ - static async triggerPostToolUse( - sessionId: string, - toolName: string, - toolResult: string, - toolParameters: Record = {}, - messageContent?: string, - messageCount: number = 0, - errorInfo?: string, - ): Promise { - return this.executeHooks({ - trigger: "post_tool_use", - context: { - session_id: sessionId, - tool_name: toolName, - tool_parameters: toolParameters, - tool_result: toolResult, - message_content: messageContent, - message_count: messageCount, - error_info: errorInfo, - metadata: { - timestamp: new Date().toISOString(), - tool_name: toolName, - has_error: errorInfo ? "true" : "false", - }, - }, - }); - } - - /** - * 触发会话停止钩子 - */ - static async triggerStop( - sessionId: string, - messageCount: number, - metadata: Record = {}, - ): Promise { - return this.executeHooks({ - trigger: "stop", - context: { - session_id: sessionId, - message_count: messageCount, - metadata: { - timestamp: new Date().toISOString(), - session_end: "true", - ...metadata, - }, - }, - }); - } - - /** - * 创建自定义钩子规则 - */ - static createCustomRule( - id: string, - name: string, - description: string, - trigger: HookTrigger, - conditions: HookCondition[], - actions: HookAction[], - priority: number = 100, - ): HookRule { - return { - id, - name, - description, - trigger, - conditions, - actions, - enabled: true, - priority, - created_at: Date.now(), - }; - } - - /** - * 创建重要发现自动保存规则 - */ - static createImportantFindingRule(): HookRule { - return this.createCustomRule( - "important-finding-auto-save", - "重要发现自动保存", - "检测到重要信息时自动保存到 findings.md", - "post_tool_use", - [{ message_contains: "重要" }, { message_contains: "发现" }], - [ - { - save_finding: { - title: "重要发现 (自动检测)", - content: "检测到重要信息,已自动保存", - tags: ["重要", "自动保存"], - priority: 4, - }, - }, - ], - 1, - ); - } - - /** - * 创建错误自动记录规则 - */ - static createErrorAutoRecordRule(): HookRule { - return this.createCustomRule( - "error-auto-record", - "错误自动记录", - "检测到错误时自动记录到错误日志", - "post_tool_use", - [{ message_contains: "错误" }], - [ - { - record_error: { - error_description: "检测到错误", - attempted_solution: "正在尝试解决", - }, - }, - ], - 1, - ); - } - - /** - * 创建 2-Action 规则 - */ - static create2ActionRule(): HookRule { - return this.createCustomRule( - "2-action-rule", - "2-Action 规则", - "每2次视觉操作后自动保存发现", - "post_tool_use", - [{ tool_name_contains: "view" }], - [ - { - save_finding: { - title: "2-Action 规则触发", - content: "视觉操作完成,自动保存发现", - tags: ["2-Action规则", "视觉操作"], - priority: 3, - }, - }, - ], - 2, - ); - } - - /** - * 批量添加默认钩子规则 - */ - static async addDefaultRules(): Promise { - const rules = [ - this.createImportantFindingRule(), - this.createErrorAutoRecordRule(), - this.create2ActionRule(), - ]; - - for (const rule of rules) { - try { - await this.addHookRule(rule); - } catch (error) { - console.warn(`添加钩子规则失败: ${rule.name}`, error); - } - } - } -} - -export default ToolHooksAPI; diff --git a/src/lib/api/unified-chat.test.ts b/src/lib/api/unified-chat.test.ts deleted file mode 100644 index c75c06c47..000000000 --- a/src/lib/api/unified-chat.test.ts +++ /dev/null @@ -1,145 +0,0 @@ -import { beforeEach, describe, expect, it, vi } from "vitest"; -import { safeInvoke } from "@/lib/dev-bridge"; -import { - configureProvider, - createSession, - deleteSession, - generateEventName, - getMessages, - getSession, - listSessions, - parseStreamEvent, - renameSession, - sendMessage, - stopGeneration, -} from "./unified-chat"; - -vi.mock("@/lib/dev-bridge", () => ({ - safeInvoke: vi.fn(), -})); - -describe("unified-chat API", () => { - beforeEach(() => { - vi.clearAllMocks(); - }); - - it("应代理会话与消息命令并转换消息结构", async () => { - vi.mocked(safeInvoke) - .mockResolvedValueOnce({ id: "session-1", title: "测试会话" }) - .mockResolvedValueOnce([{ id: "session-1", title: "测试会话" }]) - .mockResolvedValueOnce({ id: "session-1", title: "测试会话" }) - .mockResolvedValueOnce(true) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce([ - { - id: 1, - session_id: "session-1", - role: "assistant", - content: [{ type: "text", text: "你好" }], - created_at: "2024-01-01T00:00:00Z", - }, - ]) - .mockResolvedValueOnce(undefined) - .mockResolvedValueOnce(true) - .mockResolvedValueOnce(undefined); - - await expect(createSession({ mode: "general" } as never)).resolves.toEqual( - expect.objectContaining({ id: "session-1" }), - ); - await expect(listSessions()).resolves.toEqual([ - expect.objectContaining({ id: "session-1" }), - ]); - await expect(getSession("session-1")).resolves.toEqual( - expect.objectContaining({ id: "session-1" }), - ); - await expect(deleteSession("session-1")).resolves.toBe(true); - await expect(renameSession("session-1", "新标题")).resolves.toBeUndefined(); - await expect(getMessages("session-1")).resolves.toEqual([ - expect.objectContaining({ - sessionId: "session-1", - content: "你好", - }), - ]); - await expect( - sendMessage({ sessionId: "session-1", content: "hello" } as never), - ).resolves.toBeUndefined(); - await expect(stopGeneration("session-1")).resolves.toBe(true); - await expect( - configureProvider("session-1", "openai", "gpt-4o"), - ).resolves.toBeUndefined(); - }); - - it("应解析流事件并生成事件名", () => { - expect(parseStreamEvent({ type: "text_delta", text: "hi" })).toEqual({ - type: "text_delta", - text: "hi", - }); - expect( - parseStreamEvent({ type: "tool_end", tool_id: "tool-1", result: "done" }), - ).toEqual({ - type: "tool_end", - tool_id: "tool-1", - result: "done", - }); - expect( - parseStreamEvent({ - type: "harness_event", - kind: "artifact_created", - session_id: "session-1", - stage: "drafting", - }), - ).toEqual({ - type: "harness_event", - event: { - kind: "artifact_created", - sessionId: "session-1", - runId: undefined, - correlationId: undefined, - theme: undefined, - stage: "drafting", - summary: undefined, - artifact: undefined, - metadata: undefined, - }, - }); - expect( - parseStreamEvent({ - type: "artifact_snapshot", - artifact_id: "artifact-1", - file_path: "draft.md", - content: "# 标题", - }), - ).toEqual({ - type: "artifact_snapshot", - artifact: { - artifactId: "artifact-1", - filePath: "draft.md", - content: "# 标题", - metadata: undefined, - }, - }); - expect( - parseStreamEvent({ - type: "artifact_snapshot", - artifact: { - artifactId: "artifact-2", - filePath: "nested.md", - content: "nested content", - metadata: { complete: false }, - }, - }), - ).toEqual({ - type: "artifact_snapshot", - artifact: { - artifactId: "artifact-2", - filePath: "nested.md", - content: "nested content", - metadata: { complete: false }, - }, - }); - expect(parseStreamEvent({ type: "unknown" })).toBeNull(); - expect(generateEventName("session-1")).toMatch( - /^unified-chat-stream-session-1-/, - ); - }); -}); diff --git a/src/lib/api/unified-chat.ts b/src/lib/api/unified-chat.ts deleted file mode 100644 index b58841b52..000000000 --- a/src/lib/api/unified-chat.ts +++ /dev/null @@ -1,372 +0,0 @@ -/** - * @file unified-chat.ts - * @description 统一对话 API 封装 - * @module lib/api/unified-chat - * - * 封装所有统一对话相关的 Tauri 命令调用 - */ - -import { safeInvoke } from "@/lib/dev-bridge"; -import type { - ChatMode, - ChatMessage, - SessionResponse, - CreateSessionRequest, - SendMessageRequest, - StreamEvent, - ToolCall, - ToolEndEvent, - FinalDoneEvent, - HarnessArtifactSnapshot, - HarnessEventPayload, -} from "@/types/chat"; - -// ============================================================================ -// 会话管理 API -// ============================================================================ - -/** - * 创建新会话 - */ -export async function createSession( - request: CreateSessionRequest, -): Promise { - return safeInvoke("chat_create_session", { request }); -} - -/** - * 获取会话列表 - */ -export async function listSessions( - mode?: ChatMode, -): Promise { - return safeInvoke("chat_list_sessions", { mode }); -} - -/** - * 获取会话详情 - */ -export async function getSession(sessionId: string): Promise { - return safeInvoke("chat_get_session", { sessionId }); -} - -/** - * 删除会话 - */ -export async function deleteSession(sessionId: string): Promise { - return safeInvoke("chat_delete_session", { sessionId }); -} - -/** - * 重命名会话 - */ -export async function renameSession( - sessionId: string, - title: string, -): Promise { - return safeInvoke("chat_rename_session", { sessionId, title }); -} - -// ============================================================================ -// 消息管理 API -// ============================================================================ - -/** - * 获取会话消息列表 - */ -export async function getMessages( - sessionId: string, - limit?: number, -): Promise { - const messages = await safeInvoke< - Array<{ - id: number; - session_id: string; - role: string; - content: unknown; - tool_calls?: unknown; - tool_call_id?: string; - metadata?: unknown; - created_at: string; - }> - >("chat_get_messages", { sessionId, limit }); - - // 转换后端格式为前端格式 - return messages.map(convertBackendMessage); -} - -/** - * 发送消息(流式) - */ -export async function sendMessage(request: SendMessageRequest): Promise { - return safeInvoke("chat_send_message", { request }); -} - -/** - * 停止生成 - */ -export async function stopGeneration(sessionId: string): Promise { - return safeInvoke("chat_stop_generation", { sessionId }); -} - -/** - * 配置会话的 Provider - */ -export async function configureProvider( - sessionId: string, - providerType: string, - model: string, -): Promise { - return safeInvoke("chat_configure_provider", { - sessionId, - providerType, - model, - }); -} - -// ============================================================================ -// 辅助函数 -// ============================================================================ - -/** - * 转换后端消息格式为前端格式 - */ -function convertBackendMessage(msg: { - id: number; - session_id: string; - role: string; - content: unknown; - tool_calls?: unknown; - tool_call_id?: string; - metadata?: unknown; - created_at: string; -}): ChatMessage { - // 提取文本内容 - let textContent = ""; - if (typeof msg.content === "string") { - textContent = msg.content; - } else if (Array.isArray(msg.content)) { - textContent = msg.content - .filter( - (part): part is { type: "text"; text: string } => - typeof part === "object" && part !== null && part.type === "text", - ) - .map((part) => part.text) - .join("\n"); - } else if (typeof msg.content === "object" && msg.content !== null) { - // 尝试从对象中提取文本 - const contentObj = msg.content as Record; - if (typeof contentObj.text === "string") { - textContent = contentObj.text; - } - } - - return { - id: msg.id, - sessionId: msg.session_id, - role: msg.role as ChatMessage["role"], - content: textContent, - contentBlocks: Array.isArray(msg.content) - ? msg.content.map(convertContentBlock) - : [{ type: "text" as const, text: textContent }], - toolCalls: msg.tool_calls ? convertToolCalls(msg.tool_calls) : undefined, - toolCallId: msg.tool_call_id, - status: "complete", - metadata: msg.metadata as ChatMessage["metadata"], - createdAt: msg.created_at, - }; -} - -/** - * 转换内容块 - */ -function convertContentBlock( - block: unknown, -): NonNullable[number] { - if (typeof block !== "object" || block === null) { - return { type: "text", text: String(block) }; - } - - const b = block as Record; - - switch (b.type) { - case "text": - return { type: "text", text: String(b.text || "") }; - case "image": - return { type: "image", url: String(b.url || ""), alt: b.alt as string }; - case "file": - return { - type: "file", - path: String(b.path || ""), - name: String(b.name || ""), - }; - case "canvas": - return { - type: "canvas", - canvasType: String(b.canvasType || ""), - content: String(b.content || ""), - }; - default: - return { type: "text", text: JSON.stringify(block) }; - } -} - -/** - * 转换工具调用 - */ -function convertToolCalls(toolCalls: unknown): ChatMessage["toolCalls"] { - if (!Array.isArray(toolCalls)) return undefined; - - return toolCalls.map((tc) => { - const call = tc as Record; - return { - id: String(call.id || ""), - name: String(call.name || ""), - arguments: call.arguments as string, - status: (call.status as ToolCall["status"]) || "completed", - result: call.result as ToolCall["result"], - }; - }); -} - -/** - * 解析流式事件 - */ -export function parseStreamEvent(payload: unknown): StreamEvent | null { - if (typeof payload !== "object" || payload === null) { - return null; - } - - const event = payload as Record; - const type = event.type as string; - - switch (type) { - case "TextDelta": - case "text_delta": - return { - type: "text_delta", - text: String(event.text || event.content || ""), - }; - - case "ThinkingDelta": - case "thinking_delta": - return { - type: "thinking_delta", - text: String(event.text || event.content || ""), - }; - - case "ToolStart": - case "tool_start": - return { - type: "tool_start", - tool_id: String(event.tool_id || event.id || ""), - tool_name: String(event.tool_name || event.name || ""), - arguments: event.arguments as string, - }; - - case "ToolEnd": - case "tool_end": - return { - type: "tool_end", - tool_id: String(event.tool_id || event.id || ""), - result: event.result as ToolEndEvent["result"], - }; - - case "HarnessEvent": - case "harness_event": - return { - type: "harness_event", - event: { - kind: String(event.kind || event.event_kind || "unknown"), - sessionId: - (event.session_id as string) || (event.sessionId as string), - runId: (event.run_id as string) || (event.runId as string), - correlationId: - (event.correlation_id as string) || - (event.correlationId as string), - theme: event.theme as string, - stage: event.stage as string, - summary: event.summary as string, - artifact: (event.artifact || - (event.snapshot as HarnessArtifactSnapshot) || - undefined) as HarnessArtifactSnapshot | undefined, - metadata: event.metadata as HarnessEventPayload["metadata"], - }, - }; - - case "ArtifactSnapshot": - case "artifact_snapshot": - { - const nestedArtifact = - event.artifact && typeof event.artifact === "object" - ? (event.artifact as Record) - : undefined; - return { - type: "artifact_snapshot", - artifact: { - artifactId: String( - nestedArtifact?.artifactId || - nestedArtifact?.artifact_id || - event.artifact_id || - event.artifactId || - event.id || - "artifact-unknown", - ), - filePath: - (nestedArtifact?.filePath as string | undefined) || - (nestedArtifact?.file_path as string | undefined) || - (event.file_path as string | undefined) || - (event.filePath as string | undefined), - content: - (nestedArtifact?.content as string | undefined) || - (event.content as string | undefined), - metadata: - (nestedArtifact?.metadata as Record | undefined) || - (event.metadata as Record | undefined), - }, - }; - } - - case "ActionRequired": - case "action_required": - return { - type: "action_required", - request_id: String(event.request_id || ""), - action_type: String(event.action_type || ""), - tool_name: event.tool_name as string, - arguments: event.arguments as string, - prompt: event.prompt as string, - questions: event.questions as unknown[], - requested_schema: event.requested_schema, - }; - - case "Done": - case "done": - return { type: "done" }; - - case "FinalDone": - case "final_done": - return { - type: "final_done", - usage: event.usage as FinalDoneEvent["usage"], - }; - - case "Error": - case "error": - return { - type: "error", - message: String(event.message || "Unknown error"), - }; - - default: - console.warn("[parseStreamEvent] 未知事件类型:", type); - return null; - } -} - -/** - * 生成唯一事件名称 - */ -export function generateEventName(sessionId: string): string { - return `unified-chat-stream-${sessionId}-${Date.now()}`; -} diff --git a/src/lib/governance/agentCommandCatalog.json b/src/lib/governance/agentCommandCatalog.json index d7d5a0980..626039db0 100644 --- a/src/lib/governance/agentCommandCatalog.json +++ b/src/lib/governance/agentCommandCatalog.json @@ -63,22 +63,22 @@ "legacyCommandSurfaceMonitors": [ { "id": "agent-create-session-compat-command", - "classification": "compat", - "description": "agent_create_session compat 命令前端边界", + "classification": "dead-candidate", + "description": "已零引用的旧 agent_create_session 命令边界", "commands": ["agent_create_session"], - "allowedPaths": ["src/lib/api/agentRuntime.ts"] + "allowedPaths": [] }, { "id": "agent-session-message-legacy-command", - "classification": "deprecated", - "description": "旧 agent session message 命令前端边界", + "classification": "dead-candidate", + "description": "已零引用的旧 agent session message 命令边界", "commands": ["agent_get_session_messages"], - "allowedPaths": ["src/lib/api/agentRuntime.ts"] + "allowedPaths": [] }, { "id": "agent-session-compat-commands", - "classification": "deprecated", - "description": "旧 agent session compat 命令前端边界", + "classification": "dead-candidate", + "description": "已零引用的旧 agent session compat 命令边界", "commands": [ "agent_list_sessions", "agent_get_session", @@ -91,8 +91,8 @@ "legacyHelperSurfaceMonitors": [ { "id": "agent-legacy-session-api-helpers", - "classification": "deprecated", - "description": "旧 Agent session compat helper 直连回流", + "classification": "dead-candidate", + "description": "已零引用的旧 Agent session compat helper 回流", "helpers": [ "createAgentSession", "listAgentSessions", @@ -102,24 +102,24 @@ "deleteAgentSession", "generateAgentTitle" ], - "allowedPaths": ["src/lib/api/agentRuntime.ts"] + "allowedPaths": [] }, { "id": "agent-legacy-stream-action-helpers", - "classification": "deprecated", - "description": "旧 Aster stream/action helper 直连回流", + "classification": "dead-candidate", + "description": "已零引用的旧 Aster stream/action helper 回流", "helpers": [ "sendAsterMessageStream", "confirmAsterAction", "submitAsterElicitationResponse", "stopAsterSession" ], - "allowedPaths": ["src/lib/api/agentRuntime.ts"] + "allowedPaths": [] }, { "id": "aster-session-helper-direct-usage", - "classification": "deprecated", - "description": "前端 direct Aster session helper 回流", + "classification": "dead-candidate", + "description": "已零引用的前端 direct Aster session helper 回流", "helpers": [ "createAsterSession", "listAsterSessions", @@ -127,7 +127,7 @@ "deleteAsterSession", "renameAsterSession" ], - "allowedPaths": ["src/lib/api/agentRuntime.ts"] + "allowedPaths": [] } ] } diff --git a/src/lib/model/inferModelCapabilities.test.ts b/src/lib/model/inferModelCapabilities.test.ts new file mode 100644 index 000000000..283aa2828 --- /dev/null +++ b/src/lib/model/inferModelCapabilities.test.ts @@ -0,0 +1,41 @@ +import { describe, expect, it } from "vitest"; +import { + inferModelCapabilities, + inferVisionCapability, +} from "./inferModelCapabilities"; + +describe("inferModelCapabilities", () => { + it("应将 gpt-5.4 识别为支持视觉的模型", () => { + expect( + inferVisionCapability({ + modelId: "gpt-5.4", + providerId: "codex", + }), + ).toBe(true); + }); + + it("应避免将生图模型误判为视觉聊天模型", () => { + expect( + inferVisionCapability({ + modelId: "gemini-3-pro-image-preview", + providerId: "gemini", + }), + ).toBe(false); + }); + + it("应保留 thinking 模型的推理能力推断", () => { + expect( + inferModelCapabilities({ + modelId: "gpt-5.4-thinking", + providerId: "openai", + }), + ).toMatchObject({ + vision: true, + reasoning: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + }); + }); +}); diff --git a/src/lib/model/inferModelCapabilities.ts b/src/lib/model/inferModelCapabilities.ts new file mode 100644 index 000000000..41a50707a --- /dev/null +++ b/src/lib/model/inferModelCapabilities.ts @@ -0,0 +1,105 @@ +import type { ModelCapabilities } from "@/lib/types/modelRegistry"; + +const REASONING_TOKEN_PATTERN = /(^|[._/-])(thinking|reasoning)(?=$|[._/-])/i; +const VISION_HINT_PATTERN = + /\b(vision|multimodal|multi-modal|omni|image-input|image understanding)\b/i; +const NON_VISION_PATTERN = + /\b(embedding|embed|rerank|tts|stt|transcribe|transcription|speech|audio|moderation)\b/i; +const IMAGE_GENERATION_PATTERN = + /\b(imagen|dall-e|dalle|stable[ -]?diffusion|sdxl|sd3|midjourney|mj|flux|image[ -]?generation|image-gen|image-preview)\b/i; +const OPENAI_VISION_PATTERN = + /\b(gpt-5(?:[._/-]|\b)|gpt-4o(?:[._/-]|\b)|gpt-4\.1(?:[._/-]|\b)|gpt-4\.5(?:[._/-]|\b)|gpt-5.*codex)\b/i; +const GEMINI_VISION_PATTERN = /\bgemini(?:[._/-]|\b)/i; +const CLAUDE_VISION_PATTERN = /\bclaude(?:[._/-]|\b)/i; +const QWEN_VISION_PATTERN = /\bqwen(?:[._/-]|\b).*(vl|vision)|\bqvq\b/i; +const GLM_VISION_PATTERN = /\bglm-[\w.-]*v[\w.-]*\b/i; + +const normalize = (value?: string | null): string => + (value || "").trim().toLowerCase(); + +function buildSearchText(parts: Array): string { + return parts + .map((part) => normalize(part)) + .filter(Boolean) + .join(" "); +} + +export function inferReasoningCapability(modelId: string): boolean { + return REASONING_TOKEN_PATTERN.test(modelId.trim().toLowerCase()); +} + +export function inferVisionCapability(params: { + modelId: string; + providerId?: string | null; + family?: string | null; + description?: string | null; +}): boolean { + const { modelId, providerId, family, description } = params; + const text = buildSearchText([modelId, family, description]); + const provider = normalize(providerId); + + if (!text) { + return false; + } + + if (NON_VISION_PATTERN.test(text) || IMAGE_GENERATION_PATTERN.test(text)) { + return false; + } + + if (VISION_HINT_PATTERN.test(text)) { + return true; + } + + if (OPENAI_VISION_PATTERN.test(text)) { + return true; + } + + if (provider === "codex" || provider === "openai") { + return OPENAI_VISION_PATTERN.test(text); + } + + if (provider === "gemini") { + return GEMINI_VISION_PATTERN.test(text); + } + + if (provider === "anthropic" || provider === "claude") { + return CLAUDE_VISION_PATTERN.test(text); + } + + if (provider === "qwen" || provider === "alibaba") { + return QWEN_VISION_PATTERN.test(text); + } + + if (provider === "zhipuai") { + return GLM_VISION_PATTERN.test(text); + } + + return ( + GEMINI_VISION_PATTERN.test(text) || + CLAUDE_VISION_PATTERN.test(text) || + QWEN_VISION_PATTERN.test(text) || + GLM_VISION_PATTERN.test(text) + ); +} + +export function inferModelCapabilities(params: { + modelId: string; + providerId?: string | null; + family?: string | null; + description?: string | null; +}): ModelCapabilities { + const { modelId, providerId, family, description } = params; + return { + vision: inferVisionCapability({ + modelId, + providerId, + family, + description, + }), + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: inferReasoningCapability(modelId), + }; +} diff --git a/src/lib/model/visionModelResolver.test.ts b/src/lib/model/visionModelResolver.test.ts new file mode 100644 index 000000000..39e68758a --- /dev/null +++ b/src/lib/model/visionModelResolver.test.ts @@ -0,0 +1,176 @@ +import { describe, expect, it } from "vitest"; +import type { EnhancedModelMetadata } from "@/lib/types/modelRegistry"; +import { resolveVisionModel } from "./visionModelResolver"; + +function createModel( + id: string, + overrides: Partial = {}, +): EnhancedModelMetadata { + return { + id, + display_name: id, + provider_id: "zhipuai", + provider_name: "Zhipu AI", + family: id, + tier: "pro", + capabilities: { + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, + }, + pricing: null, + limits: { + context_length: null, + max_output_tokens: null, + requests_per_minute: null, + tokens_per_minute: null, + }, + status: "active", + release_date: "2026-01-01", + is_latest: false, + description: null, + source: "local", + created_at: 0, + updated_at: 0, + ...overrides, + }; +} + +describe("resolveVisionModel", () => { + it("当前模型已支持视觉时应保持不变", () => { + const models = [ + createModel("glm-4.6v-flash", { + capabilities: { + vision: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, + }, + }), + ]; + + const result = resolveVisionModel({ + currentModelId: "glm-4.6v-flash", + models, + }); + + expect(result).toEqual({ + targetModelId: "glm-4.6v-flash", + switched: false, + reason: "already_vision", + }); + }); + + it("当前模型未收录但模型名可推断支持视觉时应保持不变", () => { + const models = [ + createModel("gpt-5.3-codex", { + provider_id: "codex", + provider_name: "Codex", + capabilities: { + vision: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: true, + }, + }), + ]; + + const result = resolveVisionModel({ + currentModelId: "gpt-5.4", + models, + }); + + expect(result).toEqual({ + targetModelId: "gpt-5.4", + switched: false, + reason: "already_vision", + }); + }); + + it("应优先选择支持视觉的聊天模型,而不是纯生图模型", () => { + const models = [ + createModel("gemini-3-pro-image-preview", { + family: "gemini-3-pro-image", + capabilities: { + vision: true, + tools: false, + streaming: true, + json_mode: false, + function_calling: false, + reasoning: false, + }, + description: "image generation model", + is_latest: true, + }), + createModel("glm-4.6v-flash", { + family: "glm-4.6v", + tier: "mini", + capabilities: { + vision: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, + }, + release_date: "2026-02-01", + is_latest: true, + }), + createModel("glm-4.7", { + family: "glm-4.7", + capabilities: { + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: true, + }, + }), + ]; + + const result = resolveVisionModel({ + currentModelId: "glm-4.7", + models, + }); + + expect(result.targetModelId).toBe("glm-4.6v-flash"); + expect(result.switched).toBe(true); + expect(result.reason).toBe("fallback_latest"); + }); + + it("没有可用视觉聊天模型时应返回 no_vision_model", () => { + const models = [ + createModel("glm-4.7"), + createModel("gemini-3-pro-image-preview", { + capabilities: { + vision: true, + tools: false, + streaming: true, + json_mode: false, + function_calling: false, + reasoning: false, + }, + description: "image generation model", + }), + ]; + + const result = resolveVisionModel({ + currentModelId: "glm-4.7", + models, + }); + + expect(result).toEqual({ + targetModelId: "glm-4.7", + switched: false, + reason: "no_vision_model", + }); + }); +}); diff --git a/src/lib/model/visionModelResolver.ts b/src/lib/model/visionModelResolver.ts new file mode 100644 index 000000000..72000140e --- /dev/null +++ b/src/lib/model/visionModelResolver.ts @@ -0,0 +1,187 @@ +import type { EnhancedModelMetadata } from "@/lib/types/modelRegistry"; +import { inferVisionCapability } from "./inferModelCapabilities"; + +export type VisionResolveReason = + | "already_vision" + | "matched" + | "fallback_latest" + | "no_vision_model"; + +export interface VisionResolveResult { + targetModelId: string; + switched: boolean; + reason: VisionResolveReason; +} + +interface ResolveVisionModelParams { + currentModelId: string; + models: EnhancedModelMetadata[]; +} + +const IMAGE_GENERATION_KEYWORDS = [ + "imagen", + "dall-e", + "stable-diffusion", + "stable diffusion", + "sdxl", + "sd3", + "midjourney", + "mj", + "flux", + "image generation", + "image-gen", +]; + +const TIER_WEIGHT: Record = { + mini: 1, + pro: 2, + max: 3, +}; + +const normalize = (value?: string | null): string => + (value || "").trim().toLowerCase(); + +const findModelMeta = ( + modelId: string, + models: EnhancedModelMetadata[], +): EnhancedModelMetadata | undefined => { + const normalizedId = normalize(modelId); + return models.find((model) => normalize(model.id) === normalizedId); +}; + +const buildSearchText = (model: EnhancedModelMetadata): string => + [ + model.id, + model.display_name, + model.family || "", + model.description || "", + ] + .join(" ") + .toLowerCase(); + +const isLikelyImageGenerationModel = ( + model: EnhancedModelMetadata, +): boolean => { + const text = buildSearchText(model); + if (!IMAGE_GENERATION_KEYWORDS.some((keyword) => text.includes(keyword))) { + return false; + } + + return ( + !model.capabilities.tools && + !model.capabilities.function_calling && + !model.capabilities.json_mode + ); +}; + +const supportsVision = ( + model: EnhancedModelMetadata | undefined, + fallbackModelId?: string, +): boolean => { + if (model?.capabilities.vision) { + return true; + } + + if (!fallbackModelId) { + return false; + } + + return inferVisionCapability({ + modelId: fallbackModelId, + providerId: model?.provider_id, + family: model?.family, + description: model?.description, + }); +}; + +const capabilityScore = (model: EnhancedModelMetadata): number => { + let score = 0; + if (model.capabilities.tools) score += 5; + if (model.capabilities.function_calling) score += 4; + if (model.capabilities.json_mode) score += 3; + if (model.capabilities.reasoning) score += 2; + if (model.capabilities.streaming) score += 1; + return score; +}; + +function compareReleaseDateDesc( + left: EnhancedModelMetadata, + right: EnhancedModelMetadata, +): number { + if (left.release_date && right.release_date) { + return right.release_date.localeCompare(left.release_date); + } + if (left.release_date && !right.release_date) return -1; + if (!left.release_date && right.release_date) return 1; + return 0; +} + +export function resolveVisionModel( + params: ResolveVisionModelParams, +): VisionResolveResult { + const { currentModelId, models } = params; + const currentModel = findModelMeta(currentModelId, models); + + if (supportsVision(currentModel, currentModelId)) { + return { + targetModelId: currentModel?.id || currentModelId, + switched: false, + reason: "already_vision", + }; + } + + const currentFamily = normalize(currentModel?.family); + const candidates = models.filter( + (model) => model.capabilities.vision && !isLikelyImageGenerationModel(model), + ); + + if (candidates.length === 0) { + return { + targetModelId: currentModelId, + switched: false, + reason: "no_vision_model", + }; + } + + const sortedCandidates = [...candidates].sort((left, right) => { + const leftSameFamily = currentFamily.length > 0 && normalize(left.family) === currentFamily; + const rightSameFamily = + currentFamily.length > 0 && normalize(right.family) === currentFamily; + if (leftSameFamily !== rightSameFamily) { + return leftSameFamily ? -1 : 1; + } + + const capabilityDelta = capabilityScore(right) - capabilityScore(left); + if (capabilityDelta !== 0) { + return capabilityDelta; + } + + if (left.is_latest !== right.is_latest) { + return left.is_latest ? -1 : 1; + } + + const tierDelta = TIER_WEIGHT[right.tier] - TIER_WEIGHT[left.tier]; + if (tierDelta !== 0) { + return tierDelta; + } + + const releaseDelta = compareReleaseDateDesc(left, right); + if (releaseDelta !== 0) { + return releaseDelta; + } + + return left.id.localeCompare(right.id); + }); + + const target = sortedCandidates[0]; + const reason = + currentFamily.length > 0 && normalize(target.family) === currentFamily + ? "matched" + : "fallback_latest"; + + return { + targetModelId: target.id, + switched: normalize(target.id) !== normalize(currentModelId), + reason, + }; +} diff --git a/src/lib/tauri-mock/core.ts b/src/lib/tauri-mock/core.ts index 59840e2bf..26072bb90 100644 --- a/src/lib/tauri-mock/core.ts +++ b/src/lib/tauri-mock/core.ts @@ -724,6 +724,9 @@ const defaultMocks: Record = { whereCandidates: [], supplementalSearchDirs: ["/opt/homebrew/bin", "/usr/local/bin"], supplementalCommandCandidates: [], + gitWhereCandidates: [], + gitSupplementalSearchDirs: [], + gitSupplementalCommandCandidates: [], }, tempArtifacts: [], }), diff --git a/src/lib/workflow/threeStageWorkflow.ts b/src/lib/workflow/threeStageWorkflow.ts deleted file mode 100644 index 22493d9c5..000000000 --- a/src/lib/workflow/threeStageWorkflow.ts +++ /dev/null @@ -1,402 +0,0 @@ -/** - * 三阶段工作流管理器 - * - * 基于 planning-with-files 的核心机制,实现: - * - Pre-Action → Action → Post-Action 三阶段工作流 - * - 自动化上下文工程和错误学习 - * - 2-Action 规则和 3次错误协议 - */ - -import { ContextMemoryAPI } from "../api/contextMemory"; -import { ToolHooksAPI } from "../api/toolHooks"; - -export interface WorkflowPhase { - number: number; - name: string; - status: "pending" | "in_progress" | "complete"; - tasks: string[]; - notes?: string; -} - -export interface WorkflowConfig { - sessionId: string; - projectName: string; - goal: string; - phases: WorkflowPhase[]; -} - -export interface ActionContext { - sessionId: string; - actionType: string; - actionDescription: string; - toolName?: string; - toolParameters?: Record; - messageCount: number; -} - -/** - * 三阶段工作流管理器 - */ -export class ThreeStageWorkflowManager { - private sessionId: string; - private visualOperationCount: number = 0; - private errorAttempts: Map = new Map(); - - constructor(sessionId: string) { - this.sessionId = sessionId; - } - - /** - * 初始化工作流 - */ - async initializeWorkflow(config: WorkflowConfig): Promise { - // 触发会话开始钩子 - await ToolHooksAPI.triggerSessionStart(this.sessionId, { - project_name: config.projectName, - goal: config.goal, - }); - - // 创建任务计划 - const taskPlanContent = this.generateTaskPlanContent(config); - await ContextMemoryAPI.saveTaskPlan( - this.sessionId, - `任务计划: ${config.projectName}`, - taskPlanContent, - 5, - ); - - // 创建初始发现记录 - await ContextMemoryAPI.saveFinding( - this.sessionId, - "工作流初始化", - `三阶段工作流已初始化\n项目: ${config.projectName}\n目标: ${config.goal}`, - ["初始化", "工作流"], - 3, - ); - - // 记录初始进度 - await ContextMemoryAPI.logProgress( - this.sessionId, - "工作流启动", - `三阶段工作流已启动,共 ${config.phases.length} 个阶段`, - ); - } - - /** - * Pre-Action 阶段:执行操作前的上下文刷新 - */ - async preAction(context: ActionContext): Promise { - // 触发 Pre-Tool-Use 钩子 - await ToolHooksAPI.triggerPreToolUse( - context.sessionId, - context.toolName || context.actionType, - context.toolParameters || {}, - context.actionDescription, - context.messageCount, - ); - - // 获取当前记忆上下文 - const memoryContext = await ContextMemoryAPI.getMemoryContext( - context.sessionId, - ); - - // 检查是否应该避免该操作(3次错误协议) - const shouldAvoid = await ContextMemoryAPI.shouldAvoidOperation( - context.sessionId, - context.actionDescription, - ); - - if (shouldAvoid) { - const warning = `⚠️ 3次错误协议警告: 该操作已失败3次,建议更换方法\n操作: ${context.actionDescription}`; - - await ContextMemoryAPI.recordError({ - session_id: context.sessionId, - error_description: `重复失败操作: ${context.actionDescription}`, - attempted_solution: "触发3次错误协议,建议更换方法", - }); - - return `${warning}\n\n当前上下文:\n${memoryContext}`; - } - - // 记录上下文刷新 - await ContextMemoryAPI.logProgress( - context.sessionId, - "Pre-Action 上下文刷新", - `准备执行: ${context.actionDescription}`, - ); - - return `🔄 Pre-Action 上下文刷新完成\n\n准备执行: ${context.actionDescription}\n\n当前记忆上下文:\n${memoryContext}`; - } - - /** - * Action 阶段:执行实际操作 - */ - async executeAction( - context: ActionContext, - actionResult: string, - ): Promise { - // 记录操作执行 - await ContextMemoryAPI.logProgress( - context.sessionId, - `执行操作: ${context.actionType}`, - `操作描述: ${context.actionDescription}\n结果: ${actionResult.substring(0, 200)}${actionResult.length > 200 ? "..." : ""}`, - ); - - // 如果是视觉操作,增加计数 - if (this.isVisualOperation(context.actionType)) { - this.visualOperationCount++; - } - } - - /** - * Post-Action 阶段:操作后的状态更新 - */ - async postAction( - context: ActionContext, - actionResult: string, - error?: string, - ): Promise { - let message = "📝 Post-Action 状态更新:\n\n"; - - // 处理错误情况 - if (error) { - const errorKey = context.actionDescription; - const attemptCount = (this.errorAttempts.get(errorKey) || 0) + 1; - this.errorAttempts.set(errorKey, attemptCount); - - const { shouldAvoid } = await ContextMemoryAPI.recordErrorWithCheck( - context.sessionId, - error, - `尝试次数: ${attemptCount}`, - context.actionDescription, - ); - - message += `🚨 错误记录 (第${attemptCount}次尝试): ${error}\n`; - - if (shouldAvoid) { - message += `⚠️ 已达到3次错误限制,建议更换方法\n`; - } - } - - // 触发 Post-Tool-Use 钩子 - await ToolHooksAPI.triggerPostToolUse( - context.sessionId, - context.toolName || context.actionType, - actionResult, - context.toolParameters || {}, - context.actionDescription, - context.messageCount, - error, - ); - - // 应用 2-Action 规则 - if (this.visualOperationCount >= 2) { - await this.apply2ActionRule(actionResult); - message += `🎯 2-Action 规则已应用 (视觉操作计数: ${this.visualOperationCount})\n`; - this.visualOperationCount = 0; // 重置计数 - } - - // 提醒更新状态 - message += `\n💡 提醒:\n`; - message += `- 如果完成了某个阶段,请更新任务计划状态\n`; - message += `- 有新发现请记录到 findings.md\n`; - message += `- 重要进展请更新 progress.md\n`; - - return message; - } - - /** - * 应用 2-Action 规则 - */ - private async apply2ActionRule(actionResult: string): Promise { - const timestamp = new Date().toLocaleTimeString(); - const finding = `2-Action 规则触发 (${timestamp})\n\n最近操作结果:\n${actionResult.substring(0, 500)}${actionResult.length > 500 ? "..." : ""}`; - - await ContextMemoryAPI.apply2ActionRule(this.sessionId, finding); - } - - /** - * 更新阶段状态 - */ - async updatePhaseStatus( - phaseNumber: number, - status: "pending" | "in_progress" | "complete", - notes?: string, - ): Promise { - const statusText = { - pending: "待开始", - in_progress: "进行中", - complete: "已完成", - }[status]; - - await ContextMemoryAPI.saveTaskPlan( - this.sessionId, - `阶段 ${phaseNumber} 状态更新`, - `阶段 ${phaseNumber} 状态已更新为: ${statusText}${notes ? `\n备注: ${notes}` : ""}`, - 4, - ); - - await ContextMemoryAPI.logProgress( - this.sessionId, - `阶段 ${phaseNumber} 状态更新`, - `状态: ${statusText}${notes ? `\n备注: ${notes}` : ""}`, - ); - } - - /** - * 记录重要发现 - */ - async recordFinding( - title: string, - content: string, - tags: string[] = [], - ): Promise { - await ContextMemoryAPI.saveFinding( - this.sessionId, - title, - content, - ["发现", ...tags], - 4, - ); - } - - /** - * 记录决策 - */ - async recordDecision(decision: string, rationale: string): Promise { - await ContextMemoryAPI.saveFinding( - this.sessionId, - `决策: ${decision}`, - `决策内容: ${decision}\n\n决策理由:\n${rationale}`, - ["决策", "重要"], - 5, - ); - } - - /** - * 检查任务完成状态 - */ - async checkCompletion(): Promise<{ isComplete: boolean; summary: string }> { - const stats = await ContextMemoryAPI.getMemoryStats(this.sessionId); - const memories = await ContextMemoryAPI.getSessionMemories(this.sessionId); - - // 简单的完成度检查逻辑 - const taskPlanMemories = memories.filter( - (m) => m.file_type === "task_plan", - ); - const hasCompletedPhases = taskPlanMemories.some( - (m) => m.content.includes("已完成") || m.content.includes("complete"), - ); - - const summary = - `📊 任务完成状态检查:\n\n` + - `- 活跃记忆: ${stats.active_memories} 个\n` + - `- 未解决错误: ${stats.unresolved_errors} 个\n` + - `- 已解决错误: ${stats.resolved_errors} 个\n` + - `- 是否有已完成阶段: ${hasCompletedPhases ? "是" : "否"}\n\n` + - `${stats.unresolved_errors > 0 ? "⚠️ 仍有未解决的错误需要处理" : "✅ 无未解决错误"}`; - - return { - isComplete: hasCompletedPhases && stats.unresolved_errors === 0, - summary, - }; - } - - /** - * 结束工作流 - */ - async finalizeWorkflow(): Promise { - const { isComplete, summary } = await this.checkCompletion(); - - // 触发停止钩子 - await ToolHooksAPI.triggerStop(this.sessionId, 0, { - workflow_complete: isComplete.toString(), - }); - - // 保存会话摘要 - await ContextMemoryAPI.saveFinding( - this.sessionId, - "工作流会话摘要", - `三阶段工作流已结束\n\n${summary}`, - ["摘要", "会话结束"], - 5, - ); - - return `🎉 三阶段工作流已结束\n\n${summary}`; - } - - /** - * 生成任务计划内容 - */ - private generateTaskPlanContent(config: WorkflowConfig): string { - let content = `# 任务计划: ${config.projectName}\n\n`; - content += `## 目标\n${config.goal}\n\n`; - content += `## 当前阶段\n阶段 1\n\n`; - content += `## 阶段列表\n\n`; - - config.phases.forEach((phase) => { - content += `### 阶段 ${phase.number}: ${phase.name}\n`; - phase.tasks.forEach((task) => { - content += `- [ ] ${task}\n`; - }); - content += `- **状态**: ${phase.status}\n\n`; - }); - - content += `## 关键问题\n`; - content += `1. [需要回答的重要问题]\n`; - content += `2. [另一个关键问题]\n\n`; - - content += `## 已做决策\n`; - content += `| 决策 | 理由 |\n`; - content += `|------|------|\n`; - content += `| | |\n\n`; - - content += `## 遇到的错误\n`; - content += `| 错误 | 尝试次数 | 解决方案 |\n`; - content += `|------|----------|----------|\n`; - content += `| | 1 | |\n\n`; - - content += `## 注意事项\n`; - content += `- **2-Action 规则**: 每2次视觉操作后立即保存发现\n`; - content += `- **3次错误协议**: 永不重复相同的失败操作\n`; - content += `- **上下文刷新**: 重要决策前重新阅读计划文件\n`; - - return content; - } - - /** - * 判断是否为视觉操作 - */ - private isVisualOperation(actionType: string): boolean { - const visualActions = [ - "view", - "read", - "browse", - "search", - "screenshot", - "image", - ]; - return visualActions.some((action) => - actionType.toLowerCase().includes(action), - ); - } - - /** - * 获取会话统计 - */ - async getSessionStats(): Promise<{ - memoryStats: any; - visualOperationCount: number; - errorAttempts: Record; - }> { - const memoryStats = await ContextMemoryAPI.getMemoryStats(this.sessionId); - - return { - memoryStats, - visualOperationCount: this.visualOperationCount, - errorAttempts: Object.fromEntries(this.errorAttempts), - }; - } -} - -export default ThreeStageWorkflowManager; diff --git a/src/types/chat.ts b/src/types/chat.ts index 186efdf74..2de7ded47 100644 --- a/src/types/chat.ts +++ b/src/types/chat.ts @@ -359,74 +359,3 @@ export type StreamEvent = | DoneEvent | FinalDoneEvent | ErrorEvent; - -// ============================================================================ -// Hook 配置类型 -// ============================================================================ - -/** useUnifiedChat 配置选项 */ -export interface UseUnifiedChatOptions { - /** 对话模式 */ - mode: ChatMode; - /** 初始会话 ID(可选) */ - sessionId?: string; - /** 系统提示词(可选) */ - systemPrompt?: string; - /** Provider 类型(可选) */ - providerType?: string; - /** 模型名称(可选) */ - model?: string; - /** 画布内容更新回调 */ - onCanvasUpdate?: (path: string, content: string) => void; - /** 文件写入回调 */ - onWriteFile?: (content: string, fileName: string) => void; - /** Harness 配置 */ - harnessConfig?: Record; - /** Harness 事件回调 */ - onHarnessEvent?: (event: HarnessEventPayload) => void; - /** 产物更新回调 */ - onArtifactUpdate?: (artifact: HarnessArtifactSnapshot) => void; - /** 错误回调 */ - onError?: (error: ChatError) => void; -} - -/** useUnifiedChat 返回值 */ -export interface UseUnifiedChatReturn { - // 状态 - /** 当前会话 */ - session: ChatSession | null; - /** 消息列表 */ - messages: ChatMessage[]; - /** 是否正在加载 */ - isLoading: boolean; - /** 是否正在发送 */ - isSending: boolean; - /** 错误信息 */ - error: ChatError | null; - - // 会话操作 - /** 创建新会话 */ - createSession: (options?: Partial) => Promise; - /** 加载会话 */ - loadSession: (sessionId: string) => Promise; - /** 删除会话 */ - deleteSession: (sessionId?: string) => Promise; - /** 重命名会话 */ - renameSession: (title: string, sessionId?: string) => Promise; - - // 消息操作 - /** 发送消息 */ - sendMessage: ( - content: string, - images?: ImageInput[], - webSearch?: boolean, - ) => Promise; - /** 停止生成 */ - stopGeneration: () => Promise; - /** 清空消息 */ - clearMessages: () => void; - - // Provider 配置 - /** 配置 Provider */ - configureProvider: (providerType: string, model: string) => Promise; -}