Skip to main content

gateway_core/
stream.rs

1#[derive(Debug, Clone, PartialEq, Eq)]
2pub struct SseEvent {
3    pub event: Option<String>,
4    pub data: String,
5}
6
7#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
8pub enum StreamParseError {
9    #[error("SSE buffer exceeded {0} bytes")]
10    BufferLimit(usize),
11    #[error("stream ended with an incomplete SSE event")]
12    Incomplete,
13}
14
15pub struct SseDecoder {
16    buffer: String,
17    max_buffer_bytes: usize,
18}
19
20impl Default for SseDecoder {
21    fn default() -> Self {
22        Self::new(1024 * 1024)
23    }
24}
25
26impl SseDecoder {
27    pub fn new(max_buffer_bytes: usize) -> Self {
28        Self {
29            buffer: String::new(),
30            max_buffer_bytes,
31        }
32    }
33
34    pub fn push(&mut self, chunk: &str) -> Result<Vec<SseEvent>, StreamParseError> {
35        self.buffer.push_str(chunk);
36        if self.buffer.len() > self.max_buffer_bytes {
37            return Err(StreamParseError::BufferLimit(self.max_buffer_bytes));
38        }
39        let mut events = Vec::new();
40        while let Some(end) = event_end(&self.buffer) {
41            let block = self.buffer[..end].replace('\r', "");
42            let delimiter_len = if self.buffer[end..].starts_with("\r\n\r\n") {
43                4
44            } else {
45                2
46            };
47            self.buffer.drain(..end + delimiter_len);
48            if let Some(event) = parse_event(&block) {
49                events.push(event);
50            }
51        }
52        Ok(events)
53    }
54
55    pub fn finish(self) -> Result<(), StreamParseError> {
56        if self.buffer.trim().is_empty() {
57            Ok(())
58        } else {
59            Err(StreamParseError::Incomplete)
60        }
61    }
62}
63
64fn event_end(buffer: &str) -> Option<usize> {
65    match (buffer.find("\n\n"), buffer.find("\r\n\r\n")) {
66        (Some(a), Some(b)) => Some(a.min(b)),
67        (Some(a), None) => Some(a),
68        (None, Some(b)) => Some(b),
69        (None, None) => None,
70    }
71}
72
73fn parse_event(block: &str) -> Option<SseEvent> {
74    let mut event = None;
75    let mut data = Vec::new();
76    for line in block.lines() {
77        if line.starts_with(':') {
78            continue;
79        }
80        if let Some(value) = line.strip_prefix("event:") {
81            event = Some(value.trim_start().to_owned());
82        } else if let Some(value) = line.strip_prefix("data:") {
83            data.push(value.trim_start());
84        }
85    }
86    (!data.is_empty()).then(|| SseEvent {
87        event,
88        data: data.join("\n"),
89    })
90}
91
92#[cfg(test)]
93mod tests {
94    use super::*;
95
96    #[test]
97    fn parses_fragmented_and_multiline_events() {
98        let mut decoder = SseDecoder::default();
99        assert!(
100            decoder
101                .push("event: delta\r\ndata: {\"a\":")
102                .unwrap()
103                .is_empty()
104        );
105        let events = decoder.push("1}\r\ndata: tail\r\n\r\n").unwrap();
106        assert_eq!(events[0].event.as_deref(), Some("delta"));
107        assert_eq!(events[0].data, "{\"a\":1}\ntail");
108        decoder.finish().unwrap();
109    }
110}