use std::fmt;
use std::time::Duration;
use bytes::{Bytes, BytesMut};
use serde::Serialize;
use topcoat_core::{
context::Cx,
error::{Error, Result},
};
use crate::headers;
#[derive(Clone, Debug, Default)]
#[must_use]
pub struct Event {
comment: Option<String>,
kind: Option<String>,
data: Option<String>,
id: Option<String>,
retry: Option<Duration>,
}
impl Event {
pub fn new() -> Self {
Self::default()
}
pub fn data(mut self, data: impl Into<String>) -> Self {
self.data = Some(data.into());
self
}
pub fn json_data<T>(self, value: &T) -> Result<Self>
where
T: Serialize + ?Sized,
{
Ok(self.data(serde_json::to_string(value).map_err(Error::from)?))
}
pub fn event(mut self, event: impl Into<String>) -> Self {
self.kind = Some(event.into());
self
}
pub fn id(mut self, id: impl Into<String>) -> Self {
self.id = Some(id.into());
self
}
pub fn retry(mut self, retry: Duration) -> Self {
self.retry = Some(retry);
self
}
pub fn comment(mut self, comment: impl Into<String>) -> Self {
self.comment = Some(comment.into());
self
}
pub(super) fn serialize(&self) -> Result<Bytes> {
let mut buffer = BytesMut::new();
if let Some(comment) = &self.comment {
for line in lines(comment) {
put_field(&mut buffer, "", line);
}
}
if let Some(kind) = &self.kind {
if kind.contains(['\r', '\n']) {
return Err(InvalidEventError::new("the event type contains a line break").into());
}
put_field(&mut buffer, "event", kind);
}
if let Some(data) = &self.data {
for line in lines(data) {
put_field(&mut buffer, "data", line);
}
}
if let Some(id) = &self.id {
if id.contains(['\r', '\n', '\0']) {
return Err(InvalidEventError::new(
"the event id contains a line break or null character",
)
.into());
}
put_field(&mut buffer, "id", id);
}
if let Some(retry) = &self.retry {
put_field(&mut buffer, "retry", &retry.as_millis().to_string());
}
buffer.extend_from_slice(b"\n");
Ok(buffer.freeze())
}
}
#[derive(Debug)]
pub struct InvalidEventError {
description: &'static str,
}
impl InvalidEventError {
fn new(description: &'static str) -> Self {
Self { description }
}
#[must_use]
pub fn description(&self) -> &str {
self.description
}
}
impl fmt::Display for InvalidEventError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "invalid server-sent event: {}", self.description)
}
}
impl std::error::Error for InvalidEventError {}
#[inline]
#[must_use]
pub fn last_event_id(cx: &Cx) -> Option<&str> {
headers(cx).get("last-event-id")?.to_str().ok()
}
fn put_field(buffer: &mut BytesMut, name: &str, value: &str) {
buffer.extend_from_slice(name.as_bytes());
buffer.extend_from_slice(b": ");
buffer.extend_from_slice(value.as_bytes());
buffer.extend_from_slice(b"\n");
}
fn lines(value: &str) -> impl Iterator<Item = &str> {
let mut rest = Some(value);
std::iter::from_fn(move || {
let current = rest.take()?;
let Some(index) = current.find(['\r', '\n']) else {
return Some(current);
};
let terminator = if current[index..].starts_with("\r\n") {
2
} else {
1
};
rest = Some(¤t[index + terminator..]);
Some(¤t[..index])
})
}
#[cfg(test)]
mod tests {
use http::Request;
use topcoat_core::context::CxTestBuilder;
use super::*;
fn serialize(event: &Event) -> String {
String::from_utf8(event.serialize().unwrap().to_vec()).unwrap()
}
#[test]
fn data_becomes_a_data_field() {
assert_eq!(serialize(&Event::new().data("hi")), "data: hi\n\n");
}
#[test]
fn an_empty_event_is_a_blank_line() {
assert_eq!(serialize(&Event::new()), "\n");
}
#[test]
fn every_field_is_serialized() {
let event = Event::new()
.comment("a comment")
.event("tick")
.data("hi")
.id("1")
.retry(Duration::from_secs(2));
assert_eq!(
serialize(&event),
": a comment\nevent: tick\ndata: hi\nid: 1\nretry: 2000\n\n"
);
}
#[test]
fn multi_line_data_becomes_one_field_per_line() {
assert_eq!(
serialize(&Event::new().data("a\nb\r\nc\rd")),
"data: a\ndata: b\ndata: c\ndata: d\n\n"
);
assert_eq!(
serialize(&Event::new().data("trailing\n")),
"data: trailing\ndata: \n\n"
);
}
#[test]
fn multi_line_comments_become_one_comment_per_line() {
assert_eq!(serialize(&Event::new().comment("a\nb")), ": a\n: b\n\n");
}
#[test]
fn json_data_serializes_the_value() {
let event = Event::new()
.json_data(&serde_json::json!({ "count": 3 }))
.unwrap();
assert_eq!(serialize(&event), "data: {\"count\":3}\n\n");
}
#[test]
fn line_breaks_in_the_event_type_are_an_error() {
let error = Event::new().event("a\nb").serialize().unwrap_err();
assert!(error.downcast_ref::<InvalidEventError>().is_some());
}
#[test]
fn line_breaks_and_null_in_the_id_are_an_error() {
for id in ["a\nb", "a\rb", "a\0b"] {
let error = Event::new().id(id).serialize().unwrap_err();
assert!(error.downcast_ref::<InvalidEventError>().is_some());
}
}
#[test]
fn last_event_id_reads_the_header() {
let request = Request::builder()
.uri("/events")
.header("last-event-id", "42")
.body(())
.unwrap();
let (parts, ()) = request.into_parts();
let cx = CxTestBuilder::new().request_context(parts).build();
assert_eq!(last_event_id(&cx), Some("42"));
}
#[test]
fn a_missing_last_event_id_is_none() {
let (parts, ()) = Request::builder()
.uri("/events")
.body(())
.unwrap()
.into_parts();
let cx = CxTestBuilder::new().request_context(parts).build();
assert_eq!(last_event_id(&cx), None);
}
}