Skip to main content

matrix_sdk_sqlite/
state_store.rs

1use std::{
2    borrow::Cow,
3    collections::{BTreeMap, BTreeSet, HashMap},
4    fmt, iter,
5    path::{Path, PathBuf},
6    str::FromStr as _,
7    sync::Arc,
8};
9
10use async_trait::async_trait;
11use deadpool::managed::PoolConfig;
12use matrix_sdk_base::{
13    MinimalRoomMemberEvent, ROOM_VERSION_FALLBACK, ROOM_VERSION_RULES_FALLBACK, RoomInfo,
14    RoomMemberships, RoomState, StateChanges, StateStore, StateStoreDataKey, StateStoreDataValue,
15    deserialized_responses::{DisplayName, RawAnySyncOrStrippedState, SyncOrStrippedState},
16    store::{
17        ChildTransactionId, DependentQueuedRequest, DependentQueuedRequestKind, QueueWedgeError,
18        QueuedRequest, QueuedRequestKind, RoomLoadSettings, SentRequestKey,
19        StoredThreadSubscription, ThreadSubscriptionStatus, migration_helpers::RoomInfoV1,
20    },
21};
22use matrix_sdk_store_encryption::StoreCipher;
23use ruma::{
24    CanonicalJsonObject, EventId, MilliSecondsSinceUnixEpoch, OwnedEventId, OwnedRoomId,
25    OwnedTransactionId, OwnedUserId, RoomId, TransactionId, UInt, UserId,
26    canonical_json::{RedactedBecause, redact},
27    events::{
28        AnyGlobalAccountDataEvent, AnyRoomAccountDataEvent, AnySyncStateEvent,
29        GlobalAccountDataEventType, RoomAccountDataEventType, StateEventType,
30        presence::PresenceEvent,
31        receipt::{Receipt, ReceiptThread, ReceiptType},
32        room::{
33            create::RoomCreateEventContent,
34            member::{StrippedRoomMemberEvent, SyncRoomMemberEvent},
35        },
36    },
37    profile::{UserProfile, UserProfileUpdate},
38    serde::Raw,
39};
40use rusqlite::{OptionalExtension, Transaction};
41use serde::{Deserialize, Serialize};
42use tokio::{
43    fs,
44    sync::{Mutex, OwnedMutexGuard},
45};
46use tracing::{debug, instrument, warn};
47
48use crate::{
49    OpenStoreError, RuntimeConfig, Secret, SqliteStoreConfig,
50    connection::{self, Connection as SqliteAsyncConn, Pool as SqlitePool, SqliteConnections},
51    error::{Error, Result},
52    utils::{
53        EncryptableStore, Key, SqliteAsyncConnExt, SqliteKeyValueStoreAsyncConnExt,
54        SqliteKeyValueStoreConnExt,
55    },
56};
57
58mod keys {
59    // Tables
60    pub const KV_BLOB: &str = "kv_blob";
61    pub const ROOM_INFO: &str = "room_info";
62    pub const STATE_EVENT: &str = "state_event";
63    pub const GLOBAL_ACCOUNT_DATA: &str = "global_account_data";
64    pub const ROOM_ACCOUNT_DATA: &str = "room_account_data";
65    pub const MEMBER: &str = "member";
66    pub const PROFILE: &str = "profile";
67    pub const RECEIPT: &str = "receipt";
68    pub const DISPLAY_NAME: &str = "display_name";
69    pub const SEND_QUEUE: &str = "send_queue_events";
70    pub const DEPENDENTS_SEND_QUEUE: &str = "dependent_send_queue_events";
71    pub const THREAD_SUBSCRIPTIONS: &str = "thread_subscriptions";
72    pub const GLOBAL_PROFILES: &str = "global_profiles";
73}
74
75/// The filename used for the SQLITE database file used by the state store.
76pub const DATABASE_NAME: &str = "matrix-sdk-state.sqlite3";
77
78/// An SQLite-based state store.
79#[derive(Clone)]
80pub struct SqliteStateStore {
81    store_cipher: Option<Arc<StoreCipher>>,
82
83    /// `Some` when active, `None` when closed.
84    connections: Arc<Mutex<Option<SqliteConnections>>>,
85
86    /// Retained so we can rebuild the pool on reopen.
87    db_path: PathBuf,
88
89    /// Retained so we can rebuild the pool on reopen.
90    pool_config: PoolConfig,
91
92    /// Retained so we can re-apply runtime config on reopen.
93    runtime_config: RuntimeConfig,
94}
95
96#[cfg(not(tarpaulin_include))]
97impl fmt::Debug for SqliteStateStore {
98    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
99        f.debug_struct("SqliteStateStore").finish_non_exhaustive()
100    }
101}
102
103impl SqliteStateStore {
104    /// Open the SQLite-based state store at the given path using the given
105    /// given passphrase to encrypt private data.
106    pub async fn open(
107        path: impl AsRef<Path>,
108        passphrase: Option<&str>,
109    ) -> Result<Self, OpenStoreError> {
110        Self::open_with_config(&SqliteStoreConfig::new(path).passphrase(passphrase)).await
111    }
112
113    /// Open the SQLite-based state store at the given path using the given
114    /// key to encrypt private data.
115    pub async fn open_with_key(
116        path: impl AsRef<Path>,
117        key: Option<&[u8]>,
118    ) -> Result<Self, OpenStoreError> {
119        Self::open_with_config(&SqliteStoreConfig::new(path).key(key)).await
120    }
121
122    /// Open the SQLite-based state store with the config open config.
123    pub async fn open_with_config(config: &SqliteStoreConfig) -> Result<Self, OpenStoreError> {
124        fs::create_dir_all(&config.path).await.map_err(OpenStoreError::CreateDir)?;
125
126        let pool = config.build_pool_of_connections(DATABASE_NAME)?;
127        let pool_config = config.pool_config;
128        let runtime_config = config.runtime_config;
129
130        let this =
131            Self::open_with_pool(pool, config.secret.clone(), pool_config, runtime_config).await?;
132        this.read().await?.apply_runtime_config(runtime_config).await?;
133
134        Ok(this)
135    }
136
137    /// Create an SQLite-based state store using the given SQLite database pool.
138    /// The given secret will be used to encrypt private data.
139    pub(crate) async fn open_with_pool(
140        pool: SqlitePool,
141        secret: Option<Secret>,
142        pool_config: PoolConfig,
143        runtime_config: RuntimeConfig,
144    ) -> Result<Self, OpenStoreError> {
145        let db_path = pool.manager().database_path.clone();
146        let conn = pool.get().await?;
147
148        let mut version = conn.db_version().await?;
149
150        if version == 0 {
151            init(&conn).await?;
152            version = 1;
153        }
154
155        let store_cipher = match secret {
156            Some(s) => Some(Arc::new(conn.get_or_create_store_cipher(s).await?)),
157            None => None,
158        };
159        let this = Self {
160            store_cipher,
161            connections: Arc::new(Mutex::new(Some(SqliteConnections {
162                pool,
163                write_connection: Arc::new(Mutex::new(conn)),
164            }))),
165            db_path,
166            pool_config,
167            runtime_config,
168        };
169        this.run_migrations(version, None).await?;
170
171        this.read().await?.wal_checkpoint().await;
172
173        Ok(this)
174    }
175
176    /// Run database migrations from the given `from` version to the given `to`
177    /// version
178    ///
179    /// If `to` is `None`, the current database version will be used.
180    async fn run_migrations(&self, from: u8, to: Option<u8>) -> Result<()> {
181        if to == Some(1) {
182            return Ok(());
183        }
184
185        let conn = self.write().await?;
186
187        if from < 2 {
188            debug!("Upgrading database to version 2");
189            let this = self.clone();
190            conn.with_transaction(move |txn| {
191                // Create new table.
192                txn.execute_batch(include_str!(
193                    "../migrations/state_store/002_a_create_new_room_info.sql"
194                ))?;
195
196                // Migrate data to new table.
197                for data in txn
198                    .prepare("SELECT data FROM room_info")?
199                    .query_map((), |row| row.get::<_, Vec<u8>>(0))?
200                {
201                    let data = data?;
202                    let room_info: RoomInfoV1 = this.deserialize_json(&data)?;
203
204                    let room_id = this.encode_key(keys::ROOM_INFO, room_info.room_id());
205                    let state = this
206                        .encode_key(keys::ROOM_INFO, serde_json::to_string(&room_info.state())?);
207                    txn.prepare_cached(
208                        "INSERT OR REPLACE INTO new_room_info (room_id, state, data)
209                         VALUES (?, ?, ?)",
210                    )?
211                    .execute((room_id, state, data))?;
212                }
213
214                // Replace old table.
215                txn.execute_batch(include_str!(
216                    "../migrations/state_store/002_b_replace_room_info.sql"
217                ))?;
218
219                txn.set_db_version(2)?;
220                Result::<_, Error>::Ok(())
221            })
222            .await?;
223        }
224
225        if to == Some(2) {
226            return Ok(());
227        }
228
229        // Migration to v3: RoomInfo format has changed.
230        if from < 3 {
231            debug!("Upgrading database to version 3");
232            let this = self.clone();
233            conn.with_transaction(move |txn| {
234                // Migrate data .
235                for data in txn
236                    .prepare("SELECT data FROM room_info")?
237                    .query_map((), |row| row.get::<_, Vec<u8>>(0))?
238                {
239                    let data = data?;
240                    let room_info_v1: RoomInfoV1 = this.deserialize_json(&data)?;
241
242                    // Get the `m.room.create` event from the room state.
243                    let room_id = this.encode_key(keys::STATE_EVENT, room_info_v1.room_id());
244                    let event_type =
245                        this.encode_key(keys::STATE_EVENT, StateEventType::RoomCreate.to_string());
246                    let create_res = txn
247                        .prepare(
248                            "SELECT stripped, data FROM state_event
249                             WHERE room_id = ? AND event_type = ?",
250                        )?
251                        .query_one([room_id, event_type], |row| {
252                            Ok((row.get::<_, bool>(0)?, row.get::<_, Vec<u8>>(1)?))
253                        })
254                        .optional()?;
255
256                    let create = create_res.and_then(|(stripped, data)| {
257                        let create = if stripped {
258                            SyncOrStrippedState::<RoomCreateEventContent>::Stripped(
259                                this.deserialize_json(&data).ok()?,
260                            )
261                        } else {
262                            SyncOrStrippedState::Sync(this.deserialize_json(&data).ok()?)
263                        };
264                        Some(create)
265                    });
266
267                    let migrated_room_info = room_info_v1.migrate(create.as_ref());
268
269                    let data = this.serialize_json(&migrated_room_info)?;
270                    let room_id = this.encode_key(keys::ROOM_INFO, migrated_room_info.room_id());
271                    txn.prepare_cached("UPDATE room_info SET data = ? WHERE room_id = ?")?
272                        .execute((data, room_id))?;
273                }
274
275                txn.set_db_version(3)?;
276                Result::<_, Error>::Ok(())
277            })
278            .await?;
279        }
280
281        if to == Some(3) {
282            return Ok(());
283        }
284
285        if from < 4 {
286            debug!("Upgrading database to version 4");
287            conn.with_transaction(move |txn| {
288                // Create new table.
289                txn.execute_batch(include_str!("../migrations/state_store/003_send_queue.sql"))?;
290                txn.set_db_version(4)
291            })
292            .await?;
293        }
294
295        if to == Some(4) {
296            return Ok(());
297        }
298
299        if from < 5 {
300            debug!("Upgrading database to version 5");
301            conn.with_transaction(move |txn| {
302                // Create new table.
303                txn.execute_batch(include_str!(
304                    "../migrations/state_store/004_send_queue_with_roomid_value.sql"
305                ))?;
306                txn.set_db_version(4)
307            })
308            .await?;
309        }
310
311        if to == Some(5) {
312            return Ok(());
313        }
314
315        if from < 6 {
316            debug!("Upgrading database to version 6");
317            conn.with_transaction(move |txn| {
318                // Create new table.
319                txn.execute_batch(include_str!(
320                    "../migrations/state_store/005_send_queue_dependent_events.sql"
321                ))?;
322                txn.set_db_version(6)
323            })
324            .await?;
325        }
326
327        if to == Some(6) {
328            return Ok(());
329        }
330
331        if from < 7 {
332            debug!("Upgrading database to version 7");
333            conn.with_transaction(move |txn| {
334                // Drop media table.
335                txn.execute_batch(include_str!("../migrations/state_store/006_drop_media.sql"))?;
336                txn.set_db_version(7)
337            })
338            .await?;
339        }
340
341        if to == Some(7) {
342            return Ok(());
343        }
344
345        if from < 8 {
346            debug!("Upgrading database to version 8");
347            // Replace all existing wedged events with a generic error.
348            let error = QueueWedgeError::GenericApiError {
349                msg: "local echo failed to send in a previous session".into(),
350            };
351            let default_err = self.serialize_value(&error)?;
352
353            conn.with_transaction(move |txn| {
354                // Update send queue table to persist the wedge reason if any.
355                txn.execute_batch(include_str!("../migrations/state_store/007_a_send_queue_wedge_reason.sql"))?;
356
357                // Migrate the data, add a generic error for currently wedged events
358
359                for wedged_entries in txn
360                    .prepare("SELECT room_id, transaction_id FROM send_queue_events WHERE wedged = 1")?
361                    .query_map((), |row| {
362                        Ok(
363                            (row.get::<_, Vec<u8>>(0)?,row.get::<_, String>(1)?)
364                        )
365                    })? {
366
367                    let (room_id, transaction_id) = wedged_entries?;
368
369                    txn.prepare_cached("UPDATE send_queue_events SET wedge_reason = ? WHERE room_id = ? AND transaction_id = ?")?
370                        .execute((default_err.clone(), room_id, transaction_id))?;
371                }
372
373
374                // Clean up the table now that data is migrated
375                txn.execute_batch(include_str!("../migrations/state_store/007_b_send_queue_clean.sql"))?;
376
377                txn.set_db_version(8)
378            })
379                .await?;
380        }
381
382        if to == Some(8) {
383            return Ok(());
384        }
385
386        if from < 9 {
387            debug!("Upgrading database to version 9");
388            conn.with_transaction(move |txn| {
389                // Run the migration.
390                txn.execute_batch(include_str!("../migrations/state_store/008_send_queue.sql"))?;
391                txn.set_db_version(9)
392            })
393            .await?;
394        }
395
396        if to == Some(9) {
397            return Ok(());
398        }
399
400        if from < 10 {
401            debug!("Upgrading database to version 10");
402            conn.with_transaction(move |txn| {
403                // Run the migration.
404                txn.execute_batch(include_str!(
405                    "../migrations/state_store/009_send_queue_priority.sql"
406                ))?;
407                txn.set_db_version(10)
408            })
409            .await?;
410        }
411
412        if to == Some(10) {
413            return Ok(());
414        }
415
416        if from < 11 {
417            debug!("Upgrading database to version 11");
418            conn.with_transaction(move |txn| {
419                // Run the migration.
420                txn.execute_batch(include_str!(
421                    "../migrations/state_store/010_send_queue_enqueue_time.sql"
422                ))?;
423                txn.set_db_version(11)
424            })
425            .await?;
426        }
427
428        if to == Some(11) {
429            return Ok(());
430        }
431
432        if from < 12 {
433            debug!("Upgrading database to version 12");
434            // Defragment the DB and optimize its size on the filesystem.
435            // This should have been run in the migration for version 7, to reduce the size
436            // of the DB as we removed the media cache.
437            conn.vacuum().await?;
438            conn.set_kv("version", vec![12]).await?;
439        }
440
441        if to == Some(12) {
442            return Ok(());
443        }
444
445        if from < 13 {
446            debug!("Upgrading database to version 13");
447            conn.with_transaction(move |txn| {
448                // Run the migration.
449                txn.execute_batch(include_str!(
450                    "../migrations/state_store/011_thread_subscriptions.sql"
451                ))?;
452                txn.set_db_version(13)
453            })
454            .await?;
455        }
456
457        if to == Some(13) {
458            return Ok(());
459        }
460
461        if from < 14 {
462            debug!("Upgrading database to version 14");
463            conn.with_transaction(move |txn| {
464                // Run the migration.
465                txn.execute_batch(include_str!(
466                    "../migrations/state_store/012_thread_subscriptions_bumpstamp.sql"
467                ))?;
468                txn.set_db_version(14)
469            })
470            .await?;
471        }
472
473        if to == Some(14) {
474            return Ok(());
475        }
476
477        if from < 15 {
478            debug!("Upgrading database to version 15");
479            conn.with_transaction(move |txn| {
480                // Run the migration.
481                txn.execute_batch(include_str!(
482                    "../migrations/state_store/013_send_queue_new_parent_key_format.sql"
483                ))?;
484                txn.set_db_version(15)
485            })
486            .await?;
487        }
488
489        if to == Some(15) {
490            return Ok(());
491        }
492
493        if from < 16 {
494            debug!("Upgrading database to version 16");
495            conn.with_transaction(move |txn| {
496                // Run the migration.
497                txn.execute_batch(include_str!(
498                    "../migrations/state_store/014_global_profiles.sql"
499                ))?;
500                txn.set_db_version(16)
501            })
502            .await?;
503        }
504
505        if to == Some(16) {
506            return Ok(());
507        }
508
509        Ok(())
510    }
511
512    fn encode_state_store_data_key(&self, key: StateStoreDataKey<'_>) -> Key {
513        let key_s = match key {
514            StateStoreDataKey::SyncToken => Cow::Borrowed(StateStoreDataKey::SYNC_TOKEN),
515            StateStoreDataKey::SupportedVersions => {
516                Cow::Borrowed(StateStoreDataKey::SUPPORTED_VERSIONS)
517            }
518            StateStoreDataKey::WellKnown => Cow::Borrowed(StateStoreDataKey::WELL_KNOWN),
519            StateStoreDataKey::Filter(f) => {
520                Cow::Owned(format!("{}:{f}", StateStoreDataKey::FILTER))
521            }
522            StateStoreDataKey::UserAvatarUrl(u) => {
523                Cow::Owned(format!("{}:{u}", StateStoreDataKey::USER_AVATAR_URL))
524            }
525            StateStoreDataKey::RecentlyVisitedRooms(b) => {
526                Cow::Owned(format!("{}:{b}", StateStoreDataKey::RECENTLY_VISITED_ROOMS))
527            }
528            StateStoreDataKey::UtdHookManagerData => {
529                Cow::Borrowed(StateStoreDataKey::UTD_HOOK_MANAGER_DATA)
530            }
531            StateStoreDataKey::OneTimeKeyAlreadyUploaded => {
532                Cow::Borrowed(StateStoreDataKey::ONE_TIME_KEY_ALREADY_UPLOADED)
533            }
534            StateStoreDataKey::ComposerDraft(room_id, thread_root) => {
535                if let Some(thread_root) = thread_root {
536                    Cow::Owned(format!(
537                        "{}:{room_id}:{thread_root}",
538                        StateStoreDataKey::COMPOSER_DRAFT
539                    ))
540                } else {
541                    Cow::Owned(format!("{}:{room_id}", StateStoreDataKey::COMPOSER_DRAFT))
542                }
543            }
544            StateStoreDataKey::SeenKnockRequests(room_id) => {
545                Cow::Owned(format!("{}:{room_id}", StateStoreDataKey::SEEN_KNOCK_REQUESTS))
546            }
547            StateStoreDataKey::ThreadSubscriptionsCatchupTokens => {
548                Cow::Borrowed(StateStoreDataKey::THREAD_SUBSCRIPTIONS_CATCHUP_TOKENS)
549            }
550            StateStoreDataKey::HomeserverCapabilities => {
551                Cow::Borrowed(StateStoreDataKey::HOMESERVER_CAPABILITIES)
552            }
553        };
554
555        self.encode_key(keys::KV_BLOB, &*key_s)
556    }
557
558    fn encode_presence_key(&self, user_id: &UserId) -> Key {
559        self.encode_key(keys::KV_BLOB, format!("presence:{user_id}"))
560    }
561
562    fn encode_custom_key(&self, key: &[u8]) -> Key {
563        let mut full_key = b"custom:".to_vec();
564        full_key.extend(key);
565        self.encode_key(keys::KV_BLOB, full_key)
566    }
567
568    /// Acquire a connection for executing read operations.
569    /// Returns `StoreClosed` if closed.
570    #[instrument(skip_all)]
571    async fn read(&self) -> Result<SqliteAsyncConn> {
572        let pool = {
573            let guard = self.connections.lock().await;
574            let conns = guard.as_ref().ok_or(Error::StoreClosed)?;
575            conns.pool.clone()
576        };
577        Ok(pool.get().await?)
578    }
579
580    /// Acquire a connection for executing write operations.
581    /// Returns `StoreClosed` if closed.
582    #[instrument(skip_all)]
583    async fn write(&self) -> Result<OwnedMutexGuard<SqliteAsyncConn>> {
584        let write_conn = {
585            let guard = self.connections.lock().await;
586            let conns = guard.as_ref().ok_or(Error::StoreClosed)?;
587            conns.write_connection.clone()
588        };
589        Ok(write_conn.lock_owned().await)
590    }
591
592    fn remove_maybe_stripped_room_data(
593        &self,
594        txn: &Transaction<'_>,
595        room_id: &RoomId,
596        stripped: bool,
597    ) -> rusqlite::Result<()> {
598        let state_event_room_id = self.encode_key(keys::STATE_EVENT, room_id);
599        txn.remove_room_state_events(&state_event_room_id, Some(stripped))?;
600
601        let member_room_id = self.encode_key(keys::MEMBER, room_id);
602        txn.remove_room_members(&member_room_id, Some(stripped))
603    }
604
605    pub async fn vacuum(&self) -> Result<()> {
606        self.write().await?.vacuum().await
607    }
608
609    pub async fn get_db_size(&self) -> Result<Option<usize>> {
610        let read_conn = self.read().await?;
611        Ok(Some(read_conn.get_db_size().await?))
612    }
613}
614
615impl EncryptableStore for SqliteStateStore {
616    fn get_cypher(&self) -> Option<&StoreCipher> {
617        self.store_cipher.as_deref()
618    }
619}
620
621/// Initialize the database.
622async fn init(conn: &SqliteAsyncConn) -> Result<()> {
623    // First turn on WAL mode, this can't be done in the transaction, it fails with
624    // the error message: "cannot change into wal mode from within a transaction".
625    conn.execute_batch("PRAGMA journal_mode = wal;").await?;
626    conn.with_transaction(|txn| {
627        txn.execute_batch(include_str!("../migrations/state_store/001_init.sql"))?;
628        txn.set_db_version(1)?;
629
630        Ok(())
631    })
632    .await
633}
634
635trait SqliteConnectionStateStoreExt {
636    fn set_kv_blob(&self, key: &[u8], value: &[u8]) -> rusqlite::Result<()>;
637
638    fn set_global_account_data(&self, event_type: &[u8], data: &[u8]) -> rusqlite::Result<()>;
639
640    fn set_room_account_data(
641        &self,
642        room_id: &[u8],
643        event_type: &[u8],
644        data: &[u8],
645    ) -> rusqlite::Result<()>;
646    fn remove_room_account_data(&self, room_id: &[u8]) -> rusqlite::Result<()>;
647
648    fn set_room_info(&self, room_id: &[u8], state: &[u8], data: &[u8]) -> rusqlite::Result<()>;
649    fn get_room_info(&self, room_id: &[u8]) -> rusqlite::Result<Option<Vec<u8>>>;
650    fn remove_room_info(&self, room_id: &[u8]) -> rusqlite::Result<()>;
651
652    fn set_state_event(
653        &self,
654        room_id: &[u8],
655        event_type: &[u8],
656        state_key: &[u8],
657        stripped: bool,
658        event_id: Option<&[u8]>,
659        data: &[u8],
660    ) -> rusqlite::Result<()>;
661    fn get_state_event_by_id(
662        &self,
663        room_id: &[u8],
664        event_id: &[u8],
665    ) -> rusqlite::Result<Option<Vec<u8>>>;
666    fn remove_room_state_events(
667        &self,
668        room_id: &[u8],
669        stripped: Option<bool>,
670    ) -> rusqlite::Result<()>;
671
672    fn set_member(
673        &self,
674        room_id: &[u8],
675        user_id: &[u8],
676        membership: &[u8],
677        stripped: bool,
678        data: &[u8],
679    ) -> rusqlite::Result<()>;
680    fn remove_room_members(&self, room_id: &[u8], stripped: Option<bool>) -> rusqlite::Result<()>;
681
682    fn set_profile(&self, room_id: &[u8], user_id: &[u8], data: &[u8]) -> rusqlite::Result<()>;
683    fn remove_room_profiles(&self, room_id: &[u8]) -> rusqlite::Result<()>;
684    fn remove_room_profile(&self, room_id: &[u8], user_id: &[u8]) -> rusqlite::Result<()>;
685
686    fn set_receipt(
687        &self,
688        room_id: &[u8],
689        user_id: &[u8],
690        receipt_type: &[u8],
691        thread_id: &[u8],
692        event_id: &[u8],
693        data: &[u8],
694    ) -> rusqlite::Result<()>;
695    fn remove_room_receipts(&self, room_id: &[u8]) -> rusqlite::Result<()>;
696
697    fn set_display_name(&self, room_id: &[u8], name: &[u8], data: &[u8]) -> rusqlite::Result<()>;
698    fn remove_display_name(&self, room_id: &[u8], name: &[u8]) -> rusqlite::Result<()>;
699    fn remove_room_display_names(&self, room_id: &[u8]) -> rusqlite::Result<()>;
700    fn remove_room_send_queue(&self, room_id: &[u8]) -> rusqlite::Result<()>;
701    fn remove_room_dependent_send_queue(&self, room_id: &[u8]) -> rusqlite::Result<()>;
702}
703
704impl SqliteConnectionStateStoreExt for rusqlite::Connection {
705    fn set_kv_blob(&self, key: &[u8], value: &[u8]) -> rusqlite::Result<()> {
706        self.execute("INSERT OR REPLACE INTO kv_blob VALUES (?, ?)", (key, value))?;
707        Ok(())
708    }
709
710    fn set_global_account_data(&self, event_type: &[u8], data: &[u8]) -> rusqlite::Result<()> {
711        self.prepare_cached(
712            "INSERT OR REPLACE INTO global_account_data (event_type, data)
713             VALUES (?, ?)",
714        )?
715        .execute((event_type, data))?;
716        Ok(())
717    }
718
719    fn set_room_account_data(
720        &self,
721        room_id: &[u8],
722        event_type: &[u8],
723        data: &[u8],
724    ) -> rusqlite::Result<()> {
725        self.prepare_cached(
726            "INSERT OR REPLACE INTO room_account_data (room_id, event_type, data)
727             VALUES (?, ?, ?)",
728        )?
729        .execute((room_id, event_type, data))?;
730        Ok(())
731    }
732
733    fn remove_room_account_data(&self, room_id: &[u8]) -> rusqlite::Result<()> {
734        self.prepare(
735            "DELETE FROM room_account_data
736             WHERE room_id = ?",
737        )?
738        .execute((room_id,))?;
739        Ok(())
740    }
741
742    fn set_room_info(&self, room_id: &[u8], state: &[u8], data: &[u8]) -> rusqlite::Result<()> {
743        self.prepare_cached(
744            "INSERT OR REPLACE INTO room_info (room_id, state, data)
745             VALUES (?, ?, ?)",
746        )?
747        .execute((room_id, state, data))?;
748        Ok(())
749    }
750
751    fn get_room_info(&self, room_id: &[u8]) -> rusqlite::Result<Option<Vec<u8>>> {
752        self.query_one("SELECT data FROM room_info WHERE room_id = ?", (room_id,), |row| row.get(0))
753            .optional()
754    }
755
756    /// Remove the room info for the given room.
757    fn remove_room_info(&self, room_id: &[u8]) -> rusqlite::Result<()> {
758        self.prepare_cached("DELETE FROM room_info WHERE room_id = ?")?.execute((room_id,))?;
759        Ok(())
760    }
761
762    fn set_state_event(
763        &self,
764        room_id: &[u8],
765        event_type: &[u8],
766        state_key: &[u8],
767        stripped: bool,
768        event_id: Option<&[u8]>,
769        data: &[u8],
770    ) -> rusqlite::Result<()> {
771        self.prepare_cached(
772            "INSERT OR REPLACE
773             INTO state_event (room_id, event_type, state_key, stripped, event_id, data)
774             VALUES (?, ?, ?, ?, ?, ?)",
775        )?
776        .execute((room_id, event_type, state_key, stripped, event_id, data))?;
777        Ok(())
778    }
779
780    fn get_state_event_by_id(
781        &self,
782        room_id: &[u8],
783        event_id: &[u8],
784    ) -> rusqlite::Result<Option<Vec<u8>>> {
785        self.query_one(
786            "SELECT data FROM state_event WHERE room_id = ? AND event_id = ?",
787            (room_id, event_id),
788            |row| row.get(0),
789        )
790        .optional()
791    }
792
793    /// Remove state events for the given room.
794    ///
795    /// If `stripped` is `Some()`, only removes state events for the given
796    /// stripped state. Otherwise, state events are removed regardless of the
797    /// stripped state.
798    fn remove_room_state_events(
799        &self,
800        room_id: &[u8],
801        stripped: Option<bool>,
802    ) -> rusqlite::Result<()> {
803        if let Some(stripped) = stripped {
804            self.prepare_cached("DELETE FROM state_event WHERE room_id = ? AND stripped = ?")?
805                .execute((room_id, stripped))?;
806        } else {
807            self.prepare_cached("DELETE FROM state_event WHERE room_id = ?")?
808                .execute((room_id,))?;
809        }
810        Ok(())
811    }
812
813    fn set_member(
814        &self,
815        room_id: &[u8],
816        user_id: &[u8],
817        membership: &[u8],
818        stripped: bool,
819        data: &[u8],
820    ) -> rusqlite::Result<()> {
821        self.prepare_cached(
822            "INSERT OR REPLACE
823             INTO member (room_id, user_id, membership, stripped, data)
824             VALUES (?, ?, ?, ?, ?)",
825        )?
826        .execute((room_id, user_id, membership, stripped, data))?;
827        Ok(())
828    }
829
830    /// Remove members for the given room.
831    ///
832    /// If `stripped` is `Some()`, only removes members for the given stripped
833    /// state. Otherwise, members are removed regardless of the stripped state.
834    fn remove_room_members(&self, room_id: &[u8], stripped: Option<bool>) -> rusqlite::Result<()> {
835        if let Some(stripped) = stripped {
836            self.prepare_cached("DELETE FROM member WHERE room_id = ? AND stripped = ?")?
837                .execute((room_id, stripped))?;
838        } else {
839            self.prepare_cached("DELETE FROM member WHERE room_id = ?")?.execute((room_id,))?;
840        }
841        Ok(())
842    }
843
844    fn set_profile(&self, room_id: &[u8], user_id: &[u8], data: &[u8]) -> rusqlite::Result<()> {
845        self.prepare_cached(
846            "INSERT OR REPLACE
847             INTO profile (room_id, user_id, data)
848             VALUES (?, ?, ?)",
849        )?
850        .execute((room_id, user_id, data))?;
851        Ok(())
852    }
853
854    fn remove_room_profiles(&self, room_id: &[u8]) -> rusqlite::Result<()> {
855        self.prepare("DELETE FROM profile WHERE room_id = ?")?.execute((room_id,))?;
856        Ok(())
857    }
858
859    fn remove_room_profile(&self, room_id: &[u8], user_id: &[u8]) -> rusqlite::Result<()> {
860        self.prepare("DELETE FROM profile WHERE room_id = ? AND user_id = ?")?
861            .execute((room_id, user_id))?;
862        Ok(())
863    }
864
865    fn set_receipt(
866        &self,
867        room_id: &[u8],
868        user_id: &[u8],
869        receipt_type: &[u8],
870        thread: &[u8],
871        event_id: &[u8],
872        data: &[u8],
873    ) -> rusqlite::Result<()> {
874        self.prepare_cached(
875            "INSERT OR REPLACE
876             INTO receipt (room_id, user_id, receipt_type, thread, event_id, data)
877             VALUES (?, ?, ?, ?, ?, ?)",
878        )?
879        .execute((room_id, user_id, receipt_type, thread, event_id, data))?;
880        Ok(())
881    }
882
883    fn remove_room_receipts(&self, room_id: &[u8]) -> rusqlite::Result<()> {
884        self.prepare("DELETE FROM receipt WHERE room_id = ?")?.execute((room_id,))?;
885        Ok(())
886    }
887
888    fn set_display_name(&self, room_id: &[u8], name: &[u8], data: &[u8]) -> rusqlite::Result<()> {
889        self.prepare_cached(
890            "INSERT OR REPLACE
891             INTO display_name (room_id, name, data)
892             VALUES (?, ?, ?)",
893        )?
894        .execute((room_id, name, data))?;
895        Ok(())
896    }
897
898    fn remove_display_name(&self, room_id: &[u8], name: &[u8]) -> rusqlite::Result<()> {
899        self.prepare("DELETE FROM display_name WHERE room_id = ? AND name = ?")?
900            .execute((room_id, name))?;
901        Ok(())
902    }
903
904    fn remove_room_display_names(&self, room_id: &[u8]) -> rusqlite::Result<()> {
905        self.prepare("DELETE FROM display_name WHERE room_id = ?")?.execute((room_id,))?;
906        Ok(())
907    }
908
909    fn remove_room_send_queue(&self, room_id: &[u8]) -> rusqlite::Result<()> {
910        self.prepare("DELETE FROM send_queue_events WHERE room_id = ?")?.execute((room_id,))?;
911        Ok(())
912    }
913
914    fn remove_room_dependent_send_queue(&self, room_id: &[u8]) -> rusqlite::Result<()> {
915        self.prepare("DELETE FROM dependent_send_queue_events WHERE room_id = ?")?
916            .execute((room_id,))?;
917        Ok(())
918    }
919}
920
921#[async_trait]
922trait SqliteObjectStateStoreExt: SqliteAsyncConnExt {
923    async fn get_kv_blob(&self, key: Key) -> Result<Option<Vec<u8>>> {
924        Ok(self
925            .query_one("SELECT value FROM kv_blob WHERE key = ?", (key,), |row| row.get(0))
926            .await
927            .optional()?)
928    }
929
930    async fn get_kv_blobs(&self, keys: Vec<Key>) -> Result<Vec<Vec<u8>>> {
931        let keys_length = keys.len();
932
933        self.chunk_large_query_over(keys, Some(keys_length), |txn, keys| {
934            let sql =
935                format!("SELECT value FROM kv_blob WHERE key IN ({})", keys.host_parameters());
936
937            let params = rusqlite::params_from_iter(keys);
938
939            Ok(txn
940                .prepare(&sql)?
941                .query(params)?
942                .mapped(|row| row.get(0))
943                .collect::<Result<_, _>>()?)
944        })
945        .await
946    }
947
948    async fn set_kv_blob(&self, key: Key, value: Vec<u8>) -> Result<()>;
949
950    async fn delete_kv_blob(&self, key: Key) -> Result<()> {
951        self.execute("DELETE FROM kv_blob WHERE key = ?", (key,)).await?;
952        Ok(())
953    }
954
955    async fn get_room_infos(&self, room_id: Option<Key>) -> Result<Vec<Vec<u8>>> {
956        Ok(match room_id {
957            None => {
958                self.prepare("SELECT data FROM room_info", move |mut stmt| {
959                    stmt.query_map((), |row| row.get(0))?.collect()
960                })
961                .await?
962            }
963
964            Some(room_id) => {
965                self.prepare("SELECT data FROM room_info WHERE room_id = ?", move |mut stmt| {
966                    stmt.query((room_id,))?.mapped(|row| row.get(0)).collect()
967                })
968                .await?
969            }
970        })
971    }
972
973    async fn get_maybe_stripped_state_events_for_keys(
974        &self,
975        room_id: Key,
976        event_type: Key,
977        state_keys: Vec<Key>,
978    ) -> Result<Vec<(bool, Vec<u8>)>> {
979        self.chunk_large_query_over(state_keys, None, move |txn, state_keys| {
980            let sql = format!(
981                "SELECT stripped, data FROM state_event
982                 WHERE room_id = ? AND event_type = ? AND state_key IN ({})",
983                state_keys.host_parameters()
984            );
985
986            let params = rusqlite::params_from_iter(
987                [room_id.clone(), event_type.clone()].into_iter().chain(state_keys),
988            );
989
990            Ok(txn
991                .prepare(&sql)?
992                .query(params)?
993                .mapped(|row| Ok((row.get(0)?, row.get(1)?)))
994                .collect::<Result<_, _>>()?)
995        })
996        .await
997    }
998
999    async fn get_maybe_stripped_state_events(
1000        &self,
1001        room_id: Key,
1002        event_type: Key,
1003    ) -> Result<Vec<(bool, Vec<u8>)>> {
1004        Ok(self
1005            .prepare(
1006                "SELECT stripped, data FROM state_event
1007                 WHERE room_id = ? AND event_type = ?",
1008                |mut stmt| {
1009                    stmt.query((room_id, event_type))?
1010                        .mapped(|row| Ok((row.get(0)?, row.get(1)?)))
1011                        .collect()
1012                },
1013            )
1014            .await?)
1015    }
1016
1017    async fn get_profiles(
1018        &self,
1019        room_id: Key,
1020        user_ids: Vec<Key>,
1021    ) -> Result<Vec<(Vec<u8>, Vec<u8>)>> {
1022        let user_ids_length = user_ids.len();
1023
1024        self.chunk_large_query_over(user_ids, Some(user_ids_length), move |txn, user_ids| {
1025            let sql = format!(
1026                "SELECT user_id, data FROM profile WHERE room_id = ? AND user_id IN ({})",
1027                user_ids.host_parameters(),
1028            );
1029
1030            let params = rusqlite::params_from_iter(iter::once(room_id.clone()).chain(user_ids));
1031
1032            Ok(txn
1033                .prepare(&sql)?
1034                .query(params)?
1035                .mapped(|row| Ok((row.get(0)?, row.get(1)?)))
1036                .collect::<Result<_, _>>()?)
1037        })
1038        .await
1039    }
1040
1041    async fn get_global_profiles(&self, user_ids: Vec<Key>) -> Result<Vec<(Vec<u8>, Vec<u8>)>> {
1042        let user_ids_length = user_ids.len();
1043
1044        self.chunk_large_query_over(user_ids, Some(user_ids_length), move |txn, user_ids| {
1045            let sql = format!(
1046                "SELECT user_id, profile_data FROM global_profiles WHERE user_id IN ({})",
1047                user_ids.host_parameters(),
1048            );
1049
1050            let params = rusqlite::params_from_iter(user_ids);
1051
1052            Ok(txn
1053                .prepare(&sql)?
1054                .query(params)?
1055                .mapped(|row| Ok((row.get(0)?, row.get(1)?)))
1056                .collect::<Result<_, _>>()?)
1057        })
1058        .await
1059    }
1060
1061    async fn get_user_ids(&self, room_id: Key, memberships: Vec<Key>) -> Result<Vec<Vec<u8>>> {
1062        let res = if memberships.is_empty() {
1063            self.prepare("SELECT data FROM member WHERE room_id = ?", |mut stmt| {
1064                stmt.query((room_id,))?.mapped(|row| row.get(0)).collect()
1065            })
1066            .await?
1067        } else {
1068            self.chunk_large_query_over(memberships, None, move |txn, memberships| {
1069                let sql = format!(
1070                    "SELECT data FROM member WHERE room_id = ? AND membership IN ({})",
1071                    memberships.host_parameters(),
1072                );
1073
1074                let params =
1075                    rusqlite::params_from_iter(iter::once(room_id.clone()).chain(memberships));
1076
1077                Ok(txn
1078                    .prepare(&sql)?
1079                    .query(params)?
1080                    .mapped(|row| row.get(0))
1081                    .collect::<Result<_, _>>()?)
1082            })
1083            .await?
1084        };
1085
1086        Ok(res)
1087    }
1088
1089    async fn get_global_account_data(&self, event_type: Key) -> Result<Option<Vec<u8>>> {
1090        Ok(self
1091            .query_one(
1092                "SELECT data FROM global_account_data WHERE event_type = ?",
1093                (event_type,),
1094                |row| row.get(0),
1095            )
1096            .await
1097            .optional()?)
1098    }
1099
1100    async fn get_room_account_data(
1101        &self,
1102        room_id: Key,
1103        event_type: Key,
1104    ) -> Result<Option<Vec<u8>>> {
1105        Ok(self
1106            .query_one(
1107                "SELECT data FROM room_account_data WHERE room_id = ? AND event_type = ?",
1108                (room_id, event_type),
1109                |row| row.get(0),
1110            )
1111            .await
1112            .optional()?)
1113    }
1114
1115    async fn get_display_names(
1116        &self,
1117        room_id: Key,
1118        names: Vec<Key>,
1119    ) -> Result<Vec<(Vec<u8>, Vec<u8>)>> {
1120        let names_length = names.len();
1121
1122        self.chunk_large_query_over(names, Some(names_length), move |txn, names| {
1123            let sql = format!(
1124                "SELECT name, data FROM display_name WHERE room_id = ? AND name IN ({})",
1125                names.host_parameters()
1126            );
1127
1128            let params = rusqlite::params_from_iter(iter::once(room_id.clone()).chain(names));
1129
1130            Ok(txn
1131                .prepare(&sql)?
1132                .query(params)?
1133                .mapped(|row| Ok((row.get(0)?, row.get(1)?)))
1134                .collect::<Result<_, _>>()?)
1135        })
1136        .await
1137    }
1138
1139    async fn get_user_receipt(
1140        &self,
1141        room_id: Key,
1142        receipt_type: Key,
1143        receipt_thread: Key,
1144        user_id: Key,
1145    ) -> Result<Option<Vec<u8>>> {
1146        Ok(self
1147            .query_one(
1148                "SELECT data FROM receipt
1149                 WHERE room_id = ? AND receipt_type = ? AND thread = ? and user_id = ?",
1150                (room_id, receipt_type, receipt_thread, user_id),
1151                |row| row.get(0),
1152            )
1153            .await
1154            .optional()?)
1155    }
1156
1157    async fn get_event_receipts(
1158        &self,
1159        room_id: Key,
1160        receipt_type: Key,
1161        thread: Key,
1162        event_id: Key,
1163    ) -> Result<Vec<Vec<u8>>> {
1164        Ok(self
1165            .prepare(
1166                "SELECT data FROM receipt
1167                 WHERE room_id = ? AND receipt_type = ? AND thread = ? and event_id = ?",
1168                |mut stmt| {
1169                    stmt.query((room_id, receipt_type, thread, event_id))?
1170                        .mapped(|row| row.get(0))
1171                        .collect()
1172                },
1173            )
1174            .await?)
1175    }
1176}
1177
1178#[async_trait]
1179impl SqliteObjectStateStoreExt for SqliteAsyncConn {
1180    async fn set_kv_blob(&self, key: Key, value: Vec<u8>) -> Result<()> {
1181        Ok(self.interact(move |conn| conn.set_kv_blob(&key, &value)).await.unwrap()?)
1182    }
1183}
1184
1185#[async_trait]
1186impl StateStore for SqliteStateStore {
1187    type Error = Error;
1188
1189    async fn get_kv_data(&self, key: StateStoreDataKey<'_>) -> Result<Option<StateStoreDataValue>> {
1190        self.read()
1191            .await?
1192            .get_kv_blob(self.encode_state_store_data_key(key))
1193            .await?
1194            .map(|data| {
1195                Ok(match key {
1196                    StateStoreDataKey::SyncToken => {
1197                        StateStoreDataValue::SyncToken(self.deserialize_value(&data)?)
1198                    }
1199                    StateStoreDataKey::SupportedVersions => {
1200                        StateStoreDataValue::SupportedVersions(self.deserialize_value(&data)?)
1201                    }
1202                    StateStoreDataKey::WellKnown => {
1203                        StateStoreDataValue::WellKnown(self.deserialize_value(&data)?)
1204                    }
1205                    StateStoreDataKey::Filter(_) => {
1206                        StateStoreDataValue::Filter(self.deserialize_value(&data)?)
1207                    }
1208                    StateStoreDataKey::UserAvatarUrl(_) => {
1209                        StateStoreDataValue::UserAvatarUrl(self.deserialize_value(&data)?)
1210                    }
1211                    StateStoreDataKey::RecentlyVisitedRooms(_) => {
1212                        StateStoreDataValue::RecentlyVisitedRooms(self.deserialize_value(&data)?)
1213                    }
1214                    StateStoreDataKey::UtdHookManagerData => {
1215                        StateStoreDataValue::UtdHookManagerData(self.deserialize_value(&data)?)
1216                    }
1217                    StateStoreDataKey::OneTimeKeyAlreadyUploaded => {
1218                        StateStoreDataValue::OneTimeKeyAlreadyUploaded
1219                    }
1220                    StateStoreDataKey::ComposerDraft(_, _) => {
1221                        StateStoreDataValue::ComposerDraft(self.deserialize_value(&data)?)
1222                    }
1223                    StateStoreDataKey::SeenKnockRequests(_) => {
1224                        StateStoreDataValue::SeenKnockRequests(self.deserialize_value(&data)?)
1225                    }
1226                    StateStoreDataKey::ThreadSubscriptionsCatchupTokens => {
1227                        StateStoreDataValue::ThreadSubscriptionsCatchupTokens(
1228                            self.deserialize_value(&data)?,
1229                        )
1230                    }
1231                    StateStoreDataKey::HomeserverCapabilities => {
1232                        StateStoreDataValue::HomeserverCapabilities(self.deserialize_value(&data)?)
1233                    }
1234                })
1235            })
1236            .transpose()
1237    }
1238
1239    async fn set_kv_data(
1240        &self,
1241        key: StateStoreDataKey<'_>,
1242        value: StateStoreDataValue,
1243    ) -> Result<()> {
1244        let serialized_value = match key {
1245            StateStoreDataKey::SyncToken => self.serialize_value(
1246                &value.into_sync_token().expect("Session data not a sync token"),
1247            )?,
1248            StateStoreDataKey::SupportedVersions => self.serialize_value(
1249                &value
1250                    .into_supported_versions()
1251                    .expect("Session data not containing supported versions"),
1252            )?,
1253            StateStoreDataKey::WellKnown => self.serialize_value(
1254                &value.into_well_known().expect("Session data not containing well-known"),
1255            )?,
1256            StateStoreDataKey::Filter(_) => {
1257                self.serialize_value(&value.into_filter().expect("Session data not a filter"))?
1258            }
1259            StateStoreDataKey::UserAvatarUrl(_) => self.serialize_value(
1260                &value.into_user_avatar_url().expect("Session data not an user avatar url"),
1261            )?,
1262            StateStoreDataKey::RecentlyVisitedRooms(_) => self.serialize_value(
1263                &value.into_recently_visited_rooms().expect("Session data not breadcrumbs"),
1264            )?,
1265            StateStoreDataKey::UtdHookManagerData => self.serialize_value(
1266                &value.into_utd_hook_manager_data().expect("Session data not UtdHookManagerData"),
1267            )?,
1268            StateStoreDataKey::OneTimeKeyAlreadyUploaded => {
1269                self.serialize_value(&true).expect("We should be able to serialize a boolean")
1270            }
1271            StateStoreDataKey::ComposerDraft(_, _) => self.serialize_value(
1272                &value.into_composer_draft().expect("Session data not a composer draft"),
1273            )?,
1274            StateStoreDataKey::SeenKnockRequests(_) => self.serialize_value(
1275                &value
1276                    .into_seen_knock_requests()
1277                    .expect("Session data is not a set of seen knock request ids"),
1278            )?,
1279            StateStoreDataKey::ThreadSubscriptionsCatchupTokens => self.serialize_value(
1280                &value
1281                    .into_thread_subscriptions_catchup_tokens()
1282                    .expect("Session data is not a list of thread subscription catchup tokens"),
1283            )?,
1284            StateStoreDataKey::HomeserverCapabilities => self.serialize_value(
1285                &value
1286                    .into_homeserver_capabilities()
1287                    .expect("Session data is not the homeserver capabilities"),
1288            )?,
1289        };
1290
1291        self.write()
1292            .await?
1293            .set_kv_blob(self.encode_state_store_data_key(key), serialized_value)
1294            .await
1295    }
1296
1297    async fn remove_kv_data(&self, key: StateStoreDataKey<'_>) -> Result<()> {
1298        self.write().await?.delete_kv_blob(self.encode_state_store_data_key(key)).await
1299    }
1300
1301    async fn save_changes(&self, changes: &StateChanges) -> Result<()> {
1302        let changes = changes.to_owned();
1303        let this = self.clone();
1304        self.write()
1305            .await?
1306            .with_transaction(move |txn| {
1307                let StateChanges {
1308                    sync_token,
1309                    account_data,
1310                    presence,
1311                    profiles,
1312                    profiles_to_delete,
1313                    state,
1314                    room_account_data,
1315                    room_infos,
1316                    receipts,
1317                    redactions,
1318                    stripped_state,
1319                    ambiguity_maps,
1320                    global_profiles,
1321                } = changes;
1322
1323                if let Some(sync_token) = sync_token {
1324                    let key = this.encode_state_store_data_key(StateStoreDataKey::SyncToken);
1325                    let value = this.serialize_value(&sync_token)?;
1326                    txn.set_kv_blob(&key, &value)?;
1327                }
1328
1329                for (event_type, event) in account_data {
1330                    let event_type =
1331                        this.encode_key(keys::GLOBAL_ACCOUNT_DATA, event_type.to_string());
1332                    let data = this.serialize_json(&event)?;
1333                    txn.set_global_account_data(&event_type, &data)?;
1334                }
1335
1336                for (room_id, events) in room_account_data {
1337                    let room_id = this.encode_key(keys::ROOM_ACCOUNT_DATA, room_id);
1338                    for (event_type, event) in events {
1339                        let event_type =
1340                            this.encode_key(keys::ROOM_ACCOUNT_DATA, event_type.to_string());
1341                        let data = this.serialize_json(&event)?;
1342                        txn.set_room_account_data(&room_id, &event_type, &data)?;
1343                    }
1344                }
1345
1346                for (user_id, event) in presence {
1347                    let key = this.encode_presence_key(&user_id);
1348                    let value = this.serialize_json(&event)?;
1349                    txn.set_kv_blob(&key, &value)?;
1350                }
1351
1352                for (room_id, room_info) in room_infos {
1353                    let stripped = room_info.state() == RoomState::Invited;
1354                    // Remove non-stripped data for stripped rooms and vice-versa.
1355                    this.remove_maybe_stripped_room_data(txn, &room_id, !stripped)?;
1356
1357                    let room_id = this.encode_key(keys::ROOM_INFO, room_id);
1358                    let state = this
1359                        .encode_key(keys::ROOM_INFO, serde_json::to_string(&room_info.state())?);
1360                    let data = this.serialize_json(&room_info)?;
1361                    txn.set_room_info(&room_id, &state, &data)?;
1362                }
1363
1364                for (room_id, user_ids) in profiles_to_delete {
1365                    let room_id = this.encode_key(keys::PROFILE, room_id);
1366                    for user_id in user_ids {
1367                        let user_id = this.encode_key(keys::PROFILE, user_id);
1368                        txn.remove_room_profile(&room_id, &user_id)?;
1369                    }
1370                }
1371
1372                for (room_id, state_event_types) in state {
1373                    let profiles = profiles.get(&room_id);
1374                    let encoded_room_id = this.encode_key(keys::STATE_EVENT, &room_id);
1375
1376                    for (event_type, state_events) in state_event_types {
1377                        let encoded_event_type =
1378                            this.encode_key(keys::STATE_EVENT, event_type.to_string());
1379
1380                        for (state_key, raw_state_event) in state_events {
1381                            let encoded_state_key = this.encode_key(keys::STATE_EVENT, &state_key);
1382                            let data = this.serialize_json(&raw_state_event)?;
1383
1384                            let event_id: Option<String> =
1385                                raw_state_event.get_field("event_id").ok().flatten();
1386                            let encoded_event_id =
1387                                event_id.as_ref().map(|e| this.encode_key(keys::STATE_EVENT, e));
1388
1389                            txn.set_state_event(
1390                                &encoded_room_id,
1391                                &encoded_event_type,
1392                                &encoded_state_key,
1393                                false,
1394                                encoded_event_id.as_deref(),
1395                                &data,
1396                            )?;
1397
1398                            if event_type == StateEventType::RoomMember {
1399                                let member_event = match raw_state_event
1400                                    .deserialize_as_unchecked::<SyncRoomMemberEvent>()
1401                                {
1402                                    Ok(ev) => ev,
1403                                    Err(e) => {
1404                                        debug!(event_id, "Failed to deserialize member event: {e}");
1405                                        continue;
1406                                    }
1407                                };
1408
1409                                let encoded_room_id = this.encode_key(keys::MEMBER, &room_id);
1410                                let user_id = this.encode_key(keys::MEMBER, &state_key);
1411                                let membership = this
1412                                    .encode_key(keys::MEMBER, member_event.membership().as_str());
1413                                let data = this.serialize_value(&state_key)?;
1414
1415                                txn.set_member(
1416                                    &encoded_room_id,
1417                                    &user_id,
1418                                    &membership,
1419                                    false,
1420                                    &data,
1421                                )?;
1422
1423                                if let Some(profile) =
1424                                    profiles.and_then(|p| p.get(member_event.state_key()))
1425                                {
1426                                    let room_id = this.encode_key(keys::PROFILE, &room_id);
1427                                    let user_id = this.encode_key(keys::PROFILE, &state_key);
1428                                    let data = this.serialize_json(&profile)?;
1429                                    txn.set_profile(&room_id, &user_id, &data)?;
1430                                }
1431                            }
1432                        }
1433                    }
1434                }
1435
1436                for (room_id, stripped_state_event_types) in stripped_state {
1437                    let encoded_room_id = this.encode_key(keys::STATE_EVENT, &room_id);
1438
1439                    for (event_type, stripped_state_events) in stripped_state_event_types {
1440                        let encoded_event_type =
1441                            this.encode_key(keys::STATE_EVENT, event_type.to_string());
1442
1443                        for (state_key, raw_stripped_state_event) in stripped_state_events {
1444                            let encoded_state_key = this.encode_key(keys::STATE_EVENT, &state_key);
1445                            let data = this.serialize_json(&raw_stripped_state_event)?;
1446                            txn.set_state_event(
1447                                &encoded_room_id,
1448                                &encoded_event_type,
1449                                &encoded_state_key,
1450                                true,
1451                                None,
1452                                &data,
1453                            )?;
1454
1455                            if event_type == StateEventType::RoomMember {
1456                                let member_event = match raw_stripped_state_event
1457                                    .deserialize_as_unchecked::<StrippedRoomMemberEvent>(
1458                                ) {
1459                                    Ok(ev) => ev,
1460                                    Err(e) => {
1461                                        debug!("Failed to deserialize stripped member event: {e}");
1462                                        continue;
1463                                    }
1464                                };
1465
1466                                let room_id = this.encode_key(keys::MEMBER, &room_id);
1467                                let user_id = this.encode_key(keys::MEMBER, &state_key);
1468                                let membership = this.encode_key(
1469                                    keys::MEMBER,
1470                                    member_event.content.membership.as_str(),
1471                                );
1472                                let data = this.serialize_value(&state_key)?;
1473
1474                                txn.set_member(&room_id, &user_id, &membership, true, &data)?;
1475                            }
1476                        }
1477                    }
1478                }
1479
1480                for (room_id, receipt_event) in receipts {
1481                    let room_id = this.encode_key(keys::RECEIPT, room_id);
1482
1483                    for (event_id, receipt_types) in receipt_event {
1484                        let encoded_event_id = this.encode_key(keys::RECEIPT, &event_id);
1485
1486                        for (receipt_type, receipt_users) in receipt_types {
1487                            let receipt_type =
1488                                this.encode_key(keys::RECEIPT, receipt_type.as_str());
1489
1490                            for (user_id, receipt) in receipt_users {
1491                                let encoded_user_id = this.encode_key(keys::RECEIPT, &user_id);
1492                                // We cannot have a NULL primary key so we rely on serialization
1493                                // instead of the string representation.
1494                                let thread = this.encode_key(
1495                                    keys::RECEIPT,
1496                                    rmp_serde::to_vec_named(&receipt.thread)?,
1497                                );
1498                                let data = this.serialize_json(&ReceiptData {
1499                                    receipt,
1500                                    event_id: event_id.clone(),
1501                                    user_id,
1502                                })?;
1503
1504                                txn.set_receipt(
1505                                    &room_id,
1506                                    &encoded_user_id,
1507                                    &receipt_type,
1508                                    &thread,
1509                                    &encoded_event_id,
1510                                    &data,
1511                                )?;
1512                            }
1513                        }
1514                    }
1515                }
1516
1517                for (room_id, redactions) in redactions {
1518                    let make_redaction_rules = || {
1519                        let encoded_room_id = this.encode_key(keys::ROOM_INFO, &room_id);
1520                        txn.get_room_info(&encoded_room_id)
1521                            .ok()
1522                            .flatten()
1523                            .and_then(|v| this.deserialize_json::<RoomInfo>(&v).ok())
1524                            .map(|info| info.room_version_rules_or_default())
1525                            .unwrap_or_else(|| {
1526                                warn!(
1527                                    ?room_id,
1528                                    "Unable to get the room version rules, defaulting to rules for room version {ROOM_VERSION_FALLBACK}"
1529                                );
1530                                ROOM_VERSION_RULES_FALLBACK
1531                            }).redaction
1532                    };
1533
1534                    let encoded_room_id = this.encode_key(keys::STATE_EVENT, &room_id);
1535                    let mut redaction_rules = None;
1536
1537                    for (event_id, redaction) in redactions {
1538                        let event_id = this.encode_key(keys::STATE_EVENT, event_id);
1539
1540                        if let Some(Ok(raw_event)) = txn
1541                            .get_state_event_by_id(&encoded_room_id, &event_id)?
1542                            .map(|value| this.deserialize_json::<Raw<AnySyncStateEvent>>(&value))
1543                        {
1544                            let event = raw_event.deserialize()?;
1545                            let redacted = redact(
1546                                raw_event.deserialize_as::<CanonicalJsonObject>()?,
1547                                redaction_rules.get_or_insert_with(make_redaction_rules),
1548                                Some(RedactedBecause::from_raw_event(&redaction)?),
1549                            )
1550                            .map_err(Error::Redaction)?;
1551                            let data = this.serialize_json(&redacted)?;
1552
1553                            let event_type =
1554                                this.encode_key(keys::STATE_EVENT, event.event_type().to_string());
1555                            let state_key = this.encode_key(keys::STATE_EVENT, event.state_key());
1556
1557                            txn.set_state_event(
1558                                &encoded_room_id,
1559                                &event_type,
1560                                &state_key,
1561                                false,
1562                                Some(&event_id),
1563                                &data,
1564                            )?;
1565                        }
1566                    }
1567                }
1568
1569                for (room_id, display_names) in ambiguity_maps {
1570                    let room_id = this.encode_key(keys::DISPLAY_NAME, room_id);
1571
1572                    for (name, user_ids) in display_names {
1573                        let encoded_name = this.encode_key(
1574                            keys::DISPLAY_NAME,
1575                            name.as_normalized_str().unwrap_or_else(|| name.as_raw_str()),
1576                        );
1577                        let data = this.serialize_json(&user_ids)?;
1578
1579                        if user_ids.is_empty() {
1580                            txn.remove_display_name(&room_id, &encoded_name)?;
1581
1582                            // We can't do a migration to merge the previously distinct buckets of
1583                            // user IDs since the display names themselves are hashed before they
1584                            // are persisted in the store. So the store will always retain two
1585                            // buckets: one for raw display names and one for normalised ones.
1586                            //
1587                            // We therefore do the next best thing, which is a sort of a soft
1588                            // migration: we fetch both the raw and normalised buckets, then merge
1589                            // the user IDs contained in them into a separate, temporary merged
1590                            // bucket. The SDK then operates on the merged buckets exclusively. See
1591                            // the comment in `get_users_with_display_names` for details.
1592                            //
1593                            // If the merged bucket is empty, that must mean that both the raw and
1594                            // normalised buckets were also empty, so we can remove both from the
1595                            // store.
1596                            let raw_name = this.encode_key(keys::DISPLAY_NAME, name.as_raw_str());
1597                            txn.remove_display_name(&room_id, &raw_name)?;
1598                        } else {
1599                            // We only create new buckets with the normalized display name.
1600                            txn.set_display_name(&room_id, &encoded_name, &data)?;
1601                        }
1602                    }
1603                }
1604
1605                for (raw_user_id, profile_update) in global_profiles {
1606                    let user_id = this.encode_key(keys::GLOBAL_PROFILES, &raw_user_id);
1607                    match profile_update {
1608                        UserProfileUpdate::Updated(profile_changes) => {
1609                            let existing_data: Option<Vec<u8>> = txn
1610                                .prepare_cached(
1611                                    "SELECT profile_data FROM global_profiles WHERE user_id = ?",
1612                                )?
1613                                .query_one([&user_id], |row| row.get(0))
1614                                .optional()?;
1615
1616                            let mut profile: UserProfile = existing_data
1617                                .map(|data| this.deserialize_json(&data))
1618                                .transpose()?
1619                                .unwrap_or_default();
1620                            profile.apply(profile_changes);
1621
1622                            let serialized = this.serialize_json(&profile)?;
1623                            txn.prepare_cached(
1624                                "INSERT OR REPLACE INTO global_profiles (user_id, profile_data) VALUES (?, ?)",
1625                            )?
1626                            .execute((&user_id, serialized))?;
1627                        }
1628                        // The user left all shared rooms, so drop their stored profile.
1629                        UserProfileUpdate::Dropped => {
1630                            txn.prepare_cached("DELETE FROM global_profiles WHERE user_id = ?")?
1631                                .execute([&user_id])?;
1632                        }
1633                        _ => {
1634                            warn!(%raw_user_id, "Unhandled UserProfileUpdate variant; ignoring");
1635                        }
1636                    }
1637                }
1638
1639                Ok::<_, Error>(())
1640            })
1641            .await?;
1642
1643        Ok(())
1644    }
1645
1646    async fn get_presence_event(&self, user_id: &UserId) -> Result<Option<Raw<PresenceEvent>>> {
1647        self.read()
1648            .await?
1649            .get_kv_blob(self.encode_presence_key(user_id))
1650            .await?
1651            .map(|data| self.deserialize_json(&data))
1652            .transpose()
1653    }
1654
1655    async fn get_presence_events(
1656        &self,
1657        user_ids: &[OwnedUserId],
1658    ) -> Result<Vec<Raw<PresenceEvent>>> {
1659        if user_ids.is_empty() {
1660            return Ok(Vec::new());
1661        }
1662
1663        let user_ids = user_ids.iter().map(|u| self.encode_presence_key(u)).collect();
1664        self.read()
1665            .await?
1666            .get_kv_blobs(user_ids)
1667            .await?
1668            .into_iter()
1669            .map(|data| self.deserialize_json(&data))
1670            .collect()
1671    }
1672
1673    async fn get_state_event(
1674        &self,
1675        room_id: &RoomId,
1676        event_type: StateEventType,
1677        state_key: &str,
1678    ) -> Result<Option<RawAnySyncOrStrippedState>> {
1679        Ok(self
1680            .get_state_events_for_keys(room_id, event_type, &[state_key])
1681            .await?
1682            .into_iter()
1683            .next())
1684    }
1685
1686    async fn get_state_events(
1687        &self,
1688        room_id: &RoomId,
1689        event_type: StateEventType,
1690    ) -> Result<Vec<RawAnySyncOrStrippedState>> {
1691        let room_id = self.encode_key(keys::STATE_EVENT, room_id);
1692        let event_type = self.encode_key(keys::STATE_EVENT, event_type.to_string());
1693        self.read()
1694            .await?
1695            .get_maybe_stripped_state_events(room_id, event_type)
1696            .await?
1697            .into_iter()
1698            .map(|(stripped, data)| {
1699                let ev = if stripped {
1700                    RawAnySyncOrStrippedState::Stripped(self.deserialize_json(&data)?)
1701                } else {
1702                    RawAnySyncOrStrippedState::Sync(self.deserialize_json(&data)?)
1703                };
1704
1705                Ok(ev)
1706            })
1707            .collect()
1708    }
1709
1710    async fn get_state_events_for_keys(
1711        &self,
1712        room_id: &RoomId,
1713        event_type: StateEventType,
1714        state_keys: &[&str],
1715    ) -> Result<Vec<RawAnySyncOrStrippedState>, Self::Error> {
1716        if state_keys.is_empty() {
1717            return Ok(Vec::new());
1718        }
1719
1720        let room_id = self.encode_key(keys::STATE_EVENT, room_id);
1721        let event_type = self.encode_key(keys::STATE_EVENT, event_type.to_string());
1722        let state_keys = state_keys.iter().map(|k| self.encode_key(keys::STATE_EVENT, k)).collect();
1723        self.read()
1724            .await?
1725            .get_maybe_stripped_state_events_for_keys(room_id, event_type, state_keys)
1726            .await?
1727            .into_iter()
1728            .map(|(stripped, data)| {
1729                let ev = if stripped {
1730                    RawAnySyncOrStrippedState::Stripped(self.deserialize_json(&data)?)
1731                } else {
1732                    RawAnySyncOrStrippedState::Sync(self.deserialize_json(&data)?)
1733                };
1734
1735                Ok(ev)
1736            })
1737            .collect()
1738    }
1739
1740    async fn get_profile(
1741        &self,
1742        room_id: &RoomId,
1743        user_id: &UserId,
1744    ) -> Result<Option<MinimalRoomMemberEvent>> {
1745        let room_id = self.encode_key(keys::PROFILE, room_id);
1746        let user_ids = vec![self.encode_key(keys::PROFILE, user_id)];
1747
1748        self.read()
1749            .await?
1750            .get_profiles(room_id, user_ids)
1751            .await?
1752            .into_iter()
1753            .next()
1754            .map(|(_, data)| self.deserialize_json(&data))
1755            .transpose()
1756    }
1757
1758    async fn get_profiles<'a>(
1759        &self,
1760        room_id: &RoomId,
1761        user_ids: &'a [OwnedUserId],
1762    ) -> Result<BTreeMap<&'a UserId, MinimalRoomMemberEvent>> {
1763        if user_ids.is_empty() {
1764            return Ok(BTreeMap::new());
1765        }
1766
1767        let room_id = self.encode_key(keys::PROFILE, room_id);
1768        let mut user_ids_map = user_ids
1769            .iter()
1770            .map(|u| (self.encode_key(keys::PROFILE, u), u.as_ref()))
1771            .collect::<BTreeMap<_, _>>();
1772        let user_ids = user_ids_map.keys().cloned().collect();
1773
1774        self.read()
1775            .await?
1776            .get_profiles(room_id, user_ids)
1777            .await?
1778            .into_iter()
1779            .map(|(user_id, data)| {
1780                Ok((
1781                    user_ids_map
1782                        .remove(user_id.as_slice())
1783                        .expect("returned user IDs were requested"),
1784                    self.deserialize_json(&data)?,
1785                ))
1786            })
1787            .collect()
1788    }
1789
1790    async fn get_user_ids(
1791        &self,
1792        room_id: &RoomId,
1793        membership: RoomMemberships,
1794    ) -> Result<Vec<OwnedUserId>> {
1795        let room_id = self.encode_key(keys::MEMBER, room_id);
1796        let memberships = membership
1797            .as_vec()
1798            .into_iter()
1799            .map(|m| self.encode_key(keys::MEMBER, m.as_str()))
1800            .collect();
1801        self.read()
1802            .await?
1803            .get_user_ids(room_id, memberships)
1804            .await?
1805            .iter()
1806            .map(|data| self.deserialize_value(data))
1807            .collect()
1808    }
1809
1810    async fn get_room_infos(&self, room_load_settings: &RoomLoadSettings) -> Result<Vec<RoomInfo>> {
1811        self.read()
1812            .await?
1813            .get_room_infos(match room_load_settings {
1814                RoomLoadSettings::All => None,
1815                RoomLoadSettings::One(room_id) => Some(self.encode_key(keys::ROOM_INFO, room_id)),
1816            })
1817            .await?
1818            .into_iter()
1819            .map(|data| self.deserialize_json(&data))
1820            .collect()
1821    }
1822
1823    async fn get_users_with_display_name(
1824        &self,
1825        room_id: &RoomId,
1826        display_name: &DisplayName,
1827    ) -> Result<BTreeSet<OwnedUserId>> {
1828        let room_id = self.encode_key(keys::DISPLAY_NAME, room_id);
1829        let names = vec![self.encode_key(
1830            keys::DISPLAY_NAME,
1831            display_name.as_normalized_str().unwrap_or_else(|| display_name.as_raw_str()),
1832        )];
1833
1834        Ok(self
1835            .read()
1836            .await?
1837            .get_display_names(room_id, names)
1838            .await?
1839            .into_iter()
1840            .next()
1841            .map(|(_, data)| self.deserialize_json(&data))
1842            .transpose()?
1843            .unwrap_or_default())
1844    }
1845
1846    async fn get_users_with_display_names<'a>(
1847        &self,
1848        room_id: &RoomId,
1849        display_names: &'a [DisplayName],
1850    ) -> Result<HashMap<&'a DisplayName, BTreeSet<OwnedUserId>>> {
1851        let mut result = HashMap::new();
1852
1853        if display_names.is_empty() {
1854            return Ok(result);
1855        }
1856
1857        let room_id = self.encode_key(keys::DISPLAY_NAME, room_id);
1858        let mut names_map = display_names
1859            .iter()
1860            .flat_map(|display_name| {
1861                // We encode the display name as the `raw_str()` and the normalized string.
1862                //
1863                // This is for compatibility reasons since:
1864                //  1. Previously "Alice" and "alice" were considered to be distinct display
1865                //     names, while we now consider them to be the same so we need to merge the
1866                //     previously distinct buckets of user IDs.
1867                //  2. We can't do a migration to merge the previously distinct buckets of user
1868                //     IDs since the display names itself are hashed before they are persisted
1869                //     in the store.
1870                let raw =
1871                    (self.encode_key(keys::DISPLAY_NAME, display_name.as_raw_str()), display_name);
1872                let normalized = display_name.as_normalized_str().map(|normalized| {
1873                    (self.encode_key(keys::DISPLAY_NAME, normalized), display_name)
1874                });
1875
1876                iter::once(raw).chain(normalized)
1877            })
1878            .collect::<BTreeMap<_, _>>();
1879        let names = names_map.keys().cloned().collect();
1880
1881        for (name, data) in self.read().await?.get_display_names(room_id, names).await?.into_iter()
1882        {
1883            let display_name =
1884                names_map.remove(name.as_slice()).expect("returned display names were requested");
1885            let user_ids: BTreeSet<_> = self.deserialize_json(&data)?;
1886
1887            result.entry(display_name).or_insert_with(BTreeSet::new).extend(user_ids);
1888        }
1889
1890        Ok(result)
1891    }
1892
1893    async fn get_account_data_event(
1894        &self,
1895        event_type: GlobalAccountDataEventType,
1896    ) -> Result<Option<Raw<AnyGlobalAccountDataEvent>>> {
1897        let event_type = self.encode_key(keys::GLOBAL_ACCOUNT_DATA, event_type.to_string());
1898        self.read()
1899            .await?
1900            .get_global_account_data(event_type)
1901            .await?
1902            .map(|value| self.deserialize_json(&value))
1903            .transpose()
1904    }
1905
1906    async fn get_room_account_data_event(
1907        &self,
1908        room_id: &RoomId,
1909        event_type: RoomAccountDataEventType,
1910    ) -> Result<Option<Raw<AnyRoomAccountDataEvent>>> {
1911        let room_id = self.encode_key(keys::ROOM_ACCOUNT_DATA, room_id);
1912        let event_type = self.encode_key(keys::ROOM_ACCOUNT_DATA, event_type.to_string());
1913        self.read()
1914            .await?
1915            .get_room_account_data(room_id, event_type)
1916            .await?
1917            .map(|value| self.deserialize_json(&value))
1918            .transpose()
1919    }
1920
1921    async fn get_user_room_receipt_event(
1922        &self,
1923        room_id: &RoomId,
1924        receipt_type: ReceiptType,
1925        receipt_thread: &ReceiptThread,
1926        user_id: &UserId,
1927    ) -> Result<Option<(OwnedEventId, Receipt)>> {
1928        let room_id = self.encode_key(keys::RECEIPT, room_id);
1929        let receipt_type = self.encode_key(keys::RECEIPT, receipt_type.to_string());
1930        // We cannot have a NULL primary key so we rely on serialization instead of the
1931        // string representation.
1932        let receipt_thread =
1933            self.encode_key(keys::RECEIPT, rmp_serde::to_vec_named(receipt_thread)?);
1934        let user_id = self.encode_key(keys::RECEIPT, user_id);
1935
1936        self.read()
1937            .await?
1938            .get_user_receipt(room_id, receipt_type, receipt_thread, user_id)
1939            .await?
1940            .map(|value| {
1941                self.deserialize_json::<ReceiptData>(&value).map(|d| (d.event_id, d.receipt))
1942            })
1943            .transpose()
1944    }
1945
1946    async fn get_event_room_receipt_events(
1947        &self,
1948        room_id: &RoomId,
1949        receipt_type: ReceiptType,
1950        receipt_thread: &ReceiptThread,
1951        event_id: &EventId,
1952    ) -> Result<Vec<(OwnedUserId, Receipt)>> {
1953        let room_id = self.encode_key(keys::RECEIPT, room_id);
1954        let receipt_type = self.encode_key(keys::RECEIPT, receipt_type.to_string());
1955        // We cannot have a NULL primary key so we rely on serialization instead of the
1956        // string representation.
1957        let receipt_thread =
1958            self.encode_key(keys::RECEIPT, rmp_serde::to_vec_named(receipt_thread)?);
1959        let event_id = self.encode_key(keys::RECEIPT, event_id);
1960
1961        self.read()
1962            .await?
1963            .get_event_receipts(room_id, receipt_type, receipt_thread, event_id)
1964            .await?
1965            .iter()
1966            .map(|value| {
1967                self.deserialize_json::<ReceiptData>(value).map(|d| (d.user_id, d.receipt))
1968            })
1969            .collect()
1970    }
1971
1972    async fn get_custom_value(&self, key: &[u8]) -> Result<Option<Vec<u8>>> {
1973        self.read().await?.get_kv_blob(self.encode_custom_key(key)).await
1974    }
1975
1976    async fn set_custom_value_no_read(&self, key: &[u8], value: Vec<u8>) -> Result<()> {
1977        let conn = self.write().await?;
1978        let key = self.encode_custom_key(key);
1979        conn.set_kv_blob(key, value).await?;
1980        Ok(())
1981    }
1982
1983    async fn set_custom_value(&self, key: &[u8], value: Vec<u8>) -> Result<Option<Vec<u8>>> {
1984        let conn = self.write().await?;
1985        let key = self.encode_custom_key(key);
1986        let previous = conn.get_kv_blob(key.clone()).await?;
1987        conn.set_kv_blob(key, value).await?;
1988        Ok(previous)
1989    }
1990
1991    async fn remove_custom_value(&self, key: &[u8]) -> Result<Option<Vec<u8>>> {
1992        let conn = self.write().await?;
1993        let key = self.encode_custom_key(key);
1994        let previous = conn.get_kv_blob(key.clone()).await?;
1995        if previous.is_some() {
1996            conn.delete_kv_blob(key).await?;
1997        }
1998        Ok(previous)
1999    }
2000
2001    async fn remove_room(&self, room_id: &RoomId) -> Result<()> {
2002        let this = self.clone();
2003        let room_id = room_id.to_owned();
2004
2005        let conn = self.write().await?;
2006
2007        conn.with_transaction(move |txn| -> Result<()> {
2008            let room_info_room_id = this.encode_key(keys::ROOM_INFO, &room_id);
2009            txn.remove_room_info(&room_info_room_id)?;
2010
2011            let state_event_room_id = this.encode_key(keys::STATE_EVENT, &room_id);
2012            txn.remove_room_state_events(&state_event_room_id, None)?;
2013
2014            let member_room_id = this.encode_key(keys::MEMBER, &room_id);
2015            txn.remove_room_members(&member_room_id, None)?;
2016
2017            let profile_room_id = this.encode_key(keys::PROFILE, &room_id);
2018            txn.remove_room_profiles(&profile_room_id)?;
2019
2020            let room_account_data_room_id = this.encode_key(keys::ROOM_ACCOUNT_DATA, &room_id);
2021            txn.remove_room_account_data(&room_account_data_room_id)?;
2022
2023            let receipt_room_id = this.encode_key(keys::RECEIPT, &room_id);
2024            txn.remove_room_receipts(&receipt_room_id)?;
2025
2026            let display_name_room_id = this.encode_key(keys::DISPLAY_NAME, &room_id);
2027            txn.remove_room_display_names(&display_name_room_id)?;
2028
2029            let send_queue_room_id = this.encode_key(keys::SEND_QUEUE, &room_id);
2030            txn.remove_room_send_queue(&send_queue_room_id)?;
2031
2032            let dependent_send_queue_room_id =
2033                this.encode_key(keys::DEPENDENTS_SEND_QUEUE, &room_id);
2034            txn.remove_room_dependent_send_queue(&dependent_send_queue_room_id)?;
2035
2036            let thread_subscriptions_room_id =
2037                this.encode_key(keys::THREAD_SUBSCRIPTIONS, &room_id);
2038            txn.execute(
2039                "DELETE FROM thread_subscriptions WHERE room_id = ?",
2040                (thread_subscriptions_room_id,),
2041            )?;
2042
2043            Ok(())
2044        })
2045        .await?;
2046
2047        conn.vacuum().await
2048    }
2049
2050    async fn save_send_queue_request(
2051        &self,
2052        room_id: &RoomId,
2053        transaction_id: OwnedTransactionId,
2054        created_at: MilliSecondsSinceUnixEpoch,
2055        content: QueuedRequestKind,
2056        priority: usize,
2057    ) -> Result<(), Self::Error> {
2058        let room_id_key = self.encode_key(keys::SEND_QUEUE, room_id);
2059        let room_id_value = self.serialize_value(&room_id.to_owned())?;
2060
2061        let content = self.serialize_json(&content)?;
2062        // The transaction id is used both as a key (in remove/update) and a value (as
2063        // it's useful for the callers), so we keep it as is, and neither hash
2064        // it (with encode_key) or encrypt it (through serialize_value). After
2065        // all, it carries no personal information, so this is considered fine.
2066
2067        let created_at_ts: u64 = created_at.0.into();
2068        self.write()
2069            .await?
2070            .with_transaction(move |txn| {
2071                txn.prepare_cached("INSERT INTO send_queue_events (room_id, room_id_val, transaction_id, content, priority, created_at) VALUES (?, ?, ?, ?, ?, ?)")?.execute((room_id_key, room_id_value, transaction_id.to_string(), content, priority, created_at_ts))?;
2072                Ok(())
2073            })
2074            .await
2075    }
2076
2077    async fn update_send_queue_request(
2078        &self,
2079        room_id: &RoomId,
2080        transaction_id: &TransactionId,
2081        content: QueuedRequestKind,
2082    ) -> Result<bool, Self::Error> {
2083        let room_id = self.encode_key(keys::SEND_QUEUE, room_id);
2084
2085        let content = self.serialize_json(&content)?;
2086        // See comment in [`Self::save_send_queue_request`] to understand why the
2087        // transaction id is neither encrypted or hashed.
2088        let transaction_id = transaction_id.to_string();
2089
2090        let num_updated = self.write()
2091            .await?
2092            .with_transaction(move |txn| {
2093                txn.prepare_cached("UPDATE send_queue_events SET wedge_reason = NULL, content = ? WHERE room_id = ? AND transaction_id = ?")?.execute((content, room_id, transaction_id))
2094            })
2095            .await?;
2096
2097        Ok(num_updated > 0)
2098    }
2099
2100    async fn remove_send_queue_request(
2101        &self,
2102        room_id: &RoomId,
2103        transaction_id: &TransactionId,
2104    ) -> Result<bool, Self::Error> {
2105        let room_id = self.encode_key(keys::SEND_QUEUE, room_id);
2106
2107        // See comment in `save_send_queue_request`.
2108        let transaction_id = transaction_id.to_string();
2109
2110        let num_deleted = self
2111            .write()
2112            .await?
2113            .with_transaction(move |txn| {
2114                txn.prepare_cached(
2115                    "DELETE FROM send_queue_events WHERE room_id = ? AND transaction_id = ?",
2116                )?
2117                .execute((room_id, &transaction_id))
2118            })
2119            .await?;
2120
2121        Ok(num_deleted > 0)
2122    }
2123
2124    async fn load_send_queue_requests(
2125        &self,
2126        room_id: &RoomId,
2127    ) -> Result<Vec<QueuedRequest>, Self::Error> {
2128        let room_id = self.encode_key(keys::SEND_QUEUE, room_id);
2129
2130        // Note: ROWID is always present and is an auto-incremented integer counter. We
2131        // want to maintain the insertion order, so we can sort using it.
2132        // Note 2: transaction_id is not encoded, see why in `save_send_queue_request`.
2133        let res: Vec<(String, Vec<u8>, Option<Vec<u8>>, usize, Option<u64>)> = self
2134            .read()
2135            .await?
2136            .prepare(
2137                "SELECT transaction_id, content, wedge_reason, priority, created_at FROM send_queue_events WHERE room_id = ? ORDER BY priority DESC, ROWID",
2138                |mut stmt| {
2139                    stmt.query((room_id,))?
2140                        .mapped(|row| Ok((row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?, row.get(4)?)))
2141                        .collect()
2142                },
2143            )
2144            .await?;
2145
2146        let mut requests = Vec::with_capacity(res.len());
2147
2148        for entry in res {
2149            let created_at = entry
2150                .4
2151                .and_then(UInt::new)
2152                .map_or_else(MilliSecondsSinceUnixEpoch::now, MilliSecondsSinceUnixEpoch);
2153
2154            requests.push(QueuedRequest {
2155                transaction_id: entry.0.into(),
2156                kind: self.deserialize_json(&entry.1)?,
2157                error: entry.2.map(|v| self.deserialize_value(&v)).transpose()?,
2158                priority: entry.3,
2159                created_at,
2160            });
2161        }
2162
2163        Ok(requests)
2164    }
2165
2166    async fn update_send_queue_request_status(
2167        &self,
2168        room_id: &RoomId,
2169        transaction_id: &TransactionId,
2170        error: Option<QueueWedgeError>,
2171    ) -> Result<(), Self::Error> {
2172        let room_id = self.encode_key(keys::SEND_QUEUE, room_id);
2173
2174        // See comment in `save_send_queue_request`.
2175        let transaction_id = transaction_id.to_string();
2176
2177        // Serialize the error to json bytes (encrypted if option is enabled) if set.
2178        let error_value = error.map(|e| self.serialize_value(&e)).transpose()?;
2179
2180        self.write()
2181            .await?
2182            .with_transaction(move |txn| {
2183                txn.prepare_cached("UPDATE send_queue_events SET wedge_reason = ? WHERE room_id = ? AND transaction_id = ?")?.execute((error_value, room_id, transaction_id))?;
2184                Ok(())
2185            })
2186            .await
2187    }
2188
2189    async fn load_rooms_with_unsent_requests(&self) -> Result<Vec<OwnedRoomId>, Self::Error> {
2190        // If the values were not encrypted, we could use `SELECT DISTINCT` here, but we
2191        // have to manually do the deduplication: indeed, for all X, encrypt(X)
2192        // != encrypted(X), since we use a nonce in the encryption process.
2193
2194        let res: Vec<Vec<u8>> = self
2195            .read()
2196            .await?
2197            .prepare("SELECT room_id_val FROM send_queue_events", |mut stmt| {
2198                stmt.query(())?.mapped(|row| row.get(0)).collect()
2199            })
2200            .await?;
2201
2202        // So we collect the results into a `BTreeSet` to perform the deduplication, and
2203        // then rejigger that into a vector.
2204        Ok(res
2205            .into_iter()
2206            .map(|entry| self.deserialize_value(&entry))
2207            .collect::<Result<BTreeSet<OwnedRoomId>, _>>()?
2208            .into_iter()
2209            .collect())
2210    }
2211
2212    async fn save_dependent_queued_request(
2213        &self,
2214        room_id: &RoomId,
2215        parent_txn_id: &TransactionId,
2216        own_txn_id: ChildTransactionId,
2217        created_at: MilliSecondsSinceUnixEpoch,
2218        content: DependentQueuedRequestKind,
2219    ) -> Result<()> {
2220        let room_id = self.encode_key(keys::DEPENDENTS_SEND_QUEUE, room_id);
2221        let content = self.serialize_json(&content)?;
2222
2223        // See comment in `save_send_queue_request`.
2224        let parent_txn_id = parent_txn_id.to_string();
2225        let own_txn_id = own_txn_id.to_string();
2226
2227        let created_at_ts: u64 = created_at.0.into();
2228        self.write()
2229            .await?
2230            .with_transaction(move |txn| {
2231                txn.prepare_cached(
2232                    r#"INSERT INTO dependent_send_queue_events
2233                         (room_id, parent_transaction_id, own_transaction_id, content, created_at)
2234                       VALUES (?, ?, ?, ?, ?)"#,
2235                )?
2236                .execute((
2237                    room_id,
2238                    parent_txn_id,
2239                    own_txn_id,
2240                    content,
2241                    created_at_ts,
2242                ))?;
2243                Ok(())
2244            })
2245            .await
2246    }
2247
2248    async fn update_dependent_queued_request(
2249        &self,
2250        room_id: &RoomId,
2251        own_transaction_id: &ChildTransactionId,
2252        new_content: DependentQueuedRequestKind,
2253    ) -> Result<bool> {
2254        let room_id = self.encode_key(keys::DEPENDENTS_SEND_QUEUE, room_id);
2255        let content = self.serialize_json(&new_content)?;
2256
2257        // See comment in `save_send_queue_request`.
2258        let own_txn_id = own_transaction_id.to_string();
2259
2260        let num_updated = self
2261            .write()
2262            .await?
2263            .with_transaction(move |txn| {
2264                txn.prepare_cached(
2265                    r#"UPDATE dependent_send_queue_events
2266                       SET content = ?
2267                       WHERE own_transaction_id = ?
2268                       AND room_id = ?"#,
2269                )?
2270                .execute((content, own_txn_id, room_id))
2271            })
2272            .await?;
2273
2274        if num_updated > 1 {
2275            return Err(Error::InconsistentUpdate);
2276        }
2277
2278        Ok(num_updated == 1)
2279    }
2280
2281    async fn mark_dependent_queued_requests_as_ready(
2282        &self,
2283        room_id: &RoomId,
2284        parent_txn_id: &TransactionId,
2285        parent_key: SentRequestKey,
2286    ) -> Result<usize> {
2287        let room_id = self.encode_key(keys::DEPENDENTS_SEND_QUEUE, room_id);
2288        let parent_key = self.serialize_json(&parent_key)?;
2289
2290        // See comment in `save_send_queue_request`.
2291        let parent_txn_id = parent_txn_id.to_string();
2292
2293        self.write()
2294            .await?
2295            .with_transaction(move |txn| {
2296                Ok(txn.prepare_cached(
2297                    "UPDATE dependent_send_queue_events SET parent_key = ? WHERE parent_transaction_id = ? and room_id = ?",
2298                )?
2299                .execute((parent_key, parent_txn_id, room_id))?)
2300            })
2301            .await
2302    }
2303
2304    async fn remove_dependent_queued_request(
2305        &self,
2306        room_id: &RoomId,
2307        txn_id: &ChildTransactionId,
2308    ) -> Result<bool> {
2309        let room_id = self.encode_key(keys::DEPENDENTS_SEND_QUEUE, room_id);
2310
2311        // See comment in `save_send_queue_request`.
2312        let txn_id = txn_id.to_string();
2313
2314        let num_deleted = self
2315            .write()
2316            .await?
2317            .with_transaction(move |txn| {
2318                txn.prepare_cached(
2319                    "DELETE FROM dependent_send_queue_events WHERE own_transaction_id = ? AND room_id = ?",
2320                )?
2321                .execute((txn_id, room_id))
2322            })
2323            .await?;
2324
2325        Ok(num_deleted > 0)
2326    }
2327
2328    async fn load_dependent_queued_requests(
2329        &self,
2330        room_id: &RoomId,
2331    ) -> Result<Vec<DependentQueuedRequest>> {
2332        let room_id = self.encode_key(keys::DEPENDENTS_SEND_QUEUE, room_id);
2333
2334        // Note: transaction_id is not encoded, see why in `save_send_queue_request`.
2335        let res: Vec<(String, String, Option<Vec<u8>>, Vec<u8>, Option<u64>)> = self
2336            .read()
2337            .await?
2338            .prepare(
2339                "SELECT own_transaction_id, parent_transaction_id, parent_key, content, created_at FROM dependent_send_queue_events WHERE room_id = ? ORDER BY ROWID",
2340                |mut stmt| {
2341                    stmt.query((room_id,))?
2342                        .mapped(|row| Ok((row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?, row.get(4)?)))
2343                        .collect()
2344                },
2345            )
2346            .await?;
2347
2348        let mut dependent_events = Vec::with_capacity(res.len());
2349
2350        for entry in res {
2351            let created_at = entry
2352                .4
2353                .and_then(UInt::new)
2354                .map_or_else(MilliSecondsSinceUnixEpoch::now, MilliSecondsSinceUnixEpoch);
2355
2356            dependent_events.push(DependentQueuedRequest {
2357                own_transaction_id: entry.0.into(),
2358                parent_transaction_id: entry.1.into(),
2359                parent_key: entry.2.map(|json| self.deserialize_json(&json)).transpose()?,
2360                kind: self.deserialize_json(&entry.3)?,
2361                created_at,
2362            });
2363        }
2364
2365        Ok(dependent_events)
2366    }
2367
2368    async fn upsert_thread_subscriptions(
2369        &self,
2370        updates: Vec<(&RoomId, &EventId, StoredThreadSubscription)>,
2371    ) -> Result<(), Self::Error> {
2372        let values: Vec<_> = updates
2373            .into_iter()
2374            .map(|(room_id, thread_id, subscription)| {
2375                (
2376                    self.encode_key(keys::THREAD_SUBSCRIPTIONS, room_id),
2377                    self.encode_key(keys::THREAD_SUBSCRIPTIONS, thread_id),
2378                    subscription.status.as_str(),
2379                    subscription.bump_stamp,
2380                )
2381            })
2382            .collect();
2383
2384        self.write()
2385            .await?
2386            .with_transaction(move |txn| {
2387                let mut txn = txn.prepare_cached(
2388                    "INSERT INTO thread_subscriptions (room_id, event_id, status, bump_stamp)
2389                    VALUES (?, ?, ?, ?)
2390                    ON CONFLICT (room_id, event_id) DO UPDATE
2391                    SET
2392                        status =
2393                            CASE
2394                                WHEN thread_subscriptions.bump_stamp IS NULL THEN EXCLUDED.status
2395                                WHEN EXCLUDED.bump_stamp IS NULL THEN EXCLUDED.status
2396                                WHEN thread_subscriptions.bump_stamp < EXCLUDED.bump_stamp THEN EXCLUDED.status
2397                                ELSE thread_subscriptions.status
2398                            END,
2399                        bump_stamp =
2400                            CASE
2401                                WHEN thread_subscriptions.bump_stamp IS NULL THEN EXCLUDED.bump_stamp
2402                                WHEN EXCLUDED.bump_stamp IS NULL THEN thread_subscriptions.bump_stamp
2403                                WHEN thread_subscriptions.bump_stamp < EXCLUDED.bump_stamp THEN EXCLUDED.bump_stamp
2404                                ELSE thread_subscriptions.bump_stamp
2405                            END",
2406                )?;
2407
2408                for value in values {
2409                    txn.execute(value)?;
2410                }
2411
2412                Result::<_, Error>::Ok(())
2413            })
2414            .await?;
2415
2416        Ok(())
2417    }
2418
2419    async fn load_thread_subscription(
2420        &self,
2421        room_id: &RoomId,
2422        thread_id: &EventId,
2423    ) -> Result<Option<StoredThreadSubscription>, Self::Error> {
2424        let room_id = self.encode_key(keys::THREAD_SUBSCRIPTIONS, room_id);
2425        let thread_id = self.encode_key(keys::THREAD_SUBSCRIPTIONS, thread_id);
2426
2427        Ok(self
2428            .read()
2429            .await?
2430            .query_one(
2431                "SELECT status, bump_stamp FROM thread_subscriptions WHERE room_id = ? AND event_id = ?",
2432                (room_id, thread_id),
2433                |row| Ok((row.get::<_, String>(0)?, row.get::<_, Option<u64>>(1)?))
2434            )
2435            .await
2436            .optional()?
2437            .map(|(status, bump_stamp)| -> Result<_, Self::Error> {
2438                let status = ThreadSubscriptionStatus::from_str(&status).map_err(|_| {
2439                    Error::InvalidData { details: format!("Invalid thread status: {status}") }
2440                })?;
2441                Ok(StoredThreadSubscription { status, bump_stamp })
2442            })
2443            .transpose()?)
2444    }
2445
2446    async fn remove_thread_subscription(
2447        &self,
2448        room_id: &RoomId,
2449        thread_id: &EventId,
2450    ) -> Result<(), Self::Error> {
2451        let room_id = self.encode_key(keys::THREAD_SUBSCRIPTIONS, room_id);
2452        let thread_id = self.encode_key(keys::THREAD_SUBSCRIPTIONS, thread_id);
2453
2454        self.write()
2455            .await?
2456            .execute(
2457                "DELETE FROM thread_subscriptions WHERE room_id = ? AND event_id = ?",
2458                (room_id, thread_id),
2459            )
2460            .await?;
2461
2462        Ok(())
2463    }
2464
2465    async fn get_global_profile(
2466        &self,
2467        user_id: &UserId,
2468    ) -> Result<Option<UserProfile>, Self::Error> {
2469        self.read()
2470            .await?
2471            .get_global_profiles(vec![self.encode_key(keys::GLOBAL_PROFILES, user_id)])
2472            .await?
2473            .into_iter()
2474            .next()
2475            .map(|(_, data)| self.deserialize_json(&data))
2476            .transpose()
2477    }
2478
2479    async fn get_global_profiles<'a>(
2480        &self,
2481        user_ids: &'a [OwnedUserId],
2482    ) -> Result<BTreeMap<&'a UserId, UserProfile>, Self::Error> {
2483        if user_ids.is_empty() {
2484            return Ok(BTreeMap::new());
2485        }
2486
2487        let mut user_ids_map = user_ids
2488            .iter()
2489            .map(|u| (self.encode_key(keys::GLOBAL_PROFILES, u), u.as_ref()))
2490            .collect::<BTreeMap<_, _>>();
2491        let user_ids = user_ids_map.keys().cloned().collect();
2492
2493        self.read()
2494            .await?
2495            .get_global_profiles(user_ids)
2496            .await?
2497            .into_iter()
2498            .map(|(user_id, data)| {
2499                Ok((
2500                    user_ids_map
2501                        .remove(user_id.as_slice())
2502                        .expect("returned user IDs were requested"),
2503                    self.deserialize_json(&data)?,
2504                ))
2505            })
2506            .collect()
2507    }
2508
2509    async fn optimize(&self) -> Result<(), Self::Error> {
2510        Ok(self.vacuum().await?)
2511    }
2512
2513    async fn get_size(&self) -> Result<Option<usize>, Self::Error> {
2514        self.get_db_size().await
2515    }
2516
2517    async fn close(&self) -> Result<(), Self::Error> {
2518        connection::close_connections(&self.connections, "State store").await;
2519        Ok(())
2520    }
2521
2522    async fn reopen(&self) -> Result<(), Self::Error> {
2523        connection::reopen_connections(
2524            &self.connections,
2525            self.db_path.clone(),
2526            self.pool_config,
2527            self.runtime_config,
2528        )
2529        .await?;
2530        Ok(())
2531    }
2532}
2533
2534#[derive(Debug, Clone, Serialize, Deserialize)]
2535struct ReceiptData {
2536    receipt: Receipt,
2537    event_id: OwnedEventId,
2538    user_id: OwnedUserId,
2539}
2540
2541#[cfg(test)]
2542mod tests {
2543    use std::sync::{
2544        LazyLock,
2545        atomic::{AtomicU32, Ordering::SeqCst},
2546    };
2547
2548    use matrix_sdk_base::{StateStore, StoreError, statestore_integration_tests};
2549    use tempfile::{TempDir, tempdir};
2550
2551    use super::SqliteStateStore;
2552
2553    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2554    static NUM: AtomicU32 = AtomicU32::new(0);
2555
2556    async fn get_store() -> Result<impl StateStore, StoreError> {
2557        let name = NUM.fetch_add(1, SeqCst).to_string();
2558        let tmpdir_path = TMP_DIR.path().join(name);
2559
2560        tracing::info!("using store @ {}", tmpdir_path.to_str().unwrap());
2561
2562        Ok(SqliteStateStore::open(tmpdir_path.to_str().unwrap(), None).await.unwrap())
2563    }
2564
2565    statestore_integration_tests!();
2566}
2567
2568#[cfg(test)]
2569mod encrypted_tests {
2570    use std::{
2571        path::PathBuf,
2572        sync::{
2573            LazyLock,
2574            atomic::{AtomicU32, Ordering::SeqCst},
2575        },
2576    };
2577
2578    use base64::Engine as _;
2579    use matrix_sdk_base::{StateStore, StoreError, statestore_integration_tests};
2580    use matrix_sdk_test::async_test;
2581    use tempfile::{TempDir, tempdir};
2582
2583    use super::SqliteStateStore;
2584    use crate::{Base64Variant, SqliteStoreConfig, utils::SqliteAsyncConnExt};
2585
2586    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2587    static NUM: AtomicU32 = AtomicU32::new(0);
2588
2589    fn new_state_store_workspace() -> PathBuf {
2590        let name = NUM.fetch_add(1, SeqCst).to_string();
2591        TMP_DIR.path().join(name)
2592    }
2593
2594    async fn get_store() -> Result<impl StateStore, StoreError> {
2595        let tmpdir_path = new_state_store_workspace();
2596
2597        tracing::info!("using store @ {}", tmpdir_path.to_str().unwrap());
2598
2599        Ok(SqliteStateStore::open(tmpdir_path.to_str().unwrap(), Some("default_test_password"))
2600            .await
2601            .unwrap())
2602    }
2603
2604    /// The two passphrase methods are interchangeable in both directions.
2605    #[async_test]
2606    async fn test_high_entropy_passphrase_migrates_a_passphrase_store() {
2607        const KEY: &[u8; 32] = b"a randomly generated passphrase ";
2608        let tmpdir_path = new_state_store_workspace();
2609
2610        let passphrase = base64::prelude::BASE64_STANDARD.encode(KEY);
2611
2612        let config = SqliteStoreConfig::new(&tmpdir_path).passphrase(Some(&passphrase));
2613        drop(SqliteStateStore::open_with_config(&config).await.unwrap());
2614
2615        // Migrates and caches the copy...
2616        let config = SqliteStoreConfig::new(&tmpdir_path)
2617            .high_entropy_passphrase(Some(KEY), Base64Variant::Padded);
2618        drop(SqliteStateStore::open_with_config(&config).await.unwrap());
2619
2620        // ...which the next open uses.
2621        drop(SqliteStateStore::open_with_config(&config).await.unwrap());
2622
2623        // The `cipher` entry was replaced, so the old passphrase can't work anymore.
2624        let config = SqliteStoreConfig::new(&tmpdir_path).passphrase(Some(&passphrase));
2625        drop(
2626            SqliteStateStore::open_with_config(&config)
2627                .await
2628                .expect_err("The old passphrase-only method shouldn't work anymore"),
2629        );
2630
2631        // The `cipher` entry was replaced, so now only high entropy or key work.
2632        let config = SqliteStoreConfig::new(&tmpdir_path)
2633            .high_entropy_passphrase(Some(KEY), Base64Variant::Padded);
2634        drop(
2635            SqliteStateStore::open_with_config(&config)
2636                .await
2637                .expect("The high-entropy method should continue to work"),
2638        );
2639
2640        let config = SqliteStoreConfig::new(&tmpdir_path).key(Some(KEY));
2641        drop(
2642            SqliteStateStore::open_with_config(&config).await.expect("The key should work as well"),
2643        );
2644
2645        let config = SqliteStoreConfig::new(&tmpdir_path).high_entropy_passphrase(
2646            Some(b"wrong passphrase can't work 1234"),
2647            Base64Variant::Padded,
2648        );
2649        assert!(SqliteStateStore::open_with_config(&config).await.is_err());
2650        let config = SqliteStoreConfig::new(&tmpdir_path).passphrase(Some("wrong"));
2651        assert!(SqliteStateStore::open_with_config(&config).await.is_err());
2652    }
2653
2654    #[async_test]
2655    async fn test_pool_size() {
2656        let tmpdir_path = new_state_store_workspace();
2657        let store_open_config = SqliteStoreConfig::new(tmpdir_path).pool_max_size(42);
2658
2659        let store = SqliteStateStore::open_with_config(&store_open_config).await.unwrap();
2660
2661        let guard = store.connections.lock().await;
2662        assert_eq!(guard.as_ref().unwrap().pool.status().max_size, 42);
2663    }
2664
2665    #[async_test]
2666    async fn test_cache_size() {
2667        let tmpdir_path = new_state_store_workspace();
2668        let store_open_config = SqliteStoreConfig::new(tmpdir_path).cache_size(1500);
2669
2670        let store = SqliteStateStore::open_with_config(&store_open_config).await.unwrap();
2671
2672        let conn = store.read().await.unwrap();
2673        let cache_size =
2674            conn.query_row("PRAGMA cache_size", (), |row| row.get::<_, i32>(0)).await.unwrap();
2675
2676        // The value passed to `SqliteStoreConfig` is in bytes. Check it is
2677        // converted to kibibytes. Also, it must be a negative value because it
2678        // _is_ the size in kibibytes, not in page size.
2679        assert_eq!(cache_size, -(1500 / 1024));
2680    }
2681
2682    #[async_test]
2683    async fn test_journal_size_limit() {
2684        let tmpdir_path = new_state_store_workspace();
2685        let store_open_config = SqliteStoreConfig::new(tmpdir_path).journal_size_limit(1500);
2686
2687        let store = SqliteStateStore::open_with_config(&store_open_config).await.unwrap();
2688
2689        let conn = store.read().await.unwrap();
2690        let journal_size_limit = conn
2691            .query_row("PRAGMA journal_size_limit", (), |row| row.get::<_, u32>(0))
2692            .await
2693            .unwrap();
2694
2695        // The value passed to `SqliteStoreConfig` is in bytes. It stays in
2696        // bytes in SQLite.
2697        assert_eq!(journal_size_limit, 1500);
2698    }
2699
2700    statestore_integration_tests!();
2701}
2702
2703#[cfg(test)]
2704mod migration_tests {
2705    use std::{
2706        path::{Path, PathBuf},
2707        sync::{
2708            Arc, LazyLock,
2709            atomic::{AtomicU32, Ordering::SeqCst},
2710        },
2711    };
2712
2713    use as_variant::as_variant;
2714    use matrix_sdk_base::{
2715        RoomState, StateStore,
2716        media::{MediaFormat, MediaRequestParameters},
2717        store::{
2718            ChildTransactionId, DependentQueuedRequestKind, RoomLoadSettings,
2719            SerializableEventContent,
2720        },
2721        sync::UnreadNotificationsCount,
2722    };
2723    use matrix_sdk_test::async_test;
2724    use ruma::{
2725        EventId, MilliSecondsSinceUnixEpoch, OwnedTransactionId, RoomId, TransactionId, UserId,
2726        events::{
2727            StateEventType,
2728            room::{MediaSource, create::RoomCreateEventContent, message::RoomMessageEventContent},
2729        },
2730        room_id, server_name, user_id,
2731    };
2732    use rusqlite::Transaction;
2733    use serde::{Deserialize, Serialize};
2734    use serde_json::json;
2735    use tempfile::{TempDir, tempdir};
2736    use tokio::{fs, sync::Mutex};
2737    use zeroize::Zeroizing;
2738
2739    use super::{DATABASE_NAME, SqliteStateStore, init, keys};
2740    use crate::{
2741        OpenStoreError, Secret, SqliteStoreConfig, connection,
2742        error::{Error, Result},
2743        utils::{EncryptableStore as _, SqliteAsyncConnExt, SqliteKeyValueStoreAsyncConnExt},
2744    };
2745
2746    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2747    static NUM: AtomicU32 = AtomicU32::new(0);
2748    const SECRET: &str = "secret";
2749
2750    fn new_path() -> PathBuf {
2751        let name = NUM.fetch_add(1, SeqCst).to_string();
2752        TMP_DIR.path().join(name)
2753    }
2754
2755    async fn create_fake_db(path: &Path, version: u8) -> Result<SqliteStateStore> {
2756        let config = SqliteStoreConfig::new(path);
2757
2758        fs::create_dir_all(&config.path).await.map_err(OpenStoreError::CreateDir).unwrap();
2759
2760        let pool = config.build_pool_of_connections(DATABASE_NAME).unwrap();
2761        let db_path = pool.manager().database_path.clone();
2762        let conn = pool.get().await?;
2763
2764        init(&conn).await?;
2765
2766        let store_cipher = Some(Arc::new(
2767            conn.get_or_create_store_cipher(Secret::PassPhrase(Zeroizing::new(SECRET.to_owned())))
2768                .await
2769                .unwrap(),
2770        ));
2771        let this = SqliteStateStore {
2772            store_cipher,
2773            connections: Arc::new(Mutex::new(Some(connection::SqliteConnections {
2774                pool,
2775                write_connection: Arc::new(Mutex::new(conn)),
2776            }))),
2777            db_path,
2778            pool_config: deadpool::managed::PoolConfig::default(),
2779            runtime_config: crate::RuntimeConfig::default(),
2780        };
2781        this.run_migrations(1, Some(version)).await?;
2782
2783        Ok(this)
2784    }
2785
2786    fn room_info_v1_json(
2787        room_id: &RoomId,
2788        state: RoomState,
2789        name: Option<&str>,
2790        creator: Option<&UserId>,
2791    ) -> serde_json::Value {
2792        // Test with name set or not.
2793        let name_content = match name {
2794            Some(name) => json!({ "name": name }),
2795            None => json!({ "name": null }),
2796        };
2797        // Test with creator set or not.
2798        let create_content = match creator {
2799            Some(creator) => RoomCreateEventContent::new_v1(creator.to_owned()),
2800            None => RoomCreateEventContent::new_v11(),
2801        };
2802
2803        json!({
2804            "room_id": room_id,
2805            "room_type": state,
2806            "notification_counts": UnreadNotificationsCount::default(),
2807            "summary": {
2808                "heroes": [],
2809                "joined_member_count": 0,
2810                "invited_member_count": 0,
2811            },
2812            "members_synced": false,
2813            "base_info": {
2814                "dm_targets": [],
2815                "max_power_level": 100,
2816                "name": {
2817                    "Original": {
2818                        "content": name_content,
2819                    },
2820                },
2821                "create": {
2822                    "Original": {
2823                        "content": create_content,
2824                    }
2825                }
2826            },
2827        })
2828    }
2829
2830    #[async_test]
2831    pub async fn test_migrating_v1_to_v2() {
2832        let path = new_path();
2833        // Create and populate db.
2834        {
2835            let db = create_fake_db(&path, 1).await.unwrap();
2836            let conn = db.read().await.unwrap();
2837
2838            let this = db.clone();
2839            conn.with_transaction(move |txn| {
2840                for i in 0..5 {
2841                    let room_id = RoomId::parse(format!("!room_{i}:localhost")).unwrap();
2842                    let (state, stripped) =
2843                        if i < 3 { (RoomState::Joined, false) } else { (RoomState::Invited, true) };
2844                    let info = room_info_v1_json(&room_id, state, None, None);
2845
2846                    let room_id = this.encode_key(keys::ROOM_INFO, room_id);
2847                    let data = this.serialize_json(&info)?;
2848
2849                    txn.prepare_cached(
2850                        "INSERT INTO room_info (room_id, stripped, data)
2851                         VALUES (?, ?, ?)",
2852                    )?
2853                    .execute((room_id, stripped, data))?;
2854                }
2855
2856                Result::<_, Error>::Ok(())
2857            })
2858            .await
2859            .unwrap();
2860        }
2861
2862        // This transparently migrates to the latest version.
2863        let store = SqliteStateStore::open(path, Some(SECRET)).await.unwrap();
2864
2865        // Check all room infos are there.
2866        assert_eq!(store.get_room_infos(&RoomLoadSettings::default()).await.unwrap().len(), 5);
2867    }
2868
2869    // Add a room in version 2 format of the state store.
2870    fn add_room_v2(
2871        this: &SqliteStateStore,
2872        txn: &Transaction<'_>,
2873        room_id: &RoomId,
2874        name: Option<&str>,
2875        create_creator: Option<&UserId>,
2876        create_sender: Option<&UserId>,
2877    ) -> Result<(), Error> {
2878        let room_info_json = room_info_v1_json(room_id, RoomState::Joined, name, create_creator);
2879
2880        let encoded_room_id = this.encode_key(keys::ROOM_INFO, room_id);
2881        let encoded_state =
2882            this.encode_key(keys::ROOM_INFO, serde_json::to_string(&RoomState::Joined)?);
2883        let data = this.serialize_json(&room_info_json)?;
2884
2885        txn.prepare_cached(
2886            "INSERT INTO room_info (room_id, state, data)
2887             VALUES (?, ?, ?)",
2888        )?
2889        .execute((encoded_room_id, encoded_state, data))?;
2890
2891        // Test with or without `m.room.create` event in the room state.
2892        let Some(create_sender) = create_sender else {
2893            return Ok(());
2894        };
2895
2896        let create_content = match create_creator {
2897            Some(creator) => RoomCreateEventContent::new_v1(creator.to_owned()),
2898            None => RoomCreateEventContent::new_v11(),
2899        };
2900
2901        let event_id = EventId::new_v1(server_name!("dummy.local"));
2902        let create_event = json!({
2903            "content": create_content,
2904            "event_id": event_id,
2905            "sender": create_sender.to_owned(),
2906            "origin_server_ts": MilliSecondsSinceUnixEpoch::now(),
2907            "state_key": "",
2908            "type": "m.room.create",
2909            "unsigned": {},
2910        });
2911
2912        let encoded_room_id = this.encode_key(keys::STATE_EVENT, room_id);
2913        let encoded_event_type =
2914            this.encode_key(keys::STATE_EVENT, StateEventType::RoomCreate.to_string());
2915        let encoded_state_key = this.encode_key(keys::STATE_EVENT, "");
2916        let stripped = false;
2917        let encoded_event_id = this.encode_key(keys::STATE_EVENT, event_id);
2918        let data = this.serialize_json(&create_event)?;
2919
2920        txn.prepare_cached(
2921            "INSERT
2922             INTO state_event (room_id, event_type, state_key, stripped, event_id, data)
2923             VALUES (?, ?, ?, ?, ?, ?)",
2924        )?
2925        .execute((
2926            encoded_room_id,
2927            encoded_event_type,
2928            encoded_state_key,
2929            stripped,
2930            encoded_event_id,
2931            data,
2932        ))?;
2933
2934        Ok(())
2935    }
2936
2937    #[async_test]
2938    pub async fn test_migrating_v2_to_v3() {
2939        let path = new_path();
2940
2941        // Room A: with name, creator and sender.
2942        let room_a_id = room_id!("!room_a:dummy.local");
2943        let room_a_name = "Room A";
2944        let room_a_creator = user_id!("@creator:dummy.local");
2945        // Use a different sender to check that sender is used over creator in
2946        // migration.
2947        let room_a_create_sender = user_id!("@sender:dummy.local");
2948
2949        // Room B: without name, creator and sender.
2950        let room_b_id = room_id!("!room_b:dummy.local");
2951
2952        // Room C: only with sender.
2953        let room_c_id = room_id!("!room_c:dummy.local");
2954        let room_c_create_sender = user_id!("@creator:dummy.local");
2955
2956        // Create and populate db.
2957        {
2958            let db = create_fake_db(&path, 2).await.unwrap();
2959            let conn = db.read().await.unwrap();
2960
2961            let this = db.clone();
2962            conn.with_transaction(move |txn| {
2963                add_room_v2(
2964                    &this,
2965                    txn,
2966                    room_a_id,
2967                    Some(room_a_name),
2968                    Some(room_a_creator),
2969                    Some(room_a_create_sender),
2970                )?;
2971                add_room_v2(&this, txn, room_b_id, None, None, None)?;
2972                add_room_v2(&this, txn, room_c_id, None, None, Some(room_c_create_sender))?;
2973
2974                Result::<_, Error>::Ok(())
2975            })
2976            .await
2977            .unwrap();
2978        }
2979
2980        // This transparently migrates to the latest version.
2981        let store = SqliteStateStore::open(path, Some(SECRET)).await.unwrap();
2982
2983        // Check all room infos are there.
2984        let room_infos = store.get_room_infos(&RoomLoadSettings::default()).await.unwrap();
2985        assert_eq!(room_infos.len(), 3);
2986
2987        let room_a = room_infos.iter().find(|r| r.room_id() == room_a_id).unwrap();
2988        assert_eq!(room_a.name(), Some(room_a_name));
2989        assert_eq!(room_a.creators(), Some(vec![room_a_create_sender.to_owned()]));
2990
2991        let room_b = room_infos.iter().find(|r| r.room_id() == room_b_id).unwrap();
2992        assert_eq!(room_b.name(), None);
2993        assert_eq!(room_b.creators(), None);
2994
2995        let room_c = room_infos.iter().find(|r| r.room_id() == room_c_id).unwrap();
2996        assert_eq!(room_c.name(), None);
2997        assert_eq!(room_c.creators(), Some(vec![room_c_create_sender.to_owned()]));
2998    }
2999
3000    #[async_test]
3001    pub async fn test_migrating_v7_to_v9() {
3002        let path = new_path();
3003
3004        let room_id = room_id!("!room_a:dummy.local");
3005        let wedged_event_transaction_id = TransactionId::new();
3006        let local_event_transaction_id = TransactionId::new();
3007
3008        // Create and populate db.
3009        {
3010            let db = create_fake_db(&path, 7).await.unwrap();
3011            let conn = db.read().await.unwrap();
3012
3013            let wedge_tx = wedged_event_transaction_id.clone();
3014            let local_tx = local_event_transaction_id.clone();
3015
3016            conn.with_transaction(move |txn| {
3017                add_dependent_send_queue_event_v7(
3018                    &db,
3019                    txn,
3020                    room_id,
3021                    &local_tx,
3022                    ChildTransactionId::new(),
3023                    DependentQueuedRequestKind::RedactEvent,
3024                )?;
3025                add_send_queue_event_v7(&db, txn, &wedge_tx, room_id, true)?;
3026                add_send_queue_event_v7(&db, txn, &local_tx, room_id, false)?;
3027                Result::<_, Error>::Ok(())
3028            })
3029            .await
3030            .unwrap();
3031        }
3032
3033        // This transparently migrates to the latest version, which clears up all
3034        // requests and dependent requests.
3035        let store = SqliteStateStore::open(path, Some(SECRET)).await.unwrap();
3036
3037        let requests = store.load_send_queue_requests(room_id).await.unwrap();
3038        assert!(requests.is_empty());
3039
3040        let dependent_requests = store.load_dependent_queued_requests(room_id).await.unwrap();
3041        assert!(dependent_requests.is_empty());
3042    }
3043
3044    fn add_send_queue_event_v7(
3045        this: &SqliteStateStore,
3046        txn: &Transaction<'_>,
3047        transaction_id: &TransactionId,
3048        room_id: &RoomId,
3049        is_wedged: bool,
3050    ) -> Result<(), Error> {
3051        let content =
3052            SerializableEventContent::new(&RoomMessageEventContent::text_plain("Hello").into())?;
3053
3054        let room_id_key = this.encode_key(keys::SEND_QUEUE, room_id);
3055        let room_id_value = this.serialize_value(&room_id.to_owned())?;
3056
3057        let content = this.serialize_json(&content)?;
3058
3059        txn.prepare_cached("INSERT INTO send_queue_events (room_id, room_id_val, transaction_id, content, wedged) VALUES (?, ?, ?, ?, ?)")?
3060            .execute((room_id_key, room_id_value, transaction_id.to_string(), content, is_wedged))?;
3061
3062        Ok(())
3063    }
3064
3065    fn add_dependent_send_queue_event_v7(
3066        this: &SqliteStateStore,
3067        txn: &Transaction<'_>,
3068        room_id: &RoomId,
3069        parent_txn_id: &TransactionId,
3070        own_txn_id: ChildTransactionId,
3071        content: DependentQueuedRequestKind,
3072    ) -> Result<(), Error> {
3073        let room_id_value = this.serialize_value(&room_id.to_owned())?;
3074
3075        let parent_txn_id = parent_txn_id.to_string();
3076        let own_txn_id = own_txn_id.to_string();
3077        let content = this.serialize_json(&content)?;
3078
3079        txn.prepare_cached(
3080            "INSERT INTO dependent_send_queue_events
3081                         (room_id, parent_transaction_id, own_transaction_id, content)
3082                       VALUES (?, ?, ?, ?)",
3083        )?
3084        .execute((room_id_value, parent_txn_id, own_txn_id, content))?;
3085
3086        Ok(())
3087    }
3088
3089    #[derive(Clone, Debug, Serialize, Deserialize)]
3090    pub enum LegacyDependentQueuedRequestKind {
3091        UploadFileWithThumbnail {
3092            content_type: String,
3093            cache_key: MediaRequestParameters,
3094            related_to: OwnedTransactionId,
3095        },
3096    }
3097
3098    #[async_test]
3099    pub async fn test_dependent_queued_request_variant_renaming() {
3100        let path = new_path();
3101        let db = create_fake_db(&path, 7).await.unwrap();
3102
3103        let cache_key = MediaRequestParameters {
3104            format: MediaFormat::File,
3105            source: MediaSource::Plain("https://server.local/foobar".into()),
3106        };
3107        let related_to = TransactionId::new();
3108        let request = LegacyDependentQueuedRequestKind::UploadFileWithThumbnail {
3109            content_type: "image/png".to_owned(),
3110            cache_key,
3111            related_to: related_to.clone(),
3112        };
3113
3114        let data = db
3115            .serialize_json(&request)
3116            .expect("should be able to serialize legacy dependent request");
3117        let deserialized: DependentQueuedRequestKind = db.deserialize_json(&data).expect(
3118            "should be able to deserialize dependent request from legacy dependent request",
3119        );
3120
3121        as_variant!(deserialized, DependentQueuedRequestKind::UploadFileOrThumbnail { related_to: de_related_to, .. } => {
3122            assert_eq!(de_related_to, related_to);
3123        });
3124    }
3125}
3126
3127#[cfg(test)]
3128mod close_reopen_tests {
3129    use std::sync::{
3130        LazyLock,
3131        atomic::{AtomicU32, Ordering::SeqCst},
3132    };
3133
3134    use matrix_sdk_base::StateStore;
3135    use matrix_sdk_test::async_test;
3136    use tempfile::{TempDir, tempdir};
3137
3138    use super::SqliteStateStore;
3139
3140    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
3141    static NUM: AtomicU32 = AtomicU32::new(0);
3142
3143    async fn new_store() -> SqliteStateStore {
3144        let name = NUM.fetch_add(1, SeqCst).to_string();
3145        let tmpdir_path = TMP_DIR.path().join(name);
3146        SqliteStateStore::open(tmpdir_path.to_str().unwrap(), None).await.unwrap()
3147    }
3148
3149    #[async_test]
3150    async fn test_close_completes_without_timeout() {
3151        let store = new_store().await;
3152
3153        // Close should complete quickly without hitting the 5s timeout.
3154        let start = std::time::Instant::now();
3155        store.close().await.unwrap();
3156        let elapsed = start.elapsed();
3157
3158        assert!(
3159            elapsed < std::time::Duration::from_secs(2),
3160            "close() took {elapsed:?}, expected < 2s (no timeout)"
3161        );
3162
3163        // Connections should be None after close.
3164        let guard = store.connections.lock().await;
3165        assert!(guard.is_none(), "connections should be None after close");
3166    }
3167
3168    #[async_test]
3169    async fn test_reopen_restores_connections() {
3170        let store = new_store().await;
3171
3172        store.close().await.unwrap();
3173
3174        // Connections should be None after close.
3175        {
3176            let guard = store.connections.lock().await;
3177            assert!(guard.is_none());
3178        }
3179
3180        store.reopen().await.unwrap();
3181
3182        // Connections should be Some after reopen.
3183        {
3184            let guard = store.connections.lock().await;
3185            assert!(guard.is_some(), "connections should be Some after reopen");
3186        }
3187    }
3188
3189    #[async_test]
3190    async fn test_close_is_idempotent() {
3191        let store = new_store().await;
3192
3193        // First close.
3194        store.close().await.unwrap();
3195        // Second close should also succeed (no-op).
3196        store.close().await.unwrap();
3197
3198        let guard = store.connections.lock().await;
3199        assert!(guard.is_none());
3200    }
3201
3202    #[async_test]
3203    async fn test_reopen_is_idempotent() {
3204        let store = new_store().await;
3205
3206        // Reopen on an active store should be a no-op.
3207        store.reopen().await.unwrap();
3208
3209        // Connections should still be Some.
3210        let guard = store.connections.lock().await;
3211        assert!(guard.is_some());
3212    }
3213
3214    #[async_test]
3215    async fn test_read_fails_when_closed() {
3216        let store = new_store().await;
3217        store.close().await.unwrap();
3218
3219        let err = store.get_custom_value(b"some_key").await;
3220        assert!(err.is_err(), "read should fail when closed");
3221
3222        let err_msg = err.unwrap_err().to_string();
3223        assert!(err_msg.contains("closed"), "error should mention 'closed', got: {err_msg}");
3224    }
3225
3226    #[async_test]
3227    async fn test_write_fails_when_closed() {
3228        let store = new_store().await;
3229        store.close().await.unwrap();
3230
3231        let err = store.set_custom_value(b"key", b"value".to_vec()).await;
3232        assert!(err.is_err(), "write should fail when closed");
3233
3234        let err_msg = err.unwrap_err().to_string();
3235        assert!(err_msg.contains("closed"), "error should mention 'closed', got: {err_msg}");
3236    }
3237
3238    #[async_test]
3239    async fn test_data_persists_across_close_reopen() {
3240        let store = new_store().await;
3241
3242        // Write some data.
3243        store.set_custom_value(b"test_key", b"test_value".to_vec()).await.unwrap();
3244
3245        // Verify it's there.
3246        let value = store.get_custom_value(b"test_key").await.unwrap();
3247        assert_eq!(value.as_deref(), Some(b"test_value".as_slice()));
3248
3249        // Close and reopen.
3250        store.close().await.unwrap();
3251        store.reopen().await.unwrap();
3252
3253        // Data should still be there after reopen.
3254        let value = store.get_custom_value(b"test_key").await.unwrap();
3255        assert_eq!(
3256            value.as_deref(),
3257            Some(b"test_value".as_slice()),
3258            "data should persist across close/reopen"
3259        );
3260    }
3261
3262    #[async_test]
3263    async fn test_multiple_close_reopen_cycles() {
3264        let store = new_store().await;
3265
3266        for i in 0..3 {
3267            let key = format!("key_{i}");
3268            let value = format!("value_{i}");
3269
3270            store.set_custom_value(key.as_bytes(), value.as_bytes().to_vec()).await.unwrap();
3271
3272            store.close().await.unwrap();
3273            store.reopen().await.unwrap();
3274
3275            // Verify all previously written data is still accessible.
3276            for j in 0..=i {
3277                let k = format!("key_{j}");
3278                let v = format!("value_{j}");
3279                let retrieved = store.get_custom_value(k.as_bytes()).await.unwrap();
3280                assert_eq!(
3281                    retrieved.as_deref(),
3282                    Some(v.as_bytes()),
3283                    "data for key_{j} should persist after cycle {i}"
3284                );
3285            }
3286        }
3287    }
3288
3289    #[async_test]
3290    async fn test_pool_is_fully_drained_after_close() {
3291        let store = new_store().await;
3292
3293        // Do a few reads to exercise the pool.
3294        let _ = store.get_custom_value(b"key1").await;
3295        let _ = store.get_custom_value(b"key2").await;
3296
3297        store.close().await.unwrap();
3298
3299        // After close, the connections field should be None (pool and write
3300        // connection have been fully torn down).
3301        let guard = store.connections.lock().await;
3302        assert!(guard.is_none(), "all connections should be released after close");
3303    }
3304
3305    #[async_test]
3306    async fn test_operations_work_immediately_after_reopen() {
3307        let store = new_store().await;
3308
3309        store.close().await.unwrap();
3310        store.reopen().await.unwrap();
3311
3312        // Write should work immediately.
3313        store.set_custom_value(b"after_reopen", b"works".to_vec()).await.unwrap();
3314
3315        // Read should work immediately.
3316        let value = store.get_custom_value(b"after_reopen").await.unwrap();
3317        assert_eq!(value.as_deref(), Some(b"works".as_slice()));
3318    }
3319
3320    #[async_test]
3321    async fn test_close_waits_for_held_read_connection_to_drain() {
3322        let store = new_store().await;
3323
3324        // Acquire a read connection and hold it, simulating an in-flight read.
3325        let held_conn = store.read().await.unwrap();
3326
3327        // Spawn close in a background task — it will close the pool and then
3328        // poll-wait for pool.status().size == 0 in the drain loop.
3329        let store_clone = store.clone();
3330        let close_handle = tokio::spawn(async move {
3331            store_clone.close().await.unwrap();
3332        });
3333
3334        // Give close() a moment to close the pool and enter the drain loop.
3335        tokio::time::sleep(std::time::Duration::from_millis(200)).await;
3336
3337        // The close task should still be running because we hold a connection.
3338        assert!(!close_handle.is_finished(), "close should be waiting for the held connection");
3339
3340        // Release the held connection — this lets pool.status().size drop to 0.
3341        drop(held_conn);
3342
3343        // Now close should complete promptly (well within the 5s timeout).
3344        let timeout = tokio::time::timeout(std::time::Duration::from_secs(3), close_handle).await;
3345        assert!(timeout.is_ok(), "close should complete after the held connection is released");
3346        timeout.unwrap().unwrap();
3347
3348        // Verify the store is fully closed.
3349        let guard = store.connections.lock().await;
3350        assert!(guard.is_none(), "connections should be None after close");
3351    }
3352}