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}