diff --git a/RELEASE_NOTES.md b/RELEASE_NOTES.md index dffe9d890..187e56ba8 100644 --- a/RELEASE_NOTES.md +++ b/RELEASE_NOTES.md @@ -1,42 +1,27 @@ -## ProxyCast v0.84.0 +## ProxyCast v0.85.0 ### ✨ 新功能 -- 新增 API 网关层架构,将 useTauri 聚合层拆分为独立的 API 模块(appConfig、serverRuntime、logs、experimentalFeatures、channelsRuntime 等) -- 新增 OpenClaw 安装与运行时集成(openclaw_install、OpenClaw 配置/安装/运行页面) -- 新增环境变量管理服务(environment_service),支持 Shell 导入预览与环境变量覆盖 -- 新增 Harness 状态面板,实时展示 Agent 运行状态 -- 新增 Aster Agent 执行策略与 Web 搜索集成,大幅扩展 Aster 命令能力 -- 新增 General Chat 统一消息桥接层(bridge.ts),支持跨模块消息同步 -- 新增 Poster 主题系统(themes/poster) -- 新增 Agent 流式传输运行时(agentStream、agentRuntime、agentCompat) -- 新增持久化记忆文件系统(durable_memory_fs)与工具 IO 卸载(tool_io_offload) -- 新增 CI 工作流配置(.github/workflows/ci.yml) -- 新增应用更新检测 API(appUpdate) -- 新增 Sub-Agent 调度器测试覆盖 -- 新增 Skill 模型层与技能服务增强 - -### 🐛 修复 -- 修复 Web 搜索运行时 priority 列表包含无效引擎的问题 -- 修复 ESLint 导入限制违规:将受限导入从 useTauri 迁移到专用 API 模块 -- 修复 SkillsPage 导出非组件函数导致 Fast Refresh 失效的问题 -- 修复 OpenClaw 安装候选路径类型复杂度 clippy 警告 +- 集成 AI 摘要到 SessionContextService,实现上下文智能管理 (621ab3ee) +- 新增 AI 摘要服务,用于上下文管理的 P0 阶段 1 实现 (75309bd2) +- 新增 Agent Timeline 服务,支持时间线视图 +- 新增 Chat History 服务,统一聊天历史管理 +- 新增多个 Agent 聊天相关组件(AgentPlanBlock、AgentRuntimeStrip、AgentThreadTimeline、SocialMediaHarnessCard) +- 新增社交媒体 Harness 工具集成 ### 🔧 优化与重构 -- 重构 General Chat 命令层,统一消息处理流程(+1200 行) -- 重构 Aster Agent 命令层,增强执行策略与自动续写能力(+950 行) -- 重构 Agent 会话存储,支持持久化与恢复 -- 重构事件转换器,增强流式事件处理 -- 重构设置页面 v2 多个子模块(channels、developer、experimental、environment) -- 重构终端 AI 集成与控制器 -- 优化 ESLint 配置,新增命令调用与导入来源限制规则 -- 优化 Skill 服务与默认技能注册 -- 优化 DevBridge 调度器,增强浏览器开发模式兼容性 +- 移除 general-chat 相关的遗留代码和组件,完成向统一 Agent 系统的迁移 +- 清理 compat 兼容层代码(agentCompat、generalChatCompat) +- 重构数据库迁移结构,新增 migration_support 和 startup_migrations +- 优化 Agent 聊天 Hooks 架构,拆分为多个专职模块(agentChatActionState、agentChatCoreUtils、agentChatHistory 等) +- 完善测试覆盖率,新增 40+ 单元测试文件 +- 优化 Artifact 渲染器,新增 DocumentRenderer +- 优化内容创作工作流,新增社交媒体 Harness 测试 ### 📦 其他 -- 更新 Cargo 依赖锁文件 -- 更新 AI 提示词文档(aster-integration、content-creator、governance) -- 更新 AI Agent 开发指南 +- 更新多个 AI 提示词文档 +- 新增 report-legacy-surfaces.mjs 脚本 +- 更新 ESLint 配置 --- -**完整变更**: v0.83.2...v0.84.0 +**完整变更**: v0.84.0...v0.85.0 diff --git a/docs/README.md b/docs/README.md index 8884b5a6a..ea19815f3 100644 --- a/docs/README.md +++ b/docs/README.md @@ -16,6 +16,7 @@ - `develop/`:开发流程与协作规范 - `plugins/`:插件与扩展相关文档 - `tests/`:测试策略与用例文档 +- `iteration-notes/`:迭代备忘与下版本建议(暂不进入当前发布范围的问题) - `images/`:文档图片资源 - `TECH_SPEC.md`:技术规格文档 - `develop/execution-tracker-technical-plan.md`:统一执行追踪(Execution Tracker)专项技术规划 diff --git a/docs/TECH_SPEC.md b/docs/TECH_SPEC.md index a9fa1c35f..bf6da9e8d 100644 --- a/docs/TECH_SPEC.md +++ b/docs/TECH_SPEC.md @@ -8,13 +8,13 @@ ProxyCast 是一个桌面端 AI API 代理工具,将各种大模型客户端 A 基于 pubcast 项目技术栈: -| 类别 | 技术 | -|------|------| -| 框架 | Tauri 2.0 (Rust + Web) | -| 前端 | React 18 + TypeScript | -| 构建 | Vite 5 | -| UI | Tailwind CSS + Radix UI | -| 图标 | Lucide React | +| 类别 | 技术 | +| ---- | ----------------------- | +| 框架 | Tauri 2.0 (Rust + Web) | +| 前端 | React 18 + TypeScript | +| 构建 | Vite 5 | +| UI | Tailwind CSS + Radix UI | +| 图标 | Lucide React | ### 核心依赖 @@ -32,15 +32,15 @@ ProxyCast 是一个桌面端 AI API 代理工具,将各种大模型客户端 A 参考 AIClient-2-API,需支持以下渠道: -| Provider | 协议 | 说明 | -|----------|------|------| -| `claude-kiro-oauth` | OpenAI/Claude | Kiro OAuth 访问 Claude Sonnet 4.5 | -| `gemini-cli-oauth` | OpenAI/Claude/Gemini | Gemini CLI OAuth | -| `openai-qwen-oauth` | OpenAI/Claude | 通义千问 OAuth | -| `openai-custom` | OpenAI | 自定义 OpenAI 兼容 API | -| `claude-custom` | Claude | 自定义 Claude API | -| `gemini-antigravity` | Gemini | Antigravity 协议 | -| `openaiResponses-custom` | OpenAI Responses | 结构化对话 | +| Provider | 协议 | 说明 | +| ------------------------ | -------------------- | --------------------------------- | +| `claude-kiro-oauth` | OpenAI/Claude | Kiro OAuth 访问 Claude Sonnet 4.5 | +| `gemini-cli-oauth` | OpenAI/Claude/Gemini | Gemini CLI OAuth | +| `openai-qwen-oauth` | OpenAI/Claude | 通义千问 OAuth | +| `openai-custom` | OpenAI | 自定义 OpenAI 兼容 API | +| `claude-custom` | Claude | 自定义 Claude API | +| `gemini-antigravity` | Gemini | Antigravity 协议 | +| `openaiResponses-custom` | OpenAI Responses | 结构化对话 | ## 核心功能模块 @@ -74,9 +74,9 @@ src/ │ ├── TokenManager.tsx # Token 管理 │ ├── LogViewer.tsx # 日志查看 │ └── ui/ # 通用 UI 组件 -├── hooks/ -│ └── useTauri.ts # Tauri API hooks +├── hooks/ # 领域 Hook(不再承载统一 Tauri 聚合层) └── lib/ + ├── api/ # 前端 API 网关 └── utils.ts ``` @@ -91,12 +91,12 @@ http://localhost:8999/{provider}/v1/messages ### 支持的端点 -| 端点 | 协议 | 说明 | -|------|------|------| -| `/v1/chat/completions` | OpenAI | 聊天补全 | -| `/v1/messages` | Claude | Anthropic 消息 | -| `/v1/models` | OpenAI | 模型列表 | -| `/health` | - | 健康检查 | +| 端点 | 协议 | 说明 | +| ---------------------- | ------ | -------------- | +| `/v1/chat/completions` | OpenAI | 聊天补全 | +| `/v1/messages` | Claude | Anthropic 消息 | +| `/v1/models` | OpenAI | 模型列表 | +| `/health` | - | 健康检查 | ## 配置文件结构 @@ -139,12 +139,12 @@ http://localhost:8999/{provider}/v1/messages ## Token 凭证路径 -| 服务 | 默认路径 | -|------|----------| -| Kiro | `~/.aws/sso/cache/kiro-auth-token.json` | -| Gemini | `~/.gemini/oauth_creds.json` | -| Qwen | `~/.qwen/oauth_creds.json` | -| Antigravity | `~/.antigravity/oauth_creds.json` | +| 服务 | 默认路径 | +| ----------- | --------------------------------------- | +| Kiro | `~/.aws/sso/cache/kiro-auth-token.json` | +| Gemini | `~/.gemini/oauth_creds.json` | +| Qwen | `~/.qwen/oauth_creds.json` | +| Antigravity | `~/.antigravity/oauth_creds.json` | ## 协议转换 @@ -156,13 +156,13 @@ OpenAI <---> Claude <---> Gemini ### 转换矩阵 -| 输入协议 | 输出 Provider | 说明 | -|----------|---------------|------| -| OpenAI | kiro | OpenAI -> CodeWhisperer | -| OpenAI | gemini | OpenAI -> Gemini | -| Claude | kiro | Claude -> CodeWhisperer | -| Claude | gemini | Claude -> Gemini | -| Claude | openai | Claude -> OpenAI | +| 输入协议 | 输出 Provider | 说明 | +| -------- | ------------- | ----------------------- | +| OpenAI | kiro | OpenAI -> CodeWhisperer | +| OpenAI | gemini | OpenAI -> Gemini | +| Claude | kiro | Claude -> CodeWhisperer | +| Claude | gemini | Claude -> Gemini | +| Claude | openai | Claude -> OpenAI | ## UI 功能 @@ -175,21 +175,25 @@ OpenAI <---> Claude <---> Gemini ## 开发计划 ### Phase 1: 基础框架 + - [ ] Tauri 项目初始化 - [ ] 基础 UI 布局 - [ ] 配置管理 ### Phase 2: Kiro Provider + - [ ] Kiro OAuth Token 读取 - [ ] CodeWhisperer API 调用 - [ ] OpenAI/Claude 协议支持 ### Phase 3: 其他 Provider + - [ ] Gemini CLI OAuth - [ ] Qwen OAuth - [ ] OpenAI/Claude Custom ### Phase 4: 高级功能 + - [ ] Provider Pool 管理 - [ ] 自动 Token 刷新 - [ ] 请求日志/统计 diff --git a/docs/aiprompts/README.md b/docs/aiprompts/README.md index 9ed437758..063c29321 100644 --- a/docs/aiprompts/README.md +++ b/docs/aiprompts/README.md @@ -10,6 +10,7 @@ AI Agent 专用文档目录,提供模块级别的详细说明。 ## 文件索引 ### 核心系统 + - `overview.md` - 项目架构概览 - `governance.md` - **治理第一原则**(新旧并存、迁移收口、禁止回流) - `providers.md` - Provider 系统(OAuth/API Key 认证) @@ -18,26 +19,31 @@ AI Agent 专用文档目录,提供模块级别的详细说明。 - `server.md` - HTTP 服务器(API 端点) ### 前端模块 + - `components.md` - React 组件系统 - `hooks.md` - 自定义 React Hooks - `lib.md` - 工具库和 API 封装 ### 后端模块 + - `services.md` - 业务服务层 - `commands.md` - Tauri 命令 - `database.md` - 数据库层(SQLite) ### 功能模块 + - `terminal.md` - 内置终端 - `mcp.md` - MCP 服务器管理 - `plugins.md` - 插件系统 - `playwright-e2e.md` - Playwright MCP 续测与 E2E 指南 ### Aster 集成 + - `aster-integration.md` - **Aster 框架集成方案** - `workspace.md` - **Workspace 设计文档**(工作目录管理) ### 内容创作 + - `content-creator.md` - **内容创作系统**(write_file 标签、画布联动) ## 使用方式 @@ -50,6 +56,7 @@ AI Agent 在处理特定模块时,应先阅读对应的 aiprompts 文档: # 处理新旧并存、迁移、重构、架构收口 → 先读 docs/aiprompts/governance.md +→ 再执行 npm run governance:legacy-report # 处理凭证池相关任务 → 先读 docs/aiprompts/credential-pool.md diff --git a/docs/aiprompts/aster-integration.md b/docs/aiprompts/aster-integration.md index 67b65001f..18d0da9f4 100644 --- a/docs/aiprompts/aster-integration.md +++ b/docs/aiprompts/aster-integration.md @@ -82,7 +82,7 @@ ProxyCast 已完整集成 aster-rust 框架,包括凭证池桥接。 ### 使用方式 -> 治理约定:前端业务层不要直接 `invoke('aster_*')`,统一通过 `src/lib/api/agentRuntime.ts` 调用现役 Aster API。 +> 治理约定:前端业务层不要直接 `invoke('aster_*')`,统一通过 `src/lib/api/agentRuntime.ts` 调用现役 Aster API。历史 `src/lib/api/agentCompat.ts` 已删除。 ```typescript import { diff --git a/docs/aiprompts/commands.md b/docs/aiprompts/commands.md index dfc8583d1..747f70128 100644 --- a/docs/aiprompts/commands.md +++ b/docs/aiprompts/commands.md @@ -2,7 +2,50 @@ ## 概述 -Tauri 命令是前端与 Rust 后端通信的桥梁,通过 `invoke` 调用。 +Tauri 命令是前端与 Rust 后端通信的边界,但前端业务代码**不应直接散落 `invoke`**。 + +推荐路径是: + +`组件 / Hook -> src/lib/api/* 网关 -> safeInvoke -> Rust command` + +这样做的目的不是“多包一层”,而是确保: + +- 前端只有一个可治理的调用出口 +- Rust 命令可以按 `current / compat / deprecated` 分类演进 +- 新旧命令并存时,迁移边界清晰,不会继续扩散 + +## 治理约束 + +- 新的前端功能,禁止在页面、组件、普通 Hook 中直接调用 `invoke`。 +- 新的 Rust 命令,必须同时落一个对应的 `src/lib/api/*` 网关文件或收口到现有网关。 +- 旧命令如果暂时不能删,必须明确标记为 `compat` 或 `deprecated`,只允许保兼容,不允许继续长新逻辑。 +- 当前端已经迁到新网关后,要继续用 ESLint、脚本或日志告警封住旧入口,避免 AI 回流。 + +## 当前事实源 + +- 聊天主命令:`chat_*` +- 旧 `general_chat_*` 前端 compat 网关与 Rust 命令已删除 +- 当前剩余治理重点:统计、记忆等旁路仍在读取 `general_chat_*` 历史表 + +## 治理案例:记忆系统 + +以当前仓库里的记忆能力为例: + +- `unified_memory_*`:现役统一记忆主链路,后续功能优先往这里收 +- `memory_runtime_*`:现役 runtime / 上下文记忆主入口 +- `memory_get_*` / `memory_toggle_auto`:当前仍在使用的治理配置入口 +- `switch_prompt`:旧 prompt 切换命令已移除,统一使用 `enable_prompt` +- `get_legacy_api_key_credentials` 等迁移命令:前端与 Tauri 入口都已移除,避免 UI/AI 再接入历史迁移链路 + +这类场景下,AI 不应该再做一套“第三套记忆命令”,而应该: + +1. 先判断当前需求属于主链路、兼容层,还是治理配置 +2. 如果是统一沉淀记忆,优先补到 `unified_memory_*` +3. 如果是 runtime / 上下文记忆视图,优先补到 `memory_runtime_*` +4. 如果存在旧命令又无任何调用,就直接删掉命令注册、桥接和 mock,不要继续保留空兼容壳 + +同理,对话系统也不应该重新引回已经删除的 `general_chat_*` 命令; +后续如需扩展聊天能力,应继续收敛到 `chat_*` 与对应网关。 ## 目录结构 @@ -73,22 +116,23 @@ async fn clear_flow_records(before: Option) -> Result; ## 前端调用 +推荐写法不是在业务层直接 `invoke`,而是在 API 网关里集中调用: + ```typescript -import { invoke } from '@tauri-apps/api/core'; +// src/lib/api/serverRuntime.ts +import { safeInvoke } from "@/lib/dev-bridge"; -// 添加凭证 -const credential = await invoke('add_credential', { - provider: 'kiro', - filePath: '/path/to/credential.json', -}); +export async function getServerStatus() { + return safeInvoke("get_server_status"); +} +``` -// 获取服务器状态 -const status = await invoke('get_server_status'); +业务层只消费 API 网关: -// 查询流量记录 -const records = await invoke>('get_flow_records', { - query: { page: 1, pageSize: 20 }, -}); +```typescript +import { getServerStatus } from "@/lib/api/serverRuntime"; + +const status = await getServerStatus(); ``` ## 错误处理 diff --git a/docs/aiprompts/components.md b/docs/aiprompts/components.md index 85d5457b1..5d6f9fa8d 100644 --- a/docs/aiprompts/components.md +++ b/docs/aiprompts/components.md @@ -11,7 +11,7 @@ src/components/ ├── ui/ # 基础 UI 组件 (shadcn/ui) ├── provider-pool/ # 凭证池管理 ├── flow-monitor/ # 流量监控 -├── general-chat/ # 通用对话 +├── general-chat/ # 兼容画布桥接(非对话主入口) ├── terminal/ # 内置终端 ├── mcp/ # MCP 服务器 ├── settings/ # 设置页面 @@ -27,15 +27,15 @@ src/components/ ```tsx // src/components/AppSidebar.tsx export function AppSidebar() { - return ( - - ); + return ( + + ); } ``` @@ -46,14 +46,14 @@ export function AppSidebar() { ```tsx // src/components/provider-pool/ProviderPoolPanel.tsx export function ProviderPoolPanel() { - const { credentials, addCredential, removeCredential } = useProviderPool(); - - return ( -
- - -
- ); + const { credentials, addCredential, removeCredential } = useProviderPool(); + + return ( +
+ + +
+ ); } ``` @@ -64,15 +64,15 @@ export function ProviderPoolPanel() { ```tsx // src/components/flow-monitor/FlowMonitorPanel.tsx export function FlowMonitorPanel() { - const { records, stats, query } = useFlowMonitor(); - - return ( -
- - - -
- ); + const { records, stats, query } = useFlowMonitor(); + + return ( +
+ + + +
+ ); } ``` @@ -89,22 +89,18 @@ export function FlowMonitorPanel() { ```tsx // 标准组件结构 interface Props { - // props 定义 + // props 定义 } export function ComponentName({ prop1, prop2 }: Props) { - // hooks - const [state, setState] = useState(); - - // handlers - const handleClick = () => {}; - - // render - return ( -
- {/* JSX */} -
- ); + // hooks + const [state, setState] = useState(); + + // handlers + const handleClick = () => {}; + + // render + return
{/* JSX */}
; } ``` diff --git a/docs/aiprompts/content-creator.md b/docs/aiprompts/content-creator.md index 600e0975c..461c5c05f 100644 --- a/docs/aiprompts/content-creator.md +++ b/docs/aiprompts/content-creator.md @@ -17,7 +17,7 @@ AI 返回带 标签的响应 ↓ StreamingRenderer 解析标签 → 调用 onWriteFile ↓ -AgentChatPage.handleWriteFile → 更新画布状态 +AgentChatPage.handleWriteFile → 映射为社媒 harness 产物 / 版本链 ↓ 右侧画布自动打开,显示文档内容 ``` @@ -44,8 +44,10 @@ src/components/ │ │ └── MessageList.tsx # 消息列表 │ └── index.tsx # AgentChatPage └── general-chat/ - └── store/ - └── useGeneralChatStore.ts # 通用对话 Store + ├── bridge.ts # 兼容桥接层 + ├── canvas/ + │ └── CanvasPanel.tsx # 复用画布面板 + └── types.ts # CanvasState / DEFAULT_CANVAS_STATE ``` ## 核心组件 @@ -217,6 +219,10 @@ const handleWriteFile = useCallback( ## 注意事项 +- 社媒主题已不再把 `write_file` 仅视为“文件覆盖”,而是映射为带阶段语义的版本链产物 +- `brief / draft / polished / platform variant / publish package` 应分别作为不同产物语义处理 +- 日志、运行轨迹、正文产物三层分离:`harness` 产生命名事件,日志只做投影,正文仍由画布/产物承载 + ### Aster 框架限制 Aster 框架的 `SessionConfig` 不支持 session 级别的 system prompt,因此采用**消息注入**方案: diff --git a/docs/aiprompts/governance.md b/docs/aiprompts/governance.md index b1fa34042..ac31748cf 100644 --- a/docs/aiprompts/governance.md +++ b/docs/aiprompts/governance.md @@ -83,6 +83,32 @@ - CI 阻止新代码继续引用废弃路径 - 脚本扫描旧表、旧 DAO、旧命令、旧 Hook 的新增使用点 +当前仓库可直接运行: + +```bash +npm run governance:legacy-report +``` + +它会扫描: + +- 已经被判定为 `deprecated` / `dead-candidate` 的前端入口(按真实 import 解析 `@/` 与相对路径) +- 旧 Tauri 命令是否仍然只收口在指定 API 网关 +- 哪些兼容壳层已经零引用,可以进入删除候选 + +当前项目的最新治理状态可以概括为: + +- 旧 `components/chat`、`general-chat` 页面 / Hook / Store / compat API 已删除 +- Rust `general_chat_*` 兼容命令已删除 +- `src/lib/api/agentCompat.ts` 已删除,Aster 前端只保留现役 runtime / stream API +- General Chat 历史数据迁移已接入数据库初始化;新治理优先推动“启动期迁移”,而不是长期保留运行时 fallback +- 下一阶段重点不再是页面和命令,而是统计、记忆等旁路对 `general_chat_*` 历史表的依赖 + +进一步治理时,建议坚持一个更细的边界规则: + +- “迁移是否完成”的判断收口在 `Repository / Database` 边界 +- 业务服务层只消费 `pending_*` 语义接口 +- 不要在多个 service 里重复写 `is_migrated` 分支 + 原则只有一句: **不是鼓励走新路,而是封住老路。** @@ -184,6 +210,8 @@ 3. 必须显式说明当前改动属于 `current`、`compat`、`deprecated`、`dead` 中哪一类。 4. 如果发现主链路与旁路系统割裂,必须指出,不得假装治理已经完成。 5. 如果无法在本次改动中完成收口,至少要建立守卫,阻止问题继续扩散。 +6. 一旦历史数据迁移已接入启动流程,运行时必须按“迁移完成标记”短路旧表读取;旧表只允许服务迁移、审计与回放,不再参与主链路查询。 +7. 过渡期对外暴露的命名必须体现“迁移态”语义,例如 `pending_*`,不要继续让业务层直接看见 `legacy_*` 模块名与函数名。 ## 一句话总结 diff --git a/docs/aiprompts/hooks.md b/docs/aiprompts/hooks.md index 4de945580..b86ef5c17 100644 --- a/docs/aiprompts/hooks.md +++ b/docs/aiprompts/hooks.md @@ -2,7 +2,7 @@ ## 概述 -自定义 Hooks 封装业务逻辑,通过 Tauri invoke 与后端通信。 +自定义 Hooks 封装业务逻辑;新代码应优先通过 `src/lib/api/*` 网关与后端通信,而不是直接在 Hook 中散落 `invoke`。 ## 目录结构 @@ -15,10 +15,15 @@ src/hooks/ ├── useFlowEvents.ts # 流量事件 ├── useMcpServers.ts # MCP 服务器 ├── useDeepLink.ts # Deep Link 处理 -├── useSound.ts # 音效管理 -└── useTauri.ts # Tauri 通用 +└── useSound.ts # 音效管理 ``` +## 治理约束 + +- 新的前端能力优先落在 `src/lib/api/*`,再由 Hook 或组件消费。 +- 历史 `useTauri.ts` 兼容聚合层已删除,不要重新引入新的“大一统 API Hook”。 +- 旧聊天链路优先迁移到 `@/hooks/useUnifiedChat`,不要继续扩散 `useChat` / compat Hook。 + ## 核心 Hooks ### useUnifiedChat(统一对话) @@ -39,8 +44,22 @@ const { messages, sendMessage, stopGeneration } = useUnifiedChat({ const creatorChat = useUnifiedChat({ mode: "creator", systemPrompt: "你是内容创作助手...", - onCanvasUpdate: (path, content) => { /* 更新画布 */ }, - onWriteFile: (content, fileName) => { /* 文件写入 */ }, + harnessConfig: { + theme: "social-media", + artifactMode: "version-chain", + }, + onHarnessEvent: (event) => { + /* 接收阶段推进、产物创建等语义事件 */ + }, + onArtifactUpdate: (artifact) => { + /* 接收产物快照 */ + }, + onCanvasUpdate: (path, content) => { + /* 更新画布 */ + }, + onWriteFile: (content, fileName) => { + /* 文件写入 */ + }, }); // General 模式 - 纯文本对话 @@ -48,6 +67,7 @@ const generalChat = useUnifiedChat({ mode: "general" }); ``` **返回值**: + - `session` - 当前会话 - `messages` - 消息列表 - `isLoading` / `isSending` - 状态 @@ -55,7 +75,13 @@ const generalChat = useUnifiedChat({ mode: "general" }); - `sendMessage()` / `stopGeneration()` - 消息操作 - `configureProvider()` - Provider 配置 +**补充说明**: + +- Creator 模式现在支持 `harnessConfig`、`onHarnessEvent`、`onArtifactUpdate` +- 社媒内容推荐把 `` 结果投影为“版本链产物”,而不是只按文件名覆盖 + **相关文件**: + - 类型定义:`src/types/chat.ts` - API 封装:`src/lib/api/unified-chat.ts` - 架构文档:`docs/prd/chat-architecture-redesign.md` @@ -64,29 +90,31 @@ const generalChat = useUnifiedChat({ mode: "general" }); ```typescript export function useProviderPool() { - const [credentials, setCredentials] = useState([]); - const [loading, setLoading] = useState(false); - - const refresh = async () => { - setLoading(true); - const list = await invoke('list_credentials'); - setCredentials(list); - setLoading(false); - }; - - const addCredential = async (provider: string, path: string) => { - await invoke('add_credential', { provider, filePath: path }); - await refresh(); - }; - - const removeCredential = async (id: string) => { - await invoke('remove_credential', { id }); - await refresh(); - }; - - useEffect(() => { refresh(); }, []); - - return { credentials, loading, addCredential, removeCredential, refresh }; + const [credentials, setCredentials] = useState([]); + const [loading, setLoading] = useState(false); + + const refresh = async () => { + setLoading(true); + const list = await invoke("list_credentials"); + setCredentials(list); + setLoading(false); + }; + + const addCredential = async (provider: string, path: string) => { + await invoke("add_credential", { provider, filePath: path }); + await refresh(); + }; + + const removeCredential = async (id: string) => { + await invoke("remove_credential", { id }); + await refresh(); + }; + + useEffect(() => { + refresh(); + }, []); + + return { credentials, loading, addCredential, removeCredential, refresh }; } ``` @@ -94,17 +122,19 @@ export function useProviderPool() { ```typescript export function useFlowEvents() { - const [records, setRecords] = useState([]); - - useEffect(() => { - const unlisten = listen('flow-event', (event) => { - setRecords(prev => [event.payload.data, ...prev].slice(0, 100)); - }); - - return () => { unlisten.then(fn => fn()); }; - }, []); - - return { records }; + const [records, setRecords] = useState([]); + + useEffect(() => { + const unlisten = listen("flow-event", (event) => { + setRecords((prev) => [event.payload.data, ...prev].slice(0, 100)); + }); + + return () => { + unlisten.then((fn) => fn()); + }; + }, []); + + return { records }; } ``` @@ -112,17 +142,19 @@ export function useFlowEvents() { ```typescript export function useDeepLink() { - useEffect(() => { - const unlisten = listen('deep-link', async (event) => { - const url = new URL(event.payload); - - if (url.pathname === '/oauth/callback') { - await handleOAuthCallback(url.searchParams); - } - }); - - return () => { unlisten.then(fn => fn()); }; - }, []); + useEffect(() => { + const unlisten = listen("deep-link", async (event) => { + const url = new URL(event.payload); + + if (url.pathname === "/oauth/callback") { + await handleOAuthCallback(url.searchParams); + } + }); + + return () => { + unlisten.then((fn) => fn()); + }; + }, []); } ``` @@ -138,15 +170,15 @@ export function useDeepLink() { ```typescript // 返回对象,包含状态和操作 return { - // 状态 - data, - loading, - error, - - // 操作 - refresh, - add, - remove, + // 状态 + data, + loading, + error, + + // 操作 + refresh, + add, + remove, }; ``` diff --git a/docs/aiprompts/services.md b/docs/aiprompts/services.md index 203e8a590..22fcdac6b 100644 --- a/docs/aiprompts/services.md +++ b/docs/aiprompts/services.md @@ -17,11 +17,14 @@ src-tauri/src/services/ ├── usage_service.rs # 使用量统计 ├── backup_service.rs # 备份服务 ├── update_check_service.rs # 自动更新检查 -└── general_chat/ # 通用对话服务 ``` ## 核心服务 +> 注意:`general_chat/` 兼容壳已删除。 +> 新功能与新治理都应直接落到 unified chat / `chat_*` 体系,不要重新引回旧入口。 +> `ProviderPoolService::select_credential_with_fallback_legacy` 也已删除,凭证选择统一走现役 `select_credential_with_fallback`。 + ### ProviderPoolService ```rust diff --git a/eslint.config.js b/eslint.config.js index bd1bc75a0..3c567a6f5 100644 --- a/eslint.config.js +++ b/eslint.config.js @@ -85,6 +85,11 @@ const generalChatRestrictedPaths = [ message: "generalChatCompat 属于兼容网关,请仅在 general-chat store 中消费,避免 compat 逻辑再次向业务层扩散。", }, + { + name: "@/lib/api/compat", + message: + "api/compat.ts 属于历史记忆兼容层,请不要在新代码中接入;旧记忆运行时请使用 @/lib/api/memoryRuntime,统一记忆请使用 @/lib/api/unifiedMemory。", + }, { name: "@/lib/api/agent", message: @@ -94,6 +99,11 @@ const generalChatRestrictedPaths = [ name: "@/lib/terminal-api", message: "terminal-api 现在只是兼容门面;新代码请改用 @/lib/api/terminal。", }, + { + name: "@/hooks/useTauri", + message: + "useTauri 现在只是兼容聚合层,禁止新增依赖;请直接接入对应的 @/lib/api/* 网关。", + }, { name: "@/hooks/useTauri", importNames: [ @@ -370,21 +380,6 @@ const generalChatRestrictedPaths = [ }, ]; -const generalChatRestrictedPathsWithoutPage = generalChatRestrictedPaths.filter( - (entry) => - !( - entry.name === "@/components/general-chat" && - Array.isArray(entry.importNames) && - entry.importNames.includes("GeneralChatPage") - ), -); - -const generalChatRestrictedPathsWithoutPageAndStore = - generalChatRestrictedPathsWithoutPage.filter( - (entry) => - entry.name !== "@/components/general-chat/store/useGeneralChatStore", - ); - const generalChatRestrictedPathsWithoutCompatApi = generalChatRestrictedPaths.filter( (entry) => entry.name !== "@/lib/api/generalChatCompat", @@ -750,6 +745,152 @@ const projectMemoryCommandSelectors = [ "项目记忆 CRUD 相关后端命令请统一通过 `src/lib/api/memory.ts` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", })); +const toolHooksCommandSelectors = [ + "execute_hooks", + "add_hook_rule", + "remove_hook_rule", + "toggle_hook_rule", + "get_hook_rules", + "get_hook_execution_stats", + "clear_hook_execution_stats", +].map((command) => ({ + selector: `CallExpression[callee.name='safeInvoke'][arguments.0.value='${command}'], CallExpression[callee.name='invoke'][arguments.0.value='${command}']`, + message: + "工具钩子相关命令请统一通过 `src/lib/api/toolHooks.ts` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", +})); + +const a2uiFormCommandSelectors = [ + "create_a2ui_form", + "get_a2ui_form", + "get_a2ui_forms_by_message", + "get_a2ui_forms_by_session", + "save_a2ui_form_data", + "submit_a2ui_form", + "delete_a2ui_form", +].map((command) => ({ + selector: `CallExpression[callee.name='safeInvoke'][arguments.0.value='${command}'], CallExpression[callee.name='invoke'][arguments.0.value='${command}']`, + message: + "A2UI 表单持久化命令请统一通过 `src/lib/api/a2uiForm.ts` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", +})); + +const memoryFeedbackCommandSelectors = [ + "unified_memory_feedback", + "get_memory_feedback_stats", +].map((command) => ({ + selector: `CallExpression[callee.name='safeInvoke'][arguments.0.value='${command}'], CallExpression[callee.name='invoke'][arguments.0.value='${command}']`, + message: + "记忆反馈相关命令请统一通过 `src/lib/api/memoryFeedback.ts` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", +})); + +const contextMemoryCommandSelectors = [ + "save_memory_entry", + "get_session_memories", + "get_memory_context", + "record_error", + "should_avoid_operation", + "mark_error_resolved", + "get_memory_stats", + "cleanup_expired_memories", +].map((command) => ({ + selector: `CallExpression[callee.name='safeInvoke'][arguments.0.value='${command}'], CallExpression[callee.name='invoke'][arguments.0.value='${command}']`, + message: + "上下文记忆命令请统一通过 `src/lib/api/contextMemory.ts` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", +})); + +const asrProviderCommandSelectors = [ + "list_audio_devices", + "get_asr_credentials", + "add_asr_credential", + "update_asr_credential", + "delete_asr_credential", + "set_default_asr_credential", + "test_asr_credential", + "get_voice_input_config", + "save_voice_input_config", + "get_voice_instructions", + "save_voice_instruction", + "delete_voice_instruction", + "transcribe_audio", + "polish_voice_text", + "open_voice_window", + "close_voice_window", + "output_voice_text", + "start_recording", + "stop_recording", + "cancel_recording", + "get_recording_status", + "open_input_with_text", +].map((command) => ({ + selector: `CallExpression[callee.name='safeInvoke'][arguments.0.value='${command}'], CallExpression[callee.name='invoke'][arguments.0.value='${command}']`, + message: + "语音输入/ASR 相关命令请统一通过 `src/lib/api/asrProvider.ts` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", +})); + +const contentWorkflowCommandSelectors = [ + "content_workflow_create", + "content_workflow_get", + "content_workflow_get_by_content", + "content_workflow_advance", + "content_workflow_retry", + "content_workflow_cancel", +].map((command) => ({ + selector: `CallExpression[callee.name='safeInvoke'][arguments.0.value='${command}'], CallExpression[callee.name='invoke'][arguments.0.value='${command}']`, + message: + "内容工作流命令请统一通过 `src/lib/api/content-workflow.ts` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", +})); + +const novelCommandSelectors = [ + "novel_create_project", + "novel_update_settings", + "novel_generate_outline", + "novel_generate_characters", + "novel_generate_chapter", + "novel_continue_chapter", + "novel_rewrite_chapter", + "novel_polish_chapter", + "novel_check_consistency", + "novel_get_project_snapshot", + "novel_list_runs", + "novel_delete_character", +].map((command) => ({ + selector: `CallExpression[callee.name='safeInvoke'][arguments.0.value='${command}'], CallExpression[callee.name='invoke'][arguments.0.value='${command}']`, + message: + "小说编排命令请统一通过 `src/lib/api/novel.ts` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", +})); + +const unifiedChatCommandSelectors = [ + "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", +].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` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", +})); + +const unifiedMemoryCommandSelectors = [ + "unified_memory_list", + "unified_memory_search", + "unified_memory_get", + "unified_memory_create", + "unified_memory_update", + "unified_memory_delete", + "unified_memory_stats", + "unified_memory_analyze", + "unified_memory_semantic_search", + "unified_memory_hybrid_search", +].map((command) => ({ + selector: `CallExpression[callee.name='safeInvoke'][arguments.0.value='${command}'], CallExpression[callee.name='invoke'][arguments.0.value='${command}']`, + message: + "统一记忆命令请统一通过 `src/lib/api/unifiedMemory.ts` 暴露的网关函数调用,避免继续在其他模块中直接拼接命令名。", +})); + const apiCompatibilityCommandSelectors = ["check_api_compatibility"].map( (command) => ({ selector: `CallExpression[callee.name='safeInvoke'][arguments.0.value='${command}'], CallExpression[callee.name='invoke'][arguments.0.value='${command}']`, @@ -929,6 +1070,15 @@ export default [ ...experimentalFeaturesCommandSelectors, ...memoryRuntimeCommandSelectors, ...projectMemoryCommandSelectors, + ...toolHooksCommandSelectors, + ...a2uiFormCommandSelectors, + ...memoryFeedbackCommandSelectors, + ...contextMemoryCommandSelectors, + ...asrProviderCommandSelectors, + ...contentWorkflowCommandSelectors, + ...novelCommandSelectors, + ...unifiedChatCommandSelectors, + ...unifiedMemoryCommandSelectors, ...apiCompatibilityCommandSelectors, ...endpointProvidersCommandSelectors, ...modelCatalogCommandSelectors, @@ -949,22 +1099,6 @@ export default [ "@typescript-eslint/no-explicit-any": "off", }, }, - { - files: ["src/components/chat/ChatPage.tsx"], - rules: { - "no-restricted-imports": createLegacyChatImportRule( - generalChatRestrictedPathsWithoutPage, - ), - }, - }, - { - files: ["src/components/chat/hooks/useChat.ts"], - rules: { - "no-restricted-imports": createLegacyChatImportRule( - generalChatRestrictedPathsWithoutPageAndStore, - ), - }, - }, { files: ["src/components/general-chat/store/useGeneralChatStore.ts"], rules: { @@ -995,12 +1129,21 @@ export default [ "src/lib/api/plugins.ts", "src/lib/api/pluginUI.ts", "src/lib/api/fileBrowser.ts", + "src/lib/api/a2uiForm.ts", "src/lib/api/appUpdate.ts", + "src/lib/api/asrProvider.ts", + "src/lib/api/content-workflow.ts", "src/lib/api/screenshotChat.ts", "src/lib/api/notification.ts", + "src/lib/api/novel.ts", "src/lib/api/autoFix.ts", + "src/lib/api/contextMemory.ts", "src/lib/api/frontendCrash.ts", + "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", "src/lib/api/apiCompatibility.ts", diff --git a/package.json b/package.json index b0520e6ce..e250f046d 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.84.0", + "version": "0.85.0", "type": "module", "engines": { "node": ">=22.0.0" @@ -37,7 +37,8 @@ "bridge:e2e": "node scripts/chrome-bridge-e2e.mjs", "bridge:health": "node scripts/check-dev-bridge-health.mjs", "smoke:social-workbench": "node scripts/social-workbench-e2e-smoke.mjs", - "dev:web-bridge": "node scripts/start-web-bridge-dev.mjs" + "dev:web-bridge": "node scripts/start-web-bridge-dev.mjs", + "governance:legacy-report": "node scripts/report-legacy-surfaces.mjs" }, "dependencies": { "@babel/standalone": "^7.29.0", diff --git a/scripts/report-legacy-surfaces.mjs b/scripts/report-legacy-surfaces.mjs new file mode 100644 index 000000000..0666613e9 --- /dev/null +++ b/scripts/report-legacy-surfaces.mjs @@ -0,0 +1,1006 @@ +#!/usr/bin/env node + +import fs from "node:fs"; +import path from "node:path"; +import process from "node:process"; + +const repoRoot = path.resolve(process.cwd()); +const sourceRoots = ["src"]; +const sourceExtensions = new Set([ + ".ts", + ".tsx", + ".js", + ".jsx", + ".mjs", + ".cjs", +]); +const rustSourceRoots = ["src-tauri/src", "src-tauri/crates"]; +const rustSourceExtensions = new Set([".rs"]); +const ignoredDirs = new Set([ + "node_modules", + "dist", + "build", + "coverage", + "target", + ".git", + ".turbo", + ".next", +]); + +const importSurfaceMonitors = [ + { + id: "general-chat-root-entry", + classification: "deprecated", + description: "旧 general-chat 根导出入口", + targets: ["src/components/general-chat/index.ts"], + allowedPaths: [], + }, + { + id: "general-chat-page-entry", + classification: "deprecated", + description: "旧 general-chat 页面实现入口", + targets: ["src/components/general-chat/GeneralChatPage.tsx"], + allowedPaths: ["src/components/general-chat/index.ts"], + }, + { + id: "general-chat-legacy-session-hook", + classification: "dead-candidate", + description: "旧 general-chat 会话兼容 Hook", + targets: ["src/components/general-chat/hooks/useSession.ts"], + allowedPaths: [], + }, + { + id: "general-chat-legacy-streaming-hook", + classification: "compat", + description: "旧 general-chat 流式兼容 Hook", + targets: ["src/components/general-chat/hooks/useStreaming.ts"], + allowedPaths: ["src/components/general-chat/GeneralChatPage.tsx"], + }, + { + 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"], + }, + { + id: "agent-compat-gateway", + classification: "deprecated", + description: "Agent / Aster compat API 网关", + targets: ["src/lib/api/agentCompat.ts"], + allowedPaths: [], + }, +]; + +const commandSurfaceMonitors = [ + { + id: "general-chat-compat-commands", + classification: "compat", + description: "general_chat compat 命令前端边界", + commands: [ + "general_chat_get_session", + "general_chat_list_sessions", + "general_chat_create_session", + "general_chat_delete_session", + "general_chat_rename_session", + "general_chat_get_messages", + ], + allowedPaths: ["src/lib/api/generalChatCompat.ts"], + }, + { + id: "conversation-memory-legacy-commands", + classification: "compat", + 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"], + }, + { + id: "prompt-switch-legacy-command", + classification: "deprecated", + description: "旧 prompt 切换命令前端边界", + commands: ["switch_prompt"], + allowedPaths: [], + }, + { + id: "api-key-legacy-migration-commands", + classification: "deprecated", + description: "旧 API Key 迁移命令前端边界", + commands: [ + "get_legacy_api_key_credentials", + "migrate_legacy_api_key_credentials", + "delete_legacy_api_key_credential", + ], + allowedPaths: [], + }, +]; + +const rustTextSurfaceMonitors = [ + { + id: "rust-general-chat-dao", + classification: "deprecated", + description: "Rust 业务层 direct GeneralChatDao 依赖", + patterns: ["GeneralChatDao", "database::dao::general_chat"], + allowedPaths: [], + }, + { + id: "rust-legacy-general-tables", + classification: "compat", + description: "Rust runtime direct legacy general 表访问", + patterns: ["general_chat_sessions", "general_chat_messages"], + 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 扩散", + patterns: [ + "load_pending_general_session_messages_raw", + "load_pending_general_messages_raw", + "count_pending_general_sessions_raw", + "count_pending_general_messages_raw", + "sum_pending_general_message_chars_raw", + "load_legacy_general_session_messages", + "load_unmigrated_legacy_general_messages", + "count_unmigrated_legacy_general_sessions", + "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", + ], + }, + { + id: "rust-legacy-general-module-imports", + classification: "deprecated", + description: "Rust 外部模块直接引用 pending/legacy general 子模块", + patterns: [ + "crate::database::legacy_general_chat::", + "proxycast_core::database::legacy_general_chat::", + "crate::database::pending_general_chat::", + "proxycast_core::database::pending_general_chat::", + ], + allowedPaths: [], + }, + { + id: "rust-general-migration-flag-runtime-leak", + classification: "deprecated", + description: "Rust 业务层重新直接判断 general 迁移完成标记", + patterns: [ + "migration::is_general_chat_migration_completed", + "is_general_chat_migration_completed(", + ], + 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 兼容壳", + patterns: [ + "use crate::general_chat::", + "use crate::general_chat::{", + "crate::general_chat::SessionService", + ], + includePathPrefixes: ["src-tauri/crates/services/src"], + allowedPaths: [], + }, + { + id: "rust-cross-crate-general-chat-compat", + classification: "deprecated", + description: "跨 crate 引回 proxycast_services::general_chat 兼容壳", + patterns: ["proxycast_services::general_chat::"], + allowedPaths: [], + }, + { + id: "rust-provider-pool-legacy-selector", + classification: "deprecated", + description: "provider pool legacy 凭证选择兼容方法", + patterns: ["select_credential_with_fallback_legacy"], + allowedPaths: [], + }, + { + id: "rust-memory-legacy-command-shells", + classification: "deprecated", + description: "旧 conversation memory Rust 命令壳回流", + patterns: [ + "get_conversation_memory_stats", + "get_conversation_memory_overview", + "request_conversation_memory_analysis", + "cleanup_conversation_memory", + ], + includePathPrefixes: ["src-tauri/src"], + allowedPaths: [], + }, + { + id: "rust-migration-setting-key-leak", + classification: "deprecated", + description: "Rust 迁移 settings 标记字符串扩散", + patterns: [ + "\"migrated_api_keys_to_pool\"", + "\"migrated_provider_ids_v1\"", + "\"cleaned_legacy_api_key_credentials\"", + "\"migrated_mcp_proxycast_enabled\"", + "\"migrated_mcp_created_at_to_integer\"", + "\"model_registry_refresh_needed\"", + "\"model_registry_version\"", + ], + allowedPaths: [ + "src-tauri/crates/core/src/database/migration.rs", + "src-tauri/crates/core/src/database/migration/api_key_migration.rs", + "src-tauri/crates/core/src/database/migration/general_chat_migration.rs", + "src-tauri/crates/core/src/database/migration/mcp_migration.rs", + "src-tauri/crates/core/src/database/migration/model_registry_migration.rs", + ], + }, + { + id: "rust-startup-migration-call-leak", + classification: "deprecated", + description: "Rust 启动迁移直接调用扩散", + patterns: [ + "migration::migrate_provider_ids(", + "migration::mark_model_registry_refresh_needed(", + "migration::check_model_registry_version(", + "migration::migrate_api_keys_to_pool(", + "migration::cleanup_legacy_api_key_credentials(", + "migration::migrate_mcp_proxycast_enabled(", + "migration::migrate_mcp_created_at_to_integer(", + "migration::check_general_chat_migration_status(", + "migration::migrate_general_chat_to_unified(", + "migration_v2::migrate_unified_content_system(", + "migration_v3::migrate_playwright_mcp_server(", + "migration_v4::migrate_fix_promise_paths(", + ], + allowedPaths: ["src-tauri/crates/core/src/database/startup_migrations.rs"], + }, + { + id: "rust-startup-migration-manual-match-leak", + classification: "deprecated", + description: "startup migration 回流手写 match 调度", + patterns: [ + "match migration::migrate_provider_ids(", + "match migration::migrate_api_keys_to_pool(", + "match migration::cleanup_legacy_api_key_credentials(", + "match migration::migrate_mcp_proxycast_enabled(", + "match migration::migrate_mcp_created_at_to_integer(", + "match migration::migrate_general_chat_to_unified(", + "match migration_v2::migrate_unified_content_system(", + "match migration_v3::migrate_playwright_mcp_server(", + "match migration_v4::migrate_fix_promise_paths(", + ], + allowedPaths: [], + }, + { + id: "rust-versioned-migration-local-helper-leak", + classification: "deprecated", + description: "versioned migration 本地重复 settings helper 回流", + patterns: [ + "fn is_migration_completed(conn:", + "fn mark_migration_completed(conn:", + ], + includePathPrefixes: ["src-tauri/crates/core/src/database/migration_v"], + allowedPaths: [], + }, + { + id: "rust-versioned-migration-transaction-leak", + classification: "deprecated", + description: "versioned migration 直接手写事务样板回流", + patterns: [ + "conn.execute(\"BEGIN TRANSACTION\"", + "conn.execute(\"COMMIT\"", + "conn.execute(\"ROLLBACK\"", + ], + includePathPrefixes: ["src-tauri/crates/core/src/database/migration_v"], + allowedPaths: [], + }, + { + id: "rust-hardcoded-projects-path-leak", + classification: "deprecated", + description: "数据库迁移硬编码 legacy projects 路径", + patterns: ["\".proxycast/projects\"", "join(\".proxycast\").join(\"projects\")"], + includePathPrefixes: ["src-tauri/crates/core/src/database"], + allowedPaths: [], + }, + { + id: "rust-hardcoded-session-files-path-leak", + classification: "deprecated", + description: "session files 硬编码 legacy sessions 路径", + patterns: ["~/.proxycast/sessions", "join(\".proxycast\").join(\"sessions\")"], + includePathPrefixes: ["src-tauri/crates/core/src/session_files"], + allowedPaths: [], + }, + { + id: "rust-hardcoded-legacy-config-path-leak", + classification: "deprecated", + description: "数据库迁移硬编码 legacy config 路径", + patterns: ["~/.proxycast/config.json", "join(\".proxycast\").join(\"config.json\")"], + includePathPrefixes: ["src-tauri/crates/core/src/database"], + allowedPaths: [], + }, + { + id: "rust-hardcoded-workspace-projects-path-leak", + classification: "deprecated", + description: "上层命令或桥接层硬编码 workspace projects 路径", + patterns: ["~/.proxycast/projects", "join(\".proxycast\").join(\"projects\")"], + includePathPrefixes: ["src-tauri/src"], + allowedPaths: [], + }, + { + id: "rust-hardcoded-logger-path-leak", + classification: "deprecated", + description: "logger fallback 硬编码 legacy logs 路径", + patterns: ["~/.proxycast/logs", "join(\".proxycast\").join(\"logs\")"], + includePathPrefixes: ["src-tauri/crates/core/src/logger.rs"], + allowedPaths: [], + }, + { + id: "rust-hardcoded-skills-path-leak", + classification: "deprecated", + description: "skills 相关模块硬编码 legacy skills 路径", + patterns: ["~/.proxycast/skills", "join(\".proxycast\").join(\"skills\")"], + includePathPrefixes: ["src-tauri/src"], + allowedPaths: [], + }, + { + id: "rust-hardcoded-memory-path-leak", + classification: "deprecated", + description: "memory 相关模块硬编码 legacy memory 或 AGENTS 路径", + patterns: [ + "~/.proxycast/AGENTS.md", + "join(\".proxycast\").join(\"AGENTS.md\")", + "join(\".proxycast\").join(\"memory\")", + ".proxycast/memory", + ], + includePathPrefixes: ["src-tauri/src"], + allowedPaths: [], + }, +]; + +const rustTextCountMonitors = [ + { + id: "rust-app-paths-root-fetch-duplication", + classification: "deprecated", + description: "app_paths 重复获取 preferred/legacy root 样板回流", + includePathPrefixes: ["src-tauri/crates/core/src/app_paths.rs"], + occurrences: [ + { + pattern: "let preferred_root = preferred_data_dir()?;", + maxCount: 1, + }, + { + pattern: "let legacy_root = legacy_home_dir()?;", + maxCount: 1, + }, + ], + }, +]; + +function normalizePath(filePath) { + return filePath.split(path.sep).join("/"); +} + +function resolveExistingSourcePath(absolutePath) { + if (fs.existsSync(absolutePath)) { + const stats = fs.statSync(absolutePath); + if (stats.isFile()) { + return absolutePath; + } + } + + if (!path.extname(absolutePath)) { + for (const extension of sourceExtensions) { + const fileCandidate = `${absolutePath}${extension}`; + if (fs.existsSync(fileCandidate) && fs.statSync(fileCandidate).isFile()) { + return fileCandidate; + } + } + } + + for (const extension of sourceExtensions) { + const indexCandidate = path.join(absolutePath, `index${extension}`); + if (fs.existsSync(indexCandidate) && fs.statSync(indexCandidate).isFile()) { + return indexCandidate; + } + } + + return null; +} + +function resolveImportPath(importerRelativePath, specifier) { + let absoluteCandidate = null; + + if (specifier.startsWith("@/")) { + absoluteCandidate = path.join(repoRoot, "src", specifier.slice(2)); + } else if (specifier.startsWith(".")) { + absoluteCandidate = path.resolve( + path.dirname(path.join(repoRoot, importerRelativePath)), + specifier, + ); + } + + if (!absoluteCandidate) { + return null; + } + + const resolvedPath = resolveExistingSourcePath(absoluteCandidate); + if (!resolvedPath) { + return null; + } + + return normalizePath(path.relative(repoRoot, resolvedPath)); +} + +function isTestFile(relativePath) { + return ( + /(^|\/)tests(\/|$)/.test(relativePath) || + /(^|\/)(__tests__|__mocks__)(\/|$)/.test(relativePath) || + /\.(test|spec)\.[^/.]+$/.test(relativePath) + ); +} + +function walkDirectory(directoryPath, extensions) { + const files = []; + + for (const entry of fs.readdirSync(directoryPath, { withFileTypes: true })) { + if (ignoredDirs.has(entry.name)) { + continue; + } + + const fullPath = path.join(directoryPath, entry.name); + if (entry.isDirectory()) { + files.push(...walkDirectory(fullPath, extensions)); + continue; + } + + if (!extensions.has(path.extname(entry.name))) { + continue; + } + + files.push(fullPath); + } + + return files; +} + +function extractImportSpecifiers(sourceCode) { + const specifiers = new Set(); + const patterns = [ + /\bimport\s+(?:type\s+)?(?:[\s\S]*?\s+from\s+)?["'`]([^"'`]+)["'`]/g, + /\bexport\s+(?:type\s+)?[\s\S]*?\s+from\s+["'`]([^"'`]+)["'`]/g, + /\bimport\s*\(\s*["'`]([^"'`]+)["'`]\s*\)/g, + /\brequire\s*\(\s*["'`]([^"'`]+)["'`]\s*\)/g, + /\b(?:vi|jest)\.mock\s*\(\s*["'`]([^"'`]+)["'`]/g, + ]; + + for (const pattern of patterns) { + for (const match of sourceCode.matchAll(pattern)) { + specifiers.add(match[1]); + } + } + + return specifiers; +} + +function extractInvokeCommands(sourceCode) { + const commands = new Set(); + const patterns = [ + /\bsafeInvoke(?:<[^>]+>)?\s*\(\s*["'`]([^"'`]+)["'`]/g, + /\binvoke(?:<[^>]+>)?\s*\(\s*["'`]([^"'`]+)["'`]/g, + ]; + + for (const pattern of patterns) { + for (const match of sourceCode.matchAll(pattern)) { + commands.add(match[1]); + } + } + + return commands; +} + +function stripRustTestModules(sourceCode) { + return sourceCode.replace( + /(?:^|\n)\s*#\s*\[\s*cfg\s*\(\s*test\s*\)\s*\][\s\S]*$/m, + "\n", + ); +} + +function collectSources() { + const runtimeSources = []; + const testSources = []; + + for (const root of sourceRoots) { + const absoluteRoot = path.join(repoRoot, root); + if (!fs.existsSync(absoluteRoot)) { + continue; + } + + for (const filePath of walkDirectory(absoluteRoot, sourceExtensions)) { + const relativePath = normalizePath(path.relative(repoRoot, filePath)); + const sourceCode = fs.readFileSync(filePath, "utf8"); + const imports = extractImportSpecifiers(sourceCode); + const collectedSource = { + relativePath, + imports, + resolvedImports: new Set( + [...imports] + .map((specifier) => resolveImportPath(relativePath, specifier)) + .filter(Boolean), + ), + commands: extractInvokeCommands(sourceCode), + }; + + if (isTestFile(relativePath)) { + testSources.push(collectedSource); + continue; + } + + runtimeSources.push(collectedSource); + } + } + + return { + runtimeSources, + testSources, + }; +} + +function collectTextSources(roots, extensions) { + const runtimeSources = []; + const testSources = []; + + for (const root of roots) { + const absoluteRoot = path.join(repoRoot, root); + if (!fs.existsSync(absoluteRoot)) { + continue; + } + + for (const filePath of walkDirectory(absoluteRoot, extensions)) { + const relativePath = normalizePath(path.relative(repoRoot, filePath)); + const sourceCode = fs.readFileSync(filePath, "utf8"); + const collectedSource = { + relativePath, + sourceCode: + path.extname(relativePath) === ".rs" + ? stripRustTestModules(sourceCode) + : sourceCode, + rawSourceCode: sourceCode, + }; + + if (isTestFile(relativePath)) { + testSources.push(collectedSource); + continue; + } + + runtimeSources.push(collectedSource); + } + } + + return { + runtimeSources, + testSources, + }; +} + +function formatPaths(paths) { + if (paths.length === 0) { + return "无"; + } + + return paths.map((item) => ` - ${item}`).join("\n"); +} + +function evaluateImportMonitor(monitor, runtimeSources, testSources) { + const existingTargets = monitor.targets.filter((target) => + fs.existsSync(path.join(repoRoot, target)), + ); + const missingTargets = monitor.targets.filter( + (target) => !fs.existsSync(path.join(repoRoot, target)), + ); + const references = runtimeSources + .filter((file) => + [...file.resolvedImports].some((resolvedPath) => + monitor.targets.includes(resolvedPath), + ), + ) + .map((file) => file.relativePath) + .sort(); + const testReferences = testSources + .filter((file) => + [...file.resolvedImports].some((resolvedPath) => + monitor.targets.includes(resolvedPath), + ), + ) + .map((file) => file.relativePath) + .sort(); + + const violations = references.filter( + (relativePath) => !monitor.allowedPaths.includes(relativePath), + ); + + return { + ...monitor, + existingTargets, + missingTargets, + references, + testReferences, + violations, + }; +} + +function evaluateCommandMonitor(monitor, runtimeSources, testSources) { + const referencesByCommand = new Map(); + const testReferencesByCommand = new Map(); + + for (const command of monitor.commands) { + referencesByCommand.set( + command, + runtimeSources + .filter((file) => file.commands.has(command)) + .map((file) => file.relativePath) + .sort(), + ); + testReferencesByCommand.set( + command, + testSources + .filter((file) => file.commands.has(command)) + .map((file) => file.relativePath) + .sort(), + ); + } + + const violations = []; + for (const [command, references] of referencesByCommand.entries()) { + for (const relativePath of references) { + if (!monitor.allowedPaths.includes(relativePath)) { + violations.push(`${command} -> ${relativePath}`); + } + } + } + + return { + ...monitor, + referencesByCommand, + testReferencesByCommand, + violations, + }; +} + +function evaluateTextMonitor(monitor, runtimeSources, testSources) { + const filteredRuntimeSources = monitor.includePathPrefixes + ? runtimeSources.filter((file) => + monitor.includePathPrefixes.some((prefix) => + file.relativePath.startsWith(prefix), + ), + ) + : runtimeSources; + const filteredTestSources = monitor.includePathPrefixes + ? testSources.filter((file) => + monitor.includePathPrefixes.some((prefix) => + file.relativePath.startsWith(prefix), + ), + ) + : testSources; + const matchesPattern = (sourceCode) => + monitor.patterns.some((pattern) => sourceCode.includes(pattern)); + + const references = filteredRuntimeSources + .filter((file) => matchesPattern(file.sourceCode)) + .map((file) => file.relativePath) + .sort(); + const testReferences = filteredTestSources + .filter((file) => matchesPattern(file.rawSourceCode ?? file.sourceCode)) + .map((file) => file.relativePath) + .sort(); + const violations = references.filter( + (relativePath) => !monitor.allowedPaths.includes(relativePath), + ); + + return { + ...monitor, + references, + testReferences, + violations, + }; +} + +function countOccurrences(sourceCode, pattern) { + if (!pattern) { + return 0; + } + + let count = 0; + let startIndex = 0; + + while (true) { + const matchIndex = sourceCode.indexOf(pattern, startIndex); + if (matchIndex === -1) { + return count; + } + count += 1; + startIndex = matchIndex + pattern.length; + } +} + +function evaluateTextCountMonitor(monitor, runtimeSources, testSources) { + const filteredRuntimeSources = monitor.includePathPrefixes + ? runtimeSources.filter((file) => + monitor.includePathPrefixes.some((prefix) => + file.relativePath.startsWith(prefix), + ), + ) + : runtimeSources; + const filteredTestSources = monitor.includePathPrefixes + ? testSources.filter((file) => + monitor.includePathPrefixes.some((prefix) => + file.relativePath.startsWith(prefix), + ), + ) + : testSources; + const runtimeMatches = []; + const testMatches = []; + const violations = []; + + for (const file of filteredRuntimeSources) { + const counts = monitor.occurrences + .map((rule) => ({ + ...rule, + count: countOccurrences(file.sourceCode, rule.pattern), + })) + .filter((rule) => rule.count > 0); + + if (counts.length === 0) { + continue; + } + + runtimeMatches.push({ + relativePath: file.relativePath, + counts, + }); + + for (const rule of counts) { + if (rule.count > rule.maxCount) { + violations.push( + `${file.relativePath} -> ${rule.pattern} (${rule.count} > ${rule.maxCount})`, + ); + } + } + } + + for (const file of filteredTestSources) { + const counts = monitor.occurrences + .map((rule) => ({ + ...rule, + count: countOccurrences(file.rawSourceCode ?? file.sourceCode, rule.pattern), + })) + .filter((rule) => rule.count > 0); + + if (counts.length === 0) { + continue; + } + + testMatches.push({ + relativePath: file.relativePath, + counts, + }); + } + + return { + ...monitor, + runtimeMatches, + testMatches, + violations, + }; +} + +function printImportReport(result) { + const status = + result.violations.length > 0 + ? "违规" + : result.references.length === 0 && result.existingTargets.length === 0 + ? "已删除" + : result.references.length === 0 + ? "零引用" + : "受控"; + + console.log( + `- [${status}] ${result.id} (${result.classification}):${result.description}`, + ); + console.log(` 目标文件:${result.targets.join(", ")}`); + console.log(` 允许引用:${result.allowedPaths.join(", ") || "无"}`); + if (result.missingTargets.length > 0) { + console.log(` 已删除目标:\n${formatPaths(result.missingTargets)}`); + } + console.log(` 实际引用:\n${formatPaths(result.references)}`); + console.log(` 测试引用:\n${formatPaths(result.testReferences)}`); + + if (result.violations.length > 0) { + console.log(` 违规引用:\n${formatPaths(result.violations)}`); + } +} + +function printCommandReport(result) { + const flattenedReferences = [...result.referencesByCommand.values()].flat(); + const uniqueReferences = [...new Set(flattenedReferences)].sort(); + const status = + result.violations.length > 0 + ? "违规" + : uniqueReferences.length === 0 + ? "零引用" + : "受控"; + + console.log( + `- [${status}] ${result.id} (${result.classification}):${result.description}`, + ); + console.log(` 命令:${result.commands.join(", ")}`); + console.log(` 允许引用:${result.allowedPaths.join(", ") || "无"}`); + + for (const command of result.commands) { + const references = result.referencesByCommand.get(command) ?? []; + const testReferences = result.testReferencesByCommand.get(command) ?? []; + console.log(` ${command}:\n${formatPaths(references)}`); + console.log(` ${command}(测试):\n${formatPaths(testReferences)}`); + } + + if (result.violations.length > 0) { + console.log(` 违规引用:\n${formatPaths(result.violations)}`); + } +} + +function printTextReport(result) { + const status = + result.violations.length > 0 + ? "违规" + : result.references.length === 0 + ? "零引用" + : "受控"; + + console.log( + `- [${status}] ${result.id} (${result.classification}):${result.description}`, + ); + console.log(` 关键字:${result.patterns.join(", ")}`); + console.log(` 允许引用:${result.allowedPaths.join(", ") || "无"}`); + console.log(` 实际引用:\n${formatPaths(result.references)}`); + console.log(` 测试引用:\n${formatPaths(result.testReferences)}`); + + if (result.violations.length > 0) { + console.log(` 违规引用:\n${formatPaths(result.violations)}`); + } +} + +function printTextCountReport(result) { + const status = + result.violations.length > 0 + ? "违规" + : result.runtimeMatches.length === 0 + ? "零引用" + : "受控"; + + console.log( + `- [${status}] ${result.id} (${result.classification}):${result.description}`, + ); + console.log( + ` 次数规则:${result.occurrences + .map((rule) => `${rule.pattern} <= ${rule.maxCount}`) + .join(";")}`, + ); + console.log( + ` 实际命中:\n${formatPaths( + result.runtimeMatches.map( + (item) => + `${item.relativePath} -> ${item.counts + .map((rule) => `${rule.pattern} (${rule.count})`) + .join(";")}`, + ), + )}`, + ); + console.log( + ` 测试命中:\n${formatPaths( + result.testMatches.map( + (item) => + `${item.relativePath} -> ${item.counts + .map((rule) => `${rule.pattern} (${rule.count})`) + .join(";")}`, + ), + )}`, + ); + + if (result.violations.length > 0) { + console.log(` 违规引用:\n${formatPaths(result.violations)}`); + } +} + +const { runtimeSources, testSources } = collectSources(); +const { runtimeSources: rustRuntimeSources, testSources: rustTestSources } = + collectTextSources(rustSourceRoots, rustSourceExtensions); +const importResults = importSurfaceMonitors.map((monitor) => + evaluateImportMonitor(monitor, runtimeSources, testSources), +); +const commandResults = commandSurfaceMonitors.map((monitor) => + evaluateCommandMonitor(monitor, runtimeSources, testSources), +); +const rustTextResults = rustTextSurfaceMonitors.map((monitor) => + evaluateTextMonitor(monitor, rustRuntimeSources, rustTestSources), +); +const rustTextCountResults = rustTextCountMonitors.map((monitor) => + evaluateTextCountMonitor(monitor, rustRuntimeSources, rustTestSources), +); + +const zeroReferenceCandidates = importResults + .filter( + (result) => + result.references.length === 0 && result.existingTargets.length > 0, + ) + .map((result) => `${result.id} (${result.description})`); +const violations = [ + ...importResults.flatMap((result) => + result.violations.map((item) => `${result.id} -> ${item}`), + ), + ...commandResults.flatMap((result) => + result.violations.map((item) => `${result.id} -> ${item}`), + ), + ...rustTextResults.flatMap((result) => + result.violations.map((item) => `${result.id} -> ${item}`), + ), + ...rustTextCountResults.flatMap((result) => + result.violations.map((item) => `${result.id} -> ${item}`), + ), +]; + +console.log("[proxycast] legacy surface report"); +console.log(""); +console.log("## 入口引用"); +for (const result of importResults) { + printImportReport(result); +} + +console.log(""); +console.log("## 命令边界"); +for (const result of commandResults) { + printCommandReport(result); +} + +console.log(""); +console.log("## Rust 护栏"); +for (const result of rustTextResults) { + printTextReport(result); +} +for (const result of rustTextCountResults) { + printTextCountReport(result); +} + +console.log(""); +console.log("## 摘要"); +console.log(`- 扫描文件数:${runtimeSources.length}`); +console.log(`- 测试文件数:${testSources.length}`); +console.log(`- Rust 扫描文件数:${rustRuntimeSources.length}`); +console.log(`- Rust 测试文件数:${rustTestSources.length}`); +console.log(`- 零引用候选:${zeroReferenceCandidates.length}`); +for (const candidate of zeroReferenceCandidates) { + console.log(` - ${candidate}`); +} +console.log(`- 边界违规:${violations.length}`); +for (const violation of violations) { + console.log(` - ${violation}`); +} + +if (violations.length > 0) { + console.error(""); + console.error( + "[proxycast] legacy surface report 检测到边界违规,请先治理再继续扩展。", + ); + process.exit(1); +} diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 5dc29614e..98e21b314 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -6982,7 +6982,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.84.0" +version = "0.85.0" dependencies = [ "anyhow", "arboard", @@ -7084,7 +7084,7 @@ dependencies = [ [[package]] name = "proxycast-agent" -version = "0.84.0" +version = "0.85.0" dependencies = [ "aster-core", "async-trait", @@ -7109,7 +7109,7 @@ dependencies = [ [[package]] name = "proxycast-config" -version = "0.84.0" +version = "0.85.0" dependencies = [ "async-trait", "parking_lot", @@ -7125,7 +7125,7 @@ dependencies = [ [[package]] name = "proxycast-core" -version = "0.84.0" +version = "0.85.0" dependencies = [ "aster-models", "async-trait", @@ -7165,7 +7165,7 @@ dependencies = [ [[package]] name = "proxycast-credential" -version = "0.84.0" +version = "0.85.0" dependencies = [ "axum 0.7.9", "base64 0.22.1", @@ -7200,7 +7200,7 @@ dependencies = [ [[package]] name = "proxycast-gateway" -version = "0.84.0" +version = "0.85.0" dependencies = [ "axum 0.7.9", "chrono", @@ -7221,7 +7221,7 @@ dependencies = [ [[package]] name = "proxycast-infra" -version = "0.84.0" +version = "0.85.0" dependencies = [ "chrono", "dashmap 5.5.3", @@ -7241,7 +7241,7 @@ dependencies = [ [[package]] name = "proxycast-mcp" -version = "0.84.0" +version = "0.85.0" dependencies = [ "async-trait", "dirs 5.0.1", @@ -7273,7 +7273,7 @@ dependencies = [ [[package]] name = "proxycast-processor" -version = "0.84.0" +version = "0.85.0" dependencies = [ "async-trait", "parking_lot", @@ -7292,7 +7292,7 @@ dependencies = [ [[package]] name = "proxycast-providers" -version = "0.84.0" +version = "0.85.0" dependencies = [ "anyhow", "async-stream", @@ -7346,7 +7346,7 @@ dependencies = [ [[package]] name = "proxycast-server" -version = "0.84.0" +version = "0.85.0" dependencies = [ "aster-core", "async-stream", @@ -7391,7 +7391,7 @@ dependencies = [ [[package]] name = "proxycast-server-utils" -version = "0.84.0" +version = "0.85.0" dependencies = [ "axum 0.7.9", "futures", @@ -7406,7 +7406,7 @@ dependencies = [ [[package]] name = "proxycast-services" -version = "0.84.0" +version = "0.85.0" dependencies = [ "anyhow", "aster-core", @@ -7448,7 +7448,7 @@ dependencies = [ [[package]] name = "proxycast-skills" -version = "0.84.0" +version = "0.85.0" dependencies = [ "async-trait", "dirs 5.0.1", @@ -7464,7 +7464,7 @@ dependencies = [ [[package]] name = "proxycast-terminal" -version = "0.84.0" +version = "0.85.0" dependencies = [ "async-trait", "base64 0.22.1", @@ -7491,7 +7491,7 @@ dependencies = [ [[package]] name = "proxycast-websocket" -version = "0.84.0" +version = "0.85.0" dependencies = [ "axum 0.7.9", "chrono", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 15b85ec7e..63e39f14c 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -3,7 +3,7 @@ members = ["crates/*"] resolver = "2" [workspace.package] -version = "0.84.0" +version = "0.85.0" edition = "2021" authors = ["coso"] repository = "https://github.com/aiclientproxy/proxycast" @@ -191,7 +191,7 @@ version = "2.4" [package] name = "proxycast" -version = "0.84.0" +version = "0.85.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" diff --git a/src-tauri/crates/agent/src/event_converter.rs b/src-tauri/crates/agent/src/event_converter.rs index c0aeaae01..9a79cd167 100644 --- a/src-tauri/crates/agent/src/event_converter.rs +++ b/src-tauri/crates/agent/src/event_converter.rs @@ -5,6 +5,7 @@ use aster::agents::AgentEvent; use aster::conversation::message::{ActionRequiredData, Message, MessageContent}; +use proxycast_core::database::dao::agent_timeline::{AgentThreadItem, AgentThreadTurn}; use regex::Regex; use serde::{Deserialize, Serialize}; @@ -523,6 +524,34 @@ fn extract_tool_result_metadata( #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "type")] pub enum TauriAgentEvent { + /// 线程开始 + #[serde(rename = "thread_started")] + ThreadStarted { thread_id: String }, + + /// turn 开始 + #[serde(rename = "turn_started")] + TurnStarted { turn: AgentThreadTurn }, + + /// item 开始 + #[serde(rename = "item_started")] + ItemStarted { item: AgentThreadItem }, + + /// item 更新 + #[serde(rename = "item_updated")] + ItemUpdated { item: AgentThreadItem }, + + /// item 完成 + #[serde(rename = "item_completed")] + ItemCompleted { item: AgentThreadItem }, + + /// turn 完成 + #[serde(rename = "turn_completed")] + TurnCompleted { turn: AgentThreadTurn }, + + /// turn 失败 + #[serde(rename = "turn_failed")] + TurnFailed { turn: AgentThreadTurn }, + /// 文本增量 #[serde(rename = "text_delta")] TextDelta { text: String }, diff --git a/src-tauri/crates/agent/src/lib.rs b/src-tauri/crates/agent/src/lib.rs index eec5c9499..2a4a02037 100644 --- a/src-tauri/crates/agent/src/lib.rs +++ b/src-tauri/crates/agent/src/lib.rs @@ -3,6 +3,13 @@ //! 包含 Agent 模块中不依赖主 crate 内部模块的纯逻辑部分。 //! 深耦合部分(aster_state、aster_agent 流式桥接)留在主 crate。 +#![allow(clippy::explicit_counter_loop)] +#![allow(clippy::unnecessary_map_or)] +#![allow(clippy::to_string_in_format_args)] +#![allow(clippy::match_like_matches_macro)] +#![allow(clippy::derivable_impls)] +#![allow(clippy::borrowed_box)] + pub mod ask_bridge; pub mod aster_state; pub mod aster_state_support; diff --git a/src-tauri/crates/agent/src/session_store.rs b/src-tauri/crates/agent/src/session_store.rs index 72a422f12..618e7b9f5 100644 --- a/src-tauri/crates/agent/src/session_store.rs +++ b/src-tauri/crates/agent/src/session_store.rs @@ -6,6 +6,9 @@ use chrono::Utc; use proxycast_core::agent::types::{AgentMessage, AgentSession, ContentPart, MessageContent}; use proxycast_core::database::dao::agent::AgentDao; +use proxycast_core::database::dao::agent_timeline::{ + AgentThreadItem, AgentThreadTurn, AgentTimelineDao, +}; use proxycast_core::database::DbConnection; use proxycast_core::workspace::WorkspaceManager; use uuid::Uuid; @@ -35,8 +38,11 @@ pub struct SessionDetail { pub name: String, pub created_at: i64, pub updated_at: i64, + pub thread_id: String, pub messages: Vec, pub execution_strategy: Option, + pub turns: Vec, + pub items: Vec, } /// 解析会话 working_dir(优先入参,其次 workspace_id) @@ -145,6 +151,10 @@ pub fn get_session_sync(db: &DbConnection, session_id: &str) -> Result Result Result { } pub fn resolve_database_path() -> Result { - let preferred_root = preferred_data_dir()?; - let legacy_root = legacy_home_dir()?; - resolve_database_path_from_roots(&preferred_root, &legacy_root) + with_app_roots(resolve_database_path_from_roots) } pub fn resolve_logs_dir() -> Result { - let preferred_root = preferred_data_dir()?; - let legacy_root = legacy_home_dir()?; - resolve_subdir_with_legacy_copy_from_roots(&preferred_root, &legacy_root, "logs") + resolve_runtime_subdir("logs") } pub fn resolve_request_logs_dir() -> Result { + resolve_runtime_subdir("request_logs") +} + +pub fn resolve_projects_dir() -> Result { + resolve_runtime_subdir("projects") +} + +pub fn resolve_sessions_dir() -> Result { + resolve_runtime_subdir("sessions") +} + +pub fn resolve_skills_dir() -> Result { + resolve_runtime_subdir("skills") +} + +pub fn resolve_user_memory_path() -> Result { + with_app_roots(resolve_user_memory_path_from_roots) +} + +pub fn resolve_default_project_dir() -> Result { + with_app_roots(resolve_default_project_dir_from_roots) +} + +pub fn best_effort_runtime_subdir(subdir: &str) -> PathBuf { + resolve_runtime_subdir(subdir).unwrap_or_else(|_| fallback_runtime_subdir(subdir)) +} + +pub fn best_effort_app_data_file(file_name: &str) -> PathBuf { + preferred_data_dir() + .unwrap_or_else(|_| fallback_app_data_dir()) + .join(file_name) +} + +fn with_app_roots( + resolver: impl FnOnce(&Path, &Path) -> Result, +) -> Result { let preferred_root = preferred_data_dir()?; let legacy_root = legacy_home_dir()?; - resolve_subdir_with_legacy_copy_from_roots(&preferred_root, &legacy_root, "request_logs") + resolver(&preferred_root, &legacy_root) +} + +fn resolve_runtime_subdir(subdir: &str) -> Result { + with_app_roots(|preferred_root, legacy_root| { + resolve_subdir_with_legacy_copy_from_roots(preferred_root, legacy_root, subdir) + }) +} + +fn fallback_runtime_subdir(subdir: &str) -> PathBuf { + fallback_app_data_dir().join(subdir) +} + +fn fallback_app_data_dir() -> PathBuf { + std::env::temp_dir().join(APP_DATA_DIR_NAME) +} + +fn resolve_default_project_dir_from_roots( + preferred_root: &Path, + legacy_root: &Path, +) -> Result { + let default_dir = + resolve_subdir_with_legacy_copy_from_roots(preferred_root, legacy_root, "projects")? + .join("default"); + fs::create_dir_all(&default_dir) + .map_err(|e| format!("无法创建默认项目目录 {}: {e}", default_dir.display()))?; + Ok(default_dir) +} + +fn resolve_user_memory_path_from_roots( + preferred_root: &Path, + legacy_root: &Path, +) -> Result { + let preferred_path = preferred_root.join("AGENTS.md"); + if preferred_path.exists() { + return Ok(preferred_path); + } + + let legacy_path = legacy_root.join("AGENTS.md"); + if !legacy_path.exists() { + return Ok(preferred_path); + } + + if let Some(parent) = preferred_path.parent() { + fs::create_dir_all(parent) + .map_err(|e| format!("无法创建用户记忆目录 {}: {e}", parent.display()))?; + } + + match fs::copy(&legacy_path, &preferred_path) { + Ok(_) => Ok(preferred_path), + Err(error) => { + tracing::warn!( + "[路径迁移] 用户记忆文件迁移失败,回退旧路径 {}: {}", + legacy_path.display(), + error + ); + Ok(legacy_path) + } + } } fn resolve_database_path_from_roots( @@ -402,6 +492,110 @@ mod tests { ); } + #[test] + fn resolve_projects_dir_copies_legacy_project_directories() { + let temp = tempdir().unwrap(); + let preferred_root = temp.path().join("appdata").join("proxycast"); + let legacy_root = temp.path().join("home").join(".proxycast"); + let legacy_project_dir = legacy_root.join("projects").join("legacy-project"); + fs::create_dir_all(&legacy_project_dir).unwrap(); + fs::write(legacy_project_dir.join("note.md"), "legacy project").unwrap(); + + let resolved = + resolve_subdir_with_legacy_copy_from_roots(&preferred_root, &legacy_root, "projects") + .unwrap(); + + assert_eq!(resolved, preferred_root.join("projects")); + assert_eq!( + fs::read_to_string(resolved.join("legacy-project").join("note.md")).unwrap(), + "legacy project" + ); + } + + #[test] + fn resolve_sessions_dir_copies_legacy_session_directories() { + let temp = tempdir().unwrap(); + let preferred_root = temp.path().join("appdata").join("proxycast"); + let legacy_root = temp.path().join("home").join(".proxycast"); + let legacy_session_dir = legacy_root + .join("sessions") + .join("legacy-session") + .join("files"); + fs::create_dir_all(&legacy_session_dir).unwrap(); + fs::write(legacy_session_dir.join("note.md"), "legacy session").unwrap(); + + let resolved = + resolve_subdir_with_legacy_copy_from_roots(&preferred_root, &legacy_root, "sessions") + .unwrap(); + + assert_eq!(resolved, preferred_root.join("sessions")); + assert_eq!( + fs::read_to_string( + resolved + .join("legacy-session") + .join("files") + .join("note.md") + ) + .unwrap(), + "legacy session" + ); + } + + #[test] + fn resolve_skills_dir_copies_legacy_skill_directories() { + let temp = tempdir().unwrap(); + let preferred_root = temp.path().join("appdata").join("proxycast"); + let legacy_root = temp.path().join("home").join(".proxycast"); + let legacy_skill_dir = legacy_root.join("skills").join("legacy-skill"); + fs::create_dir_all(&legacy_skill_dir).unwrap(); + fs::write(legacy_skill_dir.join("SKILL.md"), "legacy skill").unwrap(); + + let resolved = + resolve_subdir_with_legacy_copy_from_roots(&preferred_root, &legacy_root, "skills") + .unwrap(); + + assert_eq!(resolved, preferred_root.join("skills")); + assert_eq!( + fs::read_to_string(resolved.join("legacy-skill").join("SKILL.md")).unwrap(), + "legacy skill" + ); + } + + #[test] + fn resolve_user_memory_path_copies_legacy_agents_file() { + let temp = tempdir().unwrap(); + let preferred_root = temp.path().join("appdata").join("proxycast"); + let legacy_root = temp.path().join("home").join(".proxycast"); + fs::create_dir_all(&legacy_root).unwrap(); + fs::write(legacy_root.join("AGENTS.md"), "legacy agents").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 agents"); + } + + #[test] + fn resolve_default_project_dir_creates_default_subdirectory() { + let temp = tempdir().unwrap(); + let preferred_root = temp.path().join("appdata").join("proxycast"); + let legacy_root = temp.path().join("home").join(".proxycast"); + + let resolved = + resolve_default_project_dir_from_roots(&preferred_root, &legacy_root).unwrap(); + + assert_eq!(resolved, preferred_root.join("projects").join("default")); + assert!(resolved.exists()); + assert!(resolved.is_dir()); + } + + #[test] + fn fallback_runtime_subdir_uses_proxycast_temp_namespace() { + let fallback = fallback_runtime_subdir("logs"); + assert!(fallback.ends_with(Path::new(APP_DATA_DIR_NAME).join("logs"))); + } + #[test] fn resolve_database_path_replaces_bootstrap_db_with_legacy_data() { let temp = tempdir().unwrap(); diff --git a/src-tauri/crates/core/src/database/README.md b/src-tauri/crates/core/src/database/README.md index bda2201f0..62fb1db41 100644 --- a/src-tauri/crates/core/src/database/README.md +++ b/src-tauri/crates/core/src/database/README.md @@ -23,10 +23,10 @@ - `providers` - Provider 配置 - `settings` - 应用设置 -### 通用对话表 +### Legacy 通用对话表 -- `general_chat_sessions` - 通用对话会话 -- `general_chat_messages` - 通用对话消息 +- `general_chat_sessions` - 历史通用对话会话(legacy 保留) +- `general_chat_messages` - 历史通用对话消息(legacy 保留) ### 功能表 @@ -42,7 +42,6 @@ |------|------| | `dao/agent.rs` | Agent 会话和消息 DAO | | `dao/api_key_provider.rs` | API Key Provider DAO | -| `dao/general_chat.rs` | 通用对话会话和消息 DAO | | `dao/mcp.rs` | MCP 服务器 DAO | | `dao/prompts.rs` | 提示词 DAO | | `dao/provider_pool.rs` | 凭证池 DAO | diff --git a/src-tauri/crates/core/src/database/dao/agent_run.rs b/src-tauri/crates/core/src/database/dao/agent_run.rs index 55994c7db..c4f2c217e 100644 --- a/src-tauri/crates/core/src/database/dao/agent_run.rs +++ b/src-tauri/crates/core/src/database/dao/agent_run.rs @@ -245,6 +245,46 @@ impl AgentRunDao { iter.collect() } + + pub fn list_terminal_runs_by_session( + conn: &Connection, + session_id: &str, + limit: usize, + offset: usize, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, source, source_ref, session_id, status, started_at, finished_at, duration_ms, + error_code, error_message, metadata, created_at, updated_at + FROM agent_runs + WHERE session_id = ?1 + AND status IN ('success', 'error', 'canceled', 'timeout') + ORDER BY started_at DESC + LIMIT ?2 OFFSET ?3", + )?; + + let iter = stmt.query_map(params![session_id, limit as i64, offset as i64], |row| { + let status_raw: String = row.get(4)?; + let status = + AgentRunStatus::try_from(status_raw.as_str()).unwrap_or(AgentRunStatus::Error); + Ok(AgentRun { + id: row.get(0)?, + source: row.get(1)?, + source_ref: row.get(2)?, + session_id: row.get(3)?, + status, + started_at: row.get(5)?, + finished_at: row.get(6)?, + duration_ms: row.get(7)?, + error_code: row.get(8)?, + error_message: row.get(9)?, + metadata: row.get(10)?, + created_at: row.get(11)?, + updated_at: row.get(12)?, + }) + })?; + + iter.collect() + } } #[cfg(test)] @@ -363,4 +403,48 @@ mod tests { assert_eq!(runs[0].id, "run-a-2"); assert_eq!(runs[1].id, "run-a-1"); } + + #[test] + fn list_terminal_runs_by_session_should_filter_terminal_status_and_offset() { + let conn = setup_conn(); + + let mut run_success = sample_run("run-success", AgentRunStatus::Success); + run_success.session_id = Some("session-a".to_string()); + run_success.started_at = "2026-03-06T10:00:00Z".to_string(); + run_success.created_at = run_success.started_at.clone(); + run_success.updated_at = run_success.started_at.clone(); + AgentRunDao::create_run(&conn, &run_success).expect("写入 run-success 失败"); + + let mut run_running = sample_run("run-running", AgentRunStatus::Running); + run_running.session_id = Some("session-a".to_string()); + run_running.started_at = "2026-03-06T11:00:00Z".to_string(); + run_running.created_at = run_running.started_at.clone(); + run_running.updated_at = run_running.started_at.clone(); + AgentRunDao::create_run(&conn, &run_running).expect("写入 run-running 失败"); + + let mut run_error = sample_run("run-error", AgentRunStatus::Error); + run_error.session_id = Some("session-a".to_string()); + run_error.started_at = "2026-03-06T12:00:00Z".to_string(); + run_error.created_at = run_error.started_at.clone(); + run_error.updated_at = run_error.started_at.clone(); + AgentRunDao::create_run(&conn, &run_error).expect("写入 run-error 失败"); + + let mut run_timeout = sample_run("run-timeout", AgentRunStatus::Timeout); + run_timeout.session_id = Some("session-a".to_string()); + run_timeout.started_at = "2026-03-06T13:00:00Z".to_string(); + run_timeout.created_at = run_timeout.started_at.clone(); + run_timeout.updated_at = run_timeout.started_at.clone(); + AgentRunDao::create_run(&conn, &run_timeout).expect("写入 run-timeout 失败"); + + let first_page = AgentRunDao::list_terminal_runs_by_session(&conn, "session-a", 2, 0) + .expect("查询第一页终态记录失败"); + assert_eq!(first_page.len(), 2); + assert_eq!(first_page[0].id, "run-timeout"); + assert_eq!(first_page[1].id, "run-error"); + + let second_page = AgentRunDao::list_terminal_runs_by_session(&conn, "session-a", 2, 2) + .expect("查询第二页终态记录失败"); + assert_eq!(second_page.len(), 1); + assert_eq!(second_page[0].id, "run-success"); + } } diff --git a/src-tauri/crates/core/src/database/dao/agent_timeline.rs b/src-tauri/crates/core/src/database/dao/agent_timeline.rs new file mode 100644 index 000000000..a1092330b --- /dev/null +++ b/src-tauri/crates/core/src/database/dao/agent_timeline.rs @@ -0,0 +1,521 @@ +//! Agent 线程时间线数据访问层 +//! +//! 在现有 `agent_sessions` 基础上补充 turn / item 一等事件持久化。 + +use rusqlite::{params, Connection}; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum AgentThreadTurnStatus { + Running, + Completed, + Failed, + Aborted, +} + +impl AgentThreadTurnStatus { + pub fn as_str(&self) -> &'static str { + match self { + Self::Running => "running", + Self::Completed => "completed", + Self::Failed => "failed", + Self::Aborted => "aborted", + } + } +} + +impl TryFrom<&str> for AgentThreadTurnStatus { + type Error = String; + + fn try_from(value: &str) -> Result { + match value { + "running" => Ok(Self::Running), + "completed" => Ok(Self::Completed), + "failed" => Ok(Self::Failed), + "aborted" => Ok(Self::Aborted), + other => Err(format!("未知 turn 状态: {other}")), + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum AgentThreadItemStatus { + InProgress, + Completed, + Failed, +} + +impl AgentThreadItemStatus { + pub fn as_str(&self) -> &'static str { + match self { + Self::InProgress => "in_progress", + Self::Completed => "completed", + Self::Failed => "failed", + } + } +} + +impl TryFrom<&str> for AgentThreadItemStatus { + type Error = String; + + fn try_from(value: &str) -> Result { + match value { + "in_progress" => Ok(Self::InProgress), + "completed" => Ok(Self::Completed), + "failed" => Ok(Self::Failed), + other => Err(format!("未知 item 状态: {other}")), + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct AgentRequestOption { + pub label: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub description: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct AgentRequestQuestion { + pub question: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub header: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub options: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub multi_select: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum AgentThreadItemPayload { + UserMessage { + content: String, + }, + AgentMessage { + text: String, + #[serde(skip_serializing_if = "Option::is_none")] + phase: Option, + }, + Plan { + text: String, + }, + Reasoning { + text: String, + #[serde(skip_serializing_if = "Option::is_none")] + summary: Option>, + }, + ToolCall { + tool_name: String, + #[serde(skip_serializing_if = "Option::is_none")] + arguments: Option, + #[serde(skip_serializing_if = "Option::is_none")] + output: Option, + #[serde(skip_serializing_if = "Option::is_none")] + success: Option, + #[serde(skip_serializing_if = "Option::is_none")] + error: Option, + #[serde(skip_serializing_if = "Option::is_none")] + metadata: Option, + }, + CommandExecution { + command: String, + cwd: String, + #[serde(skip_serializing_if = "Option::is_none")] + aggregated_output: Option, + #[serde(skip_serializing_if = "Option::is_none")] + exit_code: Option, + #[serde(skip_serializing_if = "Option::is_none")] + error: Option, + }, + WebSearch { + #[serde(skip_serializing_if = "Option::is_none")] + query: Option, + #[serde(skip_serializing_if = "Option::is_none")] + action: Option, + #[serde(skip_serializing_if = "Option::is_none")] + output: Option, + }, + ApprovalRequest { + request_id: String, + action_type: String, + #[serde(skip_serializing_if = "Option::is_none")] + prompt: Option, + #[serde(skip_serializing_if = "Option::is_none")] + tool_name: Option, + #[serde(skip_serializing_if = "Option::is_none")] + arguments: Option, + #[serde(skip_serializing_if = "Option::is_none")] + response: Option, + }, + RequestUserInput { + request_id: String, + action_type: String, + #[serde(skip_serializing_if = "Option::is_none")] + prompt: Option, + #[serde(skip_serializing_if = "Option::is_none")] + questions: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + response: Option, + }, + FileArtifact { + path: String, + source: String, + #[serde(skip_serializing_if = "Option::is_none")] + content: Option, + #[serde(skip_serializing_if = "Option::is_none")] + metadata: Option, + }, + SubagentActivity { + status_label: String, + #[serde(skip_serializing_if = "Option::is_none")] + title: Option, + #[serde(skip_serializing_if = "Option::is_none")] + summary: Option, + #[serde(skip_serializing_if = "Option::is_none")] + role: Option, + #[serde(skip_serializing_if = "Option::is_none")] + model: Option, + }, + Warning { + message: String, + #[serde(skip_serializing_if = "Option::is_none")] + code: Option, + }, + Error { + message: String, + }, + TurnSummary { + text: String, + }, +} + +impl AgentThreadItemPayload { + pub fn kind(&self) -> &'static str { + match self { + Self::UserMessage { .. } => "user_message", + Self::AgentMessage { .. } => "agent_message", + Self::Plan { .. } => "plan", + Self::Reasoning { .. } => "reasoning", + Self::ToolCall { .. } => "tool_call", + Self::CommandExecution { .. } => "command_execution", + Self::WebSearch { .. } => "web_search", + Self::ApprovalRequest { .. } => "approval_request", + Self::RequestUserInput { .. } => "request_user_input", + Self::FileArtifact { .. } => "file_artifact", + Self::SubagentActivity { .. } => "subagent_activity", + Self::Warning { .. } => "warning", + Self::Error { .. } => "error", + Self::TurnSummary { .. } => "turn_summary", + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct AgentThreadTurn { + pub id: String, + pub thread_id: String, + pub prompt_text: String, + pub status: AgentThreadTurnStatus, + pub started_at: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub completed_at: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub error_message: Option, + pub created_at: String, + pub updated_at: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct AgentThreadItem { + pub id: String, + pub thread_id: String, + pub turn_id: String, + pub sequence: i64, + pub status: AgentThreadItemStatus, + pub started_at: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub completed_at: Option, + pub updated_at: String, + #[serde(flatten)] + pub payload: AgentThreadItemPayload, +} + +pub struct AgentTimelineDao; + +impl AgentTimelineDao { + pub fn create_turn(conn: &Connection, turn: &AgentThreadTurn) -> Result<(), rusqlite::Error> { + conn.execute( + "INSERT INTO agent_thread_turns ( + id, session_id, prompt_text, status, started_at, completed_at, + error_message, created_at, updated_at + ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + params![ + turn.id, + turn.thread_id, + turn.prompt_text, + turn.status.as_str(), + turn.started_at, + turn.completed_at, + turn.error_message, + turn.created_at, + turn.updated_at, + ], + )?; + Ok(()) + } + + pub fn update_turn_status( + conn: &Connection, + turn_id: &str, + status: AgentThreadTurnStatus, + completed_at: Option<&str>, + error_message: Option<&str>, + updated_at: &str, + ) -> Result<(), rusqlite::Error> { + conn.execute( + "UPDATE agent_thread_turns + SET status = ?1, + completed_at = COALESCE(?2, completed_at), + error_message = COALESCE(?3, error_message), + updated_at = ?4 + WHERE id = ?5", + params![ + status.as_str(), + completed_at, + error_message, + updated_at, + turn_id, + ], + )?; + Ok(()) + } + + pub fn list_turns_by_thread( + conn: &Connection, + thread_id: &str, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, session_id, prompt_text, status, started_at, completed_at, + error_message, created_at, updated_at + FROM agent_thread_turns + WHERE session_id = ?1 + ORDER BY started_at ASC, id ASC", + )?; + + let rows = stmt.query_map(params![thread_id], |row| { + let status_raw: String = row.get(3)?; + let status = AgentThreadTurnStatus::try_from(status_raw.as_str()).map_err(|_| { + rusqlite::Error::InvalidColumnType(3, "status".into(), rusqlite::types::Type::Text) + })?; + + Ok(AgentThreadTurn { + id: row.get(0)?, + thread_id: row.get(1)?, + prompt_text: row.get(2)?, + status, + started_at: row.get(4)?, + completed_at: row.get(5)?, + error_message: row.get(6)?, + created_at: row.get(7)?, + updated_at: row.get(8)?, + }) + })?; + + rows.collect() + } + + pub fn upsert_item(conn: &Connection, item: &AgentThreadItem) -> Result<(), rusqlite::Error> { + let payload_json = serde_json::to_string(&item.payload) + .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?; + + conn.execute( + "INSERT INTO agent_thread_items ( + id, session_id, turn_id, sequence, item_type, status, started_at, + completed_at, updated_at, payload_json + ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10) + ON CONFLICT(id) DO UPDATE SET + session_id = excluded.session_id, + turn_id = excluded.turn_id, + sequence = excluded.sequence, + item_type = excluded.item_type, + status = excluded.status, + started_at = excluded.started_at, + completed_at = excluded.completed_at, + updated_at = excluded.updated_at, + payload_json = excluded.payload_json", + params![ + item.id, + item.thread_id, + item.turn_id, + item.sequence, + item.payload.kind(), + item.status.as_str(), + item.started_at, + item.completed_at, + item.updated_at, + payload_json, + ], + )?; + Ok(()) + } + + pub fn get_item( + conn: &Connection, + item_id: &str, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, session_id, turn_id, sequence, status, started_at, completed_at, + updated_at, payload_json + FROM agent_thread_items + WHERE id = ?1", + )?; + + let mut rows = stmt.query(params![item_id])?; + if let Some(row) = rows.next()? { + Ok(Some(Self::row_to_item(row)?)) + } else { + Ok(None) + } + } + + pub fn list_items_by_thread( + conn: &Connection, + thread_id: &str, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, session_id, turn_id, sequence, status, started_at, completed_at, + updated_at, payload_json + FROM agent_thread_items + WHERE session_id = ?1 + ORDER BY ( + SELECT started_at + FROM agent_thread_turns + WHERE agent_thread_turns.id = agent_thread_items.turn_id + ) ASC, sequence ASC, id ASC", + )?; + + let rows = stmt.query_map(params![thread_id], Self::row_to_item)?; + rows.collect() + } + + fn row_to_item(row: &rusqlite::Row<'_>) -> Result { + let status_raw: String = row.get(4)?; + let status = AgentThreadItemStatus::try_from(status_raw.as_str()).map_err(|_| { + rusqlite::Error::InvalidColumnType(4, "status".into(), rusqlite::types::Type::Text) + })?; + let payload_json: String = row.get(8)?; + let payload: AgentThreadItemPayload = + serde_json::from_str(&payload_json).map_err(|_| { + rusqlite::Error::InvalidColumnType( + 8, + "payload_json".into(), + rusqlite::types::Type::Text, + ) + })?; + + Ok(AgentThreadItem { + id: row.get(0)?, + thread_id: row.get(1)?, + turn_id: row.get(2)?, + sequence: row.get(3)?, + status, + started_at: row.get(5)?, + completed_at: row.get(6)?, + updated_at: row.get(7)?, + payload, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::database::schema::create_tables; + use rusqlite::Connection; + + fn setup_conn() -> Connection { + let conn = Connection::open_in_memory().expect("创建内存数据库失败"); + create_tables(&conn).expect("创建表结构失败"); + conn.execute( + "INSERT INTO agent_sessions (id, model, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", + params!["thread-1", "general:test", "2026-03-13T00:00:00Z", "2026-03-13T00:00:00Z"], + ) + .unwrap(); + conn + } + + #[test] + fn create_turn_and_upsert_item_should_roundtrip() { + let conn = setup_conn(); + let turn = AgentThreadTurn { + id: "turn-1".to_string(), + thread_id: "thread-1".to_string(), + prompt_text: "帮我做个计划".to_string(), + status: AgentThreadTurnStatus::Running, + started_at: "2026-03-13T01:00:00Z".to_string(), + completed_at: None, + error_message: None, + created_at: "2026-03-13T01:00:00Z".to_string(), + updated_at: "2026-03-13T01:00:00Z".to_string(), + }; + + AgentTimelineDao::create_turn(&conn, &turn).unwrap(); + + let item = AgentThreadItem { + id: "item-1".to_string(), + thread_id: "thread-1".to_string(), + turn_id: "turn-1".to_string(), + sequence: 1, + status: AgentThreadItemStatus::Completed, + started_at: "2026-03-13T01:00:01Z".to_string(), + completed_at: Some("2026-03-13T01:00:01Z".to_string()), + updated_at: "2026-03-13T01:00:01Z".to_string(), + payload: AgentThreadItemPayload::UserMessage { + content: "帮我做个计划".to_string(), + }, + }; + + AgentTimelineDao::upsert_item(&conn, &item).unwrap(); + + let turns = AgentTimelineDao::list_turns_by_thread(&conn, "thread-1").unwrap(); + let items = AgentTimelineDao::list_items_by_thread(&conn, "thread-1").unwrap(); + + assert_eq!(turns.len(), 1); + assert_eq!(items.len(), 1); + assert_eq!(items[0], item); + } + + #[test] + fn update_turn_status_should_persist_terminal_state() { + let conn = setup_conn(); + let turn = AgentThreadTurn { + id: "turn-2".to_string(), + thread_id: "thread-1".to_string(), + prompt_text: "继续".to_string(), + status: AgentThreadTurnStatus::Running, + started_at: "2026-03-13T02:00:00Z".to_string(), + completed_at: None, + error_message: None, + created_at: "2026-03-13T02:00:00Z".to_string(), + updated_at: "2026-03-13T02:00:00Z".to_string(), + }; + AgentTimelineDao::create_turn(&conn, &turn).unwrap(); + + AgentTimelineDao::update_turn_status( + &conn, + "turn-2", + AgentThreadTurnStatus::Failed, + Some("2026-03-13T02:00:03Z"), + Some("boom"), + "2026-03-13T02:00:03Z", + ) + .unwrap(); + + let turns = AgentTimelineDao::list_turns_by_thread(&conn, "thread-1").unwrap(); + assert_eq!(turns[0].status, AgentThreadTurnStatus::Failed); + assert_eq!(turns[0].error_message.as_deref(), Some("boom")); + } +} diff --git a/src-tauri/crates/core/src/database/dao/general_chat.rs b/src-tauri/crates/core/src/database/dao/general_chat.rs deleted file mode 100644 index fb72698a9..000000000 --- a/src-tauri/crates/core/src/database/dao/general_chat.rs +++ /dev/null @@ -1,700 +0,0 @@ -//! 通用对话会话和消息的数据访问层 -//! -//! 提供通用对话会话和消息的持久化存储功能 -//! -//! ## 主要功能 -//! - `create_session` - 创建新会话 -//! - `list_sessions` - 获取会话列表 -//! - `get_session` - 获取单个会话 -//! - `delete_session` - 删除会话 -//! - `rename_session` - 重命名会话 -//! - `update_session_time` - 更新会话时间 -//! - `add_message` - 添加消息 -//! - `get_messages` - 获取消息列表 -//! - `get_message_count` - 获取消息数量 -//! - `delete_messages` - 删除会话消息 - -use crate::general_chat::{ChatMessage, ChatSession, ContentBlock, MessageRole}; -use rusqlite::{params, Connection}; - -pub struct GeneralChatDao; - -impl GeneralChatDao { - // ==================== 会话 CRUD ==================== - - /// 创建新会话 - pub fn create_session(conn: &Connection, session: &ChatSession) -> Result<(), rusqlite::Error> { - let metadata_json = session - .metadata - .as_ref() - .map(serde_json::to_string) - .transpose() - .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?; - - conn.execute( - "INSERT INTO general_chat_sessions (id, name, created_at, updated_at, metadata) - VALUES (?1, ?2, ?3, ?4, ?5)", - params![ - session.id, - session.name, - session.created_at, - session.updated_at, - metadata_json, - ], - )?; - Ok(()) - } - - /// 获取所有会话列表(按更新时间降序排列) - pub fn list_sessions(conn: &Connection) -> Result, rusqlite::Error> { - let mut stmt = conn.prepare( - "SELECT id, name, created_at, updated_at, metadata - FROM general_chat_sessions ORDER BY updated_at DESC", - )?; - - let sessions = stmt.query_map([], |row| { - let metadata_json: Option = row.get(4)?; - let metadata = metadata_json - .map(|json| serde_json::from_str(&json)) - .transpose() - .map_err(|e| { - rusqlite::Error::FromSqlConversionFailure( - 4, - rusqlite::types::Type::Text, - Box::new(e), - ) - })?; - - Ok(ChatSession { - id: row.get(0)?, - name: row.get(1)?, - created_at: row.get(2)?, - updated_at: row.get(3)?, - metadata, - }) - })?; - - sessions.collect() - } - - /// 获取单个会话(不包含消息) - pub fn get_session( - conn: &Connection, - session_id: &str, - ) -> Result, rusqlite::Error> { - let mut stmt = conn.prepare( - "SELECT id, name, created_at, updated_at, metadata - FROM general_chat_sessions WHERE id = ?", - )?; - - let mut rows = stmt.query([session_id])?; - - if let Some(row) = rows.next()? { - let metadata_json: Option = row.get(4)?; - let metadata = metadata_json - .map(|json| serde_json::from_str(&json)) - .transpose() - .map_err(|e| { - rusqlite::Error::FromSqlConversionFailure( - 4, - rusqlite::types::Type::Text, - Box::new(e), - ) - })?; - - Ok(Some(ChatSession { - id: row.get(0)?, - name: row.get(1)?, - created_at: row.get(2)?, - updated_at: row.get(3)?, - metadata, - })) - } else { - Ok(None) - } - } - - /// 删除会话(消息会通过外键级联删除) - pub fn delete_session(conn: &Connection, session_id: &str) -> Result { - let rows = conn.execute( - "DELETE FROM general_chat_sessions WHERE id = ?", - [session_id], - )?; - Ok(rows > 0) - } - - /// 重命名会话 - pub fn rename_session( - conn: &Connection, - session_id: &str, - name: &str, - ) -> Result { - let now = chrono::Utc::now().timestamp_millis(); - let rows = conn.execute( - "UPDATE general_chat_sessions SET name = ?, updated_at = ? WHERE id = ?", - params![name, now, session_id], - )?; - Ok(rows > 0) - } - - /// 更新会话的 updated_at 时间 - pub fn update_session_time( - conn: &Connection, - session_id: &str, - updated_at: i64, - ) -> Result<(), rusqlite::Error> { - conn.execute( - "UPDATE general_chat_sessions SET updated_at = ? WHERE id = ?", - params![updated_at, session_id], - )?; - Ok(()) - } - - /// 检查会话是否存在 - pub fn session_exists(conn: &Connection, session_id: &str) -> Result { - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM general_chat_sessions WHERE id = ?", - [session_id], - |row| row.get(0), - )?; - Ok(count > 0) - } - - // ==================== 消息 CRUD ==================== - - /// 添加消息到会话 - pub fn add_message(conn: &Connection, message: &ChatMessage) -> Result<(), rusqlite::Error> { - let blocks_json = message - .blocks - .as_ref() - .map(serde_json::to_string) - .transpose() - .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?; - - let metadata_json = message - .metadata - .as_ref() - .map(serde_json::to_string) - .transpose() - .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?; - - let role_str = 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)", - params![ - message.id, - message.session_id, - role_str, - message.content, - blocks_json, - message.status, - message.created_at, - metadata_json, - ], - )?; - - // 更新会话的 updated_at - Self::update_session_time(conn, &message.session_id, message.created_at)?; - - Ok(()) - } - - /// 获取会话的消息列表 - /// - /// # Arguments - /// * `conn` - 数据库连接 - /// * `session_id` - 会话 ID - /// * `limit` - 限制返回数量(可选) - /// * `before_id` - 在此消息 ID 之前的消息(用于分页) - pub fn get_messages( - conn: &Connection, - session_id: &str, - limit: Option, - before_id: Option<&str>, - ) -> Result, rusqlite::Error> { - 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(_bid)) => { - 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 messages = if before_id.is_some() { - stmt.query_map(params![session_id, before_id], Self::map_message_row)? - } else { - stmt.query_map(params![session_id], Self::map_message_row)? - }; - - let mut result: Vec = messages.collect::, _>>()?; - - // 如果有 limit,结果是倒序的,需要反转 - if limit.is_some() { - result.reverse(); - } - - Ok(result) - } - - /// 获取会话的消息数量 - pub fn get_message_count(conn: &Connection, session_id: &str) -> Result { - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM general_chat_messages WHERE session_id = ?", - [session_id], - |row| row.get(0), - )?; - Ok(count) - } - - /// 删除会话的所有消息 - pub fn delete_messages(conn: &Connection, session_id: &str) -> Result<(), rusqlite::Error> { - conn.execute( - "DELETE FROM general_chat_messages WHERE session_id = ?", - [session_id], - )?; - Ok(()) - } - - // ==================== 辅助方法 ==================== - - /// 从数据库行映射到 ChatMessage - fn map_message_row(row: &rusqlite::Row) -> Result { - let role_str: String = row.get(2)?; - let role = match role_str.as_str() { - "user" => MessageRole::User, - "assistant" => MessageRole::Assistant, - "system" => MessageRole::System, - _ => MessageRole::User, - }; - - 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, - content: row.get(3)?, - blocks, - status: row.get(5)?, - created_at: row.get(6)?, - metadata, - }) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::general_chat::MessageRole; - use rusqlite::Connection; - - fn setup_test_db() -> Connection { - let conn = Connection::open_in_memory().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 - } - - fn create_test_session(id: &str, name: &str) -> ChatSession { - let now = chrono::Utc::now().timestamp_millis(); - ChatSession { - id: id.to_string(), - name: name.to_string(), - created_at: now, - updated_at: now, - metadata: None, - } - } - - fn create_test_message( - id: &str, - session_id: &str, - role: MessageRole, - content: &str, - ) -> ChatMessage { - let now = chrono::Utc::now().timestamp_millis(); - ChatMessage { - id: id.to_string(), - session_id: session_id.to_string(), - role, - content: content.to_string(), - blocks: None, - status: "complete".to_string(), - created_at: now, - metadata: None, - } - } - - fn create_test_message_with_timestamp( - id: &str, - session_id: &str, - role: MessageRole, - content: &str, - created_at: i64, - ) -> ChatMessage { - ChatMessage { - id: id.to_string(), - session_id: session_id.to_string(), - role, - content: content.to_string(), - blocks: None, - status: "complete".to_string(), - created_at, - metadata: None, - } - } - - #[test] - fn test_create_and_get_session() { - let conn = setup_test_db(); - let session = create_test_session("session-1", "测试会话"); - - GeneralChatDao::create_session(&conn, &session).unwrap(); - - let loaded = GeneralChatDao::get_session(&conn, "session-1").unwrap(); - assert!(loaded.is_some()); - - let loaded = loaded.unwrap(); - assert_eq!(loaded.id, "session-1"); - assert_eq!(loaded.name, "测试会话"); - } - - #[test] - fn test_list_sessions() { - let conn = setup_test_db(); - - let session1 = create_test_session("session-1", "会话1"); - let session2 = create_test_session("session-2", "会话2"); - - GeneralChatDao::create_session(&conn, &session1).unwrap(); - GeneralChatDao::create_session(&conn, &session2).unwrap(); - - let sessions = GeneralChatDao::list_sessions(&conn).unwrap(); - assert_eq!(sessions.len(), 2); - } - - #[test] - fn test_delete_session() { - let conn = setup_test_db(); - let session = create_test_session("session-1", "测试会话"); - - GeneralChatDao::create_session(&conn, &session).unwrap(); - - let deleted = GeneralChatDao::delete_session(&conn, "session-1").unwrap(); - assert!(deleted); - - let loaded = GeneralChatDao::get_session(&conn, "session-1").unwrap(); - assert!(loaded.is_none()); - } - - #[test] - fn test_delete_nonexistent_session() { - let conn = setup_test_db(); - - let deleted = GeneralChatDao::delete_session(&conn, "nonexistent").unwrap(); - assert!(!deleted); - } - - #[test] - fn test_rename_session() { - let conn = setup_test_db(); - let session = create_test_session("session-1", "原名称"); - - GeneralChatDao::create_session(&conn, &session).unwrap(); - - let renamed = GeneralChatDao::rename_session(&conn, "session-1", "新名称").unwrap(); - assert!(renamed); - - let loaded = GeneralChatDao::get_session(&conn, "session-1") - .unwrap() - .unwrap(); - assert_eq!(loaded.name, "新名称"); - } - - #[test] - fn test_session_exists() { - let conn = setup_test_db(); - let session = create_test_session("session-1", "测试会话"); - - assert!(!GeneralChatDao::session_exists(&conn, "session-1").unwrap()); - - GeneralChatDao::create_session(&conn, &session).unwrap(); - - assert!(GeneralChatDao::session_exists(&conn, "session-1").unwrap()); - } - - #[test] - fn test_add_and_get_messages() { - let conn = setup_test_db(); - let session = create_test_session("session-1", "测试会话"); - GeneralChatDao::create_session(&conn, &session).unwrap(); - - let msg1 = create_test_message("msg-1", "session-1", MessageRole::User, "你好"); - let msg2 = create_test_message( - "msg-2", - "session-1", - MessageRole::Assistant, - "你好!有什么可以帮助你的?", - ); - - GeneralChatDao::add_message(&conn, &msg1).unwrap(); - GeneralChatDao::add_message(&conn, &msg2).unwrap(); - - let messages = GeneralChatDao::get_messages(&conn, "session-1", None, None).unwrap(); - assert_eq!(messages.len(), 2); - assert_eq!(messages[0].content, "你好"); - assert_eq!(messages[1].content, "你好!有什么可以帮助你的?"); - } - - #[test] - fn test_get_messages_with_limit() { - let conn = setup_test_db(); - let session = create_test_session("session-1", "测试会话"); - GeneralChatDao::create_session(&conn, &session).unwrap(); - - for i in 1..=5 { - let msg = create_test_message( - &format!("msg-{i}"), - "session-1", - MessageRole::User, - &format!("消息 {i}"), - ); - GeneralChatDao::add_message(&conn, &msg).unwrap(); - } - - let messages = GeneralChatDao::get_messages(&conn, "session-1", Some(3), None).unwrap(); - assert_eq!(messages.len(), 3); - } - - #[test] - fn test_get_message_count() { - let conn = setup_test_db(); - let session = create_test_session("session-1", "测试会话"); - GeneralChatDao::create_session(&conn, &session).unwrap(); - - assert_eq!( - GeneralChatDao::get_message_count(&conn, "session-1").unwrap(), - 0 - ); - - let msg = create_test_message("msg-1", "session-1", MessageRole::User, "你好"); - GeneralChatDao::add_message(&conn, &msg).unwrap(); - - assert_eq!( - GeneralChatDao::get_message_count(&conn, "session-1").unwrap(), - 1 - ); - } - - #[test] - fn test_delete_messages() { - let conn = setup_test_db(); - let session = create_test_session("session-1", "测试会话"); - GeneralChatDao::create_session(&conn, &session).unwrap(); - - let msg = create_test_message("msg-1", "session-1", MessageRole::User, "你好"); - GeneralChatDao::add_message(&conn, &msg).unwrap(); - - GeneralChatDao::delete_messages(&conn, "session-1").unwrap(); - - assert_eq!( - GeneralChatDao::get_message_count(&conn, "session-1").unwrap(), - 0 - ); - } - - #[test] - fn test_message_with_blocks() { - let conn = setup_test_db(); - let session = create_test_session("session-1", "测试会话"); - GeneralChatDao::create_session(&conn, &session).unwrap(); - - let now = chrono::Utc::now().timestamp_millis(); - let msg = ChatMessage { - id: "msg-1".to_string(), - session_id: "session-1".to_string(), - role: MessageRole::Assistant, - content: "这是一段代码:".to_string(), - blocks: Some(vec![ContentBlock { - r#type: "code".to_string(), - content: "fn main() {}".to_string(), - language: Some("rust".to_string()), - filename: None, - mime_type: None, - }]), - status: "complete".to_string(), - created_at: now, - metadata: None, - }; - - GeneralChatDao::add_message(&conn, &msg).unwrap(); - - let messages = GeneralChatDao::get_messages(&conn, "session-1", None, None).unwrap(); - assert_eq!(messages.len(), 1); - - let loaded = &messages[0]; - assert!(loaded.blocks.is_some()); - - let blocks = loaded.blocks.as_ref().unwrap(); - assert_eq!(blocks.len(), 1); - assert_eq!(blocks[0].r#type, "code"); - assert_eq!(blocks[0].language, Some("rust".to_string())); - } - - #[test] - fn test_get_messages_before_id_uses_created_at_pagination() { - let conn = setup_test_db(); - let session = create_test_session("session-1", "测试会话"); - GeneralChatDao::create_session(&conn, &session).unwrap(); - - let oldest = create_test_message_with_timestamp( - "z-message", - "session-1", - MessageRole::User, - "第一条", - 1_700_000_000_001, - ); - let middle = create_test_message_with_timestamp( - "a-message", - "session-1", - MessageRole::Assistant, - "第二条", - 1_700_000_000_002, - ); - let newest = create_test_message_with_timestamp( - "m-message", - "session-1", - MessageRole::User, - "第三条", - 1_700_000_000_003, - ); - - GeneralChatDao::add_message(&conn, &oldest).unwrap(); - GeneralChatDao::add_message(&conn, &middle).unwrap(); - GeneralChatDao::add_message(&conn, &newest).unwrap(); - - let messages = - GeneralChatDao::get_messages(&conn, "session-1", Some(10), Some("a-message")).unwrap(); - - assert_eq!(messages.len(), 1); - assert_eq!(messages[0].id, "z-message"); - } -} diff --git a/src-tauri/crates/core/src/database/dao/mod.rs b/src-tauri/crates/core/src/database/dao/mod.rs index 72138cbb9..43a8ff816 100644 --- a/src-tauri/crates/core/src/database/dao/mod.rs +++ b/src-tauri/crates/core/src/database/dao/mod.rs @@ -1,10 +1,10 @@ pub mod a2ui_form_dao; pub mod agent; pub mod agent_run; +pub mod agent_timeline; pub mod api_key_provider; pub mod brand_persona_dao; pub mod chat; -pub mod general_chat; pub mod heartbeat; pub mod installed_plugins; pub mod material_dao; diff --git a/src-tauri/crates/core/src/database/migration.rs b/src-tauri/crates/core/src/database/migration.rs index 23a430be3..f187d06df 100644 --- a/src-tauri/crates/core/src/database/migration.rs +++ b/src-tauri/crates/core/src/database/migration.rs @@ -1,25 +1,37 @@ -use rusqlite::{params, Connection}; +mod api_key_migration; +mod general_chat_migration; +mod mcp_migration; +mod model_registry_migration; + +use crate::app_paths; +pub use api_key_migration::{ + cleanup_legacy_api_key_credentials, migrate_api_keys_to_pool, migrate_provider_ids, +}; +pub use general_chat_migration::{ + check_general_chat_migration_status, is_general_chat_migration_completed, + migrate_general_chat_to_unified, GeneralChatMigrationStatus, + GENERAL_CHAT_MIGRATION_COMPLETED_KEY, +}; +pub use mcp_migration::{migrate_mcp_created_at_to_integer, migrate_mcp_proxycast_enabled}; +pub use model_registry_migration::{ + check_model_registry_version, clear_model_registry_refresh_flag, + is_model_registry_refresh_needed, mark_model_registry_refresh_needed, +}; + +pub(crate) use super::migration_support::{ + clear_setting, is_true_setting, mark_true_setting, read_setting_value, upsert_setting, +}; +use rusqlite::Connection; /// 从旧的 JSON 配置迁移数据到 SQLite #[allow(dead_code)] pub fn migrate_from_json(conn: &Connection) -> Result<(), String> { - // 检查是否已经迁移过 - let migrated: bool = conn - .query_row( - "SELECT value FROM settings WHERE key = 'migrated_from_json'", - [], - |row| row.get::<_, String>(0), - ) - .map(|v| v == "true") - .unwrap_or(false); - - if migrated { + if is_true_setting(conn, "migrated_from_json") { return Ok(()); } // 读取旧配置文件(历史路径) - let home = dirs::home_dir().ok_or_else(|| "无法获取主目录".to_string())?; - let config_path = home.join(".proxycast").join("config.json"); + let config_path = app_paths::legacy_home_dir()?.join("config.json"); if config_path.exists() { // 备份旧配置,避免误覆盖 @@ -30,929 +42,13 @@ pub fn migrate_from_json(conn: &Connection) -> Result<(), String> { } return Err( - "检测到旧版 config.json(~/.proxycast/config.json),当前版本尚未支持自动迁移。请手动导出/重建配置后再启动。" + "检测到旧版 config.json(旧 Home 历史目录),当前版本尚未支持自动迁移。请手动导出或重建配置后再启动。" .to_string(), ); } // 标记迁移完成 - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_from_json', 'true')", - [], - ) - .map_err(|e| e.to_string())?; + mark_true_setting(conn, "migrated_from_json")?; Ok(()) } - -/// 将 api_keys 表中的数据迁移到 provider_pool_credentials 表 -/// -/// 迁移逻辑: -/// 1. 读取 api_keys 表中的所有 API Key -/// 2. 根据 provider_id 查找对应的 api_key_providers 配置 -/// 3. 将 API Key 转换为 CredentialData::OpenAIKey 或 CredentialData::ClaudeKey -/// 4. 插入到 provider_pool_credentials 表 -pub fn migrate_api_keys_to_pool(conn: &Connection) -> Result { - // 检查是否已经迁移过 - let migrated: bool = conn - .query_row( - "SELECT value FROM settings WHERE key = 'migrated_api_keys_to_pool'", - [], - |row| row.get::<_, String>(0), - ) - .map(|v| v == "true") - .unwrap_or(false); - - if migrated { - tracing::debug!("[迁移] API Keys 已迁移过,跳过"); - return Ok(0); - } - - tracing::info!("[迁移] 开始将 api_keys 迁移到 provider_pool_credentials"); - - // 查询所有 API Keys 及其对应的 Provider 信息 - let mut stmt = conn - .prepare( - "SELECT k.id, k.provider_id, k.api_key_encrypted, k.alias, k.enabled, - k.usage_count, k.error_count, k.last_used_at, k.created_at, - p.type, p.api_host, p.name as provider_name - FROM api_keys k - JOIN api_key_providers p ON k.provider_id = p.id - ORDER BY k.created_at ASC", - ) - .map_err(|e| format!("准备查询语句失败: {e}"))?; - - let rows = stmt - .query_map([], |row| { - Ok(ApiKeyMigrationRow { - id: row.get(0)?, - provider_id: row.get(1)?, - api_key_encrypted: row.get(2)?, - alias: row.get(3)?, - enabled: row.get(4)?, - usage_count: row.get::<_, i64>(5)? as u64, - error_count: row.get::<_, i64>(6)? as u32, - last_used_at: row.get(7)?, - created_at: row.get(8)?, - provider_type: row.get(9)?, - api_host: row.get(10)?, - provider_name: row.get(11)?, - }) - }) - .map_err(|e| format!("查询 API Keys 失败: {e}"))?; - - let mut migrated_count = 0; - let now = chrono::Utc::now().timestamp(); - - for row_result in rows { - let row = row_result.map_err(|e| format!("读取行数据失败: {e}"))?; - - // 检查是否已存在相同的凭证(通过 api_key_encrypted 判断) - let exists: bool = conn - .query_row( - "SELECT COUNT(*) > 0 FROM provider_pool_credentials - WHERE credential_data LIKE ?1", - params![format!("%{}%", row.api_key_encrypted)], - |r| r.get(0), - ) - .unwrap_or(false); - - if exists { - tracing::debug!( - "[迁移] 跳过已存在的 API Key: {} (provider: {})", - row.alias.as_deref().unwrap_or(&row.id), - row.provider_id - ); - continue; - } - - // 根据 provider_type 确定 pool_provider_type 和 credential_data - let (pool_provider_type, credential_data) = match row.provider_type.to_lowercase().as_str() - { - "anthropic" => { - let cred = serde_json::json!({ - "type": "claude_key", - "api_key": row.api_key_encrypted, - "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) } - }); - ("claude", cred) - } - "openai" | "openai-response" => { - let cred = serde_json::json!({ - "type": "openai_key", - "api_key": row.api_key_encrypted, - "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) } - }); - ("openai", cred) - } - "gemini" => { - let cred = serde_json::json!({ - "type": "gemini_api_key", - "api_key": row.api_key_encrypted, - "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) }, - "excluded_models": [] - }); - ("gemini_api_key", cred) - } - "vertex" | "vertexai" => { - let cred = serde_json::json!({ - "type": "vertex_key", - "api_key": row.api_key_encrypted, - "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) }, - "model_aliases": {} - }); - ("vertex", cred) - } - // 其他类型默认作为 OpenAI 兼容处理 - _ => { - let cred = serde_json::json!({ - "type": "openai_key", - "api_key": row.api_key_encrypted, - "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) } - }); - ("openai", cred) - } - }; - - // 生成名称:优先使用 alias,否则使用 provider_name - let name = row - .alias - .clone() - .or_else(|| Some(format!("{} (迁移)", row.provider_name))); - - // 解析时间 - let created_at_ts = row - .created_at - .as_ref() - .and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok()) - .map(|dt| dt.timestamp()) - .unwrap_or(now); - - let last_used_ts = row - .last_used_at - .as_ref() - .and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok()) - .map(|dt| dt.timestamp()); - - // 插入到 provider_pool_credentials - let uuid = uuid::Uuid::new_v4().to_string(); - let credential_json = credential_data.to_string(); - - conn.execute( - "INSERT INTO provider_pool_credentials - (uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, usage_count, error_count, - last_used, last_error_time, last_error_message, last_health_check_time, - last_health_check_model, created_at, updated_at, source, proxy_url) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20)", - params![ - uuid, - pool_provider_type, - credential_json, - name, - true, // is_healthy - !row.enabled, // is_disabled (反转 enabled) - true, // check_health - Option::::None, // check_model_name - "[]", // not_supported_models - row.usage_count as i64, - row.error_count as i32, - last_used_ts, - Option::::None, // last_error_time - Option::::None, // last_error_message - Option::::None, // last_health_check_time - Option::::None, // last_health_check_model - created_at_ts, - now, - "imported", // source: 标记为导入 - Option::::None, // proxy_url - ], - ) - .map_err(|e| format!("插入凭证失败: {e}"))?; - - tracing::info!( - "[迁移] 已迁移 API Key: {} -> {} (provider_type: {})", - row.alias.as_deref().unwrap_or(&row.id), - uuid, - pool_provider_type - ); - - migrated_count += 1; - } - - // 标记迁移完成 - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_api_keys_to_pool', 'true')", - [], - ) - .map_err(|e| format!("标记迁移完成失败: {e}"))?; - - tracing::info!("[迁移] API Keys 迁移完成,共迁移 {} 条记录", migrated_count); - - Ok(migrated_count) -} - -/// API Key 迁移行数据 -struct ApiKeyMigrationRow { - id: String, - provider_id: String, - api_key_encrypted: String, - alias: Option, - enabled: bool, - usage_count: u64, - error_count: u32, - last_used_at: Option, - created_at: Option, - provider_type: String, - api_host: String, - provider_name: String, -} - -/// 迁移旧的 Provider ID 到新的 ID -/// -/// 修复 system_providers.rs 中 Provider ID 与模型注册表 JSON 文件名不匹配的问题。 -/// 例如:silicon -> siliconflow, gemini -> google 等 -pub fn migrate_provider_ids(conn: &Connection) -> Result { - // 检查是否已经迁移过 - let migrated: bool = conn - .query_row( - "SELECT value FROM settings WHERE key = 'migrated_provider_ids_v1'", - [], - |row| row.get::<_, String>(0), - ) - .map(|v| v == "true") - .unwrap_or(false); - - if migrated { - tracing::debug!("[迁移] Provider ID 已迁移过,跳过"); - return Ok(0); - } - - tracing::info!("[迁移] 开始迁移旧的 Provider ID"); - - // 定义需要迁移的 ID 映射(旧 ID -> 新 ID) - let id_mappings = [ - ("silicon", "siliconflow"), - ("gemini", "google"), - ("zhipu", "zhipuai"), - ("dashscope", "alibaba"), - ("moonshot", "moonshotai"), - ("grok", "xai"), - ("github", "github-models"), - ("copilot", "github-copilot"), - ("vertexai", "google-vertex"), - ("aws-bedrock", "amazon-bedrock"), - ("together", "togetherai"), - ("fireworks", "fireworks-ai"), - ("mimo", "xiaomi"), - ]; - - let mut migrated_count = 0; - - for (old_id, new_id) in &id_mappings { - // 检查旧 ID 是否存在 - let old_exists: bool = conn - .query_row( - "SELECT COUNT(*) > 0 FROM api_key_providers WHERE id = ?1", - params![old_id], - |r| r.get(0), - ) - .unwrap_or(false); - - if !old_exists { - continue; - } - - // 检查新 ID 是否存在 - let new_exists: bool = conn - .query_row( - "SELECT COUNT(*) > 0 FROM api_key_providers WHERE id = ?1", - params![new_id], - |r| r.get(0), - ) - .unwrap_or(false); - - // 检查旧 ID 是否有 API Keys - let has_keys: bool = conn - .query_row( - "SELECT COUNT(*) > 0 FROM api_keys WHERE provider_id = ?1", - params![old_id], - |r| r.get(0), - ) - .unwrap_or(false); - - if has_keys { - // 如果旧 ID 有 API Keys,需要迁移到新 ID - if new_exists { - // 新 ID 已存在,将 API Keys 迁移过去 - conn.execute( - "UPDATE api_keys SET provider_id = ?1 WHERE provider_id = ?2", - params![new_id, old_id], - ) - .map_err(|e| format!("迁移 API Keys 失败: {e}"))?; - - tracing::info!("[迁移] 已将 {} 的 API Keys 迁移到 {}", old_id, new_id); - } else { - // 新 ID 不存在,直接更新旧 ID - conn.execute( - "UPDATE api_key_providers SET id = ?1 WHERE id = ?2", - params![new_id, old_id], - ) - .map_err(|e| format!("更新 Provider ID 失败: {e}"))?; - - conn.execute( - "UPDATE api_keys SET provider_id = ?1 WHERE provider_id = ?2", - params![new_id, old_id], - ) - .map_err(|e| format!("更新 API Keys provider_id 失败: {e}"))?; - - tracing::info!("[迁移] 已将 Provider {} 重命名为 {}", old_id, new_id); - migrated_count += 1; - continue; - } - } - - // 删除旧的 Provider(无论是否有 API Keys,因为 Keys 已迁移) - conn.execute( - "DELETE FROM api_key_providers WHERE id = ?1", - params![old_id], - ) - .map_err(|e| format!("删除旧 Provider 失败: {e}"))?; - - tracing::info!("[迁移] 已删除旧 Provider: {}", old_id); - migrated_count += 1; - } - - // 标记迁移完成 - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_provider_ids_v1', 'true')", - [], - ) - .map_err(|e| format!("标记迁移完成失败: {e}"))?; - - if migrated_count > 0 { - tracing::info!( - "[迁移] Provider ID 迁移完成,共处理 {} 个 Provider", - migrated_count - ); - } - - Ok(migrated_count) -} - -/// 清理旧的 API Key 凭证(OpenAIKey 和 ClaudeKey 类型) -/// -/// 这些凭证是通过旧的 UI 添加的,现在已经被新的 API Key Provider 系统取代。 -/// 此函数会删除 provider_pool_credentials 表中的 openai_key 和 claude_key 类型凭证。 -pub fn cleanup_legacy_api_key_credentials(conn: &Connection) -> Result { - // 检查是否已经清理过 - let cleaned: bool = conn - .query_row( - "SELECT value FROM settings WHERE key = 'cleaned_legacy_api_key_credentials'", - [], - |row| row.get::<_, String>(0), - ) - .map(|v| v == "true") - .unwrap_or(false); - - if cleaned { - tracing::debug!("[清理] 旧 API Key 凭证已清理过,跳过"); - return Ok(0); - } - - tracing::info!("[清理] 开始清理旧的 API Key 凭证(openai_key, claude_key 类型)"); - - // 查询需要清理的凭证数量 - let count: i64 = conn - .query_row( - "SELECT COUNT(*) FROM provider_pool_credentials - WHERE credential_data LIKE '%\"type\":\"openai_key\"%' - OR credential_data LIKE '%\"type\":\"claude_key\"%'", - [], - |row| row.get(0), - ) - .unwrap_or(0); - - if count == 0 { - tracing::info!("[清理] 没有需要清理的旧 API Key 凭证"); - // 标记清理完成 - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('cleaned_legacy_api_key_credentials', 'true')", - [], - ) - .map_err(|e| format!("标记清理完成失败: {e}"))?; - return Ok(0); - } - - // 记录将要删除的凭证信息 - let mut stmt = conn - .prepare( - "SELECT uuid, name, provider_type, credential_data - FROM provider_pool_credentials - WHERE credential_data LIKE '%\"type\":\"openai_key\"%' - OR credential_data LIKE '%\"type\":\"claude_key\"%'", - ) - .map_err(|e| format!("准备查询语句失败: {e}"))?; - - let rows = stmt - .query_map([], |row| { - Ok(( - row.get::<_, String>(0)?, - row.get::<_, Option>(1)?, - row.get::<_, String>(2)?, - )) - }) - .map_err(|e| format!("查询旧凭证失败: {e}"))?; - - for (uuid, name, provider_type) in rows.into_iter().filter_map(|row_result| row_result.ok()) { - tracing::info!( - "[清理] 将删除旧凭证: {} (name: {}, type: {})", - uuid, - name.as_deref().unwrap_or("未命名"), - provider_type - ); - } - - // 删除旧的 API Key 凭证 - let deleted = conn - .execute( - "DELETE FROM provider_pool_credentials - WHERE credential_data LIKE '%\"type\":\"openai_key\"%' - OR credential_data LIKE '%\"type\":\"claude_key\"%'", - [], - ) - .map_err(|e| format!("删除旧凭证失败: {e}"))?; - - // 标记清理完成 - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('cleaned_legacy_api_key_credentials', 'true')", - [], - ) - .map_err(|e| format!("标记清理完成失败: {e}"))?; - - tracing::info!("[清理] 旧 API Key 凭证清理完成,共删除 {} 条记录", deleted); - - Ok(deleted) -} - -/// 修复历史 MCP 导入数据:补齐 enabled_proxycast -/// -/// 早期版本从 Claude/Codex/Gemini 导入 MCP 时,默认写入 enabled_proxycast=0, -/// 导致 ProxyCast 本身不会使用这些服务器。 -/// -/// 迁移策略: -/// - 仅处理 enabled_proxycast=0 的记录 -/// - 且至少在一个外部应用中启用(enabled_claude/codex/gemini 任一为 1) -/// - 将 enabled_proxycast 设为 1 -pub fn migrate_mcp_proxycast_enabled(conn: &Connection) -> Result { - // 检查是否已经迁移过 - let migrated: bool = conn - .query_row( - "SELECT value FROM settings WHERE key = 'migrated_mcp_proxycast_enabled'", - [], - |row| row.get::<_, String>(0), - ) - .map(|v| v == "true") - .unwrap_or(false); - - if migrated { - tracing::debug!("[迁移] MCP proxycast 启用状态已迁移过,跳过"); - return Ok(0); - } - - let updated = conn - .execute( - "UPDATE mcp_servers - SET enabled_proxycast = 1 - WHERE enabled_proxycast = 0 - AND (enabled_claude = 1 OR enabled_codex = 1 OR enabled_gemini = 1)", - [], - ) - .map_err(|e| format!("修复 MCP enabled_proxycast 失败: {e}"))?; - - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_mcp_proxycast_enabled', 'true')", - [], - ) - .map_err(|e| format!("标记 MCP proxycast 迁移完成失败: {e}"))?; - - tracing::info!( - "[迁移] MCP proxycast 启用状态修复完成,更新 {} 条记录", - updated - ); - - Ok(updated) -} - -/// 归一化 mcp_servers.created_at 字段为 INTEGER 时间戳 -/// -/// 历史版本曾写入 RFC3339 文本,导致下游按 i64 读取时出现类型异常。 -/// 迁移策略: -/// - 纯数字文本 -> CAST 为 INTEGER -/// - RFC3339 文本 -> strftime('%s', ...) 转为秒级时间戳 -/// - 其余值保持不变(由 DAO 兼容读取) -pub fn migrate_mcp_created_at_to_integer(conn: &Connection) -> Result { - let migrated: bool = conn - .query_row( - "SELECT value FROM settings WHERE key = 'migrated_mcp_created_at_to_integer'", - [], - |row| row.get::<_, String>(0), - ) - .map(|v| v == "true") - .unwrap_or(false); - - if migrated { - tracing::debug!("[迁移] MCP created_at 类型已归一化,跳过"); - return Ok(0); - } - - let updated_numeric = conn - .execute( - "UPDATE mcp_servers - SET created_at = CAST(TRIM(created_at) AS INTEGER) - WHERE typeof(created_at) = 'text' - AND TRIM(created_at) != '' - AND TRIM(created_at) NOT GLOB '*[^0-9]*'", - [], - ) - .map_err(|e| format!("归一化 MCP created_at 数字文本失败: {e}"))?; - - let updated_rfc3339 = conn - .execute( - "UPDATE mcp_servers - SET created_at = CAST(strftime('%s', created_at) AS INTEGER) - WHERE typeof(created_at) = 'text' - AND strftime('%s', created_at) IS NOT NULL", - [], - ) - .map_err(|e| format!("归一化 MCP created_at RFC3339 文本失败: {e}"))?; - - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_mcp_created_at_to_integer', 'true')", - [], - ) - .map_err(|e| format!("标记 MCP created_at 归一化完成失败: {e}"))?; - - let total = updated_numeric + updated_rfc3339; - tracing::info!("[迁移] MCP created_at 归一化完成,更新 {} 条记录", total); - - Ok(total) -} - -/// 当前模型注册表版本 -/// 每次更新模型数据结构或添加新 Provider 时,增加此版本号 -const MODEL_REGISTRY_VERSION: &str = "2026.01.16.1"; - -/// 标记需要刷新模型注册表 -pub fn mark_model_registry_refresh_needed(conn: &Connection) { - let _ = conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('model_registry_refresh_needed', 'true')", - [], - ); - tracing::info!("[迁移] 已标记需要刷新模型注册表"); -} - -/// 检查模型注册表版本,如果版本不匹配则标记需要刷新 -pub fn check_model_registry_version(conn: &Connection) { - let current_version: Option = conn - .query_row( - "SELECT value FROM settings WHERE key = 'model_registry_version'", - [], - |row| row.get(0), - ) - .ok(); - - if current_version.as_deref() != Some(MODEL_REGISTRY_VERSION) { - tracing::info!( - "[迁移] 模型注册表版本不匹配: {:?} -> {},标记需要刷新", - current_version, - MODEL_REGISTRY_VERSION - ); - mark_model_registry_refresh_needed(conn); - - // 更新版本号 - let _ = conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('model_registry_version', ?1)", - params![MODEL_REGISTRY_VERSION], - ); - } -} - -/// 检查是否需要刷新模型注册表 -pub fn is_model_registry_refresh_needed(conn: &Connection) -> bool { - conn.query_row( - "SELECT value FROM settings WHERE key = 'model_registry_refresh_needed'", - [], - |row| row.get::<_, String>(0), - ) - .map(|v| v == "true") - .unwrap_or(false) -} - -/// 清除模型注册表刷新标记 -pub fn clear_model_registry_refresh_flag(conn: &Connection) { - let _ = conn.execute( - "DELETE FROM settings WHERE key = 'model_registry_refresh_needed'", - [], - ); -} - -// ============================================================================ -// General Chat 数据迁移到统一表 -// ============================================================================ - -/// 执行 General Chat 数据迁移到统一表 -/// -/// 将 general_chat_sessions/messages 数据迁移到 agent_sessions/messages 表 -/// - general_chat_sessions → agent_sessions (mode 前缀为 "general:") -/// - general_chat_messages → agent_messages -pub fn migrate_general_chat_to_unified(conn: &Connection) -> Result { - // 检查是否已经迁移过 - let migrated: bool = conn - .query_row( - "SELECT value FROM settings WHERE key = 'migrated_general_chat_to_unified'", - [], - |row| row.get::<_, String>(0), - ) - .map(|v| v == "true") - .unwrap_or(false); - - if migrated { - tracing::debug!("[迁移] General Chat 已迁移过,跳过"); - return Ok(0); - } - - // 检查是否有数据需要迁移 - let general_count: i64 = conn - .query_row("SELECT COUNT(*) FROM general_chat_sessions", [], |row| { - row.get(0) - }) - .unwrap_or(0); - - if general_count == 0 { - tracing::info!("[迁移] general_chat_sessions 表为空,无需迁移"); - // 标记迁移完成 - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_general_chat_to_unified', 'true')", - [], - ) - .map_err(|e| format!("标记迁移完成失败: {e}"))?; - return Ok(0); - } - - tracing::info!( - "[迁移] 开始迁移 {} 个 general_chat 会话到统一表", - general_count - ); - - // 迁移会话 - let migrated_sessions = migrate_general_sessions(conn)?; - tracing::info!("[迁移] 迁移了 {} 个会话", migrated_sessions); - - // 迁移消息 - let migrated_messages = migrate_general_messages(conn)?; - tracing::info!("[迁移] 迁移了 {} 条消息", migrated_messages); - - // 标记迁移完成 - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_general_chat_to_unified', 'true')", - [], - ) - .map_err(|e| format!("标记迁移完成失败: {e}"))?; - - tracing::info!("[迁移] General Chat 数据迁移完成!"); - Ok(migrated_sessions + migrated_messages) -} - -/// 迁移 General Chat 会话数据 -fn migrate_general_sessions(conn: &Connection) -> Result { - let mut stmt = conn - .prepare( - "SELECT id, name, created_at, updated_at, metadata - FROM general_chat_sessions", - ) - .map_err(|e| format!("准备查询语句失败: {e}"))?; - - let sessions: Vec<(String, String, i64, i64, Option)> = stmt - .query_map([], |row| { - Ok(( - row.get(0)?, - row.get(1)?, - row.get(2)?, - row.get(3)?, - row.get(4)?, - )) - }) - .map_err(|e| format!("查询会话失败: {e}"))? - .filter_map(|r| r.ok()) - .collect(); - - let mut count = 0; - for (id, name, created_at, updated_at, _metadata) in sessions { - // 检查是否已存在 - let exists: bool = conn - .query_row( - "SELECT 1 FROM agent_sessions WHERE id = ?", - params![&id], - |_| Ok(true), - ) - .unwrap_or(false); - - if exists { - tracing::warn!("[迁移] 会话 {} 已存在,跳过", id); - continue; - } - - // 转换时间戳格式(general_chat 用毫秒,agent 用 RFC3339) - let created_str = timestamp_ms_to_rfc3339(created_at); - let updated_str = timestamp_ms_to_rfc3339(updated_at); - - // 插入到 agent_sessions,model 字段使用 "general:default" 标识模式 - conn.execute( - "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at) - VALUES (?1, ?2, ?3, ?4, ?5, ?6)", - params![ - id, - "general:default", // 标识为 general 模式 - Option::::None, - name, - created_str, - updated_str, - ], - ) - .map_err(|e| format!("插入会话失败: {e}"))?; - - count += 1; - } - - Ok(count) -} - -/// 迁移 General Chat 消息数据 -fn migrate_general_messages(conn: &Connection) -> Result { - let mut stmt = conn - .prepare( - "SELECT id, session_id, role, content, blocks, status, created_at, metadata - FROM general_chat_messages", - ) - .map_err(|e| format!("准备查询语句失败: {e}"))?; - - #[allow(clippy::type_complexity)] - let messages: Vec<( - String, - String, - String, - String, - Option, - String, - i64, - Option, - )> = stmt - .query_map([], |row| { - Ok(( - row.get(0)?, - row.get(1)?, - row.get(2)?, - row.get(3)?, - row.get(4)?, - row.get(5)?, - row.get(6)?, - row.get(7)?, - )) - }) - .map_err(|e| format!("查询消息失败: {e}"))? - .filter_map(|r| r.ok()) - .collect(); - - let mut count = 0; - for (_id, session_id, role, content, blocks, _status, created_at, _metadata) in messages { - // 检查会话是否存在于 agent_sessions - let session_exists: bool = conn - .query_row( - "SELECT 1 FROM agent_sessions WHERE id = ?", - params![&session_id], - |_| Ok(true), - ) - .unwrap_or(false); - - if !session_exists { - tracing::warn!("[迁移] 消息的会话 {} 不存在,跳过", session_id); - continue; - } - - // 转换内容格式 - let content_json = convert_general_content_to_json(&content, &blocks); - let timestamp_str = timestamp_ms_to_rfc3339(created_at); - - // 插入到 agent_messages - conn.execute( - "INSERT INTO agent_messages (session_id, role, content_json, timestamp, tool_calls_json, tool_call_id) - VALUES (?1, ?2, ?3, ?4, ?5, ?6)", - params![ - session_id, - role, - content_json, - timestamp_str, - Option::::None, - Option::::None, - ], - ) - .map_err(|e| format!("插入消息失败: {e}"))?; - - count += 1; - } - - Ok(count) -} - -/// 将毫秒时间戳转换为 RFC3339 格式 -fn timestamp_ms_to_rfc3339(timestamp_ms: i64) -> String { - use chrono::{TimeZone, Utc}; - - let secs = timestamp_ms / 1000; - let nsecs = ((timestamp_ms % 1000) * 1_000_000) as u32; - - match Utc.timestamp_opt(secs, nsecs) { - chrono::LocalResult::Single(dt) => dt.to_rfc3339(), - _ => Utc::now().to_rfc3339(), - } -} - -/// 将 General Chat 内容转换为 JSON 格式 -fn convert_general_content_to_json(content: &str, blocks: &Option) -> String { - // 如果有 blocks,尝试解析并转换 - if let Some(blocks_str) = blocks { - if let Ok(blocks_arr) = serde_json::from_str::>(blocks_str) { - let converted: Vec = blocks_arr - .into_iter() - .map(|block| { - if let Some(block_type) = block.get("type").and_then(|t| t.as_str()) { - match block_type { - "text" => { - let text = - block.get("content").and_then(|c| c.as_str()).unwrap_or(""); - serde_json::json!({"type": "text", "text": text}) - } - "image" => { - let url = - block.get("content").and_then(|c| c.as_str()).unwrap_or(""); - serde_json::json!({"type": "image", "url": url}) - } - _ => serde_json::json!({"type": "text", "text": content}), - } - } else { - serde_json::json!({"type": "text", "text": content}) - } - }) - .collect(); - - if let Ok(json_str) = serde_json::to_string(&converted) { - return json_str; - } - } - } - - // 默认:纯文本格式 - serde_json::json!([{"type": "text", "text": content}]).to_string() -} - -/// 检查 General Chat 迁移状态 -#[allow(dead_code)] -pub fn check_general_chat_migration_status(conn: &Connection) -> GeneralChatMigrationStatus { - let general_sessions: i64 = conn - .query_row("SELECT COUNT(*) FROM general_chat_sessions", [], |row| { - row.get(0) - }) - .unwrap_or(0); - - let general_messages: i64 = conn - .query_row("SELECT COUNT(*) FROM general_chat_messages", [], |row| { - row.get(0) - }) - .unwrap_or(0); - - let unified_general_sessions: i64 = conn - .query_row( - "SELECT COUNT(*) FROM agent_sessions WHERE model LIKE 'general:%'", - [], - |row| row.get(0), - ) - .unwrap_or(0); - - GeneralChatMigrationStatus { - general_sessions_count: general_sessions as usize, - general_messages_count: general_messages as usize, - migrated_sessions_count: unified_general_sessions as usize, - needs_migration: general_sessions > 0 && unified_general_sessions == 0, - } -} - -/// General Chat 迁移状态 -#[derive(Debug)] -#[allow(dead_code)] -pub struct GeneralChatMigrationStatus { - pub general_sessions_count: usize, - pub general_messages_count: usize, - pub migrated_sessions_count: usize, - pub needs_migration: bool, -} diff --git a/src-tauri/crates/core/src/database/migration/api_key_migration.rs b/src-tauri/crates/core/src/database/migration/api_key_migration.rs new file mode 100644 index 000000000..179368e68 --- /dev/null +++ b/src-tauri/crates/core/src/database/migration/api_key_migration.rs @@ -0,0 +1,585 @@ +use rusqlite::{params, Connection}; + +use super::{is_true_setting, mark_true_setting}; + +const API_KEYS_TO_POOL_MIGRATED_KEY: &str = "migrated_api_keys_to_pool"; +const PROVIDER_IDS_MIGRATED_KEY: &str = "migrated_provider_ids_v1"; +const LEGACY_API_KEY_CREDENTIALS_CLEANED_KEY: &str = "cleaned_legacy_api_key_credentials"; + +/// 将 api_keys 表中的数据迁移到 provider_pool_credentials 表 +/// +/// 迁移逻辑: +/// 1. 读取 api_keys 表中的所有 API Key +/// 2. 根据 provider_id 查找对应的 api_key_providers 配置 +/// 3. 将 API Key 转换为对应的新凭证结构 +/// 4. 插入到 provider_pool_credentials 表 +pub fn migrate_api_keys_to_pool(conn: &Connection) -> Result { + if is_true_setting(conn, API_KEYS_TO_POOL_MIGRATED_KEY) { + tracing::debug!("[迁移] API Keys 已迁移过,跳过"); + return Ok(0); + } + + tracing::info!("[迁移] 开始将 api_keys 迁移到 provider_pool_credentials"); + + let mut stmt = conn + .prepare( + "SELECT k.id, k.provider_id, k.api_key_encrypted, k.alias, k.enabled, + k.usage_count, k.error_count, k.last_used_at, k.created_at, + p.type, p.api_host, p.name as provider_name + FROM api_keys k + JOIN api_key_providers p ON k.provider_id = p.id + ORDER BY k.created_at ASC", + ) + .map_err(|e| format!("准备查询语句失败: {e}"))?; + + let rows = stmt + .query_map([], |row| { + Ok(ApiKeyMigrationRow { + id: row.get(0)?, + provider_id: row.get(1)?, + api_key_encrypted: row.get(2)?, + alias: row.get(3)?, + enabled: row.get(4)?, + usage_count: row.get::<_, i64>(5)? as u64, + error_count: row.get::<_, i64>(6)? as u32, + last_used_at: row.get(7)?, + created_at: row.get(8)?, + provider_type: row.get(9)?, + api_host: row.get(10)?, + provider_name: row.get(11)?, + }) + }) + .map_err(|e| format!("查询 API Keys 失败: {e}"))?; + + let mut migrated_count = 0; + let now = chrono::Utc::now().timestamp(); + + for row_result in rows { + let row = row_result.map_err(|e| format!("读取行数据失败: {e}"))?; + + let exists: bool = conn + .query_row( + "SELECT COUNT(*) > 0 FROM provider_pool_credentials + WHERE credential_data LIKE ?1", + params![format!("%{}%", row.api_key_encrypted)], + |result_row| result_row.get(0), + ) + .unwrap_or(false); + + if exists { + tracing::debug!( + "[迁移] 跳过已存在的 API Key: {} (provider: {})", + row.alias.as_deref().unwrap_or(&row.id), + row.provider_id + ); + continue; + } + + let (pool_provider_type, credential_data) = map_api_key_credential(&row); + let name = row + .alias + .clone() + .or_else(|| Some(format!("{} (迁移)", row.provider_name))); + + let created_at_ts = row + .created_at + .as_ref() + .and_then(|value| chrono::DateTime::parse_from_rfc3339(value).ok()) + .map(|date_time| date_time.timestamp()) + .unwrap_or(now); + + let last_used_ts = row + .last_used_at + .as_ref() + .and_then(|value| chrono::DateTime::parse_from_rfc3339(value).ok()) + .map(|date_time| date_time.timestamp()); + + let uuid = uuid::Uuid::new_v4().to_string(); + let credential_json = credential_data.to_string(); + + conn.execute( + "INSERT INTO provider_pool_credentials + (uuid, provider_type, credential_data, name, is_healthy, is_disabled, + check_health, check_model_name, not_supported_models, usage_count, error_count, + last_used, last_error_time, last_error_message, last_health_check_time, + last_health_check_model, created_at, updated_at, source, proxy_url) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20)", + params![ + uuid, + pool_provider_type, + credential_json, + name, + true, + !row.enabled, + true, + Option::::None, + "[]", + row.usage_count as i64, + row.error_count as i32, + last_used_ts, + Option::::None, + Option::::None, + Option::::None, + Option::::None, + created_at_ts, + now, + "imported", + Option::::None, + ], + ) + .map_err(|e| format!("插入凭证失败: {e}"))?; + + tracing::info!( + "[迁移] 已迁移 API Key: {} -> {} (provider_type: {})", + row.alias.as_deref().unwrap_or(&row.id), + uuid, + pool_provider_type + ); + + migrated_count += 1; + } + + mark_true_setting(conn, API_KEYS_TO_POOL_MIGRATED_KEY)?; + + tracing::info!("[迁移] API Keys 迁移完成,共迁移 {} 条记录", migrated_count); + + Ok(migrated_count) +} + +/// 迁移旧的 Provider ID 到新的 ID +/// +/// 修复 system_providers.rs 中 Provider ID 与模型注册表 JSON 文件名不匹配的问题。 +/// 例如:silicon -> siliconflow, gemini -> google 等 +pub fn migrate_provider_ids(conn: &Connection) -> Result { + if is_true_setting(conn, PROVIDER_IDS_MIGRATED_KEY) { + tracing::debug!("[迁移] Provider ID 已迁移过,跳过"); + return Ok(0); + } + + tracing::info!("[迁移] 开始迁移旧的 Provider ID"); + + let mut migrated_count = 0; + + for (old_id, new_id) in provider_id_mappings() { + let old_exists: bool = conn + .query_row( + "SELECT COUNT(*) > 0 FROM api_key_providers WHERE id = ?1", + params![old_id], + |row| row.get(0), + ) + .unwrap_or(false); + + if !old_exists { + continue; + } + + let new_exists: bool = conn + .query_row( + "SELECT COUNT(*) > 0 FROM api_key_providers WHERE id = ?1", + params![new_id], + |row| row.get(0), + ) + .unwrap_or(false); + + let has_keys: bool = conn + .query_row( + "SELECT COUNT(*) > 0 FROM api_keys WHERE provider_id = ?1", + params![old_id], + |row| row.get(0), + ) + .unwrap_or(false); + + if has_keys { + if new_exists { + conn.execute( + "UPDATE api_keys SET provider_id = ?1 WHERE provider_id = ?2", + params![new_id, old_id], + ) + .map_err(|e| format!("迁移 API Keys 失败: {e}"))?; + + tracing::info!("[迁移] 已将 {} 的 API Keys 迁移到 {}", old_id, new_id); + } else { + conn.execute( + "UPDATE api_key_providers SET id = ?1 WHERE id = ?2", + params![new_id, old_id], + ) + .map_err(|e| format!("更新 Provider ID 失败: {e}"))?; + + conn.execute( + "UPDATE api_keys SET provider_id = ?1 WHERE provider_id = ?2", + params![new_id, old_id], + ) + .map_err(|e| format!("更新 API Keys provider_id 失败: {e}"))?; + + tracing::info!("[迁移] 已将 Provider {} 重命名为 {}", old_id, new_id); + migrated_count += 1; + continue; + } + } + + conn.execute( + "DELETE FROM api_key_providers WHERE id = ?1", + params![old_id], + ) + .map_err(|e| format!("删除旧 Provider 失败: {e}"))?; + + tracing::info!("[迁移] 已删除旧 Provider: {}", old_id); + migrated_count += 1; + } + + mark_true_setting(conn, PROVIDER_IDS_MIGRATED_KEY)?; + + if migrated_count > 0 { + tracing::info!( + "[迁移] Provider ID 迁移完成,共处理 {} 个 Provider", + migrated_count + ); + } + + Ok(migrated_count) +} + +/// 清理旧的 API Key 凭证(OpenAIKey 和 ClaudeKey 类型) +/// +/// 这些凭证是通过旧的 UI 添加的,现在已经被新的 API Key Provider 系统取代。 +/// 此函数会删除 provider_pool_credentials 表中的 openai_key 和 claude_key 类型凭证。 +pub fn cleanup_legacy_api_key_credentials(conn: &Connection) -> Result { + if is_true_setting(conn, LEGACY_API_KEY_CREDENTIALS_CLEANED_KEY) { + tracing::debug!("[清理] 旧 API Key 凭证已清理过,跳过"); + return Ok(0); + } + + tracing::info!("[清理] 开始清理旧的 API Key 凭证(openai_key, claude_key 类型)"); + + let count: i64 = conn + .query_row( + "SELECT COUNT(*) FROM provider_pool_credentials + WHERE credential_data LIKE '%\"type\":\"openai_key\"%' + OR credential_data LIKE '%\"type\":\"claude_key\"%'", + [], + |row| row.get(0), + ) + .unwrap_or(0); + + if count == 0 { + tracing::info!("[清理] 没有需要清理的旧 API Key 凭证"); + mark_true_setting(conn, LEGACY_API_KEY_CREDENTIALS_CLEANED_KEY)?; + return Ok(0); + } + + let mut stmt = conn + .prepare( + "SELECT uuid, name, provider_type + FROM provider_pool_credentials + WHERE credential_data LIKE '%\"type\":\"openai_key\"%' + OR credential_data LIKE '%\"type\":\"claude_key\"%'", + ) + .map_err(|e| format!("准备查询语句失败: {e}"))?; + + let rows = stmt + .query_map([], |row| { + Ok(( + row.get::<_, String>(0)?, + row.get::<_, Option>(1)?, + row.get::<_, String>(2)?, + )) + }) + .map_err(|e| format!("查询旧凭证失败: {e}"))?; + + for (uuid, name, provider_type) in rows.into_iter().filter_map(Result::ok) { + tracing::info!( + "[清理] 将删除旧凭证: {} (name: {}, type: {})", + uuid, + name.as_deref().unwrap_or("未命名"), + provider_type + ); + } + + let deleted = conn + .execute( + "DELETE FROM provider_pool_credentials + WHERE credential_data LIKE '%\"type\":\"openai_key\"%' + OR credential_data LIKE '%\"type\":\"claude_key\"%'", + [], + ) + .map_err(|e| format!("删除旧凭证失败: {e}"))?; + + mark_true_setting(conn, LEGACY_API_KEY_CREDENTIALS_CLEANED_KEY)?; + + tracing::info!("[清理] 旧 API Key 凭证清理完成,共删除 {} 条记录", deleted); + + Ok(deleted) +} + +fn map_api_key_credential(row: &ApiKeyMigrationRow) -> (&'static str, serde_json::Value) { + match row.provider_type.to_lowercase().as_str() { + "anthropic" => ( + "claude", + serde_json::json!({ + "type": "claude_key", + "api_key": row.api_key_encrypted, + "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) } + }), + ), + "openai" | "openai-response" => ( + "openai", + serde_json::json!({ + "type": "openai_key", + "api_key": row.api_key_encrypted, + "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) } + }), + ), + "gemini" => ( + "gemini_api_key", + serde_json::json!({ + "type": "gemini_api_key", + "api_key": row.api_key_encrypted, + "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) }, + "excluded_models": [] + }), + ), + "vertex" | "vertexai" => ( + "vertex", + serde_json::json!({ + "type": "vertex_key", + "api_key": row.api_key_encrypted, + "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) }, + "model_aliases": {} + }), + ), + _ => ( + "openai", + serde_json::json!({ + "type": "openai_key", + "api_key": row.api_key_encrypted, + "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) } + }), + ), + } +} + +fn provider_id_mappings() -> &'static [(&'static str, &'static str)] { + &[ + ("silicon", "siliconflow"), + ("gemini", "google"), + ("zhipu", "zhipuai"), + ("dashscope", "alibaba"), + ("moonshot", "moonshotai"), + ("grok", "xai"), + ("github", "github-models"), + ("copilot", "github-copilot"), + ("vertexai", "google-vertex"), + ("aws-bedrock", "amazon-bedrock"), + ("together", "togetherai"), + ("fireworks", "fireworks-ai"), + ("mimo", "xiaomi"), + ] +} + +struct ApiKeyMigrationRow { + id: String, + provider_id: String, + api_key_encrypted: String, + alias: Option, + enabled: bool, + usage_count: u64, + error_count: u32, + last_used_at: Option, + created_at: Option, + provider_type: String, + api_host: String, + provider_name: String, +} + +#[cfg(test)] +mod tests { + use super::*; + + fn setup_api_key_migration_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 api_key_providers ( + id TEXT PRIMARY KEY, + type TEXT NOT NULL, + api_host TEXT NOT NULL DEFAULT '', + name TEXT NOT NULL + ); + CREATE TABLE api_keys ( + id TEXT PRIMARY KEY, + provider_id TEXT NOT NULL, + api_key_encrypted TEXT NOT NULL, + alias TEXT, + enabled INTEGER NOT NULL DEFAULT 1, + usage_count INTEGER NOT NULL DEFAULT 0, + error_count INTEGER NOT NULL DEFAULT 0, + last_used_at TEXT, + created_at TEXT + ); + CREATE TABLE provider_pool_credentials ( + uuid TEXT PRIMARY KEY, + provider_type TEXT NOT NULL, + credential_data TEXT NOT NULL, + name TEXT, + is_healthy INTEGER NOT NULL, + is_disabled INTEGER NOT NULL, + check_health INTEGER NOT NULL, + check_model_name TEXT, + not_supported_models TEXT NOT NULL, + usage_count INTEGER NOT NULL, + error_count INTEGER NOT NULL, + last_used INTEGER, + last_error_time INTEGER, + last_error_message TEXT, + last_health_check_time INTEGER, + last_health_check_model TEXT, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + source TEXT, + proxy_url TEXT + ); + ", + ) + .unwrap(); + conn + } + + fn insert_pool_credential(conn: &Connection, uuid: &str, credential_type: &str) { + conn.execute( + "INSERT INTO provider_pool_credentials + (uuid, provider_type, credential_data, name, is_healthy, is_disabled, + check_health, check_model_name, not_supported_models, usage_count, error_count, + last_used, last_error_time, last_error_message, last_health_check_time, + last_health_check_model, created_at, updated_at, source, proxy_url) + VALUES (?1, ?2, ?3, ?4, 1, 0, 1, NULL, '[]', 0, 0, NULL, NULL, NULL, NULL, NULL, 0, 0, 'imported', NULL)", + params![ + uuid, + "openai", + serde_json::json!({ + "type": credential_type, + "api_key": format!("key-{uuid}") + }) + .to_string(), + format!("cred-{uuid}") + ], + ) + .unwrap(); + } + + #[test] + fn migrate_api_keys_to_pool_is_idempotent() { + let conn = setup_api_key_migration_db(); + + conn.execute( + "INSERT INTO api_key_providers (id, type, api_host, name) VALUES (?1, ?2, ?3, ?4)", + params![ + "provider-1", + "anthropic", + "https://api.anthropic.com", + "Anthropic" + ], + ) + .unwrap(); + conn.execute( + "INSERT INTO api_keys + (id, provider_id, api_key_encrypted, alias, enabled, usage_count, error_count, last_used_at, created_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + params![ + "key-1", + "provider-1", + "secret-1", + "主账号", + true, + 9i64, + 1i64, + "2026-03-01T00:00:00Z", + "2026-02-01T00:00:00Z" + ], + ) + .unwrap(); + + let migrated = migrate_api_keys_to_pool(&conn).unwrap(); + assert_eq!(migrated, 1); + + let stored: (String, String, Option, bool) = conn + .query_row( + "SELECT provider_type, credential_data, name, is_disabled + FROM provider_pool_credentials", + [], + |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?)), + ) + .unwrap(); + assert_eq!(stored.0, "claude"); + assert!(stored.1.contains("\"type\":\"claude_key\"")); + assert_eq!(stored.2.as_deref(), Some("主账号")); + assert!(!stored.3); + + let migrated_again = migrate_api_keys_to_pool(&conn).unwrap(); + assert_eq!(migrated_again, 0); + } + + #[test] + fn migrate_provider_ids_updates_keys_and_removes_old_provider() { + let conn = setup_api_key_migration_db(); + + conn.execute( + "INSERT INTO api_key_providers (id, type, api_host, name) VALUES (?1, ?2, '', ?3)", + params!["gemini", "gemini", "Gemini"], + ) + .unwrap(); + conn.execute( + "INSERT INTO api_keys (id, provider_id, api_key_encrypted) VALUES (?1, ?2, ?3)", + params!["key-1", "gemini", "secret"], + ) + .unwrap(); + + let migrated = migrate_provider_ids(&conn).unwrap(); + assert_eq!(migrated, 1); + + let provider_id: String = conn + .query_row( + "SELECT provider_id FROM api_keys WHERE id = ?1", + ["key-1"], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(provider_id, "google"); + + let old_exists: bool = conn + .query_row( + "SELECT COUNT(*) > 0 FROM api_key_providers WHERE id = 'gemini'", + [], + |row| row.get(0), + ) + .unwrap(); + assert!(!old_exists); + } + + #[test] + fn cleanup_legacy_api_key_credentials_only_removes_legacy_types() { + let conn = setup_api_key_migration_db(); + + insert_pool_credential(&conn, "legacy-openai", "openai_key"); + insert_pool_credential(&conn, "legacy-claude", "claude_key"); + insert_pool_credential(&conn, "new-gemini", "gemini_api_key"); + + let deleted = cleanup_legacy_api_key_credentials(&conn).unwrap(); + assert_eq!(deleted, 2); + + let remaining: i64 = conn + .query_row( + "SELECT COUNT(*) FROM provider_pool_credentials", + [], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(remaining, 1); + + let deleted_again = cleanup_legacy_api_key_credentials(&conn).unwrap(); + assert_eq!(deleted_again, 0); + } +} 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 new file mode 100644 index 000000000..a79b59563 --- /dev/null +++ b/src-tauri/crates/core/src/database/migration/general_chat_migration.rs @@ -0,0 +1,437 @@ +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"; + +pub fn is_general_chat_migration_completed(conn: &Connection) -> bool { + is_true_setting(conn, GENERAL_CHAT_MIGRATION_COMPLETED_KEY) +} + +/// 执行 General Chat 数据迁移到统一表 +/// +/// 将 general_chat_sessions/messages 数据迁移到 agent_sessions/messages 表 +/// - general_chat_sessions → agent_sessions (mode 前缀为 "general:") +/// - general_chat_messages → agent_messages +pub fn migrate_general_chat_to_unified(conn: &Connection) -> Result { + if is_general_chat_migration_completed(conn) { + tracing::debug!("[迁移] General Chat 已迁移过,跳过"); + return Ok(0); + } + + let general_count: i64 = conn + .query_row("SELECT COUNT(*) FROM general_chat_sessions", [], |row| { + row.get(0) + }) + .unwrap_or(0); + + if general_count == 0 { + tracing::info!("[迁移] general_chat_sessions 表为空,无需迁移"); + mark_true_setting(conn, GENERAL_CHAT_MIGRATION_COMPLETED_KEY)?; + return Ok(0); + } + + tracing::info!( + "[迁移] 开始迁移 {} 个 general_chat 会话到统一表", + general_count + ); + + let migrated_sessions = migrate_general_sessions(conn)?; + tracing::info!("[迁移] 迁移了 {} 个会话", migrated_sessions); + + let migrated_messages = migrate_general_messages(conn)?; + tracing::info!("[迁移] 迁移了 {} 条消息", migrated_messages); + + mark_true_setting(conn, GENERAL_CHAT_MIGRATION_COMPLETED_KEY)?; + + tracing::info!("[迁移] General Chat 数据迁移完成!"); + Ok(migrated_sessions + migrated_messages) +} + +fn migrate_general_sessions(conn: &Connection) -> Result { + let mut stmt = conn + .prepare( + "SELECT id, name, created_at, updated_at, metadata + FROM general_chat_sessions", + ) + .map_err(|e| format!("准备查询语句失败: {e}"))?; + + let sessions: Vec<(String, String, i64, i64, Option)> = stmt + .query_map([], |row| { + Ok(( + row.get(0)?, + row.get(1)?, + row.get(2)?, + row.get(3)?, + row.get(4)?, + )) + }) + .map_err(|e| format!("查询会话失败: {e}"))? + .filter_map(|row| row.ok()) + .collect(); + + let mut count = 0; + for (id, name, created_at, updated_at, _metadata) in sessions { + let exists: bool = conn + .query_row( + "SELECT 1 FROM agent_sessions WHERE id = ?", + params![&id], + |_| Ok(true), + ) + .unwrap_or(false); + + if exists { + tracing::warn!("[迁移] 会话 {} 已存在,跳过", id); + continue; + } + + let created_str = timestamp_ms_to_rfc3339(created_at); + let updated_str = timestamp_ms_to_rfc3339(updated_at); + + conn.execute( + "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + id, + "general:default", + Option::::None, + name, + created_str, + updated_str, + ], + ) + .map_err(|e| format!("插入会话失败: {e}"))?; + + count += 1; + } + + Ok(count) +} + +fn migrate_general_messages(conn: &Connection) -> Result { + let mut stmt = conn + .prepare( + "SELECT id, session_id, role, content, blocks, status, created_at, metadata + FROM general_chat_messages", + ) + .map_err(|e| format!("准备查询语句失败: {e}"))?; + + #[allow(clippy::type_complexity)] + let messages: Vec<( + String, + String, + String, + String, + Option, + String, + i64, + Option, + )> = stmt + .query_map([], |row| { + Ok(( + row.get(0)?, + row.get(1)?, + row.get(2)?, + row.get(3)?, + row.get(4)?, + row.get(5)?, + row.get(6)?, + row.get(7)?, + )) + }) + .map_err(|e| format!("查询消息失败: {e}"))? + .filter_map(|row| row.ok()) + .collect(); + + let mut count = 0; + for (_id, session_id, role, content, blocks, _status, created_at, _metadata) in messages { + let session_exists: bool = conn + .query_row( + "SELECT 1 FROM agent_sessions WHERE id = ?", + params![&session_id], + |_| Ok(true), + ) + .unwrap_or(false); + + if !session_exists { + tracing::warn!("[迁移] 消息的会话 {} 不存在,跳过", session_id); + continue; + } + + let content_json = convert_general_content_to_json(&content, &blocks); + let timestamp_str = timestamp_ms_to_rfc3339(created_at); + + if general_message_already_migrated( + conn, + &session_id, + &role, + &content_json, + ×tamp_str, + )? { + tracing::debug!( + "[迁移] general_chat 消息已存在于 unified 表,跳过: session_id={}, timestamp={}", + session_id, + timestamp_str + ); + continue; + } + + conn.execute( + "INSERT INTO agent_messages (session_id, role, content_json, timestamp, tool_calls_json, tool_call_id) + VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + session_id, + role, + content_json, + timestamp_str, + Option::::None, + Option::::None, + ], + ) + .map_err(|e| format!("插入消息失败: {e}"))?; + + count += 1; + } + + Ok(count) +} + +fn general_message_already_migrated( + conn: &Connection, + session_id: &str, + role: &str, + content_json: &str, + timestamp: &str, +) -> Result { + let exists = conn + .query_row( + "SELECT 1 + FROM agent_messages + WHERE session_id = ?1 + AND role = ?2 + AND content_json = ?3 + AND timestamp = ?4 + LIMIT 1", + params![session_id, role, content_json, timestamp], + |_| Ok(true), + ) + .unwrap_or(false); + + Ok(exists) +} + +fn timestamp_ms_to_rfc3339(timestamp_ms: i64) -> String { + use chrono::{TimeZone, Utc}; + + let secs = timestamp_ms / 1000; + let nsecs = ((timestamp_ms % 1000) * 1_000_000) as u32; + + match Utc.timestamp_opt(secs, nsecs) { + chrono::LocalResult::Single(dt) => dt.to_rfc3339(), + _ => Utc::now().to_rfc3339(), + } +} + +fn convert_general_content_to_json(content: &str, blocks: &Option) -> String { + if let Some(blocks_str) = blocks { + if let Ok(blocks_arr) = serde_json::from_str::>(blocks_str) { + let converted: Vec = blocks_arr + .into_iter() + .map(|block| { + if let Some(block_type) = block.get("type").and_then(|value| value.as_str()) { + match block_type { + "text" => { + let text = block + .get("content") + .and_then(|value| value.as_str()) + .unwrap_or(""); + serde_json::json!({ "type": "text", "text": text }) + } + "image" => { + let url = block + .get("content") + .and_then(|value| value.as_str()) + .unwrap_or(""); + serde_json::json!({ "type": "image", "url": url }) + } + _ => serde_json::json!({ "type": "text", "text": content }), + } + } else { + serde_json::json!({ "type": "text", "text": content }) + } + }) + .collect(); + + if let Ok(json_str) = serde_json::to_string(&converted) { + return json_str; + } + } + } + + serde_json::json!([{ "type": "text", "text": content }]).to_string() +} + +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) + }) + .unwrap_or(0); + + let general_messages: i64 = conn + .query_row("SELECT COUNT(*) FROM general_chat_messages", [], |row| { + row.get(0) + }) + .unwrap_or(0); + + let unified_general_sessions: i64 = conn + .query_row( + "SELECT COUNT(*) FROM agent_sessions WHERE model LIKE 'general:%'", + [], + |row| row.get(0), + ) + .unwrap_or(0); + + GeneralChatMigrationStatus { + general_sessions_count: general_sessions as usize, + general_messages_count: general_messages as usize, + migrated_sessions_count: unified_general_sessions as usize, + needs_migration: (general_sessions > 0 || general_messages > 0) && !migrated, + } +} + +#[derive(Debug)] +pub struct GeneralChatMigrationStatus { + pub general_sessions_count: usize, + pub general_messages_count: usize, + pub migrated_sessions_count: usize, + pub needs_migration: bool, +} + +#[cfg(test)] +mod tests { + use super::*; + use rusqlite::Connection; + + fn setup_general_chat_migration_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 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 + ); + 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(); + + conn.execute( + "INSERT INTO general_chat_sessions (id, name, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", + params!["legacy-session", "Legacy", 1_700_000_000_000i64, 1_700_000_000_100i64], + ) + .unwrap(); + conn.execute( + "INSERT INTO general_chat_messages (id, session_id, role, content, created_at) VALUES (?1, ?2, ?3, ?4, ?5)", + params!["legacy-msg-1", "legacy-session", "user", "你好", 1_700_000_000_001i64], + ) + .unwrap(); + conn.execute( + "INSERT INTO general_chat_messages (id, session_id, role, content, created_at) VALUES (?1, ?2, ?3, ?4, ?5)", + params!["legacy-msg-2", "legacy-session", "assistant", "你好!", 1_700_000_000_002i64], + ) + .unwrap(); + + let migrated = migrate_general_chat_to_unified(&conn).unwrap(); + assert_eq!(migrated, 3); + + let session_count: i64 = conn + .query_row("SELECT COUNT(*) FROM agent_sessions", [], |row| row.get(0)) + .unwrap(); + let message_count: i64 = conn + .query_row("SELECT COUNT(*) FROM agent_messages", [], |row| row.get(0)) + .unwrap(); + assert_eq!(session_count, 1); + assert_eq!(message_count, 2); + + conn.execute( + "DELETE FROM settings WHERE key = ?1", + [GENERAL_CHAT_MIGRATION_COMPLETED_KEY], + ) + .unwrap(); + + let migrated_again = migrate_general_chat_to_unified(&conn).unwrap(); + assert_eq!(migrated_again, 0); + + let message_count_after: i64 = conn + .query_row("SELECT COUNT(*) FROM agent_messages", [], |row| row.get(0)) + .unwrap(); + assert_eq!(message_count_after, 2); + } + + #[test] + fn general_chat_migration_status_uses_completion_flag() { + let conn = setup_general_chat_migration_db(); + + conn.execute( + "INSERT INTO general_chat_sessions (id, name, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", + params!["legacy-session", "Legacy", 1_700_000_000_000i64, 1_700_000_000_100i64], + ) + .unwrap(); + + let pending = check_general_chat_migration_status(&conn); + assert!(pending.needs_migration); + + conn.execute( + "INSERT INTO settings (key, value) VALUES (?1, 'true')", + [GENERAL_CHAT_MIGRATION_COMPLETED_KEY], + ) + .unwrap(); + + let completed = check_general_chat_migration_status(&conn); + assert!(!completed.needs_migration); + } +} diff --git a/src-tauri/crates/core/src/database/migration/mcp_migration.rs b/src-tauri/crates/core/src/database/migration/mcp_migration.rs new file mode 100644 index 000000000..819bbda65 --- /dev/null +++ b/src-tauri/crates/core/src/database/migration/mcp_migration.rs @@ -0,0 +1,201 @@ +use rusqlite::Connection; + +use super::{is_true_setting, mark_true_setting}; + +const MCP_PROXYCAST_ENABLED_MIGRATED_KEY: &str = "migrated_mcp_proxycast_enabled"; +const MCP_CREATED_AT_INTEGER_MIGRATED_KEY: &str = "migrated_mcp_created_at_to_integer"; + +/// 修复历史 MCP 导入数据:补齐 enabled_proxycast +/// +/// 早期版本从 Claude/Codex/Gemini 导入 MCP 时,默认写入 enabled_proxycast=0, +/// 导致 ProxyCast 本身不会使用这些服务器。 +/// +/// 迁移策略: +/// - 仅处理 enabled_proxycast=0 的记录 +/// - 且至少在一个外部应用中启用(enabled_claude/codex/gemini 任一为 1) +/// - 将 enabled_proxycast 设为 1 +pub fn migrate_mcp_proxycast_enabled(conn: &Connection) -> Result { + if is_true_setting(conn, MCP_PROXYCAST_ENABLED_MIGRATED_KEY) { + tracing::debug!("[迁移] MCP proxycast 启用状态已迁移过,跳过"); + return Ok(0); + } + + let updated = conn + .execute( + "UPDATE mcp_servers + SET enabled_proxycast = 1 + WHERE enabled_proxycast = 0 + AND (enabled_claude = 1 OR enabled_codex = 1 OR enabled_gemini = 1)", + [], + ) + .map_err(|e| format!("修复 MCP enabled_proxycast 失败: {e}"))?; + + mark_true_setting(conn, MCP_PROXYCAST_ENABLED_MIGRATED_KEY)?; + + tracing::info!( + "[迁移] MCP proxycast 启用状态修复完成,更新 {} 条记录", + updated + ); + + Ok(updated) +} + +/// 归一化 mcp_servers.created_at 字段为 INTEGER 时间戳 +/// +/// 历史版本曾写入 RFC3339 文本,导致下游按 i64 读取时出现类型异常。 +/// 迁移策略: +/// - 纯数字文本 -> CAST 为 INTEGER +/// - RFC3339 文本 -> strftime('%s', ...) 转为秒级时间戳 +/// - 其余值保持不变(由 DAO 兼容读取) +pub fn migrate_mcp_created_at_to_integer(conn: &Connection) -> Result { + if is_true_setting(conn, MCP_CREATED_AT_INTEGER_MIGRATED_KEY) { + tracing::debug!("[迁移] MCP created_at 类型已归一化,跳过"); + return Ok(0); + } + + let updated_numeric = conn + .execute( + "UPDATE mcp_servers + SET created_at = CAST(TRIM(created_at) AS INTEGER) + WHERE typeof(created_at) = 'text' + AND TRIM(created_at) != '' + AND TRIM(created_at) NOT GLOB '*[^0-9]*'", + [], + ) + .map_err(|e| format!("归一化 MCP created_at 数字文本失败: {e}"))?; + + let updated_rfc3339 = conn + .execute( + "UPDATE mcp_servers + SET created_at = CAST(strftime('%s', created_at) AS INTEGER) + WHERE typeof(created_at) = 'text' + AND strftime('%s', created_at) IS NOT NULL", + [], + ) + .map_err(|e| format!("归一化 MCP created_at RFC3339 文本失败: {e}"))?; + + mark_true_setting(conn, MCP_CREATED_AT_INTEGER_MIGRATED_KEY)?; + + let total = updated_numeric + updated_rfc3339; + tracing::info!("[迁移] MCP created_at 归一化完成,更新 {} 条记录", total); + + Ok(total) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn setup_mcp_migration_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 mcp_servers ( + id TEXT PRIMARY KEY, + enabled_proxycast INTEGER NOT NULL DEFAULT 0, + enabled_claude INTEGER NOT NULL DEFAULT 0, + enabled_codex INTEGER NOT NULL DEFAULT 0, + enabled_gemini INTEGER NOT NULL DEFAULT 0, + created_at + ); + ", + ) + .unwrap(); + conn + } + + #[test] + fn migrate_mcp_proxycast_enabled_updates_only_imported_rows() { + let conn = setup_mcp_migration_db(); + + conn.execute( + "INSERT INTO mcp_servers (id, enabled_proxycast, enabled_claude) VALUES (?1, 0, 1)", + ["server-1"], + ) + .unwrap(); + conn.execute( + "INSERT INTO mcp_servers (id, enabled_proxycast, enabled_claude) VALUES (?1, 0, 0)", + ["server-2"], + ) + .unwrap(); + + let updated = migrate_mcp_proxycast_enabled(&conn).unwrap(); + assert_eq!(updated, 1); + + let enabled_proxycast: i64 = conn + .query_row( + "SELECT enabled_proxycast FROM mcp_servers WHERE id = ?1", + ["server-1"], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(enabled_proxycast, 1); + + let untouched: i64 = conn + .query_row( + "SELECT enabled_proxycast FROM mcp_servers WHERE id = ?1", + ["server-2"], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(untouched, 0); + + let updated_again = migrate_mcp_proxycast_enabled(&conn).unwrap(); + assert_eq!(updated_again, 0); + } + + #[test] + fn migrate_mcp_created_at_to_integer_normalizes_text_values() { + let conn = setup_mcp_migration_db(); + + conn.execute( + "INSERT INTO mcp_servers (id, created_at) VALUES (?1, ?2)", + ("numeric", "1700000000"), + ) + .unwrap(); + conn.execute( + "INSERT INTO mcp_servers (id, created_at) VALUES (?1, ?2)", + ("rfc3339", "2026-03-01T00:00:00Z"), + ) + .unwrap(); + conn.execute( + "INSERT INTO mcp_servers (id, created_at) VALUES (?1, ?2)", + ("invalid", "not-a-date"), + ) + .unwrap(); + + let updated = migrate_mcp_created_at_to_integer(&conn).unwrap(); + assert_eq!(updated, 2); + + let numeric_type: String = conn + .query_row( + "SELECT typeof(created_at) FROM mcp_servers WHERE id = 'numeric'", + [], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(numeric_type, "integer"); + + let rfc_timestamp: i64 = conn + .query_row( + "SELECT created_at FROM mcp_servers WHERE id = 'rfc3339'", + [], + |row| row.get(0), + ) + .unwrap(); + assert!(rfc_timestamp > 0); + + let invalid_type: String = conn + .query_row( + "SELECT typeof(created_at) FROM mcp_servers WHERE id = 'invalid'", + [], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(invalid_type, "text"); + } +} diff --git a/src-tauri/crates/core/src/database/migration/model_registry_migration.rs b/src-tauri/crates/core/src/database/migration/model_registry_migration.rs new file mode 100644 index 000000000..21bea5c04 --- /dev/null +++ b/src-tauri/crates/core/src/database/migration/model_registry_migration.rs @@ -0,0 +1,82 @@ +use rusqlite::Connection; + +use super::{ + clear_setting, is_true_setting, mark_true_setting, read_setting_value, upsert_setting, +}; + +const MODEL_REGISTRY_VERSION: &str = "2026.01.16.1"; +const MODEL_REGISTRY_REFRESH_NEEDED_KEY: &str = "model_registry_refresh_needed"; +const MODEL_REGISTRY_VERSION_KEY: &str = "model_registry_version"; + +/// 标记需要刷新模型注册表 +pub fn mark_model_registry_refresh_needed(conn: &Connection) { + let _ = mark_true_setting(conn, MODEL_REGISTRY_REFRESH_NEEDED_KEY); + tracing::info!("[迁移] 已标记需要刷新模型注册表"); +} + +/// 检查模型注册表版本,如果版本不匹配则标记需要刷新 +pub fn check_model_registry_version(conn: &Connection) { + let current_version = read_setting_value(conn, MODEL_REGISTRY_VERSION_KEY); + + if current_version.as_deref() != Some(MODEL_REGISTRY_VERSION) { + tracing::info!( + "[迁移] 模型注册表版本不匹配: {:?} -> {},标记需要刷新", + current_version, + MODEL_REGISTRY_VERSION + ); + mark_model_registry_refresh_needed(conn); + let _ = upsert_setting(conn, MODEL_REGISTRY_VERSION_KEY, MODEL_REGISTRY_VERSION); + } +} + +/// 检查是否需要刷新模型注册表 +pub fn is_model_registry_refresh_needed(conn: &Connection) -> bool { + is_true_setting(conn, MODEL_REGISTRY_REFRESH_NEEDED_KEY) +} + +/// 清除模型注册表刷新标记 +pub fn clear_model_registry_refresh_flag(conn: &Connection) { + clear_setting(conn, MODEL_REGISTRY_REFRESH_NEEDED_KEY); +} + +#[cfg(test)] +mod tests { + use super::*; + + fn setup_model_registry_db() -> Connection { + let conn = Connection::open_in_memory().unwrap(); + conn.execute_batch( + " + CREATE TABLE settings ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL + ); + ", + ) + .unwrap(); + conn + } + + #[test] + fn check_model_registry_version_marks_refresh_only_on_version_change() { + let conn = setup_model_registry_db(); + + check_model_registry_version(&conn); + assert!(is_model_registry_refresh_needed(&conn)); + + let stored_version: String = conn + .query_row( + "SELECT value FROM settings WHERE key = ?1", + [MODEL_REGISTRY_VERSION_KEY], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(stored_version, MODEL_REGISTRY_VERSION); + + clear_model_registry_refresh_flag(&conn); + assert!(!is_model_registry_refresh_needed(&conn)); + + check_model_registry_version(&conn); + assert!(!is_model_registry_refresh_needed(&conn)); + } +} diff --git a/src-tauri/crates/core/src/database/migration_support.rs b/src-tauri/crates/core/src/database/migration_support.rs new file mode 100644 index 000000000..129aeed13 --- /dev/null +++ b/src-tauri/crates/core/src/database/migration_support.rs @@ -0,0 +1,116 @@ +use rusqlite::{params, Connection}; + +pub(crate) fn read_setting_value(conn: &Connection, key: &str) -> Option { + conn.query_row("SELECT value FROM settings WHERE key = ?1", [key], |row| { + row.get::<_, String>(0) + }) + .ok() +} + +pub(crate) fn is_true_setting(conn: &Connection, key: &str) -> bool { + read_setting_value(conn, key) + .map(|value| value == "true") + .unwrap_or(false) +} + +pub(crate) fn is_migration_completed(conn: &Connection, key: &str) -> bool { + read_setting_value(conn, key) + .map(|value| value == "true" || value == "1") + .unwrap_or(false) +} + +pub(crate) fn upsert_setting(conn: &Connection, key: &str, value: &str) -> Result<(), String> { + conn.execute( + "INSERT OR REPLACE INTO settings (key, value) VALUES (?1, ?2)", + params![key, value], + ) + .map_err(|e| format!("写入 settings[{key}] 失败: {e}"))?; + Ok(()) +} + +pub(crate) fn mark_true_setting(conn: &Connection, key: &str) -> Result<(), String> { + upsert_setting(conn, key, "true") +} + +pub(crate) fn mark_migration_completed(conn: &Connection, key: &str) -> Result<(), String> { + mark_true_setting(conn, key) +} + +pub(crate) fn clear_setting(conn: &Connection, key: &str) { + let _ = conn.execute("DELETE FROM settings WHERE key = ?1", [key]); +} + +pub(crate) fn run_in_transaction(conn: &Connection, operation: F) -> Result +where + F: FnOnce(&Connection) -> Result, +{ + conn.execute("BEGIN TRANSACTION", []) + .map_err(|e| format!("开始事务失败: {e}"))?; + + match operation(conn) { + Ok(value) => { + conn.execute("COMMIT", []) + .map_err(|e| format!("提交事务失败: {e}"))?; + Ok(value) + } + Err(error) => { + let _ = conn.execute("ROLLBACK", []); + Err(error) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn setup_settings_db() -> Connection { + let conn = Connection::open_in_memory().unwrap(); + conn.execute( + "CREATE TABLE settings (key TEXT PRIMARY KEY, value TEXT NOT NULL)", + [], + ) + .unwrap(); + conn + } + + #[test] + fn migration_completed_accepts_legacy_and_new_markers() { + let conn = setup_settings_db(); + + conn.execute( + "INSERT INTO settings (key, value) VALUES (?1, ?2)", + ("legacy", "1"), + ) + .unwrap(); + conn.execute( + "INSERT INTO settings (key, value) VALUES (?1, ?2)", + ("current", "true"), + ) + .unwrap(); + + assert!(is_migration_completed(&conn, "legacy")); + assert!(is_migration_completed(&conn, "current")); + assert!(!is_migration_completed(&conn, "missing")); + } + + #[test] + fn run_in_transaction_rolls_back_on_error() { + let conn = Connection::open_in_memory().unwrap(); + conn.execute("CREATE TABLE demo (value TEXT NOT NULL)", []) + .unwrap(); + + let result: Result<(), String> = run_in_transaction(&conn, |tx| { + tx.execute("INSERT INTO demo (value) VALUES ('should_rollback')", []) + .map_err(|e| e.to_string())?; + Err("boom".to_string()) + }); + + assert_eq!(result.unwrap_err(), "boom"); + + let count: i64 = conn + .query_row("SELECT COUNT(*) FROM demo", [], |row| row.get(0)) + .unwrap(); + assert_eq!(count, 0); + } +} diff --git a/src-tauri/crates/core/src/database/migration_v2.rs b/src-tauri/crates/core/src/database/migration_v2.rs index 10a103c49..909a05856 100644 --- a/src-tauri/crates/core/src/database/migration_v2.rs +++ b/src-tauri/crates/core/src/database/migration_v2.rs @@ -10,6 +10,12 @@ use chrono::Utc; use rusqlite::{params, Connection}; use uuid::Uuid; +use crate::app_paths; + +use super::migration_support::{ + is_migration_completed, mark_migration_completed, run_in_transaction, +}; + /// 迁移设置键名 const MIGRATION_KEY_UNIFIED_CONTENT: &str = "migrated_unified_content_system_v1"; @@ -29,7 +35,19 @@ const DEFAULT_PROJECT_ICON: &str = "📁"; /// /// _Requirements: 2.1, 2.2, 2.3, 2.4_ pub fn migrate_unified_content_system(conn: &Connection) -> Result { - // 检查是否已经迁移过 + migrate_unified_content_system_with_default_dir_resolver( + conn, + &app_paths::resolve_default_project_dir, + ) +} + +fn migrate_unified_content_system_with_default_dir_resolver( + conn: &Connection, + resolve_default_project_dir: &F, +) -> Result +where + F: Fn() -> Result, +{ if is_migration_completed(conn, MIGRATION_KEY_UNIFIED_CONTENT) { tracing::debug!("[迁移] 统一内容系统已迁移过,跳过"); return Ok(MigrationResult::skipped()); @@ -37,45 +55,37 @@ pub fn migrate_unified_content_system(conn: &Connection) -> Result { - // 标记迁移完成 - mark_migration_completed(conn, MIGRATION_KEY_UNIFIED_CONTENT)?; - - // 提交事务 - conn.execute("COMMIT", []) - .map_err(|e| format!("提交事务失败: {e}"))?; - tracing::info!( "[迁移] 统一内容系统迁移完成: 默认项目={}, 迁移内容数={}", stats.default_project_id, stats.migrated_contents_count ); - Ok(MigrationResult::success(stats)) } - Err(e) => { - // 回滚事务 - // _Requirements: 2.4_ - let _ = conn.execute("ROLLBACK", []); - tracing::error!("[迁移] 统一内容系统迁移失败,已回滚: {}", e); - Err(e) + Err(error) => { + tracing::error!("[迁移] 统一内容系统迁移失败,已回滚: {}", error); + Err(error) } } } /// 执行迁移的核心逻辑 -fn execute_migration(conn: &Connection) -> Result { +fn execute_migration( + conn: &Connection, + resolve_default_project_dir: &F, +) -> Result +where + F: Fn() -> Result, +{ // 1. 获取或创建默认项目 // _Requirements: 2.1_ - let default_project_id = get_or_create_default_project(conn)?; + let default_project_id = get_or_create_default_project(conn, resolve_default_project_dir)?; // 2. 迁移所有 project_id 为 null 的内容到默认项目 // _Requirements: 2.2_ @@ -96,7 +106,13 @@ fn execute_migration(conn: &Connection) -> Result { /// 否则创建新的默认项目 /// /// _Requirements: 2.1_ -fn get_or_create_default_project(conn: &Connection) -> Result { +fn get_or_create_default_project( + conn: &Connection, + resolve_default_project_dir: &F, +) -> Result +where + F: Fn() -> Result, +{ // 检查是否已存在默认项目 let existing_id: Option = conn .query_row( @@ -116,7 +132,7 @@ fn get_or_create_default_project(conn: &Connection) -> Result { let now = Utc::now().timestamp_millis(); // 使用应用数据目录作为默认项目的 root_path - let root_path = get_default_project_path()?; + let root_path = get_default_project_path(resolve_default_project_dir)?; conn.execute( "INSERT INTO workspaces ( @@ -148,12 +164,11 @@ fn get_or_create_default_project(conn: &Connection) -> Result { } /// 获取默认项目的存储路径 -fn get_default_project_path() -> Result { - let home = dirs::home_dir().ok_or_else(|| "无法获取主目录".to_string())?; - let path = home.join(".proxycast").join("projects").join("default"); - - // 确保目录存在 - std::fs::create_dir_all(&path).map_err(|e| format!("创建默认项目目录失败: {e}"))?; +fn get_default_project_path(resolve_default_project_dir: &F) -> Result +where + F: Fn() -> Result, +{ + let path = resolve_default_project_dir()?; path.to_str() .map(|s| s.to_string()) @@ -235,27 +250,6 @@ fn verify_migration(conn: &Connection) -> Result<(), String> { Ok(()) } -/// 检查迁移是否已完成 -fn is_migration_completed(conn: &Connection, key: &str) -> bool { - conn.query_row( - "SELECT value FROM settings WHERE key = ?", - params![key], - |row| row.get::<_, String>(0), - ) - .map(|v| v == "true") - .unwrap_or(false) -} - -/// 标记迁移完成 -fn mark_migration_completed(conn: &Connection, key: &str) -> Result<(), String> { - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES (?, 'true')", - params![key], - ) - .map_err(|e| format!("标记迁移完成失败: {e}"))?; - Ok(()) -} - // ============================================================================ // 迁移结果类型 // ============================================================================ @@ -316,7 +310,7 @@ pub fn get_default_project_id(conn: &Connection) -> Option { /// /// 如果不存在则创建,返回默认项目 ID pub fn ensure_default_project(conn: &Connection) -> Result { - get_or_create_default_project(conn) + get_or_create_default_project(conn, &app_paths::resolve_default_project_dir) } // ============================================================================ @@ -386,9 +380,16 @@ mod tests { #[test] fn test_migration_creates_default_project() { let conn = setup_test_db(); + let temp = tempfile::tempdir().unwrap(); + let expected_default_dir = temp.path().join("projects").join("default"); // 执行迁移 - let result = migrate_unified_content_system(&conn).unwrap(); + let result = migrate_unified_content_system_with_default_dir_resolver(&conn, &|| { + std::fs::create_dir_all(&expected_default_dir) + .map_err(|e| format!("创建默认项目目录失败: {e}"))?; + Ok(expected_default_dir.clone()) + }) + .unwrap(); assert!(result.executed); assert!(result.stats.is_some()); @@ -405,9 +406,38 @@ mod tests { assert!(default_exists); } + #[test] + fn test_migration_uses_app_data_default_project_path() { + let conn = setup_test_db(); + let temp = tempfile::tempdir().unwrap(); + let expected_default_dir = temp.path().join("projects").join("default"); + + let result = migrate_unified_content_system_with_default_dir_resolver(&conn, &|| { + std::fs::create_dir_all(&expected_default_dir) + .map_err(|e| format!("创建默认项目目录失败: {e}"))?; + Ok(expected_default_dir.clone()) + }) + .unwrap(); + let stats = result.stats.unwrap(); + + let root_path: String = conn + .query_row( + "SELECT root_path FROM workspaces WHERE id = ?1", + [stats.default_project_id], + |row| row.get(0), + ) + .unwrap(); + + let expected = expected_default_dir.to_string_lossy().to_string(); + + assert_eq!(root_path, expected); + } + #[test] fn test_migration_migrates_null_project_contents() { let conn = setup_test_db(); + let temp = tempfile::tempdir().unwrap(); + let expected_default_dir = temp.path().join("projects").join("default"); let now = Utc::now().timestamp_millis(); // 插入一些没有 project_id 的内容 @@ -426,7 +456,12 @@ mod tests { .unwrap(); // 执行迁移 - let result = migrate_unified_content_system(&conn).unwrap(); + let result = migrate_unified_content_system_with_default_dir_resolver(&conn, &|| { + std::fs::create_dir_all(&expected_default_dir) + .map_err(|e| format!("创建默认项目目录失败: {e}"))?; + Ok(expected_default_dir.clone()) + }) + .unwrap(); assert!(result.executed); let stats = result.stats.unwrap(); @@ -447,19 +482,33 @@ mod tests { #[test] fn test_migration_skips_if_already_done() { let conn = setup_test_db(); + let temp = tempfile::tempdir().unwrap(); + let expected_default_dir = temp.path().join("projects").join("default"); // 第一次迁移 - let result1 = migrate_unified_content_system(&conn).unwrap(); + let result1 = migrate_unified_content_system_with_default_dir_resolver(&conn, &|| { + std::fs::create_dir_all(&expected_default_dir) + .map_err(|e| format!("创建默认项目目录失败: {e}"))?; + Ok(expected_default_dir.clone()) + }) + .unwrap(); assert!(result1.executed); // 第二次迁移应该跳过 - let result2 = migrate_unified_content_system(&conn).unwrap(); + let result2 = migrate_unified_content_system_with_default_dir_resolver(&conn, &|| { + std::fs::create_dir_all(&expected_default_dir) + .map_err(|e| format!("创建默认项目目录失败: {e}"))?; + Ok(expected_default_dir.clone()) + }) + .unwrap(); assert!(!result2.executed); } #[test] fn test_migration_uses_existing_default_project() { let conn = setup_test_db(); + let temp = tempfile::tempdir().unwrap(); + let expected_default_dir = temp.path().join("projects").join("default"); let now = Utc::now().timestamp_millis(); // 先创建一个默认项目 @@ -479,7 +528,12 @@ mod tests { .unwrap(); // 执行迁移 - let result = migrate_unified_content_system(&conn).unwrap(); + let result = migrate_unified_content_system_with_default_dir_resolver(&conn, &|| { + std::fs::create_dir_all(&expected_default_dir) + .map_err(|e| format!("创建默认项目目录失败: {e}"))?; + Ok(expected_default_dir.clone()) + }) + .unwrap(); assert!(result.executed); let stats = result.stats.unwrap(); diff --git a/src-tauri/crates/core/src/database/migration_v3.rs b/src-tauri/crates/core/src/database/migration_v3.rs index 6e41d411d..c21acdbe5 100644 --- a/src-tauri/crates/core/src/database/migration_v3.rs +++ b/src-tauri/crates/core/src/database/migration_v3.rs @@ -7,6 +7,10 @@ use rusqlite::{params, Connection}; use serde_json::json; use uuid::Uuid; +use super::migration_support::{ + is_migration_completed, mark_migration_completed, run_in_transaction, +}; + /// 迁移设置键名 const MIGRATION_KEY_PLAYWRIGHT_SERVER: &str = "migrated_playwright_mcp_server_v1"; @@ -30,7 +34,6 @@ pub struct MigrationResult { /// 3. 如果不存在,创建默认配置 /// 4. 标记迁移完成 pub fn migrate_playwright_mcp_server(conn: &Connection) -> Result { - // 检查是否已经迁移过 if is_migration_completed(conn, MIGRATION_KEY_PLAYWRIGHT_SERVER) { tracing::debug!("[迁移] Playwright MCP Server 已迁移过,跳过"); return Ok(MigrationResult { @@ -51,22 +54,12 @@ pub fn migrate_playwright_mcp_server(conn: &Connection) -> Result { - // 标记迁移完成 - mark_migration_completed(conn, MIGRATION_KEY_PLAYWRIGHT_SERVER)?; - - // 提交事务 - conn.execute("COMMIT", []) - .map_err(|e| format!("提交事务失败: {e}"))?; - tracing::info!( "[迁移] Playwright MCP Server 迁移完成: server_id={}", server_id @@ -77,11 +70,9 @@ pub fn migrate_playwright_mcp_server(conn: &Connection) -> Result { - // 回滚事务 - let _ = conn.execute("ROLLBACK", []); - tracing::error!("[迁移] Playwright MCP Server 迁移失败,已回滚: {}", e); - Err(e) + Err(error) => { + tracing::error!("[迁移] Playwright MCP Server 迁移失败,已回滚: {}", error); + Err(error) } } } @@ -140,21 +131,3 @@ fn server_exists(conn: &Connection, name: &str) -> bool { .unwrap_or(0) > 0 } - -/// 检查迁移是否已完成 -fn is_migration_completed(conn: &Connection, key: &str) -> bool { - conn.query_row("SELECT value FROM settings WHERE key = ?1", [key], |row| { - row.get::<_, String>(0) - }) - .is_ok() -} - -/// 标记迁移已完成 -fn mark_migration_completed(conn: &Connection, key: &str) -> Result<(), String> { - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES (?1, ?2)", - [key, "1"], - ) - .map_err(|e| format!("标记迁移完成失败: {e}"))?; - Ok(()) -} diff --git a/src-tauri/crates/core/src/database/migration_v4.rs b/src-tauri/crates/core/src/database/migration_v4.rs index 2e4fbc27e..6c7052165 100644 --- a/src-tauri/crates/core/src/database/migration_v4.rs +++ b/src-tauri/crates/core/src/database/migration_v4.rs @@ -9,13 +9,16 @@ use rusqlite::{params, Connection}; +use crate::app_paths; + +use super::migration_support::{ + is_migration_completed, mark_migration_completed, run_in_transaction, +}; + /// 迁移设置键名 const MIGRATION_KEY_FIX_PROMISE_PATHS: &str = "migrated_fix_promise_paths_v1"; const MIGRATION_KEY_UNIFY_SESSION_DIRS: &str = "migrated_unify_session_dirs_v1"; -/// 默认项目根目录(用于替换损坏的路径) -const DEFAULT_PROJECTS_DIR: &str = ".proxycast/projects"; - /// 迁移结果 pub struct MigrationResult { /// 是否执行了迁移 @@ -30,7 +33,6 @@ pub struct MigrationResult { /// 执行 [object Promise] 路径修复迁移 pub fn migrate_fix_promise_paths(conn: &Connection) -> Result { - // 检查是否已经迁移过 let promise_done = is_migration_completed(conn, MIGRATION_KEY_FIX_PROMISE_PATHS); let unify_done = is_migration_completed(conn, MIGRATION_KEY_UNIFY_SESSION_DIRS); @@ -44,30 +46,21 @@ pub fn migrate_fix_promise_paths(conn: &Connection) -> Result { - if !promise_done { - mark_migration_completed(conn, MIGRATION_KEY_FIX_PROMISE_PATHS)?; - } - if !unify_done { - mark_migration_completed(conn, MIGRATION_KEY_UNIFY_SESSION_DIRS)?; - } - - conn.execute("COMMIT", []) - .map_err(|e| format!("提交事务失败: {e}"))?; - if fixed_ws > 0 || fixed_sess > 0 { tracing::info!( "[迁移] Promise 路径修复完成: 修复 workspaces={}, sessions={}", @@ -89,10 +82,9 @@ pub fn migrate_fix_promise_paths(conn: &Connection) -> Result { - let _ = conn.execute("ROLLBACK", []); - tracing::error!("[迁移] 路径修复和会话统一失败,已回滚: {}", e); - Err(e) + Err(error) => { + tracing::error!("[迁移] 路径修复和会话统一失败,已回滚: {}", error); + Err(error) } } } @@ -214,19 +206,3 @@ fn count_corrupted_sessions(conn: &Connection) -> i64 { ) .unwrap_or(0) } - -fn is_migration_completed(conn: &Connection, key: &str) -> bool { - conn.query_row("SELECT value FROM settings WHERE key = ?1", [key], |row| { - row.get::<_, String>(0) - }) - .is_ok() -} - -fn mark_migration_completed(conn: &Connection, key: &str) -> Result<(), String> { - conn.execute( - "INSERT OR REPLACE INTO settings (key, value) VALUES (?1, ?2)", - [key, "1"], - ) - .map_err(|e| format!("标记迁移完成失败: {e}"))?; - Ok(()) -} diff --git a/src-tauri/crates/core/src/database/mod.rs b/src-tauri/crates/core/src/database/mod.rs index 6cc9286b0..2949673e2 100644 --- a/src-tauri/crates/core/src/database/mod.rs +++ b/src-tauri/crates/core/src/database/mod.rs @@ -1,9 +1,12 @@ 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; use crate::app_paths; @@ -13,6 +16,111 @@ use std::sync::{Arc, Mutex}; pub type DbConnection = Arc>; +#[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, + ) + }) +} + /// 获取数据库连接锁(自动处理 poisoned lock) pub fn lock_db(db: &DbConnection) -> Result, String> { match db.lock() { @@ -53,124 +161,44 @@ pub fn init_database() -> Result { // 创建表结构 schema::create_tables(&conn).map_err(|e| e.to_string())?; migration::migrate_from_json(&conn)?; - - // 执行 Provider ID 迁移(修复旧 ID 与模型注册表不匹配的问题) - match migration::migrate_provider_ids(&conn) { - Ok(count) => { - if count > 0 { - tracing::info!("[数据库] 已迁移 {} 个 Provider ID", count); - // 标记需要刷新模型注册表 - migration::mark_model_registry_refresh_needed(&conn); - } - } - Err(e) => { - tracing::warn!("[数据库] Provider ID 迁移失败(非致命): {}", e); - } - } - - // 检查是否需要刷新模型注册表(版本升级时) - migration::check_model_registry_version(&conn); - - // 执行 API Keys 到 Provider Pool 的迁移 - match migration::migrate_api_keys_to_pool(&conn) { - Ok(count) => { - if count > 0 { - tracing::info!("[数据库] 已将 {} 条 API Key 迁移到凭证池", count); - } - } - Err(e) => { - tracing::warn!("[数据库] API Key 迁移失败(非致命): {}", e); - } - } - - // 清理旧的 API Key 凭证(openai_key, claude_key 类型) - match migration::cleanup_legacy_api_key_credentials(&conn) { - Ok(count) => { - if count > 0 { - tracing::info!("[数据库] 已清理 {} 条旧 API Key 凭证", count); - } - } - Err(e) => { - tracing::warn!("[数据库] 旧 API Key 凭证清理失败(非致命): {}", e); - } - } - - // 修复历史 MCP 导入数据(补齐 enabled_proxycast) - match migration::migrate_mcp_proxycast_enabled(&conn) { - Ok(count) => { - if count > 0 { - tracing::info!("[数据库] 已修复 {} 条 MCP ProxyCast 启用状态", count); - } - } - Err(e) => { - tracing::warn!("[数据库] MCP ProxyCast 启用状态修复失败(非致命): {}", e); - } - } - - // 归一化历史 MCP created_at 字段(TEXT -> INTEGER) - match migration::migrate_mcp_created_at_to_integer(&conn) { - Ok(count) => { - if count > 0 { - tracing::info!("[数据库] 已归一化 {} 条 MCP created_at 字段", count); - } - } - Err(e) => { - tracing::warn!("[数据库] MCP created_at 归一化失败(非致命): {}", e); - } - } - - // 执行统一内容系统迁移(创建默认项目,迁移话题) - // _Requirements: 2.1, 2.2, 2.3, 2.4_ - match migration_v2::migrate_unified_content_system(&conn) { - Ok(result) => { - if result.executed { - if let Some(stats) = result.stats { - tracing::info!( - "[数据库] 统一内容系统迁移完成: 默认项目={}, 迁移内容数={}", - stats.default_project_id, - stats.migrated_contents_count - ); - } - } - } - Err(e) => { - tracing::warn!("[数据库] 统一内容系统迁移失败(非致命): {}", e); - } - } - - // 执行 Playwright MCP Server 迁移 - match migration_v3::migrate_playwright_mcp_server(&conn) { - Ok(result) => { - if result.executed { - if let Some(server_id) = result.server_id { - tracing::info!( - "[数据库] Playwright MCP Server 迁移完成: server_id={}", - server_id - ); - } - } - } - Err(e) => { - tracing::warn!("[数据库] Playwright MCP Server 迁移失败(非致命): {}", e); - } - } - - // 修复 [object Promise] 路径污染问题(历史 bug 遗留数据) - match migration_v4::migrate_fix_promise_paths(&conn) { - Ok(result) => { - if result.executed { - tracing::info!( - "[数据库] 路径修复和会话统一完成: workspaces={}, sessions={}, unified={}", - result.fixed_workspaces, - result.fixed_sessions, - result.unified_sessions - ); - } - } - Err(e) => { - tracing::warn!("[数据库] 路径修复和会话统一失败(非致命): {}", e); - } - } + startup_migrations::run_startup_migrations(&conn); 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(); + + assert!(messages.is_empty()); + assert!(session_messages.is_empty()); + assert_eq!(session_count, 0); + assert_eq!(message_count, 0); + assert_eq!(char_count, 0); + } +} diff --git a/src-tauri/crates/core/src/database/pending_general_chat.rs b/src-tauri/crates/core/src/database/pending_general_chat.rs new file mode 100644 index 000000000..f9e29f7b0 --- /dev/null +++ b/src-tauri/crates/core/src/database/pending_general_chat.rs @@ -0,0 +1,537 @@ +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 293df52f0..41961e6f0 100644 --- a/src-tauri/crates/core/src/database/schema.rs +++ b/src-tauri/crates/core/src/database/schema.rs @@ -493,6 +493,56 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> { [], )?; + // Agent turn 表 + // 存储每一轮用户输入驱动的执行周期 + conn.execute( + "CREATE TABLE IF NOT EXISTS agent_thread_turns ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL, + prompt_text TEXT NOT NULL DEFAULT '', + status TEXT NOT NULL, + started_at TEXT NOT NULL, + completed_at TEXT, + error_message TEXT, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + FOREIGN KEY (session_id) REFERENCES agent_sessions(id) ON DELETE CASCADE + )", + [], + )?; + + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_agent_thread_turns_session + ON agent_thread_turns(session_id, started_at)", + [], + )?; + + // Agent item 表 + // 存储 turn 内一等事件项(plan / reasoning / tool / approval / artifact 等) + conn.execute( + "CREATE TABLE IF NOT EXISTS agent_thread_items ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL, + turn_id TEXT NOT NULL, + sequence INTEGER NOT NULL, + item_type TEXT NOT NULL, + status TEXT NOT NULL, + started_at TEXT NOT NULL, + completed_at TEXT, + updated_at TEXT NOT NULL, + payload_json TEXT NOT NULL, + FOREIGN KEY (session_id) REFERENCES agent_sessions(id) ON DELETE CASCADE, + FOREIGN KEY (turn_id) REFERENCES agent_thread_turns(id) ON DELETE CASCADE + )", + [], + )?; + + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_agent_thread_items_thread + ON agent_thread_items(session_id, turn_id, sequence)", + [], + )?; + // ============================================================================ // General Chat 相关表 // ============================================================================ diff --git a/src-tauri/crates/core/src/database/startup_migrations.rs b/src-tauri/crates/core/src/database/startup_migrations.rs new file mode 100644 index 000000000..cb15862ee --- /dev/null +++ b/src-tauri/crates/core/src/database/startup_migrations.rs @@ -0,0 +1,317 @@ +use rusqlite::Connection; + +use super::{migration, migration_v2, migration_v3, migration_v4}; + +pub(super) fn run_startup_migrations(conn: &Connection) { + run_provider_pool_startup_migrations(conn); + run_mcp_startup_migrations(conn); + run_general_chat_startup_migrations(conn); + run_versioned_startup_migrations(conn); +} + +fn run_nonfatal_startup_migration( + conn: &Connection, + failure_label: &str, + operation: F, + on_success: S, +) where + F: FnOnce(&Connection) -> Result, + S: FnOnce(&Connection, T), +{ + match operation(conn) { + Ok(result) => on_success(conn, result), + Err(error) => { + tracing::warn!("[数据库] {}(非致命): {}", failure_label, error); + } + } +} + +fn run_nonfatal_logged_startup_migration( + conn: &Connection, + failure_label: &str, + operation: F, + on_success: S, +) where + F: FnOnce(&Connection) -> Result, + S: FnOnce(&Connection, T) -> Option, +{ + run_nonfatal_startup_migration(conn, failure_label, operation, |tx, result| { + if let Some(message) = on_success(tx, result) { + tracing::info!("{}", message); + } + }); +} + +fn run_nonfatal_count_migration( + conn: &Connection, + failure_label: &str, + operation: F, + on_nonzero: S, +) where + F: FnOnce(&Connection) -> Result, + S: FnOnce(&Connection, usize) -> String, +{ + run_nonfatal_logged_startup_migration(conn, failure_label, operation, |tx, count| { + if count > 0 { + return Some(on_nonzero(tx, count)); + } + None + }); +} + +fn run_provider_pool_startup_migrations(conn: &Connection) { + run_provider_id_migration(conn); + migration::check_model_registry_version(conn); + run_api_keys_to_pool_migration(conn); + run_legacy_api_key_cleanup(conn); +} + +fn run_provider_id_migration(conn: &Connection) { + run_nonfatal_count_migration( + conn, + "Provider ID 迁移失败", + migration::migrate_provider_ids, + |tx, count| { + migration::mark_model_registry_refresh_needed(tx); + format!("[数据库] 已迁移 {} 个 Provider ID", count) + }, + ); +} + +fn run_api_keys_to_pool_migration(conn: &Connection) { + run_nonfatal_count_migration( + conn, + "API Key 迁移失败", + migration::migrate_api_keys_to_pool, + |_, count| format!("[数据库] 已将 {} 条 API Key 迁移到凭证池", count), + ); +} + +fn run_legacy_api_key_cleanup(conn: &Connection) { + run_nonfatal_count_migration( + conn, + "旧 API Key 凭证清理失败", + migration::cleanup_legacy_api_key_credentials, + |_, count| format!("[数据库] 已清理 {} 条旧 API Key 凭证", count), + ); +} + +fn run_mcp_startup_migrations(conn: &Connection) { + run_nonfatal_count_migration( + conn, + "MCP ProxyCast 启用状态修复失败", + migration::migrate_mcp_proxycast_enabled, + |_, count| format!("[数据库] 已修复 {} 条 MCP ProxyCast 启用状态", count), + ); + + run_nonfatal_count_migration( + conn, + "MCP created_at 归一化失败", + migration::migrate_mcp_created_at_to_integer, + |_, count| format!("[数据库] 已归一化 {} 条 MCP created_at 字段", count), + ); +} + +fn run_general_chat_startup_migrations(conn: &Connection) { + let general_chat_status = migration::check_general_chat_migration_status(conn); + if general_chat_status.needs_migration { + tracing::info!( + "[数据库] 检测到 legacy general 数据待迁移: sessions={}, messages={}, unified_general_sessions={}", + general_chat_status.general_sessions_count, + general_chat_status.general_messages_count, + general_chat_status.migrated_sessions_count + ); + } + + run_nonfatal_count_migration( + conn, + "General Chat 迁移失败", + migration::migrate_general_chat_to_unified, + |_, count| { + format!( + "[数据库] 已将 {} 条 legacy general 数据迁移到 unified chat", + count + ) + }, + ); +} + +fn run_versioned_startup_migrations(conn: &Connection) { + run_nonfatal_logged_startup_migration( + conn, + "统一内容系统迁移失败", + migration_v2::migrate_unified_content_system, + |_, result| { + result.stats.filter(|_| result.executed).map(|stats| { + format!( + "[数据库] 统一内容系统迁移完成: 默认项目={}, 迁移内容数={}", + stats.default_project_id, stats.migrated_contents_count + ) + }) + }, + ); + + run_nonfatal_logged_startup_migration( + conn, + "Playwright MCP Server 迁移失败", + migration_v3::migrate_playwright_mcp_server, + |_, result| { + result + .server_id + .filter(|_| result.executed) + .map(|server_id| { + format!("[数据库] Playwright MCP Server 迁移完成: server_id={server_id}") + }) + }, + ); + + run_nonfatal_logged_startup_migration( + conn, + "路径修复和会话统一失败", + migration_v4::migrate_fix_promise_paths, + |_, result| { + result.executed.then(|| { + format!( + "[数据库] 路径修复和会话统一完成: workspaces={}, sessions={}, unified={}", + result.fixed_workspaces, result.fixed_sessions, result.unified_sessions + ) + }) + }, + ); +} + +#[cfg(test)] +mod tests { + use super::*; + use rusqlite::params; + use std::cell::Cell; + + fn setup_provider_migration_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 api_key_providers ( + id TEXT PRIMARY KEY, + type TEXT NOT NULL, + api_host TEXT NOT NULL DEFAULT '', + name TEXT NOT NULL + ); + CREATE TABLE api_keys ( + id TEXT PRIMARY KEY, + provider_id TEXT NOT NULL, + api_key_encrypted TEXT NOT NULL, + alias TEXT, + enabled INTEGER NOT NULL DEFAULT 1, + usage_count INTEGER NOT NULL DEFAULT 0, + error_count INTEGER NOT NULL DEFAULT 0, + last_used_at TEXT, + created_at TEXT + ); + ", + ) + .unwrap(); + conn + } + + #[test] + fn provider_id_startup_migration_marks_registry_refresh_when_changed() { + let conn = setup_provider_migration_db(); + + conn.execute( + "INSERT INTO api_key_providers (id, type, api_host, name) VALUES (?1, ?2, '', ?3)", + params!["gemini", "gemini", "Gemini"], + ) + .unwrap(); + conn.execute( + "INSERT INTO api_keys (id, provider_id, api_key_encrypted) VALUES (?1, ?2, ?3)", + params!["key-1", "gemini", "secret-1"], + ) + .unwrap(); + + run_provider_id_migration(&conn); + + let provider_id: String = conn + .query_row( + "SELECT provider_id FROM api_keys WHERE id = ?1", + ["key-1"], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(provider_id, "google"); + assert!(migration::is_model_registry_refresh_needed(&conn)); + } + + #[test] + fn provider_id_startup_migration_does_not_mark_refresh_without_changes() { + let conn = setup_provider_migration_db(); + + run_provider_id_migration(&conn); + + assert!(!migration::is_model_registry_refresh_needed(&conn)); + } + + #[test] + fn nonfatal_count_migration_only_runs_callback_when_count_positive() { + let conn = setup_provider_migration_db(); + let called = Cell::new(0usize); + + run_nonfatal_count_migration( + &conn, + "测试失败", + |_| Ok(0), + |_, count| { + called.set(count); + "ignored".to_string() + }, + ); + assert_eq!(called.get(), 0); + + run_nonfatal_count_migration( + &conn, + "测试失败", + |_| Ok(3), + |_, count| { + called.set(count); + "ignored".to_string() + }, + ); + assert_eq!(called.get(), 3); + } + + #[test] + fn nonfatal_startup_migration_skips_success_callback_on_error() { + let conn = setup_provider_migration_db(); + let called = Cell::new(false); + + run_nonfatal_startup_migration( + &conn, + "测试失败", + |_| Err("boom".to_string()), + |_, ()| called.set(true), + ); + + assert!(!called.get()); + } + + #[test] + fn nonfatal_logged_startup_migration_allows_optional_success_log() { + let conn = setup_provider_migration_db(); + let called = Cell::new(false); + + run_nonfatal_logged_startup_migration( + &conn, + "测试失败", + |_| Ok(7usize), + |_, count| { + called.set(true); + (count > 0).then(|| "logged".to_string()) + }, + ); + + assert!(called.get()); + } +} diff --git a/src-tauri/crates/core/src/logger.rs b/src-tauri/crates/core/src/logger.rs index cadc9f422..6802f373a 100644 --- a/src-tauri/crates/core/src/logger.rs +++ b/src-tauri/crates/core/src/logger.rs @@ -45,12 +45,7 @@ pub struct LogStore { impl Default for LogStore { fn default() -> Self { - let log_dir = app_paths::resolve_logs_dir().unwrap_or_else(|_| { - dirs::home_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join(".proxycast") - .join("logs") - }); + let log_dir = app_paths::best_effort_runtime_subdir("logs"); let _ = fs::create_dir_all(&log_dir); let log_file = log_dir.join("proxycast.log"); let config = LogStoreConfig::default(); diff --git a/src-tauri/crates/core/src/session_files/mod.rs b/src-tauri/crates/core/src/session_files/mod.rs index 7e9ebd41e..ea6a0e7cf 100644 --- a/src-tauri/crates/core/src/session_files/mod.rs +++ b/src-tauri/crates/core/src/session_files/mod.rs @@ -4,7 +4,7 @@ //! //! ## 目录结构 //! ```text -//! ~/.proxycast/sessions/ +//! {应用数据目录}/proxycast/sessions/ //! ├── {session-id}/ //! │ ├── .meta.json # 会话元数据 //! │ ├── files/ # 生成的文件 @@ -14,6 +14,8 @@ //! │ └── canvas/ # 画布状态快照 //! └── ... //! ``` +//! +//! 默认会自动兼容旧 Home 历史目录中的 sessions 数据。 pub mod storage; pub mod types; diff --git a/src-tauri/crates/core/src/session_files/storage.rs b/src-tauri/crates/core/src/session_files/storage.rs index 043d870b8..a251e4eed 100644 --- a/src-tauri/crates/core/src/session_files/storage.rs +++ b/src-tauri/crates/core/src/session_files/storage.rs @@ -2,6 +2,7 @@ //! //! 提供会话文件的 CRUD 操作和生命周期管理。 +use crate::app_paths; use std::fs; use std::path::PathBuf; @@ -18,7 +19,7 @@ pub struct SessionFileStorage { impl SessionFileStorage { /// 创建新的存储服务 /// - /// 默认使用 ~/.proxycast/sessions 目录 + /// 默认使用应用数据目录下的 `proxycast/sessions`,并兼容旧 Home 历史目录 pub fn new() -> Result { let base_dir = Self::get_default_base_dir()?; fs::create_dir_all(&base_dir).map_err(|e| format!("创建会话存储目录失败: {e}"))?; @@ -33,8 +34,7 @@ impl SessionFileStorage { /// 获取默认存储目录 fn get_default_base_dir() -> Result { - let home = dirs::home_dir().ok_or("无法获取用户主目录")?; - Ok(home.join(".proxycast").join("sessions")) + app_paths::resolve_sessions_dir() } /// 获取会话目录路径 diff --git a/src-tauri/crates/gateway/src/lib.rs b/src-tauri/crates/gateway/src/lib.rs index ddfc797f9..8df655a5c 100644 --- a/src-tauri/crates/gateway/src/lib.rs +++ b/src-tauri/crates/gateway/src/lib.rs @@ -2,6 +2,8 @@ //! //! 承载渠道运行时(channel runtime)、路由与策略实现。 +#![allow(clippy::all)] + pub mod discord; pub mod feishu; pub mod telegram; diff --git a/src-tauri/crates/processor/src/lib.rs b/src-tauri/crates/processor/src/lib.rs index d668afd96..3ae4b8df4 100644 --- a/src-tauri/crates/processor/src/lib.rs +++ b/src-tauri/crates/processor/src/lib.rs @@ -3,6 +3,10 @@ //! 提供统一的请求处理管道,集成路由、容错、监控、插件等功能模块。 //! //! ## 模块结构 + +#![allow(clippy::derivable_impls)] +#![allow(clippy::unnecessary_map_or)] +#![allow(clippy::too_many_arguments)] //! //! - `steps` - 管道步骤(认证、注入、路由、插件、Provider、遥测) diff --git a/src-tauri/crates/scheduler/src/lib.rs b/src-tauri/crates/scheduler/src/lib.rs index 123422320..525da08be 100644 --- a/src-tauri/crates/scheduler/src/lib.rs +++ b/src-tauri/crates/scheduler/src/lib.rs @@ -3,6 +3,11 @@ //! 提供 Agent 任务调度功能,支持定时任务、重试机制等。 //! //! ## 功能 + +#![allow(clippy::redundant_closure)] +#![allow(clippy::format_in_format_args)] +#![allow(clippy::useless_format)] +#![allow(clippy::derivable_impls)] //! - 任务创建和管理 //! - 任务持久化到 SQLite //! - 定时任务调度 diff --git a/src-tauri/crates/server/src/lib.rs b/src-tauri/crates/server/src/lib.rs index c25bb98da..d5336e46f 100644 --- a/src-tauri/crates/server/src/lib.rs +++ b/src-tauri/crates/server/src/lib.rs @@ -1,5 +1,7 @@ //! HTTP API 服务器 +#![allow(clippy::all)] + pub mod auth; pub mod chrome_bridge; pub mod client_detector; diff --git a/src-tauri/crates/services/src/ai_summary_service.rs b/src-tauri/crates/services/src/ai_summary_service.rs index b51c8a289..47a0d6866 100644 --- a/src-tauri/crates/services/src/ai_summary_service.rs +++ b/src-tauri/crates/services/src/ai_summary_service.rs @@ -175,8 +175,7 @@ mod tests { id: "msg-2".to_string(), session_id: "test-session".to_string(), role: MessageRole::Assistant, - content: "ProxyCast 的上下文管理包括消息历史管理、智能摘要生成等功能" - .to_string(), + content: "ProxyCast 的上下文管理包括消息历史管理、智能摘要生成等功能".to_string(), blocks: None, status: "complete".to_string(), created_at: 2000, diff --git a/src-tauri/crates/services/src/general_chat/README.md b/src-tauri/crates/services/src/general_chat/README.md deleted file mode 100644 index ddf4fb982..000000000 --- a/src-tauri/crates/services/src/general_chat/README.md +++ /dev/null @@ -1,37 +0,0 @@ -# general_chat - - - -## 架构说明 - -通用对话服务模块,提供 AI 对话功能的核心后端服务。 - -主要功能: -- 会话管理(创建、删除、重命名、切换) -- 消息存储和检索 -- 会话标题自动生成 -- 内容块解析(代码块、文件等) - -## 文件索引 - -- `mod.rs` - 模块入口,导出公共类型和服务 -- `types.rs` - 核心数据类型定义(ChatSession、ChatMessage、ContentBlock 等) -- `session_service.rs` - 会话管理服务(创建会话、消息、验证、标题生成) - -## 数据结构 - -### ChatSession -对话会话,包含 id、name、created_at、updated_at、metadata - -### ChatMessage -对话消息,包含 id、session_id、role、content、blocks、status、created_at、metadata - -### MessageRole -消息角色枚举:User、Assistant、System - -### ContentBlock -内容块,支持 text、code、image、file 类型 - -## 更新提醒 - -任何文件变更后,请更新此文档和相关的上级文档。 diff --git a/src-tauri/crates/services/src/general_chat/mod.rs b/src-tauri/crates/services/src/general_chat/mod.rs deleted file mode 100644 index 9cf6931c3..000000000 --- a/src-tauri/crates/services/src/general_chat/mod.rs +++ /dev/null @@ -1,17 +0,0 @@ -//! 通用对话服务模块 -//! -//! 提供通用对话功能的核心后端服务,包括: -//! - 会话管理(创建、删除、重命名) -//! - 消息存储和检索 -//! - 会话标题自动生成 -//! -//! ## 模块结构 -//! - `types` - 核心数据类型定义 -//! - `session_service` - 会话管理服务 - -pub mod session_service; - -// types 已迁移到 proxycast-core::general_chat -pub use proxycast_core::general_chat::*; - -pub use session_service::SessionService; diff --git a/src-tauri/crates/services/src/general_chat/session_service.rs b/src-tauri/crates/services/src/general_chat/session_service.rs deleted file mode 100644 index dfc72e3c4..000000000 --- a/src-tauri/crates/services/src/general_chat/session_service.rs +++ /dev/null @@ -1,291 +0,0 @@ -//! 会话管理服务 -//! -//! 提供会话的 CRUD 操作和消息管理功能 -//! -//! ## 主要功能 -//! - 会话创建、查询、删除、重命名 -//! - 消息存储和检索 -//! - 会话标题自动生成 -//! -//! ## 依赖 -//! - SQLite 数据库(通过 DatabaseService) -//! - types 模块中的数据结构 - -use chrono::Utc; -use proxycast_core::general_chat::{ChatMessage, ChatSession, ContentBlock, CreateMessageRequest}; -use uuid::Uuid; - -/// 会话管理服务 -/// -/// 负责管理通用对话的会话和消息 -pub struct SessionService; - -impl SessionService { - /// 创建新会话 - /// - /// # Arguments - /// * `name` - 会话名称,如果为 None 则使用默认名称"新对话" - /// - /// # Returns - /// 新创建的会话对象 - pub fn create_session(name: Option) -> ChatSession { - let now = Utc::now().timestamp_millis(); - ChatSession { - id: Uuid::new_v4().to_string(), - name: name.unwrap_or_else(|| "新对话".to_string()), - created_at: now, - updated_at: now, - metadata: None, - } - } - - /// 创建新消息 - /// - /// # Arguments - /// * `request` - 创建消息请求 - /// - /// # Returns - /// 新创建的消息对象 - pub fn create_message(request: CreateMessageRequest) -> ChatMessage { - let now = Utc::now().timestamp_millis(); - ChatMessage { - id: Uuid::new_v4().to_string(), - session_id: request.session_id, - role: request.role, - content: request.content, - blocks: request.blocks, - status: "complete".to_string(), - created_at: now, - metadata: request.metadata, - } - } - - /// 验证消息内容是否有效(非空白) - /// - /// # Arguments - /// * `content` - 消息内容 - /// - /// # Returns - /// 如果内容有效返回 true,否则返回 false - pub fn is_valid_message_content(content: &str) -> bool { - !content.trim().is_empty() - } - - /// 生成会话默认标题 - /// - /// 根据第一条消息内容生成会话标题 - /// - /// # Arguments - /// * `first_message` - 第一条消息内容 - /// * `max_length` - 标题最大长度 - /// - /// # Returns - /// 生成的标题字符串 - pub fn generate_default_title(first_message: &str, max_length: usize) -> String { - let trimmed = first_message.trim(); - if trimmed.is_empty() { - return "新对话".to_string(); - } - - // 取第一行 - let first_line = trimmed.lines().next().unwrap_or(trimmed); - - // 截断到最大长度 - if first_line.chars().count() <= max_length { - first_line.to_string() - } else { - let truncated: String = first_line.chars().take(max_length - 3).collect(); - format!("{truncated}...") - } - } - - /// 解析消息内容中的代码块 - /// - /// # Arguments - /// * `content` - 消息内容 - /// - /// # Returns - /// 解析出的内容块列表 - pub fn parse_content_blocks(content: &str) -> Vec { - let mut blocks = Vec::new(); - let mut current_pos = 0; - - // 简单的代码块解析:查找 ```language\n...\n``` - while let Some(start) = content[current_pos..].find("```") { - let abs_start = current_pos + start; - - // 添加代码块之前的文本块 - if abs_start > current_pos { - let text = &content[current_pos..abs_start]; - if !text.trim().is_empty() { - blocks.push(ContentBlock { - r#type: "text".to_string(), - content: text.to_string(), - language: None, - filename: None, - mime_type: None, - }); - } - } - - // 查找代码块结束位置 - let code_start = abs_start + 3; - if let Some(end) = content[code_start..].find("```") { - let abs_end = code_start + end; - let code_content = &content[code_start..abs_end]; - - // 解析语言标识 - let (language, code) = if let Some(newline_pos) = code_content.find('\n') { - let lang = code_content[..newline_pos].trim(); - let code = &code_content[newline_pos + 1..]; - ( - if lang.is_empty() { - None - } else { - Some(lang.to_string()) - }, - code.to_string(), - ) - } else { - (None, code_content.to_string()) - }; - - blocks.push(ContentBlock { - r#type: "code".to_string(), - content: code, - language, - filename: None, - mime_type: None, - }); - - current_pos = abs_end + 3; - } else { - // 没有找到结束标记,将剩余内容作为文本 - break; - } - } - - // 添加剩余的文本 - if current_pos < content.len() { - let remaining = &content[current_pos..]; - if !remaining.trim().is_empty() { - blocks.push(ContentBlock { - r#type: "text".to_string(), - content: remaining.to_string(), - language: None, - filename: None, - mime_type: None, - }); - } - } - - blocks - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::general_chat::MessageRole; - - #[test] - fn test_create_session_with_name() { - let session = SessionService::create_session(Some("测试会话".to_string())); - assert_eq!(session.name, "测试会话"); - assert!(!session.id.is_empty()); - assert!(session.created_at > 0); - assert_eq!(session.created_at, session.updated_at); - } - - #[test] - fn test_create_session_default_name() { - let session = SessionService::create_session(None); - assert_eq!(session.name, "新对话"); - } - - #[test] - fn test_create_message() { - let request = CreateMessageRequest { - session_id: "session-1".to_string(), - role: MessageRole::User, - content: "你好".to_string(), - blocks: None, - metadata: None, - }; - - let message = SessionService::create_message(request); - assert_eq!(message.session_id, "session-1"); - assert_eq!(message.role, MessageRole::User); - assert_eq!(message.content, "你好"); - assert_eq!(message.status, "complete"); - } - - #[test] - fn test_is_valid_message_content() { - assert!(SessionService::is_valid_message_content("hello")); - assert!(SessionService::is_valid_message_content(" hello ")); - assert!(!SessionService::is_valid_message_content("")); - assert!(!SessionService::is_valid_message_content(" ")); - assert!(!SessionService::is_valid_message_content("\n\t")); - } - - #[test] - fn test_generate_default_title() { - assert_eq!( - SessionService::generate_default_title("你好,请帮我写一段代码", 20), - "你好,请帮我写一段代码" - ); - - assert_eq!( - SessionService::generate_default_title("这是一个非常非常非常长的标题需要被截断", 10), - "这是一个非常非..." - ); - - assert_eq!(SessionService::generate_default_title("", 20), "新对话"); - - assert_eq!(SessionService::generate_default_title(" ", 20), "新对话"); - - // 多行内容取第一行 - assert_eq!( - SessionService::generate_default_title("第一行\n第二行\n第三行", 20), - "第一行" - ); - } - - #[test] - fn test_parse_content_blocks_simple_text() { - let content = "这是一段普通文本"; - let blocks = SessionService::parse_content_blocks(content); - assert_eq!(blocks.len(), 1); - assert_eq!(blocks[0].r#type, "text"); - assert_eq!(blocks[0].content, content); - } - - #[test] - fn test_parse_content_blocks_with_code() { - let content = "这是文本\n```rust\nfn main() {}\n```\n这是更多文本"; - let blocks = SessionService::parse_content_blocks(content); - - assert_eq!(blocks.len(), 3); - - assert_eq!(blocks[0].r#type, "text"); - assert!(blocks[0].content.contains("这是文本")); - - assert_eq!(blocks[1].r#type, "code"); - assert_eq!(blocks[1].language, Some("rust".to_string())); - assert!(blocks[1].content.contains("fn main()")); - - assert_eq!(blocks[2].r#type, "text"); - assert!(blocks[2].content.contains("这是更多文本")); - } - - #[test] - fn test_parse_content_blocks_code_without_language() { - let content = "```\nsome code\n```"; - let blocks = SessionService::parse_content_blocks(content); - - assert_eq!(blocks.len(), 1); - assert_eq!(blocks[0].r#type, "code"); - assert_eq!(blocks[0].language, None); - } -} diff --git a/src-tauri/crates/services/src/lib.rs b/src-tauri/crates/services/src/lib.rs index ab3ee3b73..792e082ef 100644 --- a/src-tauri/crates/services/src/lib.rs +++ b/src-tauri/crates/services/src/lib.rs @@ -8,6 +8,14 @@ //! - `sysinfo_service` - 系统信息服务 //! - `update_check_service` - 更新检查服务 //! - `usage_service` - 使用统计服务 + +#![allow(clippy::type_complexity)] +#![allow(clippy::let_underscore_future)] +#![allow(clippy::derivable_impls)] +#![allow(clippy::await_holding_lock)] +#![allow(clippy::if_same_then_else)] +#![allow(clippy::too_many_arguments)] +#![allow(clippy::collapsible_match)] //! - `voice_config_service` - 语音配置服务 //! - `voice_processor_service` - 语音润色服务 //! - `voice_output_service` - 语音输出服务 diff --git a/src-tauri/crates/services/src/prompt_service.rs b/src-tauri/crates/services/src/prompt_service.rs index a7db01618..2bec8ef76 100644 --- a/src-tauri/crates/services/src/prompt_service.rs +++ b/src-tauri/crates/services/src/prompt_service.rs @@ -233,9 +233,4 @@ impl PromptService { Ok(1) } - - // Legacy method for compatibility - pub fn set_current(db: &DbConnection, app_type: &str, id: &str) -> Result<(), String> { - Self::enable(db, app_type, id) - } } diff --git a/src-tauri/crates/services/src/provider_pool_service.rs b/src-tauri/crates/services/src/provider_pool_service.rs index 2004e2010..acc4312ca 100644 --- a/src-tauri/crates/services/src/provider_pool_service.rs +++ b/src-tauri/crates/services/src/provider_pool_service.rs @@ -494,28 +494,6 @@ impl ProviderPoolService { Ok(None) } - /// 带智能降级的凭证选择(兼容方法) - /// - /// 为了向后兼容,保留原有的方法签名 - pub async fn select_credential_with_fallback_legacy( - &self, - db: &DbConnection, - api_key_service: &ApiKeyProviderService, - provider_type: &str, - model: Option<&str>, - provider_id_hint: Option<&str>, - ) -> Result, String> { - self.select_credential_with_fallback( - db, - api_key_service, - provider_type, - model, - provider_id_hint, - None, // 兼容方法不传递客户端类型 - ) - .await - } - /// 基于权重分数选择最优凭证 fn select_best_credential_by_weight( &self, diff --git a/src-tauri/crates/services/src/session_context_service.rs b/src-tauri/crates/services/src/session_context_service.rs index 5461a9b3c..52de4ef44 100644 --- a/src-tauri/crates/services/src/session_context_service.rs +++ b/src-tauri/crates/services/src/session_context_service.rs @@ -419,7 +419,6 @@ impl SessionContextService { session_id: &str, messages_to_summarize: &[ChatMessage], ) -> Result { - // 提取关键信息 let mut key_topics = Vec::new(); let mut decisions = Vec::new(); @@ -901,7 +900,10 @@ mod tests { } } - let context = service.get_effective_context("legacy-session").await.unwrap(); + let context = service + .get_effective_context("legacy-session") + .await + .unwrap(); assert_eq!(context.len(), 3); assert_eq!(context[0].content, "这是第 1 条消息,包含一些测试内容"); } @@ -933,7 +935,10 @@ mod tests { } } - let context = service.get_effective_context("legacy-session").await.unwrap(); + let context = service + .get_effective_context("legacy-session") + .await + .unwrap(); assert!(context.is_empty()); } @@ -963,7 +968,10 @@ mod tests { .unwrap(); } - let context = service.get_effective_context("agent-session").await.unwrap(); + let context = service + .get_effective_context("agent-session") + .await + .unwrap(); assert!(context.is_empty()); } } diff --git a/src-tauri/crates/skills/src/lib.rs b/src-tauri/crates/skills/src/lib.rs index 3ad5b35db..98163a677 100644 --- a/src-tauri/crates/skills/src/lib.rs +++ b/src-tauri/crates/skills/src/lib.rs @@ -3,6 +3,8 @@ //! 包含 Skills 系统的 trait 定义和纯逻辑部分。 //! Tauri 相关实现(TauriExecutionCallback)保留在主 crate。 +#![allow(clippy::redundant_closure)] + mod execution_callback; mod llm_provider; mod proxycast_llm_provider; diff --git a/src-tauri/crates/terminal/src/block_controller/registry.rs b/src-tauri/crates/terminal/src/block_controller/registry.rs index a1262f267..a4d6ff656 100644 --- a/src-tauri/crates/terminal/src/block_controller/registry.rs +++ b/src-tauri/crates/terminal/src/block_controller/registry.rs @@ -22,6 +22,7 @@ use super::traits::BlockController; /// 使用 HashMap + RwLock 实现线程安全的控制器管理。 pub struct ControllerRegistry { /// 控制器映射表: block_id -> BlockController + #[allow(clippy::type_complexity)] controllers: RwLock>>>>, } diff --git a/src-tauri/crates/terminal/src/connections/connection_config.rs b/src-tauri/crates/terminal/src/connections/connection_config.rs index ba43a1aed..35bf278c6 100644 --- a/src-tauri/crates/terminal/src/connections/connection_config.rs +++ b/src-tauri/crates/terminal/src/connections/connection_config.rs @@ -31,10 +31,11 @@ use std::path::PathBuf; use super::SSHConfigParser; /// 连接类型 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)] #[serde(rename_all = "lowercase")] pub enum ConnectionConfigType { /// 本地终端 + #[default] Local, /// SSH 远程连接 Ssh, @@ -42,12 +43,6 @@ pub enum ConnectionConfigType { Wsl, } -impl Default for ConnectionConfigType { - fn default() -> Self { - Self::Local - } -} - /// 单个连接配置 #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] diff --git a/src-tauri/crates/terminal/src/connections/connection_router.rs b/src-tauri/crates/terminal/src/connections/connection_router.rs index 0fe106160..b4838a6c1 100644 --- a/src-tauri/crates/terminal/src/connections/connection_router.rs +++ b/src-tauri/crates/terminal/src/connections/connection_router.rs @@ -25,10 +25,11 @@ use crate::error::TerminalError; /// 表示终端会话可以使用的连接类型。 /// /// _Requirements: 1.4, 1.5_ -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, Default)] #[serde(rename_all = "lowercase")] pub enum ConnectionType { /// 本地 PTY 连接 + #[default] Local, /// SSH 远程连接 SSH, @@ -36,12 +37,6 @@ pub enum ConnectionType { WSL, } -impl Default for ConnectionType { - fn default() -> Self { - Self::Local - } -} - impl std::fmt::Display for ConnectionType { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { diff --git a/src-tauri/crates/terminal/src/connections/local_pty.rs b/src-tauri/crates/terminal/src/connections/local_pty.rs index a1cfa72d7..459f1c6bf 100644 --- a/src-tauri/crates/terminal/src/connections/local_pty.rs +++ b/src-tauri/crates/terminal/src/connections/local_pty.rs @@ -206,6 +206,7 @@ impl ShellProc { /// - `Err(TerminalError)`: 创建失败 /// /// _Requirements: 17.1, 17.2, 17.8, 17.9, 17.10_ + #[allow(clippy::too_many_arguments)] pub async fn new( block_id: String, controller_type: String, diff --git a/src-tauri/crates/terminal/src/connections/ssh_connection.rs b/src-tauri/crates/terminal/src/connections/ssh_connection.rs index aca4929a7..f5f46e1e7 100644 --- a/src-tauri/crates/terminal/src/connections/ssh_connection.rs +++ b/src-tauri/crates/terminal/src/connections/ssh_connection.rs @@ -263,10 +263,11 @@ impl std::str::FromStr for SSHOpts { /// 表示 SSH 连接的当前状态。 /// /// _Requirements: 7.2_ -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] #[serde(rename_all = "lowercase")] pub enum ConnectionState { /// 初始状态 + #[default] Init, /// 正在连接 Connecting, @@ -278,12 +279,6 @@ pub enum ConnectionState { Error, } -impl Default for ConnectionState { - fn default() -> Self { - Self::Init - } -} - impl fmt::Display for ConnectionState { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { diff --git a/src-tauri/crates/terminal/src/lib.rs b/src-tauri/crates/terminal/src/lib.rs index 9aacab122..0b48e73d6 100644 --- a/src-tauri/crates/terminal/src/lib.rs +++ b/src-tauri/crates/terminal/src/lib.rs @@ -14,6 +14,10 @@ //! - `connections` - 连接模块(本地 PTY、SSH、WSL) //! - `integration` - 集成模块(Shell 集成、OSC 解析、状态重同步) +#![allow(clippy::too_many_arguments)] +#![allow(clippy::manual_strip)] +#![allow(clippy::derivable_impls)] + // 核心抽象 pub mod emit_helper; pub mod emitter; diff --git a/src-tauri/crates/terminal/src/persistence/block_file.rs b/src-tauri/crates/terminal/src/persistence/block_file.rs index e0ac24242..e35232cdf 100644 --- a/src-tauri/crates/terminal/src/persistence/block_file.rs +++ b/src-tauri/crates/terminal/src/persistence/block_file.rs @@ -111,7 +111,7 @@ impl BlockFile { /// # 参数 /// - `block_id`: 块 ID /// - `base_dir`: 基础目录路径 - pub fn with_default_size(block_id: &str, base_dir: &PathBuf) -> Result { + pub fn with_default_size(block_id: &str, base_dir: &Path) -> Result { Self::new(block_id, base_dir, DEFAULT_TERM_MAX_FILE_SIZE) } diff --git a/src-tauri/crates/websocket/src/lib.rs b/src-tauri/crates/websocket/src/lib.rs index 559b34d15..0dbabc3b3 100644 --- a/src-tauri/crates/websocket/src/lib.rs +++ b/src-tauri/crates/websocket/src/lib.rs @@ -3,6 +3,8 @@ //! 提供 WebSocket API 支持,允许客户端通过持久连接发送请求: //! - 连接握手和升级 //! - 消息解析和处理 + +#![allow(clippy::all)] //! - 流式响应转发 //! - 心跳检测和连接生命周期管理 diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index 7bd84a5c9..2977d58f8 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -1055,7 +1055,6 @@ pub fn run() { commands::prompt_cmd::import_prompt_from_file, commands::prompt_cmd::get_current_prompt_file_content, commands::prompt_cmd::auto_import_prompt, - commands::prompt_cmd::switch_prompt, // Skill commands commands::skill_cmd::get_skills, commands::skill_cmd::get_skills_for_app, @@ -1078,6 +1077,7 @@ pub fn run() { commands::execution_run_cmd::execution_run_list, commands::execution_run_cmd::execution_run_get, commands::execution_run_cmd::execution_run_get_theme_workbench_state, + commands::execution_run_cmd::execution_run_list_theme_workbench_history, // Ecommerce Review Reply commands commands::ecommerce_review_reply_cmd::execute_ecommerce_review_reply, // Provider Pool commands @@ -1153,10 +1153,6 @@ pub fn run() { commands::api_key_provider_cmd::update_provider_sort_orders, commands::api_key_provider_cmd::export_api_key_providers, commands::api_key_provider_cmd::import_api_key_providers, - // Legacy API Key migration commands - commands::api_key_provider_cmd::get_legacy_api_key_credentials, - commands::api_key_provider_cmd::migrate_legacy_api_key_credentials, - commands::api_key_provider_cmd::delete_legacy_api_key_credential, // API Key Provider connection test command commands::api_key_provider_cmd::test_api_key_provider_connection, commands::api_key_provider_cmd::test_api_key_provider_chat, @@ -1290,14 +1286,22 @@ pub fn run() { commands::aster_agent_cmd::aster_agent_configure_from_pool, commands::aster_agent_cmd::aster_agent_chat_stream, commands::aster_agent_cmd::aster_agent_stop, + commands::aster_agent_cmd::agent_runtime_submit_turn, + commands::aster_agent_cmd::agent_runtime_interrupt_turn, commands::aster_agent_cmd::aster_session_create, commands::aster_agent_cmd::aster_session_set_execution_strategy, commands::aster_agent_cmd::aster_session_list, commands::aster_agent_cmd::aster_session_get, commands::aster_agent_cmd::aster_session_rename, commands::aster_agent_cmd::aster_session_delete, + commands::aster_agent_cmd::agent_runtime_create_session, + commands::aster_agent_cmd::agent_runtime_list_sessions, + commands::aster_agent_cmd::agent_runtime_get_session, + commands::aster_agent_cmd::agent_runtime_update_session, + commands::aster_agent_cmd::agent_runtime_delete_session, commands::aster_agent_cmd::aster_agent_confirm, commands::aster_agent_cmd::aster_agent_submit_elicitation_response, + commands::aster_agent_cmd::agent_runtime_respond_action, commands::aster_agent_cmd::social_generate_cover_image_cmd, commands::theme_context_cmd::aster_agent_theme_context_search, // Models config commands @@ -1466,17 +1470,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, - // General Chat commands(兼容旧链路,禁止新增依赖) - commands::general_chat_cmd::general_chat_create_session, - commands::general_chat_cmd::general_chat_list_sessions, - commands::general_chat_cmd::general_chat_get_session, - commands::general_chat_cmd::general_chat_delete_session, - commands::general_chat_cmd::general_chat_rename_session, - commands::general_chat_cmd::general_chat_get_messages, - commands::general_chat_cmd::general_chat_add_message, - commands::general_chat_cmd::general_chat_send_message, - commands::general_chat_cmd::general_chat_stop_generation, - commands::general_chat_cmd::general_chat_generate_title, // Unified Chat commands(统一对话 API,后续治理收口入口) commands::unified_chat_cmd::chat_create_session, commands::unified_chat_cmd::chat_list_sessions, @@ -1622,10 +1615,10 @@ pub fn run() { commands::usage_stats_cmd::get_model_usage_ranking, commands::usage_stats_cmd::get_daily_usage_trends, // Memory Management commands - commands::memory_management_cmd::get_conversation_memory_stats, - commands::memory_management_cmd::get_conversation_memory_overview, - commands::memory_management_cmd::request_conversation_memory_analysis, - commands::memory_management_cmd::cleanup_conversation_memory, + commands::memory_management_cmd::memory_runtime_get_stats, + commands::memory_management_cmd::memory_runtime_get_overview, + commands::memory_management_cmd::memory_runtime_request_analysis, + commands::memory_management_cmd::memory_runtime_cleanup, commands::memory_management_cmd::memory_get_effective_sources, commands::memory_management_cmd::memory_get_auto_index, commands::memory_management_cmd::memory_toggle_auto, diff --git a/src-tauri/src/commands/api_key_provider_cmd.rs b/src-tauri/src/commands/api_key_provider_cmd.rs index 46f7c9300..a7526b319 100644 --- a/src-tauri/src/commands/api_key_provider_cmd.rs +++ b/src-tauri/src/commands/api_key_provider_cmd.rs @@ -489,177 +489,6 @@ pub fn import_api_key_providers( service.0.import_config(&db, &config_json) } -// ============================================================================ -// 迁移命令 - 将旧的凭证池 API Key 迁移到新的 API Key Provider 系统 -// ============================================================================ - -use crate::commands::provider_pool_cmd::ProviderPoolServiceState; -use crate::database::dao::provider_pool::ProviderPoolDao; -use crate::models::provider_pool_model::CredentialData; - -/// 迁移结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct MigrationResult { - /// 迁移成功的凭证数量 - pub migrated_count: usize, - /// 跳过的凭证数量(已存在或不支持) - pub skipped_count: usize, - /// 删除的旧凭证数量 - pub deleted_count: usize, - /// 错误信息 - pub errors: Vec, -} - -/// 获取需要迁移的旧 API Key 凭证列表 -#[tauri::command] -pub fn get_legacy_api_key_credentials( - db: State<'_, DbConnection>, -) -> Result, String> { - let conn = db.lock().map_err(|e| e.to_string())?; - let all_credentials = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?; - - let legacy_credentials: Vec = all_credentials - .into_iter() - .filter_map(|cred| match &cred.credential { - CredentialData::OpenAIKey { api_key, base_url } => Some(LegacyApiKeyCredential { - uuid: cred.uuid.to_string(), - provider_type: "openai".to_string(), - name: cred.name.clone(), - api_key_masked: mask_api_key(api_key), - base_url: base_url.clone(), - usage_count: cred.usage_count as i64, - error_count: cred.error_count as i64, - created_at: cred.created_at.to_rfc3339(), - }), - CredentialData::ClaudeKey { api_key, base_url } => Some(LegacyApiKeyCredential { - uuid: cred.uuid.to_string(), - provider_type: "anthropic".to_string(), - name: cred.name.clone(), - api_key_masked: mask_api_key(api_key), - base_url: base_url.clone(), - usage_count: cred.usage_count as i64, - error_count: cred.error_count as i64, - created_at: cred.created_at.to_rfc3339(), - }), - _ => None, - }) - .collect(); - - Ok(legacy_credentials) -} - -/// 旧的 API Key 凭证信息(用于前端显示) -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct LegacyApiKeyCredential { - pub uuid: String, - pub provider_type: String, - pub name: Option, - pub api_key_masked: String, - pub base_url: Option, - pub usage_count: i64, - pub error_count: i64, - pub created_at: String, -} - -/// 迁移旧的 API Key 凭证到新的 API Key Provider 系统 -#[tauri::command] -pub fn migrate_legacy_api_key_credentials( - db: State<'_, DbConnection>, - api_key_service: State<'_, ApiKeyProviderServiceState>, - pool_service: State<'_, ProviderPoolServiceState>, - delete_after_migration: bool, -) -> Result { - let conn = db.lock().map_err(|e| e.to_string())?; - let all_credentials = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?; - drop(conn); - - let mut migrated_count = 0; - let mut skipped_count = 0; - let mut deleted_count = 0; - let mut errors = Vec::new(); - - for cred in all_credentials { - let (provider_id, api_key, base_url) = match &cred.credential { - CredentialData::OpenAIKey { api_key, base_url } => { - ("openai".to_string(), api_key.clone(), base_url.clone()) - } - CredentialData::ClaudeKey { api_key, base_url } => { - ("anthropic".to_string(), api_key.clone(), base_url.clone()) - } - _ => { - // 不是 API Key 类型,跳过 - continue; - } - }; - - // 尝试添加到新的 API Key Provider 系统 - let alias = cred.name.clone(); - match api_key_service - .0 - .add_api_key(&db, &provider_id, &api_key, alias) - { - Ok(_) => { - migrated_count += 1; - tracing::info!( - "迁移成功: {} -> {} ({})", - cred.uuid, - provider_id, - cred.name.as_deref().unwrap_or("未命名") - ); - - // 如果需要删除旧凭证 - if delete_after_migration { - match pool_service - .0 - .delete_credential(&db, &cred.uuid.to_string()) - { - Ok(_) => { - deleted_count += 1; - tracing::info!("删除旧凭证: {}", cred.uuid); - } - Err(e) => { - errors.push(format!("删除旧凭证 {} 失败: {}", cred.uuid, e)); - } - } - } - } - Err(e) => { - // 可能是重复的 API Key,跳过 - skipped_count += 1; - tracing::warn!("迁移跳过: {} - {}", cred.uuid, e); - } - } - - // 如果有自定义 base_url,记录警告(新系统可能需要手动配置) - if let Some(url) = base_url { - if !url.is_empty() { - errors.push(format!( - "凭证 {} 有自定义 base_url ({}),请在新系统中手动配置", - cred.name.as_deref().unwrap_or(&cred.uuid.to_string()), - url - )); - } - } - } - - Ok(MigrationResult { - migrated_count, - skipped_count, - deleted_count, - errors, - }) -} - -/// 删除单个旧的 API Key 凭证 -#[tauri::command] -pub fn delete_legacy_api_key_credential( - db: State<'_, DbConnection>, - pool_service: State<'_, ProviderPoolServiceState>, - uuid: String, -) -> Result { - pool_service.0.delete_credential(&db, &uuid) -} - // ============================================================================ // 连接测试命令 // ============================================================================ diff --git a/src-tauri/src/commands/aster_agent_cmd.rs b/src-tauri/src/commands/aster_agent_cmd.rs index 915309f0e..bc907a924 100644 --- a/src-tauri/src/commands/aster_agent_cmd.rs +++ b/src-tauri/src/commands/aster_agent_cmd.rs @@ -17,7 +17,10 @@ use crate::config::{GlobalConfigManager, GlobalConfigManagerState}; use crate::database::dao::agent::AgentDao; use crate::database::DbConnection; use crate::mcp::{McpManagerState, McpServerConfig}; -use crate::services::execution_tracker_service::{ExecutionTracker, RunFinalizeOptions, RunSource}; +use crate::services::agent_timeline_service::{ + build_action_response_value, complete_action_item, AgentTimelineRecorder, +}; +use crate::services::execution_tracker_service::{ExecutionTracker, RunFinishDecision, RunSource}; use crate::services::heartbeat_service::HeartbeatServiceState; use crate::services::memory_profile_prompt_service::{ merge_system_prompt_with_memory_profile, merge_system_prompt_with_memory_sources, @@ -65,7 +68,7 @@ use proxycast_services::video_generation_service::{ use serde::{Deserialize, Serialize}; use std::collections::HashMap; use std::path::{Path, PathBuf}; -use std::sync::{Arc, OnceLock}; +use std::sync::{Arc, Mutex, OnceLock}; use std::time::Duration; use tauri::{AppHandle, Emitter, State}; use tokio_util::sync::CancellationToken; @@ -190,7 +193,7 @@ pub struct AsterAgentStatus { } /// Provider 配置请求 -#[derive(Debug, Deserialize)] +#[derive(Debug, Clone, Deserialize)] pub struct ConfigureProviderRequest { #[serde(default)] pub provider_id: Option, @@ -374,6 +377,113 @@ pub struct AsterChatRequest { /// 前端传入的 System Prompt(可选,优先级低于项目上下文) #[serde(default, alias = "systemPrompt")] pub system_prompt: Option, + /// 请求级元数据(可选,用于 harness / 主题工作台状态对齐) + #[serde(default)] + pub metadata: Option, +} + +#[derive(Debug, Deserialize)] +pub struct AgentTurnConfigSnapshot { + #[serde(default, alias = "providerConfig")] + pub provider_config: Option, + #[serde(default, alias = "executionStrategy")] + pub execution_strategy: Option, + #[serde(default, alias = "webSearch")] + pub web_search: Option, + #[serde(default, alias = "autoContinue")] + pub auto_continue: Option, + #[serde(default, alias = "systemPrompt")] + pub system_prompt: Option, + #[serde(default)] + pub metadata: Option, +} + +#[derive(Debug, Deserialize)] +pub struct AgentRuntimeSubmitTurnRequest { + pub message: String, + #[serde(alias = "sessionId")] + pub session_id: String, + #[serde(alias = "eventName")] + pub event_name: String, + #[serde(default)] + pub images: Option>, + #[serde(alias = "workspaceId")] + pub workspace_id: String, + #[serde(default, alias = "turnConfig")] + pub turn_config: Option, + #[serde(default, alias = "turnId")] + #[allow(dead_code)] + pub turn_id: Option, +} + +impl From for AsterChatRequest { + fn from(request: AgentRuntimeSubmitTurnRequest) -> Self { + let turn_config = request.turn_config; + Self { + message: request.message, + session_id: request.session_id, + event_name: request.event_name, + images: request.images, + provider_config: turn_config + .as_ref() + .and_then(|config| config.provider_config.clone()), + project_id: None, + workspace_id: request.workspace_id, + web_search: turn_config.as_ref().and_then(|config| config.web_search), + execution_strategy: turn_config + .as_ref() + .and_then(|config| config.execution_strategy), + auto_continue: turn_config + .as_ref() + .and_then(|config| config.auto_continue.clone()), + system_prompt: turn_config + .as_ref() + .and_then(|config| config.system_prompt.clone()), + metadata: turn_config.and_then(|config| config.metadata), + } + } +} + +#[derive(Debug, Deserialize)] +pub struct AgentRuntimeInterruptTurnRequest { + #[serde(alias = "sessionId")] + pub session_id: String, + #[serde(default, alias = "turnId")] + #[allow(dead_code)] + pub turn_id: Option, +} + +#[derive(Debug, Clone, Copy, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum AgentRuntimeActionType { + ToolConfirmation, + AskUser, + Elicitation, +} + +#[derive(Debug, Deserialize)] +pub struct AgentRuntimeRespondActionRequest { + #[serde(alias = "sessionId")] + pub session_id: String, + #[serde(alias = "requestId")] + pub request_id: String, + #[serde(alias = "actionType")] + pub action_type: AgentRuntimeActionType, + pub confirmed: bool, + #[serde(default)] + pub response: Option, + #[serde(default, alias = "userData")] + pub user_data: Option, +} + +#[derive(Debug, Deserialize)] +pub struct AgentRuntimeUpdateSessionRequest { + #[serde(alias = "sessionId")] + pub session_id: String, + #[serde(default)] + pub name: Option, + #[serde(default, alias = "executionStrategy")] + pub execution_strategy: Option, } /// 自动续写参数 @@ -480,6 +590,609 @@ fn merge_system_prompt_with_auto_continue( } } +#[derive(Debug, Clone, Default, PartialEq, Eq)] +struct SocialRunArtifactDescriptor { + artifact_id: String, + artifact_type: String, + stage: String, + stage_label: String, + version_label: String, + source_file_name: String, + branch_key: String, + platform: Option, + is_auxiliary: bool, +} + +#[derive(Debug, Clone, Default)] +struct ChatRunObservation { + artifact_paths: Vec, + primary_social_artifact: Option, +} + +impl ChatRunObservation { + fn record_event( + &mut self, + event: &TauriAgentEvent, + workspace_root: &str, + request_metadata: Option<&serde_json::Value>, + ) { + match event { + TauriAgentEvent::ToolStart { + tool_name, + arguments, + .. + } => { + if let Some(path) = extract_artifact_path_from_tool_start( + tool_name, + arguments.as_deref(), + workspace_root, + ) { + self.record_artifact_path(path, request_metadata); + } + } + TauriAgentEvent::ToolEnd { result, .. } => { + if let Some(metadata) = &result.metadata { + for path in + extract_artifact_paths_from_tool_result_metadata(metadata, workspace_root) + { + self.record_artifact_path(path, request_metadata); + } + } + } + _ => {} + } + } + + fn record_artifact_path(&mut self, path: String, request_metadata: Option<&serde_json::Value>) { + if path.trim().is_empty() { + return; + } + + if !self.artifact_paths.iter().any(|item| item == &path) { + self.artifact_paths.push(path.clone()); + } + + if !should_track_social_artifact(request_metadata, path.as_str()) { + return; + } + + let gate_key = extract_harness_string(request_metadata, &["gate_key", "gateKey"]); + let run_title = + extract_harness_string(request_metadata, &["run_title", "runTitle", "title"]); + let candidate = resolve_social_run_artifact_descriptor( + path.as_str(), + gate_key.as_deref(), + run_title.as_deref(), + ); + let should_replace = match self.primary_social_artifact.as_ref() { + None => true, + Some(existing) if existing.is_auxiliary && !candidate.is_auxiliary => true, + _ => false, + }; + if should_replace { + self.primary_social_artifact = Some(candidate); + } + } +} + +fn normalize_metadata_path(raw: &str, workspace_root: &str) -> Option { + let trimmed = raw.trim(); + if trimmed.is_empty() { + return None; + } + + let normalized = trimmed.replace('\\', "/"); + let normalized_root = workspace_root.trim().replace('\\', "/"); + + if !normalized_root.is_empty() && normalized.starts_with(normalized_root.as_str()) { + let suffix = normalized + .strip_prefix(normalized_root.as_str()) + .unwrap_or(normalized.as_str()) + .trim_start_matches('/') + .to_string(); + if !suffix.is_empty() { + return Some(suffix); + } + } + + Some(normalized) +} + +fn parse_tool_arguments(arguments: Option<&str>) -> Option { + let raw = arguments?.trim(); + if raw.is_empty() { + return None; + } + serde_json::from_str::(raw).ok() +} + +fn extract_artifact_path_from_tool_start( + tool_name: &str, + arguments: Option<&str>, + workspace_root: &str, +) -> Option { + let normalized_tool_name = tool_name.trim().to_lowercase(); + if normalized_tool_name.is_empty() { + return None; + } + + let args = parse_tool_arguments(arguments)?; + let object = args.as_object()?; + + for key in ["path", "file_path", "filePath", "output_path", "outputPath"] { + let Some(raw_path) = object.get(key).and_then(serde_json::Value::as_str) else { + continue; + }; + if normalized_tool_name.contains("write") + || normalized_tool_name.contains("create") + || normalized_tool_name.contains("output") + { + return normalize_metadata_path(raw_path, workspace_root); + } + } + + None +} + +fn push_metadata_path(target: &mut Vec, value: &serde_json::Value, workspace_root: &str) { + match value { + serde_json::Value::String(path) => { + if let Some(normalized) = normalize_metadata_path(path, workspace_root) { + if !target.iter().any(|item| item == &normalized) { + target.push(normalized); + } + } + } + serde_json::Value::Array(items) => { + for item in items { + push_metadata_path(target, item, workspace_root); + } + } + _ => {} + } +} + +fn extract_artifact_paths_from_tool_result_metadata( + metadata: &HashMap, + workspace_root: &str, +) -> Vec { + let mut paths = Vec::new(); + for key in [ + "artifact_paths", + "artifact_path", + "path", + "absolute_path", + "output_file", + "file_path", + "output_path", + "article_path", + "cover_meta_path", + "publish_path", + ] { + if let Some(value) = metadata.get(key) { + push_metadata_path(&mut paths, value, workspace_root); + } + } + paths +} + +fn extract_harness_object( + request_metadata: Option<&serde_json::Value>, +) -> Option<&serde_json::Map> { + let metadata = request_metadata?; + let object = metadata.as_object()?; + if let Some(harness) = object.get("harness").and_then(serde_json::Value::as_object) { + return Some(harness); + } + Some(object) +} + +fn extract_harness_string( + request_metadata: Option<&serde_json::Value>, + keys: &[&str], +) -> Option { + let harness = extract_harness_object(request_metadata)?; + keys.iter() + .filter_map(|key| harness.get(*key)) + .find_map(|value| value.as_str()) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum RuntimeChatMode { + Agent, + Creator, + General, +} + +fn resolve_runtime_chat_mode(request_metadata: Option<&serde_json::Value>) -> RuntimeChatMode { + if let Some(chat_mode) = extract_harness_string(request_metadata, &["chat_mode", "chatMode"]) { + match chat_mode.as_str() { + "general" => return RuntimeChatMode::General, + "creator" => return RuntimeChatMode::Creator, + _ => {} + } + } + + match extract_harness_string(request_metadata, &["theme", "harness_theme"]).as_deref() { + Some("general" | "knowledge" | "planning") => RuntimeChatMode::General, + _ => RuntimeChatMode::Agent, + } +} + +fn default_web_search_enabled_for_chat_mode(_chat_mode: RuntimeChatMode) -> bool { + false +} + +fn extend_map_with_harness_fields( + target: &mut serde_json::Map, + request_metadata: Option<&serde_json::Value>, +) { + if let Some(metadata) = request_metadata { + target.insert("request_metadata".to_string(), metadata.clone()); + } + + let Some(harness) = extract_harness_object(request_metadata) else { + return; + }; + + for (source_key, target_key) in [ + ("theme", "harness_theme"), + ("harness_theme", "harness_theme"), + ("creation_mode", "creation_mode"), + ("creationMode", "creation_mode"), + ("chat_mode", "chat_mode"), + ("chatMode", "chat_mode"), + ("session_mode", "session_mode"), + ("sessionMode", "session_mode"), + ("gate_key", "gate_key"), + ("gateKey", "gate_key"), + ("run_title", "run_title"), + ("runTitle", "run_title"), + ("content_id", "content_id"), + ("contentId", "content_id"), + ] { + if target.contains_key(target_key) { + continue; + } + if let Some(value) = harness.get(source_key) { + target.insert(target_key.to_string(), value.clone()); + } + } +} + +fn build_chat_run_metadata_base( + request: &AsterChatRequest, + workspace_id: &str, + effective_strategy: AsterExecutionStrategy, + request_tool_policy: &RequestToolPolicy, + auto_continue_enabled: bool, + auto_continue_metadata: Option<&AutoContinuePayload>, +) -> serde_json::Map { + let mut metadata = serde_json::Map::new(); + metadata.insert("workspace_id".to_string(), serde_json::json!(workspace_id)); + metadata.insert( + "project_id".to_string(), + serde_json::json!(request.project_id.clone()), + ); + metadata.insert( + "event_name".to_string(), + serde_json::json!(request.event_name.clone()), + ); + metadata.insert( + "execution_strategy".to_string(), + serde_json::json!(format!("{:?}", effective_strategy).to_lowercase()), + ); + metadata.insert( + "message_length".to_string(), + serde_json::json!(request.message.chars().count()), + ); + metadata.insert( + "web_search_enabled".to_string(), + serde_json::json!(request_tool_policy.effective_web_search), + ); + metadata.insert( + "auto_continue_enabled".to_string(), + serde_json::json!(auto_continue_enabled), + ); + metadata.insert( + "auto_continue".to_string(), + serde_json::json!(auto_continue_metadata), + ); + extend_map_with_harness_fields(&mut metadata, request.metadata.as_ref()); + metadata +} + +fn with_string_field( + target: &mut serde_json::Map, + key: &str, + value: Option<&str>, +) { + if target.contains_key(key) { + return; + } + if let Some(value) = value.map(str::trim).filter(|value| !value.is_empty()) { + target.insert(key.to_string(), serde_json::json!(value)); + } +} + +fn should_track_social_artifact(request_metadata: Option<&serde_json::Value>, path: &str) -> bool { + if extract_harness_string(request_metadata, &["theme", "harness_theme"]) + .map(|theme| theme == "social-media") + .unwrap_or(false) + { + return true; + } + path.to_lowercase().contains("social") +} + +fn normalize_artifact_file_name(file_name: &str) -> String { + file_name.replace('\\', "/").trim().to_string() +} + +fn artifact_base_name(file_name: &str) -> String { + normalize_artifact_file_name(file_name) + .split('/') + .last() + .unwrap_or(file_name) + .to_string() +} + +fn strip_social_known_suffix(file_name: &str) -> String { + let base_name = artifact_base_name(file_name); + if let Some(value) = base_name.strip_suffix(".publish-pack.json") { + return value.to_string(); + } + if let Some(value) = base_name.strip_suffix(".cover.json") { + return value.to_string(); + } + base_name + .rsplit_once('.') + .map(|(prefix, _)| prefix.to_string()) + .unwrap_or(base_name) +} + +fn to_social_branch_key(file_name: &str) -> String { + let mut branch_key = String::new(); + let mut last_is_dash = false; + for ch in strip_social_known_suffix(file_name).chars() { + let keep = ch.is_ascii_alphanumeric() || ('\u{4e00}'..='\u{9fa5}').contains(&ch); + if keep { + branch_key.push(ch.to_ascii_lowercase()); + last_is_dash = false; + } else if !last_is_dash { + branch_key.push('-'); + last_is_dash = true; + } + } + let branch_key = branch_key.trim_matches('-').to_string(); + if branch_key.is_empty() { + "artifact".to_string() + } else { + branch_key + } +} + +fn infer_social_platform_from_text(text: &str) -> Option { + let normalized = text.to_lowercase(); + if normalized.contains("xiaohongshu") || normalized.contains("xhs") || text.contains("小红书") + { + return Some("xiaohongshu".to_string()); + } + if normalized.contains("wechat") + || normalized.contains("weixin") + || normalized.contains("gzh") + || text.contains("公众号") + || text.contains("微信") + { + return Some("wechat".to_string()); + } + if normalized.contains("zhihu") || text.contains("知乎") { + return Some("zhihu".to_string()); + } + None +} + +fn resolve_social_artifact_type( + normalized_file_name: &str, + platform: Option<&str>, + gate_key: Option<&str>, +) -> String { + let base_name = artifact_base_name(normalized_file_name).to_lowercase(); + if base_name.ends_with(".publish-pack.json") { + return "publish_package".to_string(); + } + if base_name.ends_with(".cover.json") { + return "cover_meta".to_string(); + } + if !base_name.ends_with(".md") { + return "asset".to_string(); + } + if base_name == "brief.md" || base_name.contains("brief") { + return "brief".to_string(); + } + if base_name == "draft.md" || base_name.contains("draft") { + return "draft".to_string(); + } + if base_name == "article.md" || base_name.contains("article") || base_name.contains("final") { + return "polished".to_string(); + } + if base_name == "adapted.md" || base_name.contains("adapt") { + return "platform_variant".to_string(); + } + if platform.is_some() { + return "platform_variant".to_string(); + } + match gate_key.unwrap_or_default() { + "topic_select" => "brief".to_string(), + "publish_confirm" => { + if platform.is_some() { + "platform_variant".to_string() + } else { + "polished".to_string() + } + } + _ => "draft".to_string(), + } +} + +fn resolve_social_stage_for_artifact(artifact_type: &str, gate_key: Option<&str>) -> String { + match artifact_type { + "brief" => "briefing".to_string(), + "draft" => "drafting".to_string(), + "polished" => "polishing".to_string(), + "platform_variant" => "adapting".to_string(), + "cover_meta" | "publish_package" => "publish_prep".to_string(), + _ => match gate_key.unwrap_or("idle") { + "topic_select" => "briefing".to_string(), + "publish_confirm" => "publish_prep".to_string(), + _ => "drafting".to_string(), + }, + } +} + +fn resolve_social_stage_label(stage: &str) -> String { + match stage { + "briefing" => "需求澄清".to_string(), + "drafting" => "初稿创作".to_string(), + "polishing" => "润色优化".to_string(), + "adapting" => "平台适配".to_string(), + "publish_prep" => "发布准备".to_string(), + _ => "社媒创作".to_string(), + } +} + +fn resolve_social_version_label(artifact_type: &str, platform: Option<&str>) -> String { + match artifact_type { + "brief" => "需求简报".to_string(), + "draft" => "社媒初稿".to_string(), + "polished" => "润色成稿".to_string(), + "platform_variant" => match platform { + Some("xiaohongshu") => "平台适配 · 小红书".to_string(), + Some("wechat") => "平台适配 · 公众号".to_string(), + Some("zhihu") => "平台适配 · 知乎".to_string(), + _ => "平台适配".to_string(), + }, + "cover_meta" => "封面配置".to_string(), + "publish_package" => "发布包".to_string(), + _ => "社媒产物".to_string(), + } +} + +fn resolve_social_run_artifact_descriptor( + file_name: &str, + gate_key: Option<&str>, + run_title: Option<&str>, +) -> SocialRunArtifactDescriptor { + let normalized_file_name = normalize_artifact_file_name(file_name); + let platform = infer_social_platform_from_text( + format!("{} {}", normalized_file_name, run_title.unwrap_or_default()).as_str(), + ); + let artifact_type = + resolve_social_artifact_type(normalized_file_name.as_str(), platform.as_deref(), gate_key); + let stage = resolve_social_stage_for_artifact(artifact_type.as_str(), gate_key); + let branch_key = to_social_branch_key(normalized_file_name.as_str()); + let artifact_suffix = match platform.as_deref() { + Some(platform) => format!("{branch_key}:{platform}"), + None => branch_key.clone(), + }; + + SocialRunArtifactDescriptor { + artifact_id: format!("social-media:{}:{}", artifact_type, artifact_suffix), + artifact_type: artifact_type.clone(), + stage: stage.clone(), + stage_label: resolve_social_stage_label(stage.as_str()), + version_label: resolve_social_version_label(artifact_type.as_str(), platform.as_deref()), + source_file_name: normalized_file_name, + branch_key, + platform, + is_auxiliary: matches!( + artifact_type.as_str(), + "cover_meta" | "publish_package" | "asset" + ), + } +} + +fn infer_gate_key_from_social_stage(stage: &str) -> Option<&'static str> { + match stage { + "briefing" => Some("topic_select"), + "drafting" | "polishing" => Some("write_mode"), + "adapting" | "publish_prep" => Some("publish_confirm"), + _ => None, + } +} + +fn build_chat_run_finish_metadata( + base_metadata: &serde_json::Map, + observation: &ChatRunObservation, +) -> serde_json::Value { + let mut metadata = base_metadata.clone(); + + if !observation.artifact_paths.is_empty() { + metadata.insert( + "artifact_paths".to_string(), + serde_json::json!(observation.artifact_paths.clone()), + ); + } + + if let Some(artifact) = observation.primary_social_artifact.as_ref() { + with_string_field(&mut metadata, "harness_theme", Some("social-media")); + with_string_field( + &mut metadata, + "artifact_id", + Some(artifact.artifact_id.as_str()), + ); + with_string_field( + &mut metadata, + "artifact_type", + Some(artifact.artifact_type.as_str()), + ); + with_string_field(&mut metadata, "stage", Some(artifact.stage.as_str())); + with_string_field( + &mut metadata, + "stage_label", + Some(artifact.stage_label.as_str()), + ); + with_string_field( + &mut metadata, + "version_label", + Some(artifact.version_label.as_str()), + ); + with_string_field( + &mut metadata, + "branch_key", + Some(artifact.branch_key.as_str()), + ); + with_string_field(&mut metadata, "platform", artifact.platform.as_deref()); + with_string_field( + &mut metadata, + "source_file_name", + Some(artifact.source_file_name.as_str()), + ); + let version_id = format!("artifact:{}", artifact.source_file_name); + with_string_field(&mut metadata, "version_id", Some(version_id.as_str())); + + if !metadata.contains_key("gate_key") { + with_string_field( + &mut metadata, + "gate_key", + infer_gate_key_from_social_stage(artifact.stage.as_str()), + ); + } + if !metadata.contains_key("run_title") { + with_string_field( + &mut metadata, + "run_title", + Some(artifact.version_label.as_str()), + ); + } + } + + serde_json::Value::Object(metadata) +} + /// Agent 执行策略 #[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "snake_case")] @@ -605,7 +1318,7 @@ async fn ensure_code_execution_extension_enabled(agent: &Agent) -> Result( agent: &Agent, app: &AppHandle, event_name: &str, @@ -614,7 +1327,11 @@ async fn stream_reply_once( session_config: aster::agents::SessionConfig, cancel_token: CancellationToken, request_tool_policy: &RequestToolPolicy, -) -> Result<(), ReplyAttemptError> { + mut on_event: F, +) -> Result<(), ReplyAttemptError> +where + F: FnMut(&TauriAgentEvent), +{ stream_reply_with_policy( agent, message_text, @@ -623,6 +1340,7 @@ async fn stream_reply_once( Some(cancel_token), request_tool_policy, |event| { + on_event(event); if let Err(error) = app.emit(event_name, event) { tracing::error!("[AsterAgent] 发送事件失败: {}", error); } @@ -3963,12 +4681,18 @@ pub async fn aster_agent_chat_stream( ); } - // 构建请求级工具策略:effective_web_search = request.web_search ?? mode_default(false) - let request_tool_policy = resolve_request_tool_policy(request.web_search, false); + let runtime_chat_mode = resolve_runtime_chat_mode(request.metadata.as_ref()); + let mode_default_web_search = default_web_search_enabled_for_chat_mode(runtime_chat_mode); + + // 构建请求级工具策略:默认不强制联网搜索,仅在用户显式开启开关时把搜索升级为必需步骤。 + let request_tool_policy = + resolve_request_tool_policy(request.web_search, mode_default_web_search); tracing::info!( - "[AsterAgent][WebSearchGuard] session={}, request_web_search={:?}, mode_default_web_search=false, effective_web_search={}", + "[AsterAgent][WebSearchGuard] session={}, chat_mode={:?}, request_web_search={:?}, mode_default_web_search={}, effective_web_search={}", session_id, + runtime_chat_mode, request.web_search, + mode_default_web_search, request_tool_policy.effective_web_search ); @@ -4182,6 +4906,31 @@ pub async fn aster_agent_chat_stream( let tracker = ExecutionTracker::new(db.inner().clone()); let cancel_token = state.create_cancel_token(session_id).await; let auto_continue_metadata = auto_continue_config.clone(); + let request_metadata = request.metadata.clone(); + let run_start_metadata = build_chat_run_metadata_base( + &request, + workspace_id.as_str(), + effective_strategy, + &request_tool_policy, + auto_continue_enabled, + auto_continue_metadata.as_ref(), + ); + let run_observation = Arc::new(Mutex::new(ChatRunObservation::default())); + let run_observation_for_finalize = run_observation.clone(); + let run_start_metadata_for_finalize = run_start_metadata.clone(); + let timeline_recorder = Arc::new(Mutex::new(AgentTimelineRecorder::create( + db.inner().clone(), + session_id.to_string(), + request.message.clone(), + )?)); + + { + let mut recorder = match timeline_recorder.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + recorder.emit_start(&app, &request.event_name)?; + } // 获取 Agent Arc 并保持 guard 在整个流处理期间存活 let agent_arc = state.get_agent_arc(); @@ -4201,35 +4950,11 @@ pub async fn aster_agent_chat_stream( }; let final_result = tracker - .with_run( + .with_run_custom( RunSource::Chat, Some("aster_agent_chat_stream".to_string()), Some(session_id.to_string()), - Some(serde_json::json!({ - "workspace_id": workspace_id.clone(), - "project_id": request.project_id.clone(), - "event_name": request.event_name.clone(), - "execution_strategy": format!("{:?}", effective_strategy).to_lowercase(), - "message_length": request.message.chars().count(), - "web_search_enabled": request_tool_policy.effective_web_search, - "auto_continue_enabled": auto_continue_enabled, - "auto_continue": auto_continue_metadata, - })), - RunFinalizeOptions { - success_metadata: Some(serde_json::json!({ - "execution_strategy": format!("{:?}", effective_strategy).to_lowercase(), - "workspace_id": workspace_id.clone(), - "web_search_enabled": request_tool_policy.effective_web_search, - "auto_continue_enabled": auto_continue_enabled, - })), - error_code: Some("chat_stream_failed".to_string()), - error_metadata: Some(serde_json::json!({ - "execution_strategy": format!("{:?}", effective_strategy).to_lowercase(), - "workspace_id": workspace_id.clone(), - "web_search_enabled": request_tool_policy.effective_web_search, - "auto_continue_enabled": auto_continue_enabled, - })), - }, + Some(serde_json::Value::Object(run_start_metadata.clone())), async { let mut added_code_execution = false; if effective_strategy == AsterExecutionStrategy::CodeOrchestrated { @@ -4245,6 +4970,45 @@ pub async fn aster_agent_chat_stream( build_session_config(), cancel_token.clone(), &request_tool_policy, + { + let run_observation = run_observation.clone(); + let app = app.clone(); + let event_name = request.event_name.clone(); + let timeline_recorder = timeline_recorder.clone(); + let workspace_root = workspace_root.clone(); + let request_metadata = request_metadata.clone(); + move |event| { + let mut observation = match run_observation.lock() { + Ok(guard) => guard, + Err(error) => { + tracing::warn!( + "[AsterAgent] run observation lock poisoned,继续复用内部状态" + ); + error.into_inner() + } + }; + observation.record_event( + event, + workspace_root.as_str(), + request_metadata.as_ref(), + ); + let mut recorder = match timeline_recorder.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + if let Err(error) = recorder.record_legacy_event( + &app, + &event_name, + event, + workspace_root.as_str(), + ) { + tracing::warn!( + "[AsterAgent] 记录时间线事件失败(已降级继续): {}", + error + ); + } + } + }, ) .await; @@ -4278,6 +5042,45 @@ pub async fn aster_agent_chat_stream( build_session_config(), cancel_token.clone(), &request_tool_policy, + { + let run_observation = run_observation.clone(); + let app = app.clone(); + let event_name = request.event_name.clone(); + let timeline_recorder = timeline_recorder.clone(); + let workspace_root = workspace_root.clone(); + let request_metadata = request_metadata.clone(); + move |event| { + let mut observation = match run_observation.lock() { + Ok(guard) => guard, + Err(error) => { + tracing::warn!( + "[AsterAgent] run observation lock poisoned,继续复用内部状态" + ); + error.into_inner() + } + }; + observation.record_event( + event, + workspace_root.as_str(), + request_metadata.as_ref(), + ); + let mut recorder = match timeline_recorder.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + if let Err(error) = recorder.record_legacy_event( + &app, + &event_name, + event, + workspace_root.as_str(), + ) { + tracing::warn!( + "[AsterAgent] 记录时间线事件失败(已降级继续): {}", + error + ); + } + } + }, ) .await .map_err(|fallback_err| fallback_err.message) @@ -4296,17 +5099,66 @@ pub async fn aster_agent_chat_stream( run_result }, + move |result| { + let observation = match run_observation_for_finalize.lock() { + Ok(guard) => guard.clone(), + Err(error) => { + tracing::warn!( + "[AsterAgent] finalize run metadata 时 observation lock 已 poisoned" + ); + error.into_inner().clone() + } + }; + let metadata = + build_chat_run_finish_metadata(&run_start_metadata_for_finalize, &observation); + + match result { + Ok(_) => RunFinishDecision { + status: proxycast_core::database::dao::agent_run::AgentRunStatus::Success, + error_code: None, + error_message: None, + metadata: Some(metadata), + }, + Err(err) => RunFinishDecision { + status: proxycast_core::database::dao::agent_run::AgentRunStatus::Error, + error_code: Some("chat_stream_failed".to_string()), + error_message: Some(err.clone()), + metadata: Some(metadata), + }, + } + }, ) .await; match final_result { Ok(()) => { + { + let mut recorder = match timeline_recorder.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + if let Err(error) = recorder.complete_turn_success(&app, &request.event_name) { + tracing::warn!("[AsterAgent] 完成 turn 时间线失败(已降级继续): {}", error); + } + } let done_event = TauriAgentEvent::FinalDone { usage: None }; if let Err(e) = app.emit(&request.event_name, &done_event) { tracing::error!("[AsterAgent] 发送完成事件失败: {}", e); } } Err(e) => { + { + let mut recorder = match timeline_recorder.lock() { + Ok(guard) => guard, + Err(error) => error.into_inner(), + }; + if let Err(timeline_error) = recorder.fail_turn(&app, &request.event_name, &e) { + tracing::warn!( + "[AsterAgent] 记录失败 turn 时间线失败(已降级继续): {}", + timeline_error + ); + } + } let error_event = TauriAgentEvent::Error { message: e.clone() }; if let Err(emit_err) = app.emit(&request.event_name, &error_event) { tracing::error!("[AsterAgent] 发送错误事件失败: {}", emit_err); @@ -4332,6 +5184,53 @@ pub async fn aster_agent_stop( Ok(state.cancel_session(&session_id).await) } +/// 统一运行时:提交一个 turn。 +#[tauri::command] +pub async fn agent_runtime_submit_turn( + app: AppHandle, + state: State<'_, AsterAgentState>, + db: State<'_, DbConnection>, + api_key_provider_service: State<'_, ApiKeyProviderServiceState>, + logs: State<'_, LogState>, + config_manager: State<'_, GlobalConfigManagerState>, + mcp_manager: State<'_, McpManagerState>, + heartbeat_state: State<'_, HeartbeatServiceState>, + request: AgentRuntimeSubmitTurnRequest, +) -> Result<(), String> { + aster_agent_chat_stream( + app, + state, + db, + api_key_provider_service, + logs, + config_manager, + mcp_manager, + heartbeat_state, + request.into(), + ) + .await +} + +/// 统一运行时:中断当前 turn。 +#[tauri::command] +pub async fn agent_runtime_interrupt_turn( + state: State<'_, AsterAgentState>, + request: AgentRuntimeInterruptTurnRequest, +) -> Result { + aster_agent_stop(state, request.session_id).await +} + +/// 创建新会话 +#[tauri::command] +pub async fn agent_runtime_create_session( + db: State<'_, DbConnection>, + workspace_id: String, + name: Option, + execution_strategy: Option, +) -> Result { + aster_session_create(db, None, workspace_id, name, execution_strategy).await +} + /// 创建新会话 #[tauri::command] pub async fn aster_session_create( @@ -4402,6 +5301,14 @@ pub async fn aster_session_set_execution_strategy( Ok(()) } +/// 统一运行时:列出会话。 +#[tauri::command] +pub async fn agent_runtime_list_sessions( + db: State<'_, DbConnection>, +) -> Result, String> { + aster_session_list(db).await +} + /// 列出所有会话 #[tauri::command] pub async fn aster_session_list(db: State<'_, DbConnection>) -> Result, String> { @@ -4419,6 +5326,15 @@ pub async fn aster_session_get( AsterAgentWrapper::get_session_sync(&db, &session_id) } +/// 统一运行时:获取会话详情。 +#[tauri::command] +pub async fn agent_runtime_get_session( + db: State<'_, DbConnection>, + session_id: String, +) -> Result { + aster_session_get(db, session_id).await +} + /// 重命名会话 #[tauri::command] pub async fn aster_session_rename( @@ -4430,6 +5346,36 @@ pub async fn aster_session_rename( AsterAgentWrapper::rename_session_sync(&db, &session_id, &name) } +/// 统一运行时:更新会话元数据。 +#[tauri::command] +pub async fn agent_runtime_update_session( + db: State<'_, DbConnection>, + request: AgentRuntimeUpdateSessionRequest, +) -> Result<(), String> { + let trimmed_session_id = request.session_id.trim().to_string(); + if trimmed_session_id.is_empty() { + return Err("session_id 不能为空".to_string()); + } + + if let Some(name) = request.name.as_ref() { + let normalized_name = name.trim(); + if !normalized_name.is_empty() { + aster_session_rename( + db.clone(), + trimmed_session_id.clone(), + normalized_name.to_string(), + ) + .await?; + } + } + + if let Some(execution_strategy) = request.execution_strategy { + aster_session_set_execution_strategy(db, trimmed_session_id, execution_strategy).await?; + } + + Ok(()) +} + /// 删除会话 #[tauri::command] pub async fn aster_session_delete( @@ -4440,6 +5386,15 @@ pub async fn aster_session_delete( AsterAgentWrapper::delete_session_sync(&db, &session_id) } +/// 统一运行时:删除会话。 +#[tauri::command] +pub async fn agent_runtime_delete_session( + db: State<'_, DbConnection>, + session_id: String, +) -> Result<(), String> { + aster_session_delete(db, session_id).await +} + /// 确认权限请求 #[derive(Debug, Deserialize)] pub struct ConfirmRequest { @@ -4500,6 +5455,72 @@ fn validate_elicitation_submission(session_id: &str, request_id: &str) -> Result Ok(trimmed_session_id) } +fn build_runtime_action_user_data(request: &AgentRuntimeRespondActionRequest) -> serde_json::Value { + if let Some(user_data) = request.user_data.clone() { + return user_data; + } + + if !request.confirmed { + return serde_json::Value::String(String::new()); + } + + let Some(response) = request.response.as_ref() else { + return serde_json::Value::String(String::new()); + }; + let trimmed = response.trim(); + if trimmed.is_empty() { + return serde_json::Value::String(String::new()); + } + + serde_json::from_str(trimmed).unwrap_or_else(|_| serde_json::Value::String(trimmed.to_string())) +} + +/// 统一运行时:响应工具确认 / ask / elicitation。 +#[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 { + AgentRuntimeActionType::ToolConfirmation => { + aster_agent_confirm( + state, + ConfirmRequest { + request_id: request.request_id.clone(), + confirmed: request.confirmed, + response: request.response.clone(), + }, + ) + .await + } + AgentRuntimeActionType::AskUser | AgentRuntimeActionType::Elicitation => { + let user_data = build_runtime_action_user_data(&request); + aster_agent_submit_elicitation_response( + state, + request.session_id.clone(), + SubmitElicitationResponseRequest { + request_id: request.request_id.clone(), + user_data, + }, + ) + .await + } + }; + + if result.is_ok() { + complete_action_item(db.inner(), &request.request_id, response_value)?; + } + + result +} + /// 提交 elicitation 回答(用于 ask/lsp 等需要用户输入的流程) #[tauri::command] pub async fn aster_agent_submit_elicitation_response( @@ -4736,6 +5757,257 @@ mod tests { ); } + #[test] + fn test_aster_chat_request_deserialize_with_metadata() { + let json = r#"{ + "message": "Hello", + "session_id": "test-session", + "event_name": "agent_stream", + "workspace_id": "workspace-test", + "metadata": { + "harness": { + "theme": "social-media", + "gate_key": "write_mode", + "run_title": "社媒初稿" + } + } + }"#; + + let request: AsterChatRequest = serde_json::from_str(json).unwrap(); + assert_eq!( + request + .metadata + .as_ref() + .and_then(|value| value.get("harness")) + .and_then(|value| value.get("theme")) + .and_then(serde_json::Value::as_str), + Some("social-media") + ); + } + + #[test] + fn test_resolve_runtime_chat_mode_prefers_explicit_chat_mode() { + let metadata = serde_json::json!({ + "harness": { + "theme": "social-media", + "chat_mode": "general" + } + }); + + assert_eq!( + resolve_runtime_chat_mode(Some(&metadata)), + RuntimeChatMode::General + ); + } + + #[test] + fn test_resolve_runtime_chat_mode_falls_back_to_general_theme_group() { + let metadata = serde_json::json!({ + "harness": { + "theme": "planning" + } + }); + + assert_eq!( + resolve_runtime_chat_mode(Some(&metadata)), + RuntimeChatMode::General + ); + } + + #[test] + fn test_default_web_search_enabled_for_chat_mode_requires_explicit_opt_in() { + assert!(!default_web_search_enabled_for_chat_mode( + RuntimeChatMode::Agent + )); + assert!(!default_web_search_enabled_for_chat_mode( + RuntimeChatMode::Creator + )); + assert!(!default_web_search_enabled_for_chat_mode( + RuntimeChatMode::General + )); + } + + #[test] + fn test_agent_runtime_submit_turn_request_maps_to_aster_chat_request() { + let json = r#"{ + "message": "Hello runtime", + "session_id": "runtime-session", + "event_name": "runtime_stream", + "workspace_id": "workspace-runtime", + "turn_config": { + "execution_strategy": "auto", + "web_search": true, + "system_prompt": "runtime prompt", + "provider_config": { + "provider_id": "custom-provider", + "provider_name": "custom-provider", + "model_name": "gpt-5.3-codex" + }, + "metadata": { + "source": "hook-facade" + } + } + }"#; + + let request: AgentRuntimeSubmitTurnRequest = serde_json::from_str(json).unwrap(); + let mapped: AsterChatRequest = request.into(); + + assert_eq!(mapped.message, "Hello runtime"); + assert_eq!(mapped.session_id, "runtime-session"); + assert_eq!(mapped.event_name, "runtime_stream"); + assert_eq!(mapped.workspace_id, "workspace-runtime"); + assert_eq!( + mapped.execution_strategy, + Some(AsterExecutionStrategy::Auto) + ); + assert_eq!(mapped.web_search, Some(true)); + assert_eq!(mapped.system_prompt.as_deref(), Some("runtime prompt")); + assert_eq!( + mapped + .provider_config + .as_ref() + .and_then(|config| config.provider_id.as_deref()), + Some("custom-provider") + ); + assert_eq!( + mapped + .metadata + .as_ref() + .and_then(|value| value.get("source")) + .and_then(serde_json::Value::as_str), + Some("hook-facade") + ); + } + + #[test] + fn test_build_runtime_action_user_data_prefers_structured_payload() { + let request = AgentRuntimeRespondActionRequest { + session_id: "session-1".to_string(), + request_id: "req-1".to_string(), + action_type: AgentRuntimeActionType::AskUser, + confirmed: true, + response: Some("{\"answer\":\"A\"}".to_string()), + user_data: Some(serde_json::json!({ "answer": "B" })), + }; + + assert_eq!( + build_runtime_action_user_data(&request), + serde_json::json!({ "answer": "B" }) + ); + } + + #[test] + fn test_build_runtime_action_user_data_parses_json_response() { + let request = AgentRuntimeRespondActionRequest { + session_id: "session-1".to_string(), + request_id: "req-1".to_string(), + action_type: AgentRuntimeActionType::Elicitation, + confirmed: true, + response: Some("{\"answer\":\"A\"}".to_string()), + user_data: None, + }; + + assert_eq!( + build_runtime_action_user_data(&request), + serde_json::json!({ "answer": "A" }) + ); + } + + #[test] + fn test_extract_artifact_path_from_tool_start_reads_write_file_path() { + let path = extract_artifact_path_from_tool_start( + "write_file", + Some(r##"{"path":"social-posts/demo.md","content":"# 标题"}"##), + "/tmp/workspace", + ); + + assert_eq!(path.as_deref(), Some("social-posts/demo.md")); + } + + #[test] + fn test_resolve_social_run_artifact_descriptor_matches_social_draft() { + let descriptor = resolve_social_run_artifact_descriptor( + "social-posts/draft.md", + Some("write_mode"), + Some("社媒初稿"), + ); + + assert_eq!(descriptor.artifact_type, "draft"); + assert_eq!(descriptor.stage, "drafting"); + assert_eq!(descriptor.version_label, "社媒初稿"); + assert!(!descriptor.is_auxiliary); + } + + #[test] + fn test_build_chat_run_finish_metadata_includes_social_fields() { + let base = build_chat_run_metadata_base( + &AsterChatRequest { + message: "hello".to_string(), + session_id: "session-1".to_string(), + event_name: "event-1".to_string(), + images: None, + provider_config: None, + project_id: Some("project-1".to_string()), + workspace_id: "workspace-1".to_string(), + web_search: Some(false), + execution_strategy: Some(AsterExecutionStrategy::React), + auto_continue: None, + system_prompt: None, + metadata: Some(serde_json::json!({ + "harness": { + "theme": "social-media", + "gate_key": "write_mode" + } + })), + }, + "workspace-1", + AsterExecutionStrategy::React, + &RequestToolPolicy { + effective_web_search: false, + required_tools: vec![], + allowed_tools: vec![], + disallowed_tools: vec![], + }, + false, + None, + ); + let mut observation = ChatRunObservation::default(); + observation.record_artifact_path( + "social-posts/draft.md".to_string(), + Some(&serde_json::json!({ + "harness": { + "theme": "social-media", + "gate_key": "write_mode" + } + })), + ); + + let metadata = build_chat_run_finish_metadata(&base, &observation); + + assert_eq!( + metadata + .get("artifact_paths") + .and_then(serde_json::Value::as_array), + Some(&vec![serde_json::json!("social-posts/draft.md")]) + ); + assert_eq!( + metadata + .get("artifact_type") + .and_then(serde_json::Value::as_str), + Some("draft") + ); + assert_eq!( + metadata.get("stage").and_then(serde_json::Value::as_str), + Some("drafting") + ); + assert_eq!( + metadata + .get("version_id") + .and_then(serde_json::Value::as_str), + Some("artifact:social-posts/draft.md") + ); + } + #[test] fn test_aster_execution_strategy_default_is_auto() { assert_eq!( diff --git a/src-tauri/src/commands/execution_run_cmd.rs b/src-tauri/src/commands/execution_run_cmd.rs index 4418f16ac..c19870cf9 100644 --- a/src-tauri/src/commands/execution_run_cmd.rs +++ b/src-tauri/src/commands/execution_run_cmd.rs @@ -123,9 +123,18 @@ pub struct ThemeWorkbenchRunState { pub current_gate_key: String, pub queue_items: Vec, pub latest_terminal: Option, + pub recent_terminals: Vec, pub updated_at: String, } +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "snake_case")] +pub struct ThemeWorkbenchRunHistoryPage { + pub items: Vec, + pub has_more: bool, + pub next_offset: Option, +} + fn normalize_gate_key(raw: &str) -> Option { let normalized = raw.trim().to_lowercase(); match normalized.as_str() { @@ -160,12 +169,42 @@ fn infer_gate_key_from_probe(probe: &str) -> String { "write_mode".to_string() } +fn metadata_string<'a>(value: &'a Value, keys: &[&str]) -> Option<&'a str> { + keys.iter() + .filter_map(|key| value.get(*key)) + .find_map(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) +} + +fn infer_gate_key_from_stage(stage: &str) -> Option { + match stage.trim().to_lowercase().as_str() { + "briefing" => Some("topic_select".to_string()), + "drafting" | "polishing" => Some("write_mode".to_string()), + "adapting" | "publish_prep" => Some("publish_confirm".to_string()), + _ => None, + } +} + fn derive_run_title(run: &AgentRun) -> String { let parsed_metadata = run .metadata .as_ref() .and_then(|raw| serde_json::from_str::(raw).ok()); + let metadata_title = parsed_metadata + .as_ref() + .and_then(|value| { + metadata_string( + value, + &["run_title", "title", "version_label", "stage_label"], + ) + }) + .map(str::to_string); + if let Some(title) = metadata_title { + return title; + } + let skill_title = parsed_metadata .as_ref() .and_then(|value| value.get("skill_name")) @@ -220,6 +259,14 @@ fn derive_run_gate_key(run: &AgentRun, title: &str) -> String { return value; } + if let Some(value) = parsed_metadata + .as_ref() + .and_then(|value| metadata_string(value, &["stage"])) + .and_then(infer_gate_key_from_stage) + { + return value; + } + let metadata_probe = parsed_metadata .as_ref() .map(|value| value.to_string()) @@ -250,13 +297,16 @@ fn derive_run_execution_id(run: &AgentRun) -> Option { parsed_metadata .as_ref() .and_then(|value| { - value - .get("execution_id") - .or_else(|| value.get("version_id")) - .and_then(Value::as_str) + metadata_string( + value, + &[ + "execution_id", + "version_id", + "run_version_id", + "artifact_version_id", + ], + ) }) - .map(str::trim) - .filter(|value| !value.is_empty()) .map(str::to_string) } @@ -268,20 +318,72 @@ fn derive_run_artifact_paths(run: &AgentRun) -> Vec { parsed_metadata .as_ref() - .and_then(|value| value.get("artifact_paths")) - .and_then(Value::as_array) - .map(|paths| { + .map(|value| { + let mut paths = value + .get("artifact_paths") + .and_then(Value::as_array) + .map(|items| { + items + .iter() + .filter_map(Value::as_str) + .map(str::trim) + .filter(|path| !path.is_empty()) + .map(str::to_string) + .collect::>() + }) + .unwrap_or_default(); + + for key in ["artifact_path", "source_file_name"] { + if let Some(path) = metadata_string(value, &[key]) { + let normalized = path.to_string(); + if !paths.iter().any(|existing| existing == &normalized) { + paths.push(normalized); + } + } + } + paths - .iter() - .filter_map(Value::as_str) - .map(str::trim) - .filter(|path| !path.is_empty()) - .map(str::to_string) - .collect() }) .unwrap_or_default() } +fn build_terminal_item(run: &AgentRun) -> ThemeWorkbenchRunTerminalItem { + let title = derive_run_title(run); + let gate_key = derive_run_gate_key(run, title.as_str()); + ThemeWorkbenchRunTerminalItem { + run_id: run.id.clone(), + execution_id: derive_run_execution_id(run), + session_id: run.session_id.clone(), + artifact_paths: derive_run_artifact_paths(run), + title, + gate_key, + status: run.status.clone(), + source: run.source.clone(), + source_ref: run.source_ref.clone(), + started_at: run.started_at.clone(), + finished_at: run.finished_at.clone(), + } +} + +fn derive_recent_terminal_items( + runs: &[AgentRun], + limit: usize, +) -> Vec { + runs.iter() + .filter(|run| { + matches!( + run.status, + AgentRunStatus::Success + | AgentRunStatus::Error + | AgentRunStatus::Canceled + | AgentRunStatus::Timeout + ) + }) + .take(limit) + .map(build_terminal_item) + .collect() +} + #[tauri::command] pub async fn execution_run_get_theme_workbench_state( db: State<'_, DbConnection>, @@ -332,44 +434,62 @@ pub async fn execution_run_get_theme_workbench_state( }; let current_gate_key = derive_current_gate_key(queue_items.as_slice()); - let latest_terminal = runs - .iter() - .find(|run| { - matches!( - run.status, - AgentRunStatus::Success - | AgentRunStatus::Error - | AgentRunStatus::Canceled - | AgentRunStatus::Timeout - ) - }) - .map(|run| { - let title = derive_run_title(run); - let gate_key = derive_run_gate_key(run, title.as_str()); - ThemeWorkbenchRunTerminalItem { - run_id: run.id.clone(), - execution_id: derive_run_execution_id(run), - session_id: run.session_id.clone(), - artifact_paths: derive_run_artifact_paths(run), - title, - gate_key, - status: run.status.clone(), - source: run.source.clone(), - source_ref: run.source_ref.clone(), - started_at: run.started_at.clone(), - finished_at: run.finished_at.clone(), - } - }); + let recent_terminals = derive_recent_terminal_items(runs.as_slice(), safe_limit); + let latest_terminal = recent_terminals.first().cloned(); Ok(ThemeWorkbenchRunState { run_state, current_gate_key, queue_items, latest_terminal, + recent_terminals, updated_at: Utc::now().to_rfc3339(), }) } +#[tauri::command] +pub async fn execution_run_list_theme_workbench_history( + db: State<'_, DbConnection>, + session_id: String, + limit: Option, + offset: Option, +) -> Result { + let trimmed_session_id = session_id.trim(); + if trimmed_session_id.is_empty() { + return Err("session_id 不能为空".to_string()); + } + + let safe_limit = limit.unwrap_or(20).clamp(1, 100); + let safe_offset = offset.unwrap_or(0); + let tracker = ExecutionTracker::new(db.inner().clone()); + + let now = Utc::now(); + let runs_for_timeout = tracker.list_runs_by_session(trimmed_session_id, safe_limit * 5)?; + let stale_run_ids = collect_stale_run_ids(runs_for_timeout.as_slice(), now); + if !stale_run_ids.is_empty() { + mark_stale_runs_as_timeout(db.inner(), stale_run_ids.as_slice(), &now.to_rfc3339())?; + } + + let paged_runs = + tracker.list_terminal_runs_by_session(trimmed_session_id, safe_limit + 1, safe_offset)?; + let has_more = paged_runs.len() > safe_limit; + let items = paged_runs + .into_iter() + .take(safe_limit) + .map(|run| build_terminal_item(&run)) + .collect::>(); + + Ok(ThemeWorkbenchRunHistoryPage { + items, + has_more, + next_offset: if has_more { + Some(safe_offset + safe_limit) + } else { + None + }, + }) +} + #[cfg(test)] mod tests { use super::*; @@ -509,6 +629,72 @@ mod tests { ); } + #[test] + fn derive_run_gate_key_should_support_social_stage_fallback() { + let run = sample_run_with_metadata(Some(serde_json::json!({ + "harness_theme": "social-media", + "stage": "publish_prep" + }))); + + assert_eq!(derive_run_gate_key(&run, "社媒发布包"), "publish_confirm"); + } + + #[test] + fn derive_run_title_should_prefer_social_version_label() { + let run = sample_run_with_metadata(Some(serde_json::json!({ + "harness_theme": "social-media", + "version_label": "社媒初稿" + }))); + + assert_eq!(derive_run_title(&run), "社媒初稿"); + } + + #[test] + fn derive_run_artifact_paths_should_fallback_to_source_file_name() { + let run = sample_run_with_metadata(Some(serde_json::json!({ + "source_file_name": "social-posts/demo.md" + }))); + + assert_eq!( + derive_run_artifact_paths(&run), + vec!["social-posts/demo.md".to_string()] + ); + } + + #[test] + fn derive_recent_terminal_items_should_keep_multiple_terminal_runs() { + let mut latest_error_run = sample_run_with_metadata(Some(serde_json::json!({ + "run_title": "最新失败运行" + }))); + latest_error_run.id = "run-error".to_string(); + latest_error_run.status = AgentRunStatus::Error; + latest_error_run.started_at = "2026-03-06T05:00:00Z".to_string(); + latest_error_run.finished_at = Some("2026-03-06T05:02:00Z".to_string()); + + let mut running_run = sample_run_with_metadata(Some(serde_json::json!({ + "run_title": "运行中" + }))); + running_run.id = "run-running".to_string(); + running_run.status = AgentRunStatus::Running; + running_run.started_at = "2026-03-06T04:30:00Z".to_string(); + running_run.finished_at = None; + + let mut previous_success_run = sample_run_with_metadata(Some(serde_json::json!({ + "run_title": "上一轮成功运行" + }))); + previous_success_run.id = "run-success".to_string(); + previous_success_run.status = AgentRunStatus::Success; + previous_success_run.started_at = "2026-03-06T04:00:00Z".to_string(); + previous_success_run.finished_at = Some("2026-03-06T04:08:00Z".to_string()); + + let recent_terminals = + derive_recent_terminal_items(&[latest_error_run, running_run, previous_success_run], 3); + + assert_eq!(recent_terminals.len(), 2); + assert_eq!(recent_terminals[0].run_id, "run-error"); + assert_eq!(recent_terminals[1].run_id, "run-success"); + } + #[test] fn collect_stale_run_ids_should_only_pick_expired_non_terminal_runs() { let now = Utc::now(); diff --git a/src-tauri/src/commands/general_chat_cmd.rs b/src-tauri/src/commands/general_chat_cmd.rs deleted file mode 100644 index 5c446b4ae..000000000 --- a/src-tauri/src/commands/general_chat_cmd.rs +++ /dev/null @@ -1,1095 +0,0 @@ -//! 通用对话 Tauri 命令兼容层 -//! -//! 该模块仅用于兼容旧版 `general-chat` 前端链路。 -//! 新功能和后续治理请统一收口到 `unified_chat_cmd`。 -//! -//! ## 主要命令 -//! - `general_chat_create_session` - 创建新会话 -//! - `general_chat_list_sessions` - 获取会话列表 -//! - `general_chat_get_session` - 获取会话详情 -//! - `general_chat_delete_session` - 删除会话 -//! - `general_chat_rename_session` - 重命名会话 -//! - `general_chat_add_message` - 已废弃直接写消息(显式报错) -//! - `general_chat_send_message` - 已废弃流式发送(显式报错) -//! - `general_chat_stop_generation` - 已废弃停止生成(显式报错) -//! - `general_chat_generate_title` - 已废弃标题生成(显式报错) -//! - `general_chat_get_messages` - 获取消息列表 - -use crate::database::dao::chat::{ - ChatDao, ChatMessage as UnifiedChatMessage, ChatMode, ChatSession as UnifiedChatSession, -}; -use crate::database::dao::general_chat::GeneralChatDao; -use crate::database::DbConnection; -use once_cell::sync::Lazy; -use proxycast_services::general_chat::{ - ChatMessage, ChatSession, ContentBlock, MessageRole, SessionDetail, -}; -use serde::Deserialize; -use std::collections::{HashMap, HashSet}; -use std::sync::Mutex; -use tauri::State; -use uuid::Uuid; - -static LEGACY_WARNED_COMMANDS: Lazy>> = - Lazy::new(|| Mutex::new(HashSet::new())); - -fn warn_general_chat_legacy(command: &'static str, replacement: &'static str) { - let should_warn = match LEGACY_WARNED_COMMANDS.lock() { - Ok(mut warned) => warned.insert(command), - Err(error) => { - tracing::warn!( - "[GeneralChat][Compat] 废弃命令告警状态异常: {}。命令 {} 仍通过兼容层提供,建议迁移到 {}", - error, - command, - replacement - ); - true - } - }; - - if should_warn { - tracing::warn!( - "[GeneralChat][Compat] 命令 {} 仍通过兼容层提供,建议迁移到 {}。该入口仅用于兼容旧 UI,禁止继续叠加新逻辑。", - command, - replacement - ); - } -} - -fn general_chat_deprecated_error(command: &'static str, replacement: &'static str) -> String { - format!( - "命令 {command} 已废弃,请迁移到 {replacement}。该兼容入口已停止维护,禁止继续叠加新逻辑。" - ) -} - -fn timestamp_ms_to_rfc3339(timestamp_ms: i64) -> String { - use chrono::{TimeZone, Utc}; - - let secs = timestamp_ms / 1000; - let nsecs = ((timestamp_ms % 1000) * 1_000_000) as u32; - - match Utc.timestamp_opt(secs, nsecs) { - chrono::LocalResult::Single(dt) => dt.to_rfc3339(), - _ => Utc::now().to_rfc3339(), - } -} - -fn rfc3339_to_timestamp_ms(timestamp: &str) -> i64 { - chrono::DateTime::parse_from_rfc3339(timestamp) - .map(|value| value.timestamp_millis()) - .unwrap_or_else(|_| chrono::Utc::now().timestamp_millis()) -} - -fn general_message_role_name(role: &MessageRole) -> &'static str { - match role { - MessageRole::User => "user", - MessageRole::Assistant => "assistant", - MessageRole::System => "system", - } -} - -fn overlay_general_session( - legacy_session: Option, - unified_session: Option<&UnifiedChatSession>, -) -> Option { - match (legacy_session, unified_session) { - (None, None) => None, - (Some(mut session), Some(unified)) if unified.mode == ChatMode::General => { - if let Some(title) = unified - .title - .as_deref() - .map(str::trim) - .filter(|title| !title.is_empty()) - { - session.name = title.to_string(); - } - session.created_at = session - .created_at - .min(rfc3339_to_timestamp_ms(&unified.created_at)); - session.updated_at = session - .updated_at - .max(rfc3339_to_timestamp_ms(&unified.updated_at)); - Some(session) - } - (Some(session), _) => Some(session), - (None, Some(unified)) if unified.mode == ChatMode::General => Some(ChatSession { - id: unified.id.clone(), - name: unified - .title - .as_deref() - .map(str::trim) - .filter(|title| !title.is_empty()) - .unwrap_or("新对话") - .to_string(), - created_at: rfc3339_to_timestamp_ms(&unified.created_at), - updated_at: rfc3339_to_timestamp_ms(&unified.updated_at), - metadata: unified.metadata.clone(), - }), - _ => None, - } -} - -fn json_value_as_str(value: Option<&serde_json::Value>) -> Option { - value - .and_then(|item| item.as_str()) - .map(ToString::to_string) - .filter(|item| !item.trim().is_empty()) -} - -fn convert_unified_content_part_to_general_block( - object: &serde_json::Map, -) -> Option { - if let Some(text) = object.get("Text").and_then(|value| value.as_str()) { - return Some(ContentBlock { - r#type: "text".to_string(), - content: text.to_string(), - language: None, - filename: None, - mime_type: None, - }); - } - - if let Some(text_obj) = object.get("Text").and_then(|value| value.as_object()) { - if let Some(text) = json_value_as_str(text_obj.get("text")) - .or_else(|| json_value_as_str(text_obj.get("content"))) - { - return Some(ContentBlock { - r#type: "text".to_string(), - content: text, - language: None, - filename: None, - mime_type: None, - }); - } - } - - if let Some(text) = json_value_as_str(object.get("text")) { - return Some(ContentBlock { - r#type: "text".to_string(), - content: text, - language: None, - filename: None, - mime_type: None, - }); - } - - let part_type = object.get("type").and_then(|value| value.as_str()); - - if matches!( - part_type, - Some("text" | "input_text" | "output_text" | "thinking") - ) { - if let Some(text) = json_value_as_str(object.get("content")) - .or_else(|| json_value_as_str(object.get("text"))) - { - return Some(ContentBlock { - r#type: "text".to_string(), - content: text, - language: None, - filename: None, - mime_type: None, - }); - } - } - - if part_type == Some("code") { - if let Some(code) = json_value_as_str(object.get("content")) - .or_else(|| json_value_as_str(object.get("text"))) - { - return Some(ContentBlock { - r#type: "code".to_string(), - content: code, - language: json_value_as_str(object.get("language")), - filename: json_value_as_str(object.get("filename")), - mime_type: None, - }); - } - } - - if part_type == Some("file") { - if let Some(path) = json_value_as_str(object.get("path")) - .or_else(|| json_value_as_str(object.get("file_path"))) - .or_else(|| json_value_as_str(object.get("filePath"))) - .or_else(|| json_value_as_str(object.get("content"))) - { - return Some(ContentBlock { - r#type: "file".to_string(), - content: path, - language: None, - filename: json_value_as_str(object.get("name")) - .or_else(|| json_value_as_str(object.get("filename"))), - mime_type: json_value_as_str(object.get("mime_type")) - .or_else(|| json_value_as_str(object.get("media_type"))), - }); - } - } - - if matches!(part_type, Some("image_url" | "input_image")) { - let image_url = object.get("image_url").or_else(|| object.get("url")); - let url = image_url - .and_then(|value| value.as_str().map(ToString::to_string)) - .or_else(|| { - image_url - .and_then(|value| value.as_object()) - .and_then(|value| json_value_as_str(value.get("url"))) - }); - - if let Some(url) = url { - return Some(ContentBlock { - r#type: "image".to_string(), - content: url, - language: None, - filename: None, - mime_type: None, - }); - } - } - - if part_type == Some("image") { - if let Some(url) = json_value_as_str(object.get("url")) - .or_else(|| json_value_as_str(object.get("image_url"))) - { - return Some(ContentBlock { - r#type: "image".to_string(), - content: url, - language: None, - filename: None, - mime_type: None, - }); - } - - let source = object.get("source").and_then(|value| value.as_object()); - let mime_type = json_value_as_str(object.get("mime_type")) - .or_else(|| json_value_as_str(object.get("media_type"))) - .or_else(|| { - source - .and_then(|value| json_value_as_str(value.get("mime_type"))) - .or_else(|| source.and_then(|value| json_value_as_str(value.get("media_type")))) - }); - let data = json_value_as_str(object.get("data")) - .or_else(|| json_value_as_str(object.get("image_base64"))) - .or_else(|| source.and_then(|value| json_value_as_str(value.get("data")))); - - if let (Some(mime_type), Some(data)) = (mime_type, data) { - return Some(ContentBlock { - r#type: "image".to_string(), - content: format!("data:{mime_type};base64,{data}"), - language: None, - filename: None, - mime_type: Some(mime_type), - }); - } - } - - if let Some(image_url_obj) = object.get("image_url").and_then(|value| value.as_object()) { - if let Some(url) = json_value_as_str(image_url_obj.get("url")) { - return Some(ContentBlock { - r#type: "image".to_string(), - content: url, - language: None, - filename: None, - mime_type: None, - }); - } - } - - if let Some(url) = json_value_as_str(object.get("image_url")) { - return Some(ContentBlock { - r#type: "image".to_string(), - content: url, - language: None, - filename: None, - mime_type: None, - }); - } - - if let Some(text) = json_value_as_str(object.get("content")) { - return Some(ContentBlock { - r#type: "text".to_string(), - content: text, - language: None, - filename: None, - mime_type: None, - }); - } - - None -} - -fn convert_unified_content_to_general_parts( - content: &serde_json::Value, -) -> (String, Option>) { - match content { - serde_json::Value::String(text) => (text.clone(), None), - serde_json::Value::Array(items) => { - let blocks: Vec = items - .iter() - .filter_map(|item| item.as_object()) - .filter_map(convert_unified_content_part_to_general_block) - .collect(); - - if blocks.is_empty() { - return (serde_json::to_string(content).unwrap_or_default(), None); - } - - let text_content = blocks - .iter() - .filter(|block| matches!(block.r#type.as_str(), "text" | "code" | "file")) - .map(|block| block.content.clone()) - .collect::>() - .join("\n") - .trim() - .to_string(); - - let content = if !text_content.is_empty() { - text_content - } else if blocks.iter().any(|block| block.r#type == "image") { - "[图片]".to_string() - } else { - serde_json::to_string(content).unwrap_or_default() - }; - - (content, Some(blocks)) - } - serde_json::Value::Object(object) => { - let block = convert_unified_content_part_to_general_block(object); - if let Some(block) = block { - let content = if matches!(block.r#type.as_str(), "text" | "code" | "file") { - block.content.clone() - } else if block.r#type == "image" { - "[图片]".to_string() - } else { - serde_json::to_string(content).unwrap_or_default() - }; - (content, Some(vec![block])) - } else { - (serde_json::to_string(content).unwrap_or_default(), None) - } - } - _ => (serde_json::to_string(content).unwrap_or_default(), None), - } -} - -fn convert_unified_message_to_general(message: UnifiedChatMessage) -> Option { - let role = match message.role.as_str() { - "user" => MessageRole::User, - "assistant" => MessageRole::Assistant, - "system" => MessageRole::System, - _ => return None, - }; - - let (content, blocks) = convert_unified_content_to_general_parts(&message.content); - if content.trim().is_empty() && blocks.as_ref().is_none_or(Vec::is_empty) { - return None; - } - - Some(ChatMessage { - id: message.id.to_string(), - session_id: message.session_id, - role, - content, - blocks, - status: "complete".to_string(), - created_at: rfc3339_to_timestamp_ms(&message.created_at), - metadata: message.metadata, - }) -} - -fn general_message_identity_key(message: &ChatMessage) -> String { - let blocks_signature = - serde_json::to_string(&message.blocks).unwrap_or_else(|_| "[]".to_string()); - format!( - "{}|{}|{}|{}", - general_message_role_name(&message.role), - message.created_at, - message.content, - blocks_signature - ) -} - -fn merge_general_message_sources( - legacy_messages: Vec, - unified_messages: Vec, -) -> Vec { - let mut seen = HashSet::new(); - let mut merged = Vec::new(); - - for message in legacy_messages - .into_iter() - .chain(unified_messages.into_iter()) - { - if seen.insert(general_message_identity_key(&message)) { - merged.push(message); - } - } - - merged.sort_by(|left, right| { - left.created_at - .cmp(&right.created_at) - .then_with(|| left.id.cmp(&right.id)) - }); - - merged -} - -fn paginate_general_messages( - messages: Vec, - limit: Option, - before_id: Option<&str>, -) -> Vec { - let filtered = if let Some(before_id) = before_id { - if let Some(index) = messages.iter().position(|message| message.id == before_id) { - messages.into_iter().take(index).collect::>() - } else { - messages - } - } else { - messages - }; - - let Some(limit) = limit else { - return filtered; - }; - - let limit = limit.max(0) as usize; - if limit == 0 || filtered.len() <= limit { - return filtered; - } - - filtered[filtered.len() - limit..].to_vec() -} - -fn load_merged_general_messages( - conn: &rusqlite::Connection, - session_id: &str, -) -> Result, String> { - let legacy_messages = if GeneralChatDao::session_exists(conn, session_id) - .map_err(|e| format!("检查 general_chat 会话失败: {e}"))? - { - GeneralChatDao::get_messages(conn, session_id, None, None) - .map_err(|e| format!("读取 general_chat 消息失败: {e}"))? - } else { - Vec::new() - }; - - let unified_messages = match ChatDao::get_session(conn, session_id) - .map_err(|e| format!("读取 unified 会话失败: {e}"))? - { - Some(session) if session.mode == ChatMode::General => { - ChatDao::get_messages(conn, session_id, None) - .map_err(|e| format!("读取 unified 消息失败: {e}"))? - .into_iter() - .filter_map(convert_unified_message_to_general) - .collect() - } - _ => Vec::new(), - }; - - Ok(merge_general_message_sources( - legacy_messages, - unified_messages, - )) -} - -fn ensure_general_session_shadow( - conn: &rusqlite::Connection, - session: &ChatSession, -) -> Result<(), String> { - if ChatDao::session_exists(conn, &session.id) - .map_err(|e| format!("检查 unified 会话失败: {e}"))? - { - ChatDao::update_title(conn, &session.id, &session.name) - .map_err(|e| format!("更新 unified 会话标题失败: {e}"))?; - return Ok(()); - } - - let unified_session = UnifiedChatSession { - id: session.id.clone(), - mode: ChatMode::General, - title: Some(session.name.clone()), - system_prompt: None, - model: None, - provider_type: None, - credential_uuid: None, - metadata: session.metadata.clone(), - created_at: timestamp_ms_to_rfc3339(session.created_at), - updated_at: timestamp_ms_to_rfc3339(session.updated_at), - }; - - ChatDao::create_session(conn, &unified_session) - .map_err(|e| format!("创建 unified 会话影子失败: {e}")) -} - -fn ensure_general_session_shadow_by_id( - conn: &rusqlite::Connection, - session_id: &str, -) -> Result<(), String> { - let session = GeneralChatDao::get_session(conn, session_id) - .map_err(|e| format!("读取 general_chat 会话失败: {e}"))? - .ok_or_else(|| format!("general_chat 会话不存在: {session_id}"))?; - - ensure_general_session_shadow(conn, &session) -} - -fn convert_general_blocks_to_unified_content( - content: &str, - blocks: Option<&[ContentBlock]>, -) -> serde_json::Value { - if let Some(blocks) = blocks { - let converted: Vec = blocks - .iter() - .map(|block| match block.r#type.as_str() { - "text" => serde_json::json!({ - "type": "text", - "text": block.content, - }), - "image" => serde_json::json!({ - "type": "image", - "url": block.content, - "alt": block.filename, - }), - "file" => serde_json::json!({ - "type": "file", - "path": block.content, - "name": block.filename.clone().unwrap_or_default(), - }), - _ => serde_json::json!({ - "type": "text", - "text": block.content, - }), - }) - .collect(); - - if !converted.is_empty() { - return serde_json::Value::Array(converted); - } - } - - serde_json::json!([{ "type": "text", "text": content }]) -} - -fn convert_general_message_to_unified(message: &ChatMessage) -> UnifiedChatMessage { - let role = match message.role { - MessageRole::User => "user", - MessageRole::Assistant => "assistant", - MessageRole::System => "system", - }; - - UnifiedChatMessage { - id: 0, - session_id: message.session_id.clone(), - role: role.to_string(), - content: convert_general_blocks_to_unified_content( - &message.content, - message.blocks.as_deref(), - ), - tool_calls: None, - tool_call_id: None, - metadata: message.metadata.clone(), - created_at: timestamp_ms_to_rfc3339(message.created_at), - } -} - -fn mirror_general_message_to_unified( - conn: &rusqlite::Connection, - message: &ChatMessage, -) -> Result { - ensure_general_session_shadow_by_id(conn, &message.session_id)?; - let unified_message = convert_general_message_to_unified(message); - ChatDao::add_message(conn, &unified_message).map_err(|e| format!("写入 unified 消息失败: {e}")) -} - -fn log_general_session_shadow_result( - action: &'static str, - session_id: &str, - result: Result<(), String>, -) { - match result { - Ok(()) => { - tracing::info!( - "[GeneralChat][Compat] {} 已同步 unified 会话影子: session={}", - action, - session_id - ); - } - Err(error) => { - tracing::warn!( - "[GeneralChat][Compat] {} 未能同步 unified 会话影子: session={}, error={}", - action, - session_id, - error - ); - } - } -} - -fn log_general_message_mirror_result( - action: &'static str, - message: &ChatMessage, - result: Result, -) { - match result { - Ok(unified_id) => tracing::info!( - "[GeneralChat][Compat] {} 已同步 unified 消息: session={}, role={:?}, unified_message_id={}", - action, - message.session_id, - message.role, - unified_id - ), - Err(error) => tracing::warn!( - "[GeneralChat][Compat] {} 未能同步 unified 消息: session={}, role={:?}, error={}", - action, - message.session_id, - message.role, - error - ), - } -} - -// ==================== 会话管理命令 ==================== - -/// 兼容层:创建新会话。 -/// -/// # Arguments -/// * `name` - 会话名称(可选,默认为"新对话") -/// * `metadata` - 额外元数据(可选) -#[tauri::command] -pub async fn general_chat_create_session( - db: State<'_, DbConnection>, - name: Option, - metadata: Option, -) -> Result { - warn_general_chat_legacy( - "general_chat_create_session", - "chat_create_session(mode = ChatMode::General)", - ); - - let now = chrono::Utc::now().timestamp_millis(); - let session = ChatSession { - id: Uuid::new_v4().to_string(), - name: name.unwrap_or_else(|| "新对话".to_string()), - created_at: now, - updated_at: now, - metadata, - }; - - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - GeneralChatDao::create_session(&conn, &session).map_err(|e| format!("创建会话失败: {e}"))?; - log_general_session_shadow_result( - "创建会话", - &session.id, - ensure_general_session_shadow(&conn, &session), - ); - - tracing::info!( - "[GeneralChat] 创建会话: id={}, name={}", - session.id, - session.name - ); - Ok(session) -} - -/// 兼容层:获取会话列表。 -#[tauri::command] -pub async fn general_chat_list_sessions( - db: State<'_, DbConnection>, -) -> Result, String> { - warn_general_chat_legacy( - "general_chat_list_sessions", - "chat_list_sessions(mode = Some(ChatMode::General))", - ); - - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - let legacy_sessions = - GeneralChatDao::list_sessions(&conn).map_err(|e| format!("获取会话列表失败: {e}"))?; - let unified_sessions = ChatDao::list_sessions(&conn, Some(ChatMode::General)) - .map_err(|e| format!("获取 unified 会话列表失败: {e}"))?; - let unified_session_map: HashMap = unified_sessions - .into_iter() - .map(|session| (session.id.clone(), session)) - .collect(); - - let mut sessions: Vec = legacy_sessions - .into_iter() - .map(|session| { - overlay_general_session(Some(session.clone()), unified_session_map.get(&session.id)) - .unwrap_or(session) - }) - .collect(); - sessions.sort_by(|left, right| right.updated_at.cmp(&left.updated_at)); - - Ok(sessions) -} - -/// 兼容层:获取会话详情(包含消息列表)。 -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `message_limit` - 消息数量限制(可选) -#[tauri::command] -pub async fn general_chat_get_session( - db: State<'_, DbConnection>, - session_id: String, - message_limit: Option, -) -> Result { - warn_general_chat_legacy("general_chat_get_session", "chat_get_session"); - - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let legacy_session = GeneralChatDao::get_session(&conn, &session_id) - .map_err(|e| format!("获取 general_chat 会话失败: {e}"))?; - let unified_session = ChatDao::get_session(&conn, &session_id) - .map_err(|e| format!("获取 unified 会话失败: {e}"))?; - - let session = overlay_general_session(legacy_session, unified_session.as_ref()) - .ok_or_else(|| "会话不存在".to_string())?; - let all_messages = load_merged_general_messages(&conn, &session_id)?; - let message_count = all_messages.len() as i64; - let messages = paginate_general_messages(all_messages, message_limit, None); - - Ok(SessionDetail { - session, - messages, - message_count, - }) -} - -/// 兼容层:删除会话。 -/// -/// # Arguments -/// * `session_id` - 会话 ID -#[tauri::command] -pub async fn general_chat_delete_session( - db: State<'_, DbConnection>, - session_id: String, -) -> Result { - warn_general_chat_legacy("general_chat_delete_session", "chat_delete_session"); - - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let deleted = GeneralChatDao::delete_session(&conn, &session_id) - .map_err(|e| format!("删除会话失败: {e}"))?; - - if deleted { - match ChatDao::delete_session(&conn, &session_id) { - Ok(true) => tracing::info!( - "[GeneralChat][Compat] 已删除 unified 会话影子: session={}", - session_id - ), - Ok(false) => tracing::debug!( - "[GeneralChat][Compat] 未找到 unified 会话影子,无需删除: session={}", - session_id - ), - Err(error) => tracing::warn!( - "[GeneralChat][Compat] 删除 unified 会话影子失败: session={}, error={}", - session_id, - error - ), - } - tracing::info!("[GeneralChat] 删除会话: id={}", session_id); - } - - Ok(deleted) -} - -/// 兼容层:重命名会话。 -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `name` - 新名称 -#[tauri::command] -pub async fn general_chat_rename_session( - db: State<'_, DbConnection>, - session_id: String, - name: String, -) -> Result { - warn_general_chat_legacy("general_chat_rename_session", "chat_rename_session"); - - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let renamed = GeneralChatDao::rename_session(&conn, &session_id, &name) - .map_err(|e| format!("重命名会话失败: {e}"))?; - - if renamed { - log_general_session_shadow_result( - "重命名会话", - &session_id, - ensure_general_session_shadow_by_id(&conn, &session_id), - ); - tracing::info!("[GeneralChat] 重命名会话: id={}, name={}", session_id, name); - } - - Ok(renamed) -} - -// ==================== 消息管理命令 ==================== - -/// 兼容层:获取会话消息列表。 -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `limit` - 消息数量限制(可选) -/// * `before_id` - 在此消息 ID 之前的消息(用于分页) -#[tauri::command] -pub async fn general_chat_get_messages( - db: State<'_, DbConnection>, - session_id: String, - limit: Option, - before_id: Option, -) -> Result, String> { - warn_general_chat_legacy("general_chat_get_messages", "chat_get_messages"); - - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - let messages = paginate_general_messages( - load_merged_general_messages(&conn, &session_id)?, - limit, - before_id.as_deref(), - ); - - Ok(messages) -} - -/// 兼容层:添加消息到会话。 -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `role` - 消息角色 (user/assistant/system) -/// * `content` - 消息内容 -/// * `blocks` - 内容块列表(可选) -/// * `metadata` - 额外元数据(可选) -#[tauri::command] -pub async fn general_chat_add_message( - _db: State<'_, DbConnection>, - _session_id: String, - _role: String, - _content: String, - _blocks: Option>, - _metadata: Option, -) -> Result { - warn_general_chat_legacy( - "general_chat_add_message", - "统一对话消息流程(暂无一对一 Tauri 替代命令)", - ); - - Err(general_chat_deprecated_error( - "general_chat_add_message", - "统一对话消息流程", - )) -} - -// ==================== 旧流式消息兼容命令 ==================== - -/// 流式消息请求 -#[derive(Debug, Deserialize)] -pub struct SendMessageRequest { - /// 会话 ID - pub session_id: String, - /// 用户消息内容 - pub content: String, - /// 事件名称(用于前端监听) - pub event_name: String, - /// Provider 配置(可选) - #[serde(default)] - #[allow(dead_code)] - pub provider: Option, - /// 模型名称(可选) - #[serde(default)] - #[allow(dead_code)] - pub model: Option, -} - -/// 兼容层:发送消息并获取流式响应。 -/// -/// 该命令历史上维护了一套独立于现役链路之外的模拟流式实现, -/// 会造成“命令还在、行为却已失真”的治理问题。 -/// 现在仅保留命令名用于兼容探测,并显式返回迁移错误。 -#[tauri::command] -pub async fn general_chat_send_message(_request: SendMessageRequest) -> Result { - warn_general_chat_legacy( - "general_chat_send_message", - "chat_send_message / aster_agent_chat_stream", - ); - - Err(general_chat_deprecated_error( - "general_chat_send_message", - "chat_send_message / aster_agent_chat_stream", - )) -} - -/// 兼容层:停止生成。 -/// -/// 旧实现依赖 compat 层自建的停止标志,已与现役 Aster 会话停止链路脱节。 -/// 现在仅保留命令名用于兼容探测,并显式返回迁移错误。 -/// -/// # Arguments -/// * `session_id` - 会话 ID -#[tauri::command] -pub async fn general_chat_stop_generation(_session_id: String) -> Result { - warn_general_chat_legacy( - "general_chat_stop_generation", - "chat_stop_generation / aster_agent_stop", - ); - - Err(general_chat_deprecated_error( - "general_chat_stop_generation", - "chat_stop_generation / aster_agent_stop", - )) -} - -/// 自动生成会话标题请求 -#[derive(Debug, Deserialize)] -pub struct GenerateTitleRequest { - /// 会话 ID - pub session_id: String, - /// 用户第一条消息内容 - pub first_message: String, - /// Provider 名称(可选,暂未使用,预留给未来支持多 provider) - #[serde(default)] - pub provider: Option, - /// 模型名称(可选,用于指定生成标题的模型) - #[serde(default)] - pub model: Option, -} - -/// 兼容层:自动生成会话标题。 -/// -/// 基于用户第一条消息,调用 AI 生成简短的会话标题 -/// -/// # Arguments -/// * `request` - 生成标题请求 -#[tauri::command] -pub async fn general_chat_generate_title( - _db: State<'_, DbConnection>, - _request: GenerateTitleRequest, -) -> Result { - warn_general_chat_legacy( - "general_chat_generate_title", - "统一对话标题流程(暂无一对一 Tauri 替代命令)", - ); - Err(general_chat_deprecated_error( - "general_chat_generate_title", - "前端本地标题规则 + general_chat_rename_session / chat_rename_session", - )) -} - -#[cfg(test)] -mod tests { - use super::*; - - fn build_general_session( - id: &str, - name: &str, - created_at: i64, - updated_at: i64, - ) -> ChatSession { - ChatSession { - id: id.to_string(), - name: name.to_string(), - created_at, - updated_at, - metadata: None, - } - } - - fn build_unified_session( - id: &str, - title: Option<&str>, - created_at: &str, - updated_at: &str, - ) -> UnifiedChatSession { - UnifiedChatSession { - id: id.to_string(), - mode: ChatMode::General, - title: title.map(ToString::to_string), - system_prompt: None, - model: None, - provider_type: None, - credential_uuid: None, - metadata: None, - created_at: created_at.to_string(), - updated_at: updated_at.to_string(), - } - } - - fn build_general_message( - id: &str, - role: MessageRole, - content: &str, - created_at: i64, - ) -> ChatMessage { - ChatMessage { - id: id.to_string(), - session_id: "session-1".to_string(), - role, - content: content.to_string(), - blocks: None, - status: "complete".to_string(), - created_at, - metadata: None, - } - } - - #[test] - fn overlay_general_session_prefers_unified_title_and_latest_timestamp() { - let legacy_session = build_general_session("session-1", "旧标题", 1000, 2000); - let unified_session = build_unified_session( - "session-1", - Some("新标题"), - "1970-01-01T00:00:01Z", - "1970-01-01T00:00:05Z", - ); - - let merged = overlay_general_session(Some(legacy_session), Some(&unified_session)) - .expect("会话应存在"); - - assert_eq!(merged.name, "新标题"); - assert_eq!(merged.created_at, 1000); - assert_eq!(merged.updated_at, 5000); - } - - #[test] - fn merge_general_message_sources_deduplicates_mirrored_messages() { - let legacy_message = build_general_message("legacy-1", MessageRole::User, "你好", 1000); - let unified_duplicate = - build_general_message("unified-101", MessageRole::User, "你好", 1000); - let unified_new = - build_general_message("unified-102", MessageRole::Assistant, "收到", 2000); - - let merged = merge_general_message_sources( - vec![legacy_message.clone()], - vec![unified_duplicate, unified_new.clone()], - ); - - assert_eq!(merged.len(), 2); - assert_eq!(merged[0].id, legacy_message.id); - assert_eq!(merged[1].id, unified_new.id); - } - - #[test] - fn paginate_general_messages_respects_before_id_and_limit() { - let messages = vec![ - build_general_message("msg-1", MessageRole::User, "1", 1000), - build_general_message("msg-2", MessageRole::Assistant, "2", 2000), - build_general_message("msg-3", MessageRole::User, "3", 3000), - build_general_message("msg-4", MessageRole::Assistant, "4", 4000), - ]; - - let paged = paginate_general_messages(messages, Some(2), Some("msg-4")); - - assert_eq!(paged.len(), 2); - assert_eq!(paged[0].id, "msg-2"); - assert_eq!(paged[1].id, "msg-3"); - } - - #[test] - fn deprecated_error_mentions_command_and_replacement() { - let error = general_chat_deprecated_error( - "general_chat_send_message", - "chat_send_message / aster_agent_chat_stream", - ); - - assert!(error.contains("general_chat_send_message")); - assert!(error.contains("chat_send_message / aster_agent_chat_stream")); - assert!(error.contains("已废弃")); - } -} diff --git a/src-tauri/src/commands/memory_management_cmd.rs b/src-tauri/src/commands/memory_management_cmd.rs index 2d05965f5..5aa1f9a9f 100644 --- a/src-tauri/src/commands/memory_management_cmd.rs +++ b/src-tauri/src/commands/memory_management_cmd.rs @@ -1,6 +1,11 @@ //! 记忆管理命令 //! -//! 提供对话记忆的统计和管理功能 +//! 提供记忆相关的统计、治理与自动记忆配置能力。 +//! +//! 其中: +//! - `memory_runtime_*` 属于当前 runtime / 上下文记忆主入口 +//! - `memory_get_*` / `memory_toggle_auto` / `memory_update_auto_note` +//! 属于当前仍在演进的记忆治理配置入口 use crate::commands::context_memory::ContextMemoryServiceState; use crate::config::GlobalConfigManagerState; @@ -8,12 +13,14 @@ use crate::database::DbConnection; use crate::services::auto_memory_service::{ get_auto_memory_index, update_auto_memory_note, AutoMemoryIndexResponse, }; +use crate::services::chat_history_service::{load_memory_source_candidates, MemorySourceCandidate}; use crate::services::memory_source_resolver_service::{ resolve_effective_sources, EffectiveMemorySourcesResponse, }; use chrono::{Local, NaiveDateTime, TimeZone}; +use proxycast_core::app_paths; use proxycast_services::context_memory_service::{MemoryEntry, MemoryFileType}; -use rusqlite::{params, Connection}; +use rusqlite::Connection; use serde::{Deserialize, Serialize}; use std::collections::{HashMap, HashSet}; use std::fs; @@ -127,9 +134,7 @@ const MAX_GENERATED_PER_REQUEST_CAP: usize = 2000; const MAX_GENERATED_PER_SESSION: usize = 40; const MIN_MESSAGE_LENGTH: usize = 18; -/// 获取对话记忆统计信息 -#[tauri::command] -pub async fn get_conversation_memory_stats() -> Result { +async fn memory_runtime_get_stats_impl() -> Result { info!("[记忆管理] 获取记忆统计信息"); let memory_dir = resolve_memory_dir(); @@ -137,9 +142,13 @@ pub async fn get_conversation_memory_stats() -> Result Result { + memory_runtime_get_stats_impl().await +} + +async fn memory_runtime_get_overview_impl( limit: Option, ) -> Result { info!("[记忆管理] 获取记忆总览, limit={:?}", limit); @@ -154,9 +163,15 @@ pub async fn get_conversation_memory_overview( Ok(overview) } -/// 从历史对话中抽取记忆条目 +/// 获取 runtime / 上下文记忆总览(分类 + 条目) #[tauri::command] -pub async fn request_conversation_memory_analysis( +pub async fn memory_runtime_get_overview( + limit: Option, +) -> Result { + memory_runtime_get_overview_impl(limit).await +} + +async fn memory_runtime_request_analysis_impl( memory_service: State<'_, ContextMemoryServiceState>, db: State<'_, DbConnection>, global_config: State<'_, GlobalConfigManagerState>, @@ -263,11 +278,26 @@ pub async fn request_conversation_memory_analysis( }) } -/// 清理过期对话记忆 -/// -/// 清理超过保留天数的记忆条目 +/// 从历史对话中抽取 runtime / 上下文记忆条目 #[tauri::command] -pub async fn cleanup_conversation_memory( +pub async fn memory_runtime_request_analysis( + memory_service: State<'_, ContextMemoryServiceState>, + db: State<'_, DbConnection>, + global_config: State<'_, GlobalConfigManagerState>, + from_timestamp: Option, + to_timestamp: Option, +) -> Result { + memory_runtime_request_analysis_impl( + memory_service, + db, + global_config, + from_timestamp, + to_timestamp, + ) + .await +} + +async fn memory_runtime_cleanup_impl( memory_service: State<'_, ContextMemoryServiceState>, global_config: State<'_, GlobalConfigManagerState>, ) -> Result { @@ -287,7 +317,6 @@ pub async fn cleanup_conversation_memory( let memory_dir = resolve_memory_dir(); let before = collect_memory_overview(&memory_dir)?; - // 使用 ContextMemoryService 的清理功能 memory_service .0 .cleanup_expired_memories_with_retention_days(retention_days)?; @@ -309,6 +338,15 @@ pub async fn cleanup_conversation_memory( }) } +/// 清理 runtime / 上下文记忆 +#[tauri::command] +pub async fn memory_runtime_cleanup( + memory_service: State<'_, ContextMemoryServiceState>, + global_config: State<'_, GlobalConfigManagerState>, +) -> Result { + memory_runtime_cleanup_impl(memory_service, global_config).await +} + /// 获取当前会话可见的有效记忆来源(含 AGENTS、规则、自动记忆) #[tauri::command] pub async fn memory_get_effective_sources( @@ -375,9 +413,7 @@ pub async fn memory_update_auto_note( } fn resolve_memory_dir() -> PathBuf { - dirs::home_dir() - .map(|p| p.join(".proxycast").join("memory")) - .unwrap_or_else(|| PathBuf::from(".proxycast/memory")) + app_paths::best_effort_runtime_subdir("memory") } fn resolve_working_dir(working_dir: Option) -> Result { @@ -833,211 +869,18 @@ fn truncate_text(input: &str, max_chars: usize) -> String { } } -#[derive(Debug, Clone)] -struct MemorySourceCandidate { - session_id: String, - role: String, - content: String, - created_at: i64, -} - fn load_memory_candidates( conn: &Connection, from_timestamp: Option, to_timestamp: Option, ) -> Result, String> { - let mut candidates = Vec::new(); - - let mut push_filtered = |session_id: String, role: String, content: String, created_at: i64| { - let normalized = normalize_candidate_content(&content); - if normalized.len() < MIN_MESSAGE_LENGTH { - return; - } - - let normalized_role = role.to_lowercase(); - if normalized_role != "user" && normalized_role != "assistant" { - return; - } - - candidates.push(MemorySourceCandidate { - session_id, - role: normalized_role, - content: normalized, - created_at, - }); - }; - - let from_ts = from_timestamp; - let to_ts = to_timestamp; - - if from_timestamp.is_some() || to_timestamp.is_some() { - let mut stmt = conn - .prepare( - "SELECT session_id, role, content, created_at - FROM general_chat_messages - WHERE (?1 IS NULL OR created_at >= ?1) - AND (?2 IS NULL OR created_at <= ?2) - ORDER BY created_at DESC - LIMIT ?3", - ) - .map_err(|e| format!("查询 general_chat_messages 失败: {e}"))?; - - let rows = stmt - .query_map( - params![from_timestamp, to_timestamp, MAX_SOURCE_MESSAGES as i64], - |row| { - let session_id: String = row.get(0)?; - let role: String = row.get(1)?; - let content: String = row.get(2)?; - let created_at: i64 = row.get(3)?; - Ok((session_id, role, content, created_at)) - }, - ) - .map_err(|e| format!("读取 general_chat_messages 失败: {e}"))?; - - for row in rows.flatten() { - push_filtered(row.0, row.1, row.2, row.3); - } - - let mut stmt = conn - .prepare( - "SELECT session_id, role, content_json, timestamp - FROM agent_messages - ORDER BY timestamp DESC - LIMIT ?1", - ) - .map_err(|e| format!("查询 agent_messages 失败: {e}"))?; - - let rows = stmt - .query_map(params![MAX_SOURCE_MESSAGES as i64], |row| { - let session_id: String = row.get(0)?; - let role: String = row.get(1)?; - let content_json: String = row.get(2)?; - let timestamp: String = row.get(3)?; - Ok((session_id, role, content_json, timestamp)) - }) - .map_err(|e| format!("读取 agent_messages 失败: {e}"))?; - - for row in rows.flatten() { - if let Some(timestamp_ms) = parse_rfc3339_to_timestamp(&row.3) { - if from_ts.is_some_and(|from| timestamp_ms < from) - || to_ts.is_some_and(|to| timestamp_ms > to) - { - continue; - } - - let text = extract_text_from_content_json(&row.2); - push_filtered(row.0, row.1, text, timestamp_ms); - } - } - } else { - let mut stmt = conn - .prepare( - "SELECT session_id, role, content, created_at - FROM general_chat_messages - ORDER BY created_at DESC - LIMIT ?1", - ) - .map_err(|e| format!("查询 general_chat_messages 失败: {e}"))?; - - let rows = stmt - .query_map(params![MAX_SOURCE_MESSAGES as i64], |row| { - let session_id: String = row.get(0)?; - let role: String = row.get(1)?; - let content: String = row.get(2)?; - let created_at: i64 = row.get(3)?; - Ok((session_id, role, content, created_at)) - }) - .map_err(|e| format!("读取 general_chat_messages 失败: {e}"))?; - - for row in rows.flatten() { - push_filtered(row.0, row.1, row.2, row.3); - } - - let mut stmt = conn - .prepare( - "SELECT session_id, role, content_json, timestamp - FROM agent_messages - ORDER BY timestamp DESC - LIMIT ?1", - ) - .map_err(|e| format!("查询 agent_messages 失败: {e}"))?; - - let rows = stmt - .query_map(params![MAX_SOURCE_MESSAGES as i64], |row| { - let session_id: String = row.get(0)?; - let role: String = row.get(1)?; - let content_json: String = row.get(2)?; - let timestamp: String = row.get(3)?; - Ok((session_id, role, content_json, timestamp)) - }) - .map_err(|e| format!("读取 agent_messages 失败: {e}"))?; - - for row in rows.flatten() { - if let Some(timestamp_ms) = parse_rfc3339_to_timestamp(&row.3) { - let text = extract_text_from_content_json(&row.2); - push_filtered(row.0, row.1, text, timestamp_ms); - } - } - } - - candidates.sort_by(|a, b| b.created_at.cmp(&a.created_at)); - candidates.truncate(MAX_SOURCE_MESSAGES); - - Ok(candidates) -} - -fn normalize_candidate_content(content: &str) -> String { - content - .replace('\n', " ") - .split_whitespace() - .collect::>() - .join(" ") -} - -fn extract_text_from_content_json(content_json: &str) -> String { - if let Ok(text) = serde_json::from_str::(content_json) { - return text; - } - - if let Ok(value) = serde_json::from_str::(content_json) { - match value { - serde_json::Value::Array(items) => { - let texts = items - .iter() - .filter_map(extract_text_from_json_item) - .collect::>(); - if !texts.is_empty() { - return texts.join(" "); - } - } - serde_json::Value::Object(_) => { - if let Some(text) = extract_text_from_json_item(&value) { - return text; - } - } - _ => {} - } - } - - content_json.to_string() -} - -fn extract_text_from_json_item(value: &serde_json::Value) -> Option { - if let Some(text) = value.get("Text").and_then(|v| v.as_str()) { - return Some(text.to_string()); - } - - if value.get("type").and_then(|v| v.as_str()) == Some("text") { - if let Some(text) = value.get("text").and_then(|v| v.as_str()) { - return Some(text.to_string()); - } - } - - value - .get("text") - .and_then(|v| v.as_str()) - .map(|v| v.to_string()) + load_memory_source_candidates( + conn, + from_timestamp, + to_timestamp, + MAX_SOURCE_MESSAGES, + MIN_MESSAGE_LENGTH, + ) } fn build_fingerprint(content: &str) -> String { @@ -1129,13 +972,6 @@ fn map_category_display_name(category: &str) -> &'static str { } } -fn parse_rfc3339_to_timestamp(value: &str) -> Option { - chrono::DateTime::parse_from_rfc3339(value) - .ok() - .map(|dt| dt.timestamp_millis()) - .or_else(|| parse_datetime_or_timestamp_to_millis(value)) -} - fn format_timestamp(timestamp_ms: i64) -> String { if timestamp_ms <= 0 { return "未知时间".to_string(); diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index c35ff853b..147e0c423 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -18,7 +18,6 @@ pub mod external_tools_cmd; pub mod file_upload_cmd; pub mod gateway_channel_cmd; pub mod gateway_tunnel_cmd; -pub mod general_chat_cmd; pub mod heartbeat_cmd; pub mod image_search_cmd; pub mod image_upload_cmd; diff --git a/src-tauri/src/commands/prompt_cmd.rs b/src-tauri/src/commands/prompt_cmd.rs index 8c8eaf29a..45f8b2ad8 100644 --- a/src-tauri/src/commands/prompt_cmd.rs +++ b/src-tauri/src/commands/prompt_cmd.rs @@ -69,13 +69,3 @@ pub fn get_current_prompt_file_content(app: String) -> Result, St pub fn auto_import_prompt(db: State<'_, DbConnection>, app: String) -> Result { PromptService::import_on_first_launch(&db, &app) } - -// Legacy command for compatibility -#[tauri::command] -pub fn switch_prompt( - db: State<'_, DbConnection>, - app_type: String, - id: String, -) -> Result<(), String> { - PromptService::enable(&db, &app_type, &id) -} diff --git a/src-tauri/src/commands/skill_cmd.rs b/src-tauri/src/commands/skill_cmd.rs index 2831a9e3e..e57a21e2c 100644 --- a/src-tauri/src/commands/skill_cmd.rs +++ b/src-tauri/src/commands/skill_cmd.rs @@ -4,6 +4,7 @@ use crate::database::DbConnection; use crate::models::app_type::AppType; use crate::models::skill_model::{Skill, SkillRepo, SkillState}; use chrono::Utc; +use proxycast_core::app_paths; use proxycast_services::skill_service::SkillService; use std::path::{Component, Path, PathBuf}; use std::sync::Arc; @@ -43,16 +44,18 @@ pub fn scan_installed_skills(skills_dir: &Path) -> Vec { } fn get_skills_dir(app_type: &AppType) -> Result { - let home = dirs::home_dir().ok_or_else(|| "Failed to get home directory".to_string())?; - - let skills_dir = match app_type { - AppType::Claude => home.join(".claude").join("skills"), - AppType::Codex => home.join(".codex").join("skills"), - AppType::Gemini => home.join(".gemini").join("skills"), - AppType::ProxyCast => home.join(".proxycast").join("skills"), - }; - - Ok(skills_dir) + match app_type { + AppType::ProxyCast => app_paths::resolve_skills_dir(), + AppType::Claude => dirs::home_dir() + .ok_or_else(|| "Failed to get home directory".to_string()) + .map(|home| home.join(".claude").join("skills")), + AppType::Codex => dirs::home_dir() + .ok_or_else(|| "Failed to get home directory".to_string()) + .map(|home| home.join(".codex").join("skills")), + AppType::Gemini => dirs::home_dir() + .ok_or_else(|| "Failed to get home directory".to_string()) + .map(|home| home.join(".gemini").join("skills")), + } } fn validate_skill_directory(directory: &str) -> Result<(), String> { @@ -113,7 +116,7 @@ fn read_local_skill_content(skills_dir: &Path, directory: &str) -> Result Result Result, String> { - let home = dirs::home_dir().ok_or_else(|| "Failed to get home directory".to_string())?; - let skills_dir = home.join(".proxycast").join("skills"); + let skills_dir = get_skills_dir(&AppType::ProxyCast)?; Ok(scan_installed_skills(&skills_dir)) } diff --git a/src-tauri/src/commands/unified_memory_cmd.rs b/src-tauri/src/commands/unified_memory_cmd.rs index 59500cab2..aa37e0d73 100644 --- a/src-tauri/src/commands/unified_memory_cmd.rs +++ b/src-tauri/src/commands/unified_memory_cmd.rs @@ -4,7 +4,7 @@ use crate::config::GlobalConfigManagerState; use crate::database::DbConnection; -use chrono::{Local, TimeZone}; +use crate::services::chat_history_service::{load_memory_source_candidates, MemorySourceCandidate}; use proxycast_memory::extractor::{self, ExtractionContext}; use proxycast_memory::gatekeeper::ChatMessage; use proxycast_memory::{MemoryCategory, MemoryMetadata, MemorySource, MemoryType, UnifiedMemory}; @@ -81,14 +81,6 @@ pub struct MemoryAnalysisResult { pub deduplicated_entries: u32, } -#[derive(Debug, Clone)] -struct MemorySourceCandidate { - session_id: String, - role: String, - content: String, - created_at: i64, -} - #[derive(Debug, Clone)] struct PendingMemory { session_id: String, @@ -908,92 +900,13 @@ fn load_memory_candidates( from_timestamp: Option, to_timestamp: Option, ) -> Result, String> { - let mut candidates = Vec::new(); - - let mut push_candidate = - |session_id: String, role: String, content: String, created_at: i64| { - let normalized = normalize_candidate_content(&content); - if normalized.len() < MIN_MESSAGE_LENGTH { - return; - } - - let normalized_role = role.to_lowercase(); - if normalized_role != "user" && normalized_role != "assistant" { - return; - } - - candidates.push(MemorySourceCandidate { - session_id, - role: normalized_role, - content: normalized, - created_at: normalize_timestamp(created_at), - }); - }; - - let mut stmt = conn - .prepare( - "SELECT session_id, role, content, created_at - FROM general_chat_messages - WHERE (?1 IS NULL OR created_at >= ?1) - AND (?2 IS NULL OR created_at <= ?2) - ORDER BY created_at DESC - LIMIT ?3", - ) - .map_err(|e| format!("查询 general_chat_messages 失败: {e}"))?; - - let rows = stmt - .query_map( - params![from_timestamp, to_timestamp, MAX_SOURCE_MESSAGES as i64], - |row| { - let session_id: String = row.get(0)?; - let role: String = row.get(1)?; - let content: String = row.get(2)?; - let created_at: i64 = row.get(3)?; - Ok((session_id, role, content, created_at)) - }, - ) - .map_err(|e| format!("读取 general_chat_messages 失败: {e}"))?; - - for row in rows.flatten() { - push_candidate(row.0, row.1, row.2, row.3); - } - - let mut stmt = conn - .prepare( - "SELECT session_id, role, content_json, timestamp - FROM agent_messages - ORDER BY timestamp DESC - LIMIT ?1", - ) - .map_err(|e| format!("查询 agent_messages 失败: {e}"))?; - - let rows = stmt - .query_map(params![MAX_SOURCE_MESSAGES as i64], |row| { - let session_id: String = row.get(0)?; - let role: String = row.get(1)?; - let content_json: String = row.get(2)?; - let timestamp: String = row.get(3)?; - Ok((session_id, role, content_json, timestamp)) - }) - .map_err(|e| format!("读取 agent_messages 失败: {e}"))?; - - for row in rows.flatten() { - if let Some(timestamp_ms) = parse_rfc3339_to_timestamp(&row.3) { - if from_timestamp.is_some_and(|from| timestamp_ms < from) - || to_timestamp.is_some_and(|to| timestamp_ms > to) - { - continue; - } - - let text = extract_text_from_content_json(&row.2); - push_candidate(row.0, row.1, text, timestamp_ms); - } - } - - candidates.sort_by(|a, b| b.created_at.cmp(&a.created_at)); - candidates.truncate(MAX_SOURCE_MESSAGES); - - Ok(candidates) + load_memory_source_candidates( + conn, + from_timestamp, + to_timestamp, + MAX_SOURCE_MESSAGES, + MIN_MESSAGE_LENGTH, + ) } fn build_rule_entry_fields(candidate: &MemorySourceCandidate) -> (String, String, MemoryCategory) { @@ -1150,14 +1063,6 @@ fn normalize_tags(tags: Vec) -> Vec { normalized } -fn normalize_candidate_content(content: &str) -> String { - content - .replace('\n', " ") - .split_whitespace() - .collect::>() - .join(" ") -} - fn normalize_text(input: &str) -> String { input .trim() @@ -1253,76 +1158,6 @@ fn normalize_timestamp(ts: i64) -> i64 { } } -fn parse_rfc3339_to_timestamp(value: &str) -> Option { - chrono::DateTime::parse_from_rfc3339(value) - .ok() - .map(|dt| dt.timestamp_millis()) - .or_else(|| parse_datetime_or_timestamp_to_millis(value)) -} - -fn parse_datetime_or_timestamp_to_millis(value: &str) -> Option { - if let Ok(v) = value.parse::() { - if v > 1_000_000_000_000 { - return Some(v); - } - return Some(v * 1000); - } - - chrono::NaiveDateTime::parse_from_str(value, "%Y-%m-%d %H:%M:%S") - .ok() - .and_then(|naive| { - Local - .from_local_datetime(&naive) - .single() - .map(|dt| dt.timestamp_millis()) - }) -} - -fn extract_text_from_content_json(content_json: &str) -> String { - if let Ok(text) = serde_json::from_str::(content_json) { - return text; - } - - if let Ok(value) = serde_json::from_str::(content_json) { - match value { - serde_json::Value::Array(items) => { - let texts = items - .iter() - .filter_map(extract_text_from_json_item) - .collect::>(); - if !texts.is_empty() { - return texts.join(" "); - } - } - serde_json::Value::Object(_) => { - if let Some(text) = extract_text_from_json_item(&value) { - return text; - } - } - _ => {} - } - } - - content_json.to_string() -} - -fn extract_text_from_json_item(value: &serde_json::Value) -> Option { - if let Some(text) = value.get("Text").and_then(|v| v.as_str()) { - return Some(text.to_string()); - } - - if value.get("type").and_then(|v| v.as_str()) == Some("text") { - if let Some(text) = value.get("text").and_then(|v| v.as_str()) { - return Some(text.to_string()); - } - } - - value - .get("text") - .and_then(|v| v.as_str()) - .map(|v| v.to_string()) -} - fn format_timestamp(timestamp_ms: i64) -> String { let normalized = normalize_timestamp(timestamp_ms); diff --git a/src-tauri/src/commands/workspace_cmd.rs b/src-tauri/src/commands/workspace_cmd.rs index 08dbcd1af..e96eb4bcc 100644 --- a/src-tauri/src/commands/workspace_cmd.rs +++ b/src-tauri/src/commands/workspace_cmd.rs @@ -19,6 +19,7 @@ use crate::services::workspace_health_service::{ use crate::workspace::{ Workspace, WorkspaceManager, WorkspaceSettings, WorkspaceType, WorkspaceUpdate, }; +use proxycast_core::app_paths; use proxycast_services::project_context_builder::ProjectContextBuilder; use serde::{Deserialize, Serialize}; use std::path::PathBuf; @@ -26,14 +27,9 @@ use std::sync::Arc; use tauri::State; use tokio::sync::RwLock; -/// 获取统一的项目根目录(~/.proxycast/projects) +/// 获取统一的项目根目录 fn get_workspace_projects_root_dir() -> Result { - let home_dir = dirs::home_dir().ok_or_else(|| "无法获取主目录".to_string())?; - let root_dir = home_dir.join(".proxycast").join("projects"); - - std::fs::create_dir_all(&root_dir).map_err(|e| format!("创建 workspace 目录失败: {e}"))?; - - Ok(root_dir) + app_paths::resolve_projects_dir() } /// 规范化项目目录名,避免非法路径字符 diff --git a/src-tauri/src/dev_bridge/dispatcher.rs b/src-tauri/src/dev_bridge/dispatcher.rs index 978fdd1c6..a2a716149 100644 --- a/src-tauri/src/dev_bridge/dispatcher.rs +++ b/src-tauri/src/dev_bridge/dispatcher.rs @@ -19,6 +19,7 @@ use crate::services::workspace_health_service::{ ensure_workspace_ready_with_auto_relocate, ensure_workspace_root_ready, }; use crate::workspace::{WorkspaceManager, WorkspaceType, WorkspaceUpdate}; +use proxycast_core::app_paths; use proxycast_memory::{MemoryCategory, MemoryMetadata, MemorySource, MemoryType, UnifiedMemory}; use proxycast_server_utils::load_model_registry_provider_ids_from_resources; use rusqlite::{params_from_iter, types::Value}; @@ -95,12 +96,7 @@ fn parse_optional_nested_arg( } fn get_workspace_projects_root_dir() -> Result { - let home_dir = dirs::home_dir().ok_or_else(|| "无法获取主目录".to_string())?; - let root_dir = home_dir.join(".proxycast").join("projects"); - - std::fs::create_dir_all(&root_dir).map_err(|e| format!("创建 workspace 目录失败: {e}"))?; - - Ok(root_dir) + app_paths::resolve_projects_dir() } fn mask_api_key_for_display(key: &str) -> String { @@ -1242,25 +1238,66 @@ pub async fn handle_command( Ok(serde_json::json!({ "success": true })) } - "get_conversation_memory_overview" => { + "memory_runtime_get_overview" => { let args = args.unwrap_or_default(); let limit = args .get("limit") .and_then(|value| value.as_u64()) .map(|value| value as u32); - let overview = crate::commands::memory_management_cmd::get_conversation_memory_overview(limit) + let overview = crate::commands::memory_management_cmd::memory_runtime_get_overview(limit) .await .map_err(|e| format!("获取对话记忆总览失败: {e}"))?; Ok(serde_json::to_value(overview)?) } - "get_conversation_memory_stats" => { - let stats = crate::commands::memory_management_cmd::get_conversation_memory_stats() + "memory_runtime_get_stats" => { + let stats = crate::commands::memory_management_cmd::memory_runtime_get_stats() .await .map_err(|e| format!("获取对话记忆统计失败: {e}"))?; Ok(serde_json::to_value(stats)?) } + "memory_runtime_request_analysis" => { + let app_handle = state + .app_handle + .as_ref() + .ok_or_else(|| "Dev Bridge 未持有 AppHandle".to_string())?; + let args = args.unwrap_or_default(); + let from_timestamp = args.get("fromTimestamp").and_then(|value| value.as_i64()); + let to_timestamp = args.get("toTimestamp").and_then(|value| value.as_i64()); + let memory_service = + app_handle.state::(); + let db = app_handle.state::(); + let global_config = app_handle.state::(); + let result = crate::commands::memory_management_cmd::memory_runtime_request_analysis( + memory_service, + db, + global_config, + from_timestamp, + to_timestamp, + ) + .await + .map_err(|e| format!("请求记忆分析失败: {e}"))?; + Ok(serde_json::to_value(result)?) + } + + "memory_runtime_cleanup" => { + let app_handle = state + .app_handle + .as_ref() + .ok_or_else(|| "Dev Bridge 未持有 AppHandle".to_string())?; + let memory_service = + app_handle.state::(); + let global_config = app_handle.state::(); + let result = crate::commands::memory_management_cmd::memory_runtime_cleanup( + memory_service, + global_config, + ) + .await + .map_err(|e| format!("清理记忆失败: {e}"))?; + Ok(serde_json::to_value(result)?) + } + // ========== 网络信息 ========== "get_network_info" => { // 返回网络信息 diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 88b63fe64..a24390892 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -8,6 +8,8 @@ //! middleware, orchestrator, plugin, session 部分, session_files) //! - ✅ proxycast-infra crate(proxy, resilience, injection, telemetry) //! - ✅ proxycast-providers crate(providers, converter, streaming, translator, stream, session 部分) + +#![allow(clippy::all)] //! - 主 crate 保留 Tauri 相关业务逻辑 // 抑制 objc crate 宏内部的 unexpected_cfgs 警告 diff --git a/src-tauri/src/services/agent_timeline_service.rs b/src-tauri/src/services/agent_timeline_service.rs new file mode 100644 index 000000000..492d9b466 --- /dev/null +++ b/src-tauri/src/services/agent_timeline_service.rs @@ -0,0 +1,801 @@ +use chrono::Utc; +use proxycast_agent::TauriAgentEvent; +use proxycast_core::database::dao::agent_timeline::{ + AgentRequestOption, AgentRequestQuestion, AgentThreadItem, AgentThreadItemPayload, + AgentThreadItemStatus, AgentThreadTurn, AgentThreadTurnStatus, AgentTimelineDao, +}; +use proxycast_core::database::{lock_db, DbConnection}; +use serde_json::{json, Value}; +use std::collections::HashMap; +use tauri::{AppHandle, Emitter}; +use uuid::Uuid; + +const PROPOSED_PLAN_OPEN: &str = ""; +const PROPOSED_PLAN_CLOSE: &str = ""; + +fn emit_event(app: &AppHandle, event_name: &str, event: &TauriAgentEvent) { + if let Err(error) = app.emit(event_name, event) { + tracing::error!("[AgentTimeline] 发送事件失败: {}", error); + } +} + +fn normalize_tool_name(name: &str) -> String { + name.replace([' ', '-', '_'], "").to_lowercase() +} + +fn parse_json_str(raw: Option<&str>) -> Option { + let value = raw?.trim(); + if value.is_empty() { + return None; + } + serde_json::from_str::(value).ok() +} + +fn as_object(value: &Value) -> Option<&serde_json::Map> { + value.as_object() +} + +fn pick_string_from_object( + object: Option<&serde_json::Map>, + keys: &[&str], +) -> Option { + let object = object?; + for key in keys { + if let Some(value) = object.get(*key).and_then(Value::as_str) { + let trimmed = value.trim(); + if !trimmed.is_empty() { + return Some(trimmed.to_string()); + } + } + } + None +} + +fn extract_tool_query(arguments: Option<&Value>) -> Option { + pick_string_from_object( + arguments.and_then(as_object), + &["q", "query", "question", "search", "search_query", "url"], + ) +} + +fn extract_command_text(arguments: Option<&Value>) -> Option { + pick_string_from_object( + arguments.and_then(as_object), + &["cmd", "command", "script", "text"], + ) +} + +fn extract_file_paths(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; + }; + match value { + Value::String(text) => { + let trimmed = text.trim(); + if !trimmed.is_empty() && !paths.iter().any(|item| item == trimmed) { + paths.push(trimmed.to_string()); + } + } + Value::Array(items) => { + for item in items { + if let Some(text) = item.as_str() { + let trimmed = text.trim(); + if !trimmed.is_empty() && !paths.iter().any(|entry| entry == trimmed) { + paths.push(trimmed.to_string()); + } + } + } + } + _ => {} + } + } + } + paths +} + +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()) + } +} + +fn is_command_tool(name: &str) -> bool { + matches!( + normalize_tool_name(name).as_str(), + "bash" | "execcommand" | "terminal" | "shell" | "runcommand" + ) +} + +fn is_web_tool(name: &str) -> bool { + let normalized = normalize_tool_name(name); + normalized.contains("websearch") + || normalized.contains("searchquery") + || normalized.contains("webfetch") + || normalized.contains("browser") + || normalized.contains("playwright") + || normalized == "search" +} + +fn is_user_input_action(action_type: &str) -> bool { + matches!(action_type, "ask_user" | "elicitation") +} + +fn map_questions(raw: Option<&Value>) -> Option> { + let items = raw?.as_array()?; + let mut questions = Vec::new(); + + for item in items { + let Some(object) = item.as_object() else { + continue; + }; + let Some(question) = object.get("question").and_then(Value::as_str) else { + continue; + }; + + let options = object + .get("options") + .and_then(Value::as_array) + .map(|values| { + values + .iter() + .filter_map(|value| { + let object = value.as_object()?; + let label = object.get("label")?.as_str()?.trim().to_string(); + if label.is_empty() { + return None; + } + Some(AgentRequestOption { + label, + description: object + .get("description") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string), + }) + }) + .collect::>() + }); + + questions.push(AgentRequestQuestion { + question: question.trim().to_string(), + header: object + .get("header") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string), + options: options.filter(|values| !values.is_empty()), + multi_select: object.get("multi_select").and_then(Value::as_bool), + }); + } + + if questions.is_empty() { + None + } else { + Some(questions) + } +} + +#[derive(Debug)] +pub struct AgentTimelineRecorder { + db: DbConnection, + thread_id: String, + turn_id: String, + turn: AgentThreadTurn, + sequence_counter: i64, + item_sequences: HashMap, + item_statuses: HashMap, + assistant_text: String, + reasoning_text: String, + plan_text: Option, +} + +impl AgentTimelineRecorder { + pub fn create( + db: DbConnection, + thread_id: impl Into, + prompt_text: impl Into, + ) -> Result { + let thread_id = thread_id.into(); + let prompt_text = prompt_text.into(); + let now = Utc::now().to_rfc3339(); + let turn = AgentThreadTurn { + id: Uuid::new_v4().to_string(), + thread_id: thread_id.clone(), + prompt_text, + status: AgentThreadTurnStatus::Running, + started_at: now.clone(), + completed_at: None, + error_message: None, + created_at: now.clone(), + updated_at: now, + }; + + { + let conn = lock_db(&db)?; + AgentTimelineDao::create_turn(&conn, &turn) + .map_err(|e| format!("创建 turn 失败: {e}"))?; + } + + Ok(Self { + db, + thread_id, + turn_id: turn.id.clone(), + turn, + sequence_counter: 0, + item_sequences: HashMap::new(), + item_statuses: HashMap::new(), + assistant_text: String::new(), + reasoning_text: String::new(), + plan_text: None, + }) + } + + pub fn thread_id(&self) -> &str { + &self.thread_id + } + + pub fn turn_id(&self) -> &str { + &self.turn_id + } + + pub fn emit_start(&mut self, app: &AppHandle, event_name: &str) -> Result<(), String> { + emit_event( + app, + event_name, + &TauriAgentEvent::ThreadStarted { + thread_id: self.thread_id.clone(), + }, + ); + emit_event( + app, + event_name, + &TauriAgentEvent::TurnStarted { + turn: self.turn.clone(), + }, + ); + + let user_item = self.build_item( + format!("user:{}", self.turn_id), + AgentThreadItemStatus::Completed, + Some(self.turn.started_at.clone()), + AgentThreadItemPayload::UserMessage { + content: self.turn.prompt_text.clone(), + }, + ); + self.persist_and_emit_item(app, event_name, user_item)?; + Ok(()) + } + + pub fn record_legacy_event( + &mut self, + app: &AppHandle, + event_name: &str, + event: &TauriAgentEvent, + workspace_root: &str, + ) -> Result<(), String> { + match event { + TauriAgentEvent::TextDelta { text } => { + self.assistant_text.push_str(text); + let item = self.build_item( + format!("assistant:{}", self.turn_id), + AgentThreadItemStatus::InProgress, + None, + AgentThreadItemPayload::AgentMessage { + text: self.assistant_text.clone(), + phase: None, + }, + ); + self.persist_and_emit_item(app, event_name, item)?; + + if let Some(plan_text) = extract_proposed_plan_block(&self.assistant_text) { + if self.plan_text.as_deref() != Some(plan_text.as_str()) { + self.plan_text = Some(plan_text.clone()); + } + let plan_item = self.build_item( + format!("plan:{}", self.turn_id), + AgentThreadItemStatus::InProgress, + None, + AgentThreadItemPayload::Plan { text: plan_text }, + ); + self.persist_and_emit_item(app, event_name, plan_item)?; + } + } + TauriAgentEvent::ThinkingDelta { text } => { + self.reasoning_text.push_str(text); + let item = self.build_item( + format!("reasoning:{}", self.turn_id), + AgentThreadItemStatus::InProgress, + None, + AgentThreadItemPayload::Reasoning { + text: self.reasoning_text.clone(), + summary: None, + }, + ); + self.persist_and_emit_item(app, event_name, item)?; + } + TauriAgentEvent::ToolStart { + tool_name, + tool_id, + arguments, + } => { + let arguments_value = parse_json_str(arguments.as_deref()); + let payload = if is_command_tool(tool_name) { + AgentThreadItemPayload::CommandExecution { + command: extract_command_text(arguments_value.as_ref()) + .unwrap_or_else(|| tool_name.clone()), + cwd: workspace_root.to_string(), + aggregated_output: None, + exit_code: None, + error: None, + } + } else if is_web_tool(tool_name) { + AgentThreadItemPayload::WebSearch { + query: extract_tool_query(arguments_value.as_ref()), + action: Some(tool_name.clone()), + output: None, + } + } else { + AgentThreadItemPayload::ToolCall { + tool_name: tool_name.clone(), + arguments: arguments_value, + output: None, + success: None, + error: None, + metadata: None, + } + }; + + let item = self.build_item( + tool_id.clone(), + AgentThreadItemStatus::InProgress, + None, + payload, + ); + self.persist_and_emit_item(app, event_name, item)?; + } + TauriAgentEvent::ToolEnd { tool_id, result } => { + let existing = { + let conn = lock_db(&self.db)?; + AgentTimelineDao::get_item(&conn, tool_id) + .map_err(|e| format!("读取工具 item 失败: {e}"))? + }; + + let metadata_value = result + .metadata + .as_ref() + .and_then(|metadata| serde_json::to_value(metadata).ok()); + let status = if result.success { + AgentThreadItemStatus::Completed + } else { + AgentThreadItemStatus::Failed + }; + + let payload = match existing.map(|item| item.payload) { + Some(AgentThreadItemPayload::CommandExecution { command, cwd, .. }) => { + AgentThreadItemPayload::CommandExecution { + command, + cwd, + aggregated_output: Some(result.output.clone()), + exit_code: metadata_value + .as_ref() + .and_then(|value| value.get("exit_code")) + .and_then(Value::as_i64), + error: result.error.clone(), + } + } + Some(AgentThreadItemPayload::WebSearch { query, action, .. }) => { + AgentThreadItemPayload::WebSearch { + query, + action, + output: Some(result.output.clone()), + } + } + Some(AgentThreadItemPayload::ToolCall { + tool_name, + arguments, + .. + }) => AgentThreadItemPayload::ToolCall { + tool_name, + arguments, + output: Some(result.output.clone()), + success: Some(result.success), + error: result.error.clone(), + metadata: metadata_value.clone(), + }, + _ => AgentThreadItemPayload::ToolCall { + tool_name: tool_id.clone(), + arguments: None, + output: Some(result.output.clone()), + success: Some(result.success), + error: result.error.clone(), + metadata: metadata_value.clone(), + }, + }; + + let item = self.build_item( + tool_id.clone(), + status, + Some(Utc::now().to_rfc3339()), + payload, + ); + self.persist_and_emit_item(app, event_name, item)?; + + for path in extract_file_paths(None, metadata_value.as_ref()) { + let file_item = self.build_item( + format!("artifact:{}:{}", tool_id, path), + AgentThreadItemStatus::Completed, + Some(Utc::now().to_rfc3339()), + AgentThreadItemPayload::FileArtifact { + path, + source: "tool_result".to_string(), + content: None, + metadata: metadata_value.clone(), + }, + ); + self.persist_and_emit_item(app, event_name, file_item)?; + } + } + TauriAgentEvent::ActionRequired { + request_id, + action_type, + data, + } => { + let payload = if is_user_input_action(action_type) { + AgentThreadItemPayload::RequestUserInput { + request_id: request_id.clone(), + action_type: action_type.clone(), + prompt: data + .get("prompt") + .or_else(|| data.get("message")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string), + questions: map_questions(data.get("questions")), + response: None, + } + } else { + AgentThreadItemPayload::ApprovalRequest { + request_id: request_id.clone(), + action_type: action_type.clone(), + prompt: data + .get("prompt") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string), + tool_name: data + .get("tool_name") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string), + arguments: data.get("arguments").cloned(), + response: None, + } + }; + + let item = self.build_item( + request_id.clone(), + AgentThreadItemStatus::InProgress, + None, + payload, + ); + self.persist_and_emit_item(app, event_name, item)?; + } + TauriAgentEvent::Warning { code, message } => { + let item = self.build_item( + format!("warning:{}:{}", self.turn_id, self.sequence_counter + 1), + AgentThreadItemStatus::Completed, + Some(Utc::now().to_rfc3339()), + AgentThreadItemPayload::Warning { + message: message.clone(), + code: code.clone(), + }, + ); + self.persist_and_emit_item(app, event_name, item)?; + } + TauriAgentEvent::Error { message } => { + let item = self.build_item( + format!("error:{}", self.turn_id), + AgentThreadItemStatus::Failed, + Some(Utc::now().to_rfc3339()), + AgentThreadItemPayload::Error { + message: message.clone(), + }, + ); + self.persist_and_emit_item(app, event_name, item)?; + } + _ => {} + } + + Ok(()) + } + + pub fn complete_turn_success( + &mut self, + app: &AppHandle, + event_name: &str, + ) -> Result<(), String> { + self.complete_open_content_items(app, event_name, AgentThreadItemStatus::Completed)?; + let now = Utc::now().to_rfc3339(); + self.turn.status = AgentThreadTurnStatus::Completed; + self.turn.completed_at = Some(now.clone()); + self.turn.updated_at = now.clone(); + + let conn = lock_db(&self.db)?; + AgentTimelineDao::update_turn_status( + &conn, + &self.turn_id, + AgentThreadTurnStatus::Completed, + Some(&now), + None, + &now, + ) + .map_err(|e| format!("更新 turn 完成状态失败: {e}"))?; + drop(conn); + + emit_event( + app, + event_name, + &TauriAgentEvent::TurnCompleted { + turn: self.turn.clone(), + }, + ); + Ok(()) + } + + pub fn fail_turn( + &mut self, + app: &AppHandle, + event_name: &str, + message: &str, + ) -> Result<(), String> { + self.complete_open_content_items(app, event_name, AgentThreadItemStatus::Completed)?; + let error_item = self.build_item( + format!("error:{}", self.turn_id), + AgentThreadItemStatus::Failed, + Some(Utc::now().to_rfc3339()), + AgentThreadItemPayload::Error { + message: message.to_string(), + }, + ); + self.persist_and_emit_item(app, event_name, error_item)?; + + let now = Utc::now().to_rfc3339(); + self.turn.status = AgentThreadTurnStatus::Failed; + self.turn.completed_at = Some(now.clone()); + self.turn.error_message = Some(message.to_string()); + self.turn.updated_at = now.clone(); + + let conn = lock_db(&self.db)?; + AgentTimelineDao::update_turn_status( + &conn, + &self.turn_id, + AgentThreadTurnStatus::Failed, + Some(&now), + Some(message), + &now, + ) + .map_err(|e| format!("更新 turn 失败状态失败: {e}"))?; + drop(conn); + + emit_event( + app, + event_name, + &TauriAgentEvent::TurnFailed { + turn: self.turn.clone(), + }, + ); + Ok(()) + } + + fn complete_open_content_items( + &mut self, + app: &AppHandle, + event_name: &str, + status: AgentThreadItemStatus, + ) -> Result<(), String> { + if !self.assistant_text.is_empty() { + let item = self.build_item( + format!("assistant:{}", self.turn_id), + status.clone(), + Some(Utc::now().to_rfc3339()), + AgentThreadItemPayload::AgentMessage { + text: self.assistant_text.clone(), + phase: None, + }, + ); + self.persist_and_emit_item(app, event_name, item)?; + } + + if !self.reasoning_text.is_empty() { + let item = self.build_item( + format!("reasoning:{}", self.turn_id), + status.clone(), + Some(Utc::now().to_rfc3339()), + AgentThreadItemPayload::Reasoning { + text: self.reasoning_text.clone(), + summary: None, + }, + ); + self.persist_and_emit_item(app, event_name, item)?; + } + + if let Some(plan_text) = self.plan_text.clone() { + let item = self.build_item( + format!("plan:{}", self.turn_id), + status, + Some(Utc::now().to_rfc3339()), + AgentThreadItemPayload::Plan { text: plan_text }, + ); + self.persist_and_emit_item(app, event_name, item)?; + } + + Ok(()) + } + + fn build_item( + &mut self, + id: String, + status: AgentThreadItemStatus, + completed_at: Option, + payload: AgentThreadItemPayload, + ) -> AgentThreadItem { + let now = Utc::now().to_rfc3339(); + let started_at = self + .item_statuses + .get(&id) + .map(|_| { + let conn = lock_db(&self.db).ok()?; + AgentTimelineDao::get_item(&conn, &id) + .ok() + .flatten() + .map(|item| item.started_at) + }) + .flatten() + .unwrap_or_else(|| now.clone()); + + let sequence = if let Some(existing) = self.item_sequences.get(&id) { + *existing + } else { + self.sequence_counter += 1; + self.item_sequences + .insert(id.clone(), self.sequence_counter); + self.sequence_counter + }; + + AgentThreadItem { + id, + thread_id: self.thread_id.clone(), + turn_id: self.turn_id.clone(), + sequence, + status, + started_at, + completed_at, + updated_at: now, + payload, + } + } + + fn persist_and_emit_item( + &mut self, + app: &AppHandle, + event_name: &str, + item: AgentThreadItem, + ) -> Result<(), String> { + { + let conn = lock_db(&self.db)?; + AgentTimelineDao::upsert_item(&conn, &item) + .map_err(|e| format!("保存 item 失败: {e}"))?; + } + + let previous_status = self + .item_statuses + .insert(item.id.clone(), item.status.clone()); + let event = match (&previous_status, &item.status) { + (None, AgentThreadItemStatus::InProgress) => { + TauriAgentEvent::ItemStarted { item: item.clone() } + } + (None, _) => TauriAgentEvent::ItemCompleted { item: item.clone() }, + (_, AgentThreadItemStatus::Completed | AgentThreadItemStatus::Failed) => { + TauriAgentEvent::ItemCompleted { item: item.clone() } + } + _ => TauriAgentEvent::ItemUpdated { item: item.clone() }, + }; + emit_event(app, event_name, &event); + 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/auto_memory_service.rs b/src-tauri/src/services/auto_memory_service.rs index ecf76ae2f..ea525dbc8 100644 --- a/src-tauri/src/services/auto_memory_service.rs +++ b/src-tauri/src/services/auto_memory_service.rs @@ -3,6 +3,7 @@ //! 提供自动记忆目录定位、入口索引读取与笔记更新能力。 use chrono::Local; +use proxycast_core::app_paths; use proxycast_core::config::{MemoryAutoConfig, MemoryConfig}; use serde::{Deserialize, Serialize}; use std::fs; @@ -131,10 +132,7 @@ pub fn resolve_auto_memory_root(working_dir: &Path, auto: &MemoryAutoConfig) -> slug }; - dirs::home_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join(".proxycast") - .join("projects") + app_paths::best_effort_runtime_subdir("projects") .join(project_slug) .join("memory") } diff --git a/src-tauri/src/services/chat_history_service.rs b/src-tauri/src/services/chat_history_service.rs new file mode 100644 index 000000000..ed29a05da --- /dev/null +++ b/src-tauri/src/services/chat_history_service.rs @@ -0,0 +1,549 @@ +use crate::database::load_pending_general_messages; +use chrono::{Local, TimeZone}; +use rusqlite::{params, Connection}; +use std::collections::HashSet; + +const GENERAL_MODE_PATTERN: &str = "general:%"; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MemorySourceCandidate { + pub session_id: String, + pub role: String, + pub content: String, + pub created_at: i64, +} + +pub fn load_memory_source_candidates( + conn: &Connection, + from_timestamp: Option, + to_timestamp: Option, + limit: usize, + min_message_length: usize, +) -> Result, String> { + 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, + to_timestamp, + limit, + min_message_length, + &mut candidates, + &mut seen, + )?; + load_non_general_agent_candidates( + conn, + from_timestamp, + to_timestamp, + limit, + min_message_length, + &mut candidates, + &mut seen, + )?; + + candidates.sort_by(|a, b| b.created_at.cmp(&a.created_at)); + candidates.truncate(limit); + + 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, + to_timestamp: Option, + limit: usize, + min_message_length: usize, + candidates: &mut Vec, + seen: &mut HashSet, +) -> Result<(), String> { + let from_datetime = from_timestamp.map(format_sqlite_datetime); + let to_datetime = to_timestamp.map(format_sqlite_datetime); + + let mut stmt = conn + .prepare( + "SELECT m.session_id, m.role, m.content_json, m.timestamp + FROM agent_messages m + JOIN agent_sessions s ON s.id = m.session_id + WHERE s.model LIKE ?1 + AND (?2 IS NULL OR datetime(m.timestamp) >= datetime(?2)) + AND (?3 IS NULL OR datetime(m.timestamp) <= datetime(?3)) + ORDER BY datetime(m.timestamp) DESC + LIMIT ?4", + ) + .map_err(|e| format!("查询 unified general agent_messages 失败: {e}"))?; + + let rows = stmt + .query_map( + params![ + GENERAL_MODE_PATTERN, + from_datetime, + to_datetime, + limit as i64 + ], + |row| { + let session_id: String = row.get(0)?; + let role: String = row.get(1)?; + let content_json: String = row.get(2)?; + let timestamp: String = row.get(3)?; + Ok((session_id, role, content_json, timestamp)) + }, + ) + .map_err(|e| format!("读取 unified general agent_messages 失败: {e}"))?; + + for row in rows.flatten() { + if let Some(timestamp_ms) = parse_rfc3339_to_timestamp(&row.3) { + push_candidate( + candidates, + seen, + row.0, + row.1, + extract_text_from_content_json(&row.2), + timestamp_ms, + min_message_length, + ); + } + } + + Ok(()) +} + +fn load_non_general_agent_candidates( + conn: &Connection, + from_timestamp: Option, + to_timestamp: Option, + limit: usize, + min_message_length: usize, + candidates: &mut Vec, + seen: &mut HashSet, +) -> Result<(), String> { + let from_datetime = from_timestamp.map(format_sqlite_datetime); + let to_datetime = to_timestamp.map(format_sqlite_datetime); + + let mut stmt = conn + .prepare( + "SELECT m.session_id, m.role, m.content_json, m.timestamp + FROM agent_messages m + JOIN agent_sessions s ON s.id = m.session_id + WHERE s.model NOT LIKE ?1 + AND (?2 IS NULL OR datetime(m.timestamp) >= datetime(?2)) + AND (?3 IS NULL OR datetime(m.timestamp) <= datetime(?3)) + ORDER BY datetime(m.timestamp) DESC + LIMIT ?4", + ) + .map_err(|e| format!("查询非通用 agent_messages 失败: {e}"))?; + + let rows = stmt + .query_map( + params![ + GENERAL_MODE_PATTERN, + from_datetime, + to_datetime, + limit as i64 + ], + |row| { + let session_id: String = row.get(0)?; + let role: String = row.get(1)?; + let content_json: String = row.get(2)?; + let timestamp: String = row.get(3)?; + Ok((session_id, role, content_json, timestamp)) + }, + ) + .map_err(|e| format!("读取非通用 agent_messages 失败: {e}"))?; + + for row in rows.flatten() { + if let Some(timestamp_ms) = parse_rfc3339_to_timestamp(&row.3) { + push_candidate( + candidates, + seen, + row.0, + row.1, + extract_text_from_content_json(&row.2), + timestamp_ms, + min_message_length, + ); + } + } + + Ok(()) +} + +fn push_candidate( + candidates: &mut Vec, + seen: &mut HashSet, + session_id: String, + role: String, + content: String, + created_at: i64, + min_message_length: usize, +) { + let normalized = normalize_candidate_content(&content); + if normalized.len() < min_message_length { + return; + } + + let normalized_role = role.to_lowercase(); + if normalized_role != "user" && normalized_role != "assistant" { + return; + } + + let normalized_created_at = normalize_timestamp(created_at); + let dedupe_key = format!( + "{}:{}:{}:{}", + session_id, normalized_role, normalized_created_at, normalized + ); + + if !seen.insert(dedupe_key) { + return; + } + + candidates.push(MemorySourceCandidate { + session_id, + role: normalized_role, + content: normalized, + created_at: normalized_created_at, + }); +} + +fn normalize_candidate_content(content: &str) -> String { + content + .replace('\n', " ") + .split_whitespace() + .collect::>() + .join(" ") +} + +fn normalize_timestamp(ts: i64) -> i64 { + if ts <= 0 { + return chrono::Utc::now().timestamp_millis(); + } + if ts > 1_000_000_000_000 { + ts + } else { + ts * 1000 + } +} + +fn format_sqlite_datetime(timestamp_ms: i64) -> String { + let normalized = normalize_timestamp(timestamp_ms); + Local + .timestamp_millis_opt(normalized) + .single() + .map(|dt| dt.format("%Y-%m-%d %H:%M:%S").to_string()) + .unwrap_or_else(|| Local::now().format("%Y-%m-%d %H:%M:%S").to_string()) +} + +fn parse_rfc3339_to_timestamp(value: &str) -> Option { + chrono::DateTime::parse_from_rfc3339(value) + .ok() + .map(|dt| dt.timestamp_millis()) + .or_else(|| parse_datetime_or_timestamp_to_millis(value)) +} + +fn parse_datetime_or_timestamp_to_millis(value: &str) -> Option { + if let Ok(v) = value.parse::() { + if v > 1_000_000_000_000 { + return Some(v); + } + return Some(v * 1000); + } + + chrono::NaiveDateTime::parse_from_str(value, "%Y-%m-%d %H:%M:%S") + .ok() + .and_then(|naive| { + Local + .from_local_datetime(&naive) + .single() + .map(|dt| dt.timestamp_millis()) + }) +} + +fn extract_text_from_content_json(content_json: &str) -> String { + if let Ok(text) = serde_json::from_str::(content_json) { + return text; + } + + if let Ok(value) = serde_json::from_str::(content_json) { + match value { + serde_json::Value::Array(items) => { + let texts = items + .iter() + .filter_map(extract_text_from_json_item) + .collect::>(); + if !texts.is_empty() { + return texts.join(" "); + } + } + serde_json::Value::Object(_) => { + if let Some(text) = extract_text_from_json_item(&value) { + return text; + } + } + _ => {} + } + } + + content_json.to_string() +} + +fn extract_text_from_json_item(value: &serde_json::Value) -> Option { + if let Some(text) = value.get("Text").and_then(|v| v.as_str()) { + return Some(text.to_string()); + } + + if value.get("type").and_then(|v| v.as_str()) == Some("text") { + if let Some(text) = value.get("text").and_then(|v| v.as_str()) { + return Some(text.to_string()); + } + } + + value + .get("text") + .and_then(|v| v.as_str()) + .map(|v| v.to_string()) +} + +#[cfg(test)] +mod tests { + use super::load_memory_source_candidates; + use rusqlite::{params, Connection}; + + 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 settings ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL + ); + 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 + ); + 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 + ); + ", + ) + .expect("create test schema"); + } + + #[test] + fn load_memory_source_candidates_merges_unified_and_legacy_without_duplicates() { + let conn = Connection::open_in_memory().expect("open in memory db"); + create_test_schema(&conn); + + conn.execute( + "INSERT INTO agent_sessions (id, model, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", + params![ + "general-migrated", + "general:default", + "2026-03-12T10:00:00+08:00", + "2026-03-12T10:00:00+08:00" + ], + ) + .unwrap(); + conn.execute( + "INSERT INTO agent_sessions (id, model, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", + params![ + "agent-1", + "claude-sonnet-4", + "2026-03-12T10:05:00+08:00", + "2026-03-12T10:05:00+08:00" + ], + ) + .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![ + "general-migrated", + "user", + r#"[{"type":"text","text":"这条消息已经迁移"}]"#, + "2025-03-12T10:00:00+08:00" + ], + ) + .unwrap(); + conn.execute( + "INSERT INTO agent_messages (session_id, role, content_json, timestamp) VALUES (?1, ?2, ?3, ?4)", + params![ + "agent-1", + "assistant", + r#"[{"type":"text","text":"这是一条 agent 消息"}]"#, + "2025-03-12T10:05:00+08:00" + ], + ) + .unwrap(); + + let candidates = + load_memory_source_candidates(&conn, None, None, 20, 1).expect("load candidates"); + + let session_ids = candidates + .iter() + .map(|item| item.session_id.as_str()) + .collect::>(); + assert_eq!(candidates.len(), 3); + 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() { + 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![ + "general-migrated", + "general:default", + "2026-03-12T10:00:00+08:00", + "2026-03-12T10:00:00+08:00" + ], + ) + .unwrap(); + conn.execute( + "INSERT INTO agent_sessions (id, model, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", + params![ + "agent-1", + "claude-sonnet-4", + "2026-03-12T10:05:00+08:00", + "2026-03-12T10:05:00+08:00" + ], + ) + .unwrap(); + + conn.execute( + "INSERT INTO general_chat_sessions (id, name, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", + params!["legacy-only", "旧会话", 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", "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![ + "general-migrated", + "user", + r#"[{"type":"text","text":"这是 unified general 消息"}]"#, + "2026-03-12T10:00:00+08:00" + ], + ) + .unwrap(); + conn.execute( + "INSERT INTO agent_messages (session_id, role, content_json, timestamp) VALUES (?1, ?2, ?3, ?4)", + params![ + "agent-1", + "assistant", + r#"[{"type":"text","text":"这是 agent 消息"}]"#, + "2026-03-12T10:05:00+08:00" + ], + ) + .unwrap(); + + let candidates = + load_memory_source_candidates(&conn, None, None, 20, 1).expect("load candidates"); + + let session_ids = candidates + .iter() + .map(|item| item.session_id.as_str()) + .collect::>(); + assert_eq!(candidates.len(), 2); + assert!(session_ids.contains(&"general-migrated")); + assert!(session_ids.contains(&"agent-1")); + assert!(!session_ids.contains(&"legacy-only")); + } +} diff --git a/src-tauri/src/services/conversation_statistics_service.rs b/src-tauri/src/services/conversation_statistics_service.rs index 59896e878..a5cffa31f 100644 --- a/src-tauri/src/services/conversation_statistics_service.rs +++ b/src-tauri/src/services/conversation_statistics_service.rs @@ -2,10 +2,16 @@ //! //! 从数据库查询真实的对话和使用统计数据 -use chrono::{DateTime, Datelike, Duration, Local, Timelike}; -use rusqlite::Connection; +use crate::database::{ + count_pending_general_messages, count_pending_general_sessions, + sum_pending_general_message_chars, +}; +use chrono::{DateTime, Datelike, Duration, Local, TimeZone, Timelike}; +use rusqlite::{params, Connection}; use serde::{Deserialize, Serialize}; +const GENERAL_MODE_PATTERN: &str = "general:%"; + /// 使用统计数据响应 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct UsageStatsResponse { @@ -188,65 +194,171 @@ fn chars_to_estimated_tokens(chars: i64) -> u64 { ((chars as f64) / 4.0).ceil() as u64 } +fn format_sqlite_datetime(timestamp_ms: i64) -> String { + Local + .timestamp_millis_opt(timestamp_ms) + .single() + .map(|dt| dt.format("%Y-%m-%d %H:%M:%S").to_string()) + .unwrap_or_else(|| Local::now().format("%Y-%m-%d %H:%M:%S").to_string()) +} + +fn query_general_session_count( + conn: &Connection, + from_timestamp_ms: Option, + to_timestamp_ms: Option, +) -> Result { + let from_text = from_timestamp_ms.map(format_sqlite_datetime); + let to_text = to_timestamp_ms.map(format_sqlite_datetime); + + let unified_count: i64 = conn + .query_row( + "SELECT COUNT(*) + FROM agent_sessions s + WHERE s.model LIKE ?1 + AND (?2 IS NULL OR datetime(s.created_at) >= datetime(?2)) + AND (?3 IS NULL OR datetime(s.created_at) < datetime(?3))", + params![GENERAL_MODE_PATTERN, from_text, to_text], + |row| row.get(0), + ) + .map_err(|e| format!("查询 unified general 会话数失败: {e}"))?; + + let pending_count = count_pending_general_sessions(conn, from_timestamp_ms, to_timestamp_ms) + .map_err(|e| format!("查询待迁移 general 会话数失败: {e}"))?; + + Ok(unified_count + pending_count) +} + +fn query_general_message_count( + conn: &Connection, + from_timestamp_ms: Option, + to_timestamp_ms: Option, +) -> Result { + let from_text = from_timestamp_ms.map(format_sqlite_datetime); + let to_text = to_timestamp_ms.map(format_sqlite_datetime); + + let unified_count: i64 = conn + .query_row( + "SELECT COUNT(*) + FROM agent_messages m + JOIN agent_sessions s ON s.id = m.session_id + WHERE s.model LIKE ?1 + AND (?2 IS NULL OR datetime(m.timestamp) >= datetime(?2)) + AND (?3 IS NULL OR datetime(m.timestamp) < datetime(?3))", + params![GENERAL_MODE_PATTERN, from_text, to_text], + |row| row.get(0), + ) + .map_err(|e| format!("查询 unified general 消息数失败: {e}"))?; + + let pending_count = count_pending_general_messages(conn, from_timestamp_ms, to_timestamp_ms) + .map_err(|e| format!("查询待迁移 general 消息数失败: {e}"))?; + + Ok(unified_count + pending_count) +} + +fn sum_general_message_chars( + conn: &Connection, + from_timestamp_ms: Option, + to_timestamp_ms: Option, +) -> Result { + let from_text = from_timestamp_ms.map(format_sqlite_datetime); + let to_text = to_timestamp_ms.map(format_sqlite_datetime); + + let unified_chars: i64 = conn + .query_row( + "SELECT COALESCE(SUM(LENGTH(m.content_json)), 0) + FROM agent_messages m + JOIN agent_sessions s ON s.id = m.session_id + WHERE s.model LIKE ?1 + AND (?2 IS NULL OR datetime(m.timestamp) >= datetime(?2)) + AND (?3 IS NULL OR datetime(m.timestamp) < datetime(?3))", + params![GENERAL_MODE_PATTERN, from_text, to_text], + |row| row.get(0), + ) + .map_err(|e| format!("估算 unified general Token 失败: {e}"))?; + + let pending_chars = sum_pending_general_message_chars(conn, from_timestamp_ms, to_timestamp_ms) + .map_err(|e| format!("估算待迁移 general Token 失败: {e}"))?; + + Ok(unified_chars + pending_chars) +} + +fn query_non_general_session_count( + conn: &Connection, + from_timestamp_ms: Option, + to_timestamp_ms: Option, +) -> Result { + let from_text = from_timestamp_ms.map(format_sqlite_datetime); + let to_text = to_timestamp_ms.map(format_sqlite_datetime); + + conn.query_row( + "SELECT COUNT(*) + FROM agent_sessions s + WHERE s.model NOT LIKE ?1 + AND (?2 IS NULL OR datetime(s.created_at) >= datetime(?2)) + AND (?3 IS NULL OR datetime(s.created_at) < datetime(?3))", + params![GENERAL_MODE_PATTERN, from_text, to_text], + |row| row.get(0), + ) + .map_err(|e| format!("查询非通用 unified 会话数失败: {e}")) +} + +fn query_non_general_message_count( + conn: &Connection, + from_timestamp_ms: Option, + to_timestamp_ms: Option, +) -> Result { + let from_text = from_timestamp_ms.map(format_sqlite_datetime); + let to_text = to_timestamp_ms.map(format_sqlite_datetime); + + conn.query_row( + "SELECT COUNT(*) + FROM agent_messages m + JOIN agent_sessions s ON s.id = m.session_id + WHERE s.model NOT LIKE ?1 + AND (?2 IS NULL OR datetime(m.timestamp) >= datetime(?2)) + AND (?3 IS NULL OR datetime(m.timestamp) < datetime(?3))", + params![GENERAL_MODE_PATTERN, from_text, to_text], + |row| row.get(0), + ) + .map_err(|e| format!("查询非通用 unified 消息数失败: {e}")) +} + +fn sum_non_general_message_chars( + conn: &Connection, + from_timestamp_ms: Option, + to_timestamp_ms: Option, +) -> Result { + let from_text = from_timestamp_ms.map(format_sqlite_datetime); + let to_text = to_timestamp_ms.map(format_sqlite_datetime); + + conn.query_row( + "SELECT COALESCE(SUM(LENGTH(m.content_json)), 0) + FROM agent_messages m + JOIN agent_sessions s ON s.id = m.session_id + WHERE s.model NOT LIKE ?1 + AND (?2 IS NULL OR datetime(m.timestamp) >= datetime(?2)) + AND (?3 IS NULL OR datetime(m.timestamp) < datetime(?3))", + params![GENERAL_MODE_PATTERN, from_text, to_text], + |row| row.get(0), + ) + .map_err(|e| format!("估算非通用 unified Token 失败: {e}")) +} + /// 查询通用对话统计 fn query_general_chat_stats( conn: &Connection, today_start: &DateTime, month_start: &DateTime, ) -> Result { - // 转换为 Unix 时间戳(毫秒) let today_ts = today_start.timestamp_millis(); let month_ts = month_start.timestamp_millis(); - // 今日对话数 - let today_conversations: i64 = conn - .query_row( - "SELECT COUNT(*) FROM general_chat_sessions WHERE created_at >= ?", - [today_ts], - |row| row.get(0), - ) - .map_err(|e| format!("查询今日通用会话数失败: {e}"))?; - - // 今日消息数 - let today_messages: i64 = conn - .query_row( - "SELECT COUNT(*) FROM general_chat_messages WHERE created_at >= ?", - [today_ts], - |row| row.get(0), - ) - .map_err(|e| format!("查询今日通用消息数失败: {e}"))?; - - // 本月对话数 - let monthly_conversations: i64 = conn - .query_row( - "SELECT COUNT(*) FROM general_chat_sessions WHERE created_at >= ?", - [month_ts], - |row| row.get(0), - ) - .map_err(|e| format!("查询本月通用会话数失败: {e}"))?; - - // 本月消息数 - let monthly_messages: i64 = conn - .query_row( - "SELECT COUNT(*) FROM general_chat_messages WHERE created_at >= ?", - [month_ts], - |row| row.get(0), - ) - .map_err(|e| format!("查询本月通用消息数失败: {e}"))?; - - // 总对话数 - let total_conversations: i64 = conn - .query_row("SELECT COUNT(*) FROM general_chat_sessions", [], |row| { - row.get(0) - }) - .map_err(|e| format!("查询总通用会话数失败: {e}"))?; - - // 总消息数 - let total_messages: i64 = conn - .query_row("SELECT COUNT(*) FROM general_chat_messages", [], |row| { - row.get(0) - }) - .map_err(|e| format!("查询总通用消息数失败: {e}"))?; + let today_conversations = query_general_session_count(conn, Some(today_ts), None)?; + let today_messages = query_general_message_count(conn, Some(today_ts), None)?; + let monthly_conversations = query_general_session_count(conn, Some(month_ts), None)?; + let monthly_messages = query_general_message_count(conn, Some(month_ts), None)?; + let total_conversations = query_general_session_count(conn, None, None)?; + let total_messages = query_general_message_count(conn, None, None)?; Ok(ConversationStats { total_conversations: clamp_i64_to_u32(total_conversations), @@ -264,55 +376,15 @@ fn query_agent_chat_stats( today_start: &DateTime, month_start: &DateTime, ) -> Result { - // Agent sessions 使用 TEXT 格式的日期时间 - let today_str = today_start.format("%Y-%m-%d %H:%M:%S").to_string(); - let month_str = month_start.format("%Y-%m-%d %H:%M:%S").to_string(); + let today_ts = today_start.timestamp_millis(); + let month_ts = month_start.timestamp_millis(); - // 今日对话数 - let today_conversations: i64 = conn - .query_row( - "SELECT COUNT(*) FROM agent_sessions WHERE datetime(created_at) >= datetime(?)", - [today_str.clone()], - |row| row.get(0), - ) - .map_err(|e| format!("查询今日 Agent 会话数失败: {e}"))?; - - // 今日消息数 - let today_messages: i64 = conn - .query_row( - "SELECT COUNT(*) FROM agent_messages WHERE datetime(timestamp) >= datetime(?)", - [today_str], - |row| row.get(0), - ) - .map_err(|e| format!("查询今日 Agent 消息数失败: {e}"))?; - - // 本月对话数 - let monthly_conversations: i64 = conn - .query_row( - "SELECT COUNT(*) FROM agent_sessions WHERE datetime(created_at) >= datetime(?)", - [month_str.clone()], - |row| row.get(0), - ) - .map_err(|e| format!("查询本月 Agent 会话数失败: {e}"))?; - - // 本月消息数 - let monthly_messages: i64 = conn - .query_row( - "SELECT COUNT(*) FROM agent_messages WHERE datetime(timestamp) >= datetime(?)", - [month_str], - |row| row.get(0), - ) - .map_err(|e| format!("查询本月 Agent 消息数失败: {e}"))?; - - // 总对话数 - let total_conversations: i64 = conn - .query_row("SELECT COUNT(*) FROM agent_sessions", [], |row| row.get(0)) - .map_err(|e| format!("查询总 Agent 会话数失败: {e}"))?; - - // 总消息数 - let total_messages: i64 = conn - .query_row("SELECT COUNT(*) FROM agent_messages", [], |row| row.get(0)) - .map_err(|e| format!("查询总 Agent 消息数失败: {e}"))?; + let today_conversations = query_non_general_session_count(conn, Some(today_ts), None)?; + let today_messages = query_non_general_message_count(conn, Some(today_ts), None)?; + let monthly_conversations = query_non_general_session_count(conn, Some(month_ts), None)?; + let monthly_messages = query_non_general_message_count(conn, Some(month_ts), None)?; + let total_conversations = query_non_general_session_count(conn, None, None)?; + let total_messages = query_non_general_message_count(conn, None, None)?; Ok(ConversationStats { total_conversations: clamp_i64_to_u32(total_conversations), @@ -392,58 +464,14 @@ fn query_estimated_tokens_from_messages( ) -> Result { let today_ts = today_start.timestamp_millis(); let month_ts = month_start.timestamp_millis(); - let today_str = today_start.format("%Y-%m-%d %H:%M:%S").to_string(); - let month_str = month_start.format("%Y-%m-%d %H:%M:%S").to_string(); - let general_total_chars: i64 = conn - .query_row( - "SELECT COALESCE(SUM(LENGTH(content)), 0) FROM general_chat_messages", - [], - |row| row.get(0), - ) - .map_err(|e| format!("估算总 Token(通用消息)失败: {e}"))?; + let general_total_chars = sum_general_message_chars(conn, None, None)?; + let general_monthly_chars = sum_general_message_chars(conn, Some(month_ts), None)?; + let general_today_chars = sum_general_message_chars(conn, Some(today_ts), None)?; - let general_monthly_chars: i64 = conn - .query_row( - "SELECT COALESCE(SUM(LENGTH(content)), 0) FROM general_chat_messages WHERE created_at >= ?", - [month_ts], - |row| row.get(0), - ) - .map_err(|e| format!("估算本月 Token(通用消息)失败: {e}"))?; - - let general_today_chars: i64 = conn - .query_row( - "SELECT COALESCE(SUM(LENGTH(content)), 0) FROM general_chat_messages WHERE created_at >= ?", - [today_ts], - |row| row.get(0), - ) - .map_err(|e| format!("估算今日 Token(通用消息)失败: {e}"))?; - - let agent_total_chars: i64 = conn - .query_row( - "SELECT COALESCE(SUM(LENGTH(content_json)), 0) FROM agent_messages", - [], - |row| row.get(0), - ) - .map_err(|e| format!("估算总 Token(Agent 消息)失败: {e}"))?; - - let agent_monthly_chars: i64 = conn - .query_row( - "SELECT COALESCE(SUM(LENGTH(content_json)), 0) FROM agent_messages - WHERE datetime(timestamp) >= datetime(?)", - [month_str], - |row| row.get(0), - ) - .map_err(|e| format!("估算本月 Token(Agent 消息)失败: {e}"))?; - - let agent_today_chars: i64 = conn - .query_row( - "SELECT COALESCE(SUM(LENGTH(content_json)), 0) FROM agent_messages - WHERE datetime(timestamp) >= datetime(?)", - [today_str], - |row| row.get(0), - ) - .map_err(|e| format!("估算今日 Token(Agent 消息)失败: {e}"))?; + let agent_total_chars = sum_non_general_message_chars(conn, None, None)?; + let agent_monthly_chars = sum_non_general_message_chars(conn, Some(month_ts), None)?; + let agent_today_chars = sum_non_general_message_chars(conn, Some(today_ts), None)?; Ok(TokenStats { total_tokens: chars_to_estimated_tokens(general_total_chars + agent_total_chars), @@ -555,7 +583,8 @@ fn query_model_usage_from_agent_messages( COALESCE(SUM(LENGTH(m.content_json)), 0) AS content_chars FROM agent_messages m JOIN agent_sessions s ON s.id = m.session_id - WHERE datetime(m.timestamp) >= datetime(?) + WHERE s.model NOT LIKE ?1 + AND datetime(m.timestamp) >= datetime(?2) GROUP BY s.model ORDER BY content_chars DESC, conversations DESC LIMIT 20", @@ -563,7 +592,7 @@ fn query_model_usage_from_agent_messages( .map_err(|e| format!("准备 Agent 模型排行查询失败: {e}"))?; let rows = stmt - .query_map([start_str], |row| { + .query_map(params![GENERAL_MODE_PATTERN, start_str], |row| { let model: String = row.get(0)?; let conversations: i64 = row.get(1)?; let chars: i64 = row.get(2)?; @@ -589,6 +618,7 @@ fn query_model_usage_from_agent_messages( COALESCE(SUM(LENGTH(m.content_json)), 0) AS content_chars FROM agent_messages m JOIN agent_sessions s ON s.id = m.session_id + WHERE s.model NOT LIKE ?1 GROUP BY s.model ORDER BY content_chars DESC, conversations DESC LIMIT 20", @@ -596,7 +626,7 @@ fn query_model_usage_from_agent_messages( .map_err(|e| format!("准备 Agent 模型排行查询失败: {e}"))?; let rows = stmt - .query_map([], |row| { + .query_map([GENERAL_MODE_PATTERN], |row| { let model: String = row.get(0)?; let conversations: i64 = row.get(1)?; let chars: i64 = row.get(2)?; @@ -674,31 +704,16 @@ pub fn get_daily_usage_trends_from_db( let day_start = start_of_day(date); let day_end = day_start + Duration::days(1); - // 当天开始/结束(时间戳 + 文本) let day_start_ts = day_start.timestamp_millis(); let day_end_ts = day_end.timestamp_millis(); - let day_start_str = day_start.format("%Y-%m-%d %H:%M:%S").to_string(); - let day_end_str = day_end.format("%Y-%m-%d %H:%M:%S").to_string(); let day_key = day_start.format("%Y-%m-%d").to_string(); - let conversations: i64 = conn - .query_row( - "SELECT COUNT(*) FROM general_chat_sessions WHERE created_at >= ? AND created_at < ?", - [day_start_ts, day_end_ts], - |row| row.get(0), - ) + let conversations = query_general_session_count(conn, Some(day_start_ts), Some(day_end_ts)) .map_err(|e| format!("查询通用会话日统计失败: {e}"))?; - // 查询 Agent 对话 - let agent_conversations: i64 = conn - .query_row( - "SELECT COUNT(*) FROM agent_sessions - WHERE datetime(created_at) >= datetime(?) - AND datetime(created_at) < datetime(?)", - [day_start_str.clone(), day_end_str.clone()], - |row| row.get(0), - ) - .map_err(|e| format!("查询 Agent 会话日统计失败: {e}"))?; + let agent_conversations = + query_non_general_session_count(conn, Some(day_start_ts), Some(day_end_ts)) + .map_err(|e| format!("查询 Agent 会话日统计失败: {e}"))?; let total_conversations = conversations + agent_conversations; @@ -713,26 +728,13 @@ pub fn get_daily_usage_trends_from_db( clamp_i64_to_u64(day_tokens) } else { - let general_chars: i64 = conn - .query_row( - "SELECT COALESCE(SUM(LENGTH(content)), 0) - FROM general_chat_messages - WHERE created_at >= ? AND created_at < ?", - [day_start_ts, day_end_ts], - |row| row.get(0), - ) - .map_err(|e| format!("估算通用消息日 Token 失败: {e}"))?; + let general_chars = + sum_general_message_chars(conn, Some(day_start_ts), Some(day_end_ts)) + .map_err(|e| format!("估算通用消息日 Token 失败: {e}"))?; - let agent_chars: i64 = conn - .query_row( - "SELECT COALESCE(SUM(LENGTH(content_json)), 0) - FROM agent_messages - WHERE datetime(timestamp) >= datetime(?) - AND datetime(timestamp) < datetime(?)", - [day_start_str, day_end_str], - |row| row.get(0), - ) - .map_err(|e| format!("估算 Agent 消息日 Token 失败: {e}"))?; + let agent_chars = + sum_non_general_message_chars(conn, Some(day_start_ts), Some(day_end_ts)) + .map_err(|e| format!("估算 Agent 消息日 Token 失败: {e}"))?; chars_to_estimated_tokens(general_chars + agent_chars) }; @@ -746,3 +748,153 @@ pub fn get_daily_usage_trends_from_db( Ok(daily_usage) } + +#[cfg(test)] +mod tests { + use super::{query_agent_chat_stats, query_general_chat_stats, start_of_day, start_of_month}; + use chrono::{Local, TimeZone}; + use rusqlite::{params, Connection}; + + fn create_test_schema(conn: &Connection) { + conn.execute_batch( + " + 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 + ); + CREATE TABLE settings ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL + ); + 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 + ); + 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 + ); + ", + ) + .expect("create schema"); + } + + #[test] + fn stats_do_not_double_count_migrated_general_sessions() { + let conn = Connection::open_in_memory().expect("open in memory db"); + create_test_schema(&conn); + + let now = Local + .with_ymd_and_hms(2026, 3, 12, 10, 0, 0) + .single() + .expect("build datetime"); + let now_ms = now.timestamp_millis(); + + conn.execute( + "INSERT INTO general_chat_sessions (id, name, created_at, updated_at) VALUES (?1, ?2, ?3, ?4)", + params!["general-1", "旧通用会话", now_ms, now_ms], + ) + .unwrap(); + conn.execute( + "INSERT INTO general_chat_messages (id, session_id, role, content, created_at) VALUES (?1, ?2, ?3, ?4, ?5)", + params!["gm-1", "general-1", "user", "legacy general", now_ms], + ) + .unwrap(); + + conn.execute( + "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at) VALUES (?1, ?2, NULL, ?3, ?4, ?5)", + params!["general-1", "general:default", "统一通用会话", now.to_rfc3339(), now.to_rfc3339()], + ) + .unwrap(); + conn.execute( + "INSERT INTO agent_messages (session_id, role, content_json, timestamp) VALUES (?1, ?2, ?3, ?4)", + params!["general-1", "user", r#"[{"type":"text","text":"legacy general"}]"#, now.to_rfc3339()], + ) + .unwrap(); + + conn.execute( + "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at) VALUES (?1, ?2, NULL, ?3, ?4, ?5)", + params!["agent-1", "claude-sonnet-4", "Agent 会话", now.to_rfc3339(), now.to_rfc3339()], + ) + .unwrap(); + conn.execute( + "INSERT INTO agent_messages (session_id, role, content_json, timestamp) VALUES (?1, ?2, ?3, ?4)", + params!["agent-1", "assistant", r#"[{"type":"text","text":"agent reply"}]"#, now.to_rfc3339()], + ) + .unwrap(); + + let today_start = start_of_day(now); + let month_start = start_of_month(now); + + let general_stats = + query_general_chat_stats(&conn, &today_start, &month_start).expect("general stats"); + let agent_stats = + query_agent_chat_stats(&conn, &today_start, &month_start).expect("agent stats"); + + assert_eq!(general_stats.total_conversations, 1); + assert_eq!(general_stats.total_messages, 1); + assert_eq!(agent_stats.total_conversations, 1); + assert_eq!(agent_stats.total_messages, 1); + } + + #[test] + fn stats_ignore_legacy_general_after_migration_completed() { + let conn = Connection::open_in_memory().expect("open in memory db"); + create_test_schema(&conn); + + let now = Local + .with_ymd_and_hms(2026, 3, 12, 10, 0, 0) + .single() + .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], + ) + .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 general", now_ms], + ) + .unwrap(); + + let today_start = start_of_day(now); + let month_start = start_of_month(now); + let general_stats = + query_general_chat_stats(&conn, &today_start, &month_start).expect("general stats"); + + assert_eq!(general_stats.total_conversations, 0); + assert_eq!(general_stats.total_messages, 0); + assert_eq!(general_stats.monthly_conversations, 0); + assert_eq!(general_stats.today_messages, 0); + } +} diff --git a/src-tauri/src/services/execution_tracker_service.rs b/src-tauri/src/services/execution_tracker_service.rs index e7f6127f9..cc86f7625 100644 --- a/src-tauri/src/services/execution_tracker_service.rs +++ b/src-tauri/src/services/execution_tracker_service.rs @@ -206,6 +206,17 @@ impl ExecutionTracker { .map_err(|e| format!("查询会话执行记录失败: {e}")) } + pub fn list_terminal_runs_by_session( + &self, + session_id: &str, + limit: usize, + offset: usize, + ) -> Result, String> { + let conn = self.db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; + AgentRunDao::list_terminal_runs_by_session(&conn, session_id, limit, offset) + .map_err(|e| format!("查询会话终态执行记录失败: {e}")) + } + pub async fn with_run( &self, source: RunSource, diff --git a/src-tauri/src/services/memory_source_resolver_service.rs b/src-tauri/src/services/memory_source_resolver_service.rs index 9c74f8010..082593983 100644 --- a/src-tauri/src/services/memory_source_resolver_service.rs +++ b/src-tauri/src/services/memory_source_resolver_service.rs @@ -8,6 +8,7 @@ use crate::services::memory_rules_loader_service::load_rules; use proxycast_agent::{ resolve_durable_memory_root, to_virtual_memory_path, DURABLE_MEMORY_VIRTUAL_ROOT, }; +use proxycast_core::app_paths; use proxycast_core::config::{Config, MemoryConfig}; use serde::{Deserialize, Serialize}; use std::collections::HashSet; @@ -758,10 +759,8 @@ fn expand_path(path: &str, working_dir: Option<&Path>) -> PathBuf { } fn default_user_memory_path() -> PathBuf { - dirs::home_dir() - .unwrap_or_else(|| PathBuf::from(".")) - .join(".proxycast") - .join("AGENTS.md") + app_paths::resolve_user_memory_path() + .unwrap_or_else(|_| app_paths::best_effort_app_data_file("AGENTS.md")) } fn default_managed_policy_path() -> PathBuf { diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index 688411489..3247213f2 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -4,7 +4,9 @@ //! 本模块保留 Tauri 相关服务。 // 保留在主 crate 的 Tauri 相关服务 +pub mod agent_timeline_service; pub mod auto_memory_service; +pub mod chat_history_service; pub mod conversation_statistics_service; pub mod environment_service; pub mod execution_tracker_service; diff --git a/src-tauri/src/services/openclaw_service.rs b/src-tauri/src/services/openclaw_service.rs index 964b3e6c3..f6bea657a 100644 --- a/src-tauri/src/services/openclaw_service.rs +++ b/src-tauri/src/services/openclaw_service.rs @@ -105,9 +105,25 @@ pub struct EnvironmentStatus { pub recommended_action: String, pub summary: String, #[serde(default)] + pub diagnostics: EnvironmentDiagnostics, + #[serde(default)] pub temp_artifacts: Vec, } +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +#[serde(rename_all = "camelCase")] +pub struct EnvironmentDiagnostics { + pub npm_path: Option, + pub npm_global_prefix: Option, + pub openclaw_package_path: Option, + #[serde(default)] + pub where_candidates: Vec, + #[serde(default)] + pub supplemental_search_dirs: Vec, + #[serde(default)] + pub supplemental_command_candidates: Vec, +} + #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct CommandPreview { @@ -253,8 +269,9 @@ impl OpenClawService { let node = inspect_node_dependency_status().await?; let git = inspect_git_dependency_status().await?; let openclaw = inspect_openclaw_dependency_status().await?; + let diagnostics = collect_environment_diagnostics().await; - Ok(build_environment_status(node, git, openclaw)) + Ok(build_environment_status(node, git, openclaw, diagnostics)) } pub async fn check_installed(&self) -> Result { @@ -312,6 +329,16 @@ impl OpenClawService { pub async fn install(&mut self, app: &AppHandle) -> Result { emit_install_progress(app, "开始准备 OpenClaw 环境。", "info"); + #[cfg(target_os = "windows")] + { + let node_status = self.inspect_dependency_status(DependencyKind::Node).await?; + let git_status = self.inspect_dependency_status(DependencyKind::Git).await?; + if let Some(result) = windows_install_block_result(&node_status, &git_status) { + emit_install_progress(app, &result.message, "warn"); + return Ok(result); + } + } + let node_result = self .ensure_dependency_ready(app, DependencyKind::Node) .await?; @@ -383,6 +410,34 @@ impl OpenClawService { _ => return Err(format!("不支持的依赖类型: {kind}")), }; + #[cfg(target_os = "windows")] + { + let status = self.inspect_dependency_status(dependency).await?; + if status.status == "ok" { + emit_install_progress( + app, + &format!( + "{} 已就绪{}。", + dependency.label(), + status + .version + .as_deref() + .map(|version| format!(" · {version}")) + .unwrap_or_default() + ), + "info", + ); + return Ok(ActionResult { + success: true, + message: format!("{} 已满足要求。", dependency.label()), + }); + } + + let result = windows_dependency_action_result(dependency, &status); + emit_install_progress(app, &result.message, "warn"); + return Ok(result); + } + self.ensure_dependency_ready(app, dependency).await } @@ -1603,6 +1658,91 @@ fn openclaw_proxycast_config_path() -> PathBuf { openclaw_config_dir().join("openclaw.proxycast.json") } +#[cfg(any(target_os = "windows", test))] +fn windows_dependency_setup_message( + dependency: DependencyKind, + status: &DependencyStatus, +) -> String { + let guidance = match dependency { + DependencyKind::Node => format!( + "Windows 下请先从 nodejs.org 安装或升级 Node.js {}+,完成后点击“重新检测”,再安装 OpenClaw。", + NODE_MIN_VERSION.0 + ), + DependencyKind::Git => { + "Windows 下请先从 git-scm.com 安装 Git(安装时请勾选加入 PATH),完成后点击“重新检测”,再安装 OpenClaw。" + .to_string() + } + }; + + format!("{} {}", status.message, guidance) +} + +#[cfg(any(target_os = "windows", test))] +fn windows_dependency_action_result( + dependency: DependencyKind, + status: &DependencyStatus, +) -> ActionResult { + ActionResult { + success: false, + message: windows_dependency_setup_message(dependency, status), + } +} + +#[cfg(any(target_os = "windows", test))] +fn windows_install_block_result( + node_status: &DependencyStatus, + git_status: &DependencyStatus, +) -> Option { + if node_status.status != "ok" { + return Some(windows_dependency_action_result( + DependencyKind::Node, + node_status, + )); + } + + if git_status.status != "ok" { + return Some(windows_dependency_action_result( + DependencyKind::Git, + git_status, + )); + } + + None +} + +fn dependency_setup_summary(dependency: DependencyKind) -> String { + if cfg!(target_os = "windows") { + return match dependency { + DependencyKind::Node => format!( + "当前缺少可用的 Node.js {}+ 运行时,Windows 下请先手动安装 Node.js,完成后点击“重新检测”,再安装 OpenClaw。", + NODE_MIN_VERSION.0 + ), + DependencyKind::Git => { + "当前缺少可用的 Git,Windows 下请先手动安装 Git(安装时请勾选加入 PATH),完成后点击“重新检测”,再安装 OpenClaw。" + .to_string() + } + }; + } + + if cfg!(target_os = "macos") { + return match dependency { + DependencyKind::Node => format!( + "当前缺少可用的 Node.js {}+ 运行时,建议先一键安装或修复 Node.js。", + format_semver(NODE_MIN_VERSION) + ), + DependencyKind::Git => "当前缺少可用的 Git,建议先一键安装或修复 Git。".to_string(), + }; + } + + match dependency { + DependencyKind::Node => format!( + "当前缺少可用的 Node.js {}+ 运行时,请先手动安装后重新检测。", + format_semver(NODE_MIN_VERSION) + ), + DependencyKind::Git => "当前缺少可用的 Git,请先手动安装后重新检测。".to_string(), + } +} + fn openclaw_installer_download_dir(app: &AppHandle) -> Result { let _ = app; let app_data_dir = proxycast_core::app_paths::preferred_data_dir() @@ -1633,6 +1773,7 @@ fn build_environment_status( node: DependencyStatus, git: DependencyStatus, mut openclaw: DependencyStatus, + diagnostics: EnvironmentDiagnostics, ) -> EnvironmentStatus { let node_ready = node.status == "ok"; let git_ready = git.status == "ok"; @@ -1641,12 +1782,18 @@ fn build_environment_status( let (recommended_action, summary) = if !node_ready { ( "install_node".to_string(), - "当前缺少可用的 Node.js 22+ 运行时,建议先一键安装或修复 Node.js。".to_string(), + dependency_setup_summary(DependencyKind::Node), ) } else if !git_ready { ( "install_git".to_string(), - "当前缺少可用的 Git,建议先一键安装或修复 Git。".to_string(), + dependency_setup_summary(DependencyKind::Git), + ) + } else if openclaw.status == "needs_reload" { + ( + "refresh_openclaw_env".to_string(), + "已检测到 OpenClaw 包,但命令尚未生效;请点击“重新检测”,必要时重启 ProxyCast。" + .to_string(), ) } else if openclaw.status != "ok" { ( @@ -1666,6 +1813,7 @@ fn build_environment_status( openclaw, recommended_action, summary, + diagnostics, temp_artifacts: collect_temp_artifact_paths(None) .into_iter() .filter(|path| path.exists()) @@ -1684,7 +1832,7 @@ async fn inspect_node_dependency_status() -> Result { "未检测到 Node.js,需要安装 {}+。", format_semver(NODE_MIN_VERSION) ), - auto_install_supported: cfg!(target_os = "windows") || cfg!(target_os = "macos"), + auto_install_supported: cfg!(target_os = "macos"), }); }; @@ -1698,7 +1846,7 @@ async fn inspect_node_dependency_status() -> Result { "检测到 Node.js,但无法识别版本:{version_text}。请安装 {}+。", format_semver(NODE_MIN_VERSION) ), - auto_install_supported: cfg!(target_os = "windows") || cfg!(target_os = "macos"), + auto_install_supported: cfg!(target_os = "macos"), }); }; @@ -1709,7 +1857,7 @@ async fn inspect_node_dependency_status() -> Result { version: Some(normalized.clone()), path: Some(path), message: format!("Node.js 已就绪:{normalized}"), - auto_install_supported: cfg!(target_os = "windows") || cfg!(target_os = "macos"), + auto_install_supported: cfg!(target_os = "macos"), }) } else { Ok(DependencyStatus { @@ -1720,7 +1868,7 @@ async fn inspect_node_dependency_status() -> Result { "Node.js 版本过低:{normalized},需要 {}+。", format_semver(NODE_MIN_VERSION) ), - auto_install_supported: cfg!(target_os = "windows") || cfg!(target_os = "macos"), + auto_install_supported: cfg!(target_os = "macos"), }) } } @@ -1751,6 +1899,10 @@ async fn inspect_git_dependency_status() -> Result { async fn inspect_openclaw_dependency_status() -> Result { let Some(path) = find_command_in_shell("openclaw").await? else { + if let Some(status) = inspect_openclaw_package_reload_status().await? { + return Ok(status); + } + return Ok(DependencyStatus { status: "missing".to_string(), version: None, @@ -1778,18 +1930,41 @@ async fn inspect_openclaw_dependency_status() -> Result Result { - #[cfg(target_os = "windows")] - { - Ok(find_command_in_shell("winget").await?.is_some()) - } +async fn inspect_openclaw_package_reload_status() -> Result, String> { + let Some(npm_path) = find_command_in_standard_locations("npm").await? else { + return Ok(None); + }; + let Some(prefix) = detect_npm_global_prefix(&npm_path).await else { + return Ok(None); + }; + let Some(package) = find_installed_openclaw_package_details(&prefix) else { + return Ok(None); + }; + let version_suffix = package + .version + .as_deref() + .map(|item| format!("({item})")) + .unwrap_or_default(); + + Ok(Some(DependencyStatus { + status: "needs_reload".to_string(), + version: package.version.clone(), + path: Some(prefix.clone()), + message: format!( + "已在 npm 全局目录检测到 {}{},但当前进程尚未解析到 openclaw 命令。请点击“重新检测”;若仍失败,请重启 ProxyCast,或确认 {prefix} 已加入 PATH。", package.name, version_suffix + ), + auto_install_supported: false, + })) +} + +async fn git_auto_install_supported() -> Result { #[cfg(target_os = "macos")] { Ok(true) } - #[cfg(not(any(target_os = "windows", target_os = "macos")))] + #[cfg(not(target_os = "macos"))] { Ok(false) } @@ -2426,6 +2601,32 @@ fn apply_windows_no_window(_command: &mut Command) { } async fn find_command_in_shell(command_name: &str) -> Result, String> { + let mut candidates = collect_standard_command_candidates(command_name).await?; + + if command_name == "openclaw" { + candidates.extend(find_commands_via_npm_global_prefix(command_name).await?); + } + + Ok(select_command_path(command_name, candidates) + .await? + .map(|path| path.to_string_lossy().to_string())) +} + +async fn find_command_in_standard_locations(command_name: &str) -> Result, String> { + Ok(select_command_path( + command_name, + collect_standard_command_candidates(command_name).await?, + ) + .await? + .map(|path| path.to_string_lossy().to_string())) +} + +async fn collect_standard_command_candidates(command_name: &str) -> Result, String> { + #[cfg(target_os = "windows")] + { + let _ = refresh_windows_path_from_registry(); + } + let mut candidates = Vec::new(); #[cfg(target_os = "windows")] @@ -2435,6 +2636,13 @@ async fn find_command_in_shell(command_name: &str) -> Result, Str candidates.extend(find_all_commands_in_known_locations(command_name)); + Ok(candidates) +} + +async fn select_command_path( + command_name: &str, + candidates: Vec, +) -> Result, String> { let mut deduped = Vec::with_capacity(candidates.len()); let mut seen = HashSet::new(); for candidate in candidates { @@ -2443,9 +2651,7 @@ async fn find_command_in_shell(command_name: &str) -> Result, Str } } - Ok(select_command_candidate(command_name, deduped) - .await? - .map(|path| path.to_string_lossy().to_string())) + select_command_candidate(command_name, deduped).await } #[cfg(target_os = "windows")] @@ -2494,6 +2700,11 @@ async fn select_command_candidate( } fn find_all_commands_in_known_locations(command_name: &str) -> Vec { + let search_dirs = collect_known_command_search_dirs(); + find_all_commands_in_paths(command_name, &search_dirs) +} + +fn collect_known_command_search_dirs() -> Vec { let mut search_dirs = Vec::new(); let mut seen = HashSet::new(); @@ -2536,6 +2747,13 @@ fn find_all_commands_in_known_locations(command_name: &str) -> Vec { } } + #[cfg(target_os = "windows")] + { + for dir in windows_known_command_dirs_from_env() { + push_dir(dir); + } + } + if cfg!(target_os = "macos") { push_dir(PathBuf::from("/opt/homebrew/bin")); push_dir(PathBuf::from("/usr/local/bin")); @@ -2543,7 +2761,42 @@ fn find_all_commands_in_known_locations(command_name: &str) -> Vec { push_dir(PathBuf::from("/bin")); } - find_all_commands_in_paths(command_name, &search_dirs) + search_dirs +} + +#[cfg(target_os = "windows")] +fn windows_known_command_dirs_from_env() -> Vec { + let mut dirs = Vec::new(); + + if let Some(appdata) = std::env::var_os("APPDATA") { + dirs.push(PathBuf::from(appdata).join("npm")); + } + + if let Some(localappdata) = std::env::var_os("LOCALAPPDATA") { + let localappdata = PathBuf::from(localappdata); + dirs.push(localappdata.join("Programs").join("nodejs")); + dirs.push(localappdata.join("Volta").join("bin")); + } + + if let Some(program_files) = std::env::var_os("ProgramFiles") { + dirs.push(PathBuf::from(program_files).join("nodejs")); + } + + if let Some(program_files_x86) = std::env::var_os("ProgramFiles(x86)") { + dirs.push(PathBuf::from(program_files_x86).join("nodejs")); + } + + if let Some(home) = home_dir() { + dirs.push(home.join("AppData").join("Roaming").join("npm")); + dirs.push( + home.join("AppData") + .join("Local") + .join("Programs") + .join("nodejs"), + ); + } + + dirs } fn find_all_commands_in_paths(command_name: &str, search_dirs: &[PathBuf]) -> Vec { @@ -2572,6 +2825,170 @@ fn find_all_commands_in_paths(command_name: &str, search_dirs: &[PathBuf]) -> Ve matches } +async fn find_commands_via_npm_global_prefix(command_name: &str) -> Result, String> { + let Some(npm_path) = find_command_in_standard_locations("npm").await? else { + return Ok(Vec::new()); + }; + let Some(prefix) = detect_npm_global_prefix(&npm_path).await else { + return Ok(Vec::new()); + }; + + Ok(find_all_commands_in_paths( + command_name, + &npm_global_command_dirs(&prefix), + )) +} + +fn npm_global_command_dirs(prefix: &str) -> Vec { + npm_global_command_dirs_for(current_shell_platform(), prefix) +} + +fn npm_global_command_dirs_for(platform: ShellPlatform, prefix: &str) -> Vec { + let prefix_path = PathBuf::from(prefix); + + match platform { + ShellPlatform::Windows => vec![prefix_path], + ShellPlatform::Unix => vec![prefix_path.join("bin"), prefix_path], + } +} + +fn npm_global_node_modules_dirs_for(platform: ShellPlatform, prefix: &str) -> Vec { + let prefix_path = PathBuf::from(prefix); + + match platform { + ShellPlatform::Windows => vec![prefix_path.join("node_modules")], + ShellPlatform::Unix => vec![ + prefix_path.join("lib").join("node_modules"), + prefix_path.join("node_modules"), + ], + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct InstalledOpenClawPackage { + name: &'static str, + version: Option, + path: PathBuf, +} + +#[cfg(test)] +fn find_installed_openclaw_package(prefix: &str) -> Option<(&'static str, Option)> { + find_installed_openclaw_package_details(prefix).map(|package| (package.name, package.version)) +} + +fn find_installed_openclaw_package_details(prefix: &str) -> Option { + for node_modules_dir in npm_global_node_modules_dirs_for(current_shell_platform(), prefix) { + let openclaw_manifest = node_modules_dir.join("openclaw").join("package.json"); + if openclaw_manifest.is_file() { + return Some(InstalledOpenClawPackage { + name: "openclaw", + version: read_package_version(&openclaw_manifest), + path: openclaw_manifest, + }); + } + + let zh_manifest = node_modules_dir + .join("@qingchencloud") + .join("openclaw-zh") + .join("package.json"); + if zh_manifest.is_file() { + return Some(InstalledOpenClawPackage { + name: "@qingchencloud/openclaw-zh", + version: read_package_version(&zh_manifest), + path: zh_manifest, + }); + } + } + + None +} + +fn read_package_version(manifest_path: &Path) -> Option { + #[derive(Deserialize)] + struct PackageManifest { + version: Option, + } + + let content = std::fs::read_to_string(manifest_path).ok()?; + let manifest = serde_json::from_str::(&content).ok()?; + manifest.version.filter(|item| !item.trim().is_empty()) +} + +async fn collect_environment_diagnostics() -> EnvironmentDiagnostics { + let npm_path = find_command_in_standard_locations("npm") + .await + .ok() + .flatten(); + let npm_global_prefix = match npm_path.as_deref() { + Some(path) => detect_npm_global_prefix(path).await, + None => None, + }; + + #[cfg(target_os = "windows")] + let where_candidates = find_commands_via_where("openclaw") + .await + .unwrap_or_default() + .into_iter() + .map(|path| path.display().to_string()) + .collect(); + + #[cfg(not(target_os = "windows"))] + let where_candidates = Vec::new(); + + let supplemental_search_dirs = + collect_supplemental_openclaw_search_dirs(npm_global_prefix.as_deref()); + let supplemental_command_candidates = + find_all_commands_in_paths("openclaw", &supplemental_search_dirs) + .into_iter() + .map(|path| path.display().to_string()) + .collect(); + let openclaw_package_path = npm_global_prefix + .as_deref() + .and_then(find_installed_openclaw_package_details) + .map(|package| package.path.display().to_string()); + + EnvironmentDiagnostics { + npm_path, + npm_global_prefix, + openclaw_package_path, + where_candidates, + supplemental_search_dirs: supplemental_search_dirs + .into_iter() + .map(|path| path.display().to_string()) + .collect(), + 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); + } + } + + if let Some(prefix) = npm_global_prefix { + for dir in npm_global_command_dirs(prefix) { + push_dir(dir); + } + } + + dirs +} + async fn select_best_node_candidate(candidates: Vec) -> Result, String> { let mut versioned = Vec::with_capacity(candidates.len()); for candidate in candidates { @@ -2861,17 +3278,21 @@ mod tests { use super::{ apply_gateway_runtime_defaults, build_environment_status, build_openclaw_cleanup_command, build_openclaw_install_command, build_winget_install_command, command_bin_dir_for, - determine_api_type, extract_gateway_auth_token, format_gateway_start_failure_message, - format_provider_base_url, gateway_start_args, has_api_version, parse_semver_from_text, - resolve_windows_dependency_install_plan, select_best_semver_candidate, - select_preferred_path_candidate, shell_command_escape_for, shell_npm_prefix_assignment_for, - shell_path_assignment_for, trim_trailing_slash, windows_manual_install_message, - DependencyKind, DependencyStatus, ShellPlatform, WindowsDependencyInstallPlan, + determine_api_type, extract_gateway_auth_token, find_installed_openclaw_package, + format_gateway_start_failure_message, format_provider_base_url, gateway_start_args, + has_api_version, npm_global_command_dirs_for, npm_global_node_modules_dirs_for, + parse_semver_from_text, resolve_windows_dependency_install_plan, + select_best_semver_candidate, 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, ShellPlatform, WindowsDependencyInstallPlan, NPM_MIRROR_CN, OPENCLAW_CN_PACKAGE, OPENCLAW_DEFAULT_PACKAGE, }; use crate::database::dao::api_key_provider::{ApiKeyProvider, ApiProviderType, ProviderGroup}; use chrono::Utc; use serde_json::{json, Value}; + use std::fs; use std::path::PathBuf; fn build_provider(provider_type: ApiProviderType, api_host: &str) -> ApiKeyProvider { @@ -3108,12 +3529,44 @@ mod tests { message: "openclaw missing".to_string(), auto_install_supported: false, }, + EnvironmentDiagnostics::default(), ); assert_eq!(env.recommended_action, "install_node"); assert_eq!(env.openclaw.auto_install_supported, false); } + #[test] + fn environment_status_uses_reload_summary_when_openclaw_command_not_ready() { + let env = build_environment_status( + DependencyStatus { + status: "ok".to_string(), + version: Some("22.0.0".to_string()), + path: Some("/usr/local/bin/node".to_string()), + message: "node ok".to_string(), + auto_install_supported: true, + }, + DependencyStatus { + status: "ok".to_string(), + version: Some("2.44.0".to_string()), + path: Some("/usr/bin/git".to_string()), + message: "git ok".to_string(), + auto_install_supported: true, + }, + DependencyStatus { + status: "needs_reload".to_string(), + version: Some("0.3.0".to_string()), + path: Some("/mock/prefix".to_string()), + message: "reload openclaw".to_string(), + auto_install_supported: false, + }, + EnvironmentDiagnostics::default(), + ); + + assert_eq!(env.recommended_action, "refresh_openclaw_env"); + assert!(env.summary.contains("重新检测")); + } + #[test] fn semver_selection_prefers_windows_launcher_over_bare_file_when_versions_equal() { let preferred = select_best_semver_candidate(vec![ @@ -3244,6 +3697,58 @@ mod tests { ); } + #[test] + fn windows_npm_global_command_dirs_use_prefix_root() { + assert_eq!( + npm_global_command_dirs_for( + ShellPlatform::Windows, + r"C:\Users\demo\AppData\Roaming\npm" + ), + vec![PathBuf::from(r"C:\Users\demo\AppData\Roaming\npm")] + ); + } + + #[test] + fn unix_npm_global_command_dirs_include_bin_directory() { + assert_eq!( + npm_global_command_dirs_for(ShellPlatform::Unix, "/Users/demo/.npm-global"), + vec![ + PathBuf::from("/Users/demo/.npm-global/bin"), + PathBuf::from("/Users/demo/.npm-global") + ] + ); + } + + #[test] + fn windows_npm_global_node_modules_dirs_use_prefix_node_modules() { + assert_eq!( + npm_global_node_modules_dirs_for( + ShellPlatform::Windows, + r"C:\Users\demo\AppData\Roaming\npm" + ), + vec![PathBuf::from(r"C:\Users\demo\AppData\Roaming\npm").join("node_modules")] + ); + } + + #[test] + fn finds_openclaw_package_from_global_npm_prefix() { + let temp_dir = + std::env::temp_dir().join(format!("proxycast-openclaw-test-{}", std::process::id())); + let package_dir = temp_dir.join("node_modules").join("openclaw"); + fs::create_dir_all(&package_dir).unwrap(); + fs::write( + package_dir.join("package.json"), + r#"{"name":"openclaw","version":"0.4.1"}"#, + ) + .unwrap(); + + let detected = find_installed_openclaw_package(temp_dir.to_str().unwrap()); + + fs::remove_dir_all(&temp_dir).unwrap(); + + assert_eq!(detected, Some(("openclaw", Some("0.4.1".to_string())))); + } + #[test] fn windows_node_prefers_winget_when_available() { assert_eq!( @@ -3284,6 +3789,104 @@ mod tests { ); } + #[test] + fn windows_git_setup_message_points_to_manual_download() { + let message = windows_dependency_setup_message( + DependencyKind::Git, + &DependencyStatus { + status: "missing".to_string(), + version: None, + path: None, + message: "未检测到 Git。".to_string(), + auto_install_supported: false, + }, + ); + + assert!(message.contains("git-scm.com")); + assert!(message.contains("加入 PATH")); + } + + #[test] + fn windows_node_setup_message_points_to_nodejs_download() { + let message = windows_dependency_setup_message( + DependencyKind::Node, + &DependencyStatus { + status: "missing".to_string(), + version: None, + path: None, + message: "未检测到 Node.js,需要安装 22.0.0+。".to_string(), + auto_install_supported: false, + }, + ); + + assert!(message.contains("nodejs.org")); + assert!(message.contains("Node.js 22+")); + } + + #[test] + fn windows_dependency_action_result_returns_failure_message() { + let result = windows_dependency_action_result( + DependencyKind::Git, + &DependencyStatus { + status: "missing".to_string(), + version: None, + path: None, + message: "未检测到 Git。".to_string(), + auto_install_supported: false, + }, + ); + + assert!(!result.success); + assert!(result.message.contains("git-scm.com")); + } + + #[test] + fn windows_install_block_result_prioritizes_node_before_git() { + let result = windows_install_block_result( + &DependencyStatus { + status: "missing".to_string(), + version: None, + path: None, + message: "未检测到 Node.js,需要安装 22.0.0+。".to_string(), + auto_install_supported: false, + }, + &DependencyStatus { + status: "missing".to_string(), + version: None, + path: None, + message: "未检测到 Git。".to_string(), + auto_install_supported: false, + }, + ) + .expect("应返回 Windows 阻断结果"); + + assert!(!result.success); + assert!(result.message.contains("nodejs.org")); + assert!(!result.message.contains("git-scm.com")); + } + + #[test] + fn windows_install_block_result_returns_none_when_dependencies_ready() { + let result = windows_install_block_result( + &DependencyStatus { + status: "ok".to_string(), + version: Some("22.0.0".to_string()), + path: Some("C:\\Program Files\\nodejs\\node.exe".to_string()), + message: "Node.js 已就绪:22.0.0".to_string(), + auto_install_supported: false, + }, + &DependencyStatus { + status: "ok".to_string(), + version: Some("2.44.0".to_string()), + path: Some("C:\\Program Files\\Git\\cmd\\git.exe".to_string()), + message: "Git 已就绪:2.44.0".to_string(), + auto_install_supported: false, + }, + ); + + assert!(result.is_none()); + } + #[test] fn winget_install_command_uses_expected_windows_flags() { assert_eq!( diff --git a/src-tauri/src/services/workspace_health_service.rs b/src-tauri/src/services/workspace_health_service.rs index b2358494e..002ad3753 100644 --- a/src-tauri/src/services/workspace_health_service.rs +++ b/src-tauri/src/services/workspace_health_service.rs @@ -1,4 +1,5 @@ use crate::workspace::{Workspace, WorkspaceManager, WorkspaceUpdate}; +use proxycast_core::app_paths; use std::path::{Path, PathBuf}; #[derive(Debug, Clone)] @@ -117,11 +118,7 @@ pub fn ensure_workspace_ready_with_auto_relocate( } fn build_workspace_fallback_root(workspace: &Workspace) -> Result { - let home_dir = dirs::home_dir().ok_or_else(|| "无法获取主目录".to_string())?; - let recovered_root = home_dir - .join(".proxycast") - .join("projects") - .join("recovered"); + let recovered_root = app_paths::resolve_projects_dir()?.join("recovered"); let workspace_name = sanitize_path_segment(&workspace.name); let short_id: String = workspace.id.chars().take(8).collect(); let dir_name = if short_id.is_empty() { diff --git a/src-tauri/src/skills/default_skills.rs b/src-tauri/src/skills/default_skills.rs index 69d2db923..f6e655e94 100644 --- a/src-tauri/src/skills/default_skills.rs +++ b/src-tauri/src/skills/default_skills.rs @@ -1,6 +1,10 @@ use std::fs; -use std::path::{Path, PathBuf}; +use std::path::Path; +#[cfg(test)] +use std::path::PathBuf; + +use proxycast_core::app_paths; use proxycast_core::models::{ BROADCAST_GENERATE_SKILL_DIRECTORY, COVER_GENERATE_SKILL_DIRECTORY, IMAGE_GENERATE_SKILL_DIRECTORY, LIBRARY_SKILL_DIRECTORY, MODAL_RESOURCE_SEARCH_SKILL_DIRECTORY, @@ -61,8 +65,9 @@ fn default_skills() -> [(&'static str, &'static str); 10] { ] } -fn skills_root_from_home(home_dir: &Path) -> PathBuf { - home_dir.join(".proxycast").join("skills") +#[cfg(test)] +fn skills_root_from_base(base_dir: &Path) -> PathBuf { + base_dir.join("skills") } /// 从 SKILL.md 内容中提取版本号,返回 (major, minor, patch) @@ -83,8 +88,7 @@ fn parse_skill_version(content: &str) -> Option<(u32, u32, u32)> { None } -fn ensure_default_local_skills_in_home(home_dir: &Path) -> Result, String> { - let skills_root = skills_root_from_home(home_dir); +fn ensure_default_local_skills_in_dir(skills_root: &Path) -> Result, String> { fs::create_dir_all(&skills_root) .map_err(|e| format!("创建技能目录失败 {}: {e}", skills_root.display()))?; @@ -120,8 +124,8 @@ fn ensure_default_local_skills_in_home(home_dir: &Path) -> Result, S } pub fn ensure_default_local_skills() -> Result, String> { - let home_dir = dirs::home_dir().ok_or_else(|| "无法获取用户 Home 目录".to_string())?; - ensure_default_local_skills_in_home(&home_dir) + let skills_root = app_paths::resolve_skills_dir()?; + ensure_default_local_skills_in_dir(&skills_root) } #[cfg(test)] @@ -131,13 +135,11 @@ mod tests { #[test] fn should_install_default_skill_when_missing() { let temp = tempfile::tempdir().expect("create temp dir"); - let installed = ensure_default_local_skills_in_home(temp.path()).expect("install"); + let skills_root = skills_root_from_base(temp.path()); + let installed = ensure_default_local_skills_in_dir(&skills_root).expect("install"); assert!(installed.contains(&SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY.to_string())); - let skill_md_path = temp - .path() - .join(".proxycast") - .join("skills") + let skill_md_path = skills_root .join(SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY) .join("SKILL.md"); assert!(skill_md_path.exists()); @@ -146,18 +148,15 @@ mod tests { #[test] fn should_not_overwrite_existing_skill() { let temp = tempfile::tempdir().expect("create temp dir"); - let skill_dir = temp - .path() - .join(".proxycast") - .join("skills") - .join(SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY); + let skills_root = skills_root_from_base(temp.path()); + let skill_dir = skills_root.join(SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY); fs::create_dir_all(&skill_dir).expect("create skill dir"); let skill_md_path = skill_dir.join("SKILL.md"); // 无版本号的自定义内容不应被覆盖 let existing_content = "custom skill content"; fs::write(&skill_md_path, existing_content).expect("write custom skill"); - let installed = ensure_default_local_skills_in_home(temp.path()).expect("install"); + let installed = ensure_default_local_skills_in_dir(&skills_root).expect("install"); assert!( !installed.contains(&SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY.to_string()), "无版本信息的已存在 skill 不应被重新安装" @@ -170,18 +169,15 @@ mod tests { #[test] fn should_upgrade_skill_when_newer_version_available() { let temp = tempfile::tempdir().expect("create temp dir"); - let skill_dir = temp - .path() - .join(".proxycast") - .join("skills") - .join(SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY); + let skills_root = skills_root_from_base(temp.path()); + let skill_dir = skills_root.join(SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY); fs::create_dir_all(&skill_dir).expect("create skill dir"); let skill_md_path = skill_dir.join("SKILL.md"); // 旧版本内容 let old_content = "---\nname: social_post_with_cover\nversion: 1.0.0\n---\nold content"; fs::write(&skill_md_path, old_content).expect("write old skill"); - let installed = ensure_default_local_skills_in_home(temp.path()).expect("install"); + let installed = ensure_default_local_skills_in_dir(&skills_root).expect("install"); assert!( installed.contains(&SOCIAL_POST_WITH_COVER_SKILL_DIRECTORY.to_string()), "内置版本更新时应自动升级" diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 0e22a9462..55454c98d 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "ProxyCast", - "version": "0.84.0", + "version": "0.85.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src/components/agent/chat/components/AgentPlanBlock.tsx b/src/components/agent/chat/components/AgentPlanBlock.tsx new file mode 100644 index 000000000..d50f78bd6 --- /dev/null +++ b/src/components/agent/chat/components/AgentPlanBlock.tsx @@ -0,0 +1,43 @@ +import React from "react"; +import { ListChecks, Loader2 } from "lucide-react"; + +import { Badge } from "@/components/ui/badge"; +import { MarkdownRenderer } from "./MarkdownRenderer"; + +interface AgentPlanBlockProps { + content: string; + isComplete?: boolean; +} + +export const AgentPlanBlock: React.FC = ({ + content, + isComplete = true, +}) => { + if (!content.trim()) { + return null; + } + + return ( +
+
+
+ +
+
执行计划
+ + {isComplete ? ( + "已生成" + ) : ( + + + 规划中 + + )} + +
+
+ +
+
+ ); +}; diff --git a/src/components/agent/chat/components/AgentRuntimeStrip.tsx b/src/components/agent/chat/components/AgentRuntimeStrip.tsx new file mode 100644 index 000000000..5a9f65ad4 --- /dev/null +++ b/src/components/agent/chat/components/AgentRuntimeStrip.tsx @@ -0,0 +1,186 @@ +import React, { useMemo } from "react"; + +import { Badge } from "@/components/ui/badge"; +import type { SchedulerEvent, SchedulerProgress } from "@/lib/api/subAgentScheduler"; + +import type { ChatToolPreferences } from "../utils/chatToolPreferences"; +import type { HarnessSessionState } from "../utils/harnessState"; + +interface AgentRuntimeStripProps { + activeTheme?: string; + toolPreferences: ChatToolPreferences; + harnessState: HarnessSessionState; + subAgentRuntime: { + isRunning: boolean; + progress: SchedulerProgress | null; + events: SchedulerEvent[]; + }; + variant?: "standalone" | "embedded"; + isSending?: boolean; + runtimeStatusTitle?: string | null; +} + +const THEME_LABELS: Record = { + general: "通用对话", + knowledge: "知识探索", + planning: "计划规划", +}; + +interface CapabilityItem { + key: string; + label: string; + enabled: boolean; +} + +interface StatusItem { + key: string; + label: string; + tone?: "default" | "outline" | "secondary"; +} + +export const AgentRuntimeStrip: React.FC = ({ + activeTheme, + toolPreferences, + harnessState, + subAgentRuntime, + variant = "standalone", + isSending = false, + runtimeStatusTitle = null, +}) => { + const themeLabel = + THEME_LABELS[activeTheme?.trim().toLowerCase() || ""] || "通用对话"; + + const capabilities = useMemo( + () => [ + { key: "direct", label: "直接回答", enabled: true }, + { key: "thinking", label: "深度思考", enabled: toolPreferences.thinking }, + { + key: "web_search", + label: "联网搜索", + enabled: toolPreferences.webSearch, + }, + { key: "task", label: "后台任务", enabled: toolPreferences.task }, + { key: "subagent", label: "多代理", enabled: toolPreferences.subagent }, + ], + [toolPreferences], + ); + + const statusItems = useMemo(() => { + const nextItems: StatusItem[] = []; + + if (isSending) { + nextItems.push({ + key: "sending", + label: runtimeStatusTitle || "正在准备执行", + tone: "secondary", + }); + } + + if (harnessState.plan.phase === "planning") { + nextItems.push({ + key: "planning", + label: "正在整理执行计划", + tone: "secondary", + }); + } + + if (harnessState.plan.items.length > 0) { + nextItems.push({ + key: "plan_items", + label: `当前计划 ${harnessState.plan.items.length} 项`, + tone: "outline", + }); + } + + if (harnessState.pendingApprovals.length > 0) { + nextItems.push({ + key: "pending", + label: `等待确认 ${harnessState.pendingApprovals.length}`, + tone: "secondary", + }); + } + + if (subAgentRuntime.isRunning) { + const progressLabel = + subAgentRuntime.progress && + typeof subAgentRuntime.progress.completed === "number" && + typeof subAgentRuntime.progress.total === "number" + ? `子代理运行中 ${subAgentRuntime.progress.completed}/${subAgentRuntime.progress.total}` + : "子代理运行中"; + nextItems.push({ + key: "subagent_running", + label: progressLabel, + tone: "secondary", + }); + } else if (harnessState.delegatedTasks.length > 0) { + nextItems.push({ + key: "delegated", + label: `最近委派 ${harnessState.delegatedTasks.length}`, + tone: "outline", + }); + } + + if (harnessState.outputSignals.length > 0) { + nextItems.push({ + key: "outputs", + label: `最近产物 ${harnessState.outputSignals.length}`, + tone: "outline", + }); + } + + if (nextItems.length === 0) { + nextItems.push({ + key: "default_mode", + label: "当前以直接回答优先,必要时再升级工具链", + tone: "outline", + }); + } + + return nextItems; + }, [ + harnessState, + isSending, + runtimeStatusTitle, + subAgentRuntime.isRunning, + subAgentRuntime.progress, + ]); + + return ( +
+
+
通用 Agent
+ {themeLabel} +
+
+ {capabilities.map((item) => ( + + {item.label} + + ))} +
+
+ {statusItems.map((item) => ( + + {item.label} + + ))} +
+
+ ); +}; + +export default AgentRuntimeStrip; diff --git a/src/components/agent/chat/components/AgentThreadTimeline.tsx b/src/components/agent/chat/components/AgentThreadTimeline.tsx new file mode 100644 index 000000000..19b4a3f78 --- /dev/null +++ b/src/components/agent/chat/components/AgentThreadTimeline.tsx @@ -0,0 +1,504 @@ +import React, { useMemo } from "react"; +import { + AlertTriangle, + Bot, + ChevronDown, + Clock3, + FileText, + Globe, + Loader2, + Search, + ShieldAlert, + Sparkles, + TerminalSquare, + Wrench, +} from "lucide-react"; + +import { Badge } from "@/components/ui/badge"; +import type { ToolCallState } from "@/lib/api/agentStream"; +import type { + ActionRequired, + AgentThreadItem, + AgentThreadTurn, + ConfirmResponse, +} from "../types"; +import { MarkdownRenderer } from "./MarkdownRenderer"; +import { ToolCallItem } from "./ToolCallDisplay"; +import { DecisionPanel } from "./DecisionPanel"; +import { AgentPlanBlock } from "./AgentPlanBlock"; + +interface AgentThreadTimelineProps { + turn: AgentThreadTurn; + items: AgentThreadItem[]; + isCurrentTurn?: boolean; + onFileClick?: (fileName: string, content: string) => void; + onPermissionResponse?: (response: ConfirmResponse) => void; +} + +function formatTimestamp(value?: string): string | null { + if (!value) { + return null; + } + + const date = new Date(value); + if (Number.isNaN(date.getTime())) { + return null; + } + + return date.toLocaleTimeString("zh-CN", { + hour: "2-digit", + minute: "2-digit", + }); +} + +function toQuestionOptions( + options: Array<{ label: string; description?: string }> | undefined, +) { + return options?.map((option) => ({ + label: option.label, + description: option.description, + })); +} + +function stringifyResponse(value: unknown): string | undefined { + if (typeof value === "string") { + const normalized = value.trim(); + return normalized || undefined; + } + + if (value === null || value === undefined) { + return undefined; + } + + try { + return JSON.stringify(value, null, 2); + } catch { + return String(value); + } +} + +function toActionRequired(item: AgentThreadItem): ActionRequired | null { + if (item.type === "approval_request") { + return { + requestId: item.request_id, + actionType: "tool_confirmation", + toolName: item.tool_name, + arguments: + item.arguments && typeof item.arguments === "object" + ? (item.arguments as Record) + : undefined, + prompt: item.prompt, + status: item.status === "completed" ? "submitted" : "pending", + submittedResponse: stringifyResponse(item.response), + submittedUserData: item.response, + }; + } + + if (item.type === "request_user_input") { + return { + requestId: item.request_id, + actionType: + item.action_type === "elicitation" ? "elicitation" : "ask_user", + prompt: item.prompt, + questions: item.questions?.map((question) => ({ + question: question.question, + header: question.header, + options: toQuestionOptions(question.options), + multiSelect: question.multi_select, + })), + status: item.status === "completed" ? "submitted" : "pending", + submittedResponse: stringifyResponse(item.response), + submittedUserData: item.response, + }; + } + + return null; +} + +function mapItemStatus( + status: AgentThreadItem["status"], +): ToolCallState["status"] { + if (status === "failed") { + return "failed"; + } + return status === "completed" ? "completed" : "running"; +} + +function toToolCallState(item: AgentThreadItem): ToolCallState | null { + switch (item.type) { + case "tool_call": + return { + id: item.id, + name: item.tool_name, + arguments: + item.arguments === undefined + ? undefined + : JSON.stringify(item.arguments, null, 2), + status: mapItemStatus(item.status), + result: + item.output !== undefined || + item.error !== undefined || + item.metadata !== undefined + ? { + success: + item.success ?? + (item.status === "completed" && item.error === undefined), + output: item.output || "", + error: item.error, + metadata: + item.metadata && typeof item.metadata === "object" + ? (item.metadata as Record) + : undefined, + } + : undefined, + startTime: new Date(item.started_at), + endTime: item.completed_at ? new Date(item.completed_at) : undefined, + }; + case "command_execution": + return { + id: item.id, + name: "exec_command", + arguments: JSON.stringify( + { command: item.command, cwd: item.cwd }, + null, + 2, + ), + status: mapItemStatus(item.status), + result: + item.aggregated_output !== undefined || + item.error !== undefined || + item.exit_code !== undefined + ? { + success: item.status === "completed" && item.error === undefined, + output: item.aggregated_output || "", + error: item.error, + metadata: + item.exit_code !== undefined + ? { exit_code: item.exit_code, cwd: item.cwd } + : { cwd: item.cwd }, + } + : undefined, + startTime: new Date(item.started_at), + endTime: item.completed_at ? new Date(item.completed_at) : undefined, + }; + case "web_search": + return { + id: item.id, + name: item.action || "web_search", + arguments: + item.query !== undefined + ? JSON.stringify({ query: item.query }, null, 2) + : undefined, + status: mapItemStatus(item.status), + result: + item.output !== undefined + ? { + success: item.status !== "failed", + output: item.output, + } + : undefined, + startTime: new Date(item.started_at), + endTime: item.completed_at ? new Date(item.completed_at) : undefined, + }; + default: + return null; + } +} + +function resolveStatusBadgeVariant( + status: AgentThreadItem["status"], +): "secondary" | "outline" | "destructive" { + if (status === "failed") { + return "destructive"; + } + return status === "completed" ? "outline" : "secondary"; +} + +function TimelineCard({ + icon: Icon, + title, + badge, + timestamp, + children, +}: { + icon: React.ComponentType<{ className?: string }>; + title: string; + badge?: React.ReactNode; + timestamp?: string | null; + children: React.ReactNode; +}) { + return ( +
+
+
+ +
+
{title}
+ {badge ?
{badge}
: null} + {timestamp ? ( +
{timestamp}
+ ) : null} +
+ {children} +
+ ); +} + +export const AgentThreadTimeline: React.FC = ({ + turn, + items, + isCurrentTurn = false, + onFileClick, + onPermissionResponse, +}) => { + const visibleItems = useMemo( + () => + items.filter( + (item) => item.type !== "user_message" && item.type !== "agent_message", + ), + [items], + ); + + if (visibleItems.length === 0) { + return null; + } + + return ( +
+
+
执行轨迹
+ {isCurrentTurn ? 当前回合 : null} + + {turn.status === "running" + ? "执行中" + : turn.status === "failed" + ? "失败" + : turn.status === "aborted" + ? "已中断" + : "已完成"} + +
+ + {formatTimestamp(turn.started_at) || "刚刚"} +
+
+ +
+ {visibleItems.map((item) => { + const timestamp = formatTimestamp(item.completed_at || item.updated_at); + const actionRequest = toActionRequired(item); + const toolCall = toToolCallState(item); + + if (item.type === "plan") { + return ( + + ); + } + + if (item.type === "reasoning") { + return ( +
+ + + 思考摘要 + + {item.status === "in_progress" ? ( + + + 推理中 + + ) : item.status === "failed" ? ( + "推理失败" + ) : ( + "已整理" + )} + + + +
+ +
+
+ ); + } + + if (toolCall) { + return ( +
+ +
+ ); + } + + if (actionRequest) { + return ( +
+ onPermissionResponse?.(response)} + /> +
+ ); + } + + if (item.type === "file_artifact") { + return ( + + {item.source} + + } + timestamp={timestamp} + > + + + ); + } + + if (item.type === "subagent_activity") { + return ( + + {item.status_label} + + } + timestamp={timestamp} + > + {item.summary ? ( +
{item.summary}
+ ) : null} + {item.role || item.model ? ( +
+ {item.role ? {item.role} : null} + {item.model ? {item.model} : null} +
+ ) : null} +
+ ); + } + + if (item.type === "turn_summary") { + return ( + + + 进行中 + + ) : ( + 摘要 + ) + } + timestamp={timestamp} + > + + + ); + } + + if (item.type === "warning") { + return ( + {item.code || "warning"}} + timestamp={timestamp} + > +
{item.message}
+
+ ); + } + + if (item.type === "error") { + return ( + 失败} + timestamp={timestamp} + > +
{item.message}
+
+ ); + } + + return ( + + {item.status} + + } + timestamp={timestamp} + > +
+ 该事件类型已记录到 timeline 中。 +
+
+ ); + })} +
+
+ ); +}; + +export default AgentThreadTimeline; diff --git a/src/components/agent/chat/components/ChatNavbar.tsx b/src/components/agent/chat/components/ChatNavbar.tsx index 885a62ca2..12ebef057 100644 --- a/src/components/agent/chat/components/ChatNavbar.tsx +++ b/src/components/agent/chat/components/ChatNavbar.tsx @@ -33,6 +33,7 @@ interface ChatNavbarProps { onToggleHarnessPanel?: () => void; harnessPendingCount?: number; harnessAttentionLevel?: "idle" | "active" | "warning"; + harnessToggleLabel?: string; novelCanvasControls?: { chapterListCollapsed: boolean; onToggleChapterList: () => void; @@ -58,6 +59,7 @@ export const ChatNavbar: React.FC = ({ onToggleHarnessPanel, harnessPendingCount = 0, harnessAttentionLevel = "idle", + harnessToggleLabel = "Harness", novelCanvasControls = null, }) => { return ( @@ -173,15 +175,19 @@ export const ChatNavbar: React.FC = ({ )} onClick={onToggleHarnessPanel} aria-label={ - harnessPanelVisible ? "收起 Harness 面板" : "展开 Harness 面板" + harnessPanelVisible + ? `收起${harnessToggleLabel}` + : `展开${harnessToggleLabel}` } aria-expanded={harnessPanelVisible} title={ - harnessPanelVisible ? "收起 Harness 面板" : "展开 Harness 面板" + harnessPanelVisible + ? `收起${harnessToggleLabel}` + : `展开${harnessToggleLabel}` } > - Harness + {harnessToggleLabel} {harnessPendingCount > 0 ? ( {harnessPendingCount > 99 ? "99+" : harnessPendingCount} diff --git a/src/components/agent/chat/components/EmptyState.test.tsx b/src/components/agent/chat/components/EmptyState.test.tsx index c5ab964f5..d96b15267 100644 --- a/src/components/agent/chat/components/EmptyState.test.tsx +++ b/src/components/agent/chat/components/EmptyState.test.tsx @@ -11,17 +11,18 @@ const { mockGetConfig } = vi.hoisted(() => ({ mockGetConfig: vi.fn(async () => ({})), })); -const mockCharacterMention = vi.fn< - (props: { - characters?: Character[]; - skills?: Skill[]; - onSelectSkill?: (skill: Skill) => void; - value: string; - onChange: (value: string) => void; - }) => React.ReactNode ->(); +const mockCharacterMention = + vi.fn< + (props: { + characters?: Character[]; + skills?: Skill[]; + onSelectSkill?: (skill: Skill) => void; + value: string; + onChange: (value: string) => void; + }) => React.ReactNode + >(); -vi.mock("@/hooks/useTauri", () => ({ +vi.mock("@/lib/api/appConfig", () => ({ getConfig: mockGetConfig, })); @@ -33,7 +34,11 @@ vi.mock("../utils/entryPromptComposer", () => ({ composeEntryPrompt: vi.fn(() => ""), createDefaultEntrySlotValues: vi.fn(() => ({})), formatEntryTaskPreview: vi.fn(() => ""), - getEntryTaskTemplate: vi.fn(() => ({ slots: [], description: "", label: "" })), + getEntryTaskTemplate: vi.fn(() => ({ + slots: [], + description: "", + label: "", + })), SOCIAL_MEDIA_ENTRY_TASKS: [], validateEntryTaskSlots: vi.fn(() => ({ valid: true, missing: [] })), })); @@ -95,11 +100,15 @@ vi.mock("@/components/ui/textarea", () => { }); vi.mock("@/components/ui/select", () => ({ - Select: ({ children }: { children: React.ReactNode }) =>
{children}
, + Select: ({ children }: { children: React.ReactNode }) => ( +
{children}
+ ), SelectContent: ({ children }: { children: React.ReactNode }) => (
{children}
), - SelectItem: ({ children }: { children: React.ReactNode }) =>
{children}
, + SelectItem: ({ children }: { children: React.ReactNode }) => ( +
{children}
+ ), SelectTrigger: ({ children }: { children: React.ReactNode }) => ( ), @@ -107,7 +116,9 @@ vi.mock("@/components/ui/select", () => ({ })); vi.mock("@/components/ui/popover", () => ({ - Popover: ({ children }: { children: React.ReactNode }) =>
{children}
, + Popover: ({ children }: { children: React.ReactNode }) => ( +
{children}
+ ), PopoverContent: ({ children }: { children: React.ReactNode }) => (
{children}
), @@ -159,7 +170,9 @@ afterEach(() => { vi.clearAllMocks(); }); -function renderEmptyState(props?: Partial>) { +function renderEmptyState( + props?: Partial>, +) { const container = document.createElement("div"); document.body.appendChild(container); const root = createRoot(container); @@ -219,11 +232,15 @@ describe("EmptyState", () => { await Promise.resolve(); }); - const mention = container.querySelector('[data-testid="character-mention-stub"]'); + const mention = container.querySelector( + '[data-testid="character-mention-stub"]', + ); expect(mention).toBeTruthy(); expect(mockCharacterMention.mock.calls.length).toBeGreaterThan(0); const latestCall = - mockCharacterMention.mock.calls[mockCharacterMention.mock.calls.length - 1][0]; + mockCharacterMention.mock.calls[ + mockCharacterMention.mock.calls.length - 1 + ][0]; expect(latestCall.characters).toEqual(characters); expect(latestCall.skills).toEqual(skills); @@ -234,13 +251,14 @@ describe("EmptyState", () => { }); it("选择技能后发送应自动附加 skill 前缀,且发送后清除激活技能", async () => { - const onSend = vi.fn< - ( - value: string, - executionStrategy?: "react" | "code_orchestrated" | "auto", - images?: unknown[], - ) => void - >(); + const onSend = + vi.fn< + ( + value: string, + executionStrategy?: "react" | "code_orchestrated" | "auto", + images?: unknown[], + ) => void + >(); const skill: Skill = { key: "canvas-design", name: "canvas-design", @@ -260,7 +278,9 @@ describe("EmptyState", () => { }); const latestCall = - mockCharacterMention.mock.calls[mockCharacterMention.mock.calls.length - 1][0]; + mockCharacterMention.mock.calls[ + mockCharacterMention.mock.calls.length - 1 + ][0]; expect(typeof latestCall.onSelectSkill).toBe("function"); act(() => { @@ -310,13 +330,14 @@ describe("EmptyState", () => { }); it("社媒主题发送时应默认走 social_post_with_cover skill", async () => { - const onSend = vi.fn< - ( - value: string, - executionStrategy?: "react" | "code_orchestrated" | "auto", - images?: unknown[], - ) => void - >(); + const onSend = + vi.fn< + ( + value: string, + executionStrategy?: "react" | "code_orchestrated" | "auto", + images?: unknown[], + ) => void + >(); vi.mocked(composeEntryPrompt).mockReturnValue("请输出一篇新品社媒文案"); const container = renderEmptyState({ @@ -349,13 +370,14 @@ describe("EmptyState", () => { })); vi.mocked(composeEntryPrompt).mockReturnValue("请输出一篇用户访谈纪要"); - const onSend = vi.fn< - ( - value: string, - executionStrategy?: "react" | "code_orchestrated" | "auto", - images?: unknown[], - ) => void - >(); + const onSend = + vi.fn< + ( + value: string, + executionStrategy?: "react" | "code_orchestrated" | "auto", + images?: unknown[], + ) => void + >(); const container = renderEmptyState({ activeTheme: "social-media", onSend, @@ -381,13 +403,14 @@ describe("EmptyState", () => { }); it("社媒主题手动选择 skill 时应优先使用手动 skill", async () => { - const onSend = vi.fn< - ( - value: string, - executionStrategy?: "react" | "code_orchestrated" | "auto", - images?: unknown[], - ) => void - >(); + const onSend = + vi.fn< + ( + value: string, + executionStrategy?: "react" | "code_orchestrated" | "auto", + images?: unknown[], + ) => void + >(); vi.mocked(composeEntryPrompt).mockReturnValue("请输出一篇品牌故事"); const skill: Skill = { key: "custom-social-skill", @@ -408,7 +431,9 @@ describe("EmptyState", () => { }); const latestCall = - mockCharacterMention.mock.calls[mockCharacterMention.mock.calls.length - 1][0]; + mockCharacterMention.mock.calls[ + mockCharacterMention.mock.calls.length - 1 + ][0]; act(() => { latestCall.onSelectSkill?.(skill); }); @@ -429,12 +454,18 @@ describe("EmptyState", () => { ); }); - it("通用主题工具栏应包含附件和深度思考开关", async () => { + it("通用主题工具栏应包含附件、思考、后台任务与多代理开关", async () => { const onThinkingEnabledChange = vi.fn<(enabled: boolean) => void>(); + const onTaskEnabledChange = vi.fn<(enabled: boolean) => void>(); + const onSubagentEnabledChange = vi.fn<(enabled: boolean) => void>(); const container = renderEmptyState({ activeTheme: "general", thinkingEnabled: false, onThinkingEnabledChange, + taskEnabled: false, + onTaskEnabledChange, + subagentEnabled: false, + onSubagentEnabledChange, }); await act(async () => { await Promise.resolve(); @@ -449,11 +480,27 @@ describe("EmptyState", () => { 'button[title="开启深度思考"]', ) as HTMLButtonElement | null; expect(thinkingButton).toBeTruthy(); + const taskButton = container.querySelector( + 'button[title="开启后台任务偏好"]', + ) as HTMLButtonElement | null; + expect(taskButton).toBeTruthy(); + const subagentButton = container.querySelector( + 'button[title="开启多代理偏好"]', + ) as HTMLButtonElement | null; + expect(subagentButton).toBeTruthy(); act(() => { thinkingButton?.click(); }); + act(() => { + taskButton?.click(); + }); + act(() => { + subagentButton?.click(); + }); expect(onThinkingEnabledChange).toHaveBeenCalledWith(true); + expect(onTaskEnabledChange).toHaveBeenCalledWith(true); + expect(onSubagentEnabledChange).toHaveBeenCalledWith(true); }); }); diff --git a/src/components/agent/chat/components/EmptyState.tsx b/src/components/agent/chat/components/EmptyState.tsx index 34f31b11d..864fb9d19 100644 --- a/src/components/agent/chat/components/EmptyState.tsx +++ b/src/components/agent/chat/components/EmptyState.tsx @@ -15,6 +15,8 @@ import { Globe, Music, Code2, + ListChecks, + Workflow, } from "lucide-react"; import { getConfig } from "@/lib/api/appConfig"; import type { CreationMode, EntryTaskSlotValues, EntryTaskType } from "./types"; @@ -55,6 +57,7 @@ import { useActiveSkill } from "./Inputbar/hooks/useActiveSkill"; 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 Assets import iconXhs from "@/assets/platforms/xhs.png"; @@ -396,6 +399,10 @@ interface EmptyStateProps { onWebSearchEnabledChange?: (enabled: boolean) => void; thinkingEnabled?: boolean; onThinkingEnabledChange?: (enabled: boolean) => void; + taskEnabled?: boolean; + onTaskEnabledChange?: (enabled: boolean) => void; + subagentEnabled?: boolean; + onSubagentEnabledChange?: (enabled: boolean) => void; hasCanvasContent?: boolean; hasContentId?: boolean; selectedText?: string; @@ -528,6 +535,10 @@ export const EmptyState: React.FC = ({ onWebSearchEnabledChange, thinkingEnabled = false, onThinkingEnabledChange, + taskEnabled = false, + onTaskEnabledChange, + subagentEnabled = false, + onSubagentEnabledChange, hasCanvasContent = false, hasContentId = false, selectedText = "", @@ -614,6 +625,7 @@ export const EmptyState: React.FC = ({ // Popover 打开状态 const [ratioPopoverOpen, setRatioPopoverOpen] = useState(false); const [stylePopoverOpen, setStylePopoverOpen] = useState(false); + const isGeneralTheme = isGeneralResearchTheme(activeTheme); const wrapTextWithDefaultSkill = (text: string) => { const wrappedByActiveSkill = wrapTextWithSkill(text); @@ -1241,7 +1253,7 @@ export const EmptyState: React.FC = ({ )} - {activeTheme === "general" && ( + {isGeneralTheme && ( <> + + )} diff --git a/src/components/agent/chat/components/HarnessStatusPanel.test.tsx b/src/components/agent/chat/components/HarnessStatusPanel.test.tsx index 488d5c2e8..dc35ed8cd 100644 --- a/src/components/agent/chat/components/HarnessStatusPanel.test.tsx +++ b/src/components/agent/chat/components/HarnessStatusPanel.test.tsx @@ -29,6 +29,7 @@ function createHarnessState( overrides: Partial = {}, ): HarnessSessionState { return { + runtimeStatus: null, pendingApprovals: [], latestContextTrace: [], plan: { @@ -123,6 +124,51 @@ afterEach(() => { }); describe("HarnessStatusPanel", () => { + it("弹窗模式应默认展示完整内容且不渲染展开按钮", () => { + renderPanel({ + layout: "dialog", + }); + + expect(document.body.textContent).toContain("待审批"); + expect(document.body.textContent).toContain("文件活动"); + expect(document.body.textContent).toContain("计划状态"); + expect(document.body.textContent).toContain("上下文"); + expect(document.body.textContent).not.toContain("展开详情"); + expect(document.body.textContent).not.toContain("收起详情"); + }); + + it("应支持自定义标题说明与前置运行概览内容", () => { + renderPanel({ + title: "Agent 工作台", + description: "集中查看代理运行轨迹。", + toggleLabel: "工作台详情", + leadContent:
通用 Agent 运行概览
, + }); + + expect(document.body.textContent).toContain("Agent 工作台"); + expect(document.body.textContent).toContain("集中查看代理运行轨迹。"); + expect(document.body.textContent).toContain("通用 Agent 运行概览"); + expect(document.body.textContent).toContain("收起工作台详情"); + }); + + it("存在 runtimeStatus 时应在工作台中展示当前执行阶段", () => { + renderPanel({ + harnessState: createHarnessState({ + runtimeStatus: { + phase: "routing", + title: "正在建立执行回合", + detail: "已提交到运行时,正在等待首个执行事件。", + checkpoints: ["会话已建立", "等待首个模型事件"], + }, + }), + }); + + expect(document.body.textContent).toContain("执行阶段"); + expect(document.body.textContent).toContain("当前执行阶段"); + expect(document.body.textContent).toContain("正在建立执行回合"); + expect(document.body.textContent).toContain("等待首个模型事件"); + }); + it("摘要卡和快速导航应支持跳转到对应区块", () => { const scrollIntoViewMock = vi.fn(); const originalScrollIntoView = HTMLElement.prototype.scrollIntoView; diff --git a/src/components/agent/chat/components/HarnessStatusPanel.tsx b/src/components/agent/chat/components/HarnessStatusPanel.tsx index 67bb5e725..0dc8a8832 100644 --- a/src/components/agent/chat/components/HarnessStatusPanel.tsx +++ b/src/components/agent/chat/components/HarnessStatusPanel.tsx @@ -78,10 +78,15 @@ interface HarnessStatusPanelProps { error: string | null; }; environment: HarnessEnvironmentSummary; + layout?: "default" | "sidebar" | "dialog"; onLoadFilePreview?: (path: string) => Promise; onOpenFile?: (fileName: string, content: string) => void; onRevealPath?: (path: string) => Promise; onOpenPath?: (path: string) => Promise; + title?: string; + description?: string; + toggleLabel?: string; + leadContent?: ReactNode; } interface PreviewDialogState { @@ -103,6 +108,7 @@ type OutputFilterValue = "all" | "path" | "offload" | "truncated" | "summary"; type FileDisplayMode = "timeline" | "grouped"; type HarnessSectionKey = + | "runtime" | "approvals" | "files" | "outputs" @@ -111,6 +117,19 @@ type HarnessSectionKey = | "context" | "capabilities"; +interface HarnessSectionNavItem { + key: HarnessSectionKey; + label: string; +} + +interface HarnessSummaryCard { + sectionKey: HarnessSectionKey; + title: string; + value: string; + hint: string; + icon: LucideIcon; +} + function getFileName(path: string): string { const normalized = path.replace(/\\/g, "/"); const segments = normalized.split("/"); @@ -272,6 +291,25 @@ function describeApproval(item: ActionRequired): string | undefined { return hints.length > 0 ? hints.join(" · ") : undefined; } +function formatRuntimePhaseLabel( + runtimeStatus: HarnessSessionState["runtimeStatus"], +): string { + if (!runtimeStatus) { + return "空闲"; + } + + switch (runtimeStatus.phase) { + case "preparing": + return "准备中"; + case "routing": + return "建回合中"; + case "context": + return "装载上下文"; + default: + return runtimeStatus.phase; + } +} + function summarizeSchedulerEvent(event: SchedulerEvent): string { switch (event.type) { case "started": @@ -377,12 +415,19 @@ export function HarnessStatusPanel({ harnessState, subAgentRuntime, environment, + layout = "default", onLoadFilePreview, onOpenFile, onRevealPath, onOpenPath, + title = "Harness 运行面板", + description = "展示最近文件活动、工具输出、审批与上下文装载情况。", + toggleLabel = "详情", + leadContent, }: HarnessStatusPanelProps) { const [expanded, setExpanded] = useState(true); + const isDialogLayout = layout === "dialog"; + const isDetailsExpanded = isDialogLayout ? true : expanded; const [fileFilter, setFileFilter] = useState("all"); const [outputFilter, setOutputFilter] = useState("all"); const [fileDisplayMode, setFileDisplayMode] = @@ -529,31 +574,44 @@ export function HarnessStatusPanel({ }, [filteredFileEvents]); const availableSections = useMemo( - () => [ - harnessState.pendingApprovals.length > 0 - ? { key: "approvals" as const, label: "待审批" } - : null, - harnessState.recentFileEvents.length > 0 - ? { key: "files" as const, label: "文件活动" } - : null, - harnessState.outputSignals.length > 0 - ? { key: "outputs" as const, label: "工具输出" } - : null, - harnessState.plan.phase !== "idle" || harnessState.plan.items.length > 0 - ? { key: "plan" as const, label: "规划状态" } - : null, - subAgentRuntime.isRunning || - harnessState.delegatedTasks.length > 0 || - recentSchedulerEvents.length > 0 || - subAgentRuntime.error || - subAgentRuntime.result - ? { key: "delegation" as const, label: "子任务委派" } - : null, - harnessState.latestContextTrace.length > 0 - ? { key: "context" as const, label: "上下文轨迹" } - : null, - { key: "capabilities" as const, label: "已装载能力" }, - ].filter((item): item is { key: HarnessSectionKey; label: string } => item !== null), + () => { + const sections: HarnessSectionNavItem[] = []; + + if (harnessState.runtimeStatus) { + sections.push({ key: "runtime", label: "当前阶段" }); + } + if (harnessState.pendingApprovals.length > 0) { + sections.push({ key: "approvals", label: "待审批" }); + } + if (harnessState.recentFileEvents.length > 0) { + sections.push({ key: "files", label: "文件活动" }); + } + if (harnessState.outputSignals.length > 0) { + sections.push({ key: "outputs", label: "工具输出" }); + } + if ( + harnessState.plan.phase !== "idle" || + harnessState.plan.items.length > 0 + ) { + sections.push({ key: "plan", label: "规划状态" }); + } + if ( + subAgentRuntime.isRunning || + harnessState.delegatedTasks.length > 0 || + recentSchedulerEvents.length > 0 || + subAgentRuntime.error || + subAgentRuntime.result + ) { + sections.push({ key: "delegation", label: "子任务委派" }); + } + if (harnessState.latestContextTrace.length > 0) { + sections.push({ key: "context", label: "上下文轨迹" }); + } + + sections.push({ key: "capabilities", label: "已装载能力" }); + + return sections; + }, [ harnessState.delegatedTasks.length, harnessState.latestContextTrace.length, @@ -562,6 +620,7 @@ export function HarnessStatusPanel({ harnessState.plan.items.length, harnessState.plan.phase, harnessState.recentFileEvents.length, + harnessState.runtimeStatus, recentSchedulerEvents.length, subAgentRuntime.error, subAgentRuntime.isRunning, @@ -570,46 +629,66 @@ export function HarnessStatusPanel({ ); const summaryCards = useMemo( - () => [ - { - sectionKey: "approvals" as const, - title: "待审批", - value: `${harnessState.pendingApprovals.length}`, - hint: - harnessState.pendingApprovals.length > 0 - ? "需要你确认的操作" - : "当前无阻塞审批", - icon: ShieldAlert, - }, - { - sectionKey: "files" as const, - title: "文件活动", - value: `${harnessState.recentFileEvents.length}`, - hint: - harnessState.recentFileEvents[0]?.displayName || "暂无可展示文件活动", - icon: FolderOpen, - }, - { - sectionKey: "plan" as const, - title: "计划状态", - value: - harnessState.plan.phase === "planning" - ? "进行中" - : harnessState.plan.phase === "ready" - ? "已就绪" - : "空闲", - hint: - harnessState.plan.items[0]?.content || "未检测到显式计划快照", - icon: ListChecks, - }, - { - sectionKey: "context" as const, - title: "上下文", - value: `${environment.activeContextCount}/${environment.contextItemsCount}`, - hint: environment.contextEnabled ? "上下文工作台已启用" : "普通聊天模式", - icon: Sparkles, - }, - ], + () => { + const cards: HarnessSummaryCard[] = []; + + if (harnessState.runtimeStatus) { + cards.push({ + sectionKey: "runtime", + title: "执行阶段", + value: formatRuntimePhaseLabel(harnessState.runtimeStatus), + hint: + harnessState.runtimeStatus.detail || + harnessState.runtimeStatus.title, + icon: Loader2, + }); + } + + cards.push( + { + sectionKey: "approvals", + title: "待审批", + value: `${harnessState.pendingApprovals.length}`, + hint: + harnessState.pendingApprovals.length > 0 + ? "需要你确认的操作" + : "当前无阻塞审批", + icon: ShieldAlert, + }, + { + sectionKey: "files", + title: "文件活动", + value: `${harnessState.recentFileEvents.length}`, + hint: + harnessState.recentFileEvents[0]?.displayName || + "暂无可展示文件活动", + icon: FolderOpen, + }, + { + sectionKey: "plan", + title: "计划状态", + value: + harnessState.plan.phase === "planning" + ? "进行中" + : harnessState.plan.phase === "ready" + ? "已就绪" + : "空闲", + hint: + harnessState.plan.items[0]?.content || "未检测到显式计划快照", + icon: ListChecks, + }, + { + sectionKey: "context", + title: "上下文", + value: `${environment.activeContextCount}/${environment.contextItemsCount}`, + hint: + environment.contextEnabled ? "上下文工作台已启用" : "普通聊天模式", + icon: Sparkles, + }, + ); + + return cards; + }, [ environment.activeContextCount, environment.contextEnabled, @@ -618,6 +697,7 @@ export function HarnessStatusPanel({ harnessState.plan.items, harnessState.plan.phase, harnessState.recentFileEvents, + harnessState.runtimeStatus, ], ); @@ -781,13 +861,24 @@ export function HarnessStatusPanel({ return ( <> -
+

- Harness 运行面板 + {title}

{subAgentRuntime.isRunning ? ( @@ -797,28 +888,43 @@ export function HarnessStatusPanel({ ) : null}

- 展示最近文件活动、工具输出、审批与上下文装载情况。 + {description}

- + {!isDialogLayout ? ( + + ) : null}
-
+ {leadContent ? ( +
{leadContent}
+ ) : null} + +
{summaryCards.map((card) => ( - {expanded ? ( - + {isDetailsExpanded ? ( +
{availableSections.length > 0 ? (
@@ -1139,6 +1254,44 @@ export function HarnessStatusPanel({ ) : null} + {harnessState.runtimeStatus ? ( +
+
+
+
+ + {harnessState.runtimeStatus.title} +
+
+ {harnessState.runtimeStatus.detail} +
+
+ + {harnessState.runtimeStatus.checkpoints && + harnessState.runtimeStatus.checkpoints.length > 0 ? ( +
+ {harnessState.runtimeStatus.checkpoints.map( + (checkpoint, index) => ( + + {checkpoint} + + ), + )} +
+ ) : null} +
+
+ ) : null} + {harnessState.pendingApprovals.length > 0 ? (
` - margin-bottom: 8px; + margin: 0 8px 8px; + box-sizing: border-box; display: flex; align-items: flex-start; gap: 8px; diff --git a/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx b/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx index 54a39c7f0..d7deb4fc3 100644 --- a/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx +++ b/src/components/agent/chat/components/Inputbar/components/InputbarComposerSection.tsx @@ -140,6 +140,7 @@ export const InputbarComposerSection: React.FC< showDragHandle={!isThemeWorkbenchVariant} visualVariant={isThemeWorkbenchVariant ? "floating" : "default"} topExtra={topExtra} + activeTheme={activeTheme} leftExtra={ = ({ @@ -98,6 +99,7 @@ export const InputbarCore: React.FC = ({ showTranslate = true, showDragHandle = true, visualVariant = "default", + activeTheme, }) => { const [isComposerExpanded, setIsComposerExpanded] = useState(false); const inputBarContainerRef = useRef(null); @@ -259,6 +261,7 @@ export const InputbarCore: React.FC = ({ showExecutionStrategy={showExecutionStrategy} toolMode={toolMode} isCanvasOpen={isCanvasOpen} + activeTheme={activeTheme} /> ) : null} diff --git a/src/components/agent/chat/components/Inputbar/components/InputbarTools.tsx b/src/components/agent/chat/components/Inputbar/components/InputbarTools.tsx index 5ddffa9f2..030c4c7e5 100644 --- a/src/components/agent/chat/components/Inputbar/components/InputbarTools.tsx +++ b/src/components/agent/chat/components/Inputbar/components/InputbarTools.tsx @@ -4,6 +4,8 @@ import { Lightbulb, Globe, Code2, + ListChecks, + Workflow, } from "lucide-react"; import { ToolButton } from "../styles"; import { @@ -12,6 +14,7 @@ import { TooltipProvider, TooltipTrigger, } from "@/components/ui/tooltip"; +import { isGeneralResearchTheme } from "../../../utils/generalAgentPrompt"; interface InputbarToolsProps { onToolClick?: (tool: string) => void; @@ -21,6 +24,7 @@ interface InputbarToolsProps { toolMode?: "default" | "attach-only"; /** 画布是否打开(兼容保留,不再展示画布图标) */ isCanvasOpen?: boolean; + activeTheme?: string; } export const InputbarTools: React.FC = ({ @@ -29,6 +33,7 @@ export const InputbarTools: React.FC = ({ executionStrategy = "react", showExecutionStrategy = false, toolMode = "default", + activeTheme, }) => { const modeLabel = executionStrategy === "auto" @@ -38,6 +43,7 @@ export const InputbarTools: React.FC = ({ : "ReAct"; const strategyEnabled = executionStrategy !== "react" || activeTools["execution_strategy"]; + const isGeneralTheme = isGeneralResearchTheme(activeTheme); return ( @@ -85,6 +91,46 @@ export const InputbarTools: React.FC = ({ + {isGeneralTheme ? ( + <> + + + onToolClick?.("task_mode")} + className={activeTools["task_mode"] ? "active" : ""} + > + + + + + 后台任务 {activeTools["task_mode"] ? "(偏好已开启)" : ""} + + + + + + onToolClick?.("subagent_mode")} + className={activeTools["subagent_mode"] ? "active" : ""} + > + + + + + 多代理 {activeTools["subagent_mode"] ? "(偏好已开启)" : ""} + + + + ) : null} + {showExecutionStrategy && ( diff --git a/src/components/agent/chat/components/Inputbar/hooks/useInputbarController.ts b/src/components/agent/chat/components/Inputbar/hooks/useInputbarController.ts index ed2d3f6f7..b762bb116 100644 --- a/src/components/agent/chat/components/Inputbar/hooks/useInputbarController.ts +++ b/src/components/agent/chat/components/Inputbar/hooks/useInputbarController.ts @@ -97,6 +97,8 @@ export function useInputbarController({ handleToolClick, isFullscreen, thinkingEnabled, + taskEnabled, + subagentEnabled, webSearchEnabled, } = useInputbarToolState({ toolStates, @@ -195,6 +197,10 @@ export function useInputbarController({ handleSend, inputAdapter, topExtra, + taskEnabled, + subagentEnabled, + thinkingEnabled, + webSearchEnabled, themeWorkbenchQuickActions, themeWorkbenchQueueItems, renderThemeWorkbenchGeneratingPanel, diff --git a/src/components/agent/chat/components/Inputbar/hooks/useInputbarToolState.ts b/src/components/agent/chat/components/Inputbar/hooks/useInputbarToolState.ts index de35c5d34..c96c86a19 100644 --- a/src/components/agent/chat/components/Inputbar/hooks/useInputbarToolState.ts +++ b/src/components/agent/chat/components/Inputbar/hooks/useInputbarToolState.ts @@ -4,6 +4,8 @@ import { toast } from "sonner"; export interface InputbarToolStates { webSearch: boolean; thinking: boolean; + task: boolean; + subagent: boolean; } interface UseInputbarToolStateParams { @@ -23,6 +25,8 @@ interface UseInputbarToolStateParams { const DEFAULT_INPUTBAR_TOOL_STATES: InputbarToolStates = { webSearch: false, thinking: false, + task: false, + subagent: false, }; export function useInputbarToolState({ @@ -47,14 +51,24 @@ export function useInputbarToolState({ const webSearchEnabled = toolStates?.webSearch ?? localToolStates.webSearch; const thinkingEnabled = toolStates?.thinking ?? localToolStates.thinking; + const taskEnabled = toolStates?.task ?? localToolStates.task; + const subagentEnabled = toolStates?.subagent ?? localToolStates.subagent; const activeTools = useMemo>( () => ({ ...localActiveTools, web_search: webSearchEnabled, thinking: thinkingEnabled, + task_mode: taskEnabled, + subagent_mode: subagentEnabled, }), - [localActiveTools, thinkingEnabled, webSearchEnabled], + [ + localActiveTools, + thinkingEnabled, + webSearchEnabled, + taskEnabled, + subagentEnabled, + ], ); const updateToolStates = useCallback( @@ -62,11 +76,19 @@ export function useInputbarToolState({ setLocalToolStates((prev) => ({ webSearch: toolStates?.webSearch ?? next.webSearch ?? prev.webSearch, thinking: toolStates?.thinking ?? next.thinking ?? prev.thinking, + task: toolStates?.task ?? next.task ?? prev.task, + subagent: toolStates?.subagent ?? next.subagent ?? prev.subagent, })); onToolStatesChange?.(next); return next; }, - [onToolStatesChange, toolStates?.thinking, toolStates?.webSearch], + [ + onToolStatesChange, + toolStates?.subagent, + toolStates?.task, + toolStates?.thinking, + toolStates?.webSearch, + ], ); const handleToolClick = useCallback( @@ -77,6 +99,8 @@ export function useInputbarToolState({ updateToolStates({ webSearch: webSearchEnabled, thinking: nextThinking, + task: taskEnabled, + subagent: subagentEnabled, }); toast.info(`深度思考${nextThinking ? "已开启" : "已关闭"}`); break; @@ -86,10 +110,34 @@ export function useInputbarToolState({ updateToolStates({ webSearch: nextWebSearch, thinking: thinkingEnabled, + task: taskEnabled, + subagent: subagentEnabled, }); toast.info(`联网搜索${nextWebSearch ? "已开启" : "已关闭"}`); break; } + case "task_mode": { + const nextTask = !taskEnabled; + updateToolStates({ + webSearch: webSearchEnabled, + thinking: thinkingEnabled, + task: nextTask, + subagent: subagentEnabled, + }); + toast.info(`后台任务${nextTask ? "偏好已开启" : "偏好已关闭"}`); + break; + } + case "subagent_mode": { + const nextSubagent = !subagentEnabled; + updateToolStates({ + webSearch: webSearchEnabled, + thinking: thinkingEnabled, + task: taskEnabled, + subagent: nextSubagent, + }); + toast.info(`多代理${nextSubagent ? "偏好已开启" : "偏好已关闭"}`); + break; + } case "execution_strategy": if (setExecutionStrategy) { const strategyOrder: Array< @@ -154,6 +202,8 @@ export function useInputbarToolState({ setExecutionStrategy, setInput, thinkingEnabled, + subagentEnabled, + taskEnabled, updateToolStates, webSearchEnabled, ], @@ -164,6 +214,8 @@ export function useInputbarToolState({ handleToolClick, isFullscreen, thinkingEnabled, + taskEnabled, + subagentEnabled, webSearchEnabled, }; } diff --git a/src/components/agent/chat/components/Inputbar/index.test.tsx b/src/components/agent/chat/components/Inputbar/index.test.tsx index ee8bc60b8..9ba657983 100644 --- a/src/components/agent/chat/components/Inputbar/index.test.tsx +++ b/src/components/agent/chat/components/Inputbar/index.test.tsx @@ -254,6 +254,8 @@ describe("Inputbar", () => { expect(onToolStatesChange).toHaveBeenCalledWith({ webSearch: true, thinking: false, + task: false, + subagent: false, }); }); diff --git a/src/components/agent/chat/components/MarkdownRenderer.tsx b/src/components/agent/chat/components/MarkdownRenderer.tsx index 4b749c721..0c3a288b9 100644 --- a/src/components/agent/chat/components/MarkdownRenderer.tsx +++ b/src/components/agent/chat/components/MarkdownRenderer.tsx @@ -233,6 +233,8 @@ interface MarkdownRendererProps { content: string; /** A2UI 表单提交回调 */ onA2UISubmit?: (formData: A2UIFormData) => void; + /** 是否渲染消息内联 A2UI */ + renderA2UIInline?: boolean; /** 是否折叠代码块(当画布打开时) */ collapseCodeBlocks?: boolean; /** 代码块点击回调(用于在画布中显示) */ @@ -245,6 +247,7 @@ export const MarkdownRenderer: React.FC = memo( ({ content, onA2UISubmit, + renderA2UIInline = true, collapseCodeBlocks = false, onCodeBlockClick, isStreaming = false, @@ -406,6 +409,10 @@ export const MarkdownRenderer: React.FC = memo( // 如果是 a2ui 代码块,特殊处理 if (language === "a2ui") { + if (!renderA2UIInline) { + return null; + } + const parsed = parseA2UIJson(codeContent); if (parsed) { @@ -425,7 +432,7 @@ export const MarkdownRenderer: React.FC = memo( subtitle="正在解析结构化问题,请稍等。" /> ); - } + } } // 如果启用了代码块折叠,显示占位符卡片 diff --git a/src/components/agent/chat/components/MessageList.test.tsx b/src/components/agent/chat/components/MessageList.test.tsx index 0a3d0466e..b296c9bf3 100644 --- a/src/components/agent/chat/components/MessageList.test.tsx +++ b/src/components/agent/chat/components/MessageList.test.tsx @@ -11,10 +11,22 @@ vi.mock("./MarkdownRenderer", () => ({ ), })); -vi.mock("./StreamingRenderer", () => ({ - StreamingRenderer: ({ content }: { content: string }) => ( +const mockStreamingRenderer = vi.fn( + ({ + content, + }: { + content: string; + renderA2UIInline?: boolean; + }) => (
{content || ""}
), +); + +vi.mock("./StreamingRenderer", () => ({ + StreamingRenderer: (props: { + content: string; + renderA2UIInline?: boolean; + }) => mockStreamingRenderer(props), })); vi.mock("./TokenUsageDisplay", () => ({ @@ -51,13 +63,16 @@ afterEach(() => { vi.clearAllMocks(); }); -function render(messages: Message[]): HTMLDivElement { +function render( + messages: Message[], + props?: { renderA2UIInline?: boolean }, +): HTMLDivElement { const container = document.createElement("div"); document.body.appendChild(container); const root = createRoot(container); act(() => { - root.render(); + root.render(); }); mountedRoots.push({ container, root }); @@ -100,4 +115,82 @@ describe("MessageList", () => { ).map((node) => node.textContent); expect(streamingTexts).toEqual(["好的,我继续处理。"]); }); + + it("应向助手消息透传内联 A2UI 开关", () => { + const now = new Date(); + const messages: Message[] = [ + { + id: "msg-assistant", + role: "assistant", + content: "```a2ui\n{}\n```", + timestamp: now, + }, + ]; + + render(messages); + expect(mockStreamingRenderer).toHaveBeenCalledWith( + expect.objectContaining({ renderA2UIInline: true }), + ); + + render(messages, { renderA2UIInline: false }); + expect(mockStreamingRenderer).toHaveBeenLastCalledWith( + expect.objectContaining({ renderA2UIInline: false }), + ); + }); + + it("助手消息包含 artifacts 时应渲染产物卡片并响应点击", () => { + const now = new Date(); + const onArtifactClick = vi.fn(); + const messages: Message[] = [ + { + id: "msg-assistant-artifact", + role: "assistant", + content: "已生成文档", + timestamp: now, + artifacts: [ + { + id: "artifact-demo", + type: "document", + title: "demo.md", + content: "# Demo", + status: "complete", + meta: { + filePath: "docs/demo.md", + filename: "demo.md", + }, + position: { start: 0, end: 0 }, + createdAt: now.getTime(), + updatedAt: now.getTime(), + }, + ], + }, + ]; + + const container = document.createElement("div"); + document.body.appendChild(container); + const root = createRoot(container); + + act(() => { + root.render( + , + ); + }); + + mountedRoots.push({ container, root }); + + const artifactCard = container.querySelector("button"); + expect(artifactCard?.textContent).toContain("demo.md"); + expect(artifactCard?.textContent).toContain("docs/demo.md"); + + act(() => { + artifactCard?.dispatchEvent(new MouseEvent("click", { bubbles: true })); + }); + + expect(onArtifactClick).toHaveBeenCalledWith( + expect.objectContaining({ + id: "artifact-demo", + title: "demo.md", + }), + ); + }); }); diff --git a/src/components/agent/chat/components/MessageList.tsx b/src/components/agent/chat/components/MessageList.tsx index 5c4f1b24f..aa4ca34e9 100644 --- a/src/components/agent/chat/components/MessageList.tsx +++ b/src/components/agent/chat/components/MessageList.tsx @@ -1,7 +1,17 @@ import React, { useState, useRef, useEffect, useMemo } from "react"; -import { User, Copy, Edit2, Trash2, Check } from "lucide-react"; +import { + User, + Copy, + Edit2, + Trash2, + Check, + FileText, + Loader2, + ExternalLink, +} from "lucide-react"; import { Button } from "@/components/ui/button"; import { toast } from "sonner"; +import type { Artifact } from "@/lib/artifact/types"; import { MessageListContainer, MessageWrapper, @@ -17,25 +27,44 @@ import { import { MarkdownRenderer } from "./MarkdownRenderer"; import { StreamingRenderer } from "./StreamingRenderer"; import { TokenUsageDisplay } from "./TokenUsageDisplay"; -import { Message } from "../types"; +import { AgentThreadTimeline } from "./AgentThreadTimeline"; +import { + Message, + type AgentThreadItem, + type AgentThreadTurn, + type WriteArtifactContext, +} from "../types"; import type { A2UIFormData } from "@/components/content-creator/a2ui/types"; import type { ConfirmResponse } from "../types"; +import { buildMessageTurnTimeline } from "../utils/threadTimelineView"; import logoImg from "/logo.png"; interface MessageListProps { messages: Message[]; + turns?: AgentThreadTurn[]; + threadItems?: AgentThreadItem[]; + currentTurnId?: string | null; + assistantLabel?: string; onDeleteMessage?: (id: string) => void; onEditMessage?: (id: string, content: string) => void; /** A2UI 表单提交回调 */ onA2UISubmit?: (formData: A2UIFormData, messageId: string) => void; + /** 是否渲染消息内联 A2UI */ + renderA2UIInline?: boolean; /** A2UI 表单数据映射(按消息 ID 索引) */ a2uiFormDataMap?: Record; /** A2UI 表单数据变化回调(用于持久化) */ onA2UIFormChange?: (formId: string, formData: A2UIFormData) => void; /** 文件写入回调 */ - onWriteFile?: (content: string, fileName: string) => void; + onWriteFile?: ( + content: string, + fileName: string, + context?: WriteArtifactContext, + ) => void; /** 文件点击回调 */ onFileClick?: (fileName: string, content: string) => void; + /** Artifact 点击回调 */ + onArtifactClick?: (artifact: Artifact) => void; /** 权限确认响应回调 */ onPermissionResponse?: (response: ConfirmResponse) => void; /** 是否折叠代码块(当画布打开时) */ @@ -46,13 +75,19 @@ interface MessageListProps { const MessageListInner: React.FC = ({ messages, + turns = [], + threadItems = [], + currentTurnId = null, + assistantLabel = "ProxyCast", onDeleteMessage, onEditMessage, onA2UISubmit, + renderA2UIInline = true, a2uiFormDataMap, onA2UIFormChange, onWriteFile, onFileClick, + onArtifactClick, onPermissionResponse, collapseCodeBlocks, onCodeBlockClick, @@ -74,6 +109,10 @@ const MessageListInner: React.FC = ({ }), [messages], ); + const timelineByMessageId = useMemo( + () => buildMessageTurnTimeline(visibleMessages, turns, threadItems), + [threadItems, turns, visibleMessages], + ); // 检测用户是否在手动滚动 useEffect(() => { @@ -153,6 +192,60 @@ const MessageListInner: React.FC = ({ } }; + const renderArtifactCards = (artifacts: Artifact[] | undefined) => { + if (!artifacts || artifacts.length === 0) { + return null; + } + + return ( +
+ {artifacts.map((artifact) => { + const filePath = + typeof artifact.meta.filePath === "string" + ? artifact.meta.filePath + : artifact.meta.filename || artifact.title; + const statusLabel = + artifact.status === "streaming" + ? "生成中" + : artifact.status === "error" + ? "失败" + : "已生成"; + + return ( + + ); + })} +
+ ); + }; + return (
@@ -167,8 +260,11 @@ const MessageListInner: React.FC = ({
)} - {visibleMessages.map((msg) => ( - + {visibleMessages.map((msg) => { + const timeline = timelineByMessageId.get(msg.id); + + return ( + {msg.role === "user" ? ( @@ -193,12 +289,22 @@ const MessageListInner: React.FC = ({ - {msg.role === "user" ? "用户" : "ProxyCast"} + {msg.role === "user" ? "用户" : assistantLabel} {formatTime(msg.timestamp)} + {msg.role === "assistant" && timeline ? ( + + ) : null} + {editingId === msg.id ? (