Skip to main content

gproxy_transform/envelope/
sse.rs

1use bytes::Bytes;
2
3use crate::TransformError;
4
5#[derive(Debug)]
6pub struct SseFrame {
7    pub event: Option<String>,
8    pub data: String,
9}
10
11impl SseFrame {
12    pub fn typed<T: serde::Serialize>(
13        event: Option<&str>,
14        value: &T,
15    ) -> Result<Bytes, TransformError> {
16        Ok(Self::encode(event, &serde_json::to_string(value)?))
17    }
18
19    pub fn encode(event: Option<&str>, data: &str) -> Bytes {
20        let mut output = String::new();
21        if let Some(event) = event {
22            output.push_str("event: ");
23            output.push_str(event);
24            output.push('\n');
25        }
26        for line in data.lines() {
27            output.push_str("data: ");
28            output.push_str(line);
29            output.push('\n');
30        }
31        output.push('\n');
32        Bytes::from(output)
33    }
34}
35
36#[derive(Default)]
37pub struct SseDecoder {
38    buffer: Vec<u8>,
39}
40
41impl SseDecoder {
42    pub fn push(&mut self, chunk: &[u8]) -> Result<Vec<SseFrame>, TransformError> {
43        self.buffer.extend_from_slice(chunk);
44        if self.buffer.len() > 100 * 1024 * 1024 {
45            return Err(TransformError::shape("SSE", "frame exceeds 100 MiB"));
46        }
47        let mut frames = Vec::new();
48        while let Some((end, delimiter)) = delimiter(&self.buffer) {
49            let raw = self.buffer.drain(..end + delimiter).collect::<Vec<_>>();
50            if let Some(frame) = parse(&raw[..end])? {
51                frames.push(frame);
52            }
53        }
54        Ok(frames)
55    }
56
57    pub fn finish(&mut self) -> Result<Option<SseFrame>, TransformError> {
58        if self.buffer.is_empty() {
59            return Ok(None);
60        }
61        let raw = std::mem::take(&mut self.buffer);
62        parse(&raw)
63    }
64}
65
66fn delimiter(buffer: &[u8]) -> Option<(usize, usize)> {
67    let lf = find(buffer, b"\n\n").map(|index| (index, 2));
68    let crlf = find(buffer, b"\r\n\r\n").map(|index| (index, 4));
69    match (lf, crlf) {
70        (Some(left), Some(right)) => Some(if left.0 <= right.0 { left } else { right }),
71        (left, right) => left.or(right),
72    }
73}
74
75fn parse(raw: &[u8]) -> Result<Option<SseFrame>, TransformError> {
76    let text =
77        std::str::from_utf8(raw).map_err(|_| TransformError::shape("SSE", "frame is not UTF-8"))?;
78    let mut event = None;
79    let mut data = Vec::new();
80    for line in text.lines() {
81        if let Some(value) = line.strip_prefix("event:") {
82            event = Some(value.trim_start().to_owned());
83        } else if let Some(value) = line.strip_prefix("data:") {
84            data.push(value.strip_prefix(' ').unwrap_or(value));
85        }
86    }
87    Ok((!data.is_empty()).then(|| SseFrame {
88        event,
89        data: data.join("\n"),
90    }))
91}
92
93fn find(haystack: &[u8], needle: &[u8]) -> Option<usize> {
94    haystack
95        .windows(needle.len())
96        .position(|candidate| candidate == needle)
97}