Skip to main content

rightkit_http/
sse.rs

1//! Server-Sent Events decoder (WHATWG EventSource framing): LF, CR and CRLF line
2//! endings, multi-line `data`, comments, `id`, `retry`, BOM, bounded line size.
3
4use std::collections::VecDeque;
5use std::io::Read;
6
7use crate::error::HttpError;
8
9#[derive(Debug, Clone, Default, PartialEq, Eq)]
10pub struct SseEvent {
11    /// Event type; "message" when the stream gave none.
12    pub event: String,
13    pub data: String,
14    pub id: Option<String>,
15    pub retry_ms: Option<u64>,
16}
17
18/// Push-based decoder: feed bytes, receive completed events.
19#[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    /// Decoder bounding each line; events are unbounded (use [`Self::with_limits`]).
42    pub fn new(max_line_bytes: usize) -> Self {
43        Self::with_limits(max_line_bytes, usize::MAX)
44    }
45
46    /// Decoder bounding each line and the accumulated bytes of one event
47    /// (event name plus all data lines, one separator byte per line).
48    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; // second half of CRLF
71                }
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
158/// Blocking iterator of events over any reader. An event cut off by EOF is
159/// dropped, per the EventSource spec.
160pub 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}