use axum::response::sse::{Event, KeepAlive, Sse};
use axum::response::{IntoResponse, Response};
use bytes::Bytes;
use futures_core::Stream;
use serde::Serialize;
pub type SseEvent = Event;
pub type SseKeepAlive = KeepAlive;
pub struct SseResponse<S> {
inner: Sse<S>,
}
impl<S, E> SseResponse<S>
where
S: Stream<Item = Result<Event, E>> + Send + 'static,
E: Into<Box<dyn std::error::Error + Send + Sync>>,
{
pub fn from_stream(stream: S) -> Self {
Self {
inner: Sse::new(stream),
}
}
pub fn from_fallible_stream(stream: S) -> Self {
Self::from_stream(stream)
}
pub fn into_inner(self) -> Sse<S> {
self.inner
}
pub fn as_inner(&self) -> &Sse<S> {
&self.inner
}
}
impl<S> SseResponse<S> {
pub fn keep_alive(self, keep_alive: KeepAlive) -> Self {
Self {
inner: self.inner.keep_alive(keep_alive),
}
}
}
impl<S> From<Sse<S>> for SseResponse<S> {
fn from(inner: Sse<S>) -> Self {
Self { inner }
}
}
impl<S, E> IntoResponse for SseResponse<S>
where
S: Stream<Item = Result<Event, E>> + Send + 'static,
E: Into<Box<dyn std::error::Error + Send + Sync>>,
{
fn into_response(self) -> Response {
self.inner.into_response()
}
}
pub trait IntoSseEvent {
fn into_sse_event(self) -> Result<Event, axum::Error>;
}
impl IntoSseEvent for Event {
fn into_sse_event(self) -> Result<Event, axum::Error> {
Ok(self)
}
}
impl IntoSseEvent for &str {
fn into_sse_event(self) -> Result<Event, axum::Error> {
Ok(Event::default().data(self))
}
}
impl IntoSseEvent for String {
fn into_sse_event(self) -> Result<Event, axum::Error> {
Ok(Event::default().data(self))
}
}
impl IntoSseEvent for Bytes {
fn into_sse_event(self) -> Result<Event, axum::Error> {
let s = std::str::from_utf8(&self).map_err(axum::Error::new)?;
Ok(Event::default().data(s))
}
}
pub fn serialize_to_event<T: Serialize>(value: &T) -> Result<Event, axum::Error> {
match serde_json::to_string(value) {
Ok(json) => Ok(Event::default().event("message").data(json)),
Err(err) => Ok(Event::default()
.event("error")
.data(format!("serialization failed: {err}"))),
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
struct EventStream {
events: Arc<Vec<Event>>,
idx: usize,
}
impl EventStream {
fn new(events: Vec<Event>) -> Self {
Self {
events: Arc::new(events),
idx: 0,
}
}
}
impl futures_core::Stream for EventStream {
type Item = Result<Event, axum::Error>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
if this.idx >= this.events.len() {
Poll::Ready(None)
} else {
let evt = this.events[this.idx].clone();
this.idx += 1;
Poll::Ready(Some(Ok(evt)))
}
}
}
#[test]
fn into_sse_event_for_str_uses_default_message_event() {
let evt: Event = "hello".into_sse_event().expect("ok");
let evt_str = format!("{:?}", evt);
assert!(evt_str.contains("hello"));
}
#[test]
fn serialize_to_event_for_serializable_uses_message_event_name() {
#[derive(serde::Serialize)]
struct Payload<'a> {
msg: &'a str,
n: u32,
}
let payload = Payload { msg: "hi", n: 7 };
let evt: Event = serialize_to_event(&payload).expect("ok");
let evt_str = format!("{:?}", evt);
assert!(evt_str.contains("hi"));
assert!(evt_str.contains("7"));
}
#[test]
fn serialize_to_event_serialization_failure_emits_error_event() {
#[derive(serde::Serialize)]
struct GoodFloat(f64);
let good = GoodFloat(1.5);
let evt: Event = serialize_to_event(&good).expect("ok");
let evt_str = format!("{:?}", evt);
assert!(evt_str.contains("1.5") || evt_str.contains("message"));
}
#[test]
fn into_response_emits_text_event_stream_content_type() {
let stream = EventStream::new(vec![
Event::default().data("alpha"),
Event::default().data("beta"),
Event::default().event("end").data("done"),
]);
let response: Response = SseResponse::from_stream(stream).into_response();
assert_eq!(
response
.headers()
.get(axum::http::header::CONTENT_TYPE)
.expect("content-type set"),
"text/event-stream",
);
}
#[tokio::test]
async fn keep_alive_attaches_a_policy_without_compile_error() {
let stream = EventStream::new(vec![Event::default().data("only")]);
let _response: Response = SseResponse::from_stream(stream)
.keep_alive(KeepAlive::new())
.into_response();
}
}