1use std::{
8 collections::{HashMap, VecDeque},
9 sync::{Arc, Weak},
10};
11
12use bytes::Bytes;
13use futures::{FutureExt, StreamExt, future::BoxFuture, stream::FuturesUnordered};
14use ruma::OwnedServerName;
15use tokio::sync::watch::channel;
16use tuwunel_core::{debug_warn, implement, trace, utils::math::effective_cap};
17
18use super::{
19 Failure, Msg, Opts, Outcome, Service,
20 error::Attempted,
21 inflight::{Inflight, Key, SharedResult},
22};
23
24type FetchFuture<'a> = BoxFuture<'a, (Key, SharedResult)>;
27type FetchFutures<'a> = FuturesUnordered<FetchFuture<'a>>;
28
29#[implement(Service)]
34pub(super) async fn run_worker(self: Arc<Self>) {
35 let mut inflight: HashMap<Key, Inflight> = HashMap::new();
36 let mut pending: VecDeque<Msg> = VecDeque::new();
37 let mut futures: FetchFutures<'_> = FuturesUnordered::new();
38
39 self.work_loop(&mut inflight, &mut pending, &mut futures)
40 .await;
41}
42
43#[implement(Service)]
44async fn work_loop<'a>(
45 &'a self,
46 inflight: &mut HashMap<Key, Inflight>,
47 pending: &mut VecDeque<Msg>,
48 futures: &mut FetchFutures<'a>,
49) {
50 let rx = self.channel.1.clone();
51 while !rx.is_closed() {
52 while let Ok(msg) = rx.try_recv() {
54 self.on_request(msg, inflight, pending, futures);
55 }
56
57 tokio::select! {
58 Some((key, result)) = futures.next() =>
59 self.on_complete(key, result, inflight, pending, futures),
60 msg = rx.recv_async() => match msg {
61 | Ok(msg) => self.on_request(msg, inflight, pending, futures),
62 | Err(_) => break,
63 },
64 }
65 }
66}
67
68#[implement(Service)]
69fn on_request<'a>(
70 &'a self,
71 msg: Msg,
72 inflight: &mut HashMap<Key, Inflight>,
73 pending: &mut VecDeque<Msg>,
74 futures: &FetchFutures<'a>,
75) {
76 let Some(entry) = inflight.get_mut(&msg.key) else {
77 if futures.len() >= self.capacity {
79 pending.push_back(msg);
80 } else {
81 self.dispatch(msg, inflight, futures);
82 }
83
84 return;
85 };
86
87 match entry.interest.upgrade() {
88 | Some(strong) => {
90 msg.reply
91 .send((entry.tx.subscribe(), strong))
92 .ok();
93 },
94 | None => {
97 let interest = Arc::new(());
98 entry.interest = Arc::downgrade(&interest);
99 msg.reply
100 .send((entry.tx.subscribe(), interest))
101 .ok();
102 },
103 }
104}
105
106#[implement(Service)]
107fn dispatch<'a>(
108 &'a self,
109 msg: Msg,
110 inflight: &mut HashMap<Key, Inflight>,
111 futures: &FetchFutures<'a>,
112) {
113 let Msg { key, reply } = msg;
114 let interest = Arc::new(());
115 let (tx, rx) = channel(None);
116
117 if reply.send((rx, interest.clone())).is_err() {
119 return;
120 }
121
122 let opts = key.opts();
123 let weak = Arc::downgrade(&interest);
124 inflight.insert(key.clone(), Inflight {
125 tx,
126 interest: weak.clone(),
127 opts: opts.clone(),
128 });
129
130 self.push_attempt(futures, key, opts, weak);
131}
132
133#[implement(Service)]
137fn push_attempt<'a>(
138 &'a self,
139 futures: &FetchFutures<'a>,
140 key: Key,
141 opts: Arc<Opts>,
142 weak: Weak<()>,
143) {
144 futures.push(async move { (key, self.run_attempts(&opts, &weak).await) }.boxed());
145}
146
147#[implement(Service)]
148fn on_complete<'a>(
149 &'a self,
150 key: Key,
151 result: SharedResult,
152 inflight: &mut HashMap<Key, Inflight>,
153 pending: &mut VecDeque<Msg>,
154 futures: &FetchFutures<'a>,
155) {
156 let Some(entry) = inflight.get(&key) else {
157 return;
158 };
159
160 if matches!(&result, Err(Failure::Cancelled)) && entry.interest.upgrade().is_some() {
163 let opts = entry.opts.clone();
164 let weak = entry.interest.clone();
165 self.push_attempt(futures, key, opts, weak);
166 return;
167 }
168
169 entry.tx.send(Some(result)).ok();
170 inflight.remove(&key);
171
172 while futures.len() < self.capacity {
175 let Some(msg) = pending.pop_front() else {
176 break;
177 };
178
179 self.on_request(msg, inflight, pending, futures);
180 }
181}
182
183#[implement(Service)]
184#[tracing::instrument(
185 name = "attempts",
186 level = "debug",
187 skip_all,
188 fields(
189 op = ?opts.op,
190 room_id = ?opts.room_id,
191 event_id = ?opts.event_id,
192 ),
193)]
194async fn run_attempts(&self, opts: &Opts, interest: &Weak<()>) -> SharedResult {
195 let candidates = self.select.candidates(opts).await;
196 if candidates.is_empty() {
197 return Err(Failure::NoCandidates);
198 }
199
200 let count = candidates.len();
201 let limit = opts
202 .attempt_limit
203 .map_or(count, |n| n.get().min(count));
204
205 let (config_width, config_rounds) = self
206 .services
207 .try_get()
208 .map_or((0, 0), |services| {
209 let config = &services.server.config;
210
211 (config.fetch_fanout_max_width, config.fetch_fanout_rounds)
212 });
213
214 let max_width = effective_cap(opts.fanout_max_width, config_width);
215 let max_rounds = effective_cap(opts.fanout_rounds, config_rounds);
216
217 let mut attempted: Attempted = Attempted::new();
218 let mut remaining = candidates.into_iter();
219 let mut round: usize = 0;
220
221 while attempted.len() < limit {
222 if interest.strong_count() == 0 {
223 return Err(Failure::Cancelled);
224 }
225
226 if round >= max_rounds {
227 break;
228 }
229
230 let budget = limit.saturating_sub(attempted.len());
231 let width = opts
232 .fanout_growth
233 .round_width(round)
234 .min(max_width)
235 .min(budget);
236
237 let mut racing: FuturesUnordered<_> = remaining
240 .by_ref()
241 .take(width)
242 .map(|server| self.attempt(server, opts))
243 .collect();
244
245 if racing.is_empty() {
246 break;
247 }
248
249 while let Some((server, bytes)) = racing.next().await {
250 let Some(bytes) = bytes else {
251 attempted.push(server);
252
253 if interest.strong_count() == 0 {
254 return Err(Failure::Cancelled);
255 }
256
257 continue;
258 };
259
260 trace!(%server, "fetch satisfied");
261 return Ok(Arc::new(Outcome { bytes, origin: server }));
262 }
263
264 round = round.saturating_add(1);
265 }
266
267 Err(Failure::NotFound { attempted })
268}
269
270#[implement(Service)]
274#[tracing::instrument(
275 name = "attempt",
276 level = "trace",
277 skip_all,
278 fields(%server),
279)]
280async fn attempt(
281 &self,
282 server: OwnedServerName,
283 opts: &Opts,
284) -> (OwnedServerName, Option<Bytes>) {
285 let Some(bytes) = self
286 .transport
287 .fetch_raw(opts.op, &server, opts)
288 .await
289 .inspect_err(|error| debug_warn!(%server, "fetch attempt failed: {error}"))
290 .ok()
291 else {
292 return (server, None);
293 };
294
295 let valid = self
296 .validate(opts, &bytes)
297 .await
298 .inspect_err(|error| debug_warn!(%server, "rejecting poisoned response: {error}"))
299 .is_ok();
300
301 (server, valid.then_some(bytes))
302}