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.proxy_trust.client_peer_of(&headers, socket_peer.ip());
43 let Some(permit) = state.subscriber_gate.try_admit(peer) else {
44 return XrpcError::overloaded(
45 "knot is serving its maximum number of event subscribers, retry shortly",
46 )
47 .into_response();
48 };
49 let cursor = query
50 .cursor
51 .as_deref()
52 .and_then(|raw| raw.parse::<i64>().ok())
53 .map(EventCursor::new)
54 .unwrap_or(EventCursor::START);
55 let log = Arc::clone(&state.events);
56 upgrade.on_upgrade(move |socket| async move {
57 let _permit = permit;
58 stream_events(socket, log, cursor).await;
59 })
60}
61
62enum Drained {
63 CaughtUp,
64 Limited,
65}
66
67async fn stream_events<C: Clock>(socket: WebSocket, log: Arc<EventLog<C>>, start: EventCursor) {
68 let mut head = log.subscribe();
69 let (mut sink, mut from_client) = socket.split();
70 let mut keepalive =
71 tokio::time::interval_at(tokio::time::Instant::now() + KEEPALIVE, KEEPALIVE);
72 let mut cursor = start;
73 loop {
74 match drain(&mut sink, &log, &mut cursor).await {
75 Ok(Drained::CaughtUp) => {}
76 Ok(Drained::Limited) => {
77 let close = Message::Close(Some(CloseFrame {
78 code: TRY_AGAIN_LATER,
79 reason: "drain limit reached, reconnect to continue".into(),
80 }));
81 let _ = tokio::time::timeout(WRITE_DEADLINE, sink.send(close)).await;
82 return;
83 }
84 Err(()) => return,
85 }
86 tokio::select! {
87 changed = head.changed() => {
88 if changed.is_err() {
89 return;
90 }
91 }
92 _ = keepalive.tick() => {
93 let ping = sink.send(Message::Ping(Vec::new().into()));
94 if !matches!(tokio::time::timeout(WRITE_DEADLINE, ping).await, Ok(Ok(()))) {
95 return;
96 }
97 }
98 received = from_client.next() => {
99 match received {
100 None | Some(Err(_)) | Some(Ok(Message::Close(_))) => return,
101 Some(Ok(_)) => {}
102 }
103 }
104 }
105 }
106}
107
108fn drain_bounds() -> ReplayBounds {
109 ReplayBounds::new(
110 ReplayEvents::new(DRAIN_BATCH).expect("drain event maximum is nonzero"),
111 ReplayBytes::new(DRAIN_BYTES).expect("drain byte maximum is nonzero"),
112 )
113}
114
115// who up draining they clock
116async fn drain<C: Clock>(
117 sink: &mut SplitSink<WebSocket, Message>,
118 log: &EventLog<C>,
119 cursor: &mut EventCursor,
120) -> Result<Drained, ()> {
121 let mut batches = 0;
122 loop {
123 let Replayed { events, end } = log.replay(*cursor, drain_bounds());
124 if let Some(last) = events.last() {
125 *cursor = last.created;
126 }
127 let messages: Vec<Result<Message, axum::Error>> = events
128 .iter()
129 .map(|event| {
130 Ok(Message::Text(
131 serde_json::to_string(event.as_ref())
132 .expect("wire event serializes to JSON")
133 .into(),
134 ))
135 })
136 .collect();
137 drop(events);
138 let sent = tokio::time::timeout(
139 WRITE_DEADLINE,
140 sink.send_all(&mut futures::stream::iter(messages)),
141 )
142 .await;
143 if !matches!(sent, Ok(Ok(()))) {
144 return Err(());
145 }
146 if end == BatchEnd::CaughtUp {
147 return Ok(Drained::CaughtUp);
148 }
149 batches += 1;
150 if batches == MAX_BATCHES_PER_DRAIN {
151 return Ok(Drained::Limited);
152 }
153 }
154}