Skip to main content

tuwunel_service/sending/
worker.rs

1#[cfg(test)]
2mod tests;
3
4use std::{
5	hash::{BuildHasher, BuildHasherDefault, DefaultHasher, Hash},
6	mem::take,
7	sync::Arc,
8};
9
10use futures::FutureExt;
11use ruma::ServerName;
12use tokio::task::{JoinError, JoinSet, unconstrained};
13use tuwunel_core::{
14	Result, debug, err, error, implement,
15	utils::{available_parallelism, math::usize_from_u64_truncated},
16};
17
18use super::{Destination, Msg, Service, dest::DestinationRef};
19use crate::{Args, Service as _};
20
21/// Run the sender workers to completion.
22///
23/// One worker is spawned per channel and joined; a panic among them is
24/// returned so the manager restarts the service. The suppressed-push flush
25/// tasks are then aborted and joined so a panic among them is still reported.
26#[implement(Service)]
27pub(super) async fn run(self: Arc<Self>) -> Result {
28	let senders =
29		self.channels
30			.iter()
31			.enumerate()
32			.fold(JoinSet::new(), |mut senders, (id, _)| {
33				let worker = self.clone().sender(id);
34				let worker = if self.unconstrained() {
35					unconstrained(worker).left_future()
36				} else {
37					worker.right_future()
38				};
39
40				senders.spawn_on(worker, self.server.runtime());
41				senders
42			});
43
44	let result = join_senders(senders).await;
45	let mut flushes = take(&mut *self.flushes.lock().expect("locked"));
46
47	flushes.abort_all();
48	while let Some(result) = flushes.join_next().await {
49		log_flush(result);
50	}
51
52	result
53}
54
55/// Join the sender workers, stopping the rest on the first failure or panic.
56async fn join_senders(mut senders: JoinSet<Result>) -> Result {
57	while let Some(ret) = senders.join_next_with_id().await {
58		let (id, result) = match ret {
59			| Ok((id, result)) => (id, result),
60			| Err(error) => (error.id(), Err(error.into())),
61		};
62
63		if let Err(error) = result {
64			error!(?id, ?error, "sender worker failed");
65			senders.shutdown().await;
66			return Err(error);
67		}
68
69		debug!(?id, "sender worker finished");
70	}
71
72	Ok(())
73}
74
75#[implement(Service)]
76pub(super) fn close(&self) {
77	self.flushes.lock().expect("locked").abort_all();
78
79	for (sender, _) in &self.channels {
80		sender.close();
81	}
82}
83
84#[implement(Service)]
85pub(super) fn dispatch(&self, msg: Msg) -> Result {
86	let shard = self.shard_id(&msg.dest);
87	let sender = &self
88		.channels
89		.get(shard)
90		.expect("missing sender worker channels")
91		.0;
92
93	debug_assert!(!sender.is_full(), "channel full");
94	debug_assert!(!sender.is_closed(), "channel closed");
95	sender.send(msg).map_err(|e| err!("{e}"))
96}
97
98#[implement(Service)]
99pub(super) fn shard_id(&self, dest: &Destination) -> usize {
100	shard(&dest.borrowed(), self.channels.len())
101}
102
103#[implement(Service)]
104#[inline]
105pub(super) fn federation_shard_id(&self, server: &ServerName) -> usize {
106	shard(&DestinationRef::Federation(server), self.channels.len())
107}
108
109fn shard(dest: &impl Hash, count: usize) -> usize {
110	if count <= 1 {
111		return 0;
112	}
113
114	let hash = BuildHasherDefault::<DefaultHasher>::default().hash_one(dest);
115
116	usize_from_u64_truncated(hash)
117		.overflowing_rem(count)
118		.0
119}
120
121pub(super) fn num_senders(args: &Args<'_>) -> usize {
122	const MIN_SENDERS: usize = 1;
123	let max_senders = args
124		.server
125		.metrics
126		.num_workers()
127		.min(available_parallelism());
128
129	// The config default 0 clamps to one sender; multiple senders are experimental.
130	args.server
131		.config
132		.sender_workers
133		.clamp(MIN_SENDERS, max_senders)
134}
135
136#[implement(Service)]
137pub(super) fn spawn_flush<F>(&self, flush: F)
138where
139	F: Future<Output = ()> + Send + 'static,
140{
141	let mut flushes = self.flushes.lock().expect("locked");
142
143	// Shutdown drains the set at the end of `run`; a flush spawned after that
144	// would never be joined. A restart re-enters `run` and drains again.
145	if !self.server.is_running() {
146		return;
147	}
148
149	reap_flushes(&mut flushes);
150	flushes.spawn_on(flush, self.server.runtime());
151}
152
153fn reap_flushes(flushes: &mut JoinSet<()>) {
154	while let Some(result) = flushes.try_join_next() {
155		log_flush(result);
156	}
157}
158
159// A flush that panicked is reported here or nowhere; a cancelled one is the
160// shutdown path.
161fn log_flush(result: Result<(), JoinError>) {
162	if let Err(error) = result
163		&& error.is_panic()
164	{
165		error!(?error, "Suppressed push flush panicked");
166	}
167}