fix(view): 🐛 reuse the IM client to handle AI data streams

This commit is contained in:
卡仔
2025-11-04 15:43:59 +08:00
committed by Dawn
parent d8ebde5b8c
commit bb244fd1e9
3 changed files with 188 additions and 144 deletions
+25 -65
View File
@@ -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<SseStreamEvent>,
) -> 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::<serde_json::Value>)
.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);
}
+113 -21
View File
@@ -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<B: serde::Serialize, C: serde::Serialize>(
&self,
method: http::Method,
path: &str,
body: &Option<B>,
params: &Option<C>,
extra_headers: Option<Vec<(&str, &str)>>,
) -> 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) = &params {
request_builder = request_builder.query(params);
}
// 使用 build_request 构建请求
let request_builder = self.build_request(method.clone(), path, &body, &params, None);
// 发送请求
let response = request_builder.send().await?;
let result: ApiResult<T> = 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<B: serde::Serialize, C: serde::Serialize>(
&mut self,
method: http::Method,
path: &str,
body: Option<B>,
params: Option<C>,
) -> Result<reqwest::Response, anyhow::Error> {
// 添加流式请求头
let extra_headers = Some(vec![("Accept", "text/event-stream")]);
// 使用 build_request 构建请求
let request_builder = self.build_request(method.clone(), path, &body, &params, 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"),
+50 -58
View File
@@ -736,75 +736,67 @@ export async function messageSendStream(
body: { conversationId: string; content: string; useContext?: boolean },
callbacks?: StreamCallbacks
): Promise<string> {
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<string>((resolve, reject) => {
let unlisten: (() => void) | null = null
let fullContent = ''
let isResolved = false
// 使用立即执行的 async 函数
;(async () => {
try {
// 监听 SSE 流式事件
unlisten = await listen<SseStreamEvent>('ai-stream-event', (event) => {
const { eventType, data, error, requestId: eventRequestId } = event.payload
// 创建 Channel 用于接收流式事件
const onEvent = new Channel<SseStreamEvent>()
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)
}
})
})
}