Skip to main content

sse_stream/
stream.rs

1use std::{
2    collections::VecDeque,
3    num::ParseIntError,
4    str::Utf8Error,
5    task::{ready, Context, Poll},
6};
7
8use crate::Sse;
9use bytes::Buf;
10use futures_util::{stream::MapOk, Stream, TryStreamExt};
11use http_body::{Body, Frame};
12use http_body_util::{BodyDataStream, StreamBody};
13
14const BOM_HEADER: &[u8] = b"\xEF\xBB\xBF";
15
16struct ParserState {
17    parsed: VecDeque<Sse>,
18    current: Option<Sse>,
19    unfinished_line: Vec<u8>,
20    skip_leading_lf: bool,
21    first_line: bool,
22}
23
24impl Default for ParserState {
25    fn default() -> Self {
26        Self {
27            parsed: VecDeque::new(),
28            current: None,
29            unfinished_line: Vec::new(),
30            skip_leading_lf: false,
31            first_line: true,
32        }
33    }
34}
35
36pin_project_lite::pin_project! {
37    pub struct SseStream<B: Body> {
38        #[pin]
39        body: BodyDataStream<B>,
40        parser: ParserState,
41    }
42}
43
44pub type ByteStreamBody<S, D> = StreamBody<MapOk<S, fn(D) -> Frame<D>>>;
45impl<E, S, D> SseStream<ByteStreamBody<S, D>>
46where
47    S: Stream<Item = Result<D, E>>,
48    E: std::error::Error,
49    D: Buf,
50    StreamBody<ByteStreamBody<S, D>>: Body,
51{
52    /// Alias of [`from_bytes_stream`](Self::from_bytes_stream).
53    #[deprecated(
54        since = "0.2.4",
55        note = "It's a typo, use `from_bytes_stream` instead. This method will be removed in 0.3.0"
56    )]
57    pub fn from_byte_stream(stream: S) -> Self {
58        Self::from_bytes_stream(stream)
59    }
60
61    /// Create a new [`SseStream`] from a stream of [`Bytes`](bytes::Bytes).
62    ///
63    /// This is useful when you interact with clients don't provide response body directly like reqwest.
64    pub fn from_bytes_stream(stream: S) -> Self {
65        let stream = stream.map_ok(http_body::Frame::data as fn(D) -> Frame<D>);
66        let body = StreamBody::new(stream);
67        Self {
68            body: BodyDataStream::new(body),
69            parser: ParserState::default(),
70        }
71    }
72}
73
74impl<B: Body> SseStream<B> {
75    /// Create a new [`SseStream`] from a [`Body`].
76    pub fn new(body: B) -> Self {
77        Self {
78            body: BodyDataStream::new(body),
79            parser: ParserState::default(),
80        }
81    }
82}
83
84pub enum Error {
85    Body(Box<dyn std::error::Error + Send + Sync>),
86    InvalidLine,
87    DuplicatedEventLine,
88    DuplicatedIdLine,
89    DuplicatedRetry,
90    Utf8Parse(Utf8Error),
91    IntParse(ParseIntError),
92}
93
94impl std::fmt::Display for Error {
95    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
96        match self {
97            Error::Body(e) => write!(f, "body error: {}", e),
98            Error::InvalidLine => write!(f, "invalid line"),
99            Error::DuplicatedEventLine => write!(f, "duplicated event line"),
100            Error::DuplicatedIdLine => write!(f, "duplicated id line"),
101            Error::DuplicatedRetry => write!(f, "duplicated retry line"),
102            Error::Utf8Parse(e) => write!(f, "utf8 parse error: {}", e),
103            Error::IntParse(e) => write!(f, "int parse error: {}", e),
104        }
105    }
106}
107
108impl std::fmt::Debug for Error {
109    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
110        match self {
111            Error::Body(e) => write!(f, "Body({:?})", e),
112            Error::InvalidLine => write!(f, "InvalidLine"),
113            Error::DuplicatedEventLine => write!(f, "DuplicatedEventLine"),
114            Error::DuplicatedIdLine => write!(f, "DuplicatedIdLine"),
115            Error::DuplicatedRetry => write!(f, "DuplicatedRetry"),
116            Error::Utf8Parse(e) => write!(f, "Utf8Parse({:?})", e),
117            Error::IntParse(e) => write!(f, "IntParse({:?})", e),
118        }
119    }
120}
121
122impl std::error::Error for Error {
123    fn description(&self) -> &str {
124        match self {
125            Error::Body(_) => "body error",
126            Error::InvalidLine => "invalid line",
127            Error::DuplicatedEventLine => "duplicated event line",
128            Error::DuplicatedIdLine => "duplicated id line",
129            Error::DuplicatedRetry => "duplicated retry line",
130            Error::Utf8Parse(_) => "utf8 parse error",
131            Error::IntParse(_) => "int parse error",
132        }
133    }
134
135    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
136        match self {
137            Error::Body(e) => Some(e.as_ref()),
138            Error::Utf8Parse(e) => Some(e),
139            Error::IntParse(e) => Some(e),
140            _ => None,
141        }
142    }
143}
144
145impl ParserState {
146    fn parse_line(&mut self, mut line: &[u8]) -> Result<(), Error> {
147        if self.first_line {
148            self.first_line = false;
149            line = line.strip_prefix(BOM_HEADER).unwrap_or(line);
150        }
151
152        if line.is_empty() {
153            if let Some(sse) = self.current.take() {
154                self.parsed.push_back(sse);
155            }
156            return Ok(());
157        }
158
159        let Some(colon_index) = line.iter().position(|byte| *byte == b':') else {
160            #[cfg(feature = "tracing")]
161            tracing::warn!(?line, "invalid line, missing `:`");
162            return Err(Error::InvalidLine);
163        };
164        let field_name = &line[..colon_index];
165        let field_value = &line[colon_index + 1..];
166        let field_value = field_value.strip_prefix(b" ").unwrap_or(field_value);
167
168        match field_name {
169            b"data" => {
170                let data_line = std::str::from_utf8(field_value).map_err(Error::Utf8Parse)?;
171                let event = self.current.get_or_insert_default();
172                if let Some(data) = event.data.as_mut() {
173                    data.push('\n');
174                    data.push_str(data_line);
175                } else {
176                    event.data = Some(data_line.to_owned());
177                }
178            }
179            b"event" => {
180                let event_value = std::str::from_utf8(field_value).map_err(Error::Utf8Parse)?;
181                let event = self.current.get_or_insert_default();
182                if event.event.is_some() {
183                    return Err(Error::DuplicatedEventLine);
184                }
185                event.event = Some(event_value.to_owned());
186            }
187            b"id" => {
188                // Per spec: if the id field value contains U+0000 NULL,
189                // the entire field MUST be ignored.
190                if field_value.contains(&0_u8) {
191                    #[cfg(feature = "tracing")]
192                    tracing::warn!(?line, "id field contains NULL byte, ignoring per spec");
193                    return Ok(());
194                }
195                let id_value = std::str::from_utf8(field_value).map_err(Error::Utf8Parse)?;
196                let event = self.current.get_or_insert_default();
197                if event.id.is_some() {
198                    return Err(Error::DuplicatedIdLine);
199                }
200                event.id = Some(id_value.to_owned());
201            }
202            b"retry" => {
203                let retry_value = std::str::from_utf8(field_value)
204                    .map_err(Error::Utf8Parse)?
205                    .trim_ascii()
206                    .parse::<u64>()
207                    .map_err(Error::IntParse)?;
208                let event = self.current.get_or_insert_default();
209                if event.retry.is_some() {
210                    return Err(Error::DuplicatedRetry);
211                }
212                event.retry = Some(retry_value);
213            }
214            b"" => {
215                #[cfg(feature = "tracing")]
216                {
217                    if tracing::enabled!(tracing::Level::DEBUG) {
218                        let comment = std::str::from_utf8(field_value).map_err(Error::Utf8Parse)?;
219                        tracing::debug!(?comment, "sse comment line");
220                    }
221                }
222            }
223            _ => {
224                #[cfg(feature = "tracing")]
225                tracing::warn!(line = ?field_name, "invalid line: unknown field");
226                return Err(Error::InvalidLine);
227            }
228        }
229
230        Ok(())
231    }
232
233    fn parse_complete_line(&mut self, line: &[u8]) -> Result<(), Error> {
234        // Fast path to avoid copy overhead if we don't have anything buffered.
235        if self.unfinished_line.is_empty() {
236            self.parse_line(line)
237        } else {
238            let mut complete_line = std::mem::take(&mut self.unfinished_line);
239            complete_line.extend_from_slice(line);
240            let result = self.parse_line(&complete_line);
241            // Reuse the unfinished line buffer.
242            complete_line.clear();
243            self.unfinished_line = complete_line;
244            result
245        }
246    }
247
248    fn parse_chunk(&mut self, mut bytes: &[u8]) -> Result<(), Error> {
249        if self.skip_leading_lf {
250            self.skip_leading_lf = false;
251            if bytes[0] == b'\n' {
252                bytes = &bytes[1..];
253            }
254        }
255
256        while !bytes.is_empty() {
257            let Some(line_end) = bytes.iter().position(|byte| matches!(*byte, b'\n' | b'\r'))
258            else {
259                self.unfinished_line.extend_from_slice(bytes);
260                return Ok(());
261            };
262
263            self.parse_complete_line(&bytes[..line_end])?;
264
265            let delimiter = bytes[line_end];
266            bytes = &bytes[line_end + 1..];
267            if delimiter == b'\r' {
268                if bytes.first() == Some(&b'\n') {
269                    bytes = &bytes[1..];
270                } else if bytes.is_empty() {
271                    self.skip_leading_lf = true;
272                }
273            }
274        }
275
276        Ok(())
277    }
278}
279
280impl<B: Body> Stream for SseStream<B>
281where
282    B::Error: std::error::Error + Send + Sync + 'static,
283{
284    type Item = Result<Sse, Error>;
285
286    fn poll_next(
287        mut self: std::pin::Pin<&mut Self>,
288        cx: &mut Context<'_>,
289    ) -> Poll<Option<Self::Item>> {
290        let mut this = self.as_mut().project();
291        if let Some(sse) = this.parser.parsed.pop_front() {
292            return Poll::Ready(Some(Ok(sse)));
293        }
294        loop {
295            match ready!(this.body.as_mut().poll_next(cx)) {
296                Some(Err(error)) => return Poll::Ready(Some(Err(Error::Body(Box::new(error))))),
297                None => return Poll::Ready(None),
298                Some(Ok(mut data)) => {
299                    while data.has_remaining() {
300                        let bytes = data.chunk();
301                        debug_assert!(
302                            !bytes.is_empty(),
303                            "Buf::chunk returned an empty slice with bytes remaining"
304                        );
305                        let chunk_size = bytes.len();
306                        if let Err(error) = this.parser.parse_chunk(bytes) {
307                            return Poll::Ready(Some(Err(error)));
308                        }
309                        data.advance(chunk_size);
310                    }
311
312                    if let Some(sse) = this.parser.parsed.pop_front() {
313                        return Poll::Ready(Some(Ok(sse)));
314                    }
315                }
316            }
317        }
318    }
319}