From 91fe7ff481cff505a6fb63944ece7ca977bbc8a5 Mon Sep 17 00:00:00 2001 From: coso Date: Tue, 3 Feb 2026 00:01:34 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20v0.53.0=20-=20=E7=BB=9F=E4=B8=80?= =?UTF-8?q?=E5=AF=B9=E8=AF=9D=E7=B3=BB=E7=BB=9F=E4=B8=8E=E5=86=85=E5=AE=B9?= =?UTF-8?q?=E5=88=9B=E4=BD=9C=E5=A2=9E=E5=BC=BA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Features 统一对话系统 - 新增 useUnifiedChat Hook 支持多种对话模式 (Agent/General/Creator) - 新增统一对话 API 封装 (unified-chat.ts) - 新增 Chat 类型定义系统 (chat.ts) - 支持流式响应和工具调用 内容创作增强 - 新增 Content Creator 文档和系统提示词 - 新增 writeFile 解析器支持画布输出 - 优化 General Chat 组件支持内容创作模式 - 支持多种创作主题 (社交媒体/海报/音乐/知识/规划/文档/视频/小说) 会话管理优化 - 增强 Aster Agent 状态管理 - 新增统一会话命令 (unified_chat_cmd) - 新增 Chat DAO 数据访问层 - 优化数据库迁移脚本 Bug Fixes - 修复 ESLint 警告和未使用变量问题 - 修复 React Hooks 依赖警告 Co-Authored-By: Claude Opus 4.5 --- docs/aiprompts/README.md | 6 + docs/aiprompts/content-creator.md | 233 +++++++ docs/aiprompts/hooks.md | 40 ++ package.json | 2 +- src-tauri/Cargo.lock | 32 +- src-tauri/Cargo.toml | 8 +- src-tauri/src/agent/aster_agent.rs | 7 +- src-tauri/src/agent/aster_state.rs | 124 +++- src-tauri/src/app/runner.rs | 10 + src-tauri/src/app/setup.rs | 13 + src-tauri/src/commands/agent_cmd.rs | 7 +- src-tauri/src/commands/aster_agent_cmd.rs | 35 +- src-tauri/src/commands/mod.rs | 1 + src-tauri/src/commands/unified_chat_cmd.rs | 424 +++++++++++ src-tauri/src/database/dao/agent.rs | 60 +- src-tauri/src/database/dao/chat.rs | 426 ++++++++++++ src-tauri/src/database/dao/mod.rs | 1 + src-tauri/src/database/migration.rs | 300 ++++++++ src-tauri/src/services/aster_session_store.rs | 27 + src-tauri/tauri.conf.json | 2 +- .../agent/chat/hooks/useAgentChat.ts | 20 +- src/components/agent/chat/index.tsx | 4 + .../content-creator/canvas/canvasUtils.ts | 13 +- .../content-creator/utils/systemPrompt.ts | 6 +- .../general-chat/chat/AssistantMessage.tsx | 50 +- .../general-chat/chat/ChatPanel.tsx | 12 + .../general-chat/chat/MessageItem.tsx | 4 + .../general-chat/chat/MessageList.tsx | 5 + .../general-chat/store/useGeneralChatStore.ts | 148 +++- src/hooks/README.md | 83 ++- src/hooks/useUnifiedChat.ts | 657 ++++++++++++++++++ src/lib/README.md | 3 + src/lib/api/unified-chat.ts | 315 +++++++++ src/lib/writeFile/README.md | 67 ++ src/lib/writeFile/index.ts | 8 + src/lib/writeFile/parser.ts | 205 ++++++ src/types/chat.ts | 384 ++++++++++ 37 files changed, 3645 insertions(+), 97 deletions(-) create mode 100644 docs/aiprompts/content-creator.md create mode 100644 src-tauri/src/commands/unified_chat_cmd.rs create mode 100644 src-tauri/src/database/dao/chat.rs create mode 100644 src/hooks/useUnifiedChat.ts create mode 100644 src/lib/api/unified-chat.ts create mode 100644 src/lib/writeFile/README.md create mode 100644 src/lib/writeFile/index.ts create mode 100644 src/lib/writeFile/parser.ts create mode 100644 src/types/chat.ts diff --git a/docs/aiprompts/README.md b/docs/aiprompts/README.md index 574d9e948..c82d2f638 100644 --- a/docs/aiprompts/README.md +++ b/docs/aiprompts/README.md @@ -36,6 +36,9 @@ AI Agent 专用文档目录,提供模块级别的详细说明。 - `aster-integration.md` - **Aster 框架集成方案** - `workspace.md` - **Workspace 设计文档**(工作目录管理) +### 内容创作 +- `content-creator.md` - **内容创作系统**(write_file 标签、画布联动) + ## 使用方式 AI Agent 在处理特定模块时,应先阅读对应的 aiprompts 文档: @@ -52,6 +55,9 @@ AI Agent 在处理特定模块时,应先阅读对应的 aiprompts 文档: # 处理 Workspace 相关任务 → 先读 docs/aiprompts/workspace.md + +# 处理内容创作、画布联动 +→ 先读 docs/aiprompts/content-creator.md ``` ## 更新提醒 diff --git a/docs/aiprompts/content-creator.md b/docs/aiprompts/content-creator.md new file mode 100644 index 000000000..b958f2407 --- /dev/null +++ b/docs/aiprompts/content-creator.md @@ -0,0 +1,233 @@ +# 内容创作系统 + +## 概述 + +内容创作系统支持多种主题(社媒内容、图文海报、歌词曲谱等),通过 `` 标签实现 AI 响应与右侧画布的联动。 + +## 核心架构 + +``` +用户选择主题 → AgentChatPage 生成 systemPrompt + ↓ +用户发送消息 → useAgentChat.sendMessage() + ↓ +第一条消息时注入 systemPrompt → 发送到 Aster Agent + ↓ +AI 返回带 标签的响应 + ↓ +StreamingRenderer 解析标签 → 调用 onWriteFile + ↓ +AgentChatPage.handleWriteFile → 更新画布状态 + ↓ +右侧画布自动打开,显示文档内容 +``` + +## 目录结构 + +``` +src/components/ +├── content-creator/ +│ ├── canvas/ # 画布组件 +│ │ ├── CanvasFactory.tsx # 画布工厂 +│ │ ├── canvasUtils.ts # 画布工具 +│ │ ├── document/ # 文档画布 +│ │ └── music/ # 音乐画布 +│ ├── utils/ +│ │ └── systemPrompt.ts # 系统提示词生成 +│ └── a2ui/ +│ └── parser.ts # A2UI 和 write_file 解析器 +├── agent/chat/ +│ ├── hooks/ +│ │ └── useAgentChat.ts # Agent 聊天 Hook +│ ├── components/ +│ │ ├── StreamingRenderer.tsx # 流式渲染(解析 write_file) +│ │ └── MessageList.tsx # 消息列表 +│ └── index.tsx # AgentChatPage +└── general-chat/ + └── store/ + └── useGeneralChatStore.ts # 通用对话 Store +``` + +## 核心组件 + +### 1. systemPrompt.ts - 系统提示词生成 + +根据主题和创作模式生成 AI 系统提示词。 + +```typescript +// src/components/content-creator/utils/systemPrompt.ts + +// 主题类型 +type ThemeType = "general" | "social-media" | "poster" | "music" | ...; + +// 创作模式 +type CreationMode = "guided" | "fast" | "hybrid" | "framework"; + +// 生成系统提示词 +export function generateContentCreationPrompt( + theme: ThemeType, + mode: CreationMode +): string; +``` + +**关键指令**:系统提示词要求 AI 使用 `` 标签输出内容: + +```markdown +## 文件写入格式 + +当需要输出文档内容时,使用以下标签格式: + + +内容... + + +**重要规则**: +- 标签前:先写一句引导语 +- 标签后:写完成总结 +- 标签内的内容会实时流式显示在右侧画布 +``` + +### 2. parser.ts - write_file 标签解析 + +解析 AI 响应中的 `` 标签。 + +```typescript +// src/components/content-creator/a2ui/parser.ts + +interface ParseResult { + parts: ParsedMessageContent[]; + hasA2UI: boolean; + hasWriteFile: boolean; + hasPending: boolean; +} + +// 解析 AI 响应 +export function parseAIResponse( + content: string, + isStreaming: boolean +): ParseResult; +``` + +**支持的标签类型**: +- `write_file` - 完整的文件写入 +- `pending_write_file` - 流式传输中的文件写入 + +### 3. useAgentChat.ts - systemPrompt 注入 + +在发送第一条消息时注入 systemPrompt。 + +```typescript +// src/components/agent/chat/hooks/useAgentChat.ts + +interface UseAgentChatOptions { + systemPrompt?: string; + onWriteFile?: (content: string, fileName: string) => void; +} + +// 关键逻辑:第一条消息注入 systemPrompt +const sendMessage = async (content: string, ...) => { + let messageToSend = content; + const isFirstMessage = messages.filter(m => m.role === "user").length === 0; + + if (systemPrompt && isFirstMessage) { + messageToSend = `${systemPrompt}\n\n---\n\n用户请求:${content}`; + } + + await sendAgentMessageStream(messageToSend, ...); +}; +``` + +### 4. StreamingRenderer.tsx - 流式渲染 + +解析 AI 响应并触发文件写入回调。 + +```typescript +// src/components/agent/chat/components/StreamingRenderer.tsx + +interface Props { + content: string; + isStreaming: boolean; + onWriteFile?: (content: string, fileName: string) => void; + // ... +} + +// 解析 write_file 并触发回调 +useEffect(() => { + if (!onWriteFile) return; + + for (const part of parsedContent.parts) { + if (part.type === "write_file" && part.filePath) { + onWriteFile(part.content, part.filePath); + } + } +}, [parsedContent.parts, onWriteFile]); +``` + +### 5. AgentChatPage - 画布联动 + +处理文件写入,更新画布状态。 + +```typescript +// src/components/agent/chat/index.tsx + +const handleWriteFile = useCallback((content: string, fileName: string) => { + // General 主题使用专门的画布 + if (activeTheme === "general") { + setGeneralCanvasState({ + isOpen: true, + contentType: "markdown", + content, + filename: fileName, + }); + setLayoutMode("chat-canvas"); + return; + } + + // 其他主题使用 CanvasFactory + setCanvasState(createInitialDocumentState(content)); + setLayoutMode("chat-canvas"); +}, [activeTheme]); +``` + +## 主题类型 + +| 主题 | 说明 | 文件体系 | +|------|------|----------| +| general | 通用对话 | 无固定文件 | +| social-media | 社媒内容 | brief.md → draft.md → article.md | +| poster | 图文海报 | brief.md → copywriting.md → design.md | +| music | 歌词曲谱 | song-spec.md → lyrics-draft.md → lyrics-final.txt | +| video | 短视频 | brief.md → outline.md → script.md | +| novel | 小说创作 | brief.md → outline.md → chapter.md | +| document | 办公文档 | brief.md → outline.md → draft.md | + +## 创作模式 + +| 模式 | 说明 | AI 行为 | +|------|------|---------| +| guided | 引导模式 | 通过表单逐步引导用户创作 | +| fast | 快速模式 | 收集需求后直接生成完整内容 | +| hybrid | 混合模式 | AI 写框架,用户填核心内容 | +| framework | 框架模式 | 用户提供框架,AI 按框架填充 | + +## 注意事项 + +### Aster 框架限制 + +Aster 框架的 `SessionConfig` 不支持 session 级别的 system prompt,因此采用**消息注入**方案: +- 在第一条用户消息前注入 systemPrompt +- 后续消息不再注入(避免重复) + +### 画布触发条件 + +1. AI 响应包含 `` 标签 +2. `StreamingRenderer` 解析到标签 +3. 调用 `onWriteFile` 回调 +4. `AgentChatPage` 更新画布状态 +5. `layoutMode` 切换为 `chat-canvas` + +## 相关文档 + +- [aster-integration.md](aster-integration.md) - Aster 框架集成 +- [components.md](components.md) - 组件系统 +- [hooks.md](hooks.md) - React Hooks diff --git a/docs/aiprompts/hooks.md b/docs/aiprompts/hooks.md index 38d22843b..4de945580 100644 --- a/docs/aiprompts/hooks.md +++ b/docs/aiprompts/hooks.md @@ -9,6 +9,7 @@ ``` src/hooks/ ├── index.ts # 导出入口 +├── useUnifiedChat.ts # 统一对话 Hook(新) ├── useProviderPool.ts # 凭证池管理 ├── useOAuthCredentials.ts # OAuth 凭证 ├── useFlowEvents.ts # 流量事件 @@ -20,6 +21,45 @@ src/hooks/ ## 核心 Hooks +### useUnifiedChat(统一对话) + +统一的对话 Hook,支持三种模式:Agent、General、Creator。 + +```typescript +import { useUnifiedChat } from "@/hooks/useUnifiedChat"; + +// Agent 模式 - 支持工具调用 +const { messages, sendMessage, stopGeneration } = useUnifiedChat({ + mode: "agent", + providerType: "claude", + model: "claude-sonnet-4-20250514", +}); + +// Creator 模式 - 支持画布输出 +const creatorChat = useUnifiedChat({ + mode: "creator", + systemPrompt: "你是内容创作助手...", + onCanvasUpdate: (path, content) => { /* 更新画布 */ }, + onWriteFile: (content, fileName) => { /* 文件写入 */ }, +}); + +// General 模式 - 纯文本对话 +const generalChat = useUnifiedChat({ mode: "general" }); +``` + +**返回值**: +- `session` - 当前会话 +- `messages` - 消息列表 +- `isLoading` / `isSending` - 状态 +- `createSession()` / `loadSession()` / `deleteSession()` - 会话管理 +- `sendMessage()` / `stopGeneration()` - 消息操作 +- `configureProvider()` - Provider 配置 + +**相关文件**: +- 类型定义:`src/types/chat.ts` +- API 封装:`src/lib/api/unified-chat.ts` +- 架构文档:`docs/prd/chat-architecture-redesign.md` + ### useProviderPool ```typescript diff --git a/package.json b/package.json index 137edf94d..176abcb84 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.52.0", + "version": "0.53.0", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 2ef339810..5706086c9 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -202,7 +202,7 @@ checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" [[package]] name = "aster" -version = "0.5.2" +version = "0.5.3" dependencies = [ "ahash", "anyhow", @@ -2112,7 +2112,7 @@ dependencies = [ "dtoa-short", "itoa", "matches", - "phf 0.10.1", + "phf 0.8.0", "proc-macro2", "quote", "smallvec", @@ -2128,7 +2128,7 @@ dependencies = [ "cssparser-macros", "dtoa-short", "itoa", - "phf 0.11.3", + "phf 0.8.0", "smallvec", ] @@ -3988,7 +3988,7 @@ dependencies = [ "js-sys", "log", "wasm-bindgen", - "windows-core 0.57.0", + "windows-core 0.56.0", ] [[package]] @@ -5307,7 +5307,7 @@ version = "0.7.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ff32365de1b6743cb203b710788263c44a03de03802daf96092f2da4fe6ba4d7" dependencies = [ - "proc-macro-crate 2.0.2", + "proc-macro-crate 1.3.1", "proc-macro2", "quote", "syn 2.0.114", @@ -6024,7 +6024,9 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3dfb61232e34fcb633f43d12c58f83c1df82962dcdfa565a4e866ffc17dafe12" dependencies = [ + "phf_macros 0.8.0", "phf_shared 0.8.0", + "proc-macro-hack", ] [[package]] @@ -6033,9 +6035,7 @@ version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fabbf1ead8a5bcbc20f5f8b939ee3f5b0f6f281b6ad3468b84656b658b455259" dependencies = [ - "phf_macros 0.10.0", "phf_shared 0.10.0", - "proc-macro-hack", ] [[package]] @@ -6139,12 +6139,12 @@ dependencies = [ [[package]] name = "phf_macros" -version = "0.10.0" +version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "58fdf3184dd560f160dd73922bea2d5cd6e8f064bf4b13110abd81b03697b4e0" +checksum = "7f6fde18ff429ffc8fe78e2bf7f8b7a5a5a6e2a8b58bc5a9ac69198bbda9189c" dependencies = [ - "phf_generator 0.10.0", - "phf_shared 0.10.0", + "phf_generator 0.8.0", + "phf_shared 0.8.0", "proc-macro-hack", "proc-macro2", "quote", @@ -6545,7 +6545,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8a56d757972c98b346a9b766e3f02746cde6dd1cd1d1d563472929fdd74bec4d" dependencies = [ "anyhow", - "itertools 0.14.0", + "itertools 0.12.1", "proc-macro2", "quote", "syn 2.0.114", @@ -6553,7 +6553,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.50.0" +version = "0.53.0" dependencies = [ "anyhow", "arboard", @@ -6635,7 +6635,7 @@ dependencies = [ [[package]] name = "proxycast-core" -version = "0.52.0" +version = "0.53.0" dependencies = [ "chrono", "dirs 5.0.1", @@ -6651,7 +6651,7 @@ dependencies = [ [[package]] name = "proxycast-infra" -version = "0.52.0" +version = "0.53.0" dependencies = [ "chrono", "dashmap 5.5.3", @@ -7969,7 +7969,7 @@ version = "3.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8b1fdf65dd6331831494dd616b30351c38e96e45921a27745cf98490458b90bb" dependencies = [ - "dirs 6.0.0", + "dirs 4.0.0", ] [[package]] diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 9f15e018a..477e6c833 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -3,7 +3,7 @@ members = ["crates/*"] resolver = "2" [workspace.package] -version = "0.52.0" +version = "0.53.0" edition = "2021" authors = ["you"] repository = "https://github.com/aiclientproxy/proxycast" @@ -103,9 +103,9 @@ enigo = "0.3" # Aster Agent Framework # 开发时使用本地 aster-rust,CI/CD 使用远程 GitHub 仓库 # 本地开发: path = "../../../astercloud/aster-rust/crates/aster" (相对 src-tauri/) -# CI/CD: git = "https://github.com/astercloud/aster-rust", tag = "v0.5.2" +# CI/CD: git = "https://github.com/astercloud/aster-rust", tag = "v0.5.3" # aster = { version = "0.5.1", path = "../../../astercloud/aster-rust/crates/aster" } -aster = { git = "https://github.com/astercloud/aster-rust", tag = "v0.5.2" } +aster = { git = "https://github.com/astercloud/aster-rust", tag = "v0.5.3" } # Tauri @@ -164,7 +164,7 @@ version = "2.4" [package] name = "proxycast" -version = "0.50.0" +version = "0.53.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" diff --git a/src-tauri/src/agent/aster_agent.rs b/src-tauri/src/agent/aster_agent.rs index 6c66ecc11..3926ea5e6 100644 --- a/src-tauri/src/agent/aster_agent.rs +++ b/src-tauri/src/agent/aster_agent.rs @@ -4,6 +4,7 @@ //! 处理消息发送、事件流转换和会话管理 use crate::agent::aster_state::{AsterAgentState, SessionConfigBuilder}; +use crate::database::DbConnection; use aster::conversation::message::Message; use aster::session::SessionManager; use futures::StreamExt; @@ -20,6 +21,7 @@ impl AsterAgentWrapper { /// /// # Arguments /// * `state` - Aster Agent 状态 + /// * `db` - 数据库连接 /// * `app` - Tauri AppHandle,用于发送事件 /// * `message` - 用户消息文本 /// * `session_id` - 会话 ID @@ -29,14 +31,15 @@ impl AsterAgentWrapper { /// 成功时返回 Ok(()),失败时返回错误信息 pub async fn send_message( state: &AsterAgentState, + db: &DbConnection, app: &AppHandle, message: String, session_id: String, event_name: String, ) -> Result<(), String> { - // 1. 初始化检查 + // 1. 初始化检查(使用带数据库的版本) if !state.is_initialized().await { - state.init_agent().await?; + state.init_agent_with_db(db).await?; } // 2. 创建取消令牌 diff --git a/src-tauri/src/agent/aster_state.rs b/src-tauri/src/agent/aster_state.rs index f93cc1058..79b5e66d7 100644 --- a/src-tauri/src/agent/aster_state.rs +++ b/src-tauri/src/agent/aster_state.rs @@ -3,8 +3,21 @@ //! 管理 Aster Agent 实例和相关状态 //! 提供 Tauri 应用与 Aster 框架的桥接 //! 支持从 ProxyCast 凭证池自动选择凭证 +//! +//! ## 重要:SessionStore 注入 +//! +//! 为了让 Aster Agent 的消息存储到 ProxyCast 数据库,必须在创建 Agent 时 +//! 注入 `ProxyCastSessionStore`。使用 `init_agent_with_db()` 方法而不是 `init_agent()`。 +//! +//! ## Agent 身份配置 +//! +//! 通过 Aster 框架的 `AgentIdentity` API 设置 ProxyCast 专属的 Agent 身份, +//! 包括名称、语言偏好、产品描述等。这是架构层面的正确做法, +//! 而不是简单地追加提示词。 +//! +//! 参考文档:`docs/prd/chat-architecture-redesign.md` -use aster::agents::{Agent, SessionConfig}; +use aster::agents::{Agent, AgentIdentity, SessionConfig}; use aster::model::ModelConfig; use std::sync::Arc; use tokio::sync::RwLock; @@ -14,6 +27,7 @@ use crate::agent::credential_bridge::{ create_aster_provider, AsterProviderConfig, CredentialBridge, }; use crate::database::DbConnection; +use crate::services::aster_session_store::ProxyCastSessionStore; /// Provider 配置信息 #[derive(Debug, Clone)] @@ -61,15 +75,64 @@ impl AsterAgentState { } } - /// 初始化 Agent + /// 初始化 Agent(带数据库连接) /// - /// 如果 Agent 尚未初始化,则创建新的 Agent 实例 + /// 创建 Agent 并注入 ProxyCastSessionStore,确保消息存储到 ProxyCast 数据库。 + /// 同时设置 ProxyCast 专属的 Agent 身份(名称、语言、描述)。 + /// + /// **推荐使用此方法**而不是 `init_agent()`。 + /// + /// # 参数 + /// - `db`: 数据库连接,用于创建 SessionStore + pub async fn init_agent_with_db(&self, db: &DbConnection) -> Result<(), String> { + let mut agent_guard = self.agent.write().await; + if agent_guard.is_none() { + // 创建 SessionStore + let session_store = Arc::new(ProxyCastSessionStore::new(db.clone())); + + // 创建 Agent 并注入 SessionStore + let agent = Agent::new().with_session_store(session_store); + + // 使用异步方法设置 ProxyCast 专属身份 + let identity = Self::create_proxycast_identity(); + agent.set_identity(identity).await; + + *agent_guard = Some(agent); + tracing::info!( + "[AsterAgent] Agent 初始化成功,已注入 ProxyCastSessionStore 和 ProxyCast 身份" + ); + } + Ok(()) + } + + /// 创建 ProxyCast 专属的 Agent 身份配置 + fn create_proxycast_identity() -> AgentIdentity { + AgentIdentity::new("ProxyCast 助手") + .with_language("Chinese") + .with_description( + "ProxyCast 是一个 AI 代理服务应用,帮助用户管理和使用各种 AI 模型的凭证。", + ) + .with_custom_prompt(PROXYCAST_IDENTITY_PROMPT.to_string()) + } + + /// 初始化 Agent(无数据库版本) + /// + /// **警告**:此方法创建的 Agent 不会将消息存储到 ProxyCast 数据库, + /// 消息会存储到 Aster 默认的 `~/.aster/sessions.db`。 + /// + /// 建议使用 `init_agent_with_db()` 代替。 + #[deprecated( + since = "0.1.0", + note = "请使用 init_agent_with_db() 以确保消息存储到 ProxyCast 数据库" + )] pub async fn init_agent(&self) -> Result<(), String> { let mut agent_guard = self.agent.write().await; if agent_guard.is_none() { let agent = Agent::new(); *agent_guard = Some(agent); - tracing::info!("Aster Agent initialized"); + tracing::warn!( + "[AsterAgent] Agent 初始化(无 SessionStore),消息将存储到 Aster 默认数据库" + ); } Ok(()) } @@ -77,13 +140,19 @@ impl AsterAgentState { /// 配置 Provider /// /// 根据配置创建并设置 Provider + /// + /// # 参数 + /// - `config`: Provider 配置 + /// - `session_id`: 会话 ID + /// - `db`: 数据库连接(用于初始化 Agent) pub async fn configure_provider( &self, config: ProviderConfig, session_id: &str, + db: &DbConnection, ) -> Result<(), String> { - // 确保 Agent 已初始化 - self.init_agent().await?; + // 确保 Agent 已初始化(使用带数据库的版本) + self.init_agent_with_db(db).await?; // 设置环境变量(Aster 的 provider 从环境变量读取配置) self.set_provider_env_vars(&config); @@ -135,8 +204,8 @@ impl AsterAgentState { model: &str, session_id: &str, ) -> Result { - // 确保 Agent 已初始化 - self.init_agent().await?; + // 确保 Agent 已初始化(使用带数据库的版本) + self.init_agent_with_db(db).await?; // 从凭证池选择凭证并获取配置 let aster_config = self @@ -329,6 +398,7 @@ impl AsterAgentState { pub struct SessionConfigBuilder { id: String, max_turns: Option, + system_prompt: Option, } impl SessionConfigBuilder { @@ -336,6 +406,7 @@ impl SessionConfigBuilder { Self { id: id.into(), max_turns: None, + system_prompt: None, } } @@ -344,12 +415,18 @@ impl SessionConfigBuilder { self } + pub fn system_prompt(mut self, prompt: impl Into) -> Self { + self.system_prompt = Some(prompt.into()); + self + } + pub fn build(self) -> SessionConfig { SessionConfig { id: self.id, schedule_id: None, max_turns: self.max_turns, retry_config: None, + system_prompt: self.system_prompt, } } } @@ -378,6 +455,7 @@ mod tests { let state = AsterAgentState::new(); assert!(!state.is_initialized().await); + #[allow(deprecated)] state.init_agent().await.unwrap(); assert!(state.is_initialized().await); } @@ -397,3 +475,33 @@ mod tests { assert!(!state.cancel_session(session_id).await); } } + +// ============================================================================= +// ProxyCast Agent 身份提示词 +// ============================================================================= + +/// ProxyCast 专属的 Agent 身份提示词 +/// +/// 这是完整的身份定义,会替换 Aster 框架默认的 "aster by Block" 身份。 +/// 框架的能力描述(Extensions、Response Guidelines)会自动追加。 +const PROXYCAST_IDENTITY_PROMPT: &str = r#"你是 ProxyCast 助手,一个专业、友好的 AI 技术伙伴。 + +## 关于 ProxyCast + +ProxyCast 是一个 AI 代理服务应用,帮助用户: +- 管理多个 AI 模型提供商的凭证(OpenAI、Claude、Gemini、Kiro 等) +- 通过统一的 API 接口访问不同的 AI 模型 +- 实现凭证池的负载均衡和健康检查 + +## 语言规范 + +1. **始终使用中文回复**:除非用户明确要求使用其他语言 +2. **代码注释使用中文**:生成代码时,注释应使用中文 +3. **技术术语保持原文**:API、JSON、HTTP、Token 等专业术语保持英文 + +## 交互风格 + +- 简洁专业,直接给出解决方案 +- 友好但不啰嗦,像经验丰富的技术伙伴 +- 遇到问题时,先分析原因再提供方案 +"#; diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index 0b59758c2..290c64d62 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -1228,6 +1228,16 @@ pub fn run() { 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, + commands::unified_chat_cmd::chat_get_session, + commands::unified_chat_cmd::chat_delete_session, + commands::unified_chat_cmd::chat_rename_session, + commands::unified_chat_cmd::chat_get_messages, + commands::unified_chat_cmd::chat_send_message, + commands::unified_chat_cmd::chat_stop_generation, + commands::unified_chat_cmd::chat_configure_provider, // Workspace commands commands::workspace_cmd::workspace_create, commands::workspace_cmd::workspace_list, diff --git a/src-tauri/src/app/setup.rs b/src-tauri/src/app/setup.rs index 5512fed6e..14cee3531 100644 --- a/src-tauri/src/app/setup.rs +++ b/src-tauri/src/app/setup.rs @@ -9,6 +9,7 @@ use tauri::{App, Manager}; use crate::agent::AsterAgentState; use crate::database; use crate::flow_monitor::FlowInterceptor; +use crate::services::aster_session_store::ProxyCastSessionStore; use crate::services::provider_pool_service::ProviderPoolService; use crate::services::token_cache_service::TokenCacheService; use crate::telemetry; @@ -32,6 +33,18 @@ pub fn setup_app( flow_monitor: Arc, flow_interceptor: Arc, ) -> Result<(), Box> { + // 注册全局 SessionStore(作为后备方案) + // 注意:主要的 SessionStore 注入在 AsterAgentState::init_agent_with_db() 中完成 + // 这里的全局注册是为了兼容可能直接使用 SessionManager 静态方法的代码 + let session_store = Arc::new(ProxyCastSessionStore::new(db.clone())); + tauri::async_runtime::block_on(async { + if let Err(e) = aster::session::set_global_session_store(session_store).await { + tracing::warn!("[启动] 注册全局 SessionStore 失败(可能已注册): {}", e); + } else { + tracing::info!("[启动] 全局 ProxyCastSessionStore 已注册(后备方案)"); + } + }); + // 初始化托盘管理器 match TrayManager::new(app.handle()) { Ok(tray_manager) => { diff --git a/src-tauri/src/commands/agent_cmd.rs b/src-tauri/src/commands/agent_cmd.rs index c8842245d..7879cb78d 100644 --- a/src-tauri/src/commands/agent_cmd.rs +++ b/src-tauri/src/commands/agent_cmd.rs @@ -33,6 +33,7 @@ pub struct CreateSessionResponse { pub async fn agent_start_process( agent_state: State<'_, AsterAgentState>, app_state: State<'_, AppState>, + db: State<'_, DbConnection>, _port: Option, ) -> Result { tracing::info!("[Agent] 初始化 Aster Agent"); @@ -50,7 +51,7 @@ pub async fn agent_start_process( return Err("ProxyCast API Server 未运行,请先启动服务器".to_string()); } - agent_state.init_agent().await?; + agent_state.init_agent_with_db(&db).await?; let base_url = format!("http://{}:{}", host, port); @@ -122,8 +123,8 @@ pub async fn agent_create_session( skills.as_ref().map(|s| s.len()) ); - // 初始化 Agent - agent_state.init_agent().await?; + // 初始化 Agent(使用带数据库的版本) + agent_state.init_agent_with_db(&db).await?; // 生成会话 ID let session_id = uuid::Uuid::new_v4().to_string(); diff --git a/src-tauri/src/commands/aster_agent_cmd.rs b/src-tauri/src/commands/aster_agent_cmd.rs index 50bfd9a00..e2365b2fe 100644 --- a/src-tauri/src/commands/aster_agent_cmd.rs +++ b/src-tauri/src/commands/aster_agent_cmd.rs @@ -9,6 +9,7 @@ use crate::agent::event_converter::convert_agent_event; use crate::agent::{ AsterAgentState, AsterAgentWrapper, SessionDetail, SessionInfo, TauriAgentEvent, }; +use crate::database::dao::agent::AgentDao; use crate::database::DbConnection; use aster::conversation::message::Message; use aster::session::SessionManager; @@ -83,10 +84,11 @@ pub struct ConfigureFromPoolRequest { #[tauri::command] pub async fn aster_agent_init( state: State<'_, AsterAgentState>, + db: State<'_, DbConnection>, ) -> Result { tracing::info!("[AsterAgent] 初始化 Agent"); - state.init_agent().await?; + state.init_agent_with_db(&db).await?; let provider_config = state.get_provider_config().await; @@ -105,6 +107,7 @@ pub async fn aster_agent_init( #[tauri::command] pub async fn aster_agent_configure_provider( state: State<'_, AsterAgentState>, + db: State<'_, DbConnection>, request: ConfigureProviderRequest, session_id: String, ) -> Result { @@ -123,7 +126,7 @@ pub async fn aster_agent_configure_provider( }; state - .configure_provider(config.clone(), &session_id) + .configure_provider(config.clone(), &session_id, &db) .await?; Ok(AsterAgentStatus { @@ -211,6 +214,7 @@ pub struct ImageInput { pub async fn aster_agent_chat_stream( app: AppHandle, state: State<'_, AsterAgentState>, + db: State<'_, DbConnection>, request: AsterChatRequest, ) -> Result<(), String> { tracing::info!( @@ -219,15 +223,26 @@ pub async fn aster_agent_chat_stream( request.event_name ); - // 确保 Agent 已初始化 + // 确保 Agent 已初始化(使用带数据库的版本,注入 SessionStore) if !state.is_initialized().await { - state.init_agent().await?; + state.init_agent_with_db(&db).await?; } - // 确保 session 在 Aster 数据库中存在 + // 确保 session 在数据库中存在 // 如果 session 不存在,自动创建 let session_id = ensure_session_exists(&request.session_id).await?; + // 从数据库读取 session 的 system_prompt + let system_prompt = { + let db_conn = db + .lock() + .map_err(|e| format!("获取数据库连接失败: {}", e))?; + let session = AgentDao::get_session(&db_conn, &session_id) + .map_err(|e| format!("获取 session 失败: {}", e))? + .ok_or_else(|| format!("Session 不存在: {}", session_id))?; + session.system_prompt + }; + // 如果提供了 Provider 配置,则配置 Provider if let Some(provider_config) = &request.provider_config { let config = ProviderConfig { @@ -237,7 +252,7 @@ pub async fn aster_agent_chat_stream( base_url: provider_config.base_url.clone(), credential_uuid: None, }; - state.configure_provider(config, &session_id).await?; + state.configure_provider(config, &session_id, &db).await?; } // 检查 Provider 是否已配置 @@ -251,8 +266,12 @@ pub async fn aster_agent_chat_stream( // 创建用户消息 let user_message = Message::user().with_text(&request.message); - // 创建会话配置 - let session_config = SessionConfigBuilder::new(&session_id).build(); + // 创建会话配置,包含 system_prompt + let mut session_config_builder = SessionConfigBuilder::new(&session_id); + if let Some(prompt) = system_prompt { + session_config_builder = session_config_builder.system_prompt(prompt); + } + let session_config = session_config_builder.build(); // 获取 Agent Arc 并保持 guard 在整个流处理期间存活 let agent_arc = state.get_agent_arc(); diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index dcdf185ef..447bf26d6 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -38,6 +38,7 @@ pub mod telemetry_cmd; pub mod terminal_cmd; pub mod tool_hooks; pub mod tray_cmd; +pub mod unified_chat_cmd; pub mod update_cmd; pub mod usage_cmd; pub mod websocket_cmd; diff --git a/src-tauri/src/commands/unified_chat_cmd.rs b/src-tauri/src/commands/unified_chat_cmd.rs new file mode 100644 index 000000000..a27017569 --- /dev/null +++ b/src-tauri/src/commands/unified_chat_cmd.rs @@ -0,0 +1,424 @@ +//! 统一对话命令模块 +//! +//! 提供统一的对话 API,支持多种对话模式: +//! - Agent: AI Agent 模式,支持工具调用 +//! - General: 通用对话模式,纯文本 +//! - Creator: 内容创作模式,支持画布输出 +//! +//! ## 设计原则 +//! - 单一入口:所有对话场景使用同一套 API +//! - 模式化设计:通过 ChatMode 区分不同场景 +//! - Aster 引擎:底层使用 Aster Agent 处理对话 +//! +//! ## 参考文档 +//! - `docs/prd/chat-architecture-redesign.md` + +use crate::agent::aster_state::SessionConfigBuilder; +use crate::agent::event_converter::convert_agent_event; +use crate::agent::{AsterAgentState, TauriAgentEvent}; +use crate::database::dao::chat::{ChatDao, ChatMessage, ChatMode, ChatSession}; +use crate::database::DbConnection; +use aster::conversation::message::Message; +use futures::StreamExt; +use serde::{Deserialize, Serialize}; +use tauri::{AppHandle, Emitter, State}; + +// ============================================================================ +// 请求/响应结构 +// ============================================================================ + +/// 创建会话请求 +#[derive(Debug, Deserialize)] +pub struct CreateSessionRequest { + /// 对话模式 + pub mode: ChatMode, + /// 会话标题(可选) + pub title: Option, + /// 系统提示词(可选) + pub system_prompt: Option, + /// Provider 类型(可选) + pub provider_type: Option, + /// 模型名称(可选) + pub model: Option, + /// 扩展元数据(可选) + pub metadata: Option, +} + +/// 发送消息请求 +#[derive(Debug, Deserialize)] +pub struct SendMessageRequest { + /// 会话 ID + pub session_id: String, + /// 消息内容 + pub message: String, + /// 事件名称(用于前端监听) + pub event_name: String, + /// 图片输入(可选) + pub images: Option>, +} + +/// 图片输入 +#[derive(Debug, Deserialize)] +pub struct ImageInput { + pub data: String, + pub media_type: String, +} + +/// 会话信息响应 +#[derive(Debug, Serialize)] +pub struct SessionResponse { + pub id: String, + pub mode: ChatMode, + pub title: Option, + pub model: Option, + pub created_at: String, + pub updated_at: String, + pub message_count: usize, +} + +impl From for SessionResponse { + fn from(session: ChatSession) -> Self { + Self { + id: session.id, + mode: session.mode, + title: session.title, + model: session.model, + created_at: session.created_at, + updated_at: session.updated_at, + message_count: 0, + } + } +} + +// ============================================================================ +// 会话管理命令 +// ============================================================================ + +/// 创建新会话 +/// +/// 统一的会话创建入口,支持所有对话模式 +#[tauri::command] +pub async fn chat_create_session( + db: State<'_, DbConnection>, + agent_state: State<'_, AsterAgentState>, + request: CreateSessionRequest, +) -> Result { + let now = chrono::Utc::now().to_rfc3339(); + let session_id = uuid::Uuid::new_v4().to_string(); + + // 创建会话 + let session = ChatSession { + id: session_id.clone(), + mode: request.mode, + title: request.title, + system_prompt: request.system_prompt.clone(), + model: request.model.clone(), + provider_type: request.provider_type.clone(), + credential_uuid: None, + metadata: request.metadata, + created_at: now.clone(), + updated_at: now, + }; + + // 保存到数据库 + { + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?; + ChatDao::create_session(&conn, &session).map_err(|e| format!("创建会话失败: {}", e))?; + } + + // 初始化 Aster Agent(如果是 Agent 或 Creator 模式) + if matches!(request.mode, ChatMode::Agent | ChatMode::Creator) { + agent_state.init_agent_with_db(&db).await?; + + // 如果指定了 Provider,配置它 + if let (Some(provider_type), Some(model)) = (&request.provider_type, &request.model) { + agent_state + .configure_provider_from_pool(&db, provider_type, model, &session_id) + .await?; + } + } + + tracing::info!( + "[UnifiedChat] 创建会话: id={}, mode={:?}", + session_id, + request.mode + ); + + Ok(SessionResponse::from(session)) +} + +/// 获取会话列表 +/// +/// 可选按模式过滤 +#[tauri::command] +pub async fn chat_list_sessions( + db: State<'_, DbConnection>, + mode: Option, +) -> Result, String> { + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?; + + let sessions = + ChatDao::list_sessions(&conn, mode).map_err(|e| format!("获取会话列表失败: {}", e))?; + + let mut result: Vec = Vec::new(); + for session in sessions { + let message_count = ChatDao::get_message_count(&conn, &session.id).unwrap_or(0); + let mut resp = SessionResponse::from(session); + resp.message_count = message_count; + result.push(resp); + } + + Ok(result) +} + +/// 获取会话详情 +#[tauri::command] +pub async fn chat_get_session( + db: State<'_, DbConnection>, + session_id: String, +) -> Result { + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?; + + let session = ChatDao::get_session(&conn, &session_id) + .map_err(|e| format!("获取会话失败: {}", e))? + .ok_or_else(|| "会话不存在".to_string())?; + + let message_count = ChatDao::get_message_count(&conn, &session_id).unwrap_or(0); + let mut resp = SessionResponse::from(session); + resp.message_count = message_count; + + Ok(resp) +} + +/// 删除会话 +#[tauri::command] +pub async fn chat_delete_session( + db: State<'_, DbConnection>, + session_id: String, +) -> Result { + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?; + + let deleted = + ChatDao::delete_session(&conn, &session_id).map_err(|e| format!("删除会话失败: {}", e))?; + + if deleted { + tracing::info!("[UnifiedChat] 删除会话: id={}", session_id); + } + + Ok(deleted) +} + +/// 重命名会话 +#[tauri::command] +pub async fn chat_rename_session( + db: State<'_, DbConnection>, + session_id: String, + title: String, +) -> Result<(), String> { + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?; + + ChatDao::update_title(&conn, &session_id, &title) + .map_err(|e| format!("重命名会话失败: {}", e))?; + + tracing::info!( + "[UnifiedChat] 重命名会话: id={}, title={}", + session_id, + title + ); + + Ok(()) +} + +// ============================================================================ +// 消息管理命令 +// ============================================================================ + +/// 获取会话消息列表 +#[tauri::command] +pub async fn chat_get_messages( + db: State<'_, DbConnection>, + session_id: String, + limit: Option, +) -> Result, String> { + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?; + + let messages = ChatDao::get_messages(&conn, &session_id, limit) + .map_err(|e| format!("获取消息失败: {}", e))?; + + Ok(messages) +} + +/// 发送消息并获取流式响应 +/// +/// 统一的消息发送入口,根据会话模式选择处理方式 +#[tauri::command] +pub async fn chat_send_message( + app: AppHandle, + db: State<'_, DbConnection>, + agent_state: State<'_, AsterAgentState>, + request: SendMessageRequest, +) -> Result<(), String> { + tracing::info!( + "[UnifiedChat] 发送消息: session={}, event={}", + request.session_id, + request.event_name + ); + + // 获取会话信息 + let session = { + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?; + ChatDao::get_session(&conn, &request.session_id) + .map_err(|e| format!("获取会话失败: {}", e))? + .ok_or_else(|| "会话不存在".to_string())? + }; + + // 根据模式处理 + match session.mode { + ChatMode::Agent | ChatMode::Creator => { + // 使用 Aster Agent 处理 + send_message_with_aster( + &app, + &db, + &agent_state, + &request.session_id, + &request.message, + &request.event_name, + session.system_prompt.as_deref(), + ) + .await + } + ChatMode::General => { + // 通用模式:也使用 Aster Agent,但不启用工具 + send_message_with_aster( + &app, + &db, + &agent_state, + &request.session_id, + &request.message, + &request.event_name, + session.system_prompt.as_deref(), + ) + .await + } + } +} + +/// 使用 Aster Agent 发送消息 +async fn send_message_with_aster( + app: &AppHandle, + db: &DbConnection, + agent_state: &AsterAgentState, + session_id: &str, + message: &str, + event_name: &str, + system_prompt: Option<&str>, +) -> Result<(), String> { + // 确保 Agent 已初始化 + if !agent_state.is_initialized().await { + agent_state.init_agent_with_db(db).await?; + } + + // 检查 Provider 是否已配置 + if !agent_state.is_provider_configured().await { + return Err("Provider 未配置,请先配置凭证".to_string()); + } + + // 创建取消令牌 + let cancel_token = agent_state.create_cancel_token(session_id).await; + + // 构建消息(如果有 system_prompt 且是第一条消息,注入到消息前面) + let final_message = if let Some(prompt) = system_prompt { + format!("{}\n\n{}", prompt, message) + } else { + message.to_string() + }; + + let user_message = Message::user().with_text(&final_message); + let session_config = SessionConfigBuilder::new(session_id).build(); + + // 获取 Agent 引用 + let agent_arc = agent_state.get_agent_arc(); + let guard = agent_arc.read().await; + let agent = guard.as_ref().ok_or("Agent 未初始化")?; + + // 调用 Agent + let stream_result = agent + .reply(user_message, session_config, Some(cancel_token.clone())) + .await; + + match stream_result { + Ok(mut stream) => { + while let Some(event_result) = stream.next().await { + match event_result { + Ok(agent_event) => { + let tauri_events = convert_agent_event(agent_event); + for tauri_event in tauri_events { + if let Err(e) = app.emit(event_name, &tauri_event) { + tracing::error!("[UnifiedChat] 发送事件失败: {}", e); + } + } + } + Err(e) => { + let error_event = TauriAgentEvent::Error { + message: format!("流错误: {}", e), + }; + let _ = app.emit(event_name, &error_event); + } + } + } + + // 发送完成事件 + let done_event = TauriAgentEvent::FinalDone { usage: None }; + let _ = app.emit(event_name, &done_event); + } + Err(e) => { + let error_event = TauriAgentEvent::Error { + message: format!("Agent 错误: {}", e), + }; + let _ = app.emit(event_name, &error_event); + return Err(format!("Agent 错误: {}", e)); + } + } + + // 清理取消令牌 + agent_state.remove_cancel_token(session_id).await; + + Ok(()) +} + +/// 停止生成 +#[tauri::command] +pub async fn chat_stop_generation( + agent_state: State<'_, AsterAgentState>, + session_id: String, +) -> Result { + tracing::info!("[UnifiedChat] 停止生成: session={}", session_id); + Ok(agent_state.cancel_session(&session_id).await) +} + +/// 配置会话的 Provider +#[tauri::command] +pub async fn chat_configure_provider( + db: State<'_, DbConnection>, + agent_state: State<'_, AsterAgentState>, + session_id: String, + provider_type: String, + model: String, +) -> Result<(), String> { + tracing::info!( + "[UnifiedChat] 配置 Provider: session={}, provider={}, model={}", + session_id, + provider_type, + model + ); + + // 确保 Agent 已初始化 + agent_state.init_agent_with_db(&db).await?; + + // 配置 Provider + agent_state + .configure_provider_from_pool(&db, &provider_type, &model, &session_id) + .await?; + + Ok(()) +} diff --git a/src-tauri/src/database/dao/agent.rs b/src-tauri/src/database/dao/agent.rs index 0f94e5a69..6470b53bf 100644 --- a/src-tauri/src/database/dao/agent.rs +++ b/src-tauri/src/database/dao/agent.rs @@ -5,6 +5,54 @@ use crate::agent::types::{AgentMessage, AgentSession, MessageContent, ToolCall}; use rusqlite::{params, Connection}; +/// 解析消息内容 JSON,支持多种格式 +/// +/// 支持的格式: +/// 1. Aster 格式: `[{"Text":"..."}, {"ToolRequest":...}]` +/// 2. ProxyCast 纯文本: `"string"` +/// 3. ProxyCast Parts: `[{"type":"text","text":"..."}]` +fn parse_message_content(content_json: &str) -> MessageContent { + // 尝试解析为 Aster 格式 (Vec) + if let Ok(aster_contents) = serde_json::from_str::>(content_json) { + let mut text_parts: Vec = Vec::new(); + + for item in aster_contents { + // Aster 格式: {"Text": "..."} 或 {"ToolRequest": ...} + if let Some(text) = item.get("Text").and_then(|v| v.as_str()) { + text_parts.push(text.to_string()); + } + // 也支持小写 "text" 格式 + else if let Some(text) = item.get("text").and_then(|v| v.as_str()) { + text_parts.push(text.to_string()); + } + // ProxyCast Parts 格式: {"type": "text", "text": "..."} + else if item.get("type").and_then(|v| v.as_str()) == Some("text") { + if let Some(text) = item.get("text").and_then(|v| v.as_str()) { + text_parts.push(text.to_string()); + } + } + // 忽略 ToolRequest、ToolResponse 等非文本内容 + } + + if !text_parts.is_empty() { + return MessageContent::Text(text_parts.join("\n")); + } + } + + // 尝试解析为纯文本字符串 + if let Ok(text) = serde_json::from_str::(content_json) { + return MessageContent::Text(text); + } + + // 尝试直接解析为 ProxyCast MessageContent + if let Ok(content) = serde_json::from_str::(content_json) { + return content; + } + + // 兜底:返回原始 JSON 作为文本 + MessageContent::Text(content_json.to_string()) +} + pub struct AgentDao; impl AgentDao { @@ -178,14 +226,10 @@ impl AgentDao { let tool_calls_json: Option = row.get(3)?; let tool_call_id: Option = row.get(4)?; - // 解析 JSON - let content: MessageContent = serde_json::from_str(&content_json).map_err(|e| { - rusqlite::Error::FromSqlConversionFailure( - 1, - rusqlite::types::Type::Text, - Box::new(e), - ) - })?; + // 解析 JSON - 支持多种格式 + // 1. Aster 格式: [{"Text":"..."}, {"Text":"..."}] + // 2. ProxyCast 格式: "string" 或 [{"type":"text","text":"..."}] + let content = parse_message_content(&content_json); let tool_calls: Option> = tool_calls_json .map(|json| serde_json::from_str(&json)) diff --git a/src-tauri/src/database/dao/chat.rs b/src-tauri/src/database/dao/chat.rs new file mode 100644 index 000000000..30f98ca70 --- /dev/null +++ b/src-tauri/src/database/dao/chat.rs @@ -0,0 +1,426 @@ +//! 统一对话数据访问层 +//! +//! 提供统一的会话和消息存储功能,支持多种对话模式: +//! - Agent: AI Agent 模式,支持工具调用 +//! - General: 通用对话模式,纯文本 +//! - Creator: 内容创作模式,支持画布输出 +//! +//! ## 设计原则 +//! - 单一数据源:所有对话数据统一存储 +//! - 模式化设计:通过 ChatMode 区分不同场景 +//! - 向后兼容:复用现有的 agent_sessions/agent_messages 表 + +use rusqlite::{params, Connection}; +use serde::{Deserialize, Serialize}; + +// ============================================================================ +// 数据模型 +// ============================================================================ + +/// 对话模式 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum ChatMode { + /// AI Agent 模式,支持工具调用 + Agent, + /// 通用对话模式,纯文本 + General, + /// 内容创作模式,支持画布输出 + Creator, +} + +impl Default for ChatMode { + fn default() -> Self { + Self::General + } +} + +impl std::fmt::Display for ChatMode { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + ChatMode::Agent => write!(f, "agent"), + ChatMode::General => write!(f, "general"), + ChatMode::Creator => write!(f, "creator"), + } + } +} + +impl std::str::FromStr for ChatMode { + type Err = String; + + fn from_str(s: &str) -> Result { + match s.to_lowercase().as_str() { + "agent" => Ok(ChatMode::Agent), + "general" => Ok(ChatMode::General), + "creator" => Ok(ChatMode::Creator), + _ => Err(format!("未知的对话模式: {}", s)), + } + } +} + +/// 统一会话结构 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatSession { + /// 会话 ID + pub id: String, + /// 对话模式 + pub mode: ChatMode, + /// 会话标题 + pub title: Option, + /// 系统提示词 + pub system_prompt: Option, + /// 模型名称 + pub model: Option, + /// Provider 类型 + pub provider_type: Option, + /// 凭证 UUID + pub credential_uuid: Option, + /// 扩展元数据(JSON) + pub metadata: Option, + /// 创建时间 + pub created_at: String, + /// 更新时间 + pub updated_at: String, +} + +/// 统一消息结构 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatMessage { + /// 消息 ID + pub id: i64, + /// 会话 ID + pub session_id: String, + /// 角色 (user/assistant/system/tool) + pub role: String, + /// 消息内容(JSON 格式) + pub content: serde_json::Value, + /// 工具调用信息(JSON) + pub tool_calls: Option, + /// 工具调用 ID + pub tool_call_id: Option, + /// 扩展元数据(JSON) + pub metadata: Option, + /// 创建时间 + pub created_at: String, +} + +/// 会话详情(包含消息) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatSessionDetail { + /// 会话信息 + pub session: ChatSession, + /// 消息列表 + pub messages: Vec, + /// 消息总数 + pub message_count: usize, +} + +// ============================================================================ +// 数据访问对象 +// ============================================================================ + +/// 统一对话 DAO +pub struct ChatDao; + +impl ChatDao { + // ------------------------------------------------------------------------ + // 会话管理 + // ------------------------------------------------------------------------ + + /// 创建新会话 + pub fn create_session(conn: &Connection, session: &ChatSession) -> Result<(), rusqlite::Error> { + // 使用现有的 agent_sessions 表,通过 model 字段存储 mode + // 这样可以保持向后兼容 + let mode_str = session.mode.to_string(); + let metadata_json = session + .metadata + .as_ref() + .map(|m| serde_json::to_string(m).unwrap_or_default()); + + conn.execute( + "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + session.id, + format!( + "{}:{}", + mode_str, + session.model.as_deref().unwrap_or("default") + ), + session.system_prompt, + session.title, + session.created_at, + session.updated_at, + ], + )?; + + // 如果有扩展元数据,存储到单独的字段(未来可扩展) + if metadata_json.is_some() { + tracing::debug!("[ChatDao] 会话 {} 有扩展元数据,暂存于内存", session.id); + } + + Ok(()) + } + + /// 获取会话 + pub fn get_session( + conn: &Connection, + session_id: &str, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, model, system_prompt, title, created_at, updated_at + FROM agent_sessions WHERE id = ?", + )?; + + let mut rows = stmt.query([session_id])?; + + if let Some(row) = rows.next()? { + let model_str: String = row.get(1)?; + let (mode, model) = Self::parse_mode_model(&model_str); + + Ok(Some(ChatSession { + id: row.get(0)?, + mode, + title: row.get(3)?, + system_prompt: row.get(2)?, + model, + provider_type: None, + credential_uuid: None, + metadata: None, + created_at: row.get(4)?, + updated_at: row.get(5)?, + })) + } else { + Ok(None) + } + } + + /// 获取会话列表 + pub fn list_sessions( + conn: &Connection, + mode: Option, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, model, system_prompt, title, created_at, updated_at + FROM agent_sessions ORDER BY updated_at DESC", + )?; + + let sessions: Vec = stmt + .query_map([], |row| { + let model_str: String = row.get(1)?; + let (parsed_mode, model) = Self::parse_mode_model(&model_str); + + Ok(ChatSession { + id: row.get(0)?, + mode: parsed_mode, + title: row.get(3)?, + system_prompt: row.get(2)?, + model, + provider_type: None, + credential_uuid: None, + metadata: None, + created_at: row.get(4)?, + updated_at: row.get(5)?, + }) + })? + .filter_map(|r| r.ok()) + .collect(); + + // 如果指定了模式,过滤结果 + if let Some(filter_mode) = mode { + Ok(sessions + .into_iter() + .filter(|s| s.mode == filter_mode) + .collect()) + } else { + Ok(sessions) + } + } + + /// 删除会话 + pub fn delete_session(conn: &Connection, session_id: &str) -> Result { + let rows = conn.execute("DELETE FROM agent_sessions WHERE id = ?", [session_id])?; + Ok(rows > 0) + } + + /// 更新会话标题 + pub fn update_title( + conn: &Connection, + session_id: &str, + title: &str, + ) -> Result<(), rusqlite::Error> { + let now = chrono::Utc::now().to_rfc3339(); + conn.execute( + "UPDATE agent_sessions SET title = ?, updated_at = ? WHERE id = ?", + params![title, now, session_id], + )?; + Ok(()) + } + + /// 检查会话是否存在 + pub fn session_exists(conn: &Connection, session_id: &str) -> Result { + let count: i64 = conn.query_row( + "SELECT COUNT(*) FROM agent_sessions WHERE id = ?", + [session_id], + |row| row.get(0), + )?; + Ok(count > 0) + } + + // ------------------------------------------------------------------------ + // 消息管理 + // ------------------------------------------------------------------------ + + /// 添加消息 + pub fn add_message(conn: &Connection, message: &ChatMessage) -> Result { + let content_json = serde_json::to_string(&message.content) + .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?; + + let tool_calls_json = message + .tool_calls + .as_ref() + .map(|tc| serde_json::to_string(tc)) + .transpose() + .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?; + + 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![ + message.session_id, + message.role, + content_json, + message.created_at, + tool_calls_json, + message.tool_call_id, + ], + )?; + + let id = conn.last_insert_rowid(); + + // 更新会话时间 + conn.execute( + "UPDATE agent_sessions SET updated_at = ? WHERE id = ?", + params![message.created_at, message.session_id], + )?; + + Ok(id) + } + + /// 获取会话消息 + pub fn get_messages( + conn: &Connection, + session_id: &str, + limit: Option, + ) -> Result, rusqlite::Error> { + let query = if limit.is_some() { + "SELECT id, session_id, role, content_json, timestamp, tool_calls_json, tool_call_id + FROM agent_messages WHERE session_id = ? ORDER BY id ASC LIMIT ?" + } else { + "SELECT id, session_id, role, content_json, timestamp, tool_calls_json, tool_call_id + FROM agent_messages WHERE session_id = ? ORDER BY id ASC" + }; + + let mut stmt = conn.prepare(query)?; + + let messages: Vec = if let Some(lim) = limit { + stmt.query_map(params![session_id, lim], Self::map_message_row)? + } else { + stmt.query_map([session_id], Self::map_message_row)? + } + .filter_map(|r| r.ok()) + .collect(); + + Ok(messages) + } + + /// 获取消息数量 + pub fn get_message_count( + conn: &Connection, + session_id: &str, + ) -> Result { + let count: i64 = conn.query_row( + "SELECT COUNT(*) FROM agent_messages WHERE session_id = ?", + [session_id], + |row| row.get(0), + )?; + Ok(count as usize) + } + + /// 删除会话消息 + pub fn delete_messages(conn: &Connection, session_id: &str) -> Result<(), rusqlite::Error> { + conn.execute( + "DELETE FROM agent_messages WHERE session_id = ?", + [session_id], + )?; + Ok(()) + } + + // ------------------------------------------------------------------------ + // 辅助方法 + // ------------------------------------------------------------------------ + + /// 解析 mode:model 格式的字符串 + fn parse_mode_model(model_str: &str) -> (ChatMode, Option) { + if let Some((mode_part, model_part)) = model_str.split_once(':') { + let mode = mode_part.parse().unwrap_or(ChatMode::Agent); + let model = if model_part == "default" || model_part.is_empty() { + None + } else { + Some(model_part.to_string()) + }; + (mode, model) + } else { + // 兼容旧数据:没有 mode 前缀的视为 Agent 模式 + (ChatMode::Agent, Some(model_str.to_string())) + } + } + + /// 映射消息行 + fn map_message_row(row: &rusqlite::Row) -> Result { + let content_json: String = row.get(3)?; + let tool_calls_json: Option = row.get(5)?; + + // 解析内容 JSON + let content: serde_json::Value = serde_json::from_str(&content_json).unwrap_or_else(|_| { + // 如果解析失败,包装为文本 + serde_json::json!([{"type": "text", "text": content_json}]) + }); + + let tool_calls: Option = tool_calls_json + .map(|json| serde_json::from_str(&json).ok()) + .flatten(); + + Ok(ChatMessage { + id: row.get(0)?, + session_id: row.get(1)?, + role: row.get(2)?, + content, + tool_calls, + tool_call_id: row.get(6)?, + metadata: None, + created_at: row.get(4)?, + }) + } + + /// 获取会话详情(包含消息) + pub fn get_session_detail( + conn: &Connection, + session_id: &str, + message_limit: Option, + ) -> Result, rusqlite::Error> { + let session = match Self::get_session(conn, session_id)? { + Some(s) => s, + None => return Ok(None), + }; + + let messages = Self::get_messages(conn, session_id, message_limit)?; + let message_count = Self::get_message_count(conn, session_id)?; + + Ok(Some(ChatSessionDetail { + session, + messages, + message_count, + })) + } +} diff --git a/src-tauri/src/database/dao/mod.rs b/src-tauri/src/database/dao/mod.rs index 67172dd9b..35e9c2aa4 100644 --- a/src-tauri/src/database/dao/mod.rs +++ b/src-tauri/src/database/dao/mod.rs @@ -1,5 +1,6 @@ pub mod agent; pub mod api_key_provider; +pub mod chat; pub mod general_chat; pub mod installed_plugins; pub mod mcp; diff --git a/src-tauri/src/database/migration.rs b/src-tauri/src/database/migration.rs index 5c7970c3b..7bc4f8978 100644 --- a/src-tauri/src/database/migration.rs +++ b/src-tauri/src/database/migration.rs @@ -554,3 +554,303 @@ pub fn clear_model_registry_refresh_flag(conn: &Connection) { [], ); } + +// ============================================================================ +// General Chat 数据迁移到统一表 +// ============================================================================ + +/// 执行 General Chat 数据迁移到统一表 +/// +/// 将 general_chat_sessions/messages 数据迁移到 agent_sessions/messages 表 +/// - general_chat_sessions → agent_sessions (mode 前缀为 "general:") +/// - general_chat_messages → agent_messages +pub fn migrate_general_chat_to_unified(conn: &Connection) -> Result { + // 检查是否已经迁移过 + let migrated: bool = conn + .query_row( + "SELECT value FROM settings WHERE key = 'migrated_general_chat_to_unified'", + [], + |row| row.get::<_, String>(0), + ) + .map(|v| v == "true") + .unwrap_or(false); + + if migrated { + tracing::debug!("[迁移] General Chat 已迁移过,跳过"); + return Ok(0); + } + + // 检查是否有数据需要迁移 + let general_count: i64 = conn + .query_row("SELECT COUNT(*) FROM general_chat_sessions", [], |row| { + row.get(0) + }) + .unwrap_or(0); + + if general_count == 0 { + tracing::info!("[迁移] general_chat_sessions 表为空,无需迁移"); + // 标记迁移完成 + conn.execute( + "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_general_chat_to_unified', 'true')", + [], + ) + .map_err(|e| format!("标记迁移完成失败: {}", e))?; + return Ok(0); + } + + tracing::info!( + "[迁移] 开始迁移 {} 个 general_chat 会话到统一表", + general_count + ); + + // 迁移会话 + let migrated_sessions = migrate_general_sessions(conn)?; + tracing::info!("[迁移] 迁移了 {} 个会话", migrated_sessions); + + // 迁移消息 + let migrated_messages = migrate_general_messages(conn)?; + tracing::info!("[迁移] 迁移了 {} 条消息", migrated_messages); + + // 标记迁移完成 + conn.execute( + "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_general_chat_to_unified', 'true')", + [], + ) + .map_err(|e| format!("标记迁移完成失败: {}", e))?; + + tracing::info!("[迁移] General Chat 数据迁移完成!"); + Ok(migrated_sessions + migrated_messages) +} + +/// 迁移 General Chat 会话数据 +fn migrate_general_sessions(conn: &Connection) -> Result { + let mut stmt = conn + .prepare( + "SELECT id, name, created_at, updated_at, metadata + FROM general_chat_sessions", + ) + .map_err(|e| format!("准备查询语句失败: {}", e))?; + + let sessions: Vec<(String, String, i64, i64, Option)> = stmt + .query_map([], |row| { + Ok(( + row.get(0)?, + row.get(1)?, + row.get(2)?, + row.get(3)?, + row.get(4)?, + )) + }) + .map_err(|e| format!("查询会话失败: {}", e))? + .filter_map(|r| r.ok()) + .collect(); + + let mut count = 0; + for (id, name, created_at, updated_at, _metadata) in sessions { + // 检查是否已存在 + let exists: bool = conn + .query_row( + "SELECT 1 FROM agent_sessions WHERE id = ?", + params![&id], + |_| Ok(true), + ) + .unwrap_or(false); + + if exists { + tracing::warn!("[迁移] 会话 {} 已存在,跳过", id); + continue; + } + + // 转换时间戳格式(general_chat 用毫秒,agent 用 RFC3339) + let created_str = timestamp_ms_to_rfc3339(created_at); + let updated_str = timestamp_ms_to_rfc3339(updated_at); + + // 插入到 agent_sessions,model 字段使用 "general:default" 标识模式 + conn.execute( + "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + id, + "general:default", // 标识为 general 模式 + Option::::None, + name, + created_str, + updated_str, + ], + ) + .map_err(|e| format!("插入会话失败: {}", e))?; + + count += 1; + } + + Ok(count) +} + +/// 迁移 General Chat 消息数据 +fn migrate_general_messages(conn: &Connection) -> Result { + let mut stmt = conn + .prepare( + "SELECT id, session_id, role, content, blocks, status, created_at, metadata + FROM general_chat_messages", + ) + .map_err(|e| format!("准备查询语句失败: {}", e))?; + + #[allow(clippy::type_complexity)] + let messages: Vec<( + String, + String, + String, + String, + Option, + String, + i64, + Option, + )> = stmt + .query_map([], |row| { + Ok(( + row.get(0)?, + row.get(1)?, + row.get(2)?, + row.get(3)?, + row.get(4)?, + row.get(5)?, + row.get(6)?, + row.get(7)?, + )) + }) + .map_err(|e| format!("查询消息失败: {}", e))? + .filter_map(|r| r.ok()) + .collect(); + + let mut count = 0; + for (_id, session_id, role, content, blocks, _status, created_at, _metadata) in messages { + // 检查会话是否存在于 agent_sessions + let session_exists: bool = conn + .query_row( + "SELECT 1 FROM agent_sessions WHERE id = ?", + params![&session_id], + |_| Ok(true), + ) + .unwrap_or(false); + + if !session_exists { + tracing::warn!("[迁移] 消息的会话 {} 不存在,跳过", session_id); + continue; + } + + // 转换内容格式 + let content_json = convert_general_content_to_json(&content, &blocks); + let timestamp_str = timestamp_ms_to_rfc3339(created_at); + + // 插入到 agent_messages + conn.execute( + "INSERT INTO agent_messages (session_id, role, content_json, timestamp, tool_calls_json, tool_call_id) + VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + params![ + session_id, + role, + content_json, + timestamp_str, + Option::::None, + Option::::None, + ], + ) + .map_err(|e| format!("插入消息失败: {}", e))?; + + count += 1; + } + + Ok(count) +} + +/// 将毫秒时间戳转换为 RFC3339 格式 +fn timestamp_ms_to_rfc3339(timestamp_ms: i64) -> String { + use chrono::{TimeZone, Utc}; + + let secs = timestamp_ms / 1000; + let nsecs = ((timestamp_ms % 1000) * 1_000_000) as u32; + + match Utc.timestamp_opt(secs, nsecs) { + chrono::LocalResult::Single(dt) => dt.to_rfc3339(), + _ => Utc::now().to_rfc3339(), + } +} + +/// 将 General Chat 内容转换为 JSON 格式 +fn convert_general_content_to_json(content: &str, blocks: &Option) -> String { + // 如果有 blocks,尝试解析并转换 + if let Some(blocks_str) = blocks { + if let Ok(blocks_arr) = serde_json::from_str::>(blocks_str) { + let converted: Vec = blocks_arr + .into_iter() + .map(|block| { + if let Some(block_type) = block.get("type").and_then(|t| t.as_str()) { + match block_type { + "text" => { + let text = + block.get("content").and_then(|c| c.as_str()).unwrap_or(""); + serde_json::json!({"type": "text", "text": text}) + } + "image" => { + let url = + block.get("content").and_then(|c| c.as_str()).unwrap_or(""); + serde_json::json!({"type": "image", "url": url}) + } + _ => serde_json::json!({"type": "text", "text": content}), + } + } else { + serde_json::json!({"type": "text", "text": content}) + } + }) + .collect(); + + if let Ok(json_str) = serde_json::to_string(&converted) { + return json_str; + } + } + } + + // 默认:纯文本格式 + serde_json::json!([{"type": "text", "text": content}]).to_string() +} + +/// 检查 General Chat 迁移状态 +#[allow(dead_code)] +pub fn check_general_chat_migration_status(conn: &Connection) -> GeneralChatMigrationStatus { + let general_sessions: i64 = conn + .query_row("SELECT COUNT(*) FROM general_chat_sessions", [], |row| { + row.get(0) + }) + .unwrap_or(0); + + let general_messages: i64 = conn + .query_row("SELECT COUNT(*) FROM general_chat_messages", [], |row| { + row.get(0) + }) + .unwrap_or(0); + + let unified_general_sessions: i64 = conn + .query_row( + "SELECT COUNT(*) FROM agent_sessions WHERE model LIKE 'general:%'", + [], + |row| row.get(0), + ) + .unwrap_or(0); + + GeneralChatMigrationStatus { + general_sessions_count: general_sessions as usize, + general_messages_count: general_messages as usize, + migrated_sessions_count: unified_general_sessions as usize, + needs_migration: general_sessions > 0 && unified_general_sessions == 0, + } +} + +/// General Chat 迁移状态 +#[derive(Debug)] +#[allow(dead_code)] +pub struct GeneralChatMigrationStatus { + pub general_sessions_count: usize, + pub general_messages_count: usize, + pub migrated_sessions_count: usize, + pub needs_migration: bool, +} diff --git a/src-tauri/src/services/aster_session_store.rs b/src-tauri/src/services/aster_session_store.rs index 985c842c3..4188747f2 100644 --- a/src-tauri/src/services/aster_session_store.rs +++ b/src-tauri/src/services/aster_session_store.rs @@ -173,6 +173,33 @@ impl SessionStore for ProxyCastSessionStore { .lock() .map_err(|e| anyhow!("数据库锁定失败: {}", e))?; + // 检查会话是否存在,如果不存在则自动创建 + let session_exists: bool = conn + .query_row( + "SELECT 1 FROM agent_sessions WHERE id = ?", + [session_id], + |_| Ok(true), + ) + .unwrap_or(false); + + if !session_exists { + let now = Utc::now().to_rfc3339(); + conn.execute( + "INSERT INTO agent_sessions (id, model, system_prompt, title, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?6)", + rusqlite::params![ + session_id, + "agent:default", + None::, + "新对话", + now, + now + ], + ) + .map_err(|e| anyhow!("自动创建会话失败: {}", e))?; + tracing::info!("[SessionStore] 自动创建会话: {}", session_id); + } + let role = Self::message_role_to_string(message); let content_json = serde_json::to_string(&message.content) .map_err(|e| anyhow!("序列化消息内容失败: {}", e))?; diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 9878388da..620355347 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "ProxyCast", - "version": "0.52.0", + "version": "0.53.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src/components/agent/chat/hooks/useAgentChat.ts b/src/components/agent/chat/hooks/useAgentChat.ts index abd06a648..03128fb9b 100644 --- a/src/components/agent/chat/hooks/useAgentChat.ts +++ b/src/components/agent/chat/hooks/useAgentChat.ts @@ -895,15 +895,19 @@ export function useAgentChat(options: UseAgentChatOptions = {}) { ? images.map((img) => ({ data: img.data, media_type: img.mediaType })) : undefined; + // systemPrompt 已在创建 session 时传递给后端,无需前端注入 + const messageToSend = content; + console.log("[AgentChat] 发送消息:", { - content: content.slice(0, 50), + content: messageToSend.slice(0, 100), sessionId: activeSessionId, model, provider: providerType, + hasSystemPrompt: !!systemPrompt, }); await sendAgentMessageStream( - content, + messageToSend, eventName, activeSessionId, // 传递 sessionId 以保持上下文 model || undefined, @@ -949,9 +953,16 @@ export function useAgentChat(options: UseAgentChatOptions = {}) { const switchTopic = async (topicId: string) => { if (topicId === sessionId) return; + console.log("[useAgentChat] 切换话题:", topicId); + try { // 从后端加载消息历史 const agentMessages = await getAgentSessionMessages(topicId); + console.log("[useAgentChat] 加载到消息数量:", agentMessages.length); + console.log( + "[useAgentChat] 原始消息:", + JSON.stringify(agentMessages.slice(0, 2), null, 2), + ); // 转换为前端 Message 格式 const loadedMessages: Message[] = agentMessages.map((msg, index) => { @@ -969,6 +980,10 @@ export function useAgentChat(options: UseAgentChatOptions = {}) { .join("\n"); } + console.log( + `[useAgentChat] 消息 ${index}: role=${msg.role}, content类型=${typeof msg.content}, 内容长度=${content.length}`, + ); + return { id: `${topicId}-${index}`, role: msg.role as "user" | "assistant", @@ -978,6 +993,7 @@ export function useAgentChat(options: UseAgentChatOptions = {}) { }; }); + console.log("[useAgentChat] 转换后消息数量:", loadedMessages.length); setMessages(loadedMessages); setSessionId(topicId); toast.info("已切换话题"); diff --git a/src/components/agent/chat/index.tsx b/src/components/agent/chat/index.tsx index 315d92d7d..b0e14348e 100644 --- a/src/components/agent/chat/index.tsx +++ b/src/components/agent/chat/index.tsx @@ -423,6 +423,8 @@ export function AgentChatPage({ // 包装 switchTopic,在切换话题时重置相关状态 const switchTopic = useCallback( async (topicId: string) => { + console.log("[AgentChatPage] switchTopic 包装函数被调用:", topicId); + // 先重置本地状态 setLayoutMode("chat"); setCanvasState(null); @@ -435,7 +437,9 @@ export function AgentChatPage({ restoredFilesSessionId.current = null; // 然后调用原始的 switchTopic + console.log("[AgentChatPage] 调用 originalSwitchTopic"); await originalSwitchTopic(topicId); + console.log("[AgentChatPage] originalSwitchTopic 完成"); }, [originalSwitchTopic], ); diff --git a/src/components/content-creator/canvas/canvasUtils.ts b/src/components/content-creator/canvas/canvasUtils.ts index bd68fe504..4cd52fc4a 100644 --- a/src/components/content-creator/canvas/canvasUtils.ts +++ b/src/components/content-creator/canvas/canvasUtils.ts @@ -33,15 +33,20 @@ export type CanvasType = "document" | "poster" | "music" | "script" | "novel"; /** * 主题到画布类型的映射 - * 与 ProjectType 统一后的配置 + * 所有主题都支持 document 画布,特定主题有专用画布 + * + * 设计原则: + * - 所有主题都可以触发画布(当检测到 标签时) + * - 特定主题使用专用画布类型(如 music、poster、script) + * - 通用主题(general、knowledge、planning)使用 document 画布 */ const THEME_TO_CANVAS_TYPE: Record = { - general: null, + general: "document", // 通用对话也支持文档画布 "social-media": "document", poster: "poster", music: "music", - knowledge: null, - planning: null, + knowledge: "document", // 知识探索支持文档画布 + planning: "document", // 计划规划支持文档画布 document: "document", video: "script", novel: "novel", diff --git a/src/components/content-creator/utils/systemPrompt.ts b/src/components/content-creator/utils/systemPrompt.ts index 10a9ef974..4e51c9cbb 100644 --- a/src/components/content-creator/utils/systemPrompt.ts +++ b/src/components/content-creator/utils/systemPrompt.ts @@ -91,10 +91,12 @@ function getFileWritingInstructions(theme?: ThemeType): string { 内容... -**重要**: +**重要规则**: - 这是标签格式,不是工具调用!直接写在回复文本中 - 标签内的内容会实时流式显示在右侧画布 -- 写入完成后,在对话框中简短说明即可 +- **标签前**:先写一句引导语,如"好的,我来帮你写..." +- **标签后**:写完成总结,如"✅ 文案已生成!你可以在右侧画布中查看和编辑。" +- 标签外的文字会显示在左侧对话区,标签内的内容只显示在右侧画布 `; // 根据主题类型返回对应的文件体系 diff --git a/src/components/general-chat/chat/AssistantMessage.tsx b/src/components/general-chat/chat/AssistantMessage.tsx index 99d1b46ad..6d6b8c2bf 100644 --- a/src/components/general-chat/chat/AssistantMessage.tsx +++ b/src/components/general-chat/chat/AssistantMessage.tsx @@ -4,13 +4,16 @@ * @module components/general-chat/chat/AssistantMessage * * @requirements 6.2, 6.7, 2.6, 9.2, 9.3, 9.5 + * + * 支持 标签解析,自动触发画布显示 */ -import React, { useState, useMemo } from "react"; +import React, { useState, useMemo, useEffect, useRef } from "react"; import type { Message, ContentBlock } from "../types"; import { CodeBlock } from "./CodeBlock"; import { ErrorDisplay } from "./ErrorDisplay"; import { ImageMessage } from "./ImageMessage"; +import { WriteFileParser } from "@/lib/writeFile"; interface AssistantMessageProps { /** 消息数据 */ @@ -29,21 +32,29 @@ interface AssistantMessageProps { onRetry?: () => void; /** 是否正在重试 */ isRetrying?: boolean; + /** 画布内容更新回调(用于 write_file 标签) */ + onCanvasUpdate?: (path: string, content: string, isComplete: boolean) => void; } /** * 解析 Markdown 内容,提取代码块和图片 + * 同时处理 标签,将其从显示内容中移除 */ const parseContent = (content: string): ContentBlock[] => { const blocks: ContentBlock[] = []; + + // 先移除 标签及其内容(这些内容会显示在画布中) + const writeFileResult = WriteFileParser.parse(content); + const cleanContent = writeFileResult.plainText; + const codeBlockRegex = /```(\w+)?\n([\s\S]*?)```/g; let lastIndex = 0; let match; - while ((match = codeBlockRegex.exec(content)) !== null) { + while ((match = codeBlockRegex.exec(cleanContent)) !== null) { // 添加代码块之前的文本 if (match.index > lastIndex) { - const text = content.slice(lastIndex, match.index).trim(); + const text = cleanContent.slice(lastIndex, match.index).trim(); if (text) { blocks.push({ type: "text", content: text }); } @@ -60,16 +71,16 @@ const parseContent = (content: string): ContentBlock[] => { } // 添加剩余的文本 - if (lastIndex < content.length) { - const text = content.slice(lastIndex).trim(); + if (lastIndex < cleanContent.length) { + const text = cleanContent.slice(lastIndex).trim(); if (text) { blocks.push({ type: "text", content: text }); } } // 如果没有解析出任何块,返回整个内容作为文本 - if (blocks.length === 0) { - blocks.push({ type: "text", content }); + if (blocks.length === 0 && cleanContent.trim()) { + blocks.push({ type: "text", content: cleanContent }); } return blocks; @@ -87,9 +98,11 @@ export const AssistantMessage: React.FC = ({ onRegenerate, onRetry, isRetrying = false, + onCanvasUpdate, }) => { const [showActions, setShowActions] = useState(false); const [copied, setCopied] = useState(false); + const lastCanvasUpdateRef = useRef(""); // 判断是否为错误状态 const isError = message.status === "error" && message.error; @@ -98,6 +111,29 @@ export const AssistantMessage: React.FC = ({ const displayContent = isStreaming && streamingContent ? streamingContent : message.content; + // 流式解析 标签,实时更新画布 + useEffect(() => { + if (!onCanvasUpdate || !displayContent) return; + + const result = WriteFileParser.parse(displayContent); + + // 如果有 write_file 块,更新画布 + if (result.blocks.length > 0) { + const firstBlock = result.blocks[0]; + const updateKey = `${firstBlock.path}:${firstBlock.content.length}`; + + // 避免重复更新 + if (updateKey !== lastCanvasUpdateRef.current) { + lastCanvasUpdateRef.current = updateKey; + onCanvasUpdate( + firstBlock.path, + firstBlock.content, + firstBlock.isComplete, + ); + } + } + }, [displayContent, onCanvasUpdate]); + // 解析内容块(仅用于文本内容的 Markdown 解析) const parsedContentBlocks = useMemo( () => parseContent(displayContent), diff --git a/src/components/general-chat/chat/ChatPanel.tsx b/src/components/general-chat/chat/ChatPanel.tsx index 1371345dc..6d050e5de 100644 --- a/src/components/general-chat/chat/ChatPanel.tsx +++ b/src/components/general-chat/chat/ChatPanel.tsx @@ -126,6 +126,9 @@ export const ChatPanel: React.FC = ({ const getWorkflowManager = useGeneralChatStore( (state) => state.getWorkflowManager, ); + const streamCanvasContent = useGeneralChatStore( + (state) => state.streamCanvasContent, + ); // 获取分页状态 const paginationState = sessionId @@ -250,6 +253,14 @@ export const ChatPanel: React.FC = ({ [onOpenCanvas], ); + // 处理流式画布更新(用于 write_file 标签) + const handleCanvasUpdate = useCallback( + (path: string, content: string, isComplete: boolean) => { + streamCanvasContent(path, content, isComplete); + }, + [streamCanvasContent], + ); + // 处理重试消息 const handleRetry = useCallback( async (messageId: string) => { @@ -317,6 +328,7 @@ export const ChatPanel: React.FC = ({ hasMoreMessages={hasMoreMessages} isLoadingMore={isLoadingMore} onLoadMore={handleLoadMore} + onCanvasUpdate={handleCanvasUpdate} /> )} diff --git a/src/components/general-chat/chat/MessageItem.tsx b/src/components/general-chat/chat/MessageItem.tsx index 09df39093..2a34d4726 100644 --- a/src/components/general-chat/chat/MessageItem.tsx +++ b/src/components/general-chat/chat/MessageItem.tsx @@ -29,6 +29,8 @@ interface MessageItemProps { onRetry?: () => void; /** 是否正在重试 */ isRetrying?: boolean; + /** 画布内容更新回调(用于 write_file 标签) */ + onCanvasUpdate?: (path: string, content: string, isComplete: boolean) => void; } /** @@ -43,6 +45,7 @@ export const MessageItem: React.FC = ({ onRegenerate, onRetry, isRetrying, + onCanvasUpdate, }) => { if (message.role === "user") { return ; @@ -59,6 +62,7 @@ export const MessageItem: React.FC = ({ onRegenerate={onRegenerate} onRetry={onRetry} isRetrying={isRetrying} + onCanvasUpdate={onCanvasUpdate} /> ); } diff --git a/src/components/general-chat/chat/MessageList.tsx b/src/components/general-chat/chat/MessageList.tsx index e7692415f..39930b32f 100644 --- a/src/components/general-chat/chat/MessageList.tsx +++ b/src/components/general-chat/chat/MessageList.tsx @@ -58,6 +58,8 @@ interface MessageListProps { isLoadingMore?: boolean; /** 加载更多消息回调 */ onLoadMore?: () => void; + /** 画布内容更新回调(用于 write_file 标签) */ + onCanvasUpdate?: (path: string, content: string, isComplete: boolean) => void; } /** @@ -81,6 +83,7 @@ export const MessageList: React.FC = ({ hasMoreMessages = true, isLoadingMore = false, onLoadMore, + onCanvasUpdate, }) => { const parentRef = useRef(null); const isAtBottomRef = useRef(true); @@ -231,6 +234,7 @@ export const MessageList: React.FC = ({ } onRetry={onRetry ? () => onRetry(message.id) : undefined} isRetrying={isRetrying} + onCanvasUpdate={isStreamingMessage ? onCanvasUpdate : undefined} /> ); }, @@ -243,6 +247,7 @@ export const MessageList: React.FC = ({ onRegenerate, onRetry, retryingMessageId, + onCanvasUpdate, ], ); diff --git a/src/components/general-chat/store/useGeneralChatStore.ts b/src/components/general-chat/store/useGeneralChatStore.ts index b7d7af626..27f70625a 100644 --- a/src/components/general-chat/store/useGeneralChatStore.ts +++ b/src/components/general-chat/store/useGeneralChatStore.ts @@ -32,6 +32,61 @@ import { parseApiError, } from "../types"; +// ============================================================================ +// 内容创作系统指令生成 +// ============================================================================ + +/** + * 主题名称映射 + */ +const THEME_NAMES: Record = { + general: "通用对话", + "social-media": "社媒内容", + poster: "图文海报", + music: "歌词曲谱", + knowledge: "知识探索", + planning: "计划规划", + document: "办公文档", + video: "短视频", + novel: "小说创作", +}; + +/** + * 生成内容创作系统指令 + * 告诉 AI 使用 标签输出内容 + */ +function getContentCreationInstruction(theme: string, _mode: string): string { + const themeName = THEME_NAMES[theme] || "内容创作"; + + return `【系统指令 - ${themeName}助手】 + +你是一位专业的${themeName}助手。请遵循以下输出格式: + +## 输出格式要求 + +当需要输出文档内容时,使用 标签: + + +内容... + + +**重要规则**: +1. 标签前:先写一句引导语,如"好的,我来帮你写..." +2. 标签内:放置完整的文档内容 +3. 标签后:写完成总结,如"✅ 文案已生成!" + +**示例**: +好的,我来帮你写一篇小红书探店文案。 + + +# 标题 + +正文内容... + + +✅ 文案已生成!你可以在右侧画布中查看和编辑。`; +} + // ============================================================================ // Store 状态接口 // ============================================================================ @@ -165,6 +220,12 @@ export interface GeneralChatState { updateCanvasContent: (content: string) => void; /** 设置画布编辑模式 */ setCanvasEditing: (isEditing: boolean) => void; + /** 流式更新画布内容(用于 write_file 标签) */ + streamCanvasContent: ( + path: string, + content: string, + isComplete: boolean, + ) => void; // ========== Provider 选择操作 ========== /** 设置选中的 Provider */ @@ -181,6 +242,27 @@ export interface GeneralChatState { // ========== 重置操作 ========== /** 重置 Store 到初始状态 */ reset: () => void; + + // ========== 内容创作状态 ========== + /** 当前主题类型 */ + contentTheme: + | "general" + | "social-media" + | "poster" + | "music" + | "knowledge" + | "planning" + | "document" + | "video" + | "novel"; + /** 当前创作模式 */ + contentCreationMode: "guided" | "fast" | "hybrid" | "framework"; + /** 设置内容创作主题 */ + setContentTheme: (theme: GeneralChatState["contentTheme"]) => void; + /** 设置内容创作模式 */ + setContentCreationMode: ( + mode: GeneralChatState["contentCreationMode"], + ) => void; } // ============================================================================ @@ -214,6 +296,10 @@ const initialState = { workflowManagers: {} as Record, workflowEnabled: false, workflowThreshold: 5, + + // 内容创作状态 + contentTheme: "general" as GeneralChatState["contentTheme"], + contentCreationMode: "guided" as GeneralChatState["contentCreationMode"], }; // ============================================================================ @@ -498,9 +584,25 @@ export const useGeneralChatStore = create()( try { // 调用 Tauri 命令发送消息并开始流式响应 const { invoke } = await import("@tauri-apps/api/core"); + + // 获取内容创作状态 + const { contentTheme, contentCreationMode } = get(); + + // 根据主题生成系统指令前缀 + let messageToSend = content.trim() || "请分析这张图片"; + + // 如果不是通用主题,注入系统指令 + if (contentTheme !== "general") { + const systemInstruction = getContentCreationInstruction( + contentTheme, + contentCreationMode, + ); + messageToSend = `${systemInstruction}\n\n---\n\n用户请求:${messageToSend}`; + } + await invoke("aster_agent_chat_stream", { sessionId: currentSessionId, - message: content.trim() || "请分析这张图片", + message: messageToSend, eventName: `general-chat-stream-${currentSessionId}`, images: imageData, }); @@ -1092,6 +1194,40 @@ export const useGeneralChatStore = create()( })); }, + streamCanvasContent: ( + path: string, + content: string, + _isComplete: boolean, + ) => { + const { canvas } = get(); + + // 如果画布未打开,自动打开并设置初始状态 + if (!canvas.isOpen) { + set((state) => ({ + canvas: { + isOpen: true, + contentType: "markdown", + content, + filename: path, + isEditing: false, + }, + ui: { + ...state.ui, + canvasCollapsed: false, + }, + })); + } else { + // 画布已打开,只更新内容 + set((state) => ({ + canvas: { + ...state.canvas, + content, + filename: path, + }, + })); + } + }, + // ========== Provider 选择操作实现 ========== setSelectedProvider: (providerKey: string | null) => { @@ -1259,6 +1395,16 @@ export const useGeneralChatStore = create()( const messageCount = (messages[sessionId] || []).length; return messageCount >= workflowThreshold; }, + + // ========== 内容创作操作实现 ========== + + setContentTheme: (theme) => { + set({ contentTheme: theme }); + }, + + setContentCreationMode: (mode) => { + set({ contentCreationMode: mode }); + }, }), { name: "general-chat-storage", diff --git a/src/hooks/README.md b/src/hooks/README.md index 0584de66c..4b652931e 100644 --- a/src/hooks/README.md +++ b/src/hooks/README.md @@ -1,36 +1,59 @@ -# hooks +# Hooks 目录 - - -## 架构说明 - -React 自定义 Hooks,封装业务逻辑和状态管理。 -通过 Tauri invoke 与 Rust 后端通信。 +全局共享的 React Hooks。 ## 文件索引 -- `index.ts` - Hooks 导出入口 -- `useApiKeyProvider.ts` - API Key Provider 管理 Hook(Requirements 9.1) -- `useConnectCallback.ts` - Connect 统计回调 Hook(Requirements 5.3) -- `useDeepLink.ts` - Deep Link 事件处理 Hook(Requirements 5.1, 5.2, 5.3, 5.4) -- `useErrorHandler.ts` - 错误处理 Hook -- `useFileMonitoring.ts` - 文件监控 Hook -- `useFlowActions.ts` - 流量操作 Hook -- `useFlowEvents.ts` - 流量事件 Hook -- `useFlowNotifications.ts` - 流量通知 Hook -- `useMcpServers.ts` - MCP 服务器管理 Hook -- `useOAuthCredentials.ts` - OAuth 凭证管理 Hook -- `usePrompts.ts` - Prompt 管理 Hook -- `useProviderPool.ts` - Provider 池管理 Hook -- `useProviderState.ts` - Provider 状态 Hook -- `useRelayRegistry.ts` - Relay Registry 管理 Hook(Requirements 2.1, 7.2, 7.3) -- `useSkills.ts` - 技能管理 Hook -- `useSound.ts` - 音效管理 Hook(工具调用和打字机音效) -- `useSwitch.ts` - 开关状态 Hook -- `useTauri.ts` - Tauri 通用 Hook -- `useWindowResize.ts` - 窗口大小 Hook -- `useWorkspace.ts` - Workspace 工作目录管理 Hook +| 文件 | 说明 | +|------|------| +| `useUnifiedChat.ts` | 统一对话 Hook,支持 Agent/General/Creator 三种模式 | -## 更新提醒 +## useUnifiedChat -任何文件变更后,请更新此文档和相关的上级文档。 +统一的对话逻辑 Hook,替代原有分散的 `useAgentChat` 和 `useChat`。 + +### 使用示例 + +```typescript +import { useUnifiedChat } from "@/hooks/useUnifiedChat"; + +// Agent 模式 +const agentChat = useUnifiedChat({ + mode: "agent", + providerType: "claude", + model: "claude-sonnet-4-20250514", +}); + +// 内容创作模式 +const creatorChat = useUnifiedChat({ + mode: "creator", + systemPrompt: "你是一位专业的内容创作助手...", + onCanvasUpdate: (path, content) => { + // 更新画布内容 + }, +}); + +// 通用对话模式 +const generalChat = useUnifiedChat({ + mode: "general", +}); +``` + +### 返回值 + +- `session` - 当前会话 +- `messages` - 消息列表 +- `isLoading` - 加载状态 +- `isSending` - 发送状态 +- `error` - 错误信息 +- `createSession()` - 创建会话 +- `loadSession()` - 加载会话 +- `sendMessage()` - 发送消息 +- `stopGeneration()` - 停止生成 +- `configureProvider()` - 配置 Provider + +## 相关文档 + +- 架构设计:`docs/prd/chat-architecture-redesign.md` +- 类型定义:`src/types/chat.ts` +- API 封装:`src/lib/api/unified-chat.ts` diff --git a/src/hooks/useUnifiedChat.ts b/src/hooks/useUnifiedChat.ts new file mode 100644 index 000000000..e5a38bcf5 --- /dev/null +++ b/src/hooks/useUnifiedChat.ts @@ -0,0 +1,657 @@ +/** + * @file useUnifiedChat.ts + * @description 统一对话 Hook + * @module hooks/useUnifiedChat + * + * 提供统一的对话逻辑,支持多种对话模式: + * - Agent: AI Agent 模式,支持工具调用 + * - General: 通用对话模式,纯文本 + * - Creator: 内容创作模式,支持画布输出 + * + * ## 设计原则 + * - 单一入口:所有对话场景使用同一个 Hook + * - 模式化设计:通过 mode 参数区分不同场景 + * - 统一 API:调用后端统一的 unified_chat_cmd + */ + +import { useState, useEffect, useRef, useCallback } from "react"; +import { toast } from "sonner"; +import { safeListen } from "@/lib/dev-bridge"; +import type { UnlistenFn } from "@tauri-apps/api/event"; +import * as chatApi from "@/lib/api/unified-chat"; +import type { + ChatSession, + ChatMessage, + ChatError, + ImageInput, + StreamEvent, + UseUnifiedChatOptions, + UseUnifiedChatReturn, + ToolCall, + CreateSessionRequest, +} from "@/types/chat"; + +// ============================================================================ +// 常量 +// ============================================================================ + +const STORAGE_PREFIX = "unified_chat_"; + +// ============================================================================ +// 辅助函数 +// ============================================================================ + +/** 从 localStorage 加载数据 */ +function loadFromStorage(key: string, defaultValue: T): T { + try { + const stored = localStorage.getItem(`${STORAGE_PREFIX}${key}`); + return stored ? JSON.parse(stored) : defaultValue; + } catch { + return defaultValue; + } +} + +/** 保存数据到 localStorage */ +function saveToStorage(key: string, value: unknown): void { + try { + localStorage.setItem(`${STORAGE_PREFIX}${key}`, JSON.stringify(value)); + } catch (e) { + console.error("[useUnifiedChat] 保存到 localStorage 失败:", e); + } +} + +/** 解析 API 错误 */ +function parseApiError(error: unknown): ChatError { + const message = error instanceof Error ? error.message : String(error); + + // 根据错误消息判断类型 + if (message.includes("network") || message.includes("连接")) { + return { type: "network", message, retryable: true }; + } + if ( + message.includes("auth") || + message.includes("认证") || + message.includes("401") + ) { + return { type: "auth", message, retryable: false }; + } + if ( + message.includes("rate") || + message.includes("限流") || + message.includes("429") + ) { + return { type: "rate_limit", message, retryable: true }; + } + if (message.includes("quota") || message.includes("配额")) { + return { type: "quota", message, retryable: false }; + } + + return { type: "unknown", message, retryable: true }; +} + +// ============================================================================ +// Hook 实现 +// ============================================================================ + +/** + * 统一对话 Hook + */ +export function useUnifiedChat( + options: UseUnifiedChatOptions, +): UseUnifiedChatReturn { + const { + mode, + sessionId: initialSessionId, + systemPrompt, + providerType: initialProviderType, + model: initialModel, + onCanvasUpdate, + onWriteFile, + onError, + } = options; + + // ========== 状态 ========== + const [session, setSession] = useState(null); + const [messages, setMessages] = useState([]); + const [isLoading, setIsLoading] = useState(false); + const [isSending, setIsSending] = useState(false); + const [error, setError] = useState(null); + + // Provider 配置 + const [providerType, setProviderType] = useState( + () => initialProviderType || loadFromStorage(`${mode}_provider`, "claude"), + ); + const [model, setModel] = useState( + () => initialModel || loadFromStorage(`${mode}_model`, ""), + ); + + // Refs + const unlistenRef = useRef(null); + const currentMsgIdRef = useRef(null); + const accumulatedContentRef = useRef(""); + + // ========== 会话操作 ========== + + /** 创建新会话 */ + const createSession = useCallback( + async (opts?: Partial): Promise => { + try { + setIsLoading(true); + setError(null); + + const response = await chatApi.createSession({ + mode, + title: opts?.title, + systemPrompt: opts?.systemPrompt || systemPrompt, + providerType: opts?.providerType || providerType, + model: opts?.model || model, + metadata: opts?.metadata, + }); + + const newSession: ChatSession = { + id: response.id, + mode: response.mode, + title: response.title, + model: response.model, + createdAt: response.createdAt, + updatedAt: response.updatedAt, + messageCount: response.messageCount, + }; + + setSession(newSession); + setMessages([]); + saveToStorage(`${mode}_session_id`, response.id); + + console.log("[useUnifiedChat] 创建会话成功:", response.id); + return response.id; + } catch (e) { + const chatError = parseApiError(e); + setError(chatError); + onError?.(chatError); + throw e; + } finally { + setIsLoading(false); + } + }, + [mode, systemPrompt, providerType, model, onError], + ); + + /** 加载会话 */ + const loadSession = useCallback( + async (sessionId: string): Promise => { + try { + setIsLoading(true); + setError(null); + + // 获取会话详情 + const response = await chatApi.getSession(sessionId); + const loadedSession: ChatSession = { + id: response.id, + mode: response.mode, + title: response.title, + model: response.model, + createdAt: response.createdAt, + updatedAt: response.updatedAt, + messageCount: response.messageCount, + }; + setSession(loadedSession); + + // 获取消息列表 + const loadedMessages = await chatApi.getMessages(sessionId); + setMessages(loadedMessages); + + saveToStorage(`${mode}_session_id`, sessionId); + console.log( + "[useUnifiedChat] 加载会话成功:", + sessionId, + "消息数:", + loadedMessages.length, + ); + } catch (e) { + const chatError = parseApiError(e); + setError(chatError); + onError?.(chatError); + throw e; + } finally { + setIsLoading(false); + } + }, + [mode, onError], + ); + + /** 删除会话 */ + const deleteSession = useCallback( + async (sessionId?: string): Promise => { + const targetId = sessionId || session?.id; + if (!targetId) return; + + try { + await chatApi.deleteSession(targetId); + + if (targetId === session?.id) { + setSession(null); + setMessages([]); + localStorage.removeItem(`${STORAGE_PREFIX}${mode}_session_id`); + } + + toast.success("会话已删除"); + } catch (e) { + const chatError = parseApiError(e); + setError(chatError); + onError?.(chatError); + toast.error("删除会话失败"); + } + }, + [session?.id, mode, onError], + ); + + /** 重命名会话 */ + const renameSession = useCallback( + async (title: string, sessionId?: string): Promise => { + const targetId = sessionId || session?.id; + if (!targetId) return; + + try { + await chatApi.renameSession(targetId, title); + + if (targetId === session?.id) { + setSession((prev) => (prev ? { ...prev, title } : null)); + } + } catch (e) { + const chatError = parseApiError(e); + setError(chatError); + onError?.(chatError); + toast.error("重命名失败"); + } + }, + [session?.id, onError], + ); + + // ========== 消息操作 ========== + + /** 发送消息 */ + const sendMessage = useCallback( + async (content: string, images?: ImageInput[]): Promise => { + if (!content.trim() && (!images || images.length === 0)) return; + + let activeSessionId = session?.id; + + // 如果没有会话,先创建一个 + if (!activeSessionId) { + try { + activeSessionId = await createSession(); + } catch { + return; + } + } + + // 创建用户消息 + const userMsgId = `user-${Date.now()}`; + const userMessage: ChatMessage = { + id: userMsgId, + sessionId: activeSessionId, + role: "user", + content: content.trim(), + contentBlocks: [{ type: "text", text: content.trim() }], + status: "complete", + createdAt: new Date().toISOString(), + }; + + // 创建助手消息占位符 + const assistantMsgId = `assistant-${Date.now()}`; + const assistantMessage: ChatMessage = { + id: assistantMsgId, + sessionId: activeSessionId, + role: "assistant", + content: "", + contentBlocks: [], + status: "streaming", + createdAt: new Date().toISOString(), + }; + + setMessages((prev) => [...prev, userMessage, assistantMessage]); + setIsSending(true); + currentMsgIdRef.current = assistantMsgId; + accumulatedContentRef.current = ""; + + // 设置事件监听 + const eventName = chatApi.generateEventName(activeSessionId); + + try { + const unlisten = await safeListen(eventName, (event) => { + const data = chatApi.parseStreamEvent(event.payload); + if (!data) return; + + handleStreamEvent(data, assistantMsgId); + }); + + unlistenRef.current = unlisten; + + // 构建发送的消息内容 + let messageToSend = content.trim(); + + // 如果有 systemPrompt 且是第一条消息,注入到消息前面 + const isFirstMessage = + messages.filter((m) => m.role === "user").length === 0; + if (systemPrompt && isFirstMessage) { + messageToSend = `${systemPrompt}\n\n---\n\n用户请求:${messageToSend}`; + } + + // 发送消息 + await chatApi.sendMessage({ + sessionId: activeSessionId, + message: messageToSend, + eventName, + images: images?.map((img) => ({ + data: img.data, + media_type: img.mediaType, + })), + }); + } catch (e) { + console.error("[useUnifiedChat] 发送消息失败:", e); + const chatError = parseApiError(e); + + setMessages((prev) => + prev.map((msg) => + msg.id === assistantMsgId + ? { + ...msg, + status: "error", + error: chatError, + content: chatError.message, + } + : msg, + ), + ); + + setIsSending(false); + onError?.(chatError); + + if (unlistenRef.current) { + unlistenRef.current(); + unlistenRef.current = null; + } + } + }, + // eslint-disable-next-line react-hooks/exhaustive-deps + [session?.id, messages, systemPrompt, createSession, onError], + ); + + /** 处理流式事件 */ + const handleStreamEvent = useCallback( + (event: StreamEvent, msgId: string): void => { + switch (event.type) { + case "text_delta": + accumulatedContentRef.current += event.text; + setMessages((prev) => + prev.map((msg) => + msg.id === msgId + ? { ...msg, content: accumulatedContentRef.current } + : msg, + ), + ); + + // 检查是否有 write_file 标签(用于画布) + checkWriteFileTag(accumulatedContentRef.current); + break; + + case "thinking_delta": + // 处理思考内容(可选显示) + break; + + case "tool_start": { + const newToolCall: ToolCall = { + id: event.tool_id, + name: event.tool_name, + arguments: event.arguments, + status: "running", + startTime: new Date(), + }; + + setMessages((prev) => + prev.map((msg) => + msg.id === msgId + ? { ...msg, toolCalls: [...(msg.toolCalls || []), newToolCall] } + : msg, + ), + ); + + // 检查是否是文件写入工具 + const toolName = event.tool_name.toLowerCase(); + if (toolName.includes("write") || toolName.includes("create")) { + try { + const args = JSON.parse(event.arguments || "{}"); + const filePath = args.path || args.file_path || args.filePath; + const content = args.content || args.text || ""; + if (filePath && content && onWriteFile) { + onWriteFile(content, filePath); + } + } catch { + // 忽略解析错误 + } + } + break; + } + + case "tool_end": + setMessages((prev) => + prev.map((msg) => { + if (msg.id !== msgId) return msg; + return { + ...msg, + toolCalls: msg.toolCalls?.map((tc) => + tc.id === event.tool_id + ? { + ...tc, + status: event.result.success ? "completed" : "failed", + result: event.result, + endTime: new Date(), + } + : tc, + ), + }; + }), + ); + break; + + case "done": + // 单次 API 响应完成,但工具循环可能继续 + break; + + case "final_done": + setMessages((prev) => + prev.map((msg) => + msg.id === msgId + ? { + ...msg, + status: "complete", + content: accumulatedContentRef.current || "(无响应)", + metadata: event.usage + ? { + tokens: { + input: event.usage.input_tokens, + output: event.usage.output_tokens, + }, + } + : undefined, + } + : msg, + ), + ); + setIsSending(false); + cleanup(); + break; + + case "error": { + const chatError: ChatError = { + type: "unknown", + message: event.message, + retryable: true, + }; + + setMessages((prev) => + prev.map((msg) => + msg.id === msgId + ? { + ...msg, + status: "error", + error: chatError, + content: accumulatedContentRef.current || event.message, + } + : msg, + ), + ); + setIsSending(false); + onError?.(chatError); + cleanup(); + break; + } + } + }, + // eslint-disable-next-line react-hooks/exhaustive-deps + [onWriteFile, onError], + ); + + /** 检查 write_file 标签 */ + const checkWriteFileTag = useCallback( + (content: string): void => { + // 匹配 content + const regex = /([\s\S]*?)(<\/write_file>)?/g; + let match; + + while ((match = regex.exec(content)) !== null) { + const [, path, fileContent, closeTag] = match; + const isComplete = !!closeTag; + + if (onCanvasUpdate) { + onCanvasUpdate(path, fileContent); + } + + if (isComplete && onWriteFile) { + onWriteFile(fileContent, path); + } + } + }, + [onCanvasUpdate, onWriteFile], + ); + + /** 清理资源 */ + const cleanup = useCallback((): void => { + if (unlistenRef.current) { + unlistenRef.current(); + unlistenRef.current = null; + } + currentMsgIdRef.current = null; + accumulatedContentRef.current = ""; + }, []); + + /** 停止生成 */ + const stopGeneration = useCallback(async (): Promise => { + if (!session?.id) return; + + try { + await chatApi.stopGeneration(session.id); + + // 更新当前消息状态 + if (currentMsgIdRef.current) { + setMessages((prev) => + prev.map((msg) => + msg.id === currentMsgIdRef.current + ? { + ...msg, + status: "complete", + content: accumulatedContentRef.current || "(已停止)", + } + : msg, + ), + ); + } + + setIsSending(false); + cleanup(); + } catch (e) { + console.error("[useUnifiedChat] 停止生成失败:", e); + } + }, [session?.id, cleanup]); + + /** 清空消息 */ + const clearMessages = useCallback((): void => { + setMessages([]); + setSession(null); + localStorage.removeItem(`${STORAGE_PREFIX}${mode}_session_id`); + toast.success("新对话已创建"); + }, [mode]); + + // ========== Provider 配置 ========== + + /** 配置 Provider */ + const configureProvider = useCallback( + async (newProviderType: string, newModel: string): Promise => { + setProviderType(newProviderType); + setModel(newModel); + saveToStorage(`${mode}_provider`, newProviderType); + saveToStorage(`${mode}_model`, newModel); + + // 如果有活跃会话,更新其 Provider 配置 + if (session?.id) { + try { + await chatApi.configureProvider( + session.id, + newProviderType, + newModel, + ); + } catch (e) { + console.error("[useUnifiedChat] 配置 Provider 失败:", e); + } + } + }, + [mode, session?.id], + ); + + // ========== 初始化 ========== + + useEffect(() => { + // 尝试恢复上次的会话 + const savedSessionId = + initialSessionId || loadFromStorage(`${mode}_session_id`, null); + if (savedSessionId) { + loadSession(savedSessionId).catch(() => { + // 如果加载失败,清除保存的 ID + localStorage.removeItem(`${STORAGE_PREFIX}${mode}_session_id`); + }); + } + + // 清理函数 + return () => { + cleanup(); + }; + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [mode, initialSessionId]); + + // ========== 返回值 ========== + + return { + // 状态 + session, + messages, + isLoading, + isSending, + error, + + // 会话操作 + createSession, + loadSession, + deleteSession, + renameSession, + + // 消息操作 + sendMessage, + stopGeneration, + clearMessages, + + // Provider 配置 + configureProvider, + }; +} + +export default useUnifiedChat; diff --git a/src/lib/README.md b/src/lib/README.md index 5084c72eb..a5e442a14 100644 --- a/src/lib/README.md +++ b/src/lib/README.md @@ -11,6 +11,9 @@ - `artifact/` - Artifact 系统核心库(Requirements 1.1-1.5) - `types.ts` - Artifact 类型定义 +- `writeFile/` - WriteFile 标签解析器 + - `parser.ts` - `` 标签流式解析器 + - `index.ts` - 模块导出 - `api/` - API 调用封装 - `apiKeyProvider.ts` - API Key Provider API 封装(Requirements 9.1) - `pluginUI.ts` - 插件 UI API(Requirements 3.1) diff --git a/src/lib/api/unified-chat.ts b/src/lib/api/unified-chat.ts new file mode 100644 index 000000000..9280f3a33 --- /dev/null +++ b/src/lib/api/unified-chat.ts @@ -0,0 +1,315 @@ +/** + * @file unified-chat.ts + * @description 统一对话 API 封装 + * @module lib/api/unified-chat + * + * 封装所有统一对话相关的 Tauri 命令调用 + */ + +import { invoke } from "@tauri-apps/api/core"; +import type { + ChatMode, + ChatMessage, + SessionResponse, + CreateSessionRequest, + SendMessageRequest, + StreamEvent, + ToolCall, + ToolEndEvent, + FinalDoneEvent, +} from "@/types/chat"; + +// ============================================================================ +// 会话管理 API +// ============================================================================ + +/** + * 创建新会话 + */ +export async function createSession( + request: CreateSessionRequest, +): Promise { + return invoke("chat_create_session", { request }); +} + +/** + * 获取会话列表 + */ +export async function listSessions( + mode?: ChatMode, +): Promise { + return invoke("chat_list_sessions", { mode }); +} + +/** + * 获取会话详情 + */ +export async function getSession(sessionId: string): Promise { + return invoke("chat_get_session", { sessionId }); +} + +/** + * 删除会话 + */ +export async function deleteSession(sessionId: string): Promise { + return invoke("chat_delete_session", { sessionId }); +} + +/** + * 重命名会话 + */ +export async function renameSession( + sessionId: string, + title: string, +): Promise { + return invoke("chat_rename_session", { sessionId, title }); +} + +// ============================================================================ +// 消息管理 API +// ============================================================================ + +/** + * 获取会话消息列表 + */ +export async function getMessages( + sessionId: string, + limit?: number, +): Promise { + const messages = await invoke< + Array<{ + id: number; + session_id: string; + role: string; + content: unknown; + tool_calls?: unknown; + tool_call_id?: string; + metadata?: unknown; + created_at: string; + }> + >("chat_get_messages", { sessionId, limit }); + + // 转换后端格式为前端格式 + return messages.map(convertBackendMessage); +} + +/** + * 发送消息(流式) + */ +export async function sendMessage(request: SendMessageRequest): Promise { + return invoke("chat_send_message", { request }); +} + +/** + * 停止生成 + */ +export async function stopGeneration(sessionId: string): Promise { + return invoke("chat_stop_generation", { sessionId }); +} + +/** + * 配置会话的 Provider + */ +export async function configureProvider( + sessionId: string, + providerType: string, + model: string, +): Promise { + return invoke("chat_configure_provider", { + sessionId, + providerType, + model, + }); +} + +// ============================================================================ +// 辅助函数 +// ============================================================================ + +/** + * 转换后端消息格式为前端格式 + */ +function convertBackendMessage(msg: { + id: number; + session_id: string; + role: string; + content: unknown; + tool_calls?: unknown; + tool_call_id?: string; + metadata?: unknown; + created_at: string; +}): ChatMessage { + // 提取文本内容 + let textContent = ""; + if (typeof msg.content === "string") { + textContent = msg.content; + } else if (Array.isArray(msg.content)) { + textContent = msg.content + .filter( + (part): part is { type: "text"; text: string } => + typeof part === "object" && part !== null && part.type === "text", + ) + .map((part) => part.text) + .join("\n"); + } else if (typeof msg.content === "object" && msg.content !== null) { + // 尝试从对象中提取文本 + const contentObj = msg.content as Record; + if (typeof contentObj.text === "string") { + textContent = contentObj.text; + } + } + + return { + id: msg.id, + sessionId: msg.session_id, + role: msg.role as ChatMessage["role"], + content: textContent, + contentBlocks: Array.isArray(msg.content) + ? msg.content.map(convertContentBlock) + : [{ type: "text" as const, text: textContent }], + toolCalls: msg.tool_calls ? convertToolCalls(msg.tool_calls) : undefined, + toolCallId: msg.tool_call_id, + status: "complete", + metadata: msg.metadata as ChatMessage["metadata"], + createdAt: msg.created_at, + }; +} + +/** + * 转换内容块 + */ +function convertContentBlock( + block: unknown, +): NonNullable[number] { + if (typeof block !== "object" || block === null) { + return { type: "text", text: String(block) }; + } + + const b = block as Record; + + switch (b.type) { + case "text": + return { type: "text", text: String(b.text || "") }; + case "image": + return { type: "image", url: String(b.url || ""), alt: b.alt as string }; + case "file": + return { + type: "file", + path: String(b.path || ""), + name: String(b.name || ""), + }; + case "canvas": + return { + type: "canvas", + canvasType: String(b.canvasType || ""), + content: String(b.content || ""), + }; + default: + return { type: "text", text: JSON.stringify(block) }; + } +} + +/** + * 转换工具调用 + */ +function convertToolCalls(toolCalls: unknown): ChatMessage["toolCalls"] { + if (!Array.isArray(toolCalls)) return undefined; + + return toolCalls.map((tc) => { + const call = tc as Record; + return { + id: String(call.id || ""), + name: String(call.name || ""), + arguments: call.arguments as string, + status: (call.status as ToolCall["status"]) || "completed", + result: call.result as ToolCall["result"], + }; + }); +} + +/** + * 解析流式事件 + */ +export function parseStreamEvent(payload: unknown): StreamEvent | null { + if (typeof payload !== "object" || payload === null) { + return null; + } + + const event = payload as Record; + const type = event.type as string; + + switch (type) { + case "TextDelta": + case "text_delta": + return { + type: "text_delta", + text: String(event.text || event.content || ""), + }; + + case "ThinkingDelta": + case "thinking_delta": + return { + type: "thinking_delta", + text: String(event.text || event.content || ""), + }; + + case "ToolStart": + case "tool_start": + return { + type: "tool_start", + tool_id: String(event.tool_id || event.id || ""), + tool_name: String(event.tool_name || event.name || ""), + arguments: event.arguments as string, + }; + + case "ToolEnd": + case "tool_end": + return { + type: "tool_end", + tool_id: String(event.tool_id || event.id || ""), + result: event.result as ToolEndEvent["result"], + }; + + case "ActionRequired": + case "action_required": + return { + type: "action_required", + request_id: String(event.request_id || ""), + action_type: String(event.action_type || ""), + tool_name: event.tool_name as string, + arguments: event.arguments as string, + prompt: event.prompt as string, + questions: event.questions as unknown[], + requested_schema: event.requested_schema, + }; + + case "Done": + case "done": + return { type: "done" }; + + case "FinalDone": + case "final_done": + return { + type: "final_done", + usage: event.usage as FinalDoneEvent["usage"], + }; + + case "Error": + case "error": + return { + type: "error", + message: String(event.message || "Unknown error"), + }; + + default: + console.warn("[parseStreamEvent] 未知事件类型:", type); + return null; + } +} + +/** + * 生成唯一事件名称 + */ +export function generateEventName(sessionId: string): string { + return `unified-chat-stream-${sessionId}-${Date.now()}`; +} diff --git a/src/lib/writeFile/README.md b/src/lib/writeFile/README.md new file mode 100644 index 000000000..646d1b46b --- /dev/null +++ b/src/lib/writeFile/README.md @@ -0,0 +1,67 @@ +# WriteFile 解析器模块 + +## 概述 + +本模块提供 `` 标签的流式解析功能,用于从 AI 响应中提取文件内容并实时显示在画布中。 + +## 文件索引 + +| 文件 | 说明 | +|------|------| +| `parser.ts` | WriteFile 标签解析器实现 | +| `index.ts` | 模块导出 | + +## 核心功能 + +### WriteFileParser + +流式解析 `...` 标签。 + +```typescript +import { WriteFileParser } from "@/lib/writeFile"; + +// 静态方法:一次性解析 +const result = WriteFileParser.parse(text); + +// 实例方法:流式解析 +const parser = new WriteFileParser(); +const result1 = parser.parse(partialText); +const result2 = parser.parse(moreText); +``` + +### 解析结果 + +```typescript +interface WriteFileParseResult { + blocks: WriteFileBlock[]; // 解析出的文件块 + plainText: string; // 去除标签后的纯文本 + hasStreamingBlock: boolean; // 是否有正在解析的块 +} + +interface WriteFileBlock { + path: string; // 文件路径 + content: string; // 文件内容 + isComplete: boolean; // 是否解析完成 + startIndex: number; // 起始位置 + endIndex: number; // 结束位置 +} +``` + +## 使用场景 + +1. **流式画布更新**:AI 输出 `` 标签时,实时解析并更新画布内容 +2. **内容分离**:将文件内容从聊天消息中分离,分别显示在画布和聊天区域 + +## 与系统提示词配合 + +系统提示词 (`systemPrompt.ts`) 指导 AI 使用 `` 标签输出长文内容: + +```markdown + +# 文章标题 + +内容... + +``` + +前端检测到此标签后,自动打开画布并流式显示内容。 diff --git a/src/lib/writeFile/index.ts b/src/lib/writeFile/index.ts new file mode 100644 index 000000000..40c119f72 --- /dev/null +++ b/src/lib/writeFile/index.ts @@ -0,0 +1,8 @@ +/** + * @file WriteFile 模块导出 + * @description 导出 WriteFile 解析器相关功能 + * @module lib/writeFile + */ + +export { WriteFileParser } from "./parser"; +export type { WriteFileBlock, WriteFileParseResult } from "./parser"; diff --git a/src/lib/writeFile/parser.ts b/src/lib/writeFile/parser.ts new file mode 100644 index 000000000..e5b1ae829 --- /dev/null +++ b/src/lib/writeFile/parser.ts @@ -0,0 +1,205 @@ +/** + * @file WriteFile 标签解析器 + * @description 从 AI 响应中流式解析 标签 + * @module lib/writeFile/parser + * + * 支持流式解析,实时提取文件内容用于画布显示 + */ + +/** + * WriteFile 解析结果 + */ +export interface WriteFileBlock { + /** 文件路径 */ + path: string; + /** 文件内容 */ + content: string; + /** 是否解析完成 */ + isComplete: boolean; + /** 原始标签在文本中的起始位置 */ + startIndex: number; + /** 原始标签在文本中的结束位置 */ + endIndex: number; +} + +/** + * 解析结果 + */ +export interface WriteFileParseResult { + /** 解析出的文件块列表 */ + blocks: WriteFileBlock[]; + /** 去除 write_file 标签后的纯文本 */ + plainText: string; + /** 是否有正在解析中的块 */ + hasStreamingBlock: boolean; +} + +/** + * 解析器状态 + */ +interface ParserState { + /** 当前正在解析的块 */ + currentBlock: WriteFileBlock | null; + /** 已完成的块 */ + completedBlocks: WriteFileBlock[]; + /** 纯文本部分 */ + plainTextParts: string[]; + /** 上次处理的位置 */ + lastProcessedIndex: number; +} + +/** + * WriteFile 标签解析器 + * 支持流式解析 ... 标签 + */ +export class WriteFileParser { + private state: ParserState; + + constructor() { + this.state = { + currentBlock: null, + completedBlocks: [], + plainTextParts: [], + lastProcessedIndex: 0, + }; + } + + /** + * 重置解析器状态 + */ + reset(): void { + this.state = { + currentBlock: null, + completedBlocks: [], + plainTextParts: [], + lastProcessedIndex: 0, + }; + } + + /** + * 解析文本内容(支持流式调用) + * @param text 完整的文本内容(包含之前的内容) + * @returns 解析结果 + */ + parse(text: string): WriteFileParseResult { + // 重置状态,重新解析整个文本 + this.reset(); + + let currentIndex = 0; + const openTagRegex = //gi; + const closeTag = ""; + + while (currentIndex < text.length) { + // 如果当前在块内,查找结束标签 + if (this.state.currentBlock) { + const closeIndex = text + .toLowerCase() + .indexOf(closeTag.toLowerCase(), currentIndex); + + if (closeIndex !== -1) { + // 找到结束标签,完成当前块 + this.state.currentBlock.content = text.slice( + this.state.currentBlock.startIndex + + this.getOpenTagLength(text, this.state.currentBlock.startIndex), + closeIndex, + ); + this.state.currentBlock.isComplete = true; + this.state.currentBlock.endIndex = closeIndex + closeTag.length; + this.state.completedBlocks.push(this.state.currentBlock); + this.state.currentBlock = null; + currentIndex = closeIndex + closeTag.length; + } else { + // 没有找到结束标签,内容还在流式传输中 + this.state.currentBlock.content = text.slice( + this.state.currentBlock.startIndex + + this.getOpenTagLength(text, this.state.currentBlock.startIndex), + ); + this.state.currentBlock.endIndex = text.length; + break; + } + } else { + // 查找开始标签 + openTagRegex.lastIndex = currentIndex; + const match = openTagRegex.exec(text); + + if (match) { + // 保存开始标签之前的纯文本 + if (match.index > currentIndex) { + this.state.plainTextParts.push( + text.slice(currentIndex, match.index), + ); + } + + // 创建新的块 + this.state.currentBlock = { + path: match[1], + content: "", + isComplete: false, + startIndex: match.index, + endIndex: match.index + match[0].length, + }; + currentIndex = match.index + match[0].length; + } else { + // 没有找到开始标签,剩余都是纯文本 + this.state.plainTextParts.push(text.slice(currentIndex)); + break; + } + } + } + + return this.getResult(); + } + + /** + * 获取开始标签的长度 + */ + private getOpenTagLength(text: string, startIndex: number): number { + const openTagRegex = //gi; + openTagRegex.lastIndex = startIndex; + const match = openTagRegex.exec(text); + return match ? match[0].length : 0; + } + + /** + * 获取解析结果 + */ + private getResult(): WriteFileParseResult { + const blocks = [...this.state.completedBlocks]; + + // 如果有正在解析的块,也加入结果 + if (this.state.currentBlock) { + blocks.push(this.state.currentBlock); + } + + return { + blocks, + plainText: this.state.plainTextParts.join(""), + hasStreamingBlock: this.state.currentBlock !== null, + }; + } + + /** + * 静态方法:一次性解析完整文本 + */ + static parse(text: string): WriteFileParseResult { + const parser = new WriteFileParser(); + return parser.parse(text); + } + + /** + * 检查文本是否包含 write_file 标签 + */ + static hasWriteFileTag(text: string): boolean { + return //i.test(text); + } + + /** + * 获取第一个 write_file 块(用于快速检测) + */ + static getFirstBlock(text: string): WriteFileBlock | null { + const result = WriteFileParser.parse(text); + return result.blocks.length > 0 ? result.blocks[0] : null; + } +} + +export default WriteFileParser; diff --git a/src/types/chat.ts b/src/types/chat.ts new file mode 100644 index 000000000..14db2d3c8 --- /dev/null +++ b/src/types/chat.ts @@ -0,0 +1,384 @@ +/** + * @file chat.ts + * @description 统一对话系统类型定义 + * @module types/chat + * + * 定义了统一对话系统的核心类型,支持多种对话模式: + * - Agent: AI Agent 模式,支持工具调用 + * - General: 通用对话模式,纯文本 + * - Creator: 内容创作模式,支持画布输出 + */ + +// ============================================================================ +// 基础类型 +// ============================================================================ + +/** 对话模式 */ +export type ChatMode = "agent" | "general" | "creator"; + +/** 消息角色 */ +export type MessageRole = "user" | "assistant" | "system" | "tool"; + +/** 消息状态 */ +export type MessageStatus = "pending" | "streaming" | "complete" | "error"; + +// ============================================================================ +// 会话类型 +// ============================================================================ + +/** 统一会话结构 */ +export interface ChatSession { + /** 会话 ID */ + id: string; + /** 对话模式 */ + mode: ChatMode; + /** 会话标题 */ + title?: string; + /** 系统提示词 */ + systemPrompt?: string; + /** 模型名称 */ + model?: string; + /** Provider 类型 */ + providerType?: string; + /** 凭证 UUID */ + credentialUuid?: string; + /** 扩展元数据 */ + metadata?: Record; + /** 创建时间 */ + createdAt: string; + /** 更新时间 */ + updatedAt: string; + /** 消息数量 */ + messageCount?: number; +} + +// ============================================================================ +// 消息类型 +// ============================================================================ + +/** 消息内容块类型 */ +export type ContentBlockType = + | "text" + | "image" + | "file" + | "canvas" + | "tool_result"; + +/** 文本内容块 */ +export interface TextContentBlock { + type: "text"; + text: string; +} + +/** 图片内容块 */ +export interface ImageContentBlock { + type: "image"; + url: string; + alt?: string; +} + +/** 文件内容块 */ +export interface FileContentBlock { + type: "file"; + path: string; + name: string; +} + +/** 画布内容块 */ +export interface CanvasContentBlock { + type: "canvas"; + canvasType: string; + content: string; +} + +/** 工具结果内容块 */ +export interface ToolResultContentBlock { + type: "tool_result"; + toolCallId: string; + result: unknown; +} + +/** 内容块联合类型 */ +export type ContentBlock = + | TextContentBlock + | ImageContentBlock + | FileContentBlock + | CanvasContentBlock + | ToolResultContentBlock; + +/** 工具调用状态 */ +export type ToolCallStatus = "pending" | "running" | "completed" | "failed"; + +/** 工具调用信息 */ +export interface ToolCall { + id: string; + name: string; + arguments?: string; + status: ToolCallStatus; + result?: { + success: boolean; + output?: string; + error?: string; + }; + startTime?: Date; + endTime?: Date; +} + +/** 统一消息结构 */ +export interface ChatMessage { + /** 消息 ID */ + id: string | number; + /** 会话 ID */ + sessionId: string; + /** 角色 */ + role: MessageRole; + /** 文本内容(便捷访问) */ + content: string; + /** 结构化内容块 */ + contentBlocks?: ContentBlock[]; + /** 工具调用列表 */ + toolCalls?: ToolCall[]; + /** 工具调用 ID(用于工具响应) */ + toolCallId?: string; + /** 消息状态 */ + status: MessageStatus; + /** 错误信息 */ + error?: ChatError; + /** 扩展元数据 */ + metadata?: MessageMetadata; + /** 创建时间 */ + createdAt: string | Date; +} + +/** 消息元数据 */ +export interface MessageMetadata { + /** 使用的模型 */ + model?: string; + /** Token 使用量 */ + tokens?: { + input?: number; + output?: number; + total?: number; + }; + /** 响应耗时(毫秒) */ + duration?: number; + /** 其他自定义数据 */ + [key: string]: unknown; +} + +// ============================================================================ +// 错误类型 +// ============================================================================ + +/** 错误类型 */ +export type ChatErrorType = + | "network" + | "auth" + | "rate_limit" + | "quota" + | "invalid_request" + | "server" + | "unknown"; + +/** 错误信息 */ +export interface ChatError { + type: ChatErrorType; + message: string; + code?: string; + retryable: boolean; +} + +// ============================================================================ +// 图片输入类型 +// ============================================================================ + +/** 图片输入 */ +export interface ImageInput { + /** Base64 编码的图片数据 */ + data: string; + /** MIME 类型 */ + mediaType: string; +} + +// ============================================================================ +// API 请求/响应类型 +// ============================================================================ + +/** 创建会话请求 */ +export interface CreateSessionRequest { + mode: ChatMode; + title?: string; + systemPrompt?: string; + providerType?: string; + model?: string; + metadata?: Record; +} + +/** 发送消息请求 */ +export interface SendMessageRequest { + sessionId: string; + message: string; + eventName: string; + images?: Array<{ data: string; media_type: string }>; +} + +/** 会话响应 */ +export interface SessionResponse { + id: string; + mode: ChatMode; + title?: string; + model?: string; + createdAt: string; + updatedAt: string; + messageCount: number; +} + +// ============================================================================ +// 流式事件类型 +// ============================================================================ + +/** 流式事件类型 */ +export type StreamEventType = + | "text_delta" + | "thinking_delta" + | "tool_start" + | "tool_end" + | "action_required" + | "done" + | "final_done" + | "error"; + +/** 文本增量事件 */ +export interface TextDeltaEvent { + type: "text_delta"; + text: string; +} + +/** 思考增量事件 */ +export interface ThinkingDeltaEvent { + type: "thinking_delta"; + text: string; +} + +/** 工具开始事件 */ +export interface ToolStartEvent { + type: "tool_start"; + tool_id: string; + tool_name: string; + arguments?: string; +} + +/** 工具结束事件 */ +export interface ToolEndEvent { + type: "tool_end"; + tool_id: string; + result: { + success: boolean; + output?: string; + error?: string; + }; +} + +/** 权限请求事件 */ +export interface ActionRequiredEvent { + type: "action_required"; + request_id: string; + action_type: string; + tool_name?: string; + arguments?: string; + prompt?: string; + questions?: unknown[]; + requested_schema?: unknown; +} + +/** 完成事件 */ +export interface DoneEvent { + type: "done"; +} + +/** 最终完成事件 */ +export interface FinalDoneEvent { + type: "final_done"; + usage?: { + input_tokens?: number; + output_tokens?: number; + }; +} + +/** 错误事件 */ +export interface ErrorEvent { + type: "error"; + message: string; +} + +/** 流式事件联合类型 */ +export type StreamEvent = + | TextDeltaEvent + | ThinkingDeltaEvent + | ToolStartEvent + | ToolEndEvent + | ActionRequiredEvent + | DoneEvent + | FinalDoneEvent + | ErrorEvent; + +// ============================================================================ +// Hook 配置类型 +// ============================================================================ + +/** useUnifiedChat 配置选项 */ +export interface UseUnifiedChatOptions { + /** 对话模式 */ + mode: ChatMode; + /** 初始会话 ID(可选) */ + sessionId?: string; + /** 系统提示词(可选) */ + systemPrompt?: string; + /** Provider 类型(可选) */ + providerType?: string; + /** 模型名称(可选) */ + model?: string; + /** 画布内容更新回调 */ + onCanvasUpdate?: (path: string, content: string) => void; + /** 文件写入回调 */ + onWriteFile?: (content: string, fileName: string) => void; + /** 错误回调 */ + onError?: (error: ChatError) => void; +} + +/** useUnifiedChat 返回值 */ +export interface UseUnifiedChatReturn { + // 状态 + /** 当前会话 */ + session: ChatSession | null; + /** 消息列表 */ + messages: ChatMessage[]; + /** 是否正在加载 */ + isLoading: boolean; + /** 是否正在发送 */ + isSending: boolean; + /** 错误信息 */ + error: ChatError | null; + + // 会话操作 + /** 创建新会话 */ + createSession: (options?: Partial) => Promise; + /** 加载会话 */ + loadSession: (sessionId: string) => Promise; + /** 删除会话 */ + deleteSession: (sessionId?: string) => Promise; + /** 重命名会话 */ + renameSession: (title: string, sessionId?: string) => Promise; + + // 消息操作 + /** 发送消息 */ + sendMessage: (content: string, images?: ImageInput[]) => Promise; + /** 停止生成 */ + stopGeneration: () => Promise; + /** 清空消息 */ + clearMessages: () => void; + + // Provider 配置 + /** 配置 Provider */ + configureProvider: (providerType: string, model: string) => Promise; +}