tuwunel_service/fetcher/
select.rs1use std::sync::Arc;
8
9use async_trait::async_trait;
10use futures::{Stream, StreamExt, stream::empty};
11use ruma::{EventId, OwnedServerName, RoomId, ServerName};
12use tuwunel_core::{
13 arrayvec::ArrayVec,
14 implement,
15 utils::{BoolExt, IterStream, ReadyExt, StreamTools, rand::index},
16};
17
18use super::{Op, Opts};
19use crate::{
20 federation::{Candidates, WhenAllBackedOff},
21 services::OnceServices,
22};
23
24const ROUTE_FANOUT: usize = 5;
27
28#[async_trait]
33pub(super) trait Select: Send + Sync {
34 async fn candidates(&self, opts: &Opts) -> Candidates;
39}
40
41pub(super) struct RoomCandidates {
45 pub(super) services: Arc<OnceServices>,
47}
48
49#[async_trait]
50impl Select for RoomCandidates {
51 #[tracing::instrument(
52 level = "trace",
53 skip_all,
54 fields(
55 room_id = ?opts.room_id,
56 ),
57 )]
58 async fn candidates(&self, opts: &Opts) -> Candidates {
59 if !opts.candidates.is_empty() {
60 return self.ranked_override(opts).await;
61 }
62
63 let authority = self.authority_server(opts).await;
64
65 let mxid_hosts = [
66 opts.event_id
67 .as_deref()
68 .and_then(EventId::server_name),
69 opts.room_id
70 .as_deref()
71 .and_then(RoomId::server_name),
72 ]
73 .into_iter()
74 .flatten()
75 .map(ToOwned::to_owned);
76
77 let popular = match opts.room_id.as_deref() {
78 | None => empty::<OwnedServerName>().right_stream(),
79 | Some(room_id) => self
80 .route_by_popularity(room_id)
81 .await
82 .left_stream(),
83 };
84
85 let eligible = opts
86 .hint
87 .clone()
88 .into_iter()
89 .chain(authority)
90 .stream()
91 .chain(popular)
92 .chain(mxid_hosts.stream())
93 .ready_filter(|server| self.is_eligible(server));
94
95 self.rank_unique(eligible).await
96 }
97}
98
99#[implement(RoomCandidates)]
103#[tracing::instrument(level = "trace", skip_all)]
104async fn ranked_override(&self, opts: &Opts) -> Candidates {
105 let eligible = opts
106 .hint
107 .iter()
108 .chain(opts.candidates.iter())
109 .filter(|&server| self.is_eligible(server))
110 .cloned()
111 .stream();
112
113 self.rank_unique(eligible).await
114}
115
116#[implement(RoomCandidates)]
119async fn rank_unique<S>(&self, eligible: S) -> Candidates
120where
121 S: Stream<Item = OwnedServerName> + Send,
122{
123 let ordered: Candidates = eligible
124 .ready_fold(Candidates::new(), push_unique)
125 .await;
126
127 self.services
128 .federation
129 .rank_candidates(ordered, WhenAllBackedOff::Attempt)
130 .await
131}
132
133fn push_unique(mut ordered: Candidates, server: OwnedServerName) -> Candidates {
135 if !ordered.contains(&server) {
136 ordered.push(server);
137 }
138
139 ordered
140}
141
142#[implement(RoomCandidates)]
145#[tracing::instrument(level = "trace", skip_all)]
146async fn authority_server(&self, opts: &Opts) -> Option<OwnedServerName> {
147 let room_id = opts.room_id.as_deref()?;
148
149 matches!(opts.op, Op::AuthEvent | Op::AuthChain)
150 .then_async(|| {
151 self.services
152 .state_cache
153 .most_powerful_user_server(room_id)
154 })
155 .await
156 .flatten()
157}
158
159#[implement(RoomCandidates)]
165#[tracing::instrument(level = "trace", skip_all)]
166async fn route_by_popularity<'a>(
167 &'a self,
168 room_id: &'a RoomId,
169) -> impl Stream<Item = OwnedServerName> + Send + 'a {
170 let sampled: ArrayVec<OwnedServerName, ROUTE_FANOUT> = self
171 .services
172 .state_cache
173 .room_members(room_id)
174 .sample_by(|user| user.server_name().to_owned())
175 .await;
176
177 if sampled.is_empty() {
178 return self
179 .services
180 .state_cache
181 .room_servers(room_id)
182 .map(ToOwned::to_owned)
183 .right_stream();
184 }
185
186 sampled.into_iter().stream().left_stream()
187}
188
189#[implement(RoomCandidates)]
194#[expect(dead_code)]
195async fn route_uniformly<'a>(
196 &'a self,
197 room_id: &'a RoomId,
198) -> impl Stream<Item = OwnedServerName> + Send + 'a {
199 let count = self
200 .services
201 .state_cache
202 .room_servers(room_id)
203 .count()
204 .await;
205
206 let offset = index(count);
207
208 self.services
209 .state_cache
210 .room_servers(room_id)
211 .map(ToOwned::to_owned)
212 .skip(offset)
213 .take(ROUTE_FANOUT)
214}
215
216#[implement(RoomCandidates)]
217fn is_eligible(&self, server: &ServerName) -> bool {
218 !self.services.globals.server_is_ours(server)
219 && !self
220 .services
221 .server
222 .config
223 .is_forbidden_remote_server_name(server)
224}
225
226#[cfg(test)]
227mod tests {
228 use ruma::owned_server_name;
229
230 use super::{Candidates, push_unique};
231
232 #[test]
233 fn push_unique_keeps_first_occurrence() {
234 let pool = [
235 owned_server_name!("a.test"),
236 owned_server_name!("b.test"),
237 owned_server_name!("a.test"),
238 owned_server_name!("c.test"),
239 owned_server_name!("b.test"),
240 ];
241
242 let deduped: Candidates = pool
243 .into_iter()
244 .fold(Candidates::new(), push_unique);
245
246 let names: Vec<&str> = deduped.iter().map(AsRef::as_ref).collect();
247
248 assert_eq!(names, ["a.test", "b.test", "c.test"]);
249 }
250}