1use 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
33pub 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#[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#[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#[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#[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#[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#[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#[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 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
310fn 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#[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#[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#[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#[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#[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#[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#[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#[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#[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}