1use std::collections::VecDeque;
5use std::io::Read;
6
7use crate::error::HttpError;
8
9#[derive(Debug, Clone, Default, PartialEq, Eq)]
10pub struct SseEvent {
11 pub event: String,
13 pub data: String,
14 pub id: Option<String>,
15 pub retry_ms: Option<u64>,
16}
17
18#[derive(Debug)]
20pub struct SseDecoder {
21 line: Vec<u8>,
22 prev_cr: bool,
23 first: bool,
24 event: String,
25 data: Vec<String>,
26 has_data: bool,
27 id: Option<String>,
28 retry_ms: Option<u64>,
29 max_line: usize,
30 max_event: usize,
31 event_bytes: usize,
32}
33
34impl Default for SseDecoder {
35 fn default() -> Self {
36 Self::new(1024 * 1024)
37 }
38}
39
40impl SseDecoder {
41 pub fn new(max_line_bytes: usize) -> Self {
43 Self::with_limits(max_line_bytes, usize::MAX)
44 }
45
46 pub fn with_limits(max_line_bytes: usize, max_event_bytes: usize) -> Self {
49 Self {
50 line: Vec::new(),
51 prev_cr: false,
52 first: true,
53 event: String::new(),
54 data: Vec::new(),
55 has_data: false,
56 id: None,
57 retry_ms: None,
58 max_line: max_line_bytes,
59 max_event: max_event_bytes,
60 event_bytes: 0,
61 }
62 }
63
64 pub fn push(&mut self, bytes: &[u8]) -> Result<Vec<SseEvent>, HttpError> {
65 let mut out = Vec::new();
66 for &b in bytes {
67 if self.prev_cr {
68 self.prev_cr = false;
69 if b == b'\n' {
70 continue; }
72 }
73 match b {
74 b'\n' | b'\r' => {
75 self.prev_cr = b == b'\r';
76 let line = std::mem::take(&mut self.line);
77 if let Some(ev) = self.handle_line(&line)? {
78 out.push(ev);
79 }
80 }
81 _ => {
82 self.line.push(b);
83 if self.line.len() > self.max_line {
84 return Err(HttpError::Stream("SSE line exceeds limit".into()));
85 }
86 }
87 }
88 }
89 Ok(out)
90 }
91
92 fn handle_line(&mut self, raw: &[u8]) -> Result<Option<SseEvent>, HttpError> {
93 let mut line = raw;
94 if self.first {
95 self.first = false;
96 if line.starts_with(&[0xEF, 0xBB, 0xBF]) {
97 line = &line[3..];
98 }
99 }
100 if line.is_empty() {
101 return Ok(self.dispatch());
102 }
103 let text = std::str::from_utf8(line)
104 .map_err(|_| HttpError::Stream("SSE line is not valid UTF-8".into()))?;
105 if text.starts_with(':') {
106 return Ok(None);
107 }
108 let (field, value) = match text.split_once(':') {
109 Some((f, v)) => (f, v.strip_prefix(' ').unwrap_or(v)),
110 None => (text, ""),
111 };
112 if matches!(field, "event" | "data") {
113 self.event_bytes = self.event_bytes.saturating_add(value.len() + 1);
114 if self.event_bytes > self.max_event {
115 return Err(HttpError::Stream("SSE event exceeds limit".into()));
116 }
117 }
118 match field {
119 "event" => self.event = value.to_string(),
120 "data" => {
121 self.data.push(value.to_string());
122 self.has_data = true;
123 }
124 "id" if !value.contains('\0') => self.id = Some(value.to_string()),
125 "retry" => {
126 if let Ok(ms) = value.parse::<u64>() {
127 self.retry_ms = Some(ms);
128 }
129 }
130 _ => {}
131 }
132 Ok(None)
133 }
134
135 fn dispatch(&mut self) -> Option<SseEvent> {
136 let had = self.has_data;
137 self.event_bytes = 0;
138 let event = std::mem::take(&mut self.event);
139 let data = std::mem::take(&mut self.data);
140 self.has_data = false;
141 let retry_ms = self.retry_ms.take();
142 if !had {
143 return None;
144 }
145 Some(SseEvent {
146 event: if event.is_empty() {
147 "message".into()
148 } else {
149 event
150 },
151 data: data.join("\n"),
152 id: self.id.clone(),
153 retry_ms,
154 })
155 }
156}
157
158pub struct SseReader<R: Read> {
161 reader: R,
162 decoder: SseDecoder,
163 queue: VecDeque<SseEvent>,
164 done: bool,
165}
166
167impl<R: Read> SseReader<R> {
168 pub fn new(reader: R) -> Self {
169 Self {
170 reader,
171 decoder: SseDecoder::default(),
172 queue: VecDeque::new(),
173 done: false,
174 }
175 }
176}
177
178impl<R: Read> Iterator for SseReader<R> {
179 type Item = Result<SseEvent, HttpError>;
180 fn next(&mut self) -> Option<Self::Item> {
181 let mut buf = [0u8; 8192];
182 loop {
183 if let Some(ev) = self.queue.pop_front() {
184 return Some(Ok(ev));
185 }
186 if self.done {
187 return None;
188 }
189 match self.reader.read(&mut buf) {
190 Ok(0) => {
191 self.done = true;
192 }
193 Ok(n) => match self.decoder.push(&buf[..n]) {
194 Ok(events) => self.queue.extend(events),
195 Err(e) => {
196 self.done = true;
197 return Some(Err(e));
198 }
199 },
200 Err(e) if e.kind() == std::io::ErrorKind::Interrupted => {}
201 Err(e) => {
202 self.done = true;
203 return Some(Err(match e.kind() {
204 std::io::ErrorKind::TimedOut => HttpError::Timeout(e.to_string()),
205 _ => HttpError::Transport {
206 message: e.to_string(),
207 connect_phase: false,
208 },
209 }));
210 }
211 }
212 }
213 }
214}