1mod device_changes;
2mod presence;
3mod receipts;
4
5use std::{
6 iter::repeat_with,
7 mem::replace,
8 sync::atomic::{AtomicU64, AtomicUsize, Ordering},
9 time::SystemTime,
10};
11
12use futures::{StreamExt, future::join3};
13use ruma::{ServerName, api::federation::transactions::edu::Edu};
14use tuwunel_core::{
15 Result, implement,
16 matrix::ShortRoomId,
17 trace,
18 utils::{BoolExt, ReadyExt, time::now_secs},
19};
20
21use super::{
22 DEQUEUE_LIMIT, EDU_LIMIT, NewEvents, RetryAction, TransactionStatus, TransactionStatuses,
23 split::{Split, pdu_room},
24};
25use crate::{
26 federation::ShouldAttempt,
27 sending::{
28 Destination, EduBuf, EduVec, SendingEvent, Service,
29 data::{Key, QueueItem},
30 },
31};
32
33#[derive(Default)]
38struct Selected {
39 shipped: EduVec,
40 overflow: Vec<EduBuf>,
41}
42
43#[derive(Debug, PartialEq, Eq)]
44pub(super) enum Selection {
45 Events(Vec<QueueItem>),
46 Slice(Vec<QueueItem>, Split),
47 Parked {
48 until: u64,
49 },
50 Busy,
51 Refused {
52 earliest_retry: SystemTime,
53 },
54}
55
56enum Current {
57 Ready {
58 replay: bool,
59 },
60 Split(Split),
61 Busy,
62 Refused {
63 earliest_retry: SystemTime,
64 },
65}
66
67impl Selected {
68 fn push(&mut self, edu: EduBuf, events_len: &AtomicUsize) {
73 if !self.overflow.is_empty() || events_len.fetch_add(1, Ordering::Relaxed) >= EDU_LIMIT {
74 self.overflow.push(edu);
75 } else {
76 self.shipped.push(edu);
77 }
78 }
79}
80
81#[implement(Service)]
82#[tracing::instrument(
83 name = "select",
84 level = "debug",
85 skip_all,
86 fields(
87 ?dest,
88 new_events = %new_events.len(),
89 ),
90)]
91pub(super) async fn select_events(
92 &self,
93 dest: &Destination,
94 new_events: NewEvents,
95 statuses: &mut TransactionStatuses,
96) -> Result<Selection> {
97 let retry_action = if matches!(dest, Destination::Appservice(_))
98 && new_events
99 .iter()
100 .any(|(_, event)| matches!(event, SendingEvent::Flush))
101 {
102 RetryAction::Force
103 } else {
104 RetryAction::None
105 };
106
107 let current = self
108 .select_events_current(dest, statuses, retry_action)
109 .await;
110
111 let retry = match current {
112 | Current::Busy => return Ok(Selection::Busy),
113 | Current::Refused { earliest_retry } =>
114 return Ok(Selection::Refused { earliest_retry }),
115 | Current::Ready { replay } => replay,
116 | Current::Split(split) => match self.slice(dest, split).await {
117 | Some((items, split)) => return Ok(Selection::Slice(items, split)),
118 | None => true,
119 },
120 };
121
122 if retry {
123 let active: Vec<_> = self.db.active_requests_for(dest).collect().await;
124
125 if !active.is_empty() {
126 return Ok(Selection::Events(active));
127 }
128 }
129
130 let _cork = self.db.db.cork();
131 let Destination::Federation(server) = dest else {
132 return Ok(Selection::Events(self.claim_new(new_events, &[]).await));
133 };
134
135 let selection = self
136 .federation_batch(dest, server, new_events)
137 .await;
138
139 Ok(selection)
140}
141
142#[implement(Service)]
143async fn select_events_current(
144 &self,
145 dest: &Destination,
146 statuses: &mut TransactionStatuses,
147 retry_action: RetryAction,
148) -> Current {
149 if let Destination::Federation(server) = dest
151 && let ShouldAttempt::No { earliest_retry } = self
152 .services
153 .federation
154 .should_attempt(server)
155 .await
156 {
157 return Current::Refused { earliest_retry };
158 }
159
160 let Some(status) = statuses.get_mut(dest) else {
161 statuses.insert(dest.clone(), TransactionStatus::Running { tries: 0 });
162 return Current::Ready { replay: false };
163 };
164
165 let current = self.transition(dest, status, retry_action);
166
167 if matches!(current, Current::Ready { replay: true } | Current::Split(_)) {
168 self.clear_stalled(dest);
169 }
170
171 current
172}
173
174#[implement(Service)]
178fn transition(
179 &self,
180 dest: &Destination,
181 status: &mut TransactionStatus,
182 retry_action: RetryAction,
183) -> Current {
184 let remaining = self.push_backoff_remaining(Some(&*status));
185
186 match status {
187 | TransactionStatus::Retrying { .. } if matches!(dest, Destination::Push(..)) =>
188 Current::Busy,
189 | TransactionStatus::Running { tries }
190 | TransactionStatus::RunningForceRetry { tries } => {
191 if matches!(retry_action, RetryAction::Force) {
192 *status = TransactionStatus::RunningForceRetry { tries: *tries };
193 }
194
195 Current::Busy
196 },
197 | TransactionStatus::Failed { tries, .. } => {
198 let tries = *tries;
199
200 trace!(?dest, tries, ?remaining, "Push destination remains in backoff");
201 if remaining.is_some() {
202 return Current::Busy;
203 }
204
205 *status = TransactionStatus::Retrying { tries };
206 Current::Ready { replay: true }
207 },
208 | TransactionStatus::Pending => {
209 *status = TransactionStatus::Running { tries: 0 };
210 Current::Ready { replay: true }
211 },
212 | TransactionStatus::Retrying { tries } => {
213 *status = TransactionStatus::Running { tries: *tries };
214 Current::Ready { replay: true }
215 },
216 | TransactionStatus::Splitting { tries, .. } => {
217 let tries = *tries;
218
219 launch(status, tries)
220 },
221 }
222}
223
224fn launch(status: &mut TransactionStatus, tries: u32) -> Current {
225 match replace(status, TransactionStatus::Running { tries }) {
226 | TransactionStatus::Splitting { split, .. } => Current::Split(split),
227 | _ => Current::Ready { replay: true },
228 }
229}
230
231#[implement(Service)]
232fn clear_stalled(&self, dest: &Destination) {
233 if let Destination::Federation(server) = dest {
234 self.stalled
235 .lock()
236 .expect("locked")
237 .remove(server);
238 }
239}
240
241#[implement(Service)]
247pub(super) async fn federation_batch(
248 &self,
249 dest: &Destination,
250 server: &ServerName,
251 new_events: NewEvents,
252) -> Selection {
253 let parks: Vec<_> = self.db.parks(server).collect().await;
254 let now = now_secs();
255
256 for park in parks.iter().filter(|park| park.until <= now) {
257 if let Some((items, split)) = self.slice(dest, Split::probe(park.room)).await {
258 return Selection::Slice(items, split);
259 }
260
261 self.db.unpark(server, park.room);
262 }
263
264 let skip: Vec<_> = parks.iter().map(|park| park.room).collect();
265 let items = self.claim_new(new_events, &skip).await;
266 let items = self.with_edus(dest, server, items, &skip).await;
267 let until = parks
268 .iter()
269 .map(|park| park.until)
270 .filter(|&until| until > now)
271 .min();
272
273 match until {
274 | Some(until) if items.is_empty() => Selection::Parked { until },
275 | _ => Selection::Events(items),
276 }
277}
278
279#[implement(Service)]
284async fn claim_new(&self, new_events: NewEvents, skip: &[ShortRoomId]) -> Vec<QueueItem> {
285 let items: Vec<_> = self
286 .db
287 .retain_queued(new_events)
288 .ready_filter(|item| {
289 matches!(item.1, SendingEvent::Flush).is_false()
290 && pdu_room(item).is_none_or(|room| !skip.contains(&room))
291 })
292 .collect()
293 .await;
294
295 self.db.mark_as_active(items.iter());
296 items
297}
298
299#[implement(Service)]
304async fn with_edus(
305 &self,
306 dest: &Destination,
307 server_name: &ServerName,
308 items: Vec<QueueItem>,
309 skip: &[ShortRoomId],
310) -> Vec<QueueItem> {
311 let items = if items.is_empty() {
312 self.resume_queued(dest, skip).await
313 } else {
314 items
315 };
316
317 if self
319 .db
320 .queued_except(dest, skip)
321 .take(1)
322 .count()
323 .await
324 .ne(&0)
325 {
326 return items;
327 }
328
329 let budget_used = items
330 .iter()
331 .filter(|(_, event)| matches!(event, SendingEvent::Edu(_)))
332 .count();
333
334 let edus = self.select_edus(server_name, budget_used).await;
335
336 append_edus(items, edus)
337}
338
339fn append_edus(
340 mut items: Vec<QueueItem>,
341 edus: impl Iterator<Item = QueueItem>,
342) -> Vec<QueueItem> {
343 items.extend(edus);
344 items
345}
346
347#[implement(Service)]
352#[tracing::instrument(level = "trace", skip_all)]
353pub(super) async fn resume_queued(
354 &self,
355 dest: &Destination,
356 skip: &[ShortRoomId],
357) -> Vec<QueueItem> {
358 let queued: Vec<_> = self
359 .db
360 .queued_except(dest, skip)
361 .take(DEQUEUE_LIMIT)
362 .collect()
363 .await;
364
365 self.db.mark_as_active(queued.iter());
366 queued
367}
368
369#[implement(Service)]
370#[tracing::instrument(name = "edus", level = "debug", skip_all)]
371pub(super) async fn select_edus(
372 &self,
373 server_name: &ServerName,
374 budget_used: usize,
375) -> impl Iterator<Item = QueueItem> {
376 let since = self.db.get_latest_educount(server_name).await;
377 let since_upper = self.services.globals.current_count();
378
379 if since == since_upper {
381 return keyed(Vec::new().into_iter(), EduVec::new());
382 }
383
384 let batch = (since, since_upper);
385
386 debug_assert!(batch.0 <= batch.1, "since range must not be negative");
387
388 let events_len = AtomicUsize::new(budget_used);
389 let max_edu_count = AtomicU64::new(since);
390 let device_changes =
391 self.select_edus_device_changes(server_name, batch, &max_edu_count, &events_len);
392
393 let receipts = self
394 .server
395 .config
396 .allow_outgoing_read_receipts
397 .then_async(|| {
398 self.select_edus_receipts(server_name, batch, &max_edu_count, &events_len)
399 });
400
401 let presence = self
402 .server
403 .config
404 .allow_outgoing_presence
405 .then_async(|| {
406 self.select_edus_presence(server_name, batch, &max_edu_count, &events_len)
407 });
408
409 let (device_changes, receipts, presence) = join3(device_changes, receipts, presence).await;
410 let receipts = receipts.unwrap_or_default();
411
412 let durable_len = device_changes
415 .shipped
416 .len()
417 .saturating_add(receipts.shipped.len());
418
419 let events: EduVec = device_changes
420 .shipped
421 .into_iter()
422 .chain(receipts.shipped)
423 .chain(presence.flatten())
424 .collect();
425
426 debug_assert!(budget_used.saturating_add(events.len()) <= EDU_LIMIT, "exceeded edus limit");
427
428 let overflow: Vec<SendingEvent> = device_changes
430 .overflow
431 .into_iter()
432 .chain(receipts.overflow)
433 .map(SendingEvent::Edu)
434 .collect();
435
436 if !overflow.is_empty() {
437 let dest = Destination::Federation(server_name.to_owned());
438
439 self.db
440 .queue_requests(overflow.iter().map(|event| (event, &dest)));
441 }
442
443 let keys = self
446 .db
447 .persist_active_edus(server_name, &events[..durable_len]);
448
449 let last_count = max_edu_count.load(Ordering::Acquire);
450 if last_count > since {
451 self.db
452 .set_latest_educount(server_name, last_count);
453 }
454
455 keyed(keys, events)
456}
457
458fn keyed(keys: impl Iterator<Item = Key>, edus: EduVec) -> impl Iterator<Item = QueueItem> {
463 keys.chain(repeat_with(Key::new))
464 .zip(edus)
465 .map(|(key, edu)| (key, SendingEvent::Edu(edu)))
466}
467
468pub(super) fn edu_buf(edu: &Edu) -> EduBuf {
470 let mut buf = EduBuf::new(); serde_json::to_writer(&mut buf, edu).expect("EDU serializes to JSON");
473 buf
474}