1use mls_rs_core::group::{EpochRecord, GroupState, GroupStateStorage};
6use rusqlite::{params, Connection, OptionalExtension};
7use std::{
8 fmt::Debug,
9 sync::{Arc, Mutex},
10};
11use zeroize::Zeroizing;
12
13use crate::SqLiteDataStorageError;
14
15pub(crate) const DEFAULT_EPOCH_RETENTION_LIMIT: u64 = 3;
16
17#[derive(Debug, Clone)]
18pub struct SqLiteGroupStateStorage {
26 connection: Arc<Mutex<Connection>>,
27 max_epoch_retention: u64,
28}
29
30impl SqLiteGroupStateStorage {
31 pub(crate) fn new(connection: Connection) -> SqLiteGroupStateStorage {
32 SqLiteGroupStateStorage {
33 connection: Arc::new(Mutex::new(connection)),
34 max_epoch_retention: DEFAULT_EPOCH_RETENTION_LIMIT,
35 }
36 }
37
38 pub fn with_max_epoch_retention(self, max_epoch_retention: u64) -> Self {
39 Self {
40 connection: self.connection,
41 max_epoch_retention,
42 }
43 }
44
45 pub fn group_ids(&self) -> Result<Vec<Vec<u8>>, SqLiteDataStorageError> {
47 let connection = self.connection.lock().unwrap();
48
49 let mut statement = connection
50 .prepare("SELECT group_id FROM mls_group")
51 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))?;
52
53 let res = statement
54 .query_map([], |row| row.get(0))
55 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))?
56 .try_fold(Vec::new(), |mut ids, id| {
57 ids.push(id.map_err(|e| SqLiteDataStorageError::DataConversionError(e.into()))?);
58 Ok::<_, SqLiteDataStorageError>(ids)
59 })
60 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))?;
61
62 Ok(res)
63 }
64
65 pub fn delete_group(&self, group_id: &[u8]) -> Result<(), SqLiteDataStorageError> {
67 let connection = self.connection.lock().unwrap();
68
69 connection
70 .execute(
71 "DELETE FROM mls_group WHERE group_id = ?",
72 params![group_id],
73 )
74 .map(|_| ())
75 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))
76 }
77
78 pub fn max_epoch_retention(&self) -> u64 {
79 self.max_epoch_retention
80 }
81
82 fn get_snapshot_data(
83 &self,
84 group_id: &[u8],
85 ) -> Result<Option<Vec<u8>>, SqLiteDataStorageError> {
86 let connection = self.connection.lock().unwrap();
87
88 connection
89 .query_row(
90 "SELECT snapshot FROM mls_group where group_id = ?",
91 [group_id],
92 |row| row.get::<_, Vec<u8>>(0),
93 )
94 .optional()
95 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))
96 }
97
98 fn get_epoch_data(
99 &self,
100 group_id: &[u8],
101 epoch_id: u64,
102 ) -> Result<Option<Vec<u8>>, SqLiteDataStorageError> {
103 let connection = self.connection.lock().unwrap();
104
105 connection
106 .query_row(
107 "SELECT epoch_data FROM epoch where group_id = ? AND epoch_id = ?",
108 params![
109 group_id,
110 i64::try_from(epoch_id)
111 .map_err(|_| SqLiteDataStorageError::EpochIdOverflow(epoch_id))?
112 ],
113 |row| row.get::<_, Vec<u8>>(0),
114 )
115 .optional()
116 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))
117 }
118
119 fn max_epoch_id(&self, group_id: &[u8]) -> Result<Option<u64>, SqLiteDataStorageError> {
120 let connection = self.connection.lock().unwrap();
121
122 connection
123 .query_row(
124 "SELECT MAX(epoch_id) FROM epoch WHERE group_id = ?",
125 params![group_id],
126 |row| {
127 row.get::<_, Option<i64>>(0).and_then(|opt| {
128 opt.map(|v| {
129 u64::try_from(v)
130 .map_err(|_| rusqlite::Error::IntegralValueOutOfRange(0, v))
131 })
132 .transpose()
133 })
134 },
135 )
136 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))
137 }
138
139 fn update_group_state(
140 &self,
141 group_id: &[u8],
142 group_snapshot: &[u8],
143 inserts: Vec<EpochRecord>,
144 updates: Vec<EpochRecord>,
145 ) -> Result<(), SqLiteDataStorageError> {
146 let mut max_epoch_id = None;
147
148 let mut connection = self.connection.lock().unwrap();
149 let transaction = connection
150 .transaction()
151 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))?;
152
153 transaction.execute(
155 "INSERT INTO mls_group (group_id, snapshot) VALUES (?, ?) ON CONFLICT(group_id) DO UPDATE SET snapshot=excluded.snapshot",
156 params![group_id, group_snapshot],
157 ).map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))?;
158
159 for epoch in inserts {
161 max_epoch_id = Some(epoch.id);
162
163 transaction
164 .execute(
165 "INSERT INTO epoch (group_id, epoch_id, epoch_data) VALUES (?, ?, ?)",
166 params![
167 group_id,
168 i64::try_from(epoch.id)
169 .map_err(|_| SqLiteDataStorageError::EpochIdOverflow(epoch.id))?,
170 &*epoch.data
171 ],
172 )
173 .map(|_| ())
174 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))?;
175 }
176
177 updates.into_iter().try_for_each(|epoch| {
179 transaction
180 .execute(
181 "UPDATE epoch SET epoch_data = ? WHERE group_id = ? AND epoch_id = ?",
182 params![
183 &*epoch.data,
184 group_id,
185 i64::try_from(epoch.id)
186 .map_err(|_| SqLiteDataStorageError::EpochIdOverflow(epoch.id))?
187 ],
188 )
189 .map(|_| ())
190 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))
191 })?;
192
193 if let Some(max_epoch_id) = max_epoch_id {
195 if max_epoch_id >= self.max_epoch_retention {
196 let delete_under = max_epoch_id - self.max_epoch_retention;
197
198 transaction
199 .execute(
200 "DELETE FROM epoch WHERE group_id = ? AND epoch_id <= ?",
201 params![
202 group_id,
203 i64::try_from(delete_under).map_err(|_| {
204 SqLiteDataStorageError::EpochIdOverflow(delete_under)
205 })?
206 ],
207 )
208 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))?;
209 }
210 }
211
212 transaction
214 .commit()
215 .map_err(|e| SqLiteDataStorageError::SqlEngineError(e.into()))
216 }
217}
218
219#[cfg_attr(not(mls_build_async), maybe_async::must_be_sync)]
220#[cfg_attr(mls_build_async, maybe_async::must_be_async)]
221impl GroupStateStorage for SqLiteGroupStateStorage {
222 type Error = SqLiteDataStorageError;
223
224 async fn write(
225 &mut self,
226 state: GroupState,
227 inserts: Vec<EpochRecord>,
228 updates: Vec<EpochRecord>,
229 ) -> Result<(), Self::Error> {
230 self.update_group_state(&state.id, &state.data, inserts, updates)
231 }
232
233 async fn state(&self, group_id: &[u8]) -> Result<Option<Zeroizing<Vec<u8>>>, Self::Error> {
234 let data = self.get_snapshot_data(group_id)?;
235 Ok(data.map(Into::into))
236 }
237
238 async fn max_epoch_id(&self, group_id: &[u8]) -> Result<Option<u64>, Self::Error> {
239 self.max_epoch_id(group_id)
240 }
241
242 async fn epoch(
243 &self,
244 group_id: &[u8],
245 epoch_id: u64,
246 ) -> Result<Option<Zeroizing<Vec<u8>>>, Self::Error> {
247 let data = self.get_epoch_data(group_id, epoch_id)?;
248 Ok(data.map(Into::into))
249 }
250}
251
252#[cfg(test)]
253mod tests {
254 use assert_matches::assert_matches;
255
256 use crate::{
257 SqLiteDataStorageEngine,
258 {connection_strategy::MemoryStrategy, test_utils::gen_rand_bytes},
259 };
260
261 use super::*;
262
263 fn get_test_storage() -> SqLiteGroupStateStorage {
264 SqLiteDataStorageEngine::new(MemoryStrategy)
265 .unwrap()
266 .group_state_storage()
267 .unwrap()
268 }
269
270 fn test_group_id() -> Vec<u8> {
271 gen_rand_bytes(32)
272 }
273
274 fn test_snapshot() -> Vec<u8> {
275 gen_rand_bytes(1024)
276 }
277
278 fn test_epoch(id: u64) -> EpochRecord {
279 EpochRecord {
280 data: gen_rand_bytes(256).into(),
281 id,
282 }
283 }
284
285 struct TestData {
286 storage: SqLiteGroupStateStorage,
287 snapshot: Vec<u8>,
288 group_id: Vec<u8>,
289 epoch_0: EpochRecord,
290 }
291
292 fn setup_group_storage_test() -> TestData {
293 let test_storage = get_test_storage();
294 let test_group_id = test_group_id();
295 let test_epoch_0 = test_epoch(0);
296 let test_snapshot = test_snapshot();
297
298 test_storage
299 .update_group_state(
300 &test_group_id,
301 &test_snapshot,
302 vec![test_epoch_0.clone()],
303 vec![],
304 )
305 .unwrap();
306
307 TestData {
308 storage: test_storage,
309 group_id: test_group_id,
310 epoch_0: test_epoch_0,
311 snapshot: test_snapshot,
312 }
313 }
314
315 #[test]
316 fn group_can_be_initially_stored() {
317 let test_data = setup_group_storage_test();
318
319 let snapshot = test_data
321 .storage
322 .get_snapshot_data(&test_data.group_id)
323 .unwrap();
324 assert_eq!(snapshot.unwrap(), test_data.snapshot);
325
326 let epoch = test_data
328 .storage
329 .get_epoch_data(&test_data.group_id, 0)
330 .unwrap();
331 assert_eq!(epoch.unwrap(), *test_data.epoch_0.data);
332 }
333
334 #[test]
335 fn snapshot_and_epoch_can_be_updated() {
336 let test_data = setup_group_storage_test();
337 let test_snapshot = test_snapshot();
338
339 let epoch_update = test_epoch(0);
340
341 test_data
342 .storage
343 .update_group_state(
344 &test_data.group_id,
345 &test_snapshot,
346 vec![],
347 vec![epoch_update.clone()],
348 )
349 .unwrap();
350
351 let snapshot = test_data
353 .storage
354 .get_snapshot_data(&test_data.group_id)
355 .unwrap();
356
357 assert_eq!(snapshot.unwrap(), test_snapshot);
358
359 assert_eq!(
361 test_data
362 .storage
363 .get_epoch_data(&test_data.group_id, 0)
364 .unwrap()
365 .unwrap(),
366 *epoch_update.data
367 );
368 }
369
370 #[test]
371 fn epochs_are_truncated() {
372 test_epochs_are_truncated(9);
373 test_epochs_are_truncated(DEFAULT_EPOCH_RETENTION_LIMIT);
374 }
375
376 fn test_epochs_are_truncated(n: u64) {
377 let test_data = setup_group_storage_test();
378
379 let mut test_epochs = (1..n + 1).map(test_epoch).collect::<Vec<_>>();
380
381 test_data
382 .storage
383 .update_group_state(
384 &test_data.group_id,
385 &test_snapshot(),
386 test_epochs.clone(),
387 vec![],
388 )
389 .unwrap();
390
391 test_epochs.insert(0, test_data.epoch_0);
392
393 for epoch in test_epochs {
394 let stored = test_data
395 .storage
396 .get_epoch_data(&test_data.group_id, epoch.id)
397 .unwrap();
398
399 if epoch.id <= n - DEFAULT_EPOCH_RETENTION_LIMIT {
400 assert!(stored.is_none());
401 } else {
402 assert_eq!(stored.unwrap(), *epoch.data);
403 }
404 }
405 }
406
407 #[test]
408 fn epoch_insert_update_old_epoch() {
409 let test_data = setup_group_storage_test();
410
411 test_data
412 .storage
413 .update_group_state(
414 &test_data.group_id,
415 &test_snapshot(),
416 vec![test_epoch(1)],
417 vec![],
418 )
419 .unwrap();
420
421 let test_epochs = (2..10).map(test_epoch).collect::<Vec<_>>();
422 let new_epoch_1 = test_epoch(1);
423
424 test_data
425 .storage
426 .update_group_state(
427 &test_data.group_id,
428 &test_snapshot(),
429 test_epochs.clone(),
430 vec![new_epoch_1.clone()],
431 )
432 .unwrap();
433
434 assert!(test_data
435 .storage
436 .get_epoch_data(&test_data.group_id, 1)
437 .unwrap()
438 .is_none());
439 }
440
441 #[test]
442 fn max_epoch_is_none_for_non_persisted_group() {
443 let storage = get_test_storage();
444
445 let res = storage.max_epoch_id(&[0, 1, 2]).unwrap();
446
447 assert!(res.is_none())
448 }
449
450 #[test]
451 fn max_epoch_is_none_when_no_epochs() {
452 let storage = get_test_storage();
453 let group_id = b"test";
454
455 storage
456 .update_group_state(group_id, &[0, 1, 2], vec![], vec![])
457 .unwrap();
458
459 let res = storage.max_epoch_id(group_id).unwrap();
460
461 assert!(res.is_none())
462 }
463
464 #[test]
465 fn max_epoch_can_be_calculated() {
466 let test_data = setup_group_storage_test();
467
468 test_data
469 .storage
470 .update_group_state(
471 &test_data.group_id,
472 &test_snapshot(),
473 (1..10).map(test_epoch).collect(),
474 vec![],
475 )
476 .unwrap();
477
478 assert_eq!(
479 test_data
480 .storage
481 .max_epoch_id(&test_data.group_id)
482 .unwrap()
483 .unwrap(),
484 9
485 );
486 }
487
488 #[test]
489 fn muiltiple_groups_can_exist() {
490 let test_data = setup_group_storage_test();
491
492 let new_group = test_group_id();
493 let new_group_epoch = test_epoch(0);
494
495 test_data
496 .storage
497 .update_group_state(
498 &new_group,
499 &test_snapshot(),
500 vec![new_group_epoch.clone()],
501 vec![],
502 )
503 .unwrap();
504
505 let all_groups = test_data.storage.group_ids().unwrap();
506
507 vec![test_data.group_id.clone(), new_group.clone()]
509 .into_iter()
510 .for_each(|id| {
511 assert!(all_groups.contains(&id));
512 });
513
514 assert_eq!(
515 test_data
516 .storage
517 .get_epoch_data(&new_group, 0)
518 .unwrap()
519 .unwrap(),
520 *new_group_epoch.data
521 );
522 }
523
524 #[test]
525 fn delete_group() {
526 let test_data = setup_group_storage_test();
527
528 test_data.storage.delete_group(&test_data.group_id).unwrap();
529
530 assert!(test_data.storage.group_ids().unwrap().is_empty());
531 }
532
533 #[test]
534 fn epoch_id_overflow() {
535 let storage = get_test_storage();
536 let group_id = test_group_id();
537 let snapshot = test_snapshot();
538 let overflow_epoch = EpochRecord {
539 id: u64::MAX,
540 data: gen_rand_bytes(256).into(),
541 };
542
543 let err = storage
544 .update_group_state(&group_id, &snapshot, vec![overflow_epoch], vec![])
545 .unwrap_err();
546
547 assert_matches!(err, SqLiteDataStorageError::EpochIdOverflow(u64::MAX));
548 }
549}