188 lines
8.2 KiB
Rust
188 lines
8.2 KiB
Rust
//! 原生 SSE 流式解析器 — 替代 eventsource-stream 库
|
|
//!
|
|
//! eventsource-stream 0.2 在 Windows 上对 Deepseek 等 provider
|
|
//! 的 SSE 响应解析时报 "Transport error: error decoding response body" 错误。
|
|
//!
|
|
//! 根因分析:
|
|
//! eventsource-stream 内部对 bytes_stream 做严格的 UTF-8 + SSE 协议校验,遇到以下情况
|
|
//! 即报错(且不可恢复):
|
|
//! - 流中断时未完整接收 UTF-8 字符(网络抖动常见)
|
|
//! - 缺少结束的 \n\n(连接断开常见)
|
|
//! - 非 ASCII 字符的多字节序列跨 chunk 边界
|
|
//!
|
|
//! 本解析器实现:
|
|
//! - 宽松的 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;
|
|
use std::pin::Pin;
|
|
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,
|
|
buffer: Vec<u8>,
|
|
}
|
|
|
|
impl<S> SseStream<S>
|
|
where
|
|
S: Stream<Item = Result<bytes::Bytes, reqwest::Error>> + Unpin,
|
|
{
|
|
pub fn new(inner: S) -> Self {
|
|
Self {
|
|
inner,
|
|
buffer: Vec::with_capacity(8192),
|
|
}
|
|
}
|
|
|
|
/// 从 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 = 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 data = Self::extract_data_fields(&event_text);
|
|
if !data.is_empty() {
|
|
events.push(data);
|
|
}
|
|
}
|
|
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();
|
|
for line in event_text.lines() {
|
|
if let Some(rest) = line.strip_prefix("data:") {
|
|
let rest = rest.strip_prefix(' ').unwrap_or(rest);
|
|
data_parts.push(rest);
|
|
}
|
|
// 忽略 event:/id:/retry: 等其他 SSE 字段(OpenAI/Anthropic 协议未使用)
|
|
}
|
|
data_parts.join("\n")
|
|
}
|
|
}
|
|
|
|
impl<S> Stream for SseStream<S>
|
|
where
|
|
S: Stream<Item = Result<bytes::Bytes, reqwest::Error>> + Unpin,
|
|
{
|
|
type Item = Result<Vec<SseEvent>, String>;
|
|
|
|
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
|
use futures::StreamExt;
|
|
loop {
|
|
// 先尝试从 buffer 解析完整事件
|
|
let events = self.parse_events();
|
|
if !events.is_empty() {
|
|
return Poll::Ready(Some(Ok(events)));
|
|
}
|
|
|
|
// buffer 不足以解析出完整事件,从 inner 读更多数据
|
|
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))) => {
|
|
return Poll::Ready(Some(Err(format!("SSE 流读取错误: {}", e))));
|
|
}
|
|
Poll::Ready(None) => {
|
|
// 流结束,处理 buffer 中的剩余数据(可能没有 \n\n 结束的最后一段)
|
|
if !self.buffer.is_empty() {
|
|
// 解码:优先严格 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])));
|
|
}
|
|
}
|
|
return Poll::Ready(None);
|
|
}
|
|
Poll::Pending => return Poll::Pending,
|
|
}
|
|
}
|
|
}
|
|
}
|