This repository has no description
1use std::net::SocketAddr;
2use std::sync::Arc;
3use std::time::Duration;
4
5use axum::extract::ws::{CloseFrame, Message, WebSocket, WebSocketUpgrade};
6use axum::extract::{ConnectInfo, Query, State};
7use axum::response::{IntoResponse, Response};
8use futures::stream::SplitSink;
9use futures::{SinkExt, StreamExt};
10use http::HeaderMap;
11use serde::Deserialize;
12
13use knot_events::{
14 BatchEnd, EventCursor, EventLog, ReplayBounds, ReplayBytes, ReplayEvents, Replayed,
15};
16use knot_runtime::{Clock, HttpTransport};
17
18use crate::XrpcState;
19use crate::error::XrpcError;
20
21pub(crate) const EVENTS_ROUTE: &str = "/events";
22
23const DRAIN_BATCH: usize = 100;
24const DRAIN_BYTES: usize = 4 << 20;
25const MAX_BATCHES_PER_DRAIN: usize = 1_000;
26const KEEPALIVE: Duration = Duration::from_secs(30);
27const WRITE_DEADLINE: Duration = Duration::from_secs(10);
28const TRY_AGAIN_LATER: u16 = 1013;
29
30#[derive(Deserialize)]
31pub(crate) struct EventsQuery {
32 cursor: Option<String>,
33}
34
35pub(crate) async fn events<H: HttpTransport, C: Clock>(
36 State(state): State<Arc<XrpcState<H, C>>>,
37 ConnectInfo(socket_peer): ConnectInfo<SocketAddr>,
38 headers: HeaderMap,
39 Query(query): Query<EventsQuery>,
40 upgrade: WebSocketUpgrade,
41) -> Response {
42 let peer = state
43 .trusted_proxy_header
44 .as_ref()
45 .and_then(|header| knot_types::forwarded_peer(&headers, header))
46 .unwrap_or_else(|| socket_peer.ip());
47 let Some(permit) = state.subscriber_gate.try_admit(peer) else {
48 return XrpcError::overloaded(
49 "knot is serving its maximum number of event subscribers, retry shortly",
50 )
51 .into_response();
52 };
53 let cursor = query
54 .cursor
55 .as_deref()
56 .and_then(|raw| raw.parse::<i64>().ok())
57 .map(EventCursor::new)
58 .unwrap_or(EventCursor::START);
59 let log = Arc::clone(&state.events);
60 upgrade.on_upgrade(move |socket| async move {
61 let _permit = permit;
62 stream_events(socket, log, cursor).await;
63 })
64}
65
66enum Drained {
67 CaughtUp,
68 Limited,
69}
70
71async fn stream_events<C: Clock>(socket: WebSocket, log: Arc<EventLog<C>>, start: EventCursor) {
72 let mut head = log.subscribe();
73 let (mut sink, mut from_client) = socket.split();
74 let mut keepalive =
75 tokio::time::interval_at(tokio::time::Instant::now() + KEEPALIVE, KEEPALIVE);
76 let mut cursor = start;
77 loop {
78 match drain(&mut sink, &log, &mut cursor).await {
79 Ok(Drained::CaughtUp) => {}
80 Ok(Drained::Limited) => {
81 let close = Message::Close(Some(CloseFrame {
82 code: TRY_AGAIN_LATER,
83 reason: "drain limit reached, reconnect to continue".into(),
84 }));
85 let _ = tokio::time::timeout(WRITE_DEADLINE, sink.send(close)).await;
86 return;
87 }
88 Err(()) => return,
89 }
90 tokio::select! {
91 changed = head.changed() => {
92 if changed.is_err() {
93 return;
94 }
95 }
96 _ = keepalive.tick() => {
97 let ping = sink.send(Message::Ping(Vec::new().into()));
98 if !matches!(tokio::time::timeout(WRITE_DEADLINE, ping).await, Ok(Ok(()))) {
99 return;
100 }
101 }
102 received = from_client.next() => {
103 match received {
104 None | Some(Err(_)) | Some(Ok(Message::Close(_))) => return,
105 Some(Ok(_)) => {}
106 }
107 }
108 }
109 }
110}
111
112fn drain_bounds() -> ReplayBounds {
113 ReplayBounds::new(
114 ReplayEvents::new(DRAIN_BATCH).expect("drain event maximum is nonzero"),
115 ReplayBytes::new(DRAIN_BYTES).expect("drain byte maximum is nonzero"),
116 )
117}
118
119// who up draining they clock
120async fn drain<C: Clock>(
121 sink: &mut SplitSink<WebSocket, Message>,
122 log: &EventLog<C>,
123 cursor: &mut EventCursor,
124) -> Result<Drained, ()> {
125 let mut batches = 0;
126 loop {
127 let Replayed { events, end } = log.replay(*cursor, drain_bounds());
128 if let Some(last) = events.last() {
129 *cursor = last.created;
130 }
131 let messages: Vec<Result<Message, axum::Error>> = events
132 .iter()
133 .map(|event| {
134 Ok(Message::Text(
135 serde_json::to_string(event.as_ref())
136 .expect("wire event serializes to JSON")
137 .into(),
138 ))
139 })
140 .collect();
141 drop(events);
142 let sent = tokio::time::timeout(
143 WRITE_DEADLINE,
144 sink.send_all(&mut futures::stream::iter(messages)),
145 )
146 .await;
147 if !matches!(sent, Ok(Ok(()))) {
148 return Err(());
149 }
150 if end == BatchEnd::CaughtUp {
151 return Ok(Drained::CaughtUp);
152 }
153 batches += 1;
154 if batches == MAX_BATCHES_PER_DRAIN {
155 return Ok(Drained::Limited);
156 }
157 }
158}