tuwunel_service/federation/
rank.rs1use futures::StreamExt;
8use ruma::OwnedServerName;
9use tuwunel_core::{debug_warn, implement, smallvec::SmallVec, utils::IterStream};
10
11use super::ShouldAttempt;
12
13pub type Candidates = SmallVec<[OwnedServerName; 3]>;
19
20type Verdicts = SmallVec<[(OwnedServerName, ShouldAttempt); 3]>;
23
24#[derive(Clone, Copy, Debug)]
29pub enum WhenAllBackedOff {
30 Attempt,
33
34 Fail,
36}
37
38#[implement(super::Service)]
44pub async fn rank_candidates(
45 &self,
46 eligible: Candidates,
47 when_all: WhenAllBackedOff,
48) -> Candidates {
49 let verdicts: Verdicts = eligible
50 .into_iter()
51 .stream()
52 .then(async |server| {
53 let verdict = self.should_attempt(&server).await;
54 (server, verdict)
55 })
56 .collect()
57 .await;
58
59 rank_from_verdicts(verdicts, when_all).collect()
60}
61
62fn rank_from_verdicts(
66 mut verdicts: Verdicts,
67 when_all: WhenAllBackedOff,
68) -> impl Iterator<Item = OwnedServerName> {
69 let all_backed_off = verdicts
70 .iter()
71 .all(|(_, verdict)| matches!(verdict, ShouldAttempt::No { .. }));
72
73 let keep_backed_off = all_backed_off && matches!(when_all, WhenAllBackedOff::Attempt);
74
75 if keep_backed_off && !verdicts.is_empty() {
76 debug_warn!(
77 n = verdicts.len(),
78 "All candidates backed off via peer_status; attempting anyway"
79 );
80 }
81
82 verdicts.sort_by_key(|(_, verdict)| verdict.rank());
83
84 verdicts
85 .into_iter()
86 .filter(move |(_, verdict)| {
87 keep_backed_off || !matches!(verdict, ShouldAttempt::No { .. })
88 })
89 .map(|(server, _)| server)
90}
91
92#[implement(ShouldAttempt)]
94#[inline]
95fn rank(self) -> u8 {
96 match self {
97 | ShouldAttempt::Yes => 0,
98 | ShouldAttempt::Deprioritize => 1,
99 | ShouldAttempt::No { .. } => 2,
100 }
101}
102
103#[cfg(test)]
104mod tests {
105 use std::time::SystemTime;
106
107 use ruma::{OwnedServerName, owned_server_name};
108 use tuwunel_core::smallvec::smallvec;
109
110 use super::{Verdicts, WhenAllBackedOff, rank_from_verdicts};
111 use crate::federation::ShouldAttempt;
112
113 fn no() -> ShouldAttempt { ShouldAttempt::No { earliest_retry: SystemTime::UNIX_EPOCH } }
114
115 fn names(servers: &[OwnedServerName]) -> Vec<&str> {
116 servers.iter().map(AsRef::as_ref).collect()
117 }
118
119 #[test]
120 fn all_yes_preserves_order() {
121 let verdicts: Verdicts = smallvec![
122 (owned_server_name!("a.test"), ShouldAttempt::Yes),
123 (owned_server_name!("b.test"), ShouldAttempt::Yes),
124 (owned_server_name!("c.test"), ShouldAttempt::Yes),
125 ];
126
127 let ranked: Vec<_> = rank_from_verdicts(verdicts, WhenAllBackedOff::Attempt).collect();
128
129 assert_eq!(names(&ranked), ["a.test", "b.test", "c.test"]);
130 }
131
132 #[test]
133 fn drops_backed_off_when_pool_has_alternatives() {
134 let verdicts: Verdicts = smallvec![
135 (owned_server_name!("a.test"), ShouldAttempt::Yes),
136 (owned_server_name!("b.test"), no()),
137 (owned_server_name!("c.test"), ShouldAttempt::Yes),
138 ];
139
140 let ranked: Vec<_> = rank_from_verdicts(verdicts, WhenAllBackedOff::Attempt).collect();
141
142 assert_eq!(names(&ranked), ["a.test", "c.test"]);
143 }
144
145 #[test]
146 fn all_backed_off_attempt_falls_through() {
147 let verdicts: Verdicts = smallvec![
148 (owned_server_name!("a.test"), no()),
149 (owned_server_name!("b.test"), no()),
150 ];
151
152 let ranked: Vec<_> = rank_from_verdicts(verdicts, WhenAllBackedOff::Attempt).collect();
153
154 assert_eq!(names(&ranked), ["a.test", "b.test"]);
155 }
156
157 #[test]
158 fn all_backed_off_fail_returns_empty() {
159 let verdicts: Verdicts = smallvec![
160 (owned_server_name!("a.test"), no()),
161 (owned_server_name!("b.test"), no()),
162 ];
163
164 assert!(
165 rank_from_verdicts(verdicts, WhenAllBackedOff::Fail)
166 .next()
167 .is_none()
168 );
169 }
170
171 #[test]
172 fn deprioritize_ranks_after_yes() {
173 let verdicts: Verdicts = smallvec![
174 (owned_server_name!("d.test"), ShouldAttempt::Deprioritize),
175 (owned_server_name!("y.test"), ShouldAttempt::Yes),
176 (owned_server_name!("n.test"), no()),
177 ];
178
179 let ranked: Vec<_> = rank_from_verdicts(verdicts, WhenAllBackedOff::Attempt).collect();
180
181 assert_eq!(names(&ranked), ["y.test", "d.test"]);
182 }
183}