1use crate::{Eventable, encode};
2use futures_lite::{AsyncWriteExt, Stream};
3use std::{
4 fmt::{self, Debug, Formatter},
5 future::{self, Future},
6 io,
7 pin::Pin,
8 task::Poll,
9 time::Duration,
10};
11use sync_wrapper::SyncWrapper;
12use trillium::{Conn, Handler, Info, KnownHeaderName, Status, Upgrade};
13use trillium_server_common::Runtime;
14
15pub trait SseHandler: Send + Sync + Sized + 'static {
55 type Event: Eventable;
57
58 type EventStream: Stream<Item = Self::Event> + Unpin + Send + 'static;
61
62 fn connect(&self, conn: &mut Conn) -> impl Future<Output = Self::EventStream> + Send;
67}
68
69impl<F, S, E> SseHandler for F
70where
71 F: Fn(&mut Conn) -> S + Send + Sync + 'static,
72 S: Stream<Item = E> + Unpin + Send + 'static,
73 E: Eventable,
74{
75 type Event = E;
76 type EventStream = S;
77
78 async fn connect(&self, conn: &mut Conn) -> Self::EventStream {
79 self(conn)
80 }
81}
82
83pub struct Sse<H> {
96 handler: H,
97 heartbeat: Option<Duration>,
98 runtime: Option<Runtime>,
99}
100
101pub fn sse<H: SseHandler>(sse_handler: H) -> Sse<H> {
103 Sse::new(sse_handler)
104}
105
106impl<H: SseHandler> Sse<H> {
107 pub fn new(handler: H) -> Self {
109 Self {
110 handler,
111 heartbeat: None,
112 runtime: None,
113 }
114 }
115
116 pub fn with_heartbeat(mut self, heartbeat: Duration) -> Self {
125 self.heartbeat = Some(heartbeat);
126 self
127 }
128}
129
130fn accepts_event_stream(conn: &Conn) -> bool {
137 let Some(accept) = conn.request_headers().get_str(KnownHeaderName::Accept) else {
138 return false;
139 };
140
141 accept.split(',').any(|media_range| {
142 let mut parts = media_range.split(';').map(str::trim);
143
144 let matches_range = parts
145 .next()
146 .is_some_and(|range| range.eq_ignore_ascii_case("text/event-stream"));
147
148 matches_range
149 && !parts.any(|parameter| {
150 parameter
151 .strip_prefix("q=")
152 .and_then(|q| q.parse::<f32>().ok())
153 .is_some_and(|q| q <= 0.0)
154 })
155 })
156}
157
158struct SseStream<S>(SyncWrapper<S>);
162
163const READ_ALLOWANCE: usize = 16 * 1024;
167
168enum Tick<E> {
169 Event(E),
170 Heartbeat,
171 StreamEnded,
172 ClientDisconnected,
173}
174
175async fn write_flush(upgrade: &mut Upgrade, bytes: &[u8]) -> io::Result<()> {
176 upgrade.write_all(bytes).await?;
177 upgrade.flush().await
178}
179
180async fn drive_events<S, E>(mut upgrade: Upgrade, stream: S, heartbeat: Option<(Duration, Runtime)>)
181where
182 S: Stream<Item = E> + Unpin + Send + 'static,
183 E: Eventable,
184{
185 let swansong = upgrade.swansong();
186 let mut stream = swansong.interrupt(stream);
187
188 let new_delay = |(duration, runtime): &(Duration, Runtime)| {
189 let duration = *duration;
190 let runtime = runtime.clone();
191 Box::pin(async move { runtime.delay(duration).await })
192 as Pin<Box<dyn Future<Output = ()> + Send>>
193 };
194 let mut delay = heartbeat.as_ref().map(new_delay);
195
196 loop {
197 let tick = future::poll_fn(|cx| {
198 if Pin::new(upgrade.as_mut())
199 .poll_closed(cx, READ_ALLOWANCE)
200 .is_ready()
201 {
202 return Poll::Ready(Tick::ClientDisconnected);
203 }
204
205 match Pin::new(&mut stream).poll_next(cx) {
206 Poll::Ready(Some(event)) => return Poll::Ready(Tick::Event(event)),
207 Poll::Ready(None) => return Poll::Ready(Tick::StreamEnded),
208 Poll::Pending => {}
209 }
210
211 if let Some(delay) = &mut delay
212 && delay.as_mut().poll(cx).is_ready()
213 {
214 return Poll::Ready(Tick::Heartbeat);
215 }
216
217 Poll::Pending
218 })
219 .await;
220
221 match tick {
222 Tick::ClientDisconnected => return,
223 Tick::StreamEnded => break,
224 Tick::Event(event) => {
225 delay = heartbeat.as_ref().map(new_delay);
226 let Some(encoded) = encode(&event) else {
227 continue;
228 };
229 if write_flush(&mut upgrade, encoded.as_bytes()).await.is_err() {
230 return;
231 }
232 }
233 Tick::Heartbeat => {
234 delay = heartbeat.as_ref().map(new_delay);
235 if write_flush(&mut upgrade, b":\n\n").await.is_err() {
236 return;
237 }
238 }
239 }
240 }
241
242 let _ = upgrade.close().await;
243}
244
245impl<H: SseHandler> Handler for Sse<H> {
246 async fn run(&self, mut conn: Conn) -> Conn {
247 if !accepts_event_stream(&conn) {
248 return conn;
249 }
250
251 let stream = self.handler.connect(&mut conn).await;
252
253 conn.with_state(SseStream(SyncWrapper::new(stream)))
254 .with_response_header(KnownHeaderName::ContentType, "text/event-stream")
255 .with_response_header(KnownHeaderName::CacheControl, "no-cache")
256 .with_response_header(KnownHeaderName::Connection, "close")
261 .with_status(Status::Ok)
262 .halt()
263 .upgrade()
264 }
265
266 async fn init(&mut self, info: &mut Info) {
267 self.runtime = info.shared_state::<Runtime>().cloned();
268
269 if self.heartbeat.is_some() && self.runtime.is_none() {
270 log::warn!(
271 "no runtime in shared state; sse heartbeats are disabled. this handler was \
272 probably not initialized by a trillium runtime adapter."
273 );
274 }
275 }
276
277 fn has_upgrade(&self, upgrade: &Upgrade) -> bool {
278 upgrade.state().contains::<SseStream<H::EventStream>>()
279 }
280
281 async fn upgrade(&self, mut upgrade: Upgrade) {
282 let Some(SseStream(stream)) = upgrade.state_mut().take::<SseStream<H::EventStream>>()
283 else {
284 return;
285 };
286 let stream = stream.into_inner();
287
288 let heartbeat = self.heartbeat.zip(self.runtime.clone());
289 drive_events(upgrade, stream, heartbeat).await;
290 }
291}
292
293impl<H> Debug for Sse<H>
294where
295 H: Debug,
296{
297 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
298 f.debug_struct("Sse")
299 .field("handler", &self.handler)
300 .field("heartbeat", &self.heartbeat)
301 .field("runtime", &self.runtime)
302 .finish()
303 }
304}