From dfc15bce9dcaf014eef2a64a41aaa59397f0f32f Mon Sep 17 00:00:00 2001 From: Glen De Cauwsemaecker Date: Wed, 13 Aug 2025 23:32:08 +0200 Subject: [PATCH] Support custom (binary) data to be written into SSE Event (#3425) --- axum/src/response/sse.rs | 229 +++++++++++++++++++++++---------------- 1 file changed, 135 insertions(+), 94 deletions(-) diff --git a/axum/src/response/sse.rs b/axum/src/response/sse.rs index 43cd590a..974a4939 100644 --- a/axum/src/response/sse.rs +++ b/axum/src/response/sse.rs @@ -39,7 +39,9 @@ use futures_util::stream::TryStream; use http_body::Frame; use pin_project_lite::pin_project; use std::{ - fmt, mem, + fmt::{self, Write as _}, + io::Write as _, + mem, pin::Pin, task::{ready, Context, Poll}, time::Duration, @@ -174,6 +176,27 @@ pub struct Event { flags: EventFlags, } +/// Expose [`Event`] as a [`std::fmt::Write`] +/// such that any form of data can be written as data safely. +/// +/// This also ensures that newline characters `\r` and `\n` +/// correctly trigger a split with a new `data: ` prefix. +/// +/// # Panics +/// +/// Panics if any `data` has already been written prior to the first write +/// of this [`EventDataWriter`] instance. +#[derive(Debug)] +#[must_use] +pub struct EventDataWriter { + event: Event, + + // Indicates if _this_ EventDataWriter has written data, + // this does not say anything about whether or not `event` contains + // data or not. + data_written: bool, +} + impl Event { /// Default keep-alive event pub const DEFAULT_KEEP_ALIVE: Self = Self::finalized(Bytes::from_static(b":\n\n")); @@ -185,6 +208,19 @@ impl Event { } } + /// Use this [`Event`] as a [`EventDataWriter`] to write custom data. + /// + /// - [`Self::data`] can be used as a shortcut to write `str` data + /// - [`Self::json_data`] can be used as a shortcut to write `json` data + /// + /// Turn it into an [`Event`] again using [`EventDataWriter::into_event`]. + pub fn into_data_writer(self) -> EventDataWriter { + EventDataWriter { + event: self, + data_written: false, + } + } + /// Set the event's data data field(s) (`data: `) /// /// Newlines in `data` will automatically be broken across `data: ` fields. @@ -195,25 +231,16 @@ impl Event { /// /// # Panics /// - /// - Panics if `data` contains any carriage returns, as they cannot be transmitted over SSE. - /// - Panics if `data` or `json_data` have already been called. + /// Panics if any `data` has already been written before. /// /// [`MessageEvent`'s data field]: https://developer.mozilla.org/en-US/docs/Web/API/MessageEvent/data - pub fn data(mut self, data: T) -> Self + pub fn data(self, data: T) -> Self where T: AsRef, { - if self.flags.contains(EventFlags::HAS_DATA) { - panic!("Called `Event::data` multiple times"); - } - - for line in memchr_split(b'\n', data.as_ref().as_bytes()) { - self.field("data", line); - } - - self.flags.insert(EventFlags::HAS_DATA); - - self + let mut writer = self.into_data_writer(); + let _ = writer.write_str(data.as_ref()); + writer.into_event() } /// Set the event's data field to a value serialized as unformatted JSON (`data: `). @@ -222,43 +249,31 @@ impl Event { /// /// # Panics /// - /// Panics if `data` or `json_data` have already been called. + /// Panics if any `data` has already been written before. /// /// [`MessageEvent`'s data field]: https://developer.mozilla.org/en-US/docs/Web/API/MessageEvent/data #[cfg(feature = "json")] - pub fn json_data(mut self, data: T) -> Result + pub fn json_data(self, data: T) -> Result where T: serde::Serialize, { - struct IgnoreNewLines<'a>(bytes::buf::Writer<&'a mut BytesMut>); - impl std::io::Write for IgnoreNewLines<'_> { + struct JsonWriter<'a>(&'a mut EventDataWriter); + impl std::io::Write for JsonWriter<'_> { + #[inline] fn write(&mut self, buf: &[u8]) -> std::io::Result { - let mut last_split = 0; - for delimiter in memchr::memchr2_iter(b'\n', b'\r', buf) { - self.0.write_all(&buf[last_split..delimiter])?; - last_split = delimiter + 1; - } - self.0.write_all(&buf[last_split..])?; - Ok(buf.len()) + Ok(self.0.write_buf(buf)) } - fn flush(&mut self) -> std::io::Result<()> { - self.0.flush() + Ok(()) } } - if self.flags.contains(EventFlags::HAS_DATA) { - panic!("Called `Event::json_data` multiple times"); - } - let buffer = self.buffer.as_mut(); - buffer.extend_from_slice(b"data: "); - serde_json::to_writer(IgnoreNewLines(buffer.writer()), &data) - .map_err(axum_core::Error::new)?; - buffer.put_u8(b'\n'); + let mut writer = self.into_data_writer(); - self.flags.insert(EventFlags::HAS_DATA); + let json_writer = JsonWriter(&mut writer); + serde_json::to_writer(json_writer, &data).map_err(axum_core::Error::new)?; - Ok(self) + Ok(writer.into_event()) } /// Set the event's comment field (`:`). @@ -407,6 +422,60 @@ impl Event { } } +impl EventDataWriter { + /// Consume the [`EventDataWriter`] and return the [`Event`] once again. + /// + /// In case any data was written by this instance + /// it will also write the trailing `\n` character. + pub fn into_event(self) -> Event { + let mut event = self.event; + if self.data_written { + let _ = event.buffer.as_mut().write_char('\n'); + } + event + } +} + +impl EventDataWriter { + // Assumption: underlying writer never returns an error: + // + fn write_buf(&mut self, buf: &[u8]) -> usize { + if buf.is_empty() { + return 0; + } + + let buffer = self.event.buffer.as_mut(); + + if !std::mem::replace(&mut self.data_written, true) { + if self.event.flags.contains(EventFlags::HAS_DATA) { + panic!("Called `Event::data*` multiple times"); + } + + let _ = buffer.write_str("data: "); + self.event.flags.insert(EventFlags::HAS_DATA); + } + + let mut writer = buffer.writer(); + + let mut last_split = 0; + for delimiter in memchr::memchr2_iter(b'\n', b'\r', buf) { + let _ = writer.write_all(&buf[last_split..=delimiter]); + let _ = writer.write_all(b"data: "); + last_split = delimiter + 1; + } + let _ = writer.write_all(&buf[last_split..]); + + buf.len() + } +} + +impl fmt::Write for EventDataWriter { + fn write_str(&mut self, s: &str) -> fmt::Result { + let _ = self.write_buf(s.as_bytes()); + Ok(()) + } +} + impl Default for Event { fn default() -> Self { Self { @@ -566,32 +635,6 @@ where } } -fn memchr_split(needle: u8, haystack: &[u8]) -> MemchrSplit<'_> { - MemchrSplit { - needle, - haystack: Some(haystack), - } -} - -struct MemchrSplit<'a> { - needle: u8, - haystack: Option<&'a [u8]>, -} - -impl<'a> Iterator for MemchrSplit<'a> { - type Item = &'a [u8]; - fn next(&mut self) -> Option { - let haystack = self.haystack?; - if let Some(pos) = memchr::memchr(self.needle, haystack) { - let (front, back) = haystack.split_at(pos); - self.haystack = Some(&back[1..]); - Some(front) - } else { - self.haystack.take() - } - } -} - #[cfg(test)] mod tests { use super::*; @@ -611,14 +654,40 @@ mod tests { } #[test] - fn valid_json_raw_value_chars_stripped() { + fn write_data_writer_str() { + // also confirm that nop writers do nothing :) + let mut writer = Event::default() + .into_data_writer() + .into_event() + .into_data_writer(); + writer.write_str("").unwrap(); + let mut writer = writer.into_event().into_data_writer(); + + writer.write_str("").unwrap(); + writer.write_str("moon ").unwrap(); + writer.write_str("star\nsun").unwrap(); + writer.write_str("").unwrap(); + writer.write_str("set").unwrap(); + writer.write_str("").unwrap(); + writer.write_str(" bye\r").unwrap(); + + let event = writer.into_event(); + + assert_eq!( + &*event.finalize(), + b"data: moon star\ndata: sunset bye\rdata: \n\n" + ); + } + + #[test] + fn valid_json_raw_value_chars_handled() { let json_string = "{\r\"foo\": \n\r\r \"bar\\n\"\n}"; let json_raw_value_event = Event::default() .json_data(serde_json::from_str::<&RawValue>(json_string).unwrap()) .unwrap(); assert_eq!( &*json_raw_value_event.finalize(), - format!("data: {}\n\n", json_string.replace(['\n', '\r'], "")).as_bytes() + b"data: {\rdata: \"foo\": \ndata: \rdata: \rdata: \"bar\\n\"\ndata: }\n\n" ); } @@ -763,32 +832,4 @@ mod tests { fields } - - #[test] - fn memchr_splitting() { - assert_eq!( - memchr_split(2, &[]).collect::>(), - [&[]] as [&[u8]; 1] - ); - assert_eq!( - memchr_split(2, &[2]).collect::>(), - [&[], &[]] as [&[u8]; 2] - ); - assert_eq!( - memchr_split(2, &[1]).collect::>(), - [&[1]] as [&[u8]; 1] - ); - assert_eq!( - memchr_split(2, &[1, 2]).collect::>(), - [&[1], &[]] as [&[u8]; 2] - ); - assert_eq!( - memchr_split(2, &[2, 1]).collect::>(), - [&[], &[1]] as [&[u8]; 2] - ); - assert_eq!( - memchr_split(2, &[1, 2, 2, 1]).collect::>(), - [&[1], &[], &[1]] as [&[u8]; 3] - ); - } }