Skip to main content

serverkit/
sse.rs

1use std::{
2    task::{Context, Poll},
3    time::Duration,
4};
5
6use crate::{Chunk, IntoResponse, Response, ResponseStream, StreamError, openapi::Operation};
7
8pub trait SseStream {
9    fn poll_next(
10        &mut self,
11        context: &mut Context<'_>,
12    ) -> Poll<Option<Result<SseEvent, StreamError>>>;
13}
14
15pub struct Sse<S> {
16    stream: S,
17}
18
19impl<S> Sse<S> {
20    pub fn new(stream: S) -> Self {
21        Self { stream }
22    }
23}
24
25impl<S: SseStream + 'static> IntoResponse for Sse<S> {
26    fn into_response(self) -> Response {
27        let mut response = Response::stream(
28            200,
29            EncodedSseStream {
30                stream: self.stream,
31            },
32        );
33        response
34            .headers()
35            .set("Content-Type", "text/event-stream; charset=utf-8")
36            .expect("the built-in SSE content type is valid");
37        response
38            .headers()
39            .set("Cache-Control", "no-cache")
40            .expect("the built-in SSE cache policy is valid");
41        response
42    }
43
44    fn openapi(operation: &mut Operation) {
45        operation.response(
46            200,
47            "Server-sent event stream",
48            Some("text/event-stream"),
49            None,
50        );
51    }
52}
53
54pub struct SseEvent {
55    data: String,
56    event: Option<String>,
57    id: Option<String>,
58    retry: Option<Duration>,
59}
60
61impl SseEvent {
62    pub fn data(data: impl Into<String>) -> Self {
63        Self {
64            data: data.into(),
65            event: None,
66            id: None,
67            retry: None,
68        }
69    }
70
71    pub fn event(mut self, event: impl Into<String>) -> Self {
72        self.event = Some(event.into());
73        self
74    }
75
76    pub fn id(mut self, id: impl Into<String>) -> Self {
77        self.id = Some(id.into());
78        self
79    }
80
81    pub fn retry(mut self, retry: Duration) -> Self {
82        self.retry = Some(retry);
83        self
84    }
85
86    #[cfg(test)]
87    fn encode(self) -> Vec<u8> {
88        let mut encoded = Vec::new();
89        self.encode_into(&mut encoded);
90        encoded
91    }
92
93    fn encode_into(self, encoded: &mut Vec<u8>) {
94        encoded.clear();
95
96        if let Some(event) = self.event {
97            encoded.extend_from_slice(b"event: ");
98            extend_sanitized(encoded, &event);
99            encoded.push(b'\n');
100        }
101
102        if let Some(id) = self.id {
103            encoded.extend_from_slice(b"id: ");
104            extend_sanitized(encoded, &id);
105            encoded.push(b'\n');
106        }
107
108        if let Some(retry) = self.retry {
109            encoded.extend_from_slice(b"retry: ");
110            encoded.extend_from_slice(retry.as_millis().to_string().as_bytes());
111            encoded.push(b'\n');
112        }
113
114        for line in self.data.lines() {
115            encoded.extend_from_slice(b"data: ");
116            encoded.extend_from_slice(line.as_bytes());
117            encoded.push(b'\n');
118        }
119
120        if self.data.is_empty() {
121            encoded.extend_from_slice(b"data:\n");
122        }
123
124        encoded.push(b'\n');
125    }
126}
127
128struct EncodedSseStream<S> {
129    stream: S,
130}
131
132impl<S: SseStream> ResponseStream for EncodedSseStream<S> {
133    fn poll_next(&mut self, context: &mut Context<'_>) -> Poll<Option<Result<Chunk, StreamError>>> {
134        match self.stream.poll_next(context) {
135            Poll::Ready(Some(Ok(event))) => {
136                let mut encoded = Vec::new();
137                event.encode_into(&mut encoded);
138                Poll::Ready(Some(Ok(Chunk::from(encoded))))
139            }
140            Poll::Ready(Some(Err(error))) => Poll::Ready(Some(Err(error))),
141            Poll::Ready(None) => Poll::Ready(None),
142            Poll::Pending => Poll::Pending,
143        }
144    }
145}
146
147fn extend_sanitized(encoded: &mut Vec<u8>, value: &str) {
148    encoded.extend(value.bytes().filter(|byte| !matches!(byte, b'\r' | b'\n')));
149}
150
151#[cfg(test)]
152mod tests {
153    use std::time::Duration;
154
155    use super::SseEvent;
156
157    #[test]
158    fn encodes_an_event() {
159        let event = SseEvent::data("first\nsecond")
160            .event("update")
161            .id("42")
162            .retry(Duration::from_secs(1));
163
164        assert_eq!(
165            event.encode(),
166            b"event: update\nid: 42\nretry: 1000\ndata: first\ndata: second\n\n",
167        );
168    }
169}