Skip to main content

ai_cortex_sdk/
sse.rs

1//! 服务端推送的更新通知 SSE 订阅器。
2//!
3//! 服务端通过 `GET /api/softwares/sdk-update-events`(text/event-stream,PAT 鉴权)下发更新事件。
4//! 帧格式:
5//! - update 事件:`id: <u64>\nevent: update\ndata: <UpdateEvent JSON>\n\n`
6//! - keepalive 注释:`: keepalive\n\n`(忽略)
7//! - lagged 注释:`: lagged\n\n`(忽略)
8//!
9//! 支持可选查询参数 `?software_id=<id>`(仅转发该软件的事件)和可选请求头
10//! `Last-Event-ID: <u64>`(断线重连,服务端补播 id > last 的事件)。
11
12use std::collections::VecDeque;
13use std::pin::Pin;
14use std::task::{Context, Poll};
15
16use futures_util::Stream;
17
18use crate::error::{SdkError, SdkResult};
19use crate::models::UpdateEvent;
20
21/// 底层字节 chunk 流:每个 item 是一个 chunk 的字节(已从 `bytes::Bytes` 转为 `Vec<u8>`,
22/// 避免在本 crate 显式依赖 bytes)。
23type ByteChunkStream = Pin<Box<dyn Stream<Item = Result<Vec<u8>, reqwest::Error>> + Send>>;
24
25/// SSE 更新事件流。实现 [`futures_util::Stream`],逐个产出解析后的 [`UpdateEvent`]。
26///
27/// 用法:
28/// ```ignore
29/// use futures_util::StreamExt;
30/// let mut stream = client.open_update_events(None, None).await?;
31/// while let Some(ev) = stream.next().await {
32///     match ev {
33///         Ok(e) => println!("update: {:?}", e),
34///         Err(e) => eprintln!("stream error: {e}"),
35///     }
36/// }
37/// ```
38pub struct UpdateEventStream {
39    /// 底层 chunk 流(已 pin + type-erase)
40    inner: ByteChunkStream,
41    /// 尚未凑成完整 dispatch 的字节(lossy UTF-8 累积)
42    line_buf: String,
43    /// 已解析但尚未 yield 的事件队列
44    pending: VecDeque<UpdateEvent>,
45}
46
47impl UpdateEventStream {
48    /// 从一个 chunk 流构造。`stream` 的 item 为 `Result<Vec<u8>, reqwest::Error>`。
49    pub fn new<S>(stream: S) -> Self
50    where
51        S: Stream<Item = Result<Vec<u8>, reqwest::Error>> + Send + 'static,
52    {
53        Self {
54            inner: Box::pin(stream),
55            line_buf: String::new(),
56            pending: VecDeque::new(),
57        }
58    }
59
60    /// 尝试从 buffer 中切出所有完整 dispatch(以 `\n\n` 分隔),逐个解析并入队。
61    fn drain_complete_dispatches(&mut self) {
62        // 反复切出第一个 `\n\n` 之前的内容(含分隔符)
63        while let Some(idx) = self.line_buf.find("\n\n") {
64            // 取出 dispatch(含末尾 `\n\n`),剩余留在 buffer
65            let dispatch: String = self.line_buf.drain(..idx + 2).collect();
66            if let Some(ev) = parse_dispatch(&dispatch) {
67                self.pending.push_back(ev);
68            }
69        }
70    }
71}
72
73impl Stream for UpdateEventStream {
74    type Item = SdkResult<UpdateEvent>;
75
76    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
77        // 所有字段均为 Unpin(Pin<Box<..>>/String/VecDeque 都是 Unpin),
78        // 可安全 get_mut
79        let this = self.get_mut();
80
81        loop {
82            // 1. 优先吐出已解析的事件
83            if let Some(ev) = this.pending.pop_front() {
84                return Poll::Ready(Some(Ok(ev)));
85            }
86
87            // 2. 从底层流拉取下一个 chunk
88            match this.inner.as_mut().poll_next(cx) {
89                Poll::Ready(Some(Ok(chunk))) => {
90                    // lossy UTF-8 追加(SSE 帧边界不会切在多字节字符中间,lossy 足够)
91                    this.line_buf.push_str(&String::from_utf8_lossy(&chunk));
92                    // 切出所有完整 dispatch 并入队
93                    this.drain_complete_dispatches();
94                    // 循环:若有 pending 则下一轮吐出,否则继续拉取
95                }
96                Poll::Ready(Some(Err(e))) => {
97                    return Poll::Ready(Some(Err(SdkError::RequestError(e))));
98                }
99                Poll::Ready(None) => {
100                    // 连接关闭:尝试解析 buffer 中残留的最后一段(无尾随 `\n\n`)
101                    if !this.line_buf.is_empty() {
102                        let rest: String = this.line_buf.drain(..).collect();
103                        if let Some(ev) = parse_dispatch(&rest) {
104                            this.pending.push_back(ev);
105                        }
106                    }
107                    if let Some(ev) = this.pending.pop_front() {
108                        return Poll::Ready(Some(Ok(ev)));
109                    }
110                    return Poll::Ready(None);
111                }
112                Poll::Pending => {
113                    // 底层未就绪:若有 pending 则吐出,否则挂起
114                    if let Some(ev) = this.pending.pop_front() {
115                        return Poll::Ready(Some(Ok(ev)));
116                    }
117                    return Poll::Pending;
118                }
119            }
120        }
121    }
122}
123
124/// 解析单个 SSE dispatch(可能包含 `id:`/`event:`/`data:` 行及 `:` 注释行)。
125///
126/// 规则:
127/// - `:` 开头为注释行(keepalive/lagged),忽略;
128/// - 多个 `data:` 行按 SSE 规范以 `\n` 拼接;
129/// - `id:` 行解析为 u64,覆盖事件 id(用于 Last-Event-ID 重连);
130/// - `event:` 行记录类型(当前不按类型过滤,data 能反序列化即转发);
131/// - data 为空或反序列化失败则跳过该 dispatch(不产出事件)。
132fn parse_dispatch(dispatch: &str) -> Option<UpdateEvent> {
133    let mut data_lines: Vec<String> = Vec::new();
134    let mut id: Option<u64> = None;
135    let mut _event_type: Option<String> = None;
136
137    for line in dispatch.lines() {
138        if line.is_empty() {
139            continue;
140        }
141        // 注释行(: keepalive / : lagged 等)
142        if let Some(_comment) = line.strip_prefix(':') {
143            continue;
144        }
145        if let Some(rest) = line.strip_prefix("data:") {
146            // 去掉一个可选前导空格(SSE 规范:`data: xxx` 与 `data:xxx` 都合法)
147            let rest = rest.strip_prefix(' ').unwrap_or(rest);
148            data_lines.push(rest.to_string());
149        } else if let Some(rest) = line.strip_prefix("id:") {
150            let rest = rest.strip_prefix(' ').unwrap_or(rest);
151            if let Ok(parsed) = rest.trim().parse::<u64>() {
152                id = Some(parsed);
153            }
154        } else if let Some(rest) = line.strip_prefix("event:") {
155            let rest = rest.strip_prefix(' ').unwrap_or(rest);
156            _event_type = Some(rest.to_string());
157        }
158        // 其余字段(如 retry:)忽略
159    }
160
161    if data_lines.is_empty() {
162        return None;
163    }
164
165    let data = data_lines.join("\n");
166    match serde_json::from_str::<UpdateEvent>(&data) {
167        Ok(mut ev) => {
168            // SSE 帧 id: 行优先于 JSON 内 id(id: 是 Last-Event-ID 的权威来源)
169            if let Some(id) = id {
170                ev.id = id;
171            }
172            Some(ev)
173        }
174        // 反序列化失败:跳过畸形帧(按 spec 选择 skip-with-warn,这里静默跳过)
175        Err(_) => None,
176    }
177}
178
179#[cfg(test)]
180mod tests {
181    use super::*;
182
183    #[test]
184    fn parse_single_update_event() {
185        let frame = "id: 42\nevent: update\ndata: {\"software_id\":\"s1\",\"version\":\"1.2.0\",\"platform\":\"linux\",\"channel\":\"stable\",\"force_update\":false}\n\n";
186        let ev = parse_dispatch(frame).expect("should parse");
187        assert_eq!(ev.id, 42);
188        assert_eq!(ev.software_id, "s1");
189        assert_eq!(ev.version, "1.2.0");
190        assert_eq!(ev.platform, "linux");
191        assert_eq!(ev.channel, "stable");
192        assert!(!ev.force_update);
193    }
194
195    #[test]
196    fn parse_ignores_comments() {
197        // keepalive / lagged 注释应被忽略,不产出事件
198        assert!(parse_dispatch(": keepalive\n\n").is_none());
199        assert!(parse_dispatch(": lagged\n\n").is_none());
200    }
201
202    #[test]
203    fn parse_multi_line_data() {
204        let frame = "data: {\"software_id\":\"s2\",\ndata: \"version\":\"2.0\",\"platform\":\"win\",\"channel\":\"beta\",\"force_update\":true}\n\n";
205        let ev = parse_dispatch(frame).expect("should parse");
206        assert_eq!(ev.software_id, "s2");
207        assert_eq!(ev.version, "2.0");
208        assert!(ev.force_update);
209    }
210
211    #[test]
212    fn parse_skips_malformed() {
213        // data 不是合法 JSON → 跳过(返回 None)
214        assert!(parse_dispatch("data: not-json\n\n").is_none());
215    }
216
217    #[test]
218    fn parse_id_overrides_json() {
219        // SSE id: 行优先于 JSON 内的 id
220        let frame = "id: 99\ndata: {\"id\":5,\"software_id\":\"s\",\"version\":\"1\",\"platform\":\"p\",\"channel\":\"c\",\"force_update\":false}\n\n";
221        let ev = parse_dispatch(frame).expect("should parse");
222        assert_eq!(ev.id, 99);
223    }
224}