From bb244fd1e97b4bebc55fa580e7555972f8a9f28e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=8D=A1=E4=BB=94?= <1271013637@qq.com> Date: Mon, 3 Nov 2025 16:23:31 +0800 Subject: [PATCH] fix(view): :bug: reuse the IM client to handle AI data streams --- src-tauri/src/command/ai_command.rs | 90 ++++++------------- src-tauri/src/im_request_client.rs | 134 +++++++++++++++++++++++----- src/utils/ImRequestUtils.ts | 108 +++++++++++----------- 3 files changed, 188 insertions(+), 144 deletions(-) diff --git a/src-tauri/src/command/ai_command.rs b/src-tauri/src/command/ai_command.rs index c8490f4e..a9dc448c 100644 --- a/src-tauri/src/command/ai_command.rs +++ b/src-tauri/src/command/ai_command.rs @@ -1,7 +1,8 @@ +use crate::im_request_client::ImUrl; use crate::AppData; use futures::StreamExt; use serde::{Deserialize, Serialize}; -use tauri::{AppHandle, Emitter, State}; +use tauri::{ipc::Channel, State}; use tracing::{error, info}; /// SSE 流式数据事件 @@ -33,73 +34,36 @@ pub async fn ai_message_send_stream( state: State<'_, AppData>, body: AiMessageRequest, request_id: String, - app_handle: AppHandle, + on_event: Channel, ) -> Result<(), String> { info!( "🤖 开始发送 AI 流式消息请求, conversation_id: {}, request_id: {}", body.conversation_id, request_id ); - // 获取配置和token - let (base_url, token) = { - let config = state.config.lock().await; - let rc = state.rc.lock().await; - (config.backend.base_url.clone(), rc.token.clone()) - }; + // 使用 ImRequestClient 发送流式请求 + let response = { + let mut rc = state.rc.lock().await; + let (method, path) = ImUrl::MessageSendStream.get_url(); - // 构建完整 URL - 使用硬编码的路径,因为 get_url 是私有的 - let url = format!("{}/ai/chat/message/send-stream", base_url); - info!("📡 SSE Request URL: {}", url); - - // 创建 HTTP 客户端 - let client = reqwest::Client::new(); - - // 构建请求 - let mut request = client.post(&url).json(&body); - - // 添加 token - if let Some(token) = token { - request = request.header("token", token); - } - - // 添加 Accept 头以接收 SSE - request = request.header("Accept", "text/event-stream"); - - // 发送请求 - let response = match request.send().await { - Ok(resp) => resp, - Err(e) => { - error!("❌ 发送请求失败: {}", e); - let error_event = SseStreamEvent { - event_type: "error".to_string(), - data: None, - error: Some(format!("发送请求失败: {}", e)), - request_id: request_id.clone(), - }; - let _ = app_handle.emit("ai-stream-event", error_event); - return Err(format!("发送请求失败: {}", e)); - } - }; - - // 检查响应状态 - if !response.status().is_success() { - let status = response.status(); - let error_msg = format!("请求失败,状态码: {}", status); - error!("❌ {}", error_msg); - let error_event = SseStreamEvent { - event_type: "error".to_string(), - data: None, - error: Some(error_msg.clone()), - request_id: request_id.clone(), - }; - let _ = app_handle.emit("ai-stream-event", error_event); - return Err(error_msg); - } + rc.request_stream(method, path, Some(body), None::) + .await + .map_err(|e| { + error!("❌ 发送流式请求失败: {}", e); + let error_event = SseStreamEvent { + event_type: "error".to_string(), + data: None, + error: Some(e.to_string()), + request_id: request_id.clone(), + }; + let _ = on_event.send(error_event); + e.to_string() + })? + }; // 锁在这里释放 info!("✅ SSE 连接已建立,开始监听流式数据..."); // 在后台任务中处理 SSE 事件流 - let app_handle_clone = app_handle.clone(); let request_id_clone = request_id.clone(); tokio::spawn(async move { @@ -141,9 +105,7 @@ pub async fn ai_message_send_stream( request_id: request_id_clone.clone(), }; - if let Err(e) = - app_handle_clone.emit("ai-stream-event", chunk_event) - { + if let Err(e) = on_event.send(chunk_event) { error!("❌ 发送 chunk 事件失败: {}", e); } } else if line.starts_with("data:") { @@ -162,9 +124,7 @@ pub async fn ai_message_send_stream( request_id: request_id_clone.clone(), }; - if let Err(e) = - app_handle_clone.emit("ai-stream-event", chunk_event) - { + if let Err(e) = on_event.send(chunk_event) { error!("❌ 发送 chunk 事件失败: {}", e); } } @@ -183,7 +143,7 @@ pub async fn ai_message_send_stream( request_id: request_id_clone.clone(), }; - if let Err(e) = app_handle_clone.emit("ai-stream-event", error_event) { + if let Err(e) = on_event.send(error_event) { error!("❌ 发送 error 事件失败: {}", e); } break; @@ -200,7 +160,7 @@ pub async fn ai_message_send_stream( request_id: request_id_clone.clone(), }; - if let Err(e) = app_handle_clone.emit("ai-stream-event", done_event) { + if let Err(e) = on_event.send(done_event) { error!("❌ 发送 done 事件失败: {}", e); } diff --git a/src-tauri/src/im_request_client.rs b/src-tauri/src/im_request_client.rs index f07a08ad..db152d97 100644 --- a/src-tauri/src/im_request_client.rs +++ b/src-tauri/src/im_request_client.rs @@ -48,6 +48,56 @@ impl ImRequestClient { self.base_url = base_url; } + /// 构建请求的公共方法(不发送请求) + /// + /// 提取了 URL 构建、token 添加、body/params 处理等公共逻辑 + /// + /// # 参数 + /// - `method`: HTTP 方法 + /// - `path`: API 路径 + /// - `body`: 请求体(可选) + /// - `params`: 查询参数(可选) + /// - `extra_headers`: 额外的请求头(可选) + fn build_request( + &self, + method: http::Method, + path: &str, + body: &Option, + params: &Option, + extra_headers: Option>, + ) -> reqwest::RequestBuilder { + let url = format!("{}/{}", self.base_url, path); + info!("📡 Request URL: {}, Method: {}", &url, method); + + let mut request_builder = self.client.request(method, &url); + + // 设置 token 请求头 + if let Some(token) = &self.token { + request_builder = request_builder.header("token", token); + } + + // 添加额外的请求头 + if let Some(headers) = extra_headers { + for (key, value) in headers { + request_builder = request_builder.header(key, value); + } + } + + // 设置请求体 + if let Some(body) = body { + request_builder = request_builder.json(body); + } else { + request_builder = request_builder.json(&serde_json::json!({})); + } + + // 设置查询参数 + if let Some(params) = params { + request_builder = request_builder.query(params); + } + + request_builder + } + pub async fn request< T: serde::de::DeserializeOwned, B: serde::Serialize, @@ -63,31 +113,15 @@ impl ImRequestClient { const MAX_RETRY_COUNT: u8 = 2; loop { - let url = format!("{}/{}", self.base_url, path); - info!("📡 Request URL: {}, Method: {}", &url, method.clone()); - - let mut request_builder = self.client.request(method.clone(), &url); - - // 设置请求头 - if let Some(token) = self.token.clone() { - request_builder = request_builder.header("token", token); - } - - // 设置请求体 - if let Some(body) = &body { - request_builder = request_builder.json(body); - } else { - request_builder = request_builder.json(&serde_json::json!({})); - } - - if let Some(params) = ¶ms { - request_builder = request_builder.query(params); - } + // 使用 build_request 构建请求 + let request_builder = self.build_request(method.clone(), path, &body, ¶ms, None); // 发送请求 let response = request_builder.send().await?; let result: ApiResult = response.json().await?; + let url = format!("{}/{}", self.base_url, path); + match result.code { Some(406) => { if retry_count >= MAX_RETRY_COUNT { @@ -132,6 +166,64 @@ impl ImRequestClient { } } + /// 流式请求方法(用于 SSE 等流式响应) + /// + /// 与 `request` 方法的区别: + /// 1. 添加 `Accept: text/event-stream` 请求头 + /// 2. 返回 `reqwest::Response` 而不是解析 JSON + /// 3. 不支持自动 token 刷新重试(因为流式响应无法中断重试) + /// + /// # 参数 + /// - `method`: HTTP 方法 + /// - `path`: API 路径 + /// - `body`: 请求体(可选) + /// - `params`: 查询参数(可选) + /// + /// # 返回 + /// - `Ok(Response)`: 成功返回响应对象,可用于读取流式数据 + /// - `Err`: 请求失败或状态码非 2xx + pub async fn request_stream( + &mut self, + method: http::Method, + path: &str, + body: Option, + params: Option, + ) -> Result { + // 添加流式请求头 + let extra_headers = Some(vec![("Accept", "text/event-stream")]); + + // 使用 build_request 构建请求 + let request_builder = self.build_request(method.clone(), path, &body, ¶ms, extra_headers); + + // 发送请求 + let response = request_builder.send().await?; + + // 检查响应状态(但不解析 JSON) + let status = response.status(); + if !status.is_success() { + let url = format!("{}/{}", self.base_url, path); + error!("❌ 流式请求失败,URL: {}, 状态码: {}", url, status); + + // 根据状态码返回不同的错误信息 + match status.as_u16() { + 406 => { + error!("🔄 Token expired in stream request"); + return Err(anyhow::anyhow!("token过期,请刷新后重试")); + } + 401 => { + error!("🔐 Unauthorized in stream request"); + return Err(anyhow::anyhow!("请重新登录")); + } + _ => { + return Err(anyhow::anyhow!("请求失败,状态码: {}", status)); + } + } + } + + info!("✅ 流式请求成功,开始接收流式数据"); + Ok(response) + } + pub async fn start_refresh_token(&mut self) -> Result<(), anyhow::Error> { info!("🔄 Starting token refresh"); let url = format!("{}/{}", self.base_url, ImUrl::RefreshToken.get_url().1); @@ -422,7 +514,7 @@ pub enum ImUrl { } impl ImUrl { - fn get_url(&self) -> (http::Method, &str) { + pub fn get_url(&self) -> (http::Method, &str) { match self { // Token 相关 ImUrl::Login => (http::Method::POST, "oauth/anyTenant/login"), diff --git a/src/utils/ImRequestUtils.ts b/src/utils/ImRequestUtils.ts index 8bee3d44..b84d9ed5 100644 --- a/src/utils/ImRequestUtils.ts +++ b/src/utils/ImRequestUtils.ts @@ -736,75 +736,67 @@ export async function messageSendStream( body: { conversationId: string; content: string; useContext?: boolean }, callbacks?: StreamCallbacks ): Promise { - const { invoke } = await import('@tauri-apps/api/core') - const { listen } = await import('@tauri-apps/api/event') + const { invoke, Channel } = await import('@tauri-apps/api/core') const { TauriCommand } = await import('@/enums') // 生成唯一的请求 ID const requestId = `ai-stream-${Date.now()}-${Math.random().toString(36).substring(2, 9)}` return new Promise((resolve, reject) => { - let unlisten: (() => void) | null = null let fullContent = '' let isResolved = false - // 使用立即执行的 async 函数 - ;(async () => { - try { - // 监听 SSE 流式事件 - unlisten = await listen('ai-stream-event', (event) => { - const { eventType, data, error, requestId: eventRequestId } = event.payload + // 创建 Channel 用于接收流式事件 + const onEvent = new Channel() + onEvent.onmessage = (event: SseStreamEvent) => { + const { eventType, data, error, requestId: eventRequestId } = event - // 只处理当前请求的事件 - if (eventRequestId !== requestId) { - return - } - - switch (eventType) { - case 'chunk': - if (data) { - fullContent += data - callbacks?.onChunk?.(data) - } - break - - case 'done': - if (!isResolved) { - isResolved = true - const finalContent = data || fullContent - callbacks?.onDone?.(finalContent) - if (unlisten) unlisten() - resolve(finalContent) - } - break - - case 'error': - if (!isResolved) { - isResolved = true - const errorMsg = error || '未知错误' - callbacks?.onError?.(errorMsg) - if (unlisten) unlisten() - reject(new Error(errorMsg)) - } - break - } - }) - - // 调用 Rust 后端命令发送请求 - await invoke(TauriCommand.AI_MESSAGE_SEND_STREAM, { - body, - requestId - }) - } catch (error) { - if (!isResolved) { - isResolved = true - if (unlisten) unlisten() - const errorMsg = error instanceof Error ? error.message : String(error) - callbacks?.onError?.(errorMsg) - reject(error) - } + // 只处理当前请求的事件 + if (eventRequestId !== requestId) { + return } - })() + + switch (eventType) { + case 'chunk': + if (data) { + fullContent += data + callbacks?.onChunk?.(data) + } + break + + case 'done': + if (!isResolved) { + isResolved = true + const finalContent = data || fullContent + callbacks?.onDone?.(finalContent) + resolve(finalContent) + } + break + + case 'error': + if (!isResolved) { + isResolved = true + const errorMsg = error || '未知错误' + callbacks?.onError?.(errorMsg) + reject(new Error(errorMsg)) + } + break + } + } + + // 调用 Rust 后端命令发送请求 + invoke(TauriCommand.AI_MESSAGE_SEND_STREAM, { + body, + requestId, + onEvent + }).catch((error) => { + if (!isResolved) { + isResolved = true + const errorMsg = error instanceof Error ? error.message : String(error) + callbacks?.onError?.(errorMsg) + reject(error) + } + }) }) }