1#[cfg(test)]
2mod tests;
3
4use std::{
5 fmt::Debug,
6 iter::once,
7 ops::{
8 Bound,
9 Bound::{Excluded, Included, Unbounded},
10 },
11 pin::pin,
12 sync::Arc,
13 time::Duration,
14 vec,
15};
16
17use futures::{
18 Stream, StreamExt,
19 stream::{iter, unfold},
20};
21use ruma::{OwnedServerName, ServerName, UserId};
22use tuwunel_core::{
23 Error, Result, at, implement,
24 matrix::ShortRoomId,
25 utils,
26 utils::{
27 IterStream, ReadyExt,
28 bytes::prefix_successor,
29 str_from_bytes,
30 stream::{TryIgnore, WidebandExt},
31 time::now_secs,
32 },
33};
34use tuwunel_database::{Database, Deserialized, Interfix, Map, Txn};
35
36use super::{
37 Destination, EduBuf, SendingEvent, TAG_BADGE_REFRESH, TAG_DEVICE_LIST_CHANGED, TAG_TO_DEVICE,
38};
39
40pub(super) type OutgoingItem = (Key, SendingEvent, Destination);
41pub(super) type SendingItem = (Key, SendingEvent);
42
43pub(super) type QueueItem = (Key, SendingEvent);
47pub(super) type Key = Vec<u8>;
48pub(super) type Keys = Vec<Key>;
49
50const PARK_BASE: Duration = Duration::from_hours(1);
51const PARK_LIMIT: Duration = Duration::from_hours(24);
52
53#[derive(Clone, Copy, Debug, Eq, PartialEq)]
58pub struct Park {
59 pub room: ShortRoomId,
61
62 pub until: u64,
64
65 pub count: u64,
67}
68
69type ParkRow<'a> = ((&'a ServerName, ShortRoomId), (u64, u64));
70
71pub struct Data {
78 servercurrentevent_data: Arc<Map>,
79 servernameevent_data: Arc<Map>,
80 servername_educount: Arc<Map>,
81 servershortroomid_park: Arc<Map>,
82 pub(super) db: Arc<Database>,
83 services: Arc<crate::services::OnceServices>,
84}
85
86#[implement(Data)]
87pub(super) fn new(args: &crate::Args<'_>) -> Self {
88 let db = &args.db;
89
90 Self {
91 servercurrentevent_data: db["servercurrentevent_data"].clone(),
92 servernameevent_data: db["servernameevent_data"].clone(),
93 servername_educount: db["servername_educount"].clone(),
94 servershortroomid_park: db["servershortroomid_park"].clone(),
95 db: args.db.clone(),
96 services: args.services.clone(),
97 }
98}
99
100#[implement(Data)]
101#[inline]
102pub(super) fn delete_active_request(&self, key: &[u8]) {
103 self.servercurrentevent_data.remove(key);
104}
105
106#[implement(Data)]
110pub(super) fn delete_active_requests<'a, I>(&self, keys: I)
111where
112 I: IntoIterator<Item = &'a Key>,
113{
114 keys.into_iter()
115 .filter(|key| !key.is_empty())
116 .fold(self.db.txn(), |mut txn, key| {
117 txn.del_raw(&self.servercurrentevent_data, key);
118 txn
119 })
120 .execute();
121}
122
123#[implement(Data)]
124pub(super) async fn delete_all_requests_for(&self, destination: &Destination) {
125 let prefix = destination.get_prefix();
126
127 self.servercurrentevent_data
128 .raw_keys_prefix(&prefix)
129 .ignore_err()
130 .ready_for_each(|key| self.servercurrentevent_data.remove(key))
131 .await;
132
133 self.servernameevent_data
134 .raw_keys_prefix(&prefix)
135 .ignore_err()
136 .ready_for_each(|key| self.servernameevent_data.remove(key))
137 .await;
138}
139
140#[implement(Data)]
141pub(super) fn mark_as_active<'a, I>(&self, events: I)
142where
143 I: Iterator<Item = &'a QueueItem>,
144{
145 events
146 .filter(|(key, _)| !key.is_empty())
147 .fold(self.db.txn(), |mut txn, (key, val)| {
148 txn.insert_raw(&self.servercurrentevent_data, key, val.value_bytes());
149 txn.del_raw(&self.servernameevent_data, key);
150 txn
151 })
152 .execute();
153}
154
155#[implement(Data)]
160pub(super) fn persist_active_edus(
161 &self,
162 server: &ServerName,
163 edus: &[EduBuf],
164) -> vec::IntoIter<Key> {
165 let dest = Destination::Federation(server.to_owned());
166
167 let permits: Vec<_> = edus
169 .iter()
170 .map(|_| self.services.globals.next_count())
171 .collect();
172
173 let keys: Keys = permits
174 .iter()
175 .map(|permit| dest.count_key(**permit))
176 .collect();
177
178 let items = keys
179 .iter()
180 .map(Vec::as_slice)
181 .zip(edus.iter().map(EduBuf::as_slice));
182
183 Txn::insert(&self.servercurrentevent_data, items).execute();
184
185 keys.into_iter()
186}
187
188#[implement(Data)]
193#[inline]
194pub fn active_requests(&self) -> impl Stream<Item = OutgoingItem> + Send + '_ {
195 self.servercurrentevent_data
196 .raw_stream()
197 .ignore_err()
198 .map(|(key, val)| {
199 let (dest, event) =
200 parse_servercurrentevent(key, val).expect("invalid servercurrentevent");
201
202 (key.to_vec(), event, dest)
203 })
204}
205
206#[implement(Data)]
211#[inline]
212pub fn active_requests_for(
213 &self,
214 destination: &Destination,
215) -> impl Stream<Item = SendingItem> + Send + '_ + use<'_> {
216 let prefix = destination.get_prefix();
217
218 self.servercurrentevent_data
219 .raw_stream_from(&prefix)
220 .ignore_err()
221 .ready_take_while(move |(key, _)| key.starts_with(&prefix))
222 .map(queue_item)
223}
224
225#[implement(Data)]
226pub(super) fn queue_requests<'a, I>(&self, requests: I) -> Keys
227where
228 I: Iterator<Item = (&'a SendingEvent, &'a Destination)> + Clone + Debug + Send,
229{
230 let (keys, _permits): (Keys, Vec<_>) = requests
232 .clone()
233 .map(|(event, dest)| match event {
234 | SendingEvent::Pdu(pdu_id) => (dest.event_key(pdu_id), None),
235 | _ => {
236 let permit = self.services.globals.next_count();
237
238 (dest.count_key(*permit), Some(permit))
239 },
240 })
241 .unzip();
242
243 let items = keys
244 .iter()
245 .map(Vec::as_slice)
246 .zip(requests.map(at!(0)))
247 .map(|(key, event)| (key, event.value_bytes()));
248
249 Txn::insert(&self.servernameevent_data, items).execute();
250
251 keys
252}
253
254#[implement(Data)]
259pub(super) fn retain_queued<'a, I>(
260 &'a self,
261 events: I,
262) -> impl Stream<Item = QueueItem> + Send + 'a
263where
264 I: IntoIterator<Item = QueueItem> + Send + 'a,
265 I::IntoIter: Send,
266{
267 iter(events).wide_filter_map(async |item| {
268 let key = &item.0;
269 let pending = key.is_empty()
270 || self
271 .servernameevent_data
272 .exists(key)
273 .await
274 .is_ok();
275
276 pending.then_some(item)
277 })
278}
279
280#[implement(Data)]
285pub fn queued_requests(
286 &self,
287 destination: &Destination,
288) -> impl Stream<Item = QueueItem> + Send + '_ + use<'_> {
289 let prefix = destination.get_prefix();
290
291 self.queued_range(&prefix, prefix_end(prefix.clone()))
292}
293
294#[implement(Data)]
299pub(super) fn queued_room(
300 &self,
301 destination: &Destination,
302 room: ShortRoomId,
303) -> impl Stream<Item = QueueItem> + Send + '_ + use<'_> {
304 let start = destination.count_key(room);
305 let len = start.len();
306 let end = prefix_end(start.clone());
307
308 self.queued_range(&start, end)
309 .ready_filter(move |(key, _)| key.len() > len)
310}
311
312#[implement(Data)]
317pub(super) fn queued_except<'a>(
318 &'a self,
319 destination: &'a Destination,
320 skip: &'a [ShortRoomId],
321) -> impl Stream<Item = QueueItem> + Send + 'a {
322 debug_assert!(skip.is_sorted(), "skipped rooms must be sorted");
323
324 let resumes = skip
325 .iter()
326 .map(|room| destination.count_key(room.saturating_add(1)));
327
328 let ends = skip
329 .iter()
330 .map(|&room| Included(destination.count_key(room)))
331 .chain(once(prefix_end(destination.get_prefix())));
332
333 once(destination.get_prefix())
334 .chain(resumes)
335 .zip(ends)
336 .stream()
337 .flat_map(|(start, end)| self.queued_range(&start, end))
338}
339
340#[implement(Data)]
341fn queued_range(
342 &self,
343 start: &[u8],
344 end: Bound<Key>,
345) -> impl Stream<Item = QueueItem> + Send + '_ + use<'_> {
346 self.servernameevent_data
347 .raw_stream_from(start)
348 .ignore_err()
349 .ready_take_while(move |&(key, _)| match &end {
350 | Included(end) => key <= end.as_slice(),
351 | Excluded(end) => key < end.as_slice(),
352 | Unbounded => true,
353 })
354 .map(queue_item)
355}
356
357#[implement(Data)]
362pub(super) fn demote<'a, I>(&self, items: I, park: Option<(&ServerName, Park)>)
363where
364 I: IntoIterator<Item = &'a QueueItem>,
365{
366 let txn = items
367 .into_iter()
368 .filter(|(_, event)| matches!(event, SendingEvent::Pdu(_)))
369 .fold(self.db.txn(), |mut txn, (key, _)| {
370 txn.insert_raw(&self.servernameevent_data, key, []);
371 txn.del_raw(&self.servercurrentevent_data, key);
372 txn
373 });
374
375 park.into_iter()
376 .fold(txn, |mut txn, (server, park)| {
377 txn.put(&self.servershortroomid_park, (server, park.room), (park.until, park.count));
378 txn
379 })
380 .execute();
381}
382
383#[implement(Data)]
388pub(super) async fn next_park(&self, server: &ServerName, room: ShortRoomId) -> Park {
389 let count = self
390 .servershortroomid_park
391 .qry(&(server, room))
392 .await
393 .deserialized()
394 .map_or(1, |(_, count): (u64, u64)| count.saturating_add(1));
395
396 let doublings = u32::try_from(count.saturating_sub(1)).unwrap_or(u32::MAX);
397 let hold = PARK_BASE
398 .saturating_mul(2_u32.saturating_pow(doublings))
399 .min(PARK_LIMIT);
400
401 Park {
402 room,
403 until: now_secs().saturating_add(hold.as_secs()),
404 count,
405 }
406}
407
408#[implement(Data)]
409pub(super) fn unpark(&self, server: &ServerName, room: ShortRoomId) {
410 self.servershortroomid_park.del((server, room));
411}
412
413#[implement(Data)]
417pub(super) fn parks<'a>(
418 &'a self,
419 server: &'a ServerName,
420) -> impl Stream<Item = Park> + Send + 'a {
421 self.servershortroomid_park
422 .stream_prefix(&(server, Interfix))
423 .ignore_err()
424 .map(park_row)
425 .map(at!(1))
426}
427
428#[implement(Data)]
432pub fn parked(&self) -> impl Stream<Item = (OwnedServerName, Park)> + Send + '_ {
433 self.servershortroomid_park
434 .stream()
435 .ignore_err()
436 .map(park_row)
437 .map(|(server, park)| (server.to_owned(), park))
438}
439
440fn queue_item((key, val): (&[u8], &[u8])) -> QueueItem {
441 let (_, event) = parse_servercurrentevent(key, val).expect("invalid servercurrentevent");
442
443 (key.to_vec(), event)
444}
445
446fn prefix_end(prefix: Key) -> Bound<Key> { prefix_successor(prefix).map_or(Unbounded, Excluded) }
447
448fn park_row(((server, room), (until, count)): ParkRow<'_>) -> (&ServerName, Park) {
449 (server, Park { room, until, count })
450}
451
452#[implement(Data)]
456pub(super) fn queued_badge_refresh_destinations(
457 &self,
458) -> impl Stream<Item = Destination> + Send + '_ {
459 self.servernameevent_data
460 .raw_stream_from(b"$")
461 .ignore_err()
462 .ready_take_while(|(key, _)| key.starts_with(b"$"))
463 .ready_filter_map(|(key, val)| {
464 val.eq(&[TAG_BADGE_REFRESH]).then(|| {
465 parse_servercurrentevent(key, val)
466 .expect("invalid servercurrentevent")
467 .0
468 })
469 })
470}
471
472#[implement(Data)]
477pub(super) fn queued_federation_destinations<'a, F>(
478 &'a self,
479 owns: F,
480) -> impl Stream<Item = Result<Destination>> + Send + 'a
481where
482 F: Fn(&ServerName) -> bool + Copy + Send + 'a,
483{
484 queued_destinations(move |key| self.seek_queued_key(key), owns)
485}
486
487fn queued_destinations<S, F, C>(
488 seek: S,
489 owns: C,
490) -> impl Stream<Item = Result<Destination>> + Send
491where
492 S: Fn(Key) -> F + Send,
493 F: Future<Output = Result<Option<Key>>> + Send,
494 C: Fn(&ServerName) -> bool + Copy + Send,
495{
496 unfold((Some(Key::new()), seek, owns), async move |(lower, seek, owns)| {
497 let lower = lower?;
498 let (item, next) = match seek(lower).await {
499 | Ok(None) => return None,
500 | Err(error) => (Err(error), None),
501 | Ok(Some(key)) => queued_destination(key, owns),
502 };
503
504 Some((item, (next, seek, owns)))
505 })
506 .ready_filter_map(Result::transpose)
507}
508
509fn queued_destination(
510 key: Key,
511 owns: impl Fn(&ServerName) -> bool,
512) -> (Result<Option<Destination>>, Option<Key>) {
513 match key.first() {
514 | Some(b'$') => return (Ok(None), Some(single_key(key, b'%'))),
515 | Some(b'+') => return (Ok(None), Some(single_key(key, b','))),
516 | _ => {},
517 }
518
519 let Some(end) = key.iter().position(|byte| *byte == u8::MAX) else {
520 let error = Error::bad_database("Queued federation key has no destination delimiter");
521
522 return (Err(error), Some(after_key(key)));
523 };
524
525 let prefix = truncate_key(key, end.saturating_add(1));
526 let destination = str_from_bytes(&prefix[..end])
527 .ok()
528 .and_then(|server| <&ServerName>::try_from(server).ok())
529 .ok_or_else(|| Error::bad_database("Invalid queued federation destination"))
530 .map(|server| owns(server).then(|| Destination::Federation(server.to_owned())));
531
532 (destination, prefix_successor(prefix))
533}
534
535fn single_key(mut key: Key, byte: u8) -> Key {
536 key.clear();
537 key.push(byte);
538 key
539}
540
541fn after_key(mut key: Key) -> Key {
542 key.push(0);
543 key
544}
545
546fn truncate_key(mut key: Key, len: usize) -> Key {
547 key.truncate(len);
548 key
549}
550
551#[implement(Data)]
552#[tracing::instrument(level = "trace", skip_all)]
553async fn seek_queued_key(&self, mut key: Key) -> Result<Option<Key>> {
554 let found = pin!(
556 self.servernameevent_data
557 .raw_keys_from(&key)
558 .map(|item| item.map(|bytes| replace_key(&mut key, bytes)))
559 )
560 .next()
561 .await
562 .transpose()?
563 .is_some();
564
565 Ok(found.then_some(key))
566}
567
568fn replace_key(key: &mut Key, bytes: &[u8]) {
569 key.clear();
570 key.extend_from_slice(bytes);
571}
572
573#[implement(Data)]
574pub(super) fn set_latest_educount(&self, server_name: &ServerName, last_count: u64) {
575 self.servername_educount
576 .raw_put(server_name, last_count);
577}
578
579#[implement(Data)]
584pub async fn get_latest_educount(&self, server_name: &ServerName) -> u64 {
585 self.servername_educount
586 .get(server_name)
587 .await
588 .deserialized()
589 .unwrap_or(0)
590}
591
592impl SendingEvent {
593 pub(super) fn value_bytes(&self) -> &[u8] {
598 match self {
599 | Self::Edu(bytes) | Self::ToDevice(bytes) | Self::DeviceListChanged(bytes) => bytes,
600 | Self::BadgeRefresh => &[TAG_BADGE_REFRESH],
601 | Self::Pdu(_) | Self::Flush => &[],
602 }
603 }
604}
605
606pub(super) fn parse_servercurrentevent(
607 key: &[u8],
608 value: &[u8],
609) -> Result<(Destination, SendingEvent)> {
610 Ok::<_, Error>(if key.starts_with(b"+") {
612 let mut parts = key[1..].splitn(2, |&b| b == 0xFF);
613
614 let server = parts
615 .next()
616 .expect("splitn always returns one element");
617 let event = parts
618 .next()
619 .ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
620
621 let server = utils::string_from_bytes(server).map_err(|_| {
622 Error::bad_database("Invalid server bytes in server_currenttransaction")
623 })?;
624
625 let decoded = match value {
626 | [] => SendingEvent::Pdu(event.into()),
627 | [TAG_TO_DEVICE, ..] => SendingEvent::ToDevice(value.into()),
628 | [TAG_DEVICE_LIST_CHANGED, ..] => SendingEvent::DeviceListChanged(value.into()),
629 | _ => SendingEvent::Edu(value.into()),
630 };
631
632 (Destination::Appservice(server), decoded)
633 } else if key.starts_with(b"$") {
634 let mut parts = key[1..].splitn(3, |&b| b == 0xFF);
635
636 let user = parts
637 .next()
638 .expect("splitn always returns one element");
639
640 let user_string = str_from_bytes(user)
641 .map_err(|_| Error::bad_database("Invalid user string in servercurrentevent"))?;
642
643 let user_id = UserId::parse(user_string)
644 .map_err(|_| Error::bad_database("Invalid user id in servercurrentevent"))?;
645
646 let pushkey = parts
647 .next()
648 .ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
649 let pushkey_string = utils::string_from_bytes(pushkey)
650 .map_err(|_| Error::bad_database("Invalid pushkey in servercurrentevent"))?;
651
652 let event = parts
653 .next()
654 .ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
655
656 (Destination::Push(user_id, pushkey_string), match value {
657 | [] => SendingEvent::Pdu(event.into()),
658 | [tag] if *tag == TAG_BADGE_REFRESH => SendingEvent::BadgeRefresh,
659 | _ => SendingEvent::Edu(value.into()),
660 })
661 } else {
662 let mut parts = key.splitn(2, |&b| b == 0xFF);
663
664 let server = parts
665 .next()
666 .expect("splitn always returns one element");
667 let event = parts
668 .next()
669 .ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
670
671 let server = utils::string_from_bytes(server).map_err(|_| {
672 Error::bad_database("Invalid server bytes in server_currenttransaction")
673 })?;
674
675 (
676 Destination::Federation(OwnedServerName::parse(&server).map_err(|_| {
677 Error::bad_database("Invalid server string in server_currenttransaction")
678 })?),
679 if value.is_empty() {
680 SendingEvent::Pdu(event.into())
681 } else {
682 SendingEvent::Edu(value.into())
683 },
684 )
685 })
686}