Skip to main content

tuwunel_service/sending/
data.rs

1use std::{fmt::Debug, sync::Arc};
2
3use futures::{Stream, StreamExt, stream::iter};
4use ruma::{OwnedServerName, ServerName, UserId};
5use tuwunel_core::{
6	Error, Result, at, utils,
7	utils::{ReadyExt, stream::TryIgnore},
8};
9use tuwunel_database::{Database, Deserialized, Map, Txn};
10
11use super::{
12	Destination, EduBuf, SendingEvent, TAG_BADGE_REFRESH, TAG_DEVICE_LIST_CHANGED, TAG_TO_DEVICE,
13};
14
15pub(super) type OutgoingItem = (Key, SendingEvent, Destination);
16pub(super) type SendingItem = (Key, SendingEvent);
17pub(super) type QueueItem = (Key, SendingEvent);
18pub(super) type Key = Vec<u8>;
19
20pub struct Data {
21	servercurrentevent_data: Arc<Map>,
22	servernameevent_data: Arc<Map>,
23	servername_educount: Arc<Map>,
24	pub(super) db: Arc<Database>,
25	services: Arc<crate::services::OnceServices>,
26}
27
28impl Data {
29	pub(super) fn new(args: &crate::Args<'_>) -> Self {
30		let db = &args.db;
31		Self {
32			servercurrentevent_data: db["servercurrentevent_data"].clone(),
33			servernameevent_data: db["servernameevent_data"].clone(),
34			servername_educount: db["servername_educount"].clone(),
35			db: args.db.clone(),
36			services: args.services.clone(),
37		}
38	}
39
40	#[inline]
41	pub(super) fn delete_active_request(&self, key: &[u8]) {
42		self.servercurrentevent_data.remove(key);
43	}
44
45	pub(super) async fn delete_all_active_requests_for(&self, destination: &Destination) {
46		let prefix = destination.get_prefix();
47		self.servercurrentevent_data
48			.raw_keys_prefix(&prefix)
49			.ignore_err()
50			.ready_for_each(|key| self.servercurrentevent_data.remove(key))
51			.await;
52	}
53
54	pub(super) async fn delete_all_requests_for(&self, destination: &Destination) {
55		let prefix = destination.get_prefix();
56		self.servercurrentevent_data
57			.raw_keys_prefix(&prefix)
58			.ignore_err()
59			.ready_for_each(|key| self.servercurrentevent_data.remove(key))
60			.await;
61
62		self.servernameevent_data
63			.raw_keys_prefix(&prefix)
64			.ignore_err()
65			.ready_for_each(|key| self.servernameevent_data.remove(key))
66			.await;
67	}
68
69	pub(super) fn mark_as_active<'a, I>(&self, events: I)
70	where
71		I: Iterator<Item = &'a QueueItem>,
72	{
73		events
74			.filter(|(key, _)| !key.is_empty())
75			.fold(self.db.txn(), |mut txn, (key, val)| {
76				txn.insert_raw(&self.servercurrentevent_data, key, val.value_bytes());
77				txn.del_raw(&self.servernameevent_data, key);
78				txn
79			})
80			.execute();
81	}
82
83	/// Write composed EDUs straight into the active set, keyed by fresh counts;
84	/// unlike `mark_as_active` there is no queue row to delete.
85	pub(super) fn persist_active_edus(&self, server: &ServerName, edus: &[EduBuf]) {
86		let prefix = Destination::Federation(server.to_owned()).get_prefix();
87
88		let items = edus.iter().map(|edu| {
89			let mut key = prefix.clone();
90			let count = self.services.globals.next_count();
91			let count = count.to_be_bytes();
92			key.extend(&count);
93
94			(key, edu.as_slice())
95		});
96
97		Txn::insert(&self.servercurrentevent_data, items).execute();
98	}
99
100	#[inline]
101	pub fn active_requests(&self) -> impl Stream<Item = OutgoingItem> + Send + '_ {
102		self.servercurrentevent_data
103			.raw_stream()
104			.ignore_err()
105			.map(|(key, val)| {
106				let (dest, event) =
107					parse_servercurrentevent(key, val).expect("invalid servercurrentevent");
108
109				(key.to_vec(), event, dest)
110			})
111	}
112
113	#[inline]
114	pub fn active_requests_for(
115		&self,
116		destination: &Destination,
117	) -> impl Stream<Item = SendingItem> + Send + '_ + use<'_> {
118		let prefix = destination.get_prefix();
119		self.servercurrentevent_data
120			.raw_stream_from(&prefix)
121			.ignore_err()
122			.ready_take_while(move |(key, _)| key.starts_with(&prefix))
123			.map(|(key, val)| {
124				let (_, event) =
125					parse_servercurrentevent(key, val).expect("invalid servercurrentevent");
126
127				(key.to_vec(), event)
128			})
129	}
130
131	pub(super) fn queue_requests<'a, I>(&self, requests: I) -> Vec<Vec<u8>>
132	where
133		I: Iterator<Item = (&'a SendingEvent, &'a Destination)> + Clone + Debug + Send,
134	{
135		let keys: Vec<_> = requests
136			.clone()
137			.map(|(event, dest)| {
138				let mut key = dest.get_prefix();
139				if let SendingEvent::Pdu(value) = event {
140					key.extend(value.as_ref());
141				} else {
142					let count = self.services.globals.next_count();
143					let count = count.to_be_bytes();
144					key.extend(&count);
145				}
146
147				key
148			})
149			.collect();
150
151		let items = keys
152			.iter()
153			.map(Vec::as_slice)
154			.zip(requests.map(at!(0)))
155			.map(|(key, event)| (key, event.value_bytes()));
156
157		Txn::insert(&self.servernameevent_data, items).execute();
158
159		keys
160	}
161
162	/// Yields only pending queue items.
163	///
164	/// Empty-key payload wakes always pass because they have no durable row. A
165	/// wake can outlive its row after a completed drain delivered the event.
166	pub(super) fn retain_queued<'a, I>(
167		&'a self,
168		events: I,
169	) -> impl Stream<Item = QueueItem> + Send + 'a
170	where
171		I: IntoIterator<Item = QueueItem> + Send + 'a,
172		I::IntoIter: Send,
173	{
174		iter(events).filter_map(async |item| {
175			let key = &item.0;
176			let exists = async || {
177				self.servernameevent_data
178					.exists(key)
179					.await
180					.is_ok()
181			};
182
183			(key.is_empty() || exists().await).then_some(item)
184		})
185	}
186
187	pub fn queued_requests(
188		&self,
189		destination: &Destination,
190	) -> impl Stream<Item = QueueItem> + Send + '_ + use<'_> {
191		let prefix = destination.get_prefix();
192		self.servernameevent_data
193			.raw_stream_from(&prefix)
194			.ignore_err()
195			.ready_take_while(move |(key, _)| key.starts_with(&prefix))
196			.map(|(key, val)| {
197				let (_, event) =
198					parse_servercurrentevent(key, val).expect("invalid servercurrentevent");
199
200				(key.to_vec(), event)
201			})
202	}
203
204	/// Streams queued push destinations with a pending badge refresh.
205	///
206	/// Returned destinations are owned and may safely cross cursor advances.
207	pub(super) fn queued_badge_refresh_destinations(
208		&self,
209	) -> impl Stream<Item = Destination> + Send + '_ {
210		self.servernameevent_data
211			.raw_stream_from(b"$")
212			.ignore_err()
213			.ready_take_while(|(key, _)| key.starts_with(b"$"))
214			.ready_filter_map(|(key, val)| {
215				(val == [TAG_BADGE_REFRESH]).then(|| {
216					parse_servercurrentevent(key, val)
217						.expect("invalid servercurrentevent")
218						.0
219				})
220			})
221	}
222
223	pub(super) fn set_latest_educount(&self, server_name: &ServerName, last_count: u64) {
224		self.servername_educount
225			.raw_put(server_name, last_count);
226	}
227
228	pub async fn get_latest_educount(&self, server_name: &ServerName) -> u64 {
229		self.servername_educount
230			.get(server_name)
231			.await
232			.deserialized()
233			.unwrap_or(0)
234	}
235}
236
237pub(super) fn parse_servercurrentevent(
238	key: &[u8],
239	value: &[u8],
240) -> Result<(Destination, SendingEvent)> {
241	// Appservices start with a plus
242	Ok::<_, Error>(if key.starts_with(b"+") {
243		let mut parts = key[1..].splitn(2, |&b| b == 0xFF);
244
245		let server = parts
246			.next()
247			.expect("splitn always returns one element");
248		let event = parts
249			.next()
250			.ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
251
252		let server = utils::string_from_bytes(server).map_err(|_| {
253			Error::bad_database("Invalid server bytes in server_currenttransaction")
254		})?;
255
256		let decoded = match value {
257			| [] => SendingEvent::Pdu(event.into()),
258			| [TAG_TO_DEVICE, ..] => SendingEvent::ToDevice(value.into()),
259			| [TAG_DEVICE_LIST_CHANGED, ..] => SendingEvent::DeviceListChanged(value.into()),
260			| _ => SendingEvent::Edu(value.into()),
261		};
262
263		(Destination::Appservice(server), decoded)
264	} else if key.starts_with(b"$") {
265		let mut parts = key[1..].splitn(3, |&b| b == 0xFF);
266
267		let user = parts
268			.next()
269			.expect("splitn always returns one element");
270		let user_string = utils::str_from_bytes(user)
271			.map_err(|_| Error::bad_database("Invalid user string in servercurrentevent"))?;
272		let user_id = UserId::parse(user_string)
273			.map_err(|_| Error::bad_database("Invalid user id in servercurrentevent"))?;
274
275		let pushkey = parts
276			.next()
277			.ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
278		let pushkey_string = utils::string_from_bytes(pushkey)
279			.map_err(|_| Error::bad_database("Invalid pushkey in servercurrentevent"))?;
280
281		let event = parts
282			.next()
283			.ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
284
285		(Destination::Push(user_id, pushkey_string), match value {
286			| [] => SendingEvent::Pdu(event.into()),
287			| [tag] if *tag == TAG_BADGE_REFRESH => SendingEvent::BadgeRefresh,
288			| _ => SendingEvent::Edu(value.into()),
289		})
290	} else {
291		let mut parts = key.splitn(2, |&b| b == 0xFF);
292
293		let server = parts
294			.next()
295			.expect("splitn always returns one element");
296		let event = parts
297			.next()
298			.ok_or_else(|| Error::bad_database("Invalid bytes in servercurrentpdus."))?;
299
300		let server = utils::string_from_bytes(server).map_err(|_| {
301			Error::bad_database("Invalid server bytes in server_currenttransaction")
302		})?;
303
304		(
305			Destination::Federation(OwnedServerName::parse(&server).map_err(|_| {
306				Error::bad_database("Invalid server string in server_currenttransaction")
307			})?),
308			if value.is_empty() {
309				SendingEvent::Pdu(event.into())
310			} else {
311				SendingEvent::Edu(value.into())
312			},
313		)
314	})
315}