Skip to content

Commit 6778114

Browse files
committed
feat(ai): 内嵌 AI 对话——纯聊天(流式/停止/重发)+ 精修轨迹对话视图(REQ-224~230,v0.16.0)
1 parent 0253ee2 commit 6778114

28 files changed

Lines changed: 2605 additions & 36 deletions

‎app/src-tauri/src/ai_chat.rs‎

Lines changed: 154 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,154 @@
1+
//! AI 对话纯函数层(REQ-224/225/230,v0.16.0)。
2+
//!
3+
//! @ai-context: 本模块只含"无副作用"的原子逻辑(消息组装 / SSE 行解析 /
4+
//! 轨迹序列化 / 取消标志)——业务编排在 command 层,网络在
5+
//! ai_chat_stream.rs,存储在 db_ai_chat.rs;纯函数 AAA 单测
6+
//! (AGENTS.md §3.5 测试纪律)。
7+
//! @ai-context: 轨迹(AiTurn)= 每次 LLM 调用的提示词与回答全文——"AI 任务
8+
//! 对话视图"数据源(REQ-230 用户裁决:能看提示词和回答)。
9+
//! vision 调用只记图数占位(base64 不入库:图在本机会话图库,
10+
//! 入库会使体积翻数倍且冗余同图)。
11+
12+
use std::sync::Arc;
13+
use std::sync::atomic::{AtomicBool, Ordering};
14+
15+
use serde::{Deserialize, Serialize};
16+
17+
/// 聊天消息角色(OpenAI 兼容白名单——防任意 role 注入,防御性编程)。
18+
#[derive(Debug, Clone, PartialEq)]
19+
pub enum ChatRole {
20+
/// 保留(system 由 build_messages 单独注入——角色不落库不旁路)
21+
#[allow(dead_code)]
22+
System,
23+
User,
24+
Assistant,
25+
}
26+
27+
impl ChatRole {
28+
pub fn as_str(&self) -> &'static str {
29+
match self {
30+
ChatRole::System => "system",
31+
ChatRole::User => "user",
32+
ChatRole::Assistant => "assistant",
33+
}
34+
}
35+
}
36+
37+
/// 历史消息输入(组装前的最小结构——不耦合 DB 行,纯函数可单测)。
38+
#[derive(Debug, Clone, PartialEq)]
39+
pub struct ChatMessageInput {
40+
pub role: ChatRole,
41+
pub content: String,
42+
}
43+
44+
impl ChatMessageInput {
45+
/// 便捷构造(测试/未来上下文注入层使用)
46+
#[allow(dead_code)]
47+
pub fn user(content: impl Into<String>) -> Self {
48+
Self { role: ChatRole::User, content: content.into() }
49+
}
50+
}
51+
52+
/// 组装 OpenAI messages 数组(system 置顶;history 仅 user/assistant)。
53+
///
54+
/// @ai-context: 多轮对话历史一律截断到最近 MAX_HISTORY 条(防长会话 token
55+
/// 失控;超出后第一条 user 摘要占比仍递减——MVP 不做摘要压缩,
56+
/// 与 v0.8 精修"切片 + 片间摘要"策略同思路,后续再议)。
57+
pub const MAX_HISTORY: usize = 30;
58+
59+
pub fn build_messages(system: &str, history: &[ChatMessageInput]) -> Vec<serde_json::Value> {
60+
let mut out = Vec::with_capacity(history.len() + 1);
61+
if !system.trim().is_empty() {
62+
out.push(serde_json::json!({ "role": "system", "content": system }));
63+
}
64+
let tail_start = history.len().saturating_sub(MAX_HISTORY);
65+
for m in &history[tail_start..] {
66+
out.push(serde_json::json!({ "role": m.role.as_str(), "content": m.content }));
67+
}
68+
out
69+
}
70+
71+
/// SSE 行解析结果。
72+
#[derive(Debug, Clone, PartialEq)]
73+
pub enum SseEvent {
74+
/// 增量文本(choices[0].delta.content)
75+
Delta(String),
76+
/// 流结束标记 `data: [DONE]`
77+
Done,
78+
/// 非 data 行 / 空增量 / 畸形 JSON(跳过不失败——服务商行尾差异容错)
79+
Ignore,
80+
}
81+
82+
/// 解析一行 SSE(`data: {json}` 或 `data: [DONE]`)。
83+
///
84+
/// @ai-context: 流式响应各服务商末 chunk 的 usage/finish_reason 字段不一致
85+
/// (OpenAI 单独 usage chunk / DeepSeek 附在末 chunk / 无),
86+
/// 本解析只取 delta.content;用量另由非流式兜底与账单口径处理。
87+
pub fn parse_sse_line(line: &str) -> SseEvent {
88+
let line = line.trim();
89+
if line.is_empty() {
90+
return SseEvent::Ignore;
91+
}
92+
let Some(payload) = line.strip_prefix("data:") else {
93+
return SseEvent::Ignore;
94+
};
95+
let payload = payload.trim();
96+
if payload == "[DONE]" {
97+
return SseEvent::Done;
98+
}
99+
let Ok(v) = serde_json::from_str::<serde_json::Value>(payload) else {
100+
return SseEvent::Ignore;
101+
};
102+
let delta = v["choices"][0]["delta"]["content"].as_str().unwrap_or("");
103+
if delta.is_empty() {
104+
SseEvent::Ignore
105+
} else {
106+
SseEvent::Delta(delta.to_string())
107+
}
108+
}
109+
110+
/// 一次 LLM 调用的完整轨迹(提示词 + 回答全文)。
111+
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
112+
#[serde(rename_all = "camelCase")]
113+
pub struct AiTurn {
114+
/// 片序(1 起;失败片无轨迹——任务卡已显失败片数)
115+
pub turn: usize,
116+
/// 组装后的 system 提示词(模板构建结果,非模板原文)
117+
pub system: String,
118+
/// user 请求文本(精修=AiRefineRequest JSON;vision 附图数占位)
119+
pub user: String,
120+
/// 模型回答(结构化响应 JSON;原始返回即 JSON)
121+
pub response: String,
122+
}
123+
124+
/// 轨迹 → JSON(落库 trajectory_json 列;失败返回 None 由调用方降级不写)。
125+
pub fn trajectory_to_json(turns: &[AiTurn]) -> Option<String> {
126+
serde_json::to_string(turns).ok()
127+
}
128+
129+
/// JSON → 轨迹(旧任务/损坏数据 → None,视图诚实提示"无轨迹存档")。
130+
pub fn trajectory_from_json(s: &str) -> Option<Vec<AiTurn>> {
131+
serde_json::from_str(s).ok()
132+
}
133+
134+
/// 流取消标志(Arc 共享;chat_cancel 置位 → 流循环每 chunk 检查)。
135+
#[derive(Debug, Default, Clone)]
136+
pub struct CancelFlag(Arc<AtomicBool>);
137+
138+
impl CancelFlag {
139+
pub fn new() -> Self {
140+
Self::default()
141+
}
142+
143+
pub fn cancel(&self) {
144+
self.0.store(true, Ordering::Relaxed);
145+
}
146+
147+
pub fn is_cancelled(&self) -> bool {
148+
self.0.load(Ordering::Relaxed)
149+
}
150+
}
151+
152+
#[cfg(test)]
153+
#[path = "ai_chat_tests.rs"]
154+
mod tests;
Lines changed: 170 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,170 @@
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

Comments
 (0)