mirror of
https://github.com/HuLaSpark/HuLa.git
synced 2026-09-24 23:11:59 +08:00
fix(view): 🐛 reuse the IM client to handle AI data streams
This commit is contained in:
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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) = ¶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<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, ¶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"),
|
||||
|
||||
+50
-58
@@ -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)
|
||||
}
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user