mirror of
https://github.com/tokio-rs/axum.git
synced 2026-08-13 00:00:35 +02:00
Support custom (binary) data to be written into SSE Event (#3425)
This commit is contained in:
+135
-94
@@ -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: <content>`)
|
||||
///
|
||||
/// 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<T>(mut self, data: T) -> Self
|
||||
pub fn data<T>(self, data: T) -> Self
|
||||
where
|
||||
T: AsRef<str>,
|
||||
{
|
||||
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: <content>`).
|
||||
@@ -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<T>(mut self, data: T) -> Result<Self, axum_core::Error>
|
||||
pub fn json_data<T>(self, data: T) -> Result<Self, axum_core::Error>
|
||||
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<usize> {
|
||||
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 (`:<comment-text>`).
|
||||
@@ -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:
|
||||
// <https://docs.rs/bytes/latest/src/bytes/buf/writer.rs.html#79-82>
|
||||
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<Self::Item> {
|
||||
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::<Vec<_>>(),
|
||||
[&[]] as [&[u8]; 1]
|
||||
);
|
||||
assert_eq!(
|
||||
memchr_split(2, &[2]).collect::<Vec<_>>(),
|
||||
[&[], &[]] as [&[u8]; 2]
|
||||
);
|
||||
assert_eq!(
|
||||
memchr_split(2, &[1]).collect::<Vec<_>>(),
|
||||
[&[1]] as [&[u8]; 1]
|
||||
);
|
||||
assert_eq!(
|
||||
memchr_split(2, &[1, 2]).collect::<Vec<_>>(),
|
||||
[&[1], &[]] as [&[u8]; 2]
|
||||
);
|
||||
assert_eq!(
|
||||
memchr_split(2, &[2, 1]).collect::<Vec<_>>(),
|
||||
[&[], &[1]] as [&[u8]; 2]
|
||||
);
|
||||
assert_eq!(
|
||||
memchr_split(2, &[1, 2, 2, 1]).collect::<Vec<_>>(),
|
||||
[&[1], &[], &[1]] as [&[u8]; 3]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user