Skip to main content

tuwunel_service/sending/
data.rs

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
43/// A queued event paired with its row key.
44///
45/// An empty key marks a synthetic wake that has no durable row.
46pub(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/// A room held back from a server that keeps rejecting it.
54///
55/// The server's queued rows for the room are skipped until `until` passes,
56/// when the room is retried alone; a delivery ends the park.
57#[derive(Clone, Copy, Debug, Eq, PartialEq)]
58pub struct Park {
59	/// The room's short id.
60	pub room: ShortRoomId,
61
62	/// When the room may be sent again, in seconds since the epoch.
63	pub until: u64,
64
65	/// Consecutive parks without a delivery, each doubling the last.
66	pub count: u64,
67}
68
69type ParkRow<'a> = ((&'a ServerName, ShortRoomId), (u64, u64));
70
71/// The sending service's column families.
72///
73/// Queued rows wait in `servernameevent_data` until a transaction claims them
74/// into `servercurrentevent_data`; `servername_educount` is the per-server EDU
75/// watermark, and `servershortroomid_park` holds rooms a server keeps
76/// rejecting.
77pub 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/// Acknowledges the active rows one transaction carried.
107///
108/// Empty keys mark synthetic events, which have no row.
109#[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/// Write composed EDUs straight into the active set, keyed by fresh counts.
156///
157/// Unlike `mark_as_active` there is no queue row to delete. Yields the new
158/// keys in `edus` order for the transaction to acknowledge.
159#[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	// The permits retire their counts only once the rows have landed.
168	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/// Streams every active row across all destinations.
189///
190/// Rows are decoded as they are read; a row that fails to decode is a
191/// corrupt database and panics.
192#[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/// Streams the active rows of one destination.
207///
208/// Rows are decoded as they are read; a row that fails to decode is a
209/// corrupt database and panics.
210#[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	// The permits retire their counts only once the rows have landed.
231	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/// Yields only pending queue items.
255///
256/// Empty-key payload wakes always pass because they have no durable row. A
257/// wake can outlive its row after a completed drain delivered the event.
258#[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/// Streams the queued rows of one destination.
281///
282/// Rows are decoded as they are read; a row that fails to decode is a
283/// corrupt database and panics.
284#[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/// Streams the queued PDU rows of one room for a destination.
295///
296/// An EDU row whose count equals the room's short id has exactly the room's
297/// key prefix and is skipped.
298#[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/// Streams the queued rows of one destination, past the PDU rows of `skip`.
313///
314/// `skip` must be sorted. Each skipped room costs one seek rather than a visit
315/// per row, and the EDU row sharing a skipped room's key bytes is still yielded.
316#[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/// Moves a failed transaction's PDU rows back to the queue under their own keys.
358///
359/// EDU rows have no room, so they stay active and ride the next transaction.
360/// A park lands in the same batch as the rows it holds back.
361#[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/// Computes the park for a room that failed again.
384///
385/// Each consecutive park doubles the last hold, up to a day; the caller
386/// writes it.
387#[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/// Streams a server's parked rooms, expired ones included.
414///
415/// Rooms come in short id order, so their ids form a sorted skip list.
416#[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/// Streams every parked room with its server, expired ones included.
429///
430/// The server is owned, so an item may be kept across polls.
431#[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/// Streams queued push destinations with a pending badge refresh.
453///
454/// Returned destinations are owned and may safely cross cursor advances.
455#[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/// Streams distinct queued federation destinations belonging to this worker.
473///
474/// Each seek copies one key into the owned cursor and skips its complete
475/// destination prefix. Sigils and other shards are skipped before ownership.
476#[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	// The cursor callback copies into the owned seek buffer before its drop.
555	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/// The count of the newest EDU shipped to a server.
580///
581/// A server never shipped to reads as zero, so its first window starts at the
582/// beginning of the log.
583#[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	/// Return bytes written verbatim as the queue row value.
594	///
595	/// PDUs keep their ID in the row key and flushes are not persisted. EDU
596	/// variants own `[tag][count][body]`; a badge refresh owns only its tag.
597	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	// Appservices start with a plus
611	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}