tuwunel_service/key_backups/
mod.rs1use 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 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
202fn 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}