tuwunel_service/sending/
worker.rs1#[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#[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
55async 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 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 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
159fn 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}