mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
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:
+18
-33
@@ -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
|
||||
|
||||
@@ -16,6 +16,7 @@
|
||||
- `develop/`:开发流程与协作规范
|
||||
- `plugins/`:插件与扩展相关文档
|
||||
- `tests/`:测试策略与用例文档
|
||||
- `iteration-notes/`:迭代备忘与下版本建议(暂不进入当前发布范围的问题)
|
||||
- `images/`:文档图片资源
|
||||
- `TECH_SPEC.md`:技术规格文档
|
||||
- `develop/execution-tracker-technical-plan.md`:统一执行追踪(Execution Tracker)专项技术规划
|
||||
|
||||
+41
-37
@@ -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 刷新
|
||||
- [ ] 请求日志/统计
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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();
|
||||
```
|
||||
|
||||
## 错误处理
|
||||
|
||||
@@ -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>;
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
@@ -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,因此采用**消息注入**方案:
|
||||
|
||||
@@ -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
@@ -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,
|
||||
};
|
||||
```
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
Generated
+16
-16
@@ -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",
|
||||
|
||||
@@ -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 },
|
||||
|
||||
@@ -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};
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
×tamp_str,
|
||||
)? {
|
||||
tracing::debug!(
|
||||
"[迁移] general_chat 消息已存在于 unified 表,跳过: session_id={}, timestamp={}",
|
||||
session_id,
|
||||
timestamp_str
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
conn.execute(
|
||||
"INSERT INTO agent_messages (session_id, role, content_json, timestamp, tool_calls_json, tool_call_id)
|
||||
VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
|
||||
params![
|
||||
session_id,
|
||||
role,
|
||||
content_json,
|
||||
timestamp_str,
|
||||
Option::<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(())
|
||||
}
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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,6 +2,8 @@
|
||||
//!
|
||||
//! 承载渠道运行时(channel runtime)、路由与策略实现。
|
||||
|
||||
#![allow(clippy::all)]
|
||||
|
||||
pub mod discord;
|
||||
pub mod feishu;
|
||||
pub mod telegram;
|
||||
|
||||
@@ -3,6 +3,10 @@
|
||||
//! 提供统一的请求处理管道,集成路由、容错、监控、插件等功能模块。
|
||||
//!
|
||||
//! ## 模块结构
|
||||
|
||||
#![allow(clippy::derivable_impls)]
|
||||
#![allow(clippy::unnecessary_map_or)]
|
||||
#![allow(clippy::too_many_arguments)]
|
||||
//!
|
||||
//! - `steps` - 管道步骤(认证、注入、路由、插件、Provider、遥测)
|
||||
|
||||
|
||||
@@ -3,6 +3,11 @@
|
||||
//! 提供 Agent 任务调度功能,支持定时任务、重试机制等。
|
||||
//!
|
||||
//! ## 功能
|
||||
|
||||
#![allow(clippy::redundant_closure)]
|
||||
#![allow(clippy::format_in_format_args)]
|
||||
#![allow(clippy::useless_format)]
|
||||
#![allow(clippy::derivable_impls)]
|
||||
//! - 任务创建和管理
|
||||
//! - 任务持久化到 SQLite
|
||||
//! - 定时任务调度
|
||||
|
||||
@@ -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,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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
//! 提供 WebSocket API 支持,允许客户端通过持久连接发送请求:
|
||||
//! - 连接握手和升级
|
||||
//! - 消息解析和处理
|
||||
|
||||
#![allow(clippy::all)]
|
||||
//! - 流式响应转发
|
||||
//! - 心跳检测和连接生命周期管理
|
||||
|
||||
|
||||
+13
-20
@@ -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
@@ -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
@@ -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();
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
/// 规范化项目目录名,避免非法路径字符
|
||||
|
||||
@@ -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" => {
|
||||
// 返回网络信息
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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,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
Reference in New Issue
Block a user