Skip to main content

tuwunel_service/key_backups/
mod.rs

1use std::{cmp::Ordering, collections::BTreeMap, sync::Arc};
2
3use futures::StreamExt;
4use ruma::{
5	OwnedRoomId, RoomId, UInt, UserId,
6	api::client::backup::{BackupAlgorithm, KeyBackupData, RoomKeyBackup},
7	serde::Raw,
8};
9use tuwunel_core::{
10	Err, Result, err, implement,
11	utils::stream::{ReadyExt, TryIgnore},
12};
13use tuwunel_database::{Deserialized, Ignore, Interfix, Json, Map};
14
15pub struct Service {
16	db: Data,
17	services: Arc<crate::services::OnceServices>,
18}
19
20struct Data {
21	backupid_algorithm: Arc<Map>,
22	backupid_etag: Arc<Map>,
23	backupkeyid_backup: Arc<Map>,
24}
25
26impl crate::Service for Service {
27	fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
28		Ok(Arc::new(Self {
29			db: Data {
30				backupid_algorithm: args.db["backupid_algorithm"].clone(),
31				backupid_etag: args.db["backupid_etag"].clone(),
32				backupkeyid_backup: args.db["backupkeyid_backup"].clone(),
33			},
34			services: args.services.clone(),
35		}))
36	}
37
38	fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
39}
40
41#[implement(Service)]
42pub fn create_backup(
43	&self,
44	user_id: &UserId,
45	backup_metadata: &Raw<BackupAlgorithm>,
46) -> Result<String> {
47	let version = self.services.globals.next_count();
48	let count = self.services.globals.next_count();
49
50	let version_string = version.to_string();
51	let key = (user_id, &version_string);
52	let mut txn = self.services.db.txn();
53
54	txn.put(&self.db.backupid_algorithm, key, Json(backup_metadata));
55	txn.put(&self.db.backupid_etag, key, *count);
56	txn.execute();
57
58	Ok(version_string)
59}
60
61#[implement(Service)]
62pub async fn delete_backup(&self, user_id: &UserId, version: &str) {
63	let key = (user_id, version);
64	self.db.backupid_algorithm.del(key);
65	self.db.backupid_etag.del(key);
66
67	let key = (user_id, version, Interfix);
68	self.db
69		.backupkeyid_backup
70		.keys_prefix_raw(&key)
71		.ignore_err()
72		.ready_for_each(|outdated_key| {
73			self.db.backupkeyid_backup.remove(outdated_key);
74		})
75		.await;
76}
77
78#[implement(Service)]
79pub async fn update_backup<'a>(
80	&self,
81	user_id: &UserId,
82	version: &'a str,
83	backup_metadata: &Raw<BackupAlgorithm>,
84) -> Result<&'a str> {
85	let key = (user_id, version);
86	if self
87		.db
88		.backupid_algorithm
89		.qry(&key)
90		.await
91		.is_err()
92	{
93		return Err!(Request(NotFound("Tried to update nonexistent backup.")));
94	}
95
96	let count = self.services.globals.next_count();
97	let mut txn = self.services.db.txn();
98
99	txn.put(&self.db.backupid_etag, key, *count);
100	txn.put_raw(&self.db.backupid_algorithm, key, backup_metadata.json().get());
101	txn.execute();
102
103	Ok(version)
104}
105
106#[implement(Service)]
107pub async fn get_latest_backup_version(&self, user_id: &UserId) -> Result<String> {
108	type Key<'a> = (&'a UserId, &'a str);
109
110	let key = (user_id, Interfix);
111	let mut versions: Vec<_> = self
112		.db
113		.backupid_algorithm
114		.keys_from(&key)
115		.ignore_err()
116		.ready_take_while(|(user_id_, _): &Key<'_>| *user_id_ == user_id)
117		.ready_filter_map(|(_, version): Key<'_>| version.parse::<u64>().ok())
118		.collect()
119		.await;
120
121	versions.sort_unstable();
122	let Some(latest) = versions.last() else {
123		return Err!(Request(NotFound("No backup versions found")));
124	};
125
126	Ok(latest.to_string())
127}
128
129#[implement(Service)]
130pub async fn get_latest_backup(
131	&self,
132	user_id: &UserId,
133) -> Result<(String, Raw<BackupAlgorithm>)> {
134	let version = self.get_latest_backup_version(user_id).await?;
135
136	let key = (user_id, version.as_str());
137	self.db
138		.backupid_algorithm
139		.qry(&key)
140		.await
141		.deserialized()
142		.map(|algorithm| (version, algorithm))
143		.map_err(|e| err!(Request(NotFound("No backup found: {e}"))))
144}
145
146#[implement(Service)]
147pub async fn get_backup(&self, user_id: &UserId, version: &str) -> Result<Raw<BackupAlgorithm>> {
148	let key = (user_id, version);
149	self.db
150		.backupid_algorithm
151		.qry(&key)
152		.await
153		.deserialized()
154}
155
156#[implement(Service)]
157pub async fn add_key(
158	&self,
159	user_id: &UserId,
160	version: &str,
161	room_id: &RoomId,
162	session_id: &str,
163	key_data: &Raw<KeyBackupData>,
164) -> Result {
165	let key = (user_id, version);
166	if self
167		.db
168		.backupid_algorithm
169		.qry(&key)
170		.await
171		.is_err()
172	{
173		return Err!(Request(NotFound("Tried to update nonexistent backup.")));
174	}
175
176	// Keep the existing key unless the incoming one is preferable per MSC1219.
177	let replace = match self
178		.get_session(user_id, version, room_id, session_id)
179		.await
180	{
181		| Ok(old_key) => is_better_key(&old_key, key_data)?,
182		| Err(_) => true,
183	};
184
185	if !replace {
186		return Ok(());
187	}
188
189	let count = self.services.globals.next_count();
190	let mut txn = self.services.db.txn();
191
192	txn.put(&self.db.backupid_etag, key, *count);
193
194	let key = (user_id, version, room_id, session_id);
195
196	txn.put_raw(&self.db.backupkeyid_backup, key, key_data.json().get());
197	txn.execute();
198
199	Ok(())
200}
201
202/// Per MSC1219: prefer verified, then lower `first_message_index`, then lower
203/// `forwarded_count`; equal on all three keeps the existing key.
204fn is_better_key(old: &Raw<KeyBackupData>, new: &Raw<KeyBackupData>) -> Result<bool> {
205	let old_verified = old
206		.get_field::<bool>("is_verified")?
207		.unwrap_or_default();
208
209	let new_verified = new
210		.get_field::<bool>("is_verified")?
211		.ok_or_else(|| err!(Request(BadJson("`is_verified` field should exist"))))?;
212
213	if old_verified != new_verified {
214		return Ok(new_verified);
215	}
216
217	let old_first_message_index = old
218		.get_field::<UInt>("first_message_index")?
219		.unwrap_or(UInt::MAX);
220
221	let new_first_message_index = new
222		.get_field::<UInt>("first_message_index")?
223		.ok_or_else(|| err!(Request(BadJson("`first_message_index` field should exist"))))?;
224
225	match new_first_message_index.cmp(&old_first_message_index) {
226		| Ordering::Less => Ok(true),
227		| Ordering::Greater => Ok(false),
228		| Ordering::Equal => {
229			let old_forwarded_count = old
230				.get_field::<UInt>("forwarded_count")?
231				.unwrap_or(UInt::MAX);
232
233			let new_forwarded_count = new
234				.get_field::<UInt>("forwarded_count")?
235				.ok_or_else(|| err!(Request(BadJson("`forwarded_count` field should exist"))))?;
236
237			Ok(new_forwarded_count < old_forwarded_count)
238		},
239	}
240}
241
242#[implement(Service)]
243pub async fn count_keys(&self, user_id: &UserId, version: &str) -> usize {
244	let prefix = (user_id, version);
245	self.db
246		.backupkeyid_backup
247		.keys_prefix_raw(&prefix)
248		.count()
249		.await
250}
251
252#[implement(Service)]
253pub async fn get_etag(&self, user_id: &UserId, version: &str) -> String {
254	let key = (user_id, version);
255	self.db
256		.backupid_etag
257		.qry(&key)
258		.await
259		.deserialized::<u64>()
260		.as_ref()
261		.map(ToString::to_string)
262		.expect("Backup has no etag.")
263}
264
265#[implement(Service)]
266pub async fn get_all(
267	&self,
268	user_id: &UserId,
269	version: &str,
270) -> BTreeMap<OwnedRoomId, RoomKeyBackup> {
271	type Key<'a> = (Ignore, Ignore, &'a RoomId, &'a str);
272	type KeyVal<'a> = (Key<'a>, Raw<KeyBackupData>);
273
274	let mut rooms = BTreeMap::<OwnedRoomId, RoomKeyBackup>::new();
275	let default = || RoomKeyBackup { sessions: BTreeMap::new() };
276
277	let prefix = (user_id, version, Interfix);
278	self.db
279		.backupkeyid_backup
280		.stream_prefix(&prefix)
281		.ignore_err()
282		.ready_for_each(|((_, _, room_id, session_id), key_backup_data): KeyVal<'_>| {
283			rooms
284				.entry(room_id.into())
285				.or_insert_with(default)
286				.sessions
287				.insert(session_id.into(), key_backup_data);
288		})
289		.await;
290
291	rooms
292}
293
294#[implement(Service)]
295pub async fn get_room(
296	&self,
297	user_id: &UserId,
298	version: &str,
299	room_id: &RoomId,
300) -> BTreeMap<String, Raw<KeyBackupData>> {
301	type KeyVal<'a> = ((Ignore, Ignore, Ignore, &'a str), Raw<KeyBackupData>);
302
303	let prefix = (user_id, version, room_id, Interfix);
304	self.db
305		.backupkeyid_backup
306		.stream_prefix(&prefix)
307		.ignore_err()
308		.map(|((.., session_id), key_backup_data): KeyVal<'_>| {
309			(session_id.to_owned(), key_backup_data)
310		})
311		.collect()
312		.await
313}
314
315#[implement(Service)]
316pub async fn get_session(
317	&self,
318	user_id: &UserId,
319	version: &str,
320	room_id: &RoomId,
321	session_id: &str,
322) -> Result<Raw<KeyBackupData>> {
323	let key = (user_id, version, room_id, session_id);
324
325	self.db
326		.backupkeyid_backup
327		.qry(&key)
328		.await
329		.deserialized()
330}
331
332#[implement(Service)]
333pub async fn delete_all_keys(&self, user_id: &UserId, version: &str) {
334	let key = (user_id, version, Interfix);
335	self.db
336		.backupkeyid_backup
337		.keys_prefix_raw(&key)
338		.ignore_err()
339		.ready_for_each(|outdated_key| self.db.backupkeyid_backup.remove(outdated_key))
340		.await;
341}
342
343#[implement(Service)]
344pub async fn delete_room_keys(&self, user_id: &UserId, version: &str, room_id: &RoomId) {
345	let key = (user_id, version, room_id, Interfix);
346	self.db
347		.backupkeyid_backup
348		.keys_prefix_raw(&key)
349		.ignore_err()
350		.ready_for_each(|outdated_key| {
351			self.db.backupkeyid_backup.remove(outdated_key);
352		})
353		.await;
354}
355
356#[implement(Service)]
357pub async fn delete_room_key(
358	&self,
359	user_id: &UserId,
360	version: &str,
361	room_id: &RoomId,
362	session_id: &str,
363) {
364	let key = (user_id, version, room_id, session_id);
365	self.db
366		.backupkeyid_backup
367		.keys_prefix_raw(&key)
368		.ignore_err()
369		.ready_for_each(|outdated_key| {
370			self.db.backupkeyid_backup.remove(outdated_key);
371		})
372		.await;
373}