Skip to main content

tuwunel_service/key_backups/
mod.rs

1//! Versioned room key backup storage.
2//!
3//! The service stores encrypted session keys by user, backup version, room, and session. Backup
4//! metadata and change tags are maintained alongside the key rows for client synchronization.
5
6use std::{cmp::Ordering, collections::BTreeMap, module_path, sync::Arc};
7
8use futures::{FutureExt, Stream, StreamExt, TryStreamExt, future::try_join};
9use http::StatusCode;
10use ruma::{
11	OwnedRoomId, OwnedUserId, RoomId, UInt, UserId,
12	api::{
13		client::backup::{BackupAlgorithm, KeyBackupData, RoomKeyBackup},
14		error::{ErrorKind, WrongRoomKeysVersionErrorData},
15	},
16	serde::Raw,
17};
18use tuwunel_core::{
19	Err, Result, err, implement,
20	utils::{
21		MutexMap,
22		stream::{ReadyExt, TryIgnore},
23	},
24};
25use tuwunel_database::{Deserialized, Ignore, Interfix, Json, Map};
26
27type Key<'a> = (&'a RoomId, &'a str, &'a Raw<KeyBackupData>);
28type StoredKey<'a> = (Ignore, Ignore, &'a RoomId, &'a str);
29type StoredKeyVal<'a> = (StoredKey<'a>, Raw<KeyBackupData>);
30type StoredRoomKeyVal<'a> = ((Ignore, Ignore, Ignore, &'a str), Raw<KeyBackupData>);
31type VersionKey<'a> = (&'a UserId, &'a str);
32
33/// Stores and mutates versioned room key backups.
34///
35/// Mutations take a per-user lock so concurrent update paths cannot interleave. Read operations
36/// stream directly from the versioned database prefixes.
37pub struct Service {
38	db: Data,
39	mutex: MutexMap<OwnedUserId, ()>,
40	services: Arc<crate::services::OnceServices>,
41}
42
43struct Data {
44	backupid_algorithm: Arc<Map>,
45	backupid_etag: Arc<Map>,
46	backupkeyid_backup: Arc<Map>,
47}
48
49impl crate::Service for Service {
50	fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
51		Ok(Arc::new(Self {
52			db: Data {
53				backupid_algorithm: args.db["backupid_algorithm"].clone(),
54				backupid_etag: args.db["backupid_etag"].clone(),
55				backupkeyid_backup: args.db["backupkeyid_backup"].clone(),
56			},
57			mutex: MutexMap::new(),
58			services: args.services.clone(),
59		}))
60	}
61
62	fn name(&self) -> &str { crate::service::make_name(module_path!()) }
63}
64
65/// Creates an empty backup version and returns its identifier.
66///
67/// The version and initial change tag use separate global sequence numbers. Metadata and the change
68/// tag are committed together while the user's backup lock is held.
69///
70/// # Panics
71///
72/// Panics when dispatching either global sequence number fails.
73#[implement(Service)]
74pub async fn create_backup(
75	&self,
76	user_id: &UserId,
77	backup_metadata: &Raw<BackupAlgorithm>,
78) -> Result<String> {
79	let _backup_lock = self.mutex.lock(user_id).await;
80	let version = self.services.globals.next_count();
81	let count = self.services.globals.next_count();
82
83	let version_string = version.to_string();
84	let key = (user_id, &version_string);
85	let mut txn = self.services.db.txn();
86
87	txn.put(&self.db.backupid_algorithm, key, Json(backup_metadata));
88	txn.put(&self.db.backupid_etag, key, *count);
89	txn.execute();
90
91	Ok(version_string)
92}
93
94/// Deletes a backup version and its stored session keys.
95///
96/// Metadata and the change tag are removed before the version's key rows are scanned. Missing
97/// versions are accepted, and unreadable key rows are skipped during cleanup.
98#[implement(Service)]
99pub async fn delete_backup(&self, user_id: &UserId, version: &str) {
100	let _backup_lock = self.mutex.lock(user_id).await;
101	let key = (user_id, version);
102	self.db.backupid_algorithm.del(key);
103	self.db.backupid_etag.del(key);
104
105	let key = (user_id, version, Interfix);
106
107	self.db
108		.backupkeyid_backup
109		.keys_prefix_raw(&key)
110		.ignore_err()
111		.ready_for_each(|outdated_key| {
112			self.db.backupkeyid_backup.remove(outdated_key);
113		})
114		.await;
115}
116
117/// Replaces a backup version's metadata and advances its change tag.
118///
119/// The version must already exist. Its metadata and fresh change tag are committed together while
120/// the user's backup lock is held.
121///
122/// # Panics
123///
124/// Panics when dispatching the new global sequence number fails.
125#[implement(Service)]
126pub async fn update_backup<'a>(
127	&self,
128	user_id: &UserId,
129	version: &'a str,
130	backup_metadata: &Raw<BackupAlgorithm>,
131) -> Result<&'a str> {
132	let _backup_lock = self.mutex.lock(user_id).await;
133	let key = (user_id, version);
134	if self
135		.db
136		.backupid_algorithm
137		.qry(&key)
138		.await
139		.is_err()
140	{
141		return Err!(Request(NotFound("Tried to update nonexistent backup.")));
142	}
143
144	let count = self.services.globals.next_count();
145	let mut txn = self.services.db.txn();
146
147	txn.put(&self.db.backupid_etag, key, *count);
148	txn.put_raw(&self.db.backupid_algorithm, key, backup_metadata.json().get());
149	txn.execute();
150
151	Ok(version)
152}
153
154/// Returns a user's latest numeric backup version.
155///
156/// The greatest parseable version number is selected. Unreadable rows and nonnumeric version names
157/// are skipped, and an absent result is reported as not found.
158#[implement(Service)]
159pub async fn get_latest_backup_version(&self, user_id: &UserId) -> Result<String> {
160	let key = (user_id, Interfix);
161	let latest = self
162		.db
163		.backupid_algorithm
164		.keys_from(&key)
165		.ignore_err()
166		.ready_take_while(|(user_id_, _): &VersionKey<'_>| *user_id_ == user_id)
167		.ready_filter_map(|(_, version): VersionKey<'_>| version.parse::<u64>().ok())
168		.ready_fold(None, |latest: Option<u64>, version| {
169			Some(latest.map_or(version, |latest| latest.max(version)))
170		})
171		.await;
172
173	let Some(latest) = latest else {
174		return Err!(Request(NotFound("No backup versions found")));
175	};
176
177	Ok(latest.to_string())
178}
179
180/// Returns a user's latest backup version and its algorithm metadata.
181///
182/// The version is selected by [`Self::get_latest_backup_version`]. Missing or unreadable metadata
183/// is reported as a not-found request error.
184#[implement(Service)]
185pub async fn get_latest_backup(
186	&self,
187	user_id: &UserId,
188) -> Result<(String, Raw<BackupAlgorithm>)> {
189	let version = self.get_latest_backup_version(user_id).await?;
190
191	let key = (user_id, version.as_str());
192	self.db
193		.backupid_algorithm
194		.qry(&key)
195		.await
196		.deserialized()
197		.map(|algorithm| (version, algorithm))
198		.map_err(|e| err!(Request(NotFound("No backup found: {e}"))))
199}
200
201/// Loads the algorithm metadata for a backup version.
202///
203/// The raw value retains the stored JSON representation while constraining it to
204/// [`BackupAlgorithm`].
205#[implement(Service)]
206pub async fn get_backup(&self, user_id: &UserId, version: &str) -> Result<Raw<BackupAlgorithm>> {
207	let key = (user_id, version);
208	self.db
209		.backupid_algorithm
210		.qry(&key)
211		.await
212		.deserialized()
213}
214
215/// Adds a stream of room keys to the latest backup version.
216///
217/// The stream is drained serially while this user's backup mutations are
218/// locked. The returned count and etag describe the completed operation.
219///
220/// # Panics
221///
222/// Panics when an accepted key requires a global sequence number that cannot be allocated or
223/// persisted.
224#[implement(Service)]
225pub async fn add_keys<'a, S>(
226	&self,
227	user_id: &UserId,
228	version: &str,
229	keys: S,
230) -> Result<(usize, u64)>
231where
232	S: Stream<Item = Key<'a>> + Send,
233{
234	let _backup_lock = self.mutex.lock(user_id).await;
235
236	self.check_backup_version(user_id, version)
237		.await?;
238
239	let key = (user_id, version);
240
241	self.db
242		.backupid_algorithm
243		.qry(&key)
244		.await
245		.map_err(|_| err!(Request(NotFound("Tried to update nonexistent backup."))))?;
246
247	keys.map(Ok)
248		.try_for_each(async |(room_id, session_id, key_data)| {
249			self.add_key(user_id, version, room_id, session_id, key_data)
250				.await
251		})
252		.await?;
253
254	self.get_count_etag(user_id, version).await
255}
256
257#[implement(Service)]
258async fn check_backup_version(&self, user_id: &UserId, version: &str) -> Result {
259	let current_version = self.get_latest_backup_version(user_id).await?;
260
261	if current_version == version {
262		return Ok(());
263	}
264
265	let status = StatusCode::BAD_REQUEST;
266	let data = WrongRoomKeysVersionErrorData::new(current_version);
267	let kind = ErrorKind::WrongRoomKeysVersion(data);
268	let message =
269		"You may only manipulate the most recently created version of the backup.".into();
270
271	Err!(Request(kind, message, status))
272}
273
274#[implement(Service)]
275async fn add_key(
276	&self,
277	user_id: &UserId,
278	version: &str,
279	room_id: &RoomId,
280	session_id: &str,
281	key_data: &Raw<KeyBackupData>,
282) -> Result {
283	// Keep the existing key unless the incoming one is preferable per MSC1219.
284	let replace = match self
285		.get_session(user_id, version, room_id, session_id)
286		.await
287	{
288		| Ok(old_key) => is_better_key(&old_key, key_data)?,
289		| Err(_) => true,
290	};
291
292	if !replace {
293		return Ok(());
294	}
295
296	let key = (user_id, version);
297	let count = self.services.globals.next_count();
298	let mut txn = self.services.db.txn();
299
300	txn.put(&self.db.backupid_etag, key, *count);
301
302	let key = (user_id, version, room_id, session_id);
303
304	txn.put_raw(&self.db.backupkeyid_backup, key, key_data.json().get());
305	txn.execute();
306
307	Ok(())
308}
309
310/// Reports whether a new session key should replace the stored key.
311///
312/// MSC1219 prefers verified keys, then lower `first_message_index`, then lower
313/// `forwarded_count`. A complete tie preserves the existing key.
314fn is_better_key(old: &Raw<KeyBackupData>, new: &Raw<KeyBackupData>) -> Result<bool> {
315	let old_verified = old
316		.get_field::<bool>("is_verified")?
317		.unwrap_or_default();
318
319	let new_verified = new
320		.get_field::<bool>("is_verified")?
321		.ok_or_else(|| err!(Request(BadJson("`is_verified` field should exist"))))?;
322
323	if old_verified != new_verified {
324		return Ok(new_verified);
325	}
326
327	let old_first_message_index = old
328		.get_field::<UInt>("first_message_index")?
329		.unwrap_or(UInt::MAX);
330
331	let new_first_message_index = new
332		.get_field::<UInt>("first_message_index")?
333		.ok_or_else(|| err!(Request(BadJson("`first_message_index` field should exist"))))?;
334
335	match new_first_message_index.cmp(&old_first_message_index) {
336		| Ordering::Less => Ok(true),
337		| Ordering::Greater => Ok(false),
338		| Ordering::Equal => {
339			let old_forwarded_count = old
340				.get_field::<UInt>("forwarded_count")?
341				.unwrap_or(UInt::MAX);
342
343			let new_forwarded_count = new
344				.get_field::<UInt>("forwarded_count")?
345				.ok_or_else(|| err!(Request(BadJson("`forwarded_count` field should exist"))))?;
346
347			Ok(new_forwarded_count < old_forwarded_count)
348		},
349	}
350}
351
352/// Reads a backup's current key count and etag.
353///
354/// Mutation callers keep the user lock across this method. Other callers
355/// receive two independent current observations without snapshot semantics.
356#[implement(Service)]
357pub async fn get_count_etag(&self, user_id: &UserId, version: &str) -> Result<(usize, u64)> {
358	let count = self.count_keys(user_id, version).map(Ok);
359	let etag = self.get_etag(user_id, version);
360
361	try_join(count, etag).await
362}
363
364/// Counts the session keys stored in a backup version.
365///
366/// Every raw key row under the user's version prefix contributes to the count.
367#[implement(Service)]
368pub async fn count_keys(&self, user_id: &UserId, version: &str) -> usize {
369	let prefix = (user_id, version, Interfix);
370
371	self.db
372		.backupkeyid_backup
373		.keys_prefix_raw(&prefix)
374		.count()
375		.await
376}
377
378/// Returns the current change tag for a backup version.
379///
380/// The tag advances whenever accepted key data or metadata changes.
381#[implement(Service)]
382pub async fn get_etag(&self, user_id: &UserId, version: &str) -> Result<u64> {
383	let key = (user_id, version);
384
385	self.db
386		.backupid_etag
387		.qry(&key)
388		.await
389		.deserialized::<u64>()
390}
391
392/// Loads every session key in a backup version, grouped by room.
393///
394/// Session identifiers map to their raw key backup data within each room. Unreadable rows are
395/// omitted from the result.
396#[implement(Service)]
397pub async fn get_all(
398	&self,
399	user_id: &UserId,
400	version: &str,
401) -> BTreeMap<OwnedRoomId, RoomKeyBackup> {
402	let default = || RoomKeyBackup { sessions: BTreeMap::new() };
403	let prefix = (user_id, version, Interfix);
404
405	self.db
406		.backupkeyid_backup
407		.stream_prefix(&prefix)
408		.ignore_err()
409		.ready_fold(BTreeMap::new(), |mut rooms, row: StoredKeyVal<'_>| {
410			let ((_, _, room_id, session_id), key_backup_data) = row;
411
412			rooms
413				.entry(room_id.into())
414				.or_insert_with(default)
415				.sessions
416				.insert(session_id.into(), key_backup_data);
417
418			rooms
419		})
420		.await
421}
422
423/// Loads every backed-up session key for one room.
424///
425/// The returned map is keyed by session ID and retains each value as raw key backup data.
426/// Unreadable rows are omitted.
427#[implement(Service)]
428pub async fn get_room(
429	&self,
430	user_id: &UserId,
431	version: &str,
432	room_id: &RoomId,
433) -> BTreeMap<String, Raw<KeyBackupData>> {
434	let prefix = (user_id, version, room_id, Interfix);
435
436	self.db
437		.backupkeyid_backup
438		.stream_prefix(&prefix)
439		.ignore_err()
440		.map(|((.., session_id), key_backup_data): StoredRoomKeyVal<'_>| {
441			(session_id.to_owned(), key_backup_data)
442		})
443		.collect()
444		.await
445}
446
447/// Loads one backed-up session key.
448///
449/// The key is selected by user, backup version, room, and session ID. The stored JSON is returned
450/// without eagerly deserializing the key data.
451#[implement(Service)]
452pub async fn get_session(
453	&self,
454	user_id: &UserId,
455	version: &str,
456	room_id: &RoomId,
457	session_id: &str,
458) -> Result<Raw<KeyBackupData>> {
459	let key = (user_id, version, room_id, session_id);
460
461	self.db
462		.backupkeyid_backup
463		.qry(&key)
464		.await
465		.deserialized()
466}
467
468/// Deletes session keys from a backup version.
469///
470/// The version must exist. Rows whose scan reports an error are skipped. Deletion runs under the
471/// user's backup lock, advances the change tag, and reports a zero remaining count with the new tag.
472///
473/// # Panics
474///
475/// Panics when dispatching the new change tag fails.
476#[implement(Service)]
477pub async fn delete_all_keys(&self, user_id: &UserId, version: &str) -> Result<(usize, u64)> {
478	let _backup_lock = self.mutex.lock(user_id).await;
479
480	self.check_backup_exists(user_id, version).await?;
481
482	let key = (user_id, version, Interfix);
483	self.db
484		.backupkeyid_backup
485		.keys_prefix_raw(&key)
486		.ignore_err()
487		.ready_for_each(|outdated_key| self.db.backupkeyid_backup.remove(outdated_key))
488		.await;
489
490	let etag = self.bump_etag(user_id, version);
491
492	Ok((0, etag))
493}
494
495#[implement(Service)]
496async fn check_backup_exists(&self, user_id: &UserId, version: &str) -> Result {
497	let algorithm = self
498		.get_backup(user_id, version)
499		.map(|result| result.map(drop));
500
501	let etag = self
502		.get_etag(user_id, version)
503		.map(|result| result.map(drop));
504
505	try_join(algorithm, etag).await.map(drop)
506}
507
508#[implement(Service)]
509fn bump_etag(&self, user_id: &UserId, version: &str) -> u64 {
510	let etag = self.services.globals.next_count();
511
512	self.db
513		.backupid_etag
514		.put((user_id, version), *etag);
515
516	*etag
517}
518
519/// Deletes session keys for one room in a backup version.
520///
521/// The version must exist. Rows whose scan reports an error are skipped. Deletion runs under the
522/// user's backup lock, advances the change tag, and returns the remaining key count from the
523/// exact-version scan.
524///
525/// # Panics
526///
527/// Panics when dispatching the new change tag fails.
528#[implement(Service)]
529pub async fn delete_room_keys(
530	&self,
531	user_id: &UserId,
532	version: &str,
533	room_id: &RoomId,
534) -> Result<(usize, u64)> {
535	let _backup_lock = self.mutex.lock(user_id).await;
536
537	self.check_backup_exists(user_id, version).await?;
538
539	let key = (user_id, version, room_id, Interfix);
540
541	self.db
542		.backupkeyid_backup
543		.keys_prefix_raw(&key)
544		.ignore_err()
545		.ready_for_each(|outdated_key| self.db.backupkeyid_backup.remove(outdated_key))
546		.await;
547
548	let etag = self.bump_etag(user_id, version);
549	let count = self.count_keys(user_id, version).await;
550
551	Ok((count, etag))
552}
553
554/// Deletes one session key from a backup version.
555///
556/// The version must exist. The change tag advances even when the selected key was absent, and the
557/// returned count covers every remaining key in the version.
558///
559/// # Panics
560///
561/// Panics when dispatching the new change tag fails.
562#[implement(Service)]
563pub async fn delete_room_key(
564	&self,
565	user_id: &UserId,
566	version: &str,
567	room_id: &RoomId,
568	session_id: &str,
569) -> Result<(usize, u64)> {
570	let _backup_lock = self.mutex.lock(user_id).await;
571
572	self.check_backup_exists(user_id, version).await?;
573
574	self.db
575		.backupkeyid_backup
576		.del((user_id, version, room_id, session_id));
577
578	let etag = self.bump_etag(user_id, version);
579	let count = self.count_keys(user_id, version).await;
580
581	Ok((count, etag))
582}