Skip to main content

tuwunel_service/fetcher/
worker.rs

1//! The fetch worker loop: one task owning every in-flight fetch, lock-free.
2//!
3//! [`Service::run_worker`] coalesces incoming requests, dispatches attempts up
4//! to the capacity bound, defers the rest, and broadcasts each outcome to its
5//! subscribers.
6
7use 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
24/// One in-flight fetch, borrowing the worker for the service's lifetime and
25/// yielding its key alongside the result so the worker can route it.
26type FetchFuture<'a> = BoxFuture<'a, (Key, SharedResult)>;
27type FetchFutures<'a> = FuturesUnordered<FetchFuture<'a>>;
28
29/// Runs the single owner task for request, deferral, and in-flight state.
30///
31/// The request map, pending queue, and fetch futures are owned exclusively by
32/// this task, so no lock guards them.
33#[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		// Coalesce co-arriving callers before any completion evicts their entry.
53		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		// no in-flight request for this key: dispatch, or defer at the cap
78		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		// live callers: coalesce onto the running future
89		| Some(strong) => {
90			msg.reply
91				.send((entry.tx.subscribe(), strong))
92				.ok();
93		},
94		// every prior caller dropped and the future is draining toward
95		// Cancelled; re-arm so it revives at its next attempt boundary
96		| 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	// caller already gone: do not touch the network
118	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/// Push one attempt future onto the worker's set, yielding its key with the
134/// result so the worker can route it back. The lone construction site for the
135/// borrowed-future shape, shared by the fresh-dispatch and re-arm paths.
136#[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	// a fresh caller re-armed after the future gave up: revive from the
161	// retained opts rather than publishing a stale Cancelled
162	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	// A freed slot re-admits deferred requests through on_request so a deferred
173	// same-key pair coalesces instead of double-dispatching.
174	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		// race this round's window; the first valid response wins and dropping the
238		// set cancels the losing requests in flight
239		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/// Fetch one candidate and validate it: `Some(bytes)` on a clean response,
271/// `None` on a transport error or a poisoned body. A miss is logged, never
272/// fatal, so it cannot cancel a sibling racing the same round.
273#[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}