feat: release v0.85.0 with full pending changes

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
coso
2026-03-13 02:40:20 +08:00
co-authored by Claude Opus 4.6
parent 621ab3ee73
commit 8995ef9c68
259 changed files with 24283 additions and 18539 deletions
+18 -33
View File
@@ -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
+1
View File
@@ -16,6 +16,7 @@
- `develop/`:开发流程与协作规范
- `plugins/`:插件与扩展相关文档
- `tests/`:测试策略与用例文档
- `iteration-notes/`:迭代备忘与下版本建议(暂不进入当前发布范围的问题)
- `images/`:文档图片资源
- `TECH_SPEC.md`:技术规格文档
- `develop/execution-tracker-technical-plan.md`:统一执行追踪(Execution Tracker)专项技术规划
+41 -37
View File
@@ -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 刷新
- [ ] 请求日志/统计
+7
View File
@@ -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
+1 -1
View File
@@ -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 {
+57 -13
View File
@@ -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<i64>) -> Result<u64, String>;
## 前端调用
推荐写法不是在业务层直接 `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<CredentialInfo>('add_credential', {
provider: 'kiro',
filePath: '/path/to/credential.json',
});
export async function getServerStatus() {
return safeInvoke<ServerStatus>("get_server_status");
}
```
// 获取服务器状态
const status = await invoke<ServerStatus>('get_server_status');
业务层只消费 API 网关:
// 查询流量记录
const records = await invoke<PagedResult<FlowRecord>>('get_flow_records', {
query: { page: 1, pageSize: 20 },
});
```typescript
import { getServerStatus } from "@/lib/api/serverRuntime";
const status = await getServerStatus();
```
## 错误处理
+36 -40
View File
@@ -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 (
<aside className="w-14 bg-sidebar">
<nav className="flex flex-col items-center gap-2">
<SidebarItem icon={Home} to="/" />
<SidebarItem icon={MessageSquare} to="/chat" />
<SidebarItem icon={Settings} to="/settings" />
</nav>
</aside>
);
return (
<aside className="w-14 bg-sidebar">
<nav className="flex flex-col items-center gap-2">
<SidebarItem icon={Home} to="/" />
<SidebarItem icon={MessageSquare} to="/chat" />
<SidebarItem icon={Settings} to="/settings" />
</nav>
</aside>
);
}
```
@@ -46,14 +46,14 @@ export function AppSidebar() {
```tsx
// src/components/provider-pool/ProviderPoolPanel.tsx
export function ProviderPoolPanel() {
const { credentials, addCredential, removeCredential } = useProviderPool();
return (
<div className="space-y-4">
<CredentialList credentials={credentials} onRemove={removeCredential} />
<AddCredentialDialog onAdd={addCredential} />
</div>
);
const { credentials, addCredential, removeCredential } = useProviderPool();
return (
<div className="space-y-4">
<CredentialList credentials={credentials} onRemove={removeCredential} />
<AddCredentialDialog onAdd={addCredential} />
</div>
);
}
```
@@ -64,15 +64,15 @@ export function ProviderPoolPanel() {
```tsx
// src/components/flow-monitor/FlowMonitorPanel.tsx
export function FlowMonitorPanel() {
const { records, stats, query } = useFlowMonitor();
return (
<div className="flex flex-col h-full">
<FlowStats stats={stats} />
<FlowTable records={records} />
<FlowPagination query={query} />
</div>
);
const { records, stats, query } = useFlowMonitor();
return (
<div className="flex flex-col h-full">
<FlowStats stats={stats} />
<FlowTable records={records} />
<FlowPagination query={query} />
</div>
);
}
```
@@ -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 (
<div>
{/* JSX */}
</div>
);
// hooks
const [state, setState] = useState();
// handlers
const handleClick = () => {};
// render
return <div>{/* JSX */}</div>;
}
```
+9 -3
View File
@@ -17,7 +17,7 @@ AI 返回带 <write_file> 标签的响应
↓
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,因此采用**消息注入**方案:
+28
View File
@@ -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_*` 模块名与函数名。
## 一句话总结
+91 -59
View File
@@ -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`
- 社媒内容推荐把 `<write_file>` 结果投影为“版本链产物”,而不是只按文件名覆盖
**相关文件**:
- 类型定义:`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<Credential[]>([]);
const [loading, setLoading] = useState(false);
const refresh = async () => {
setLoading(true);
const list = await invoke<Credential[]>('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<Credential[]>([]);
const [loading, setLoading] = useState(false);
const refresh = async () => {
setLoading(true);
const list = await invoke<Credential[]>("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<FlowRecord[]>([]);
useEffect(() => {
const unlisten = listen<FlowEvent>('flow-event', (event) => {
setRecords(prev => [event.payload.data, ...prev].slice(0, 100));
});
return () => { unlisten.then(fn => fn()); };
}, []);
return { records };
const [records, setRecords] = useState<FlowRecord[]>([]);
useEffect(() => {
const unlisten = listen<FlowEvent>("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<string>('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<string>("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,
};
```
+4 -1
View File
@@ -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
+174 -31
View File
@@ -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",
+3 -2
View File
@@ -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",
File diff suppressed because it is too large Load Diff
+16 -16
View File
@@ -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",
+2 -2
View File
@@ -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"
@@ -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<T: serde::Serialize>(
#[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 },
+7
View File
@@ -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;
@@ -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<TauriMessage>,
pub execution_strategy: Option<String>,
pub turns: Vec<AgentThreadTurn>,
pub items: Vec<AgentThreadItem>,
}
/// 解析会话 working_dir(优先入参,其次 workspace_id)
@@ -145,6 +151,10 @@ pub fn get_session_sync(db: &DbConnection, session_id: &str) -> Result<SessionDe
let messages =
AgentDao::get_messages(&conn, session_id).map_err(|e| format!("获取消息失败: {e}"))?;
let turns = AgentTimelineDao::list_turns_by_thread(&conn, session_id)
.map_err(|e| format!("获取 turn 历史失败: {e}"))?;
let items = AgentTimelineDao::list_items_by_thread(&conn, session_id)
.map_err(|e| format!("获取 item 历史失败: {e}"))?;
let tauri_messages = convert_agent_messages(&messages, Some(session.model.as_str()));
@@ -163,8 +173,11 @@ pub fn get_session_sync(db: &DbConnection, session_id: &str) -> Result<SessionDe
updated_at: chrono::DateTime::parse_from_rfc3339(&session.updated_at)
.map(|dt| dt.timestamp())
.unwrap_or(0),
thread_id: session_id.to_string(),
messages: tauri_messages,
execution_strategy: session.execution_strategy,
turns,
items,
})
}
@@ -16,7 +16,6 @@ use aster::context::{
DEFAULT_CONTEXT_WINDOW_TRIGGER_RATIO, DEFAULT_TOOL_IO_PREVIEW_MAX_CHARS,
DEFAULT_TOOL_IO_PREVIEW_MAX_LINES, DEFAULT_TOOL_TOKEN_LIMIT_BEFORE_EVICT,
};
use chrono::Utc;
use proxycast_core::agent::types::AgentMessage;
use serde::Serialize;
use serde_json::{json, Map, Value};
+201 -7
View File
@@ -39,21 +39,111 @@ pub fn legacy_database_path() -> Result<PathBuf, String> {
}
pub fn resolve_database_path() -> Result<PathBuf, String> {
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<PathBuf, String> {
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<PathBuf, String> {
resolve_runtime_subdir("request_logs")
}
pub fn resolve_projects_dir() -> Result<PathBuf, String> {
resolve_runtime_subdir("projects")
}
pub fn resolve_sessions_dir() -> Result<PathBuf, String> {
resolve_runtime_subdir("sessions")
}
pub fn resolve_skills_dir() -> Result<PathBuf, String> {
resolve_runtime_subdir("skills")
}
pub fn resolve_user_memory_path() -> Result<PathBuf, String> {
with_app_roots(resolve_user_memory_path_from_roots)
}
pub fn resolve_default_project_dir() -> Result<PathBuf, String> {
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<T>(
resolver: impl FnOnce(&Path, &Path) -> Result<T, String>,
) -> Result<T, String> {
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<PathBuf, String> {
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<PathBuf, String> {
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<PathBuf, String> {
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();
+3 -4
View File
@@ -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 |
@@ -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<Vec<AgentRun>, 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");
}
}
@@ -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<Self, Self::Error> {
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<Self, Self::Error> {
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<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct AgentRequestQuestion {
pub question: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub header: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub options: Option<Vec<AgentRequestOption>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub multi_select: Option<bool>,
}
#[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<String>,
},
Plan {
text: String,
},
Reasoning {
text: String,
#[serde(skip_serializing_if = "Option::is_none")]
summary: Option<Vec<String>>,
},
ToolCall {
tool_name: String,
#[serde(skip_serializing_if = "Option::is_none")]
arguments: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
output: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
success: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
error: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
metadata: Option<serde_json::Value>,
},
CommandExecution {
command: String,
cwd: String,
#[serde(skip_serializing_if = "Option::is_none")]
aggregated_output: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
exit_code: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
error: Option<String>,
},
WebSearch {
#[serde(skip_serializing_if = "Option::is_none")]
query: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
action: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
output: Option<String>,
},
ApprovalRequest {
request_id: String,
action_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
prompt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
arguments: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
response: Option<serde_json::Value>,
},
RequestUserInput {
request_id: String,
action_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
prompt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
questions: Option<Vec<AgentRequestQuestion>>,
#[serde(skip_serializing_if = "Option::is_none")]
response: Option<serde_json::Value>,
},
FileArtifact {
path: String,
source: String,
#[serde(skip_serializing_if = "Option::is_none")]
content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
metadata: Option<serde_json::Value>,
},
SubagentActivity {
status_label: String,
#[serde(skip_serializing_if = "Option::is_none")]
title: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
summary: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
role: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
model: Option<String>,
},
Warning {
message: String,
#[serde(skip_serializing_if = "Option::is_none")]
code: Option<String>,
},
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<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub error_message: Option<String>,
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<String>,
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<Vec<AgentThreadTurn>, 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<Option<AgentThreadItem>, 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<Vec<AgentThreadItem>, 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<AgentThreadItem, rusqlite::Error> {
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"));
}
}
@@ -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<Vec<ChatSession>, 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<String> = 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<Option<ChatSession>, 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<String> = 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<bool, rusqlite::Error> {
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<bool, rusqlite::Error> {
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<bool, rusqlite::Error> {
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<i32>,
before_id: Option<&str>,
) -> Result<Vec<ChatMessage>, 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<ChatMessage> = messages.collect::<Result<Vec<_>, _>>()?;
// 如果有 limit,结果是倒序的,需要反转
if limit.is_some() {
result.reverse();
}
Ok(result)
}
/// 获取会话的消息数量
pub fn get_message_count(conn: &Connection, session_id: &str) -> Result<i64, rusqlite::Error> {
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<ChatMessage, rusqlite::Error> {
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<String> = row.get(4)?;
let blocks: Option<Vec<ContentBlock>> = 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<String> = 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");
}
}
@@ -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;
+28 -932
View File
@@ -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<usize, String> {
// 检查是否已经迁移过
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::<String>::None, // check_model_name
"[]", // not_supported_models
row.usage_count as i64,
row.error_count as i32,
last_used_ts,
Option::<i64>::None, // last_error_time
Option::<String>::None, // last_error_message
Option::<i64>::None, // last_health_check_time
Option::<String>::None, // last_health_check_model
created_at_ts,
now,
"imported", // source: 标记为导入
Option::<String>::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<String>,
enabled: bool,
usage_count: u64,
error_count: u32,
last_used_at: Option<String>,
created_at: Option<String>,
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<usize, String> {
// 检查是否已经迁移过
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<usize, String> {
// 检查是否已经清理过
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<String>>(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<usize, String> {
// 检查是否已经迁移过
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<usize, String> {
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<String> = 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<usize, String> {
// 检查是否已经迁移过
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<usize, String> {
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<String>)> = 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::<String>::None,
name,
created_str,
updated_str,
],
)
.map_err(|e| format!("插入会话失败: {e}"))?;
count += 1;
}
Ok(count)
}
/// 迁移 General Chat 消息数据
fn migrate_general_messages(conn: &Connection) -> Result<usize, String> {
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>,
String,
i64,
Option<String>,
)> = 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::<String>::None,
Option::<String>::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>) -> String {
// 如果有 blocks,尝试解析并转换
if let Some(blocks_str) = blocks {
if let Ok(blocks_arr) = serde_json::from_str::<Vec<serde_json::Value>>(blocks_str) {
let converted: Vec<serde_json::Value> = 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,
}
@@ -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<usize, String> {
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::<String>::None,
"[]",
row.usage_count as i64,
row.error_count as i32,
last_used_ts,
Option::<i64>::None,
Option::<String>::None,
Option::<i64>::None,
Option::<String>::None,
created_at_ts,
now,
"imported",
Option::<String>::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<usize, String> {
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<usize, String> {
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<String>>(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<String>,
enabled: bool,
usage_count: u64,
error_count: u32,
last_used_at: Option<String>,
created_at: Option<String>,
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<String>, 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);
}
}
@@ -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<usize, String> {
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<usize, String> {
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<String>)> = 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::<String>::None,
name,
created_str,
updated_str,
],
)
.map_err(|e| format!("插入会话失败: {e}"))?;
count += 1;
}
Ok(count)
}
fn migrate_general_messages(conn: &Connection) -> Result<usize, String> {
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>,
String,
i64,
Option<String>,
)> = 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,
&timestamp_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::<String>::None,
Option::<String>::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<bool, String> {
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>) -> String {
if let Some(blocks_str) = blocks {
if let Ok(blocks_arr) = serde_json::from_str::<Vec<serde_json::Value>>(blocks_str) {
let converted: Vec<serde_json::Value> = 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);
}
}
@@ -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<usize, String> {
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<usize, String> {
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");
}
}
@@ -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));
}
}
@@ -0,0 +1,116 @@
use rusqlite::{params, Connection};
pub(crate) fn read_setting_value(conn: &Connection, key: &str) -> Option<String> {
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<T, F>(conn: &Connection, operation: F) -> Result<T, String>
where
F: FnOnce(&Connection) -> Result<T, String>,
{
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);
}
}
@@ -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<MigrationResult, String> {
// 检查是否已经迁移过
migrate_unified_content_system_with_default_dir_resolver(
conn,
&app_paths::resolve_default_project_dir,
)
}
fn migrate_unified_content_system_with_default_dir_resolver<F>(
conn: &Connection,
resolve_default_project_dir: &F,
) -> Result<MigrationResult, String>
where
F: Fn() -> Result<std::path::PathBuf, String>,
{
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<MigrationResu
tracing::info!("[迁移] 开始执行统一内容系统迁移");
// 开始事务
conn.execute("BEGIN TRANSACTION", [])
.map_err(|e| format!("开始事务失败: {e}"))?;
// 执行迁移
let result = execute_migration(conn);
match result {
match run_in_transaction(conn, |tx| {
let stats = execute_migration(tx, resolve_default_project_dir)?;
mark_migration_completed(tx, MIGRATION_KEY_UNIFIED_CONTENT)?;
Ok(stats)
}) {
Ok(stats) => {
// 标记迁移完成
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<MigrationStats, String> {
fn execute_migration<F>(
conn: &Connection,
resolve_default_project_dir: &F,
) -> Result<MigrationStats, String>
where
F: Fn() -> Result<std::path::PathBuf, String>,
{
// 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<MigrationStats, String> {
/// 否则创建新的默认项目
///
/// _Requirements: 2.1_
fn get_or_create_default_project(conn: &Connection) -> Result<String, String> {
fn get_or_create_default_project<F>(
conn: &Connection,
resolve_default_project_dir: &F,
) -> Result<String, String>
where
F: Fn() -> Result<std::path::PathBuf, String>,
{
// 检查是否已存在默认项目
let existing_id: Option<String> = conn
.query_row(
@@ -116,7 +132,7 @@ fn get_or_create_default_project(conn: &Connection) -> Result<String, String> {
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<String, String> {
}
/// 获取默认项目的存储路径
fn get_default_project_path() -> Result<String, String> {
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<F>(resolve_default_project_dir: &F) -> Result<String, String>
where
F: Fn() -> Result<std::path::PathBuf, String>,
{
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<String> {
///
/// 如果不存在则创建,返回默认项目 ID
pub fn ensure_default_project(conn: &Connection) -> Result<String, String> {
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();
@@ -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<MigrationResult, String> {
// 检查是否已经迁移过
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<MigrationResul
});
}
// 开始事务
conn.execute("BEGIN TRANSACTION", [])
.map_err(|e| format!("开始事务失败: {e}"))?;
// 执行迁移
let result = execute_playwright_migration(conn);
match result {
match run_in_transaction(conn, |tx| {
let server_id = execute_playwright_migration(tx)?;
mark_migration_completed(tx, MIGRATION_KEY_PLAYWRIGHT_SERVER)?;
Ok(server_id)
}) {
Ok(server_id) => {
// 标记迁移完成
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<MigrationResul
server_id: Some(server_id),
})
}
Err(e) => {
// 回滚事务
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(())
}
@@ -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<MigrationResult, String> {
// 检查是否已经迁移过
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<MigrationResult, S
});
}
// 获取用户主目录
let home = dirs::home_dir()
.map(|p| p.to_string_lossy().to_string())
.unwrap_or_else(|| "/Users/unknown".to_string());
let default_path = format!("{}/{}", home, DEFAULT_PROJECTS_DIR);
let default_path = app_paths::resolve_projects_dir()?
.to_string_lossy()
.to_string();
// 开始事务
conn.execute("BEGIN TRANSACTION", [])
.map_err(|e| format!("开始事务失败: {e}"))?;
let result = execute_migration(conn, &default_path, promise_done, unify_done);
match result {
match run_in_transaction(conn, |tx| {
let result = execute_migration(tx, &default_path, promise_done, unify_done)?;
if !promise_done {
mark_migration_completed(tx, MIGRATION_KEY_FIX_PROMISE_PATHS)?;
}
if !unify_done {
mark_migration_completed(tx, MIGRATION_KEY_UNIFY_SESSION_DIRS)?;
}
Ok(result)
}) {
Ok((fixed_ws, fixed_sess, unified_sess)) => {
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<MigrationResult, S
unified_sessions: unified_sess,
})
}
Err(e) => {
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(())
}
+146 -118
View File
@@ -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<Mutex<Connection>>;
#[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<pending_general_chat::PendingGeneralMessageRow> 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<T, F>(
conn: &Connection,
empty_value: T,
query: F,
) -> Result<T, rusqlite::Error>
where
F: FnOnce(&Connection) -> Result<T, rusqlite::Error>,
{
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<Vec<PendingGeneralMessage>, 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<i64>,
to_timestamp_ms: Option<i64>,
limit: usize,
) -> Result<Vec<PendingGeneralMessage>, 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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, rusqlite::Error> {
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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, rusqlite::Error> {
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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, rusqlite::Error> {
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<std::sync::MutexGuard<'_, Connection>, String> {
match db.lock() {
@@ -53,124 +161,44 @@ pub fn init_database() -> Result<DbConnection, String> {
// 创建表结构
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);
}
}
@@ -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<i32>,
before_id: Option<&str>,
) -> Result<Vec<ChatMessage>, 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::<Result<Vec<_>, _>>()?;
if limit.is_some() {
messages.reverse();
}
Ok(messages)
}
fn has_pending_general_messages_table(conn: &Connection) -> Result<bool, rusqlite::Error> {
table_exists(conn, "general_chat_messages")
}
fn has_pending_general_sessions_table(conn: &Connection) -> Result<bool, rusqlite::Error> {
table_exists(conn, "general_chat_sessions")
}
pub(super) fn load_pending_general_session_messages_raw(
conn: &Connection,
session_id: &str,
) -> Result<Vec<PendingGeneralMessageRow>, 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<i64>,
to_timestamp_ms: Option<i64>,
limit: usize,
) -> Result<Vec<PendingGeneralMessageRow>, 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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, rusqlite::Error> {
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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, rusqlite::Error> {
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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, rusqlite::Error> {
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<PendingGeneralMessageRow, rusqlite::Error> {
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<ChatMessage, rusqlite::Error> {
let role_str: String = row.get(2)?;
let blocks_json: Option<String> = row.get(4)?;
let blocks: Option<Vec<ContentBlock>> = 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<String> = 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<bool, rusqlite::Error> {
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::<String>::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::<String>::None,
"complete",
now + index as i64,
Option::<String>::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::<Vec<_>>();
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());
}
}
@@ -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 相关表
// ============================================================================
@@ -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<T, F, S>(
conn: &Connection,
failure_label: &str,
operation: F,
on_success: S,
) where
F: FnOnce(&Connection) -> Result<T, String>,
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<T, F, S>(
conn: &Connection,
failure_label: &str,
operation: F,
on_success: S,
) where
F: FnOnce(&Connection) -> Result<T, String>,
S: FnOnce(&Connection, T) -> Option<String>,
{
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<F, S>(
conn: &Connection,
failure_label: &str,
operation: F,
on_nonzero: S,
) where
F: FnOnce(&Connection) -> Result<usize, String>,
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());
}
}
+1 -6
View File
@@ -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();
@@ -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;
@@ -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<Self, String> {
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<PathBuf, String> {
let home = dirs::home_dir().ok_or("无法获取用户主目录")?;
Ok(home.join(".proxycast").join("sessions"))
app_paths::resolve_sessions_dir()
}
/// 获取会话目录路径
+2
View File
@@ -2,6 +2,8 @@
//!
//! 承载渠道运行时(channel runtime)、路由与策略实现。
#![allow(clippy::all)]
pub mod discord;
pub mod feishu;
pub mod telegram;
+4
View File
@@ -3,6 +3,10 @@
//! 提供统一的请求处理管道,集成路由、容错、监控、插件等功能模块。
//!
//! ## 模块结构
#![allow(clippy::derivable_impls)]
#![allow(clippy::unnecessary_map_or)]
#![allow(clippy::too_many_arguments)]
//!
//! - `steps` - 管道步骤(认证、注入、路由、插件、Provider、遥测)
+5
View File
@@ -3,6 +3,11 @@
//! 提供 Agent 任务调度功能,支持定时任务、重试机制等。
//!
//! ## 功能
#![allow(clippy::redundant_closure)]
#![allow(clippy::format_in_format_args)]
#![allow(clippy::useless_format)]
#![allow(clippy::derivable_impls)]
//! - 任务创建和管理
//! - 任务持久化到 SQLite
//! - 定时任务调度
+2
View File
@@ -1,5 +1,7 @@
//! HTTP API 服务器
#![allow(clippy::all)]
pub mod auth;
pub mod chrome_bridge;
pub mod client_detector;
@@ -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,
@@ -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 类型
## 更新提醒
任何文件变更后,请更新此文档和相关的上级文档。
@@ -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;
@@ -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<String>) -> 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<ContentBlock> {
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);
}
}
+8
View File
@@ -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` - 语音输出服务
@@ -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)
}
}
@@ -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<Option<ProviderCredential>, 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,
@@ -419,7 +419,6 @@ impl SessionContextService {
session_id: &str,
messages_to_summarize: &[ChatMessage],
) -> Result<SessionSummary, String> {
// 提取关键信息
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());
}
}
+2
View File
@@ -3,6 +3,8 @@
//! 包含 Skills 系统的 trait 定义和纯逻辑部分。
//! Tauri 相关实现(TauriExecutionCallback)保留在主 crate。
#![allow(clippy::redundant_closure)]
mod execution_callback;
mod llm_provider;
mod proxycast_llm_provider;
@@ -22,6 +22,7 @@ use super::traits::BlockController;
/// 使用 HashMap + RwLock 实现线程安全的控制器管理。
pub struct ControllerRegistry {
/// 控制器映射表: block_id -> BlockController
#[allow(clippy::type_complexity)]
controllers: RwLock<HashMap<String, Arc<RwLock<Box<dyn BlockController>>>>>,
}
@@ -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")]
@@ -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 {
@@ -206,6 +206,7 @@ impl<E: TerminalEventEmitter> ShellProc<E> {
/// - `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,
@@ -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 {
+4
View File
@@ -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;
@@ -111,7 +111,7 @@ impl BlockFile {
/// # 参数
/// - `block_id`: 块 ID
/// - `base_dir`: 基础目录路径
pub fn with_default_size(block_id: &str, base_dir: &PathBuf) -> Result<Self, TerminalError> {
pub fn with_default_size(block_id: &str, base_dir: &Path) -> Result<Self, TerminalError> {
Self::new(block_id, base_dir, DEFAULT_TERM_MAX_FILE_SIZE)
}
+2
View File
@@ -3,6 +3,8 @@
//! 提供 WebSocket API 支持,允许客户端通过持久连接发送请求:
//! - 连接握手和升级
//! - 消息解析和处理
#![allow(clippy::all)]
//! - 流式响应转发
//! - 心跳检测和连接生命周期管理
+13 -20
View File
@@ -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,
@@ -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<String>,
}
/// 获取需要迁移的旧 API Key 凭证列表
#[tauri::command]
pub fn get_legacy_api_key_credentials(
db: State<'_, DbConnection>,
) -> Result<Vec<LegacyApiKeyCredential>, 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<LegacyApiKeyCredential> = 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<String>,
pub api_key_masked: String,
pub base_url: Option<String>,
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<MigrationResult, String> {
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<bool, String> {
pool_service.0.delete_credential(&db, &uuid)
}
// ============================================================================
// 连接测试命令
// ============================================================================
File diff suppressed because it is too large Load Diff
+229 -43
View File
@@ -123,9 +123,18 @@ pub struct ThemeWorkbenchRunState {
pub current_gate_key: String,
pub queue_items: Vec<ThemeWorkbenchRunTodoItem>,
pub latest_terminal: Option<ThemeWorkbenchRunTerminalItem>,
pub recent_terminals: Vec<ThemeWorkbenchRunTerminalItem>,
pub updated_at: String,
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "snake_case")]
pub struct ThemeWorkbenchRunHistoryPage {
pub items: Vec<ThemeWorkbenchRunTerminalItem>,
pub has_more: bool,
pub next_offset: Option<usize>,
}
fn normalize_gate_key(raw: &str) -> Option<String> {
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<String> {
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::<Value>(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<String> {
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<String> {
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::<Vec<_>>()
})
.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<ThemeWorkbenchRunTerminalItem> {
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<usize>,
offset: Option<usize>,
) -> Result<ThemeWorkbenchRunHistoryPage, String> {
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::<Vec<_>>();
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();
File diff suppressed because it is too large Load Diff
+60 -224
View File
@@ -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<MemoryStatsResponse, String> {
async fn memory_runtime_get_stats_impl() -> Result<MemoryStatsResponse, String> {
info!("[记忆管理] 获取记忆统计信息");
let memory_dir = resolve_memory_dir();
@@ -137,9 +142,13 @@ pub async fn get_conversation_memory_stats() -> Result<MemoryStatsResponse, Stri
Ok(overview.stats)
}
/// 获取对话记忆总览(分类 + 条目)
/// 获取 runtime / 上下文记忆统计信息
#[tauri::command]
pub async fn get_conversation_memory_overview(
pub async fn memory_runtime_get_stats() -> Result<MemoryStatsResponse, String> {
memory_runtime_get_stats_impl().await
}
async fn memory_runtime_get_overview_impl(
limit: Option<u32>,
) -> Result<MemoryOverviewResponse, String> {
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<u32>,
) -> Result<MemoryOverviewResponse, String> {
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<i64>,
to_timestamp: Option<i64>,
) -> Result<MemoryAnalysisResult, String> {
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<CleanupMemoryResult, String> {
@@ -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<CleanupMemoryResult, String> {
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<String>) -> Result<PathBuf, String> {
@@ -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<i64>,
to_timestamp: Option<i64>,
) -> Result<Vec<MemorySourceCandidate>, 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::<Vec<_>>()
.join(" ")
}
fn extract_text_from_content_json(content_json: &str) -> String {
if let Ok(text) = serde_json::from_str::<String>(content_json) {
return text;
}
if let Ok(value) = serde_json::from_str::<serde_json::Value>(content_json) {
match value {
serde_json::Value::Array(items) => {
let texts = items
.iter()
.filter_map(extract_text_from_json_item)
.collect::<Vec<_>>();
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<String> {
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<i64> {
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();
-1
View File
@@ -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;
-10
View File
@@ -69,13 +69,3 @@ pub fn get_current_prompt_file_content(app: String) -> Result<Option<String>, St
pub fn auto_import_prompt(db: State<'_, DbConnection>, app: String) -> Result<usize, String> {
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)
}
+15 -13
View File
@@ -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<String> {
}
fn get_skills_dir(app_type: &AppType) -> Result<PathBuf, String> {
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<String
/// 获取已安装的 ProxyCast Skills 目录列表
///
/// 扫描 ~/.proxycast/skills/ 目录,返回包含 SKILL.md 的子目录名列表。
/// 扫描 ProxyCast Skills 目录,返回包含 SKILL.md 的子目录名列表。
/// 这些 Skills 将被传递给 aster 用于 AI Agent 功能。
///
/// # Returns
@@ -121,8 +124,7 @@ fn read_local_skill_content(skills_dir: &Path, directory: &str) -> Result<String
/// - `Err(String)`: 错误信息
#[tauri::command]
pub async fn get_installed_proxycast_skills() -> Result<Vec<String>, 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))
}
+8 -173
View File
@@ -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<i64>,
to_timestamp: Option<i64>,
) -> Result<Vec<MemorySourceCandidate>, 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<String>) -> Vec<String> {
normalized
}
fn normalize_candidate_content(content: &str) -> String {
content
.replace('\n', " ")
.split_whitespace()
.collect::<Vec<_>>()
.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<i64> {
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<i64> {
if let Ok(v) = value.parse::<i64>() {
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::<String>(content_json) {
return text;
}
if let Ok(value) = serde_json::from_str::<serde_json::Value>(content_json) {
match value {
serde_json::Value::Array(items) => {
let texts = items
.iter()
.filter_map(extract_text_from_json_item)
.collect::<Vec<_>>();
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<String> {
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);
+3 -7
View File
@@ -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<PathBuf, String> {
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()
}
/// 规范化项目目录名,避免非法路径字符
+47 -10
View File
@@ -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<T: DeserializeOwned>(
}
fn get_workspace_projects_root_dir() -> Result<PathBuf, String> {
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::<crate::commands::context_memory::ContextMemoryServiceState>();
let db = app_handle.state::<crate::database::DbConnection>();
let global_config = app_handle.state::<crate::config::GlobalConfigManagerState>();
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::<crate::commands::context_memory::ContextMemoryServiceState>();
let global_config = app_handle.state::<crate::config::GlobalConfigManagerState>();
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" => {
// 返回网络信息
+2
View File
@@ -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 警告
@@ -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 = "<proposed_plan>";
const PROPOSED_PLAN_CLOSE: &str = "</proposed_plan>";
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<Value> {
let value = raw?.trim();
if value.is_empty() {
return None;
}
serde_json::from_str::<Value>(value).ok()
}
fn as_object(value: &Value) -> Option<&serde_json::Map<String, Value>> {
value.as_object()
}
fn pick_string_from_object(
object: Option<&serde_json::Map<String, Value>>,
keys: &[&str],
) -> Option<String> {
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<String> {
pick_string_from_object(
arguments.and_then(as_object),
&["q", "query", "question", "search", "search_query", "url"],
)
}
fn extract_command_text(arguments: Option<&Value>) -> Option<String> {
pick_string_from_object(
arguments.and_then(as_object),
&["cmd", "command", "script", "text"],
)
}
fn extract_file_paths(arguments: Option<&Value>, metadata: Option<&Value>) -> Vec<String> {
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<String> {
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<Vec<AgentRequestQuestion>> {
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::<Vec<_>>()
});
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<String, i64>,
item_statuses: HashMap<String, AgentThreadItemStatus>,
assistant_text: String,
reasoning_text: String,
plan_text: Option<String>,
}
impl AgentTimelineRecorder {
pub fn create(
db: DbConnection,
thread_id: impl Into<String>,
prompt_text: impl Into<String>,
) -> Result<Self, String> {
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<String>,
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<Value>,
) -> 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<Value> {
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()))
}
@@ -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")
}
@@ -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<i64>,
to_timestamp: Option<i64>,
limit: usize,
min_message_length: usize,
) -> Result<Vec<MemorySourceCandidate>, 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<i64>,
to_timestamp: Option<i64>,
limit: usize,
min_message_length: usize,
candidates: &mut Vec<MemorySourceCandidate>,
seen: &mut HashSet<String>,
) -> 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<i64>,
to_timestamp: Option<i64>,
limit: usize,
min_message_length: usize,
candidates: &mut Vec<MemorySourceCandidate>,
seen: &mut HashSet<String>,
) -> 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<i64>,
to_timestamp: Option<i64>,
limit: usize,
min_message_length: usize,
candidates: &mut Vec<MemorySourceCandidate>,
seen: &mut HashSet<String>,
) -> 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<MemorySourceCandidate>,
seen: &mut HashSet<String>,
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::<Vec<_>>()
.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<i64> {
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<i64> {
if let Ok(v) = value.parse::<i64>() {
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::<String>(content_json) {
return text;
}
if let Ok(value) = serde_json::from_str::<serde_json::Value>(content_json) {
match value {
serde_json::Value::Array(items) => {
let texts = items
.iter()
.filter_map(extract_text_from_json_item)
.collect::<Vec<_>>();
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<String> {
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::<Vec<_>>();
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::<Vec<_>>();
assert_eq!(candidates.len(), 2);
assert!(session_ids.contains(&"general-migrated"));
assert!(session_ids.contains(&"agent-1"));
assert!(!session_ids.contains(&"legacy-only"));
}
}
@@ -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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, String> {
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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, String> {
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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, String> {
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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, String> {
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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, String> {
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<i64>,
to_timestamp_ms: Option<i64>,
) -> Result<i64, String> {
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<Local>,
month_start: &DateTime<Local>,
) -> Result<ConversationStats, String> {
// 转换为 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<Local>,
month_start: &DateTime<Local>,
) -> Result<ConversationStats, String> {
// 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<TokenStats, String> {
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);
}
}
@@ -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<Vec<AgentRun>, 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<T, Fut>(
&self,
source: RunSource,
@@ -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 {
+2
View File
@@ -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;
+626 -23
View File
@@ -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<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
pub struct EnvironmentDiagnostics {
pub npm_path: Option<String>,
pub npm_global_prefix: Option<String>,
pub openclaw_package_path: Option<String>,
#[serde(default)]
pub where_candidates: Vec<String>,
#[serde(default)]
pub supplemental_search_dirs: Vec<String>,
#[serde(default)]
pub supplemental_command_candidates: Vec<String>,
}
#[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<BinaryInstallStatus, String> {
@@ -312,6 +329,16 @@ impl OpenClawService {
pub async fn install(&mut self, app: &AppHandle) -> Result<ActionResult, String> {
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<ActionResult> {
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<PathBuf, String> {
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<DependencyStatus, String> {
"未检测到 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<DependencyStatus, String> {
"检测到 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<DependencyStatus, String> {
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<DependencyStatus, String> {
"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<DependencyStatus, String> {
async fn inspect_openclaw_dependency_status() -> Result<DependencyStatus, String> {
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<DependencyStatus, String
})
}
async fn git_auto_install_supported() -> Result<bool, String> {
#[cfg(target_os = "windows")]
{
Ok(find_command_in_shell("winget").await?.is_some())
}
async fn inspect_openclaw_package_reload_status() -> Result<Option<DependencyStatus>, 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<bool, String> {
#[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<Option<String>, 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<Option<String>, 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<Vec<PathBuf>, 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<Option<String>, Str
candidates.extend(find_all_commands_in_known_locations(command_name));
Ok(candidates)
}
async fn select_command_path(
command_name: &str,
candidates: Vec<PathBuf>,
) -> Result<Option<PathBuf>, 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<Option<String>, 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<PathBuf> {
let search_dirs = collect_known_command_search_dirs();
find_all_commands_in_paths(command_name, &search_dirs)
}
fn collect_known_command_search_dirs() -> Vec<PathBuf> {
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<PathBuf> {
}
}
#[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<PathBuf> {
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<PathBuf> {
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<PathBuf> {
@@ -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<Vec<PathBuf>, 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<PathBuf> {
npm_global_command_dirs_for(current_shell_platform(), prefix)
}
fn npm_global_command_dirs_for(platform: ShellPlatform, prefix: &str) -> Vec<PathBuf> {
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<PathBuf> {
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<String>,
path: PathBuf,
}
#[cfg(test)]
fn find_installed_openclaw_package(prefix: &str) -> Option<(&'static str, Option<String>)> {
find_installed_openclaw_package_details(prefix).map(|package| (package.name, package.version))
}
fn find_installed_openclaw_package_details(prefix: &str) -> Option<InstalledOpenClawPackage> {
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<String> {
#[derive(Deserialize)]
struct PackageManifest {
version: Option<String>,
}
let content = std::fs::read_to_string(manifest_path).ok()?;
let manifest = serde_json::from_str::<PackageManifest>(&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<PathBuf> {
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<PathBuf>) -> Result<Option<PathBuf>, 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!(
@@ -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<PathBuf, String> {
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() {
+20 -24
View File
@@ -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<Vec<String>, String> {
let skills_root = skills_root_from_home(home_dir);
fn ensure_default_local_skills_in_dir(skills_root: &Path) -> Result<Vec<String>, 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<Vec<String>, S
}
pub fn ensure_default_local_skills() -> Result<Vec<String>, 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()),
"内置版本更新时应自动升级"
+1 -1
View File
@@ -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",
@@ -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<AgentPlanBlockProps> = ({
content,
isComplete = true,
}) => {
if (!content.trim()) {
return null;
}
return (
<div className="rounded-2xl border border-border/70 bg-muted/35 px-4 py-3">
<div className="mb-2 flex items-center gap-2">
<div className="flex h-7 w-7 items-center justify-center rounded-full bg-primary/10 text-primary">
<ListChecks className="h-4 w-4" />
</div>
<div className="text-sm font-medium text-foreground">执行计划</div>
<Badge variant={isComplete ? "outline" : "secondary"} className="ml-auto">
{isComplete ? (
"已生成"
) : (
<span className="inline-flex items-center gap-1">
<Loader2 className="h-3 w-3 animate-spin" />
规划中
</span>
)}
</Badge>
</div>
<div className="text-sm">
<MarkdownRenderer content={content} />
</div>
</div>
);
};
@@ -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<string, string> = {
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<AgentRuntimeStripProps> = ({
activeTheme,
toolPreferences,
harnessState,
subAgentRuntime,
variant = "standalone",
isSending = false,
runtimeStatusTitle = null,
}) => {
const themeLabel =
THEME_LABELS[activeTheme?.trim().toLowerCase() || ""] || "通用对话";
const capabilities = useMemo<CapabilityItem[]>(
() => [
{ 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<StatusItem[]>(() => {
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 (
<div
className={
variant === "embedded"
? "rounded-xl border border-border/70 bg-[linear-gradient(135deg,hsl(var(--background)),hsl(var(--muted)/0.35))] px-4 py-3"
: "mx-3 mb-2 mt-3 rounded-2xl border border-border/70 bg-[linear-gradient(135deg,hsl(var(--background)),hsl(var(--muted)/0.35))] px-4 py-3"
}
>
<div className="mb-2 flex flex-wrap items-center gap-2">
<div className="text-sm font-medium text-foreground">通用 Agent</div>
<Badge variant="outline">{themeLabel}</Badge>
</div>
<div className="mb-3 flex flex-wrap gap-2">
{capabilities.map((item) => (
<span
key={item.key}
className={[
"rounded-full border px-2.5 py-1 text-xs transition-colors",
item.enabled
? "border-primary/30 bg-primary/10 text-primary"
: "border-border/70 bg-background/80 text-muted-foreground",
].join(" ")}
>
{item.label}
</span>
))}
</div>
<div className="flex flex-wrap gap-2">
{statusItems.map((item) => (
<Badge key={item.key} variant={item.tone || "outline"}>
{item.label}
</Badge>
))}
</div>
</div>
);
};
export default AgentRuntimeStrip;
@@ -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<string, unknown>)
: 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<string, unknown>)
: 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 (
<div className="rounded-2xl border border-border/70 bg-background/80 px-4 py-3">
<div className="mb-2 flex items-center gap-2">
<div className="flex h-7 w-7 items-center justify-center rounded-full bg-primary/10 text-primary">
<Icon className="h-4 w-4" />
</div>
<div className="text-sm font-medium text-foreground">{title}</div>
{badge ? <div className="ml-auto">{badge}</div> : null}
{timestamp ? (
<div className="text-xs text-muted-foreground">{timestamp}</div>
) : null}
</div>
{children}
</div>
);
}
export const AgentThreadTimeline: React.FC<AgentThreadTimelineProps> = ({
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 (
<div className="mb-3 space-y-2 rounded-2xl border border-border/60 bg-muted/20 p-3">
<div className="flex flex-wrap items-center gap-2">
<div className="text-sm font-medium text-foreground">执行轨迹</div>
{isCurrentTurn ? <Badge variant="secondary">当前回合</Badge> : null}
<Badge variant="outline">
{turn.status === "running"
? "执行中"
: turn.status === "failed"
? "失败"
: turn.status === "aborted"
? "已中断"
: "已完成"}
</Badge>
<div className="ml-auto flex items-center gap-1 text-xs text-muted-foreground">
<Clock3 className="h-3.5 w-3.5" />
<span>{formatTimestamp(turn.started_at) || "刚刚"}</span>
</div>
</div>
<div className="space-y-2">
{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 (
<AgentPlanBlock
key={item.id}
content={item.text}
isComplete={item.status !== "in_progress"}
/>
);
}
if (item.type === "reasoning") {
return (
<details
key={item.id}
className="overflow-hidden rounded-2xl border border-border/70 bg-background/80"
open={item.status === "in_progress"}
>
<summary className="flex cursor-pointer items-center gap-2 px-4 py-3 text-sm font-medium text-foreground">
<Sparkles className="h-4 w-4 text-primary" />
思考摘要
<Badge
variant={resolveStatusBadgeVariant(item.status)}
className="ml-auto"
>
{item.status === "in_progress" ? (
<span className="inline-flex items-center gap-1">
<Loader2 className="h-3 w-3 animate-spin" />
推理中
</span>
) : item.status === "failed" ? (
"推理失败"
) : (
"已整理"
)}
</Badge>
<ChevronDown className="h-4 w-4 text-muted-foreground" />
</summary>
<div className="border-t border-border/70 px-4 py-3">
<MarkdownRenderer content={item.text} />
</div>
</details>
);
}
if (toolCall) {
return (
<div key={item.id} className="rounded-2xl border border-border/70 bg-background/80">
<ToolCallItem
toolCall={toolCall}
defaultExpanded={item.status === "in_progress"}
onFileClick={onFileClick}
/>
</div>
);
}
if (actionRequest) {
return (
<div key={item.id}>
<DecisionPanel
request={actionRequest}
onSubmit={(response) => onPermissionResponse?.(response)}
/>
</div>
);
}
if (item.type === "file_artifact") {
return (
<TimelineCard
key={item.id}
icon={FileText}
title="文件产物"
badge={
<Badge variant={resolveStatusBadgeVariant(item.status)}>
{item.source}
</Badge>
}
timestamp={timestamp}
>
<button
type="button"
className="w-full rounded-xl border border-border/70 bg-muted/30 px-3 py-2 text-left transition-colors hover:bg-muted/50"
onClick={() => onFileClick?.(item.path, item.content || "")}
>
<div className="text-sm font-medium text-foreground">
{item.path}
</div>
{item.content?.trim() ? (
<div className="mt-1 line-clamp-3 whitespace-pre-wrap text-xs text-muted-foreground">
{item.content}
</div>
) : (
<div className="mt-1 text-xs text-muted-foreground">
点击在画布中打开文件
</div>
)}
</button>
</TimelineCard>
);
}
if (item.type === "subagent_activity") {
return (
<TimelineCard
key={item.id}
icon={Bot}
title={item.title || "子代理协作"}
badge={
<Badge variant={resolveStatusBadgeVariant(item.status)}>
{item.status_label}
</Badge>
}
timestamp={timestamp}
>
{item.summary ? (
<div className="text-sm text-muted-foreground">{item.summary}</div>
) : null}
{item.role || item.model ? (
<div className="mt-2 flex flex-wrap gap-2">
{item.role ? <Badge variant="outline">{item.role}</Badge> : null}
{item.model ? <Badge variant="outline">{item.model}</Badge> : null}
</div>
) : null}
</TimelineCard>
);
}
if (item.type === "turn_summary") {
return (
<TimelineCard
key={item.id}
icon={Sparkles}
title={item.status === "in_progress" ? "执行准备" : "回合总结"}
badge={
item.status === "in_progress" ? (
<Badge variant="secondary" className="inline-flex items-center gap-1">
<Loader2 className="h-3 w-3 animate-spin" />
进行中
</Badge>
) : (
<Badge variant="outline">摘要</Badge>
)
}
timestamp={timestamp}
>
<MarkdownRenderer content={item.text} />
</TimelineCard>
);
}
if (item.type === "warning") {
return (
<TimelineCard
key={item.id}
icon={AlertTriangle}
title="运行提醒"
badge={<Badge variant="secondary">{item.code || "warning"}</Badge>}
timestamp={timestamp}
>
<div className="text-sm text-muted-foreground">{item.message}</div>
</TimelineCard>
);
}
if (item.type === "error") {
return (
<TimelineCard
key={item.id}
icon={ShieldAlert}
title="执行错误"
badge={<Badge variant="destructive">失败</Badge>}
timestamp={timestamp}
>
<div className="text-sm text-destructive">{item.message}</div>
</TimelineCard>
);
}
return (
<TimelineCard
key={item.id}
icon={
item.type === "web_search"
? Search
: item.type === "command_execution"
? TerminalSquare
: item.type === "approval_request"
? ShieldAlert
: item.type === "request_user_input"
? Globe
: Wrench
}
title={item.type}
badge={
<Badge variant={resolveStatusBadgeVariant(item.status)}>
{item.status}
</Badge>
}
timestamp={timestamp}
>
<div className="text-sm text-muted-foreground">
该事件类型已记录到 timeline 中。
</div>
</TimelineCard>
);
})}
</div>
</div>
);
};
export default AgentThreadTimeline;
@@ -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<ChatNavbarProps> = ({
onToggleHarnessPanel,
harnessPendingCount = 0,
harnessAttentionLevel = "idle",
harnessToggleLabel = "Harness",
novelCanvasControls = null,
}) => {
return (
@@ -173,15 +175,19 @@ export const ChatNavbar: React.FC<ChatNavbarProps> = ({
)}
onClick={onToggleHarnessPanel}
aria-label={
harnessPanelVisible ? "收起 Harness 面板" : "展开 Harness 面板"
harnessPanelVisible
? `收起${harnessToggleLabel}`
: `展开${harnessToggleLabel}`
}
aria-expanded={harnessPanelVisible}
title={
harnessPanelVisible ? "收起 Harness 面板" : "展开 Harness 面板"
harnessPanelVisible
? `收起${harnessToggleLabel}`
: `展开${harnessToggleLabel}`
}
>
<Sparkles size={14} />
<span>Harness</span>
<span>{harnessToggleLabel}</span>
{harnessPendingCount > 0 ? (
<span className="rounded-full bg-destructive px-1.5 py-0.5 text-[10px] font-medium leading-none text-destructive-foreground">
{harnessPendingCount > 99 ? "99+" : harnessPendingCount}
@@ -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 }) => <div>{children}</div>,
Select: ({ children }: { children: React.ReactNode }) => (
<div>{children}</div>
),
SelectContent: ({ children }: { children: React.ReactNode }) => (
<div>{children}</div>
),
SelectItem: ({ children }: { children: React.ReactNode }) => <div>{children}</div>,
SelectItem: ({ children }: { children: React.ReactNode }) => (
<div>{children}</div>
),
SelectTrigger: ({ children }: { children: React.ReactNode }) => (
<button type="button">{children}</button>
),
@@ -107,7 +116,9 @@ vi.mock("@/components/ui/select", () => ({
}));
vi.mock("@/components/ui/popover", () => ({
Popover: ({ children }: { children: React.ReactNode }) => <div>{children}</div>,
Popover: ({ children }: { children: React.ReactNode }) => (
<div>{children}</div>
),
PopoverContent: ({ children }: { children: React.ReactNode }) => (
<div>{children}</div>
),
@@ -159,7 +170,9 @@ afterEach(() => {
vi.clearAllMocks();
});
function renderEmptyState(props?: Partial<React.ComponentProps<typeof EmptyState>>) {
function renderEmptyState(
props?: Partial<React.ComponentProps<typeof EmptyState>>,
) {
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);
});
});
@@ -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<EmptyStateProps> = ({
onWebSearchEnabledChange,
thinkingEnabled = false,
onThinkingEnabledChange,
taskEnabled = false,
onTaskEnabledChange,
subagentEnabled = false,
onSubagentEnabledChange,
hasCanvasContent = false,
hasContentId = false,
selectedText = "",
@@ -614,6 +625,7 @@ export const EmptyState: React.FC<EmptyStateProps> = ({
// 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<EmptyStateProps> = ({
</>
)}
{activeTheme === "general" && (
{isGeneralTheme && (
<>
<Button
variant="outline"
@@ -1266,6 +1278,34 @@ export const EmptyState: React.FC<EmptyStateProps> = ({
>
<Lightbulb className="w-4 h-4 opacity-70" />
</Button>
<Button
variant="outline"
size="icon"
className={`h-8 w-8 rounded-full ml-1 bg-background shadow-sm hover:bg-muted ${
taskEnabled
? "border-emerald-500 text-emerald-600 bg-emerald-50 dark:bg-emerald-950/30"
: ""
}`}
onClick={() => onTaskEnabledChange?.(!taskEnabled)}
aria-pressed={taskEnabled}
title={taskEnabled ? "关闭后台任务偏好" : "开启后台任务偏好"}
>
<ListChecks className="w-4 h-4 opacity-70" />
</Button>
<Button
variant="outline"
size="icon"
className={`h-8 w-8 rounded-full ml-1 bg-background shadow-sm hover:bg-muted ${
subagentEnabled
? "border-fuchsia-500 text-fuchsia-600 bg-fuchsia-50 dark:bg-fuchsia-950/30"
: ""
}`}
onClick={() => onSubagentEnabledChange?.(!subagentEnabled)}
aria-pressed={subagentEnabled}
title={subagentEnabled ? "关闭多代理偏好" : "开启多代理偏好"}
>
<Workflow className="w-4 h-4 opacity-70" />
</Button>
</>
)}
@@ -29,6 +29,7 @@ function createHarnessState(
overrides: Partial<HarnessSessionState> = {},
): 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: <div>通用 Agent 运行概览</div>,
});
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;
@@ -78,10 +78,15 @@ interface HarnessStatusPanelProps {
error: string | null;
};
environment: HarnessEnvironmentSummary;
layout?: "default" | "sidebar" | "dialog";
onLoadFilePreview?: (path: string) => Promise<HarnessFilePreviewResult>;
onOpenFile?: (fileName: string, content: string) => void;
onRevealPath?: (path: string) => Promise<void>;
onOpenPath?: (path: string) => Promise<void>;
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<FileFilterValue>("all");
const [outputFilter, setOutputFilter] = useState<OutputFilterValue>("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 (
<>
<div className="mx-3 mt-2 rounded-2xl border border-border bg-muted/30">
<div
data-testid="harness-status-panel"
data-layout={layout}
className={cn(
"bg-muted/30",
layout === "sidebar"
? "rounded-xl border border-border"
: layout === "dialog"
? "overflow-hidden rounded-xl border border-border/70 bg-background"
: "mx-3 mt-2 rounded-2xl border border-border",
)}
>
<div className="flex items-center justify-between gap-3 border-b border-border px-4 py-3">
<div className="min-w-0">
<div className="flex items-center gap-2">
<Wrench className="h-4 w-4 text-muted-foreground" />
<h2 className="text-sm font-semibold text-foreground">
Harness 运行面板
{title}
</h2>
{subAgentRuntime.isRunning ? (
<Badge variant="secondary" className="gap-1">
@@ -797,28 +888,43 @@ export function HarnessStatusPanel({
) : null}
</div>
<p className="mt-1 text-xs text-muted-foreground">
展示最近文件活动、工具输出、审批与上下文装载情况。
{description}
</p>
</div>
<Button
type="button"
size="sm"
variant="ghost"
className="shrink-0"
onClick={() => setExpanded((value) => !value)}
aria-expanded={expanded}
aria-label={expanded ? "折叠 Harness 详情" : "展开 Harness 详情"}
>
{expanded ? (
<ChevronDown className="mr-1 h-4 w-4" />
) : (
<ChevronRight className="mr-1 h-4 w-4" />
)}
{expanded ? "收起详情" : "展开详情"}
</Button>
{!isDialogLayout ? (
<Button
type="button"
size="sm"
variant="ghost"
className="shrink-0"
onClick={() => setExpanded((value) => !value)}
aria-expanded={isDetailsExpanded}
aria-label={
isDetailsExpanded
? `折叠${toggleLabel}`
: `展开${toggleLabel}`
}
>
{isDetailsExpanded ? (
<ChevronDown className="mr-1 h-4 w-4" />
) : (
<ChevronRight className="mr-1 h-4 w-4" />
)}
{isDetailsExpanded ? `收起${toggleLabel}` : `展开${toggleLabel}`}
</Button>
) : null}
</div>
<div className="grid gap-3 px-4 py-4 md:grid-cols-2 xl:grid-cols-4">
{leadContent ? (
<div className="border-b border-border px-4 py-4">{leadContent}</div>
) : null}
<div
className={cn(
"grid gap-3 px-4 py-4",
layout === "sidebar" ? "grid-cols-1" : "md:grid-cols-2 xl:grid-cols-4",
)}
>
{summaryCards.map((card) => (
<SummaryCard
key={card.title}
@@ -831,8 +937,17 @@ export function HarnessStatusPanel({
))}
</div>
{expanded ? (
<ScrollArea className="max-h-[28rem] border-t border-border px-4 py-4">
{isDetailsExpanded ? (
<ScrollArea
className={cn(
"border-t border-border px-4 py-4",
layout === "sidebar"
? "max-h-[24rem]"
: layout === "dialog"
? "max-h-[58vh]"
: "max-h-[28rem]",
)}
>
<div className="space-y-4 pb-1">
{availableSections.length > 0 ? (
<div className="flex flex-wrap gap-2">
@@ -1139,6 +1254,44 @@ export function HarnessStatusPanel({
</Section>
) : null}
{harnessState.runtimeStatus ? (
<Section
sectionKey="runtime"
title="当前执行阶段"
badge={formatRuntimePhaseLabel(harnessState.runtimeStatus)}
registerRef={registerSectionRef}
>
<div className="space-y-3">
<div className="rounded-xl border border-primary/20 bg-primary/5 p-3">
<div className="flex items-center gap-2 text-sm font-medium text-foreground">
<Loader2 className="h-4 w-4 animate-spin text-primary" />
<span>{harnessState.runtimeStatus.title}</span>
</div>
<div className="mt-2 text-sm text-muted-foreground">
{harnessState.runtimeStatus.detail}
</div>
</div>
{harnessState.runtimeStatus.checkpoints &&
harnessState.runtimeStatus.checkpoints.length > 0 ? (
<div className="flex flex-wrap gap-2">
{harnessState.runtimeStatus.checkpoints.map(
(checkpoint, index) => (
<Badge
key={`${checkpoint}-${index}`}
variant="outline"
className="max-w-full whitespace-normal text-left"
>
{checkpoint}
</Badge>
),
)}
</div>
) : null}
</div>
</Section>
) : null}
{harnessState.pendingApprovals.length > 0 ? (
<Section
sectionKey="approvals"
@@ -13,7 +13,8 @@ interface A2UISubmissionNoticeProps {
}
const Container = styled.div<{ $visible: boolean }>`
margin-bottom: 8px;
margin: 0 8px 8px;
box-sizing: border-box;
display: flex;
align-items: flex-start;
gap: 8px;
@@ -140,6 +140,7 @@ export const InputbarComposerSection: React.FC<
showDragHandle={!isThemeWorkbenchVariant}
visualVariant={isThemeWorkbenchVariant ? "floating" : "default"}
topExtra={topExtra}
activeTheme={activeTheme}
leftExtra={
<InputbarModelExtra
isFullscreen={isFullscreen}
@@ -71,6 +71,7 @@ interface InputbarCoreProps {
showDragHandle?: boolean;
/** 视觉风格 */
visualVariant?: "default" | "floating";
activeTheme?: string;
}
export const InputbarCore: React.FC<InputbarCoreProps> = ({
@@ -98,6 +99,7 @@ export const InputbarCore: React.FC<InputbarCoreProps> = ({
showTranslate = true,
showDragHandle = true,
visualVariant = "default",
activeTheme,
}) => {
const [isComposerExpanded, setIsComposerExpanded] = useState(false);
const inputBarContainerRef = useRef<HTMLDivElement | null>(null);
@@ -259,6 +261,7 @@ export const InputbarCore: React.FC<InputbarCoreProps> = ({
showExecutionStrategy={showExecutionStrategy}
toolMode={toolMode}
isCanvasOpen={isCanvasOpen}
activeTheme={activeTheme}
/>
) : null}
</LeftSection>
@@ -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<InputbarToolsProps> = ({
@@ -29,6 +33,7 @@ export const InputbarTools: React.FC<InputbarToolsProps> = ({
executionStrategy = "react",
showExecutionStrategy = false,
toolMode = "default",
activeTheme,
}) => {
const modeLabel =
executionStrategy === "auto"
@@ -38,6 +43,7 @@ export const InputbarTools: React.FC<InputbarToolsProps> = ({
: "ReAct";
const strategyEnabled =
executionStrategy !== "react" || activeTools["execution_strategy"];
const isGeneralTheme = isGeneralResearchTheme(activeTheme);
return (
<TooltipProvider>
@@ -85,6 +91,46 @@ export const InputbarTools: React.FC<InputbarToolsProps> = ({
</TooltipContent>
</Tooltip>
{isGeneralTheme ? (
<>
<Tooltip>
<TooltipTrigger asChild>
<ToolButton
onClick={() => onToolClick?.("task_mode")}
className={activeTools["task_mode"] ? "active" : ""}
>
<ListChecks
className={
activeTools["task_mode"] ? "text-emerald-500" : ""
}
/>
</ToolButton>
</TooltipTrigger>
<TooltipContent side="top">
后台任务 {activeTools["task_mode"] ? "(偏好已开启)" : ""}
</TooltipContent>
</Tooltip>
<Tooltip>
<TooltipTrigger asChild>
<ToolButton
onClick={() => onToolClick?.("subagent_mode")}
className={activeTools["subagent_mode"] ? "active" : ""}
>
<Workflow
className={
activeTools["subagent_mode"] ? "text-fuchsia-500" : ""
}
/>
</ToolButton>
</TooltipTrigger>
<TooltipContent side="top">
多代理 {activeTools["subagent_mode"] ? "(偏好已开启)" : ""}
</TooltipContent>
</Tooltip>
</>
) : null}
{showExecutionStrategy && (
<Tooltip>
<TooltipTrigger asChild>
@@ -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,

Some files were not shown because too many files have changed in this diff Show More