|
| 1 | +//! AI 对话流式发送(REQ-225,v0.16.0)。 |
| 2 | +//! |
| 3 | +//! @ai-context: 纯聊天走 `stream: true` SSE——逐 delta 经 Tauri Channel 推给 |
| 4 | +//! 前端(打字效果 + 可停止)。与 ai_client.rs 的阻塞式 |
| 5 | +//! post_completions(精修/补充用)分居:流式**不自动重试** |
| 6 | +//! (重试=重复生成,聊天天然由用户"重发"触发;精修幂等才重试)。 |
| 7 | +//! @ai-context: usage 口径各家不一(末 chunk 附带 / 独立 chunk / 无)——本层 |
| 8 | +//! 只"看见则存"(最后一个带 usage 的 data 行),看不见则为 None |
| 9 | +//! (前端仅显示 token(如获知)与估算成本,不阻塞会话)。 |
| 10 | +
|
| 11 | +use std::io::{BufRead, BufReader}; |
| 12 | + |
| 13 | +use serde::Serialize; |
| 14 | + |
| 15 | +use crate::ai_chat::{CancelFlag, SseEvent, parse_sse_line}; |
| 16 | +use crate::ai_client::{AiClient, AiClientError, chat_completions_url}; |
| 17 | + |
| 18 | +/// 流式事件(Tauri Channel 载荷契约——kind 标签,前端按 kind 分发)。 |
| 19 | +#[derive(Debug, Clone, PartialEq, Serialize)] |
| 20 | +#[serde(tag = "kind", rename_all = "snake_case", rename_all_fields = "camelCase")] |
| 21 | +pub enum ChatStreamEvent { |
| 22 | + /// 增量文本(流式打字) |
| 23 | + Chunk { delta: String }, |
| 24 | + /// 流正常结束(附用量 JSON 原样——含 token/成本口径的原始数据) |
| 25 | + Done { content: String, usage_json: Option<String> }, |
| 26 | + /// 失败(AiClientError 六类归一;前端映射 + 重试按钮) |
| 27 | + Failed { error_kind: String, message: String }, |
| 28 | + /// 用户取消(content=已生成文本;消息落库标 aborted) |
| 29 | + Aborted { content: String }, |
| 30 | +} |
| 31 | + |
| 32 | +impl From<&AiClientError> for ChatStreamEvent { |
| 33 | + fn from(e: &AiClientError) -> Self { |
| 34 | + ChatStreamEvent::Failed { error_kind: e.kind().to_string(), message: e.to_string() } |
| 35 | + } |
| 36 | +} |
| 37 | + |
| 38 | +impl AiClientError { |
| 39 | + /// 错误类别标签(前端映射 + 任务失败四类契约复用)。 |
| 40 | + pub fn kind(&self) -> &'static str { |
| 41 | + match self { |
| 42 | + AiClientError::Auth(_) => "auth", |
| 43 | + AiClientError::Network(_) => "network", |
| 44 | + AiClientError::Balance(_) => "balance", |
| 45 | + AiClientError::Quota(_) => "quota", |
| 46 | + AiClientError::Server(_) => "server", |
| 47 | + AiClientError::Parse(_) => "parse", |
| 48 | + } |
| 49 | + } |
| 50 | +} |
| 51 | + |
| 52 | +/// 流式结果(content/usage/是否取消——command 层落库口径)。 |
| 53 | +#[derive(Debug)] |
| 54 | +pub struct StreamOutcome { |
| 55 | + pub content: String, |
| 56 | + pub usage_json: Option<String>, |
| 57 | + pub cancelled: bool, |
| 58 | +} |
| 59 | + |
| 60 | +/// 发送流式 chat/completions(SSE),逐 delta 回调 emit。 |
| 61 | +/// |
| 62 | +/// @ai-context: 取消语义:每读一行检查 CancelFlag(Arc 共享——chat_cancel |
| 63 | +/// 命令置位);响应头未到时的取消由 HTTP 超时兜底(不做 |
| 64 | +/// abort transport——ureq 无取消句柄,超时后命令层重试/重发)。 |
| 65 | +/// emit 为 FnMut 无返回:Channel 发送失败(前端已关)由 |
| 66 | +/// command 层静默降级(不阻断落库,数据不丢)。 |
| 67 | +pub fn stream_chat( |
| 68 | + client: &AiClient, |
| 69 | + messages: &[serde_json::Value], |
| 70 | + cancel: &CancelFlag, |
| 71 | + mut emit: impl FnMut(ChatStreamEvent), |
| 72 | +) -> Result<StreamOutcome, AiClientError> { |
| 73 | + if !client.config.is_local && client.config.api_key.trim().is_empty() { |
| 74 | + return Err(AiClientError::Auth( |
| 75 | + "未配置 API 密钥(设置页保存或配置环境变量)".to_string(), |
| 76 | + )); |
| 77 | + } |
| 78 | + let mut payload = serde_json::json!({ |
| 79 | + "model": client.config.model, |
| 80 | + "messages": messages, |
| 81 | + "temperature": 0.7, |
| 82 | + "max_tokens": client.config.max_tokens, |
| 83 | + "stream": true, |
| 84 | + }); |
| 85 | + // R1 系推理模型:关闭思考标签(与 build_chat_payload 同款约束——保持 |
| 86 | + // 输出稳定;DeepSeek-R1 在线端点会输出 reasoning_content 占 delta) |
| 87 | + if client.config.model.to_lowercase().contains("r1") { |
| 88 | + payload["no_think"] = serde_json::json!(true); |
| 89 | + } |
| 90 | + let url = chat_completions_url(&client.config.base_url); |
| 91 | + let agent = ureq::AgentBuilder::new() |
| 92 | + .timeout(std::time::Duration::from_secs(client.config.timeout_secs.max(5))) |
| 93 | + .build(); |
| 94 | + let resp = agent |
| 95 | + .post(&url) |
| 96 | + .set("Content-Type", "application/json") |
| 97 | + .set("Authorization", &format!("Bearer {}", client.config.api_key.trim())) |
| 98 | + .send_string(&payload.to_string()) |
| 99 | + .map_err(map_status)?; |
| 100 | + let reader = BufReader::new(resp.into_reader()); |
| 101 | + let mut content = String::new(); |
| 102 | + let mut usage_json: Option<String> = None; |
| 103 | + let mut cancelled = false; |
| 104 | + for line in reader.lines() { |
| 105 | + if cancel.is_cancelled() { |
| 106 | + cancelled = true; |
| 107 | + break; |
| 108 | + } |
| 109 | + let line = match line { |
| 110 | + Ok(l) => l, |
| 111 | + Err(_) => break, // 传输中途断流:以已累积内容为准(无 usage) |
| 112 | + }; |
| 113 | + match parse_sse_line(&line) { |
| 114 | + SseEvent::Delta(d) => { |
| 115 | + content.push_str(&d); |
| 116 | + emit(ChatStreamEvent::Chunk { delta: d }); |
| 117 | + } |
| 118 | + SseEvent::Done => break, |
| 119 | + SseEvent::Ignore => { |
| 120 | + // usage 可能挂在非 delta 的 data 行(OpenAI 兼容末 chunk) |
| 121 | + if let Some(usage) = extract_usage(&line) { |
| 122 | + usage_json = Some(usage); |
| 123 | + } |
| 124 | + } |
| 125 | + } |
| 126 | + } |
| 127 | + Ok(StreamOutcome { content, usage_json, cancelled }) |
| 128 | +} |
| 129 | + |
| 130 | +/// 从 data 行提取 usage(纯函数;无 usage → None)。 |
| 131 | +pub fn extract_usage(line: &str) -> Option<String> { |
| 132 | + let line = line.trim(); |
| 133 | + let payload = line.strip_prefix("data:")?.trim(); |
| 134 | + let v: serde_json::Value = serde_json::from_str(payload).ok()?; |
| 135 | + let u = v.get("usage")?; |
| 136 | + if u.is_null() { |
| 137 | + None |
| 138 | + } else { |
| 139 | + serde_json::to_string(u).ok() |
| 140 | + } |
| 141 | +} |
| 142 | + |
| 143 | +/// HTTP 状态 → AiClientError(与 post_completions 同归一口径——四下一致)。 |
| 144 | +fn map_status(e: ureq::Error) -> AiClientError { |
| 145 | + match e { |
| 146 | + ureq::Error::Status(401, _) => AiClientError::Auth( |
| 147 | + "API 密钥无效(HTTP 401)——请检查设置页密钥或环境变量".to_string(), |
| 148 | + ), |
| 149 | + ureq::Error::Status(403, _) => AiClientError::Auth( |
| 150 | + "API 密钥无权限(HTTP 403)——账号未开通该模型或额度受限".to_string(), |
| 151 | + ), |
| 152 | + ureq::Error::Status(402, _) => { |
| 153 | + AiClientError::Balance("账户余额不足(请充值或切换免费档模型)".to_string()) |
| 154 | + } |
| 155 | + ureq::Error::Status(429, _) => { |
| 156 | + AiClientError::Quota("请求过频或配额耗尽(HTTP 429)".to_string()) |
| 157 | + } |
| 158 | + ureq::Error::Status(code, _) if code >= 500 => { |
| 159 | + AiClientError::Server(format!("服务端错误 HTTP {}", code)) |
| 160 | + } |
| 161 | + ureq::Error::Status(code, _) => { |
| 162 | + AiClientError::Network(format!("请求被拒绝 HTTP {}(不重试)", code)) |
| 163 | + } |
| 164 | + e => AiClientError::Network(format!("传输错误: {}", e)), |
| 165 | + } |
| 166 | +} |
| 167 | + |
| 168 | +#[cfg(test)] |
| 169 | +#[path = "ai_chat_stream_tests.rs"] |
| 170 | +mod tests; |
0 commit comments