use axum::response::sse::{Event as AxumEvent, KeepAlive, Sse};
use axum::response::IntoResponse;
use core::convert::Infallible;
use futures::stream::{Stream, StreamExt};
#[derive(Debug, Clone)]
pub struct SseEvent {
data: String,
event: Option<String>,
id: Option<String>,
retry: Option<u64>,
}
impl SseEvent {
pub fn data(data: impl Into<String>) -> Self {
Self {
data: data.into(),
event: None,
id: None,
retry: None,
}
}
pub fn event(mut self, event: impl Into<String>) -> Self {
self.event = 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_ms: u64) -> Self {
self.retry = Some(retry_ms);
self
}
pub fn into_axum_event(self) -> Result<AxumEvent, Infallible> {
let mut event = AxumEvent::default().data(&self.data);
if let Some(name) = self.event {
event = event.event(name);
}
if let Some(id) = self.id {
event = event.id(id);
}
if let Some(retry) = self.retry {
event = event.retry(std::time::Duration::from_millis(retry));
}
Ok(event)
}
}
pub fn sse_response<S>(stream: S) -> impl IntoResponse
where
S: Stream<Item = Result<SseEvent, Infallible>> + Send + 'static,
{
let axum_stream = stream.map(|item| item.and_then(|e| e.into_axum_event()));
Sse::new(axum_stream).keep_alive(KeepAlive::default())
}
pub fn sse_response_with_interval<S>(stream: S, interval_secs: u64) -> impl IntoResponse
where
S: Stream<Item = Result<SseEvent, Infallible>> + Send + 'static,
{
let axum_stream = stream.map(|item| item.and_then(|e| e.into_axum_event()));
Sse::new(axum_stream).keep_alive(
KeepAlive::new()
.interval(std::time::Duration::from_secs(interval_secs))
.text("keep-alive"),
)
}
pub fn sse_from_events(events: Vec<SseEvent>) -> impl Stream<Item = Result<SseEvent, Infallible>> {
futures::stream::iter(events.into_iter().map(Ok))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_sse_event_data() {
let event = SseEvent::data("hello");
assert_eq!(event.data, "hello");
assert!(event.event.is_none());
assert!(event.id.is_none());
assert!(event.retry.is_none());
}
#[test]
fn test_sse_event_builder() {
let event = SseEvent::data("payload")
.event("update")
.id("123")
.retry(5000);
assert_eq!(event.data, "payload");
assert_eq!(event.event.as_deref(), Some("update"));
assert_eq!(event.id.as_deref(), Some("123"));
assert_eq!(event.retry, Some(5000));
}
#[test]
fn test_sse_event_to_axum() {
let event = SseEvent::data("test").event("ping");
let axum_event = event.into_axum_event();
assert!(axum_event.is_ok());
}
#[test]
fn test_sse_from_events() {
let events = vec![SseEvent::data("first"), SseEvent::data("second")];
let stream = sse_from_events(events);
let collected: Vec<_> = futures::executor::block_on(stream.collect());
assert_eq!(collected.len(), 2);
}
}