1mod data;
2mod dest;
3mod device;
4mod sender;
5#[cfg(test)]
6mod tests;
7mod worker;
8
9use std::{
10 collections::HashMap,
11 io::Write,
12 iter::{once, repeat_with},
13 sync::{Arc, Mutex as StdMutex},
14 time::{Duration, Instant},
15};
16
17use async_trait::async_trait;
18use futures::{Stream, StreamExt};
19use loole::unbounded;
20use ruma::{OwnedServerName, RoomId, ServerName, UserId};
21use tokio::task::JoinSet;
22use tuwunel_core::{
23 Result, Server, debug_warn, implement,
24 smallvec::SmallVec,
25 utils::{IterStream, ReadyExt, TryReadyExt, result::LogErr},
26};
27
28use self::worker::num_senders;
29pub use self::{
30 data::{Data, Park},
31 dest::Destination,
32 sender::{EDU_LIMIT, PDU_LIMIT},
33};
34use crate::rooms::timeline::RawPduId;
35
36type StalledDestinations = StdMutex<HashMap<OwnedServerName, Option<Instant>>>;
37
38pub struct Service {
44 pub db: Data,
45 server: Arc<Server>,
46 services: Arc<crate::services::OnceServices>,
47 channels: Vec<(loole::Sender<Msg>, loole::Receiver<Msg>)>,
48
49 flushes: StdMutex<JoinSet<()>>,
51
52 stalled: StalledDestinations,
54}
55
56#[expect(clippy::module_name_repetitions)]
60#[derive(Clone, Debug, PartialEq, Eq, Hash)]
61pub enum SendingEvent {
62 Pdu(RawPduId),
64
65 Edu(EduBuf),
67
68 ToDevice(EduBuf),
70
71 DeviceListChanged(EduBuf),
73
74 BadgeRefresh,
78
79 Flush,
81}
82
83#[derive(Clone, Debug, PartialEq, Eq)]
84struct Msg {
85 dest: Destination,
86 event: SendingEvent,
87 queue_id: Vec<u8>,
88}
89
90pub type EduBuf = SmallVec<[u8; EDU_BUF_CAP]>;
95
96pub type EduVec = SmallVec<[EduBuf; EDU_VEC_CAP]>;
100
101const EDU_BUF_CAP: usize = 128 - 16;
102const EDU_VEC_CAP: usize = 1;
103
104const TAG_TO_DEVICE: u8 = 0x01;
107const TAG_DEVICE_LIST_CHANGED: u8 = 0x02;
108const TAG_BADGE_REFRESH: u8 = 0x03;
109const TAG_PREFIX_LEN: usize = 1 + size_of::<u64>();
110
111#[async_trait]
112impl crate::Service for Service {
113 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
114 let channels = repeat_with(unbounded)
115 .take(num_senders(args))
116 .collect();
117
118 Ok(Arc::new(Self {
119 db: Data::new(args),
120 server: args.server.clone(),
121 services: args.services.clone(),
122 channels,
123 flushes: JoinSet::new().into(),
124 stalled: HashMap::new().into(),
125 }))
126 }
127
128 async fn worker(self: Arc<Self>) -> Result { self.run().await }
129
130 async fn interrupt(&self) { self.close(); }
131
132 fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
133
134 fn unconstrained(&self) -> bool { true }
135}
136
137#[implement(Service)]
141#[tracing::instrument(skip(self, pdu_id, user, pushkey), level = "debug")]
142pub fn send_pdu_push(&self, pdu_id: &RawPduId, user: &UserId, pushkey: String) -> Result {
143 let dest = Destination::Push(user.to_owned(), pushkey);
144 let event = SendingEvent::Pdu(*pdu_id);
145 let _cork = self.db.db.cork();
146
147 self.queue_and_dispatch(dest, event)
148}
149
150#[implement(Service)]
151fn queue_and_dispatch(&self, dest: Destination, event: SendingEvent) -> Result {
152 let queue_id = self
153 .db
154 .queue_requests(once((&event, &dest)))
155 .pop()
156 .expect("request queue key");
157
158 self.dispatch(Msg { dest, event, queue_id })
159}
160
161#[implement(Service)]
165#[tracing::instrument(level = "debug", skip(self))]
166pub async fn refresh_push_badge(&self, user_id: &UserId) -> Result {
167 self.services
168 .pusher
169 .get_pushkeys(user_id)
170 .map(Ok)
171 .ready_try_for_each(|pushkey| {
172 let dest = Destination::Push(user_id.to_owned(), pushkey.to_owned());
173
174 self.queue_and_dispatch(dest, SendingEvent::BadgeRefresh)
175 })
176 .await
177}
178
179#[implement(Service)]
183#[tracing::instrument(skip(self), level = "debug")]
184pub fn send_pdu_appservice(&self, appservice_id: String, pdu_id: RawPduId) -> Result {
185 let dest = Destination::Appservice(appservice_id);
186 let event = SendingEvent::Pdu(pdu_id);
187 let _cork = self.db.db.cork();
188
189 self.queue_and_dispatch(dest, event)
190}
191
192#[implement(Service)]
196#[tracing::instrument(skip(self, room_id, pdu_id), level = "debug")]
197pub async fn send_pdu_room(&self, room_id: &RoomId, pdu_id: &RawPduId) -> Result {
198 let servers = self
199 .services
200 .state_cache
201 .remote_room_servers(room_id);
202
203 self.send_pdu_servers(servers, pdu_id).await
204}
205
206#[implement(Service)]
210#[tracing::instrument(skip(self, servers, pdu_id), level = "debug")]
211pub async fn send_pdu_servers<'a, S>(&self, servers: S, pdu_id: &RawPduId) -> Result
212where
213 S: Stream<Item = &'a ServerName> + Send + 'a,
214{
215 self.queue_and_dispatch_servers(servers, SendingEvent::Pdu(*pdu_id))
216 .await
217}
218
219#[implement(Service)]
220async fn queue_and_dispatch_servers<'a, S>(&self, servers: S, event: SendingEvent) -> Result
221where
222 S: Stream<Item = &'a ServerName> + Send + 'a,
223{
224 let requests: Vec<_> = servers
225 .map(|server| (event.clone(), Destination::Federation(server.to_owned())))
226 .collect()
227 .await;
228
229 let _cork = self.db.db.cork();
230 let keys = self
231 .db
232 .queue_requests(requests.iter().map(|(event, dest)| (event, dest)));
233
234 requests
235 .into_iter()
236 .zip(keys)
237 .try_for_each(|((event, dest), queue_id)| self.dispatch(Msg { dest, event, queue_id }))
238}
239
240#[implement(Service)]
244#[tracing::instrument(skip(self, server, serialized), level = "debug")]
245pub fn send_edu_server(&self, server: &ServerName, serialized: EduBuf) -> Result {
246 let dest = Destination::Federation(server.to_owned());
247 let event = SendingEvent::Edu(serialized);
248 let _cork = self.db.db.cork();
249
250 self.queue_and_dispatch(dest, event)
251}
252
253#[implement(Service)]
257#[tracing::instrument(skip(self, room_id, serialized), level = "debug")]
258pub async fn send_edu_room(&self, room_id: &RoomId, serialized: EduBuf) -> Result {
259 let servers = self
260 .services
261 .state_cache
262 .remote_room_servers(room_id);
263
264 self.send_edu_servers(servers, serialized).await
265}
266
267#[implement(Service)]
271#[tracing::instrument(skip(self, servers, serialized), level = "debug")]
272pub async fn send_edu_servers<'a, S>(&self, servers: S, serialized: EduBuf) -> Result
273where
274 S: Stream<Item = &'a ServerName> + Send + 'a,
275{
276 self.queue_and_dispatch_servers(servers, SendingEvent::Edu(serialized))
277 .await
278}
279
280#[implement(Service)]
287#[expect(closure_returning_async_block)]
290#[tracing::instrument(skip(self, serializer), level = "debug")]
291pub async fn send_edu_room_appservices<'a, F>(&self, room_id: &RoomId, serializer: F) -> Result
292where
293 F: Fn(&mut dyn Write) -> Result + Send + 'a,
294 &'a F: Send + Sync,
295{
296 self.services
297 .appservice
298 .read()
299 .await
300 .values()
301 .stream()
302 .filter(|&appservice| async move {
303 if !appservice.registration.receive_ephemeral {
304 return false;
305 }
306
307 if appservice.rooms.is_match(room_id.as_str()) {
308 return true;
309 }
310
311 if self
312 .services
313 .state_cache
314 .appservice_in_room(room_id, appservice)
315 .await
316 {
317 return true;
318 }
319
320 self.services
321 .alias
322 .local_aliases_for_room(room_id)
323 .ready_any(|room_alias| appservice.aliases.is_match(room_alias.as_str()))
324 .await
325 })
326 .map(Ok)
327 .ready_try_for_each(|appservice| {
328 let mut buf = EduBuf::new(); serializer(&mut buf)?;
331 self.send_edu_appservice(appservice.registration.id.clone(), buf)
332 .log_err()
333 .ok();
334
335 Ok(())
336 })
337 .await
338}
339
340#[implement(Service)]
344#[tracing::instrument(skip(self, serialized), level = "debug")]
345pub fn send_edu_appservice(&self, appservice_id: String, serialized: EduBuf) -> Result {
346 let dest = Destination::Appservice(appservice_id);
347 let event = SendingEvent::Edu(serialized);
348 let _cork = self.db.db.cork();
349
350 self.queue_and_dispatch(dest, event)
351}
352
353#[implement(Service)]
358#[tracing::instrument(skip(self, room_id), level = "debug")]
359pub async fn flush_room(&self, room_id: &RoomId) -> Result {
360 let servers = self
361 .services
362 .state_cache
363 .remote_room_servers(room_id);
364
365 self.flush_servers(servers).await
366}
367
368#[implement(Service)]
373#[tracing::instrument(skip(self, servers), level = "debug")]
374pub async fn flush_servers<'a, S>(&self, servers: S) -> Result
375where
376 S: Stream<Item = &'a ServerName> + Send + 'a,
377{
378 servers
379 .map(ToOwned::to_owned)
380 .map(Destination::Federation)
381 .map(Ok)
382 .ready_try_for_each(|dest| self.dispatch_flush(dest))
383 .await
384}
385
386#[implement(Service)]
387fn dispatch_flush(&self, dest: Destination) -> Result {
388 self.dispatch(Msg {
389 dest,
390 event: SendingEvent::Flush,
391 queue_id: Vec::new(),
392 })
393}
394
395#[implement(Service)]
400#[tracing::instrument(skip(self), level = "debug")]
401pub fn flush_appservice(&self, appservice_id: String) -> Result {
402 self.dispatch_flush(Destination::Appservice(appservice_id))
403}
404
405#[implement(Service)]
412#[tracing::instrument(
413 level = "debug",
414 skip(self),
415 fields(
416 %server,
417 ),
418)]
419pub async fn notify_peer_alive(&self, server: &ServerName) -> bool {
420 let sad = self
421 .services
422 .federation
423 .note_peer_alive(server)
424 .await;
425
426 let replay = sad
427 || self
428 .stalled
429 .lock()
430 .expect("locked")
431 .get(server)
432 .is_some_and(|last| {
433 last.is_none_or(|last| {
434 last.elapsed() >= Duration::from_secs(self.server.config.sender_timeout)
435 })
436 });
437
438 if replay {
439 self.dispatch_flush(Destination::Federation(server.to_owned()))
440 .log_err()
441 .ok();
442 }
443
444 sad
445}
446
447#[implement(Service)]
453#[tracing::instrument(skip(self), level = "debug")]
454pub async fn cleanup_events(
455 &self,
456 appservice_id: Option<&str>,
457 user_id: Option<&UserId>,
458 push_key: Option<&str>,
459) -> Result {
460 let dest = match (appservice_id, user_id, push_key) {
461 | (None, Some(user_id), Some(push_key)) =>
462 Destination::Push(user_id.to_owned(), push_key.to_owned()),
463 | (Some(appservice_id), None, None) => Destination::Appservice(appservice_id.to_owned()),
464 | _ => {
465 debug_warn!("cleanup_events called with too many or too few arguments");
466 return Ok(());
467 },
468 };
469
470 self.db.delete_all_requests_for(&dest).await;
471
472 Ok(())
473}