tuwunel_service/federation/
peer.rs1use std::{
11 collections::BTreeMap,
12 time::{Duration, SystemTime, UNIX_EPOCH},
13};
14
15use futures::{Stream, StreamExt};
16use http::StatusCode;
17use ruma::{OwnedServerName, ServerName, api::error::ErrorBody};
18use tuwunel_core::{
19 Error, implement,
20 utils::{
21 stream::{ReadyExt, TryIgnore},
22 time::now_secs,
23 },
24};
25use tuwunel_database::Interfix;
26
27pub(super) const MAX_BACKOFF: Duration = Duration::from_hours(24);
31
32#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
36pub enum Classification {
37 #[default]
39 Transient,
40
41 Permanent,
43}
44
45impl Classification {
46 #[inline]
49 #[must_use]
50 fn from_byte(byte: u8) -> Self {
51 match byte {
52 | 1 => Self::Permanent,
53 | _ => Self::Transient,
54 }
55 }
56}
57
58impl From<Classification> for u8 {
59 #[inline]
60 fn from(c: Classification) -> Self {
61 match c {
62 | Classification::Transient => 0,
63 | Classification::Permanent => 1,
64 }
65 }
66}
67
68#[derive(Clone, Copy, Debug, Eq, PartialEq)]
73pub enum ShouldAttempt {
74 Yes,
76
77 No {
79 earliest_retry: SystemTime,
81 },
82
83 Deprioritize,
86}
87
88pub(super) struct Backoff {
93 pub(super) class: Classification,
95
96 pub(super) anchor_secs: u64,
98
99 pub(super) streak: u32,
101
102 pub(super) now: u64,
104
105 pub(super) window_secs: u64,
107
108 pub(super) grace_secs: u64,
110}
111
112#[derive(Clone, Copy)]
117pub(super) struct Streak {
118 pub(super) class: Classification,
120
121 pub(super) anchor_secs: u64,
123
124 pub(super) oldest_bucket: u64,
126
127 pub(super) latest_bucket: u64,
129}
130
131#[derive(Clone, Copy, Debug)]
136pub struct PeerBackoff {
137 pub class: Classification,
139
140 pub anchor_secs: u64,
142
143 pub oldest_secs: u64,
145
146 pub delay_secs: u64,
148}
149
150#[implement(super::Service)]
154pub async fn record_success(&self, server: &ServerName) {
155 self.statuses
156 .del_prefix(&(server, Interfix))
157 .await;
158}
159
160#[implement(super::Service)]
165#[tracing::instrument(
166 level = "trace",
167 skip(self),
168 fields(
169 %server,
170 ),
171)]
172pub async fn note_peer_alive(&self, server: &ServerName) -> bool {
173 let sad = self.peer_has_failures(server).await;
174
175 if sad {
176 self.statuses
177 .del_prefix(&(server, Interfix))
178 .await;
179 }
180
181 sad
182}
183
184#[implement(super::Service)]
188#[tracing::instrument(
189 level = "trace",
190 skip(self),
191 fields(
192 %server,
193 ),
194)]
195pub async fn peer_has_failures(&self, server: &ServerName) -> bool {
196 self.statuses
197 .stream_prefix_raw(&(server, Interfix))
198 .ignore_err()
199 .ready_any(|_| true)
200 .await
201}
202
203#[implement(super::Service)]
208pub fn record_failure(&self, server: &ServerName, classification: Classification) {
209 let mut value = [0_u8; 9];
211 value[0] = u8::from(classification);
212 value[1..].copy_from_slice(&now_secs().to_be_bytes());
213
214 self.statuses
215 .put_raw((server, self.current_bucket()), value);
216}
217
218#[implement(super::Service)]
223#[tracing::instrument(skip(self), fields(%server), level = "trace")]
224pub async fn should_attempt(&self, server: &ServerName) -> ShouldAttempt {
225 let Some(streak) = self.peer_streak(server).await else {
226 return ShouldAttempt::Yes;
227 };
228
229 attempt_verdict(&self.backoff(streak))
230}
231
232#[implement(super::Service)]
236pub async fn peer_backoff(&self, server: &ServerName) -> Option<PeerBackoff> {
237 self.peer_streak(server)
238 .await
239 .map(|streak| self.peer_backoff_from(streak))
240}
241
242#[implement(super::Service)]
247pub async fn peer_backoffs(&self) -> BTreeMap<OwnedServerName, PeerBackoff> {
248 let window_secs = self.window_secs;
249
250 self.statuses
251 .stream()
252 .ignore_err()
253 .ready_fold(
254 Vec::<(OwnedServerName, Streak)>::new(),
255 |mut runs, ((server, bucket), value): ((&ServerName, u64), &[u8])| {
256 match runs.last_mut() {
257 | Some((last, streak)) if *last == *server =>
258 *streak = fold_streak(window_secs, Some(*streak), bucket, value),
259 | _ => runs
260 .push((server.to_owned(), fold_streak(window_secs, None, bucket, value))),
261 }
262
263 runs
264 },
265 )
266 .await
267 .into_iter()
268 .map(|(server, streak)| (server, self.peer_backoff_from(streak)))
269 .collect()
270}
271
272#[implement(super::Service)]
279pub fn peer_snapshot(
280 &self,
281) -> impl Stream<Item = (&ServerName, SystemTime, Classification)> + Send + '_ {
282 self.statuses.stream().ignore_err().map(
283 move |((server, bucket), value): ((&ServerName, u64), &[u8])| {
284 (server, self.bucket_start(bucket), classify(value))
285 },
286 )
287}
288
289#[implement(super::Service)]
290#[inline]
291#[must_use]
292fn current_bucket(&self) -> u64 {
293 now_secs()
294 .checked_div(self.window_secs.max(1))
295 .unwrap_or(0)
296}
297
298#[implement(super::Service)]
300#[inline]
301#[must_use]
302fn bucket_start(&self, bucket: u64) -> SystemTime {
303 let offset = bucket.saturating_mul(self.window_secs);
304
305 UNIX_EPOCH
306 .checked_add(Duration::from_secs(offset))
307 .unwrap_or(UNIX_EPOCH)
308}
309
310#[implement(super::Service)]
311#[inline]
312#[must_use]
313fn streak(&self, latest_bucket: u64, oldest_bucket: u64) -> u32 {
314 let span = latest_bucket
315 .saturating_sub(oldest_bucket)
316 .saturating_add(1);
317
318 u32::try_from(span)
319 .unwrap_or(u32::MAX)
320 .min(self.n_max)
321}
322
323#[implement(super::Service)]
325async fn peer_streak(&self, server: &ServerName) -> Option<Streak> {
326 let window_secs = self.window_secs;
327
328 self.statuses
329 .stream_prefix(&(server, Interfix))
330 .ignore_err()
331 .ready_fold(None, |state, ((_, bucket), value): ((&ServerName, u64), &[u8])| {
332 Some(fold_streak(window_secs, state, bucket, value))
333 })
334 .await
335}
336
337#[implement(super::Service)]
339fn backoff(&self, run: Streak) -> Backoff {
340 Backoff {
341 class: run.class,
342 anchor_secs: run.anchor_secs,
343 streak: self.streak(run.latest_bucket, run.oldest_bucket),
344 now: now_secs(),
345 window_secs: self.window_secs,
346 grace_secs: self.grace.as_secs(),
347 }
348}
349
350#[implement(super::Service)]
352fn peer_backoff_from(&self, streak: Streak) -> PeerBackoff {
353 PeerBackoff {
354 class: streak.class,
355 anchor_secs: streak.anchor_secs,
356 oldest_secs: streak
357 .oldest_bucket
358 .saturating_mul(self.window_secs),
359 delay_secs: self.backoff(streak).delay_secs(),
360 }
361}
362
363#[must_use]
369pub(super) fn attempt_verdict(backoff: &Backoff) -> ShouldAttempt {
370 let earliest_secs = backoff
371 .anchor_secs
372 .saturating_add(backoff.delay_secs());
373
374 if backoff.now >= earliest_secs {
375 return ShouldAttempt::Yes;
376 }
377
378 ShouldAttempt::No {
379 earliest_retry: UNIX_EPOCH
380 .checked_add(Duration::from_secs(earliest_secs))
381 .unwrap_or_else(SystemTime::now),
382 }
383}
384
385impl Backoff {
386 #[must_use]
392 pub(super) fn delay_secs(&self) -> u64 {
393 let max_backoff = MAX_BACKOFF.as_secs();
394
395 match self.class {
396 | Classification::Permanent => max_backoff,
397 | Classification::Transient if self.streak <= 1 && self.grace_secs != 0 =>
398 self.grace_secs.min(max_backoff),
399 | Classification::Transient => self
400 .window_secs
401 .saturating_mul(u64::from(self.streak))
402 .saturating_mul(u64::from(self.streak))
403 .min(max_backoff),
404 }
405 }
406}
407
408#[must_use]
413pub(super) fn fold_streak(
414 window_secs: u64,
415 state: Option<Streak>,
416 bucket: u64,
417 value: &[u8],
418) -> Streak {
419 let anchor_secs = failure_secs(value).unwrap_or_else(|| bucket.saturating_mul(window_secs));
420
421 let oldest_bucket = state.map_or(bucket, |streak| streak.oldest_bucket);
422
423 Streak {
424 class: classify(value),
425 anchor_secs,
426 oldest_bucket,
427 latest_bucket: bucket,
428 }
429}
430
431#[inline]
432#[must_use]
433pub(super) fn classify(bytes: &[u8]) -> Classification {
438 bytes
439 .first()
440 .copied()
441 .map_or(Classification::Transient, Classification::from_byte)
442}
443
444#[must_use]
449pub(super) fn failure_secs(bytes: &[u8]) -> Option<u64> {
450 bytes
451 .get(1..9)
452 .and_then(|tail| tail.try_into().ok())
453 .map(u64::from_be_bytes)
454}
455
456#[must_use]
462pub fn is_content_rejection(error: &Error) -> bool { classify_error(error).is_none() }
463
464#[must_use]
470pub(super) fn classify_error(error: &Error) -> Option<Classification> {
471 let Error::Federation(_, response) = error else {
472 return Some(Classification::Transient);
473 };
474
475 let status = response.status_code;
476
477 match status {
478 | _ if status == StatusCode::GONE => Some(Classification::Permanent),
479 | _ if status.is_server_error() || status == StatusCode::TOO_MANY_REQUESTS =>
480 Some(Classification::Transient),
481 | _ if matches!(response.body, ErrorBody::NotJson { .. }) =>
482 Some(Classification::Transient),
483 | _ => None,
484 }
485}