修复: checkpoint 参数化 + 压缩 permit 死锁降级 + 密钥归一化 + SSE CRLF/多字节

This commit is contained in:
lxy
2026-08-01 11:30:05 +08:00
parent c4ba920cf5
commit 20ec571dbb
9 changed files with 232 additions and 42 deletions
+82 -13
View File
@@ -11,9 +11,13 @@
//! - 非 ASCII 字符的多字节序列跨 chunk 边界
//!
//! 本解析器实现:
//! - 宽松的 UTF-8 处理(用 bytes 累积,String::from_utf8_lossy 转换,不报错)
//! - SSE 协议简单解析(以 \n\n 分隔事件,data: 前缀提取)
//! - 容错:解析失败时跳过该事件继续,不中断流
//! - 宽松的 UTF-8 处理(用 bytes 累积,优先 `String::from_utf8` 严格解码保留多字节完整,
//! 失败再降级 `from_utf8_lossy` 不报错)
//! - SSE 协议简单解析(分隔符兼容 `\n\n` / `\r\n\r\n` / `\r\r` 三种行尾归一,
//! data: 前缀提取)
//! - 多字节续接:事件体只在定位到完整分隔符后才解码,跨 chunk 的多字节字符在
//! buffer 中天然拼接复原(分隔符为 ASCII,不落在多字节序列中间)
//! - 容错:解析失败时跳过该事件继续,不中断流;buffer 设 `BUF_MAX` 上限防 OOM
//! - 返回 Vec<String>(每个元素是一个事件 data 字段拼接内容)
use futures::Stream;
@@ -23,6 +27,10 @@ use std::task::{Context, Poll};
/// SSE 事件流的 data 字段内容
pub type SseEvent = String;
/// 缓冲区字节上限(防御性):正常流会被 parse_events 持续消费,
/// 仅畸形上游(持续不发分隔符)时触发裁剪,避免无界增长 OOM。
const BUF_MAX: usize = 1 << 20; // 1 MiB
/// 原生 SSE 解析器流:包装 bytes_stream,产出 Vec<SseEvent>(一次 poll 可能产出多个事件)
pub struct SseStream<S> {
inner: S,
@@ -40,19 +48,42 @@ where
}
}
/// 从 buffer 解析完整的 SSE 事件(以 \n\n 分隔),返回事件列表
/// 从 buffer 解析完整的 SSE 事件,返回事件列表
///
/// 事件分隔符兼容三种行尾归一:
/// - `\n\n` (LF LF,规范形态)
/// - `\r\n\r\n` (CRLF CRLF,部分中转站/反代发送)
/// - `\r\r` (CR CR,极少见但同源)
///
/// 取最早出现的分隔符;事件体内部的 `\r\n` 在解码后归一为 `\n`,
/// 确保 `lines()` / `strip_prefix` 等行级解析正常工作。
///
/// 多字节续接:事件体只在定位到完整分隔符后才解码,跨 chunk 的多字节
/// 字符在 buffer 中天然拼接复原(分隔符为 ASCII,不会落在多字节序列中间)。
/// 解码用 `String::from_utf8`(失败再降级 lossy),避免把合法多字节误判为残缺。
fn parse_events(&mut self) -> Vec<SseEvent> {
let mut events = Vec::new();
loop {
let sep_pos = self.buffer.windows(2).position(|w| w == b"\n\n");
if sep_pos.is_none() {
break;
// 同时识别三种分隔符,取最早出现的位置与该分隔符的字节长度
let sep = Self::find_event_boundary(&self.buffer);
let (sep_pos, sep_len) = match sep {
Some(v) => v,
None => break,
};
// 取出事件体 + 分隔符整体(buffer 前 sep_pos+sep_len 字节)
let event_bytes: Vec<u8> = self.buffer.drain(..sep_pos + sep_len).collect();
// 去掉末尾的分隔符,得到事件体
let body_end = event_bytes.len().saturating_sub(sep_len);
let body_bytes = &event_bytes[..body_end];
// 解码:优先严格 UTF-8(保留多字节完整),失败再降级 lossy(不崩)
let mut event_text = match std::str::from_utf8(body_bytes) {
Ok(s) => s.to_owned(),
Err(_) => String::from_utf8_lossy(body_bytes).into_owned(),
};
// 事件体内 CRLF / 裸 CR 归一为 LF,保证后续行级解析一致
if event_text.contains('\r') {
event_text = event_text.replace("\r\n", "\n").replace('\r', "\n");
}
let sep_pos = sep_pos.unwrap();
let event_bytes: Vec<u8> = self.buffer.drain(..sep_pos + 2).collect();
// 去掉末尾的 \n\n
let body_end = event_bytes.len().saturating_sub(2);
let event_text = String::from_utf8_lossy(&event_bytes[..body_end]);
let data = Self::extract_data_fields(&event_text);
if !data.is_empty() {
events.push(data);
@@ -61,6 +92,29 @@ where
events
}
/// 在 buffer 中查找最早的事件分隔符,返回 `(起始位置, 分隔符字节长度)`。
/// 扫描 `\r\n\r\n` / `\n\n` / `\r\r` 三种形态,取最小起始位置(取最早分隔符)。
fn find_event_boundary(buf: &[u8]) -> Option<(usize, usize)> {
// candidates: (pattern, length)
const PATTERNS: &[(&[u8], usize)] = &[
(b"\r\n\r\n", 4),
(b"\n\n", 2),
(b"\r\r", 2),
];
let mut best: Option<(usize, usize)> = None;
for &(pat, len) in PATTERNS {
// windows 匹配;找到首次出现位置
if let Some(pos) = buf.windows(pat.len()).position(|w| w == pat) {
match best {
None => best = Some((pos, len)),
Some((bpos, _)) if pos < bpos => best = Some((pos, len)),
_ => {}
}
}
}
best
}
/// 从 SSE 事件文本中提取所有 data: 行的内容,拼接为单个字符串(多个 data 行用 \n 连接)
fn extract_data_fields(event_text: &str) -> String {
let mut data_parts: Vec<&str> = Vec::new();
@@ -94,6 +148,12 @@ where
match self.inner.poll_next_unpin(cx) {
Poll::Ready(Some(Ok(chunk))) => {
self.buffer.extend_from_slice(&chunk);
// 防御:缓冲区不应无限增长。正常情况下 parse_events 会持续消费,
// 仅当上游持续不发分隔符(畸形流)时才触发,此处裁掉头部旧数据避免 OOM。
if self.buffer.len() > BUF_MAX {
let drop_n = self.buffer.len() - BUF_MAX;
self.buffer.drain(..drop_n);
}
continue;
}
Poll::Ready(Some(Err(e))) => {
@@ -102,8 +162,17 @@ where
Poll::Ready(None) => {
// 流结束,处理 buffer 中的剩余数据(可能没有 \n\n 结束的最后一段)
if !self.buffer.is_empty() {
let remaining = String::from_utf8_lossy(&self.buffer).to_string();
// 解码:优先严格 UTF-8(EOF 无下一 chunk 可拼接,
// 残缺多字节尾部只能尽力而为,降级 lossy 不崩)
let mut remaining = match String::from_utf8(self.buffer.clone()) {
Ok(s) => s,
Err(_) => String::from_utf8_lossy(&self.buffer).into_owned(),
};
self.buffer.clear();
// 同主解析路径:CRLF / 裸 CR 归一为 LF
if remaining.contains('\r') {
remaining = remaining.replace("\r\n", "\n").replace('\r', "\n");
}
let data = Self::extract_data_fields(&remaining);
if !data.is_empty() {
return Poll::Ready(Some(Ok(vec![data])));
+2 -3
View File
@@ -118,8 +118,7 @@ pub(crate) async fn resolve_provider(
老路径将在后续版本移除"
);
let base_url = base_url_str.to_string();
let api_key = api_key_str.to_string();
ensure_resolved_key("(明文注入)", &api_key)
let api_key = ensure_resolved_key("(明文注入)", api_key_str)
.map_err(anyhow::Error::msg)?;
let protocol = config
.get("protocol")
@@ -162,7 +161,7 @@ pub(crate) fn resolve_from_record(
config_model: &str,
) -> anyhow::Result<ResolvedProvider> {
let api_key = resolve_provider_secret(record);
ensure_resolved_key(&record.name, &api_key).map_err(anyhow::Error::msg)?;
let api_key = ensure_resolved_key(&record.name, &api_key).map_err(anyhow::Error::msg)?;
let default_model = if !config_model.is_empty() {
config_model.to_string()
} else if !record.default_model.is_empty() {
@@ -323,6 +323,40 @@ impl AiToolExecutionRepo {
// AiConversationRepo 的整体更新已由 impl_repo! 宏统一生成的 update_full 提供。
impl AiConversationRepo {
/// 写入对话版本化快照 checkpoint(INSERT OR IGNORE,同 id 已存在则跳过)。
///
/// 参数化绑定(替代原调用方的 format! 拼 SQL + execute_batch),防 snapshot 含引号/
/// 特殊字符致注入或损坏。列对齐 conversation_checkpoints(id, conv_id, snapshot,
/// token_total, created_at)。
pub async fn insert_checkpoint(
&self,
id: &str,
conv_id: &str,
snapshot: &str,
token_total: i64,
created_at: &str,
) -> Result<()> {
let conn = self.conn.clone();
let id = id.to_string();
let conv_id = conv_id.to_string();
let snapshot = snapshot.to_string();
let created_at = created_at.to_string();
let _ = tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
guard
.execute(
"INSERT OR IGNORE INTO conversation_checkpoints \
(id, conv_id, snapshot, token_total, created_at) \
VALUES (?1, ?2, ?3, ?4, ?5)",
params![id, conv_id, snapshot, token_total, created_at],
)
.map_err(storage_err)
})
.await
.map_err(storage_err)??;
Ok(())
}
/// 清空对话消息内容(保留 conversation 记录本身,只清 messages JSON + 清零 token 计数)
///
/// "清空对话"语义:对话壳保留(侧栏仍可见,可继续在该对话内聊),仅清空历史消息。
+52 -7
View File
@@ -200,17 +200,40 @@ pub async fn migrate_secrets_to_keyring(repo: &AiProviderRepo) -> anyhow::Result
Ok(migrated)
}
/// 校验已解析的密钥是否可用:空(含纯空白)→明确错误信息,非空→Ok
/// 校验已解析的密钥是否可用并归一化:空(含纯空白)→明确错误信息;非空→返回归一化后的 String
///
/// 归一化 = trim → 剥首尾配对引号(`"`/`'`)→ 再 trim。覆盖用户粘贴脏 key 的常见场景:
/// 复制带前后引号/换行/全角空格/尾部空白,校验能过但原样发 provider 致 401 误报"API Key 无效"。
/// 调用方应用返回值(归一化后的 key)替代原 resolved 传给 provider。
///
/// 用于消费点(build_provider 前)早失败,避免空 key 发请求吃 401,错误伪装成"API Key 无效"。
pub fn ensure_resolved_key(provider_name: &str, resolved: &str) -> Result<(), String> {
pub fn ensure_resolved_key(provider_name: &str, resolved: &str) -> Result<String, String> {
// 先 trim 判空(纯空白视为无密钥,防粘贴时只有空格)。
if resolved.trim().is_empty() {
Err(format!(
return Err(format!(
"未读取到「{}」的 API 密钥(系统钥匙串无记录或已损坏),请在设置中重新填写并保存",
provider_name
))
} else {
Ok(())
));
}
// 归一化:trim → 剥首尾引号(成对,支持 ""xxx"" / ''xxx'' 多层)→ 再 trim。
let mut cleaned = resolved.trim().to_string();
while cleaned.len() >= 2 {
let first = cleaned.chars().next().unwrap();
let last = cleaned.chars().last().unwrap();
if (first == '"' || first == '\'') && first == last {
cleaned = cleaned[first.len_utf8()..cleaned.len() - last.len_utf8()].trim().to_string();
} else {
break;
}
}
// 剥引号后可能变空(如粘贴仅一对引号)→ 视为无密钥。
if cleaned.is_empty() {
return Err(format!(
"未读取到「{}」的 API 密钥(系统钥匙串无记录或已损坏),请在设置中重新填写并保存",
provider_name
));
}
Ok(cleaned)
}
#[cfg(test)]
@@ -230,7 +253,7 @@ mod tests {
#[test]
fn ensure_resolved_key_accepts_nonempty() {
assert!(ensure_resolved_key("GLM", "sk-abc").is_ok());
assert_eq!(ensure_resolved_key("GLM", "sk-abc").unwrap(), "sk-abc");
}
#[test]
@@ -239,6 +262,28 @@ mod tests {
assert!(err.contains("我的提供商"), "错误信息应含 provider 名便于定位");
}
#[test]
fn ensure_resolved_key_trims_surrounding_whitespace() {
// 粘贴带前后空白/换行 → 归一化为干净 key
assert_eq!(ensure_resolved_key("GLM", " sk-abc \n").unwrap(), "sk-abc");
}
#[test]
fn ensure_resolved_key_strips_surrounding_quotes() {
// 粘贴带前后引号(双引号/单引号,多层嵌套)→ 剥引号
assert_eq!(ensure_resolved_key("GLM", "\"sk-abc\"").unwrap(), "sk-abc");
assert_eq!(ensure_resolved_key("GLM", "'sk-abc'").unwrap(), "sk-abc");
assert_eq!(ensure_resolved_key("GLM", "\"\"sk-abc\"\"").unwrap(), "sk-abc");
assert_eq!(ensure_resolved_key("GLM", " \"sk-abc\"\n").unwrap(), "sk-abc");
}
#[test]
fn ensure_resolved_key_rejects_only_quotes() {
// 仅一对引号(剥后为空)→ 视为无密钥
assert!(ensure_resolved_key("GLM", "\"\"").is_err());
assert!(ensure_resolved_key("GLM", " '' ").is_err());
}
#[test]
fn resolve_prefers_db_when_non_empty() {
// DB api_key 非空 → 直接返回 DB 值,不触发 keyring(兼容未迁移老库)