Skip to main content

mls_rs_provider_sqlite/
group_state.rs

1// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
2// Copyright by contributors to this project.
3// SPDX-License-Identifier: (Apache-2.0 OR MIT)
4
5use 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)]
18/// SQLite Storage for MLS group states.
19///
20/// # Limitations
21///
22/// Epoch IDs are stored as SQLite INTEGER (signed 64-bit), limiting the maximum
23/// epoch ID to [`i64::MAX`] (9,223,372,036,854,775,807). Operations with epoch IDs
24/// exceeding this value will return [`SqLiteDataStorageError::EpochIdOverflow`].
25pub 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    /// List all the group ids for groups that are stored.
46    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    /// Delete a group from storage.
66    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        // Upsert into the group table to set the most recent snapshot
154        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        // Insert new epochs as needed
160        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        // Update existing epochs as needed
178        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        // Delete old epochs as needed
194        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        // Execute the full transaction
213        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        // Attempt to fetch the snapshot
320        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        // Attempt to fetch the epoch data
327        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        // Attempt to fetch the new snapshot
352        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        // Attempt to access the epochs
360        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        // Order is not deterministic
508        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}