Skip to main content

matrix_sdk_sqlite/
crypto_store.rs

1// Copyright 2022, 2026 The Matrix.org Foundation C.I.C.
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::{
16    collections::HashMap,
17    fmt,
18    ops::Deref,
19    path::{Path, PathBuf},
20    sync::{Arc, RwLock},
21};
22
23use async_trait::async_trait;
24use deadpool::managed::PoolConfig;
25use matrix_sdk_base::cross_process_lock::CrossProcessLockGeneration;
26use matrix_sdk_crypto::{
27    Account, DeviceData, GossipRequest, GossippedSecret, SecretInfo, TrackedUser, UserIdentityData,
28    olm::{
29        InboundGroupSession, OutboundGroupSession, PickledInboundGroupSession,
30        PrivateCrossSigningIdentity, SenderDataType, Session, StaticAccountData,
31    },
32    store::{
33        CryptoStore,
34        types::{
35            BackupKeys, Changes, DehydratedDeviceKey, PendingChanges, RoomKeyCounts,
36            RoomKeyWithheldEntry, RoomPendingKeyBundleDetails, RoomSettings,
37            StoredRoomKeyBundleData,
38        },
39    },
40};
41use matrix_sdk_store_encryption::StoreCipher;
42use ruma::{
43    DeviceId, MilliSecondsSinceUnixEpoch, OwnedDeviceId, RoomId, TransactionId, UserId,
44    events::secret::request::SecretName,
45};
46use rusqlite::{OptionalExtension, named_params, params_from_iter};
47use tokio::{
48    fs,
49    sync::{Mutex, OwnedMutexGuard},
50};
51use tracing::{debug, instrument, warn};
52use vodozemac::Curve25519PublicKey;
53use zeroize::Zeroizing;
54
55use crate::{
56    OpenStoreError, RuntimeConfig, Secret, SqliteStoreConfig,
57    connection::{self, Connection as SqliteAsyncConn, Pool as SqlitePool, SqliteConnections},
58    error::{Error, Result},
59    utils::{
60        EncryptableStore, Key, SqliteAsyncConnExt, SqliteKeyValueStoreAsyncConnExt,
61        SqliteKeyValueStoreConnExt,
62    },
63};
64
65/// The database name.
66const DATABASE_NAME: &str = "matrix-sdk-crypto.sqlite3";
67
68/// An SQLite-based crypto store.
69#[derive(Clone)]
70pub struct SqliteCryptoStore {
71    store_cipher: Option<Arc<StoreCipher>>,
72
73    /// `Some` when active, `None` when closed. The outer `Mutex` serialises
74    /// close/reopen with connection access.
75    connections: Arc<Mutex<Option<SqliteConnections>>>,
76
77    /// Retained so we can rebuild the pool on reopen.
78    db_path: PathBuf,
79
80    /// Retained so we can rebuild the pool on reopen.
81    pool_config: PoolConfig,
82
83    /// Retained so we can re-apply runtime config on reopen.
84    runtime_config: RuntimeConfig,
85
86    // DB values cached in memory
87    static_account: Arc<RwLock<Option<StaticAccountData>>>,
88    save_changes_lock: Arc<Mutex<()>>,
89}
90
91#[cfg(not(tarpaulin_include))]
92impl fmt::Debug for SqliteCryptoStore {
93    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
94        f.debug_struct("SqliteCryptoStore").finish_non_exhaustive()
95    }
96}
97
98impl EncryptableStore for SqliteCryptoStore {
99    fn get_cypher(&self) -> Option<&StoreCipher> {
100        self.store_cipher.as_deref()
101    }
102}
103
104impl SqliteCryptoStore {
105    /// Create an `SqliteCryptoStore` struct without trying to create the
106    /// database or migrate to a newer version. This is only for use internally,
107    /// and for testing.
108    ///
109    /// # Arguments
110    ///
111    /// - `secret` - The secret used to encrypt the data.
112    /// - `pool` - A connection pool to use for reading from the store.
113    /// - `conn` - The connection to use for writing to the store.
114    pub(crate) async fn create_raw(
115        secret: Option<Secret>,
116        pool: SqlitePool,
117        conn: SqliteAsyncConn,
118        pool_config: PoolConfig,
119        runtime_config: RuntimeConfig,
120    ) -> Result<Self, OpenStoreError> {
121        let store_cipher = match secret {
122            Some(s) => Some(Arc::new(conn.get_or_create_store_cipher(s).await?)),
123            None => None,
124        };
125
126        let db_path = pool.manager().database_path.clone();
127
128        Ok(Self {
129            store_cipher,
130            connections: Arc::new(Mutex::new(Some(SqliteConnections {
131                pool,
132                write_connection: Arc::new(Mutex::new(conn)),
133            }))),
134            db_path,
135            pool_config,
136            runtime_config,
137            static_account: Arc::new(RwLock::new(None)),
138            save_changes_lock: Default::default(),
139        })
140    }
141
142    /// Open the SQLite-based crypto store at the given path using the given
143    /// passphrase to encrypt private data.
144    pub async fn open(
145        path: impl AsRef<Path>,
146        passphrase: Option<&str>,
147    ) -> Result<Self, OpenStoreError> {
148        Self::open_with_config(&SqliteStoreConfig::new(path).passphrase(passphrase)).await
149    }
150
151    /// Open the SQLite-based crypto store at the given path using the given key
152    /// to encrypt private data.
153    pub async fn open_with_key(
154        path: impl AsRef<Path>,
155        key: Option<&[u8]>,
156    ) -> Result<Self, OpenStoreError> {
157        Self::open_with_config(&SqliteStoreConfig::new(path).key(key)).await
158    }
159
160    /// Open the SQLite-based crypto store with the config open config.
161    pub async fn open_with_config(config: &SqliteStoreConfig) -> Result<Self, OpenStoreError> {
162        fs::create_dir_all(&config.path).await.map_err(OpenStoreError::CreateDir)?;
163
164        let pool = config.build_pool_of_connections(DATABASE_NAME)?;
165        let pool_config = config.pool_config();
166        let runtime_config = config.runtime_config();
167
168        let this =
169            Self::open_with_pool(pool, config.secret.clone(), pool_config, runtime_config).await?;
170        this.read().await?.apply_runtime_config(runtime_config).await?;
171
172        Ok(this)
173    }
174
175    /// Create an SQLite-based crypto store using the given SQLite database
176    /// pool. The given secret will be used to encrypt private data.
177    async fn open_with_pool(
178        pool: SqlitePool,
179        secret: Option<Secret>,
180        pool_config: PoolConfig,
181        runtime_config: RuntimeConfig,
182    ) -> Result<Self, OpenStoreError> {
183        let conn = pool.get().await?;
184
185        let version = conn.db_version().await?;
186        debug!("Opened sqlite store with version {}", version);
187
188        let version = initialize_store(&conn, version).await?;
189
190        let store = Self::create_raw(secret, pool, conn, pool_config, runtime_config).await?;
191
192        run_migrations(&store, version, None).await?;
193
194        store.write().await?.wal_checkpoint().await;
195
196        Ok(store)
197    }
198
199    fn deserialize_and_unpickle_inbound_group_session(
200        &self,
201        value: Vec<u8>,
202        backed_up: bool,
203    ) -> Result<InboundGroupSession> {
204        let mut pickle: PickledInboundGroupSession = self.deserialize_value(&value)?;
205
206        // The `backed_up` SQL column is the source of truth, because we update
207        // it inside `mark_inbound_group_sessions_as_backed_up` and don't update
208        // the pickled value inside the `data` column (until now, when we are
209        // puling it out of the DB).
210        pickle.backed_up = backed_up;
211
212        Ok(InboundGroupSession::from_pickle(pickle)?)
213    }
214
215    fn deserialize_key_request(&self, value: &[u8], sent_out: bool) -> Result<GossipRequest> {
216        let mut request: GossipRequest = self.deserialize_value(value)?;
217        // sent_out SQL column is source of truth, sent_out field in serialized
218        // value needed for other stores though
219        request.sent_out = sent_out;
220        Ok(request)
221    }
222
223    fn get_static_account(&self) -> Option<StaticAccountData> {
224        self.static_account.read().unwrap().clone()
225    }
226
227    /// Acquire a connection for executing read operations.
228    #[instrument(skip_all)]
229    async fn read(&self) -> Result<SqliteAsyncConn> {
230        let pool = {
231            let guard = self.connections.lock().await;
232            let conns = guard.as_ref().ok_or(Error::StoreClosed)?;
233            conns.pool.clone()
234        };
235        Ok(pool.get().await?)
236    }
237
238    /// Acquire a connection for executing write operations.
239    #[instrument(skip_all)]
240    pub(crate) async fn write(&self) -> Result<OwnedMutexGuard<SqliteAsyncConn>> {
241        let write_connection = {
242            let guard = self.connections.lock().await;
243            let conns = guard.as_ref().ok_or(Error::StoreClosed)?;
244            conns.write_connection.clone()
245        };
246        Ok(write_connection.lock_owned().await)
247    }
248}
249
250const DATABASE_VERSION: u8 = 19;
251
252/// key for the dehydrated device pickle key in the key/value table.
253const DEHYDRATED_DEVICE_PICKLE_KEY: &str = "dehydrated_device_pickle_key";
254
255/// Initialize the database to version 1
256///
257/// This must be done before creating the store cipher, because the store cipher
258/// requires the `kv` table.
259///
260/// # Arguments
261///
262/// - `conn` - The connection to use.
263/// - `version` - the current version of the database.
264pub(crate) async fn initialize_store(conn: &SqliteAsyncConn, version: u8) -> Result<u8> {
265    if version == 0 {
266        debug!("Creating database");
267    } else if version < DATABASE_VERSION {
268        debug!(version, new_version = DATABASE_VERSION, "Upgrading database");
269    } else {
270        return Ok(version);
271    }
272
273    if version < 1 {
274        debug!("Creating database");
275        // First turn on WAL mode, this can't be done in the transaction, it
276        // fails with the error message: "cannot change into wal mode from
277        // within a transaction".
278        conn.execute_batch("PRAGMA journal_mode = wal;").await?;
279        conn.with_transaction(|txn| {
280            txn.execute_batch(include_str!("../migrations/crypto_store/001_init.sql"))?;
281            txn.set_db_version(1)
282        })
283        .await?;
284        return Ok(1);
285    }
286
287    Ok(version)
288}
289
290/// Run migrations for the given version of the database.
291///
292/// # Arguments
293///
294/// - `store` - The store to run the migrations on
295/// - `version` - The current version of the database.
296/// - `max_version` - The maximum version that the database will be migrated to.
297///   Only used for testing, so will only be checked for the versions that are
298///   needed for tests.
299pub(crate) async fn run_migrations(
300    store: &SqliteCryptoStore,
301    version: u8,
302    max_version: Option<u8>,
303) -> Result<()> {
304    let conn = store.write().await?;
305
306    if version < 2 {
307        debug!("Upgrading database to version 2");
308        conn.with_transaction(|txn| {
309            txn.execute_batch(include_str!("../migrations/crypto_store/002_reset_olm_hash.sql"))?;
310            txn.set_db_version(2)
311        })
312        .await?;
313    }
314
315    if version < 3 {
316        debug!("Upgrading database to version 3");
317        conn.with_transaction(|txn| {
318            txn.execute_batch(include_str!("../migrations/crypto_store/003_room_settings.sql"))?;
319            txn.set_db_version(3)
320        })
321        .await?;
322    }
323
324    if version < 4 {
325        debug!("Upgrading database to version 4");
326        conn.with_transaction(|txn| {
327            txn.execute_batch(include_str!(
328                "../migrations/crypto_store/004_drop_outbound_group_sessions.sql"
329            ))?;
330            txn.set_db_version(4)
331        })
332        .await?;
333    }
334
335    if version < 5 {
336        debug!("Upgrading database to version 5");
337        conn.with_transaction(|txn| {
338            txn.execute_batch(include_str!("../migrations/crypto_store/005_withheld_code.sql"))?;
339            txn.set_db_version(5)
340        })
341        .await?;
342    }
343
344    if version < 6 {
345        debug!("Upgrading database to version 6");
346        conn.with_transaction(|txn| {
347            txn.execute_batch(include_str!(
348                "../migrations/crypto_store/006_drop_outbound_group_sessions.sql"
349            ))?;
350            txn.set_db_version(6)
351        })
352        .await?;
353    }
354
355    if version < 7 {
356        debug!("Upgrading database to version 7");
357        conn.with_transaction(|txn| {
358            txn.execute_batch(include_str!("../migrations/crypto_store/007_lock_leases.sql"))?;
359            txn.set_db_version(7)
360        })
361        .await?;
362    }
363
364    if version < 8 {
365        debug!("Upgrading database to version 8");
366        conn.with_transaction(|txn| {
367            txn.execute_batch(include_str!("../migrations/crypto_store/008_secret_inbox.sql"))?;
368            txn.set_db_version(8)
369        })
370        .await?;
371    }
372
373    if version < 9 {
374        debug!("Upgrading database to version 9");
375        conn.with_transaction(|txn| {
376            txn.execute_batch(include_str!(
377                "../migrations/crypto_store/009_inbound_group_session_sender_key_sender_data_type.sql"
378            ))?;
379            txn.set_db_version(9)
380        })
381        .await?;
382    }
383
384    if version < 10 {
385        debug!("Upgrading database to version 10");
386        conn.with_transaction(|txn| {
387            txn.execute_batch(include_str!(
388                "../migrations/crypto_store/010_received_room_key_bundles.sql"
389            ))?;
390            txn.set_db_version(10)
391        })
392        .await?;
393    }
394
395    if version < 11 {
396        debug!("Upgrading database to version 11");
397        conn.with_transaction(|txn| {
398            txn.execute_batch(include_str!(
399                "../migrations/crypto_store/011_received_room_key_bundles_with_curve_key.sql"
400            ))?;
401            txn.set_db_version(11)
402        })
403        .await?;
404    }
405
406    if version < 12 {
407        debug!("Upgrading database to version 12");
408        conn.with_transaction(|txn| {
409            txn.execute_batch(include_str!(
410                "../migrations/crypto_store/012_withheld_code_by_room.sql"
411            ))?;
412            txn.set_db_version(12)
413        })
414        .await?;
415    }
416
417    if version < 13 {
418        debug!("Upgrading database to version 13");
419        conn.with_transaction(|txn| {
420            txn.execute_batch(include_str!(
421                "../migrations/crypto_store/013_lease_locks_with_generation.sql"
422            ))?;
423            txn.set_db_version(13)
424        })
425        .await?;
426    }
427
428    if version < 14 {
429        debug!("Upgrading database to version 14");
430        conn.with_transaction(|txn| {
431            txn.execute_batch(include_str!(
432                "../migrations/crypto_store/014_room_key_backups_fully_downloaded.sql"
433            ))?;
434            txn.set_db_version(14)
435        })
436        .await?;
437    }
438
439    if version < 15 {
440        debug!("Upgrading database to version 15");
441        conn.with_transaction(|txn| {
442            txn.execute_batch(include_str!(
443                "../migrations/crypto_store/015_rooms_pending_key_bundle.sql"
444            ))?;
445            txn.set_db_version(15)
446        })
447        .await?;
448    }
449
450    if version < 16 {
451        debug!("Upgrading database to version 16");
452        conn.with_transaction(|txn| {
453            txn.execute_batch(include_str!(
454                "../migrations/crypto_store/016_remove_old_generation_counter.sql"
455            ))?;
456            txn.set_db_version(16)
457        })
458        .await?;
459    }
460
461    if max_version.is_some_and(|max_version| max_version < 17) {
462        return Ok(());
463    }
464
465    if version < 17 {
466        debug!("Upgrading database to version 17");
467        let store = store.clone();
468        conn.with_transaction(move |txn| {
469            txn.execute_batch(include_str!(
470                "../migrations/crypto_store/017_add_new_secrets_inbox.sql"
471            ))?;
472            let mut select_query = txn.prepare("SELECT data FROM secrets")?;
473            let mut secrets = select_query.query([])?;
474            let mut insert_query = txn.prepare(
475                "INSERT OR IGNORE INTO secrets_inbox (secret_name, secret)
476            VALUES (?1, ?2)",
477            )?;
478            while let Some(row) = secrets.next()? {
479                let Ok(secret) =
480                    store.deserialize_json::<GossippedSecret>(row.get::<_, Vec<u8>>(0)?.as_ref())
481                else {
482                    continue;
483                };
484                let Ok(encoded_secret) = store.serialize_json(&secret.event.content.secret) else {
485                    continue;
486                };
487                insert_query.execute((
488                    store.encode_key("secrets_inbox", secret.secret_name.to_string()),
489                    &encoded_secret,
490                ))?;
491            }
492            txn.execute_batch(include_str!(
493                "../migrations/crypto_store/017_drop_old_secrets_inbox.sql"
494            ))?;
495            txn.set_db_version(17)
496        })
497        .await?;
498    }
499
500    if version < 18 {
501        debug!("Upgrading database to version 18");
502        let store = store.clone();
503        conn.with_transaction(move |txn| {
504            txn.execute_batch(include_str!(
505                "../migrations/crypto_store/018_add_gossip_request_info.sql"
506            ))?;
507            let mut select_query =
508                txn.prepare("SELECT request_id, sent_out, data FROM key_requests")?;
509            let mut requests = select_query.query([])?;
510            let mut update_query =
511                txn.prepare("UPDATE OR REPLACE key_requests SET info = ?1 WHERE request_id = ?2")?;
512            while let Some(row) = requests.next()? {
513                let Ok(request) = store.deserialize_key_request(
514                    row.get::<_, Vec<u8>>(2)?.as_ref(),
515                    row.get::<_, bool>(1)?,
516                ) else {
517                    continue;
518                };
519                let info = store.encode_key("key_requests", request.info.as_key());
520                update_query.execute((info, row.get::<_, Vec<u8>>(0)?))?;
521            }
522            txn.set_db_version(18)
523        })
524        .await?;
525    }
526
527    if version < 19 {
528        debug!("Upgrading database to version 19");
529        // Remove the sliding sync `pos` value stored in the crypto store. There
530        // was recently an event cache migration that emptied the cache but
531        // never reset the `pos` value, this fixes it.
532        let user_id = store.load_account().await?.map(|account| account.user_id.clone());
533
534        conn.with_transaction(move |txn| {
535            if let Some(user_id) = user_id {
536                txn.clear_kv(&format!("sliding_sync_store::room-list::{user_id}::instance"))?;
537            }
538            txn.set_db_version(19)
539        })
540        .await?;
541    }
542
543    Ok(())
544}
545
546trait SqliteConnectionExt {
547    fn set_session(
548        &self,
549        session_id: &[u8],
550        sender_key: &[u8],
551        data: &[u8],
552    ) -> rusqlite::Result<()>;
553
554    fn set_inbound_group_session(
555        &self,
556        room_id: &[u8],
557        session_id: &[u8],
558        data: &[u8],
559        backed_up: bool,
560        sender_key: Option<&[u8]>,
561        sender_data_type: Option<u8>,
562    ) -> rusqlite::Result<()>;
563
564    fn set_outbound_group_session(&self, room_id: &[u8], data: &[u8]) -> rusqlite::Result<()>;
565
566    fn set_device(&self, user_id: &[u8], device_id: &[u8], data: &[u8]) -> rusqlite::Result<()>;
567    fn delete_device(&self, user_id: &[u8], device_id: &[u8]) -> rusqlite::Result<()>;
568
569    fn set_identity(&self, user_id: &[u8], data: &[u8]) -> rusqlite::Result<()>;
570
571    fn add_olm_hash(&self, data: &[u8]) -> rusqlite::Result<()>;
572
573    fn set_key_request(
574        &self,
575        request_id: &[u8],
576        sent_out: bool,
577        data: &[u8],
578        info: &[u8],
579    ) -> rusqlite::Result<()>;
580
581    fn set_direct_withheld(
582        &self,
583        session_id: &[u8],
584        room_id: &[u8],
585        data: &[u8],
586    ) -> rusqlite::Result<()>;
587
588    fn set_room_settings(&self, room_id: &[u8], data: &[u8]) -> rusqlite::Result<()>;
589
590    fn set_secret(&self, request_id: &[u8], data: &[u8]) -> rusqlite::Result<()>;
591
592    fn set_received_room_key_bundle(
593        &self,
594        room_id: &[u8],
595        user_id: &[u8],
596        data: &[u8],
597    ) -> rusqlite::Result<()>;
598
599    fn set_has_downloaded_all_room_keys(&self, room_id: &[u8]) -> rusqlite::Result<()>;
600
601    fn set_room_pending_key_bundle(
602        &self,
603        room_id: &[u8],
604        details: Option<&[u8]>,
605    ) -> rusqlite::Result<()>;
606}
607
608impl SqliteConnectionExt for rusqlite::Connection {
609    fn set_session(
610        &self,
611        session_id: &[u8],
612        sender_key: &[u8],
613        data: &[u8],
614    ) -> rusqlite::Result<()> {
615        self.execute(
616            "INSERT INTO session (session_id, sender_key, data)
617             VALUES (?1, ?2, ?3)
618             ON CONFLICT (session_id) DO UPDATE SET data = ?3",
619            (session_id, sender_key, data),
620        )?;
621        Ok(())
622    }
623
624    fn set_inbound_group_session(
625        &self,
626        room_id: &[u8],
627        session_id: &[u8],
628        data: &[u8],
629        backed_up: bool,
630        sender_key: Option<&[u8]>,
631        sender_data_type: Option<u8>,
632    ) -> rusqlite::Result<()> {
633        self.execute(
634            "INSERT INTO inbound_group_session (session_id, room_id, data, backed_up, sender_key, sender_data_type) \
635             VALUES (?1, ?2, ?3, ?4, ?5, ?6)
636             ON CONFLICT (session_id) DO UPDATE SET data = ?3, backed_up = ?4, sender_key = ?5, sender_data_type = ?6",
637            (session_id, room_id, data, backed_up, sender_key, sender_data_type),
638        )?;
639        Ok(())
640    }
641
642    fn set_outbound_group_session(&self, room_id: &[u8], data: &[u8]) -> rusqlite::Result<()> {
643        self.execute(
644            "INSERT INTO outbound_group_session (room_id, data) \
645             VALUES (?1, ?2)
646             ON CONFLICT (room_id) DO UPDATE SET data = ?2",
647            (room_id, data),
648        )?;
649        Ok(())
650    }
651
652    fn set_device(&self, user_id: &[u8], device_id: &[u8], data: &[u8]) -> rusqlite::Result<()> {
653        self.execute(
654            "INSERT INTO device (user_id, device_id, data) \
655             VALUES (?1, ?2, ?3)
656             ON CONFLICT (user_id, device_id) DO UPDATE SET data = ?3",
657            (user_id, device_id, data),
658        )?;
659        Ok(())
660    }
661
662    fn delete_device(&self, user_id: &[u8], device_id: &[u8]) -> rusqlite::Result<()> {
663        self.execute(
664            "DELETE FROM device WHERE user_id = ? AND device_id = ?",
665            (user_id, device_id),
666        )?;
667        Ok(())
668    }
669
670    fn set_identity(&self, user_id: &[u8], data: &[u8]) -> rusqlite::Result<()> {
671        self.execute(
672            "INSERT INTO identity (user_id, data) \
673             VALUES (?1, ?2)
674             ON CONFLICT (user_id) DO UPDATE SET data = ?2",
675            (user_id, data),
676        )?;
677        Ok(())
678    }
679
680    fn add_olm_hash(&self, data: &[u8]) -> rusqlite::Result<()> {
681        self.execute("INSERT INTO olm_hash (data) VALUES (?) ON CONFLICT DO NOTHING", (data,))?;
682        Ok(())
683    }
684
685    fn set_key_request(
686        &self,
687        request_id: &[u8],
688        sent_out: bool,
689        data: &[u8],
690        info: &[u8],
691    ) -> rusqlite::Result<()> {
692        // The first `ON CONFLICT` cause will update a request if we try to save
693        // it again. The second `ON CONFLICT` will replace an old request for
694        // the same key/secret with the new request.
695        self.execute(
696            "INSERT INTO key_requests (request_id, sent_out, data, info)
697            VALUES (?1, ?2, ?3, ?4)
698            ON CONFLICT (request_id) DO UPDATE SET sent_out = ?2, data = ?3, info = ?4
699            ON CONFLICT (info) DO UPDATE SET request_id = ?1, sent_out = ?2, data = ?3",
700            (request_id, sent_out, data, info),
701        )?;
702        Ok(())
703    }
704
705    fn set_direct_withheld(
706        &self,
707        session_id: &[u8],
708        room_id: &[u8],
709        data: &[u8],
710    ) -> rusqlite::Result<()> {
711        self.execute(
712            "INSERT INTO direct_withheld_info (session_id, room_id, data)
713            VALUES (?1, ?2, ?3)
714            ON CONFLICT (session_id) DO UPDATE SET room_id = ?2, data = ?3",
715            (session_id, room_id, data),
716        )?;
717        Ok(())
718    }
719
720    fn set_room_settings(&self, room_id: &[u8], data: &[u8]) -> rusqlite::Result<()> {
721        self.execute(
722            "INSERT INTO room_settings (room_id, data)
723            VALUES (?1, ?2)
724            ON CONFLICT (room_id) DO UPDATE SET data = ?2",
725            (room_id, data),
726        )?;
727        Ok(())
728    }
729
730    fn set_secret(&self, secret_name: &[u8], secret: &[u8]) -> rusqlite::Result<()> {
731        // Ignore duplicate values, since we may get set the same secret
732        // multiple times.
733        self.execute(
734            "INSERT OR IGNORE INTO secrets_inbox (secret_name, secret)
735            VALUES (?1, ?2)",
736            (secret_name, secret),
737        )?;
738
739        Ok(())
740    }
741
742    fn set_received_room_key_bundle(
743        &self,
744        room_id: &[u8],
745        sender_user_id: &[u8],
746        data: &[u8],
747    ) -> rusqlite::Result<()> {
748        self.execute(
749            "INSERT INTO received_room_key_bundle(room_id, sender_user_id, bundle_data)
750            VALUES (?1, ?2, ?3)
751            ON CONFLICT (room_id, sender_user_id) DO UPDATE SET bundle_data = ?3",
752            (room_id, sender_user_id, data),
753        )?;
754        Ok(())
755    }
756
757    fn set_room_pending_key_bundle(
758        &self,
759        room_id: &[u8],
760        data: Option<&[u8]>,
761    ) -> rusqlite::Result<()> {
762        if let Some(data) = data {
763            self.execute(
764                "INSERT INTO rooms_pending_key_bundle (room_id, data)
765                 VALUES (?1, ?2)
766                 ON CONFLICT (room_id) DO UPDATE SET data = ?2",
767                (room_id, data),
768            )?;
769        } else {
770            self.execute("DELETE FROM rooms_pending_key_bundle WHERE room_id = ?1", (room_id,))?;
771        }
772        Ok(())
773    }
774
775    fn set_has_downloaded_all_room_keys(&self, room_id: &[u8]) -> rusqlite::Result<()> {
776        self.execute(
777            "INSERT INTO room_key_backups_fully_downloaded(room_id)
778             VALUES (?1)
779             ON CONFLICT(room_id) DO NOTHING",
780            (room_id,),
781        )?;
782        Ok(())
783    }
784}
785
786#[async_trait]
787trait SqliteObjectCryptoStoreExt: SqliteAsyncConnExt {
788    async fn get_sessions_for_sender_key(&self, sender_key: Key) -> Result<Vec<Vec<u8>>> {
789        Ok(self
790            .prepare("SELECT data FROM session WHERE sender_key = ?", |mut stmt| {
791                stmt.query((sender_key,))?.mapped(|row| row.get(0)).collect()
792            })
793            .await?)
794    }
795
796    async fn get_inbound_group_session(
797        &self,
798        session_id: Key,
799    ) -> Result<Option<(Vec<u8>, Vec<u8>, bool)>> {
800        Ok(self
801            .query_one(
802                "SELECT room_id, data, backed_up FROM inbound_group_session WHERE session_id = ?",
803                (session_id,),
804                |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?)),
805            )
806            .await
807            .optional()?)
808    }
809
810    async fn get_inbound_group_sessions(&self) -> Result<Vec<(Vec<u8>, bool)>> {
811        Ok(self
812            .prepare("SELECT data, backed_up FROM inbound_group_session", |mut stmt| {
813                stmt.query(())?.mapped(|row| Ok((row.get(0)?, row.get(1)?))).collect()
814            })
815            .await?)
816    }
817
818    async fn get_inbound_group_session_counts(
819        &self,
820        _backup_version: Option<&str>,
821    ) -> Result<RoomKeyCounts> {
822        let total = self
823            .query_one("SELECT count(*) FROM inbound_group_session", (), |row| row.get(0))
824            .await?;
825        let backed_up = self
826            .query_one(
827                "SELECT count(*) FROM inbound_group_session WHERE backed_up = TRUE",
828                (),
829                |row| row.get(0),
830            )
831            .await?;
832        Ok(RoomKeyCounts { total, backed_up })
833    }
834
835    async fn get_inbound_group_sessions_by_room_id(
836        &self,
837        room_id: Key,
838    ) -> Result<Vec<(Vec<u8>, bool)>> {
839        Ok(self
840            .prepare(
841                "SELECT data, backed_up FROM inbound_group_session WHERE room_id = :room_id",
842                move |mut stmt| {
843                    stmt.query(named_params! {
844                        ":room_id": room_id,
845                    })?
846                    .mapped(|row| Ok((row.get(0)?, row.get(1)?)))
847                    .collect()
848                },
849            )
850            .await?)
851    }
852
853    async fn get_inbound_group_sessions_for_device_batch(
854        &self,
855        sender_key: Key,
856        sender_data_type: SenderDataType,
857        after_session_id: Option<Key>,
858        limit: usize,
859    ) -> Result<Vec<(Vec<u8>, bool)>> {
860        Ok(self
861            .prepare(
862                "
863                SELECT data, backed_up
864                FROM inbound_group_session
865                WHERE sender_key = :sender_key
866                    AND sender_data_type = :sender_data_type
867                    AND session_id > :after_session_id
868                ORDER BY session_id
869                LIMIT :limit
870                ",
871                move |mut stmt| {
872                    let sender_data_type = sender_data_type as u8;
873
874                    // If we are not provided with an `after_session_id`, use a
875                    // key which will sort before all real keys: the empty
876                    // string.
877                    let after_session_id = after_session_id.unwrap_or(Key::Plain(Vec::new()));
878
879                    stmt.query(named_params! {
880                        ":sender_key": sender_key,
881                        ":sender_data_type": sender_data_type,
882                        ":after_session_id": after_session_id,
883                        ":limit": limit,
884                    })?
885                    .mapped(|row| Ok((row.get(0)?, row.get(1)?)))
886                    .collect()
887                },
888            )
889            .await?)
890    }
891
892    async fn get_inbound_group_sessions_for_backup(&self, limit: usize) -> Result<Vec<Vec<u8>>> {
893        Ok(self
894            .prepare(
895                "SELECT data FROM inbound_group_session WHERE backed_up = FALSE LIMIT ?",
896                move |mut stmt| stmt.query((limit,))?.mapped(|row| row.get(0)).collect(),
897            )
898            .await?)
899    }
900
901    async fn mark_inbound_group_sessions_as_backed_up(&self, session_ids: Vec<Key>) -> Result<()> {
902        if session_ids.is_empty() {
903            // We are not expecting to be called with an empty list of sessions
904            warn!("No sessions to mark as backed up!");
905            return Ok(());
906        }
907
908        self.chunk_large_query_over(session_ids, None, move |txn, session_ids| {
909            // Safety: host parameters are not generated using any user input
910            // except the number of session IDs, so it is safe from injection.
911            let query = format!(
912                "UPDATE inbound_group_session SET backed_up = TRUE where session_id IN ({})",
913                session_ids.host_parameters()
914            );
915            txn.prepare(&query)?.execute(params_from_iter(session_ids))?;
916            Ok(Vec::<()>::new())
917        })
918        .await?;
919
920        Ok(())
921    }
922
923    async fn reset_inbound_group_session_backup_state(&self) -> Result<()> {
924        self.execute("UPDATE inbound_group_session SET backed_up = FALSE", ()).await?;
925        Ok(())
926    }
927
928    async fn get_outbound_group_session(&self, room_id: Key) -> Result<Option<Vec<u8>>> {
929        Ok(self
930            .query_one(
931                "SELECT data FROM outbound_group_session WHERE room_id = ?",
932                (room_id,),
933                |row| row.get(0),
934            )
935            .await
936            .optional()?)
937    }
938
939    async fn get_device(&self, user_id: Key, device_id: Key) -> Result<Option<Vec<u8>>> {
940        Ok(self
941            .query_one(
942                "SELECT data FROM device WHERE user_id = ? AND device_id = ?",
943                (user_id, device_id),
944                |row| row.get(0),
945            )
946            .await
947            .optional()?)
948    }
949
950    async fn get_user_devices(&self, user_id: Key) -> Result<Vec<Vec<u8>>> {
951        Ok(self
952            .prepare("SELECT data FROM device WHERE user_id = ?", |mut stmt| {
953                stmt.query((user_id,))?.mapped(|row| row.get(0)).collect()
954            })
955            .await?)
956    }
957
958    async fn get_user_identity(&self, user_id: Key) -> Result<Option<Vec<u8>>> {
959        Ok(self
960            .query_one("SELECT data FROM identity WHERE user_id = ?", (user_id,), |row| row.get(0))
961            .await
962            .optional()?)
963    }
964
965    async fn has_olm_hash(&self, data: Vec<u8>) -> Result<bool> {
966        Ok(self
967            .query_one("SELECT count(*) FROM olm_hash WHERE data = ?", (data,), |row| {
968                row.get::<_, i32>(0)
969            })
970            .await?
971            > 0)
972    }
973
974    async fn get_tracked_users(&self) -> Result<Vec<Vec<u8>>> {
975        Ok(self
976            .prepare("SELECT data FROM tracked_user", |mut stmt| {
977                stmt.query(())?.mapped(|row| row.get(0)).collect()
978            })
979            .await?)
980    }
981
982    async fn add_tracked_users(&self, users: Vec<(Key, Vec<u8>)>) -> Result<()> {
983        Ok(self
984            .prepare(
985                "INSERT INTO tracked_user (user_id, data) \
986                 VALUES (?1, ?2) \
987                 ON CONFLICT (user_id) DO UPDATE SET data = ?2",
988                |mut stmt| {
989                    for (user_id, data) in users {
990                        stmt.execute((user_id, data))?;
991                    }
992
993                    Ok(())
994                },
995            )
996            .await?)
997    }
998
999    async fn get_outgoing_secret_request(
1000        &self,
1001        request_id: Key,
1002    ) -> Result<Option<(Vec<u8>, bool)>> {
1003        Ok(self
1004            .query_one(
1005                "SELECT data, sent_out FROM key_requests WHERE request_id = ?",
1006                (request_id,),
1007                |row| Ok((row.get(0)?, row.get(1)?)),
1008            )
1009            .await
1010            .optional()?)
1011    }
1012
1013    async fn get_secret_request_by_info(&self, info: Key) -> Result<Option<(Vec<u8>, bool)>> {
1014        Ok(self
1015            .query_one("SELECT data, sent_out FROM key_requests WHERE info = ?", (info,), |row| {
1016                Ok((row.get(0)?, row.get(1)?))
1017            })
1018            .await
1019            .optional()?)
1020    }
1021
1022    async fn get_unsent_secret_requests(&self) -> Result<Vec<Vec<u8>>> {
1023        Ok(self
1024            .prepare("SELECT data FROM key_requests WHERE sent_out = FALSE", |mut stmt| {
1025                stmt.query(())?.mapped(|row| row.get(0)).collect()
1026            })
1027            .await?)
1028    }
1029
1030    async fn delete_key_request(&self, request_id: Key) -> Result<()> {
1031        self.execute("DELETE FROM key_requests WHERE request_id = ?", (request_id,)).await?;
1032        Ok(())
1033    }
1034
1035    async fn get_secrets_from_inbox(&self, secret_name: Key) -> Result<Vec<Vec<u8>>> {
1036        Ok(self
1037            .prepare("SELECT secret FROM secrets_inbox WHERE secret_name = ?", |mut stmt| {
1038                stmt.query((secret_name,))?.mapped(|row| row.get(0)).collect()
1039            })
1040            .await?)
1041    }
1042
1043    async fn delete_secrets_from_inbox(&self, secret_name: Key) -> Result<()> {
1044        self.execute("DELETE FROM secrets_inbox WHERE secret_name = ?", (secret_name,)).await?;
1045        Ok(())
1046    }
1047
1048    async fn get_direct_withheld_info(
1049        &self,
1050        session_id: Key,
1051        room_id: Key,
1052    ) -> Result<Option<Vec<u8>>> {
1053        Ok(self
1054            .query_one(
1055                "SELECT data FROM direct_withheld_info WHERE session_id = ?1 AND room_id = ?2",
1056                (session_id, room_id),
1057                |row| row.get(0),
1058            )
1059            .await
1060            .optional()?)
1061    }
1062
1063    async fn get_withheld_sessions_by_room_id(&self, room_id: Key) -> Result<Vec<Vec<u8>>> {
1064        Ok(self
1065            .prepare("SELECT data FROM direct_withheld_info WHERE room_id = ?1", |mut stmt| {
1066                stmt.query((room_id,))?.mapped(|row| row.get(0)).collect()
1067            })
1068            .await?)
1069    }
1070
1071    async fn get_room_settings(&self, room_id: Key) -> Result<Option<Vec<u8>>> {
1072        Ok(self
1073            .query_one("SELECT data FROM room_settings WHERE room_id = ?", (room_id,), |row| {
1074                row.get(0)
1075            })
1076            .await
1077            .optional()?)
1078    }
1079
1080    async fn get_received_room_key_bundle(
1081        &self,
1082        room_id: Key,
1083        sender_user: Key,
1084    ) -> Result<Option<Vec<u8>>> {
1085        Ok(self
1086            .query_one(
1087                "SELECT bundle_data FROM received_room_key_bundle WHERE room_id = ? AND sender_user_id = ?",
1088                (room_id, sender_user),
1089                |row| { row.get(0) },
1090            )
1091            .await
1092            .optional()?)
1093    }
1094
1095    async fn get_room_pending_key_bundle(&self, room_id: Key) -> Result<Option<Vec<u8>>> {
1096        Ok(self
1097            .query_one(
1098                "SELECT data FROM rooms_pending_key_bundle WHERE room_id = ?",
1099                (room_id,),
1100                |row| row.get(0),
1101            )
1102            .await
1103            .optional()?)
1104    }
1105
1106    async fn get_all_rooms_pending_key_bundle(&self) -> Result<Vec<Vec<u8>>> {
1107        Ok(self
1108            .query_many("SELECT data FROM rooms_pending_key_bundle", (), |row| row.get(0))
1109            .await?)
1110    }
1111
1112    async fn has_downloaded_all_room_keys(&self, room_id: Key) -> Result<bool> {
1113        Ok(self
1114            .query_row(
1115                "SELECT EXISTS (SELECT 1 FROM room_key_backups_fully_downloaded WHERE room_id = ?)",
1116                (room_id,),
1117                |row| row.get(0),
1118            )
1119            .await?)
1120    }
1121}
1122
1123#[async_trait]
1124impl SqliteObjectCryptoStoreExt for SqliteAsyncConn {}
1125
1126#[async_trait]
1127impl CryptoStore for SqliteCryptoStore {
1128    type Error = Error;
1129
1130    async fn load_account(&self) -> Result<Option<Account>> {
1131        let conn = self.read().await?;
1132        if let Some(pickle) = conn.get_kv("account").await? {
1133            let pickle = self.deserialize_value(&pickle)?;
1134
1135            let account = Account::from_pickle(pickle).map_err(|_| Error::Unpickle)?;
1136
1137            *self.static_account.write().unwrap() = Some(account.static_data().clone());
1138
1139            Ok(Some(account))
1140        } else {
1141            Ok(None)
1142        }
1143    }
1144
1145    async fn load_identity(&self) -> Result<Option<PrivateCrossSigningIdentity>> {
1146        let conn = self.read().await?;
1147        if let Some(i) = conn.get_kv("identity").await? {
1148            let pickle = self.deserialize_value(&i)?;
1149            Ok(Some(PrivateCrossSigningIdentity::from_pickle(pickle).map_err(|_| Error::Unpickle)?))
1150        } else {
1151            Ok(None)
1152        }
1153    }
1154
1155    async fn save_pending_changes(&self, changes: PendingChanges) -> Result<()> {
1156        // Serialize calls to `save_pending_changes`; there are multiple await
1157        // points below, and we're pickling data as we go, so we don't want to
1158        // invalidate data we've previously read and overwrite it in the store.
1159        // TODO: #2000 should make this lock go away, or change its shape.
1160        let _guard = self.save_changes_lock.lock().await;
1161
1162        let pickled_account = if let Some(account) = changes.account {
1163            *self.static_account.write().unwrap() = Some(account.static_data().clone());
1164            Some(account.pickle())
1165        } else {
1166            None
1167        };
1168
1169        let this = self.clone();
1170        self.write()
1171            .await?
1172            .with_transaction(move |txn| {
1173                if let Some(pickled_account) = pickled_account {
1174                    let serialized_account = this.serialize_value(&pickled_account)?;
1175                    txn.set_kv("account", &serialized_account)?;
1176                }
1177
1178                Ok::<_, Error>(())
1179            })
1180            .await?;
1181
1182        Ok(())
1183    }
1184
1185    async fn save_changes(&self, changes: Changes) -> Result<()> {
1186        // Serialize calls to `save_changes`; there are multiple await points
1187        // below, and we're pickling data as we go, so we don't want to
1188        // invalidate data we've previously read and overwrite it in the store.
1189        // TODO: #2000 should make this lock go away, or change its shape.
1190        let _guard = self.save_changes_lock.lock().await;
1191
1192        let pickled_private_identity =
1193            if let Some(i) = changes.private_identity { Some(i.pickle().await) } else { None };
1194
1195        let mut session_changes = Vec::new();
1196
1197        for session in changes.sessions {
1198            let session_id = self.encode_key("session", session.session_id());
1199            let sender_key = self.encode_key("session", session.sender_key().to_base64());
1200            let pickle = session.pickle().await;
1201            session_changes.push((session_id, sender_key, pickle));
1202        }
1203
1204        let mut inbound_session_changes = Vec::new();
1205        for session in changes.inbound_group_sessions {
1206            let room_id = self.encode_key("inbound_group_session", session.room_id().as_bytes());
1207            let session_id = self.encode_key("inbound_group_session", session.session_id());
1208            let pickle = session.pickle().await;
1209            let sender_key =
1210                self.encode_key("inbound_group_session", session.sender_key().to_base64());
1211            inbound_session_changes.push((room_id, session_id, pickle, sender_key));
1212        }
1213
1214        let mut outbound_session_changes = Vec::new();
1215        for session in changes.outbound_group_sessions {
1216            let room_id = self.encode_key("outbound_group_session", session.room_id().as_bytes());
1217            let pickle = session.pickle().await;
1218            outbound_session_changes.push((room_id, pickle));
1219        }
1220
1221        let this = self.clone();
1222        self.write()
1223            .await?
1224            .with_transaction(move |txn| {
1225                if let Some(pickled_private_identity) = &pickled_private_identity {
1226                    let serialized_private_identity =
1227                        this.serialize_value(pickled_private_identity)?;
1228                    txn.set_kv("identity", &serialized_private_identity)?;
1229                }
1230
1231                if let Some(token) = &changes.next_batch_token {
1232                    let serialized_token = this.serialize_value(token)?;
1233                    txn.set_kv("next_batch_token", &serialized_token)?;
1234                }
1235
1236                if let Some(decryption_key) = &changes.backup_decryption_key {
1237                    let serialized_decryption_key = this.serialize_value(decryption_key)?;
1238                    txn.set_kv("recovery_key_v1", &serialized_decryption_key)?;
1239                }
1240
1241                if let Some(backup_version) = &changes.backup_version {
1242                    let serialized_backup_version = this.serialize_value(backup_version)?;
1243                    txn.set_kv("backup_version_v1", &serialized_backup_version)?;
1244                }
1245
1246                if let Some(pickle_key) = &changes.dehydrated_device_pickle_key {
1247                    let serialized_pickle_key = this.serialize_value(pickle_key)?;
1248                    txn.set_kv(DEHYDRATED_DEVICE_PICKLE_KEY, &serialized_pickle_key)?;
1249                }
1250
1251                for device in changes.devices.new.iter().chain(&changes.devices.changed) {
1252                    let user_id = this.encode_key("device", device.user_id().as_bytes());
1253                    let device_id = this.encode_key("device", device.device_id().as_bytes());
1254                    let data = this.serialize_value(&device)?;
1255                    txn.set_device(&user_id, &device_id, &data)?;
1256                }
1257
1258                for device in &changes.devices.deleted {
1259                    let user_id = this.encode_key("device", device.user_id().as_bytes());
1260                    let device_id = this.encode_key("device", device.device_id().as_bytes());
1261                    txn.delete_device(&user_id, &device_id)?;
1262                }
1263
1264                for identity in changes.identities.changed.iter().chain(&changes.identities.new) {
1265                    let user_id = this.encode_key("identity", identity.user_id().as_bytes());
1266                    let data = this.serialize_value(&identity)?;
1267                    txn.set_identity(&user_id, &data)?;
1268                }
1269
1270                for (session_id, sender_key, pickle) in &session_changes {
1271                    let serialized_session = this.serialize_value(&pickle)?;
1272                    txn.set_session(session_id, sender_key, &serialized_session)?;
1273                }
1274
1275                for (room_id, session_id, pickle, sender_key) in &inbound_session_changes {
1276                    let serialized_session = this.serialize_value(&pickle)?;
1277                    txn.set_inbound_group_session(
1278                        room_id,
1279                        session_id,
1280                        &serialized_session,
1281                        pickle.backed_up,
1282                        Some(sender_key),
1283                        Some(pickle.sender_data.to_type() as u8),
1284                    )?;
1285                }
1286
1287                for (room_id, pickle) in &outbound_session_changes {
1288                    let serialized_session = this.serialize_json(&pickle)?;
1289                    txn.set_outbound_group_session(room_id, &serialized_session)?;
1290                }
1291
1292                for hash in &changes.message_hashes {
1293                    let hash = rmp_serde::to_vec(hash)?;
1294                    txn.add_olm_hash(&hash)?;
1295                }
1296
1297                for request in changes.key_requests {
1298                    let request_id = this.encode_key("key_requests", request.request_id.as_bytes());
1299                    let serialized_request = this.serialize_value(&request)?;
1300                    let serialized_info = this.encode_key("key_requests", request.info.as_key());
1301                    txn.set_key_request(
1302                        &request_id,
1303                        request.sent_out,
1304                        &serialized_request,
1305                        &serialized_info,
1306                    )?;
1307                }
1308
1309                for (room_id, data) in changes.withheld_session_info {
1310                    for (session_id, event) in data {
1311                        let session_id = this.encode_key("direct_withheld_info", session_id);
1312                        let room_id = this.encode_key("direct_withheld_info", &room_id);
1313                        let serialized_info = this.serialize_json(&event)?;
1314                        txn.set_direct_withheld(&session_id, &room_id, &serialized_info)?;
1315                    }
1316                }
1317
1318                for (room_id, settings) in changes.room_settings {
1319                    let room_id = this.encode_key("room_settings", room_id.as_bytes());
1320                    let value = this.serialize_value(&settings)?;
1321                    txn.set_room_settings(&room_id, &value)?;
1322                }
1323
1324                for secret in changes.secrets {
1325                    let secret_name =
1326                        this.encode_key("secrets_inbox", secret.secret_name.to_string());
1327                    let value = this.serialize_json(secret.secret.deref())?;
1328                    txn.set_secret(&secret_name, &value)?;
1329                }
1330
1331                for bundle in changes.received_room_key_bundles {
1332                    let room_id =
1333                        this.encode_key("received_room_key_bundle", &bundle.bundle_data.room_id);
1334                    let user_id = this.encode_key("received_room_key_bundle", &bundle.sender_user);
1335                    let value = this.serialize_value(&bundle)?;
1336                    txn.set_received_room_key_bundle(&room_id, &user_id, &value)?;
1337                }
1338
1339                for room in changes.room_key_backups_fully_downloaded {
1340                    let room_id = this.encode_key("room_key_backups_fully_downloaded", &room);
1341                    txn.set_has_downloaded_all_room_keys(&room_id)?;
1342                }
1343
1344                for (room, details) in changes.rooms_pending_key_bundle {
1345                    let room_id = this.encode_key("rooms_pending_key_bundle", &room);
1346                    let value = details.as_ref().map(|d| this.serialize_value(d)).transpose()?;
1347                    txn.set_room_pending_key_bundle(&room_id, value.as_deref())?;
1348                }
1349
1350                Ok::<_, Error>(())
1351            })
1352            .await?;
1353
1354        Ok(())
1355    }
1356
1357    async fn save_inbound_group_sessions(
1358        &self,
1359        sessions: Vec<InboundGroupSession>,
1360        backed_up_to_version: Option<&str>,
1361    ) -> matrix_sdk_crypto::store::Result<(), Self::Error> {
1362        // Sanity-check that the data in the sessions corresponds to
1363        // backed_up_version
1364        sessions.iter().for_each(|s| {
1365            let backed_up = s.backed_up();
1366            if backed_up != backed_up_to_version.is_some() {
1367                warn!(
1368                    backed_up,
1369                    backed_up_to_version,
1370                    "Session backed-up flag does not correspond to backup version setting",
1371                );
1372            }
1373        });
1374
1375        // Currently, this store doesn't save the backup version separately, so
1376        // this just delegates to save_changes.
1377        self.save_changes(Changes { inbound_group_sessions: sessions, ..Changes::default() }).await
1378    }
1379
1380    async fn get_sessions(&self, sender_key: &str) -> Result<Option<Vec<Session>>> {
1381        let device_keys = self.get_own_device().await?.as_device_keys().clone();
1382
1383        let sessions: Vec<_> = self
1384            .read()
1385            .await?
1386            .get_sessions_for_sender_key(self.encode_key("session", sender_key.as_bytes()))
1387            .await?
1388            .into_iter()
1389            .map(|bytes| {
1390                let pickle = self.deserialize_value(&bytes)?;
1391                Session::from_pickle(device_keys.clone(), pickle).map_err(|_| Error::AccountUnset)
1392            })
1393            .collect::<Result<_>>()?;
1394
1395        if sessions.is_empty() { Ok(None) } else { Ok(Some(sessions)) }
1396    }
1397
1398    #[instrument(skip(self))]
1399    async fn get_inbound_group_session(
1400        &self,
1401        room_id: &RoomId,
1402        session_id: &str,
1403    ) -> Result<Option<InboundGroupSession>> {
1404        let session_id = self.encode_key("inbound_group_session", session_id);
1405        let Some((room_id_from_db, value, backed_up)) =
1406            self.read().await?.get_inbound_group_session(session_id).await?
1407        else {
1408            return Ok(None);
1409        };
1410
1411        let room_id = self.encode_key("inbound_group_session", room_id.as_bytes());
1412        if *room_id != room_id_from_db {
1413            warn!("expected room_id for session_id doesn't match what's in the DB");
1414            return Ok(None);
1415        }
1416
1417        Ok(Some(self.deserialize_and_unpickle_inbound_group_session(value, backed_up)?))
1418    }
1419
1420    async fn get_inbound_group_sessions(&self) -> Result<Vec<InboundGroupSession>> {
1421        self.read()
1422            .await?
1423            .get_inbound_group_sessions()
1424            .await?
1425            .into_iter()
1426            .map(|(value, backed_up)| {
1427                self.deserialize_and_unpickle_inbound_group_session(value, backed_up)
1428            })
1429            .collect()
1430    }
1431
1432    async fn get_inbound_group_sessions_by_room_id(
1433        &self,
1434        room_id: &RoomId,
1435    ) -> Result<Vec<InboundGroupSession>> {
1436        let room_id = self.encode_key("inbound_group_session", room_id.as_bytes());
1437        self.read()
1438            .await?
1439            .get_inbound_group_sessions_by_room_id(room_id)
1440            .await?
1441            .into_iter()
1442            .map(|(value, backed_up)| {
1443                self.deserialize_and_unpickle_inbound_group_session(value, backed_up)
1444            })
1445            .collect()
1446    }
1447
1448    async fn get_inbound_group_sessions_for_device_batch(
1449        &self,
1450        sender_key: Curve25519PublicKey,
1451        sender_data_type: SenderDataType,
1452        after_session_id: Option<String>,
1453        limit: usize,
1454    ) -> Result<Vec<InboundGroupSession>, Self::Error> {
1455        let after_session_id =
1456            after_session_id.map(|session_id| self.encode_key("inbound_group_session", session_id));
1457        let sender_key = self.encode_key("inbound_group_session", sender_key.to_base64());
1458
1459        self.read()
1460            .await?
1461            .get_inbound_group_sessions_for_device_batch(
1462                sender_key,
1463                sender_data_type,
1464                after_session_id,
1465                limit,
1466            )
1467            .await?
1468            .into_iter()
1469            .map(|(value, backed_up)| {
1470                self.deserialize_and_unpickle_inbound_group_session(value, backed_up)
1471            })
1472            .collect()
1473    }
1474
1475    async fn inbound_group_session_counts(
1476        &self,
1477        backup_version: Option<&str>,
1478    ) -> Result<RoomKeyCounts> {
1479        Ok(self.read().await?.get_inbound_group_session_counts(backup_version).await?)
1480    }
1481
1482    async fn inbound_group_sessions_for_backup(
1483        &self,
1484        _backup_version: &str,
1485        limit: usize,
1486    ) -> Result<Vec<InboundGroupSession>> {
1487        self.read()
1488            .await?
1489            .get_inbound_group_sessions_for_backup(limit)
1490            .await?
1491            .into_iter()
1492            .map(|value| self.deserialize_and_unpickle_inbound_group_session(value, false))
1493            .collect()
1494    }
1495
1496    async fn mark_inbound_group_sessions_as_backed_up(
1497        &self,
1498        _backup_version: &str,
1499        session_ids: &[(&RoomId, &str)],
1500    ) -> Result<()> {
1501        Ok(self
1502            .write()
1503            .await?
1504            .mark_inbound_group_sessions_as_backed_up(
1505                session_ids
1506                    .iter()
1507                    .map(|(_, s)| self.encode_key("inbound_group_session", s))
1508                    .collect(),
1509            )
1510            .await?)
1511    }
1512
1513    async fn reset_backup_state(&self) -> Result<()> {
1514        Ok(self.write().await?.reset_inbound_group_session_backup_state().await?)
1515    }
1516
1517    async fn load_backup_keys(&self) -> Result<BackupKeys> {
1518        let conn = self.read().await?;
1519
1520        let backup_version = conn
1521            .get_kv("backup_version_v1")
1522            .await?
1523            .map(|value| self.deserialize_value(&value))
1524            .transpose()?;
1525
1526        let decryption_key = conn
1527            .get_kv("recovery_key_v1")
1528            .await?
1529            .map(|value| self.deserialize_value(&value))
1530            .transpose()?;
1531
1532        Ok(BackupKeys { backup_version, decryption_key })
1533    }
1534
1535    async fn load_dehydrated_device_pickle_key(&self) -> Result<Option<DehydratedDeviceKey>> {
1536        let conn = self.read().await?;
1537
1538        conn.get_kv(DEHYDRATED_DEVICE_PICKLE_KEY)
1539            .await?
1540            .map(|value| self.deserialize_value(&value))
1541            .transpose()
1542    }
1543
1544    async fn delete_dehydrated_device_pickle_key(&self) -> Result<(), Self::Error> {
1545        Ok(self.write().await?.clear_kv(DEHYDRATED_DEVICE_PICKLE_KEY).await?)
1546    }
1547    async fn get_outbound_group_session(
1548        &self,
1549        room_id: &RoomId,
1550    ) -> Result<Option<OutboundGroupSession>> {
1551        let room_id = self.encode_key("outbound_group_session", room_id.as_bytes());
1552        let Some(value) = self.read().await?.get_outbound_group_session(room_id).await? else {
1553            return Ok(None);
1554        };
1555
1556        let account_info = self.get_static_account().ok_or(Error::AccountUnset)?;
1557
1558        let pickle = self.deserialize_json(&value)?;
1559        let session = OutboundGroupSession::from_pickle(
1560            account_info.device_id,
1561            account_info.identity_keys,
1562            pickle,
1563        )
1564        .map_err(|_| Error::Unpickle)?;
1565
1566        return Ok(Some(session));
1567    }
1568
1569    async fn load_tracked_users(&self) -> Result<Vec<TrackedUser>> {
1570        self.read()
1571            .await?
1572            .get_tracked_users()
1573            .await?
1574            .iter()
1575            .map(|value| self.deserialize_value(value))
1576            .collect()
1577    }
1578
1579    async fn save_tracked_users(&self, tracked_users: &[(&UserId, bool)]) -> Result<()> {
1580        let users: Vec<(Key, Vec<u8>)> = tracked_users
1581            .iter()
1582            .map(|(u, d)| {
1583                let user_id = self.encode_key("tracked_users", u.as_bytes());
1584                let data =
1585                    self.serialize_value(&TrackedUser { user_id: (*u).into(), dirty: *d })?;
1586                Ok((user_id, data))
1587            })
1588            .collect::<Result<_>>()?;
1589
1590        Ok(self.write().await?.add_tracked_users(users).await?)
1591    }
1592
1593    async fn get_device(
1594        &self,
1595        user_id: &UserId,
1596        device_id: &DeviceId,
1597    ) -> Result<Option<DeviceData>> {
1598        let user_id = self.encode_key("device", user_id.as_bytes());
1599        let device_id = self.encode_key("device", device_id.as_bytes());
1600        Ok(self
1601            .read()
1602            .await?
1603            .get_device(user_id, device_id)
1604            .await?
1605            .map(|value| self.deserialize_value(&value))
1606            .transpose()?)
1607    }
1608
1609    async fn get_user_devices(
1610        &self,
1611        user_id: &UserId,
1612    ) -> Result<HashMap<OwnedDeviceId, DeviceData>> {
1613        let user_id = self.encode_key("device", user_id.as_bytes());
1614        self.read()
1615            .await?
1616            .get_user_devices(user_id)
1617            .await?
1618            .into_iter()
1619            .map(|value| {
1620                let device: DeviceData = self.deserialize_value(&value)?;
1621                Ok((device.device_id().to_owned(), device))
1622            })
1623            .collect()
1624    }
1625
1626    async fn get_own_device(&self) -> Result<DeviceData> {
1627        let account_info = self.get_static_account().ok_or(Error::AccountUnset)?;
1628
1629        Ok(self
1630            .get_device(&account_info.user_id, &account_info.device_id)
1631            .await?
1632            .expect("We should be able to find our own device."))
1633    }
1634
1635    async fn get_user_identity(&self, user_id: &UserId) -> Result<Option<UserIdentityData>> {
1636        let user_id = self.encode_key("identity", user_id.as_bytes());
1637        Ok(self
1638            .read()
1639            .await?
1640            .get_user_identity(user_id)
1641            .await?
1642            .map(|value| self.deserialize_value(&value))
1643            .transpose()?)
1644    }
1645
1646    async fn is_message_known(
1647        &self,
1648        message_hash: &matrix_sdk_crypto::olm::OlmMessageHash,
1649    ) -> Result<bool> {
1650        let value = rmp_serde::to_vec(message_hash)?;
1651        Ok(self.read().await?.has_olm_hash(value).await?)
1652    }
1653
1654    async fn get_outgoing_secret_requests(
1655        &self,
1656        request_id: &TransactionId,
1657    ) -> Result<Option<GossipRequest>> {
1658        let request_id = self.encode_key("key_requests", request_id.as_bytes());
1659        Ok(self
1660            .read()
1661            .await?
1662            .get_outgoing_secret_request(request_id)
1663            .await?
1664            .map(|(value, sent_out)| self.deserialize_key_request(&value, sent_out))
1665            .transpose()?)
1666    }
1667
1668    async fn get_secret_request_by_info(
1669        &self,
1670        key_info: &SecretInfo,
1671    ) -> Result<Option<GossipRequest>> {
1672        let key_info = self.encode_key("key_requests", key_info.as_key());
1673        Ok(self
1674            .read()
1675            .await?
1676            .get_secret_request_by_info(key_info)
1677            .await?
1678            .map(|(value, sent_out)| self.deserialize_key_request(&value, sent_out))
1679            .transpose()?)
1680    }
1681
1682    async fn get_unsent_secret_requests(&self) -> Result<Vec<GossipRequest>> {
1683        self.read()
1684            .await?
1685            .get_unsent_secret_requests()
1686            .await?
1687            .iter()
1688            .map(|value| {
1689                let request = self.deserialize_key_request(value, false)?;
1690                Ok(request)
1691            })
1692            .collect()
1693    }
1694
1695    async fn delete_outgoing_secret_requests(&self, request_id: &TransactionId) -> Result<()> {
1696        let request_id = self.encode_key("key_requests", request_id.as_bytes());
1697        Ok(self.write().await?.delete_key_request(request_id).await?)
1698    }
1699
1700    async fn get_secrets_from_inbox(
1701        &self,
1702        secret_name: &SecretName,
1703    ) -> Result<Vec<Zeroizing<String>>> {
1704        let secret_name = self.encode_key("secrets_inbox", secret_name.to_string());
1705
1706        self.read()
1707            .await?
1708            .get_secrets_from_inbox(secret_name)
1709            .await?
1710            .into_iter()
1711            .map(|value| self.deserialize_json(value.as_ref()).map(|value: String| value.into()))
1712            .collect()
1713    }
1714
1715    async fn delete_secrets_from_inbox(&self, secret_name: &SecretName) -> Result<()> {
1716        let secret_name = self.encode_key("secrets_inbox", secret_name.to_string());
1717        self.write().await?.delete_secrets_from_inbox(secret_name).await
1718    }
1719
1720    async fn get_withheld_info(
1721        &self,
1722        room_id: &RoomId,
1723        session_id: &str,
1724    ) -> Result<Option<RoomKeyWithheldEntry>> {
1725        let room_id = self.encode_key("direct_withheld_info", room_id);
1726        let session_id = self.encode_key("direct_withheld_info", session_id);
1727
1728        self.read()
1729            .await?
1730            .get_direct_withheld_info(session_id, room_id)
1731            .await?
1732            .map(|value| {
1733                let info = self.deserialize_json::<RoomKeyWithheldEntry>(&value)?;
1734                Ok(info)
1735            })
1736            .transpose()
1737    }
1738
1739    async fn get_withheld_sessions_by_room_id(
1740        &self,
1741        room_id: &RoomId,
1742    ) -> matrix_sdk_crypto::store::Result<Vec<RoomKeyWithheldEntry>, Self::Error> {
1743        let room_id = self.encode_key("direct_withheld_info", room_id);
1744
1745        self.read()
1746            .await?
1747            .get_withheld_sessions_by_room_id(room_id)
1748            .await?
1749            .into_iter()
1750            .map(|value| self.deserialize_json(&value))
1751            .collect()
1752    }
1753
1754    async fn get_room_settings(&self, room_id: &RoomId) -> Result<Option<RoomSettings>> {
1755        let room_id = self.encode_key("room_settings", room_id.as_bytes());
1756        let Some(value) = self.read().await?.get_room_settings(room_id).await? else {
1757            return Ok(None);
1758        };
1759
1760        let settings = self.deserialize_value(&value)?;
1761
1762        return Ok(Some(settings));
1763    }
1764
1765    async fn get_received_room_key_bundle_data(
1766        &self,
1767        room_id: &RoomId,
1768        user_id: &UserId,
1769    ) -> Result<Option<StoredRoomKeyBundleData>> {
1770        let room_id = self.encode_key("received_room_key_bundle", room_id);
1771        let user_id = self.encode_key("received_room_key_bundle", user_id);
1772        self.read()
1773            .await?
1774            .get_received_room_key_bundle(room_id, user_id)
1775            .await?
1776            .map(|value| self.deserialize_value(&value))
1777            .transpose()
1778    }
1779
1780    async fn has_downloaded_all_room_keys(&self, room_id: &RoomId) -> Result<bool> {
1781        let room_id = self.encode_key("room_key_backups_fully_downloaded", room_id);
1782        self.read().await?.has_downloaded_all_room_keys(room_id).await
1783    }
1784
1785    async fn get_pending_key_bundle_details_for_room(
1786        &self,
1787        room_id: &RoomId,
1788    ) -> Result<Option<RoomPendingKeyBundleDetails>> {
1789        let room_id = self.encode_key("rooms_pending_key_bundle", room_id.as_bytes());
1790        let Some(value) = self.read().await?.get_room_pending_key_bundle(room_id).await? else {
1791            return Ok(None);
1792        };
1793
1794        let details = self.deserialize_value(&value)?;
1795        Ok(Some(details))
1796    }
1797
1798    async fn get_all_rooms_pending_key_bundles(&self) -> Result<Vec<RoomPendingKeyBundleDetails>> {
1799        let details = self.read().await?.get_all_rooms_pending_key_bundle().await?;
1800        let room_ids = details
1801            .into_iter()
1802            .map(|value| self.deserialize_value(&value))
1803            .collect::<Result<_, _>>()?;
1804        Ok(room_ids)
1805    }
1806
1807    async fn get_custom_value(&self, key: &str) -> Result<Option<Vec<u8>>> {
1808        let Some(serialized) = self.read().await?.get_kv(key).await? else {
1809            return Ok(None);
1810        };
1811        let value = if let Some(cipher) = &self.store_cipher {
1812            let encrypted = rmp_serde::from_slice(&serialized)?;
1813            cipher.decrypt_value_data(encrypted)?
1814        } else {
1815            serialized
1816        };
1817
1818        Ok(Some(value))
1819    }
1820
1821    async fn set_custom_value(&self, key: &str, value: Vec<u8>) -> Result<()> {
1822        let serialized = if let Some(cipher) = &self.store_cipher {
1823            let encrypted = cipher.encrypt_value_data(value)?;
1824            rmp_serde::to_vec_named(&encrypted)?
1825        } else {
1826            value
1827        };
1828
1829        self.write().await?.set_kv(key, serialized).await?;
1830        Ok(())
1831    }
1832
1833    async fn remove_custom_value(&self, key: &str) -> Result<()> {
1834        let key = key.to_owned();
1835        self.write()
1836            .await?
1837            .interact(move |conn| conn.execute("DELETE FROM kv WHERE key = ?1", (&key,)))
1838            .await
1839            .unwrap()?;
1840        Ok(())
1841    }
1842
1843    #[instrument(skip(self))]
1844    async fn try_take_leased_lock(
1845        &self,
1846        lease_duration_ms: u32,
1847        key: &str,
1848        holder: &str,
1849    ) -> Result<Option<CrossProcessLockGeneration>> {
1850        let key = key.to_owned();
1851        let holder = holder.to_owned();
1852
1853        let now: u64 = MilliSecondsSinceUnixEpoch::now().get().into();
1854        let expiration = now + lease_duration_ms as u64;
1855
1856        // Learn about the `excluded` keyword in https://sqlite.org/lang_upsert.html.
1857        let generation = self
1858            .write()
1859            .await?
1860            .with_transaction(move |txn| {
1861                txn.query_one(
1862                    "INSERT INTO lease_locks (key, holder, expiration)
1863                    VALUES (?1, ?2, ?3)
1864                    ON CONFLICT (key)
1865                    DO
1866                        UPDATE SET
1867                            holder = excluded.holder,
1868                            expiration = excluded.expiration,
1869                            generation =
1870                                CASE holder
1871                                    WHEN excluded.holder THEN generation
1872                                    ELSE generation + 1
1873                                END
1874                        WHERE
1875                            holder = excluded.holder
1876                            OR expiration < ?4
1877                    RETURNING generation
1878                    ",
1879                    (key, holder, expiration, now),
1880                    |row| row.get(0),
1881                )
1882                .optional()
1883            })
1884            .await?;
1885
1886        Ok(generation)
1887    }
1888
1889    async fn next_batch_token(&self) -> Result<Option<String>, Self::Error> {
1890        let conn = self.read().await?;
1891        if let Some(token) = conn.get_kv("next_batch_token").await? {
1892            let maybe_token: Option<String> = self.deserialize_value(&token)?;
1893            Ok(maybe_token)
1894        } else {
1895            Ok(None)
1896        }
1897    }
1898
1899    async fn close(&self) -> Result<()> {
1900        connection::close_connections(&self.connections, "Crypto store").await;
1901        Ok(())
1902    }
1903
1904    async fn reopen(&self) -> Result<()> {
1905        connection::reopen_connections(
1906            &self.connections,
1907            self.db_path.clone(),
1908            self.pool_config,
1909            self.runtime_config,
1910        )
1911        .await?;
1912        Ok(())
1913    }
1914
1915    async fn get_size(&self) -> Result<Option<usize>, Self::Error> {
1916        Ok(Some(self.read().await?.get_db_size().await?))
1917    }
1918}
1919
1920#[cfg(test)]
1921mod tests {
1922    use std::{path::Path, sync::LazyLock};
1923
1924    use matrix_sdk_common::deserialized_responses::WithheldCode;
1925    use matrix_sdk_crypto::{
1926        cryptostore_integration_tests, cryptostore_integration_tests_time, olm::SenderDataType,
1927        store::CryptoStore,
1928    };
1929    use matrix_sdk_test::async_test;
1930    use ruma::{device_id, room_id, user_id};
1931    use similar_asserts::assert_eq;
1932    use tempfile::{TempDir, tempdir};
1933    use tokio::fs;
1934
1935    use super::SqliteCryptoStore;
1936    use crate::SqliteStoreConfig;
1937
1938    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
1939
1940    struct TestDb {
1941        // Needs to be kept alive because the Drop implementation for TempDir
1942        // deletes the directory.
1943        _dir: TempDir,
1944        database: SqliteCryptoStore,
1945    }
1946
1947    fn copy_db(data_path: &str) -> TempDir {
1948        let db_name = super::DATABASE_NAME;
1949
1950        let manifest_path = Path::new(env!("CARGO_MANIFEST_DIR")).join("../..");
1951        let database_path = manifest_path.join(data_path).join(db_name);
1952
1953        let tmpdir = tempdir().unwrap();
1954        let destination = tmpdir.path().join(db_name);
1955
1956        // Copy the test database to the tempdir so our test runs are
1957        // idempotent.
1958        std::fs::copy(&database_path, destination).unwrap();
1959
1960        tmpdir
1961    }
1962
1963    async fn get_test_db(data_path: &str, passphrase: Option<&str>) -> TestDb {
1964        let tmpdir = copy_db(data_path);
1965
1966        let database = SqliteCryptoStore::open(tmpdir.path(), passphrase)
1967            .await
1968            .expect("Can't open the test store");
1969
1970        TestDb { _dir: tmpdir, database }
1971    }
1972
1973    #[async_test]
1974    async fn test_pool_size() {
1975        let store_open_config =
1976            SqliteStoreConfig::new(TMP_DIR.path().join("test_pool_size")).pool_max_size(42);
1977
1978        let store = SqliteCryptoStore::open_with_config(&store_open_config).await.unwrap();
1979
1980        let guard = store.connections.lock().await;
1981        let conns = guard.as_ref().unwrap();
1982        assert_eq!(conns.pool.status().max_size, 42);
1983    }
1984
1985    /// Test that we didn't regress in our storage layer by loading data from a
1986    /// pre-filled database, or in other words use a test vector for this.
1987    #[async_test]
1988    async fn test_open_test_vector_store() {
1989        let TestDb { _dir: _, database } = get_test_db("testing/data/storage", None).await;
1990
1991        let account = database
1992            .load_account()
1993            .await
1994            .unwrap()
1995            .expect("The test database is prefilled with data, we should find an account");
1996
1997        let user_id = account.user_id();
1998        let device_id = account.device_id();
1999
2000        assert_eq!(
2001            user_id.as_str(),
2002            "@pjtest:synapse-oidc.element.dev",
2003            "The user ID should match to the one we expect."
2004        );
2005
2006        assert_eq!(
2007            device_id.as_str(),
2008            "v4TqgcuIH6",
2009            "The device ID should match to the one we expect."
2010        );
2011
2012        let device = database
2013            .get_device(user_id, device_id)
2014            .await
2015            .unwrap()
2016            .expect("Our own device should be found in the store.");
2017
2018        assert_eq!(device.device_id(), device_id);
2019        assert_eq!(device.user_id(), user_id);
2020
2021        assert_eq!(
2022            device.ed25519_key().expect("The device should have a Ed25519 key.").to_base64(),
2023            "+cxl1Gl3du5i7UJwfWnoRDdnafFF+xYdAiTYYhYLr8s"
2024        );
2025
2026        assert_eq!(
2027            device.curve25519_key().expect("The device should have a Curve25519 key.").to_base64(),
2028            "4SL9eEUlpyWSUvjljC5oMjknHQQJY7WZKo5S1KL/5VU"
2029        );
2030
2031        let identity = database
2032            .get_user_identity(user_id)
2033            .await
2034            .unwrap()
2035            .expect("The store should contain an identity.");
2036
2037        assert_eq!(identity.user_id(), user_id);
2038
2039        let identity = identity
2040            .own()
2041            .expect("The identity should be of the correct type, it should be our own identity.");
2042
2043        let master_key = identity
2044            .master_key()
2045            .get_first_key()
2046            .expect("Our own identity should have a master key");
2047
2048        assert_eq!(master_key.to_base64(), "iCUEtB1RwANeqRa5epDrblLk4mer/36sylwQ5hYY3oE");
2049    }
2050
2051    /// Test that we didn't regress in our storage layer by loading data from a
2052    /// pre-filled database, or in other words use a test vector for this.
2053    #[async_test]
2054    async fn test_open_test_vector_encrypted_store() {
2055        let TestDb { _dir: _, database } = get_test_db(
2056            "testing/data/storage/alice",
2057            Some(concat!(
2058                "/rCia2fYAJ+twCZ1Xm2mxFCYcmJdyzkdJjwtgXsziWpYS/UeNxnixuSieuwZXm+x1VsJHmWpl",
2059                "H+QIQBZpEGZtC9/S/l8xK+WOCesmET0o6yJ/KP73ofDtjBlnNpPwuHLKFpyTbyicpCgQ4UT+5E",
2060                "UBuJ08TY9Ujdf1D13k5kr5tSZUefDKKCuG1fCRqlU8ByRas1PMQsZxT2W8t7QgBrQiiGmhpo/O",
2061                "Ti4hfx97GOxncKcxTzppiYQNoHs/f15+XXQD7/oiCcqRIuUlXNsU6hRpFGmbYx2Pi1eyQViQCt",
2062                "B5dAEiSD0N8U81wXYnpynuTPtnL+hfnOJIn7Sy7mkERQeKg"
2063            )),
2064        )
2065        .await;
2066
2067        let account = database
2068            .load_account()
2069            .await
2070            .unwrap()
2071            .expect("The test database is prefilled with data, we should find an account");
2072
2073        let user_id = account.user_id();
2074        let device_id = account.device_id();
2075
2076        assert_eq!(
2077            user_id.as_str(),
2078            "@alice:localhost",
2079            "The user ID should match to the one we expect."
2080        );
2081
2082        assert_eq!(
2083            device_id.as_str(),
2084            "JVVORTHFXY",
2085            "The device ID should match to the one we expect."
2086        );
2087
2088        let tracked_users =
2089            database.load_tracked_users().await.expect("Should be tracking some users");
2090
2091        assert_eq!(tracked_users.len(), 6);
2092
2093        let known_users = vec![
2094            user_id!("@alice:localhost"),
2095            user_id!("@dehydration3:localhost"),
2096            user_id!("@eve:localhost"),
2097            user_id!("@bob:localhost"),
2098            user_id!("@malo:localhost"),
2099            user_id!("@carl:localhost"),
2100        ];
2101
2102        // load the identities
2103        for user_id in known_users {
2104            database.get_user_identity(user_id).await.expect("Should load this identity").unwrap();
2105        }
2106
2107        let carl_identity =
2108            database.get_user_identity(user_id!("@carl:localhost")).await.unwrap().unwrap();
2109
2110        assert_eq!(
2111            carl_identity.master_key().get_first_key().unwrap().to_base64(),
2112            "CdhKYYDeBDQveOioXEGWhTPCyzc63Irpar3CNyfun2Q"
2113        );
2114        assert!(!carl_identity.was_previously_verified());
2115
2116        let bob_identity =
2117            database.get_user_identity(user_id!("@bob:localhost")).await.unwrap().unwrap();
2118
2119        assert_eq!(
2120            bob_identity.master_key().get_first_key().unwrap().to_base64(),
2121            "COh2GYOJWSjem5QPRCaGp9iWV83IELG1IzLKW2S3pFY"
2122        );
2123        // Bob is verified so this flag should be set
2124        assert!(bob_identity.was_previously_verified());
2125
2126        let known_devices = vec![
2127            (device_id!("OPXQHCZSKW"), user_id!("@alice:localhost")),
2128            // a dehydrated one
2129            (
2130                device_id!("EvW+9IrGR10KVgVeZP25/KaPfx4R86FofVMcaz7VOho"),
2131                user_id!("@alice:localhost"),
2132            ),
2133            (device_id!("HEEFRFQENV"), user_id!("@alice:localhost")),
2134            (device_id!("JVVORTHFXY"), user_id!("@alice:localhost")),
2135            (device_id!("NQUWWSKKHS"), user_id!("@alice:localhost")),
2136            (device_id!("ORBLPFYCPG"), user_id!("@alice:localhost")),
2137            (device_id!("YXOWENSEGM"), user_id!("@dehydration3:localhost")),
2138            (device_id!("VXLFMYCHXC"), user_id!("@bob:localhost")),
2139            (device_id!("FDGDQAEWOW"), user_id!("@bob:localhost")),
2140            (device_id!("VXLFMYCHXC"), user_id!("@bob:localhost")),
2141            (device_id!("FDGDQAEWOW"), user_id!("@bob:localhost")),
2142            (device_id!("QKUKWJTTQC"), user_id!("@malo:localhost")),
2143            (device_id!("LOUXJECTFG"), user_id!("@malo:localhost")),
2144            (device_id!("MKKMAEVLPB"), user_id!("@carl:localhost")),
2145        ];
2146
2147        for (device_id, user_id) in known_devices {
2148            database.get_device(user_id, device_id).await.expect("Should load the device").unwrap();
2149        }
2150
2151        let known_sender_key_to_session_count = vec![
2152            ("FfYcYfDF4nWy+LHdK6CEpIMlFAQDORc30WUkghL06kM", 1),
2153            ("EvW+9IrGR10KVgVeZP25/KaPfx4R86FofVMcaz7VOho", 1),
2154            ("hAGsoA4a9M6wwEUX5Q1jux1i+tUngLi01n5AmhDoHTY", 1),
2155            ("aKqtSJymLzuoglWFwPGk1r/Vm2LE2hFESzXxn4RNjRM", 0),
2156            ("zHK1psCrgeMn0kaz8hcdvA3INyar9jg1yfrSp0p1pHo", 1),
2157            ("1QmBA316Wj5jIFRwNOti6N6Xh/vW0bsYCcR4uPfy8VQ", 1),
2158            ("g5ef2vZF3VXgSPyODIeXpyHIRkuthvLhGvd6uwYggWU", 1),
2159            ("o7hfupPd1VsNkRIvdlH6ujrEJFSKjFCGbxhAd31XxjI", 1),
2160            ("Z3RxKQLxY7xpP+ZdOGR2SiNE37SrvmRhW7GPu1UGdm8", 1),
2161            ("GDomaav8NiY3J+dNEeApJm+O0FooJ3IpVaIyJzCN4w4", 1),
2162            ("7m7fqkHyEr47V5s/KjaxtJMOr3pSHrrns2q2lWpAQi8", 0),
2163            ("9psAkPUIF8vNbWbnviX3PlwRcaeO53EHJdNtKpTY1X0", 0),
2164            ("mqanh+ztw5oRtpqYQgLGW864i6NY2zpoKMIlrcyC+Aw", 0),
2165            ("fJU/TJdbsv7tVbbpHw1Ke73ziElnM32cNhP2WIg4T10", 0),
2166            ("sUIeFeFcCZoa5IC6nJ6Vrbvztcyx09m8BBg57XKRClg", 1),
2167        ];
2168
2169        for (id, count) in known_sender_key_to_session_count {
2170            let olm_sessions =
2171                database.get_sessions(id).await.expect("Should have some olm sessions");
2172
2173            println!("### Session id: {id:?}");
2174            assert_eq!(olm_sessions.map_or(0, |v| v.len()), count);
2175        }
2176
2177        let inbound_group_sessions = database.get_inbound_group_sessions().await.unwrap();
2178        assert_eq!(inbound_group_sessions.len(), 15);
2179        let known_inbound_group_sessions = vec![
2180            (
2181                "5hNAxrLai3VI0LKBwfh3wLfksfBFWds0W1a5X5/vSXA",
2182                room_id!("!SRstFdydzrGwJYtVfm:localhost"),
2183            ),
2184            (
2185                "M6d2eU3y54gaYTbvGSlqa/xc1Az35l56Cp9sxzHWO4g",
2186                room_id!("!SRstFdydzrGwJYtVfm:localhost"),
2187            ),
2188            (
2189                "IrydwXkRk2N2AqUMIVmLL3oJgMq14R9KId0P/uSD100",
2190                room_id!("!SRstFdydzrGwJYtVfm:localhost"),
2191            ),
2192            (
2193                "Y74+l9jTo7N5UF+GQwdpgJGe4sn1+QtWITq7BxulHIE",
2194                room_id!("!SRstFdydzrGwJYtVfm:localhost"),
2195            ),
2196            (
2197                "HpJxQR57WbQGdY6w2Q+C16znVvbXGa+JvQdRoMpWbXg",
2198                room_id!("!SRstFdydzrGwJYtVfm:localhost"),
2199            ),
2200            (
2201                "Xetvi+ydFkZt8dpONGFbEusQb/Chc2V0XlLByZhsbgE",
2202                room_id!("!ZIwZcFqZVAYLAqVjfV:localhost"),
2203            ),
2204            (
2205                "wv/WN/39akyerIXczTaIpjAuLnwgXKRtbXFSEHiJqxo",
2206                room_id!("!ZIwZcFqZVAYLAqVjfV:localhost"),
2207            ),
2208            (
2209                "nA4gQwL//Cm8OdlyjABl/jChbPT/cP5V4Sd8iuE6H0s",
2210                room_id!("!ZIwZcFqZVAYLAqVjfV:localhost"),
2211            ),
2212            (
2213                "bAAgqFeRDTjfEqL6Qf/c9mk55zoNDCSlboAIRd6b0hw",
2214                room_id!("!ZIwZcFqZVAYLAqVjfV:localhost"),
2215            ),
2216            (
2217                "exPbsMMdGfAG2qmDdFtpAn+koVprfzS0Zip/RA9QRCE",
2218                room_id!("!ZIwZcFqZVAYLAqVjfV:localhost"),
2219            ),
2220            (
2221                "h+om7oSw/ZV94fcKaoe8FGXJwQXWOfKQfzbGgNWQILI",
2222                room_id!("!ZIwZcFqZVAYLAqVjfV:localhost"),
2223            ),
2224            (
2225                "ul3VXonpgk4lO2L3fEWubP/nxsTmLHqu5v8ZM9vHEcw",
2226                room_id!("!ZIwZcFqZVAYLAqVjfV:localhost"),
2227            ),
2228            (
2229                "JXY15UxC3az2mwg8uX4qwgxfvCM4aygiIWMcdNiVQoc",
2230                room_id!("!ZIwZcFqZVAYLAqVjfV:localhost"),
2231            ),
2232            (
2233                "OGB9lObr9kWUvha9tB5sMfOF/Mztk24JwQz/nwg3iFQ",
2234                room_id!("!OgRiTRMaUzLdpCeDBM:localhost"),
2235            ),
2236            (
2237                "SFkHcbxjUOYF7mUAYI/oEMDZFaXszQbCN6Jza7iemj0",
2238                room_id!("!OgRiTRMaUzLdpCeDBM:localhost"),
2239            ),
2240        ];
2241
2242        // ensure we can load them all
2243        for (session_id, room_id) in &known_inbound_group_sessions {
2244            database
2245                .get_inbound_group_session(room_id, session_id)
2246                .await
2247                .expect("Should be able to load inbound group session")
2248                .unwrap();
2249        }
2250
2251        let bob_sender_verified = database
2252            .get_inbound_group_session(
2253                room_id!("!ZIwZcFqZVAYLAqVjfV:localhost"),
2254                "exPbsMMdGfAG2qmDdFtpAn+koVprfzS0Zip/RA9QRCE",
2255            )
2256            .await
2257            .unwrap()
2258            .unwrap();
2259
2260        assert_eq!(bob_sender_verified.sender_data.to_type(), SenderDataType::SenderVerified);
2261        assert!(bob_sender_verified.backed_up());
2262        assert!(!bob_sender_verified.has_been_imported());
2263
2264        let alice_unknown_device = database
2265            .get_inbound_group_session(
2266                room_id!("!SRstFdydzrGwJYtVfm:localhost"),
2267                "IrydwXkRk2N2AqUMIVmLL3oJgMq14R9KId0P/uSD100",
2268            )
2269            .await
2270            .unwrap()
2271            .unwrap();
2272
2273        assert_eq!(alice_unknown_device.sender_data.to_type(), SenderDataType::UnknownDevice);
2274        assert!(alice_unknown_device.backed_up());
2275        assert!(alice_unknown_device.has_been_imported());
2276
2277        let carl_tofu_session = database
2278            .get_inbound_group_session(
2279                room_id!("!OgRiTRMaUzLdpCeDBM:localhost"),
2280                "OGB9lObr9kWUvha9tB5sMfOF/Mztk24JwQz/nwg3iFQ",
2281            )
2282            .await
2283            .unwrap()
2284            .unwrap();
2285
2286        assert_eq!(carl_tofu_session.sender_data.to_type(), SenderDataType::SenderUnverified);
2287        assert!(carl_tofu_session.backed_up());
2288        assert!(!carl_tofu_session.has_been_imported());
2289
2290        // Load outbound sessions
2291        database
2292            .get_outbound_group_session(room_id!("!OgRiTRMaUzLdpCeDBM:localhost"))
2293            .await
2294            .unwrap()
2295            .unwrap();
2296        database
2297            .get_outbound_group_session(room_id!("!ZIwZcFqZVAYLAqVjfV:localhost"))
2298            .await
2299            .unwrap()
2300            .unwrap();
2301        database
2302            .get_outbound_group_session(room_id!("!SRstFdydzrGwJYtVfm:localhost"))
2303            .await
2304            .unwrap()
2305            .unwrap();
2306
2307        let withheld_info = database
2308            .get_withheld_info(
2309                room_id!("!OgRiTRMaUzLdpCeDBM:localhost"),
2310                "SASgZ+EklvAF4QxJclMlDRlmL0fAMjAJJIKFMdb4Ht0",
2311            )
2312            .await
2313            .expect("This session should be withheld")
2314            .unwrap();
2315
2316        assert_eq!(withheld_info.content.withheld_code(), WithheldCode::Unverified);
2317
2318        let backup_keys = database.load_backup_keys().await.expect("backup key should be cached");
2319        assert_eq!(backup_keys.backup_version.unwrap(), "6");
2320        assert!(backup_keys.decryption_key.is_some());
2321    }
2322
2323    /// Test that we migrate the secrets inbox properly.
2324    ///
2325    /// The format for the secrets inbox changed in version 17. Previously, the
2326    /// secrets inbox stored a full `GossippedSecrets` struct. In version 17,
2327    /// the secrets inbox now stores only the secret.
2328    #[async_test]
2329    async fn test_secrets_inbox_migration() {
2330        use std::ops::Deref;
2331
2332        use matrix_sdk_crypto::{
2333            GossipRequest, GossippedSecret, SecretInfo,
2334            types::events::{
2335                olm_v1::{DecryptedSecretSendEvent, OlmV1Keys},
2336                secret_send::SecretSendContent,
2337            },
2338            vodozemac::Ed25519SecretKey,
2339        };
2340        use ruma::{TransactionId, events::secret::request::SecretName, owned_user_id};
2341
2342        use crate::utils::{EncryptableStore, SqliteAsyncConnExt};
2343
2344        // Create a database with version 16
2345        let tmpdir = tempdir().unwrap();
2346        let config = SqliteStoreConfig::new(tmpdir.path());
2347        let pool = config.build_pool_of_connections(super::DATABASE_NAME).unwrap();
2348        let conn = pool.get().await.unwrap();
2349        let version = super::initialize_store(&conn, 0).await.unwrap();
2350        let old_data_store = SqliteCryptoStore::create_raw(
2351            config.secret.clone(),
2352            pool,
2353            conn,
2354            config.pool_config(),
2355            config.runtime_config(),
2356        )
2357        .await
2358        .unwrap();
2359        super::run_migrations(&old_data_store, version, Some(16)).await.unwrap();
2360        old_data_store.write().await.unwrap().wal_checkpoint().await;
2361
2362        // Store a secret using the old format
2363        let secret = GossippedSecret {
2364            secret_name: SecretName::CrossSigningMasterKey,
2365            gossip_request: GossipRequest {
2366                request_recipient: owned_user_id!("@alice:example.com"),
2367                request_id: TransactionId::new(),
2368                info: SecretInfo::SecretRequest(SecretName::CrossSigningMasterKey),
2369                sent_out: true,
2370            },
2371            event: DecryptedSecretSendEvent {
2372                sender: owned_user_id!("@alice:example.com"),
2373                recipient: owned_user_id!("@alice:example.com"),
2374                keys: OlmV1Keys { ed25519: Ed25519SecretKey::new().public_key() },
2375                recipient_keys: OlmV1Keys { ed25519: Ed25519SecretKey::new().public_key() },
2376                sender_device_keys: None,
2377                content: SecretSendContent::new(
2378                    "abc".into(),
2379                    "It is a secret to everybody".to_owned(),
2380                ),
2381            },
2382        };
2383        let value = old_data_store.serialize_json(&secret).unwrap();
2384        old_data_store
2385            .write()
2386            .await
2387            .unwrap()
2388            .prepare("INSERT INTO secrets (secret_name, data) VALUES (?1, ?2)", |mut stmt| {
2389                stmt.execute((SecretName::CrossSigningMasterKey.to_string(), value))
2390            })
2391            .await
2392            .unwrap();
2393
2394        // After we open the store, the data will be migrated
2395        let store = SqliteCryptoStore::open_with_config(&config).await.unwrap();
2396
2397        // and we should be able to read the secrets from the inbox
2398        let secrets =
2399            store.get_secrets_from_inbox(&SecretName::CrossSigningMasterKey).await.unwrap();
2400        assert_eq!(secrets.len(), 1);
2401        assert_eq!(secrets[0].deref(), "It is a secret to everybody");
2402    }
2403
2404    /// Test that we migrate the key requests table properly.
2405    ///
2406    /// Version 18 added a new column, with a unique index on the column,
2407    /// meaning that it can now only store one request per requested secret/key.
2408    /// Test that when we migrate from an older version that has multiple
2409    /// requests for the same secret, it only keeps one.
2410    #[async_test]
2411    async fn test_key_requests_migration() {
2412        use matrix_sdk_crypto::{GossipRequest, SecretInfo};
2413        use ruma::{TransactionId, events::secret::request::SecretName, owned_user_id};
2414
2415        use crate::utils::{EncryptableStore, SqliteAsyncConnExt};
2416
2417        // Create a database with version 16
2418        let tmpdir = tempdir().unwrap();
2419        let config = SqliteStoreConfig::new(tmpdir.path());
2420        let pool = config.build_pool_of_connections(super::DATABASE_NAME).unwrap();
2421        let conn = pool.get().await.unwrap();
2422        let version = super::initialize_store(&conn, 0).await.unwrap();
2423        let old_data_store = SqliteCryptoStore::create_raw(
2424            config.secret.clone(),
2425            pool,
2426            conn,
2427            config.pool_config(),
2428            config.runtime_config(),
2429        )
2430        .await
2431        .unwrap();
2432        super::run_migrations(&old_data_store, version, Some(16)).await.unwrap();
2433        old_data_store.write().await.unwrap().wal_checkpoint().await;
2434
2435        // Store a secret using the old format
2436        let recovery_request1 = GossipRequest {
2437            request_recipient: owned_user_id!("@alice:example.com"),
2438            request_id: TransactionId::new(),
2439            info: SecretInfo::SecretRequest(SecretName::RecoveryKey),
2440            sent_out: true,
2441        };
2442        let serialized_recovery_request1 =
2443            old_data_store.serialize_value(&recovery_request1).unwrap();
2444        let recovery_request2 = GossipRequest {
2445            request_recipient: owned_user_id!("@alice:example.com"),
2446            request_id: TransactionId::new(),
2447            info: SecretInfo::SecretRequest(SecretName::RecoveryKey),
2448            sent_out: true,
2449        };
2450        let serialized_recovery_request2 =
2451            old_data_store.serialize_value(&recovery_request2).unwrap();
2452        let msk_request = GossipRequest {
2453            request_recipient: owned_user_id!("@alice:example.com"),
2454            request_id: TransactionId::new(),
2455            info: SecretInfo::SecretRequest(SecretName::CrossSigningMasterKey),
2456            sent_out: true,
2457        };
2458        let serialized_msk_request = old_data_store.serialize_value(&msk_request).unwrap();
2459        let recovery_request1_clone = recovery_request1.clone();
2460        let recovery_request2_clone = recovery_request2.clone();
2461        let msk_request_clone = msk_request.clone();
2462        old_data_store
2463            .write()
2464            .await
2465            .unwrap()
2466            .prepare(
2467                "INSERT INTO key_requests (request_id, sent_out, data) VALUES (?1, ?2, ?3)",
2468                move |mut stmt| {
2469                    stmt.execute((
2470                        old_data_store.encode_key(
2471                            "key_requests",
2472                            recovery_request1_clone.request_id.as_bytes(),
2473                        ),
2474                        recovery_request1_clone.sent_out,
2475                        serialized_recovery_request1,
2476                    ))?;
2477                    stmt.execute((
2478                        old_data_store.encode_key(
2479                            "key_requests",
2480                            recovery_request2_clone.request_id.as_bytes(),
2481                        ),
2482                        recovery_request2_clone.sent_out,
2483                        serialized_recovery_request2,
2484                    ))?;
2485                    stmt.execute((
2486                        old_data_store
2487                            .encode_key("key_requests", msk_request_clone.request_id.as_bytes()),
2488                        msk_request_clone.sent_out,
2489                        serialized_msk_request,
2490                    ))
2491                },
2492            )
2493            .await
2494            .unwrap();
2495
2496        // After we open the store, the data will be migrated
2497        let store = SqliteCryptoStore::open_with_config(&config).await.unwrap();
2498
2499        // and we should be able to read one request for the recovery key and
2500        // one request for the MSK
2501        if let Some(GossipRequest {
2502            request_id,
2503            info: SecretInfo::SecretRequest(SecretName::RecoveryKey),
2504            ..
2505        }) = store
2506            .get_secret_request_by_info(&SecretInfo::SecretRequest(SecretName::RecoveryKey))
2507            .await
2508            .unwrap()
2509        {
2510            if request_id == recovery_request1.request_id {
2511                assert!(
2512                    store
2513                        .get_outgoing_secret_requests(&recovery_request2.request_id)
2514                        .await
2515                        .unwrap()
2516                        .is_none()
2517                );
2518            } else if request_id == recovery_request2.request_id {
2519                assert!(
2520                    store
2521                        .get_outgoing_secret_requests(&recovery_request1.request_id)
2522                        .await
2523                        .unwrap()
2524                        .is_none()
2525                );
2526            } else {
2527                panic!("unexpected record found");
2528            }
2529        } else {
2530            panic!("expected to get a secret request");
2531        }
2532        if let Some(GossipRequest {
2533            request_id,
2534            info: SecretInfo::SecretRequest(SecretName::CrossSigningMasterKey),
2535            ..
2536        }) = store
2537            .get_secret_request_by_info(&SecretInfo::SecretRequest(
2538                SecretName::CrossSigningMasterKey,
2539            ))
2540            .await
2541            .unwrap()
2542        {
2543            assert_eq!(request_id, msk_request.request_id);
2544        } else {
2545            panic!("expected to get a secret request");
2546        }
2547    }
2548
2549    async fn get_store(
2550        name: &str,
2551        passphrase: Option<&str>,
2552        clear_data: bool,
2553    ) -> SqliteCryptoStore {
2554        let tmpdir_path = TMP_DIR.path().join(name);
2555
2556        if clear_data {
2557            let _ = fs::remove_dir_all(&tmpdir_path).await;
2558        }
2559
2560        SqliteCryptoStore::open(tmpdir_path.to_str().unwrap(), passphrase)
2561            .await
2562            .expect("Can't create a secret protected store")
2563    }
2564
2565    cryptostore_integration_tests!();
2566    cryptostore_integration_tests_time!();
2567}
2568
2569#[cfg(test)]
2570mod encrypted_tests {
2571    use std::sync::LazyLock;
2572
2573    use matrix_sdk_crypto::{cryptostore_integration_tests, cryptostore_integration_tests_time};
2574    use tempfile::{TempDir, tempdir};
2575    use tokio::fs;
2576
2577    use super::SqliteCryptoStore;
2578
2579    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2580
2581    async fn get_store(
2582        name: &str,
2583        passphrase: Option<&str>,
2584        clear_data: bool,
2585    ) -> SqliteCryptoStore {
2586        let tmpdir_path = TMP_DIR.path().join(name);
2587        let pass = passphrase.unwrap_or("default_test_password");
2588
2589        if clear_data {
2590            let _ = fs::remove_dir_all(&tmpdir_path).await;
2591        }
2592
2593        SqliteCryptoStore::open(tmpdir_path.to_str().unwrap(), Some(pass))
2594            .await
2595            .expect("Can't create a secret protected store")
2596    }
2597
2598    cryptostore_integration_tests!();
2599    cryptostore_integration_tests_time!();
2600}
2601
2602#[cfg(test)]
2603mod close_reopen_tests {
2604    use std::sync::LazyLock;
2605
2606    use matrix_sdk_crypto::store::CryptoStore;
2607    use matrix_sdk_test::async_test;
2608    use tempfile::{TempDir, tempdir};
2609
2610    use super::SqliteCryptoStore;
2611
2612    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2613
2614    async fn new_store(name: &str) -> SqliteCryptoStore {
2615        let tmpdir_path = TMP_DIR.path().join(name);
2616        SqliteCryptoStore::open(tmpdir_path, None).await.unwrap()
2617    }
2618
2619    #[async_test]
2620    async fn test_close_completes_without_timeout() {
2621        let store = new_store("close_no_timeout").await;
2622
2623        // Close should complete quickly without hitting the 5s timeout.
2624        let start = std::time::Instant::now();
2625        store.close().await.unwrap();
2626        let elapsed = start.elapsed();
2627
2628        assert!(
2629            elapsed < std::time::Duration::from_secs(2),
2630            "close() took {elapsed:?}, expected < 2s (no timeout)"
2631        );
2632
2633        // Connections should be None after close.
2634        let guard = store.connections.lock().await;
2635        assert!(guard.is_none(), "connections should be None after close");
2636    }
2637
2638    #[async_test]
2639    async fn test_reopen_restores_connections() {
2640        let store = new_store("reopen_restores").await;
2641
2642        store.close().await.unwrap();
2643
2644        {
2645            let guard = store.connections.lock().await;
2646            assert!(guard.is_none());
2647        }
2648
2649        store.reopen().await.unwrap();
2650
2651        {
2652            let guard = store.connections.lock().await;
2653            assert!(guard.is_some(), "connections should be Some after reopen");
2654        }
2655    }
2656
2657    #[async_test]
2658    async fn test_close_is_idempotent() {
2659        let store = new_store("close_idempotent").await;
2660
2661        store.close().await.unwrap();
2662        // Second close should be a no-op.
2663        store.close().await.unwrap();
2664
2665        let guard = store.connections.lock().await;
2666        assert!(guard.is_none());
2667    }
2668
2669    #[async_test]
2670    async fn test_reopen_is_idempotent() {
2671        let store = new_store("reopen_idempotent").await;
2672
2673        // Reopen on an active store should be a no-op.
2674        store.reopen().await.unwrap();
2675
2676        let guard = store.connections.lock().await;
2677        assert!(guard.is_some());
2678    }
2679
2680    #[async_test]
2681    async fn test_read_fails_when_closed() {
2682        let store = new_store("read_fails_closed").await;
2683        store.close().await.unwrap();
2684
2685        let err = store.load_account().await;
2686        assert!(err.is_err(), "read should fail when closed");
2687
2688        let err_msg = err.unwrap_err().to_string();
2689        assert!(err_msg.contains("closed"), "error should mention 'closed', got: {err_msg}");
2690    }
2691
2692    #[async_test]
2693    async fn test_operations_work_after_reopen() {
2694        let store = new_store("ops_after_reopen").await;
2695
2696        store.close().await.unwrap();
2697        store.reopen().await.unwrap();
2698
2699        // A read operation should work immediately after reopen.
2700        let account = store.load_account().await;
2701        assert!(account.is_ok(), "load_account should succeed after reopen");
2702        // No account was saved, so this should be None.
2703        assert!(account.unwrap().is_none());
2704    }
2705
2706    #[async_test]
2707    async fn test_multiple_close_reopen_cycles() {
2708        let store = new_store("multi_cycles").await;
2709
2710        for _ in 0..5 {
2711            store.close().await.unwrap();
2712            store.reopen().await.unwrap();
2713
2714            // After each cycle, the store should be fully operational.
2715            let account = store.load_account().await;
2716            assert!(account.is_ok(), "store should work after close/reopen cycle");
2717        }
2718    }
2719
2720    #[async_test]
2721    async fn test_pool_is_fully_drained_after_close() {
2722        let store = new_store("pool_drained").await;
2723
2724        // Do a few reads to exercise the pool.
2725        let _ = store.load_account().await;
2726        let _ = store.load_account().await;
2727
2728        store.close().await.unwrap();
2729
2730        // After close, the connections field should be None.
2731        let guard = store.connections.lock().await;
2732        assert!(guard.is_none(), "all connections should be released after close");
2733    }
2734
2735    #[async_test]
2736    async fn test_close_waits_for_held_read_connection_to_drain() {
2737        let store = new_store("held_read_drain").await;
2738
2739        // Acquire a read connection and hold it, simulating an in-flight read.
2740        let held_conn = store.read().await.unwrap();
2741
2742        // Spawn close in a background task — it will close the pool and then
2743        // poll-wait for pool.status().size == 0 in the drain loop.
2744        let store_clone = store.clone();
2745        let close_handle = tokio::spawn(async move {
2746            store_clone.close().await.unwrap();
2747        });
2748
2749        // Give close() a moment to close the pool and enter the drain loop.
2750        tokio::time::sleep(std::time::Duration::from_millis(200)).await;
2751
2752        // The close task should still be running because we hold a connection.
2753        assert!(!close_handle.is_finished(), "close should be waiting for the held connection");
2754
2755        // Release the held connection — this lets pool.status().size drop to 0.
2756        drop(held_conn);
2757
2758        // Now close should complete promptly (well within the 5s timeout).
2759        let timeout = tokio::time::timeout(std::time::Duration::from_secs(3), close_handle).await;
2760        assert!(timeout.is_ok(), "close should complete after the held connection is released");
2761        timeout.unwrap().unwrap();
2762
2763        // Verify the store is fully closed.
2764        let guard = store.connections.lock().await;
2765        assert!(guard.is_none(), "connections should be None after close");
2766    }
2767}