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 #[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 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 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 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 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 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}