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, repeat_vars,
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; 32]>,
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_row([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_row("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_row(
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_row("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_params = repeat_vars(keys.len());
935            let sql = format!("SELECT value FROM kv_blob WHERE key IN ({sql_params})");
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: Vec<Key>| {
980            let sql_params = repeat_vars(state_keys.len());
981            let sql = format!(
982                "SELECT stripped, data FROM state_event
983                 WHERE room_id = ? AND event_type = ? AND state_key IN ({sql_params})"
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_params = repeat_vars(user_ids_length);
1026            let sql = format!(
1027                "SELECT user_id, data FROM profile WHERE room_id = ? AND user_id IN ({sql_params})"
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_params = repeat_vars(user_ids_length);
1046            let sql = format!(
1047                "SELECT user_id, profile_data FROM global_profiles WHERE user_id IN ({sql_params})"
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_params = repeat_vars(memberships.len());
1070                let sql = format!(
1071                    "SELECT data FROM member WHERE room_id = ? AND membership IN ({sql_params})"
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_row(
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_row(
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_params = repeat_vars(names.len());
1124            let sql = format!(
1125                "SELECT name, data FROM display_name WHERE room_id = ? AND name IN ({sql_params})"
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        thread: Key,
1144        user_id: Key,
1145    ) -> Result<Option<Vec<u8>>> {
1146        Ok(self
1147            .query_row(
1148                "SELECT data FROM receipt
1149                 WHERE room_id = ? AND receipt_type = ? AND thread = ? and user_id = ?",
1150                (room_id, receipt_type, 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_row([&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        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 thread = self.encode_key(keys::RECEIPT, rmp_serde::to_vec_named(&thread)?);
1933        let user_id = self.encode_key(keys::RECEIPT, user_id);
1934
1935        self.read()
1936            .await?
1937            .get_user_receipt(room_id, receipt_type, thread, user_id)
1938            .await?
1939            .map(|value| {
1940                self.deserialize_json::<ReceiptData>(&value).map(|d| (d.event_id, d.receipt))
1941            })
1942            .transpose()
1943    }
1944
1945    async fn get_event_room_receipt_events(
1946        &self,
1947        room_id: &RoomId,
1948        receipt_type: ReceiptType,
1949        thread: ReceiptThread,
1950        event_id: &EventId,
1951    ) -> Result<Vec<(OwnedUserId, Receipt)>> {
1952        let room_id = self.encode_key(keys::RECEIPT, room_id);
1953        let receipt_type = self.encode_key(keys::RECEIPT, receipt_type.to_string());
1954        // We cannot have a NULL primary key so we rely on serialization instead of the
1955        // string representation.
1956        let thread = self.encode_key(keys::RECEIPT, rmp_serde::to_vec_named(&thread)?);
1957        let event_id = self.encode_key(keys::RECEIPT, event_id);
1958
1959        self.read()
1960            .await?
1961            .get_event_receipts(room_id, receipt_type, thread, event_id)
1962            .await?
1963            .iter()
1964            .map(|value| {
1965                self.deserialize_json::<ReceiptData>(value).map(|d| (d.user_id, d.receipt))
1966            })
1967            .collect()
1968    }
1969
1970    async fn get_custom_value(&self, key: &[u8]) -> Result<Option<Vec<u8>>> {
1971        self.read().await?.get_kv_blob(self.encode_custom_key(key)).await
1972    }
1973
1974    async fn set_custom_value_no_read(&self, key: &[u8], value: Vec<u8>) -> Result<()> {
1975        let conn = self.write().await?;
1976        let key = self.encode_custom_key(key);
1977        conn.set_kv_blob(key, value).await?;
1978        Ok(())
1979    }
1980
1981    async fn set_custom_value(&self, key: &[u8], value: Vec<u8>) -> Result<Option<Vec<u8>>> {
1982        let conn = self.write().await?;
1983        let key = self.encode_custom_key(key);
1984        let previous = conn.get_kv_blob(key.clone()).await?;
1985        conn.set_kv_blob(key, value).await?;
1986        Ok(previous)
1987    }
1988
1989    async fn remove_custom_value(&self, key: &[u8]) -> Result<Option<Vec<u8>>> {
1990        let conn = self.write().await?;
1991        let key = self.encode_custom_key(key);
1992        let previous = conn.get_kv_blob(key.clone()).await?;
1993        if previous.is_some() {
1994            conn.delete_kv_blob(key).await?;
1995        }
1996        Ok(previous)
1997    }
1998
1999    async fn remove_room(&self, room_id: &RoomId) -> Result<()> {
2000        let this = self.clone();
2001        let room_id = room_id.to_owned();
2002
2003        let conn = self.write().await?;
2004
2005        conn.with_transaction(move |txn| -> Result<()> {
2006            let room_info_room_id = this.encode_key(keys::ROOM_INFO, &room_id);
2007            txn.remove_room_info(&room_info_room_id)?;
2008
2009            let state_event_room_id = this.encode_key(keys::STATE_EVENT, &room_id);
2010            txn.remove_room_state_events(&state_event_room_id, None)?;
2011
2012            let member_room_id = this.encode_key(keys::MEMBER, &room_id);
2013            txn.remove_room_members(&member_room_id, None)?;
2014
2015            let profile_room_id = this.encode_key(keys::PROFILE, &room_id);
2016            txn.remove_room_profiles(&profile_room_id)?;
2017
2018            let room_account_data_room_id = this.encode_key(keys::ROOM_ACCOUNT_DATA, &room_id);
2019            txn.remove_room_account_data(&room_account_data_room_id)?;
2020
2021            let receipt_room_id = this.encode_key(keys::RECEIPT, &room_id);
2022            txn.remove_room_receipts(&receipt_room_id)?;
2023
2024            let display_name_room_id = this.encode_key(keys::DISPLAY_NAME, &room_id);
2025            txn.remove_room_display_names(&display_name_room_id)?;
2026
2027            let send_queue_room_id = this.encode_key(keys::SEND_QUEUE, &room_id);
2028            txn.remove_room_send_queue(&send_queue_room_id)?;
2029
2030            let dependent_send_queue_room_id =
2031                this.encode_key(keys::DEPENDENTS_SEND_QUEUE, &room_id);
2032            txn.remove_room_dependent_send_queue(&dependent_send_queue_room_id)?;
2033
2034            let thread_subscriptions_room_id =
2035                this.encode_key(keys::THREAD_SUBSCRIPTIONS, &room_id);
2036            txn.execute(
2037                "DELETE FROM thread_subscriptions WHERE room_id = ?",
2038                (thread_subscriptions_room_id,),
2039            )?;
2040
2041            Ok(())
2042        })
2043        .await?;
2044
2045        conn.vacuum().await
2046    }
2047
2048    async fn save_send_queue_request(
2049        &self,
2050        room_id: &RoomId,
2051        transaction_id: OwnedTransactionId,
2052        created_at: MilliSecondsSinceUnixEpoch,
2053        content: QueuedRequestKind,
2054        priority: usize,
2055    ) -> Result<(), Self::Error> {
2056        let room_id_key = self.encode_key(keys::SEND_QUEUE, room_id);
2057        let room_id_value = self.serialize_value(&room_id.to_owned())?;
2058
2059        let content = self.serialize_json(&content)?;
2060        // The transaction id is used both as a key (in remove/update) and a value (as
2061        // it's useful for the callers), so we keep it as is, and neither hash
2062        // it (with encode_key) or encrypt it (through serialize_value). After
2063        // all, it carries no personal information, so this is considered fine.
2064
2065        let created_at_ts: u64 = created_at.0.into();
2066        self.write()
2067            .await?
2068            .with_transaction(move |txn| {
2069                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))?;
2070                Ok(())
2071            })
2072            .await
2073    }
2074
2075    async fn update_send_queue_request(
2076        &self,
2077        room_id: &RoomId,
2078        transaction_id: &TransactionId,
2079        content: QueuedRequestKind,
2080    ) -> Result<bool, Self::Error> {
2081        let room_id = self.encode_key(keys::SEND_QUEUE, room_id);
2082
2083        let content = self.serialize_json(&content)?;
2084        // See comment in [`Self::save_send_queue_request`] to understand why the
2085        // transaction id is neither encrypted or hashed.
2086        let transaction_id = transaction_id.to_string();
2087
2088        let num_updated = self.write()
2089            .await?
2090            .with_transaction(move |txn| {
2091                txn.prepare_cached("UPDATE send_queue_events SET wedge_reason = NULL, content = ? WHERE room_id = ? AND transaction_id = ?")?.execute((content, room_id, transaction_id))
2092            })
2093            .await?;
2094
2095        Ok(num_updated > 0)
2096    }
2097
2098    async fn remove_send_queue_request(
2099        &self,
2100        room_id: &RoomId,
2101        transaction_id: &TransactionId,
2102    ) -> Result<bool, Self::Error> {
2103        let room_id = self.encode_key(keys::SEND_QUEUE, room_id);
2104
2105        // See comment in `save_send_queue_request`.
2106        let transaction_id = transaction_id.to_string();
2107
2108        let num_deleted = self
2109            .write()
2110            .await?
2111            .with_transaction(move |txn| {
2112                txn.prepare_cached(
2113                    "DELETE FROM send_queue_events WHERE room_id = ? AND transaction_id = ?",
2114                )?
2115                .execute((room_id, &transaction_id))
2116            })
2117            .await?;
2118
2119        Ok(num_deleted > 0)
2120    }
2121
2122    async fn load_send_queue_requests(
2123        &self,
2124        room_id: &RoomId,
2125    ) -> Result<Vec<QueuedRequest>, Self::Error> {
2126        let room_id = self.encode_key(keys::SEND_QUEUE, room_id);
2127
2128        // Note: ROWID is always present and is an auto-incremented integer counter. We
2129        // want to maintain the insertion order, so we can sort using it.
2130        // Note 2: transaction_id is not encoded, see why in `save_send_queue_request`.
2131        let res: Vec<(String, Vec<u8>, Option<Vec<u8>>, usize, Option<u64>)> = self
2132            .read()
2133            .await?
2134            .prepare(
2135                "SELECT transaction_id, content, wedge_reason, priority, created_at FROM send_queue_events WHERE room_id = ? ORDER BY priority DESC, ROWID",
2136                |mut stmt| {
2137                    stmt.query((room_id,))?
2138                        .mapped(|row| Ok((row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?, row.get(4)?)))
2139                        .collect()
2140                },
2141            )
2142            .await?;
2143
2144        let mut requests = Vec::with_capacity(res.len());
2145
2146        for entry in res {
2147            let created_at = entry
2148                .4
2149                .and_then(UInt::new)
2150                .map_or_else(MilliSecondsSinceUnixEpoch::now, MilliSecondsSinceUnixEpoch);
2151
2152            requests.push(QueuedRequest {
2153                transaction_id: entry.0.into(),
2154                kind: self.deserialize_json(&entry.1)?,
2155                error: entry.2.map(|v| self.deserialize_value(&v)).transpose()?,
2156                priority: entry.3,
2157                created_at,
2158            });
2159        }
2160
2161        Ok(requests)
2162    }
2163
2164    async fn update_send_queue_request_status(
2165        &self,
2166        room_id: &RoomId,
2167        transaction_id: &TransactionId,
2168        error: Option<QueueWedgeError>,
2169    ) -> Result<(), Self::Error> {
2170        let room_id = self.encode_key(keys::SEND_QUEUE, room_id);
2171
2172        // See comment in `save_send_queue_request`.
2173        let transaction_id = transaction_id.to_string();
2174
2175        // Serialize the error to json bytes (encrypted if option is enabled) if set.
2176        let error_value = error.map(|e| self.serialize_value(&e)).transpose()?;
2177
2178        self.write()
2179            .await?
2180            .with_transaction(move |txn| {
2181                txn.prepare_cached("UPDATE send_queue_events SET wedge_reason = ? WHERE room_id = ? AND transaction_id = ?")?.execute((error_value, room_id, transaction_id))?;
2182                Ok(())
2183            })
2184            .await
2185    }
2186
2187    async fn load_rooms_with_unsent_requests(&self) -> Result<Vec<OwnedRoomId>, Self::Error> {
2188        // If the values were not encrypted, we could use `SELECT DISTINCT` here, but we
2189        // have to manually do the deduplication: indeed, for all X, encrypt(X)
2190        // != encrypted(X), since we use a nonce in the encryption process.
2191
2192        let res: Vec<Vec<u8>> = self
2193            .read()
2194            .await?
2195            .prepare("SELECT room_id_val FROM send_queue_events", |mut stmt| {
2196                stmt.query(())?.mapped(|row| row.get(0)).collect()
2197            })
2198            .await?;
2199
2200        // So we collect the results into a `BTreeSet` to perform the deduplication, and
2201        // then rejigger that into a vector.
2202        Ok(res
2203            .into_iter()
2204            .map(|entry| self.deserialize_value(&entry))
2205            .collect::<Result<BTreeSet<OwnedRoomId>, _>>()?
2206            .into_iter()
2207            .collect())
2208    }
2209
2210    async fn save_dependent_queued_request(
2211        &self,
2212        room_id: &RoomId,
2213        parent_txn_id: &TransactionId,
2214        own_txn_id: ChildTransactionId,
2215        created_at: MilliSecondsSinceUnixEpoch,
2216        content: DependentQueuedRequestKind,
2217    ) -> Result<()> {
2218        let room_id = self.encode_key(keys::DEPENDENTS_SEND_QUEUE, room_id);
2219        let content = self.serialize_json(&content)?;
2220
2221        // See comment in `save_send_queue_request`.
2222        let parent_txn_id = parent_txn_id.to_string();
2223        let own_txn_id = own_txn_id.to_string();
2224
2225        let created_at_ts: u64 = created_at.0.into();
2226        self.write()
2227            .await?
2228            .with_transaction(move |txn| {
2229                txn.prepare_cached(
2230                    r#"INSERT INTO dependent_send_queue_events
2231                         (room_id, parent_transaction_id, own_transaction_id, content, created_at)
2232                       VALUES (?, ?, ?, ?, ?)"#,
2233                )?
2234                .execute((
2235                    room_id,
2236                    parent_txn_id,
2237                    own_txn_id,
2238                    content,
2239                    created_at_ts,
2240                ))?;
2241                Ok(())
2242            })
2243            .await
2244    }
2245
2246    async fn update_dependent_queued_request(
2247        &self,
2248        room_id: &RoomId,
2249        own_transaction_id: &ChildTransactionId,
2250        new_content: DependentQueuedRequestKind,
2251    ) -> Result<bool> {
2252        let room_id = self.encode_key(keys::DEPENDENTS_SEND_QUEUE, room_id);
2253        let content = self.serialize_json(&new_content)?;
2254
2255        // See comment in `save_send_queue_request`.
2256        let own_txn_id = own_transaction_id.to_string();
2257
2258        let num_updated = self
2259            .write()
2260            .await?
2261            .with_transaction(move |txn| {
2262                txn.prepare_cached(
2263                    r#"UPDATE dependent_send_queue_events
2264                       SET content = ?
2265                       WHERE own_transaction_id = ?
2266                       AND room_id = ?"#,
2267                )?
2268                .execute((content, own_txn_id, room_id))
2269            })
2270            .await?;
2271
2272        if num_updated > 1 {
2273            return Err(Error::InconsistentUpdate);
2274        }
2275
2276        Ok(num_updated == 1)
2277    }
2278
2279    async fn mark_dependent_queued_requests_as_ready(
2280        &self,
2281        room_id: &RoomId,
2282        parent_txn_id: &TransactionId,
2283        parent_key: SentRequestKey,
2284    ) -> Result<usize> {
2285        let room_id = self.encode_key(keys::DEPENDENTS_SEND_QUEUE, room_id);
2286        let parent_key = self.serialize_json(&parent_key)?;
2287
2288        // See comment in `save_send_queue_request`.
2289        let parent_txn_id = parent_txn_id.to_string();
2290
2291        self.write()
2292            .await?
2293            .with_transaction(move |txn| {
2294                Ok(txn.prepare_cached(
2295                    "UPDATE dependent_send_queue_events SET parent_key = ? WHERE parent_transaction_id = ? and room_id = ?",
2296                )?
2297                .execute((parent_key, parent_txn_id, room_id))?)
2298            })
2299            .await
2300    }
2301
2302    async fn remove_dependent_queued_request(
2303        &self,
2304        room_id: &RoomId,
2305        txn_id: &ChildTransactionId,
2306    ) -> Result<bool> {
2307        let room_id = self.encode_key(keys::DEPENDENTS_SEND_QUEUE, room_id);
2308
2309        // See comment in `save_send_queue_request`.
2310        let txn_id = txn_id.to_string();
2311
2312        let num_deleted = self
2313            .write()
2314            .await?
2315            .with_transaction(move |txn| {
2316                txn.prepare_cached(
2317                    "DELETE FROM dependent_send_queue_events WHERE own_transaction_id = ? AND room_id = ?",
2318                )?
2319                .execute((txn_id, room_id))
2320            })
2321            .await?;
2322
2323        Ok(num_deleted > 0)
2324    }
2325
2326    async fn load_dependent_queued_requests(
2327        &self,
2328        room_id: &RoomId,
2329    ) -> Result<Vec<DependentQueuedRequest>> {
2330        let room_id = self.encode_key(keys::DEPENDENTS_SEND_QUEUE, room_id);
2331
2332        // Note: transaction_id is not encoded, see why in `save_send_queue_request`.
2333        let res: Vec<(String, String, Option<Vec<u8>>, Vec<u8>, Option<u64>)> = self
2334            .read()
2335            .await?
2336            .prepare(
2337                "SELECT own_transaction_id, parent_transaction_id, parent_key, content, created_at FROM dependent_send_queue_events WHERE room_id = ? ORDER BY ROWID",
2338                |mut stmt| {
2339                    stmt.query((room_id,))?
2340                        .mapped(|row| Ok((row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?, row.get(4)?)))
2341                        .collect()
2342                },
2343            )
2344            .await?;
2345
2346        let mut dependent_events = Vec::with_capacity(res.len());
2347
2348        for entry in res {
2349            let created_at = entry
2350                .4
2351                .and_then(UInt::new)
2352                .map_or_else(MilliSecondsSinceUnixEpoch::now, MilliSecondsSinceUnixEpoch);
2353
2354            dependent_events.push(DependentQueuedRequest {
2355                own_transaction_id: entry.0.into(),
2356                parent_transaction_id: entry.1.into(),
2357                parent_key: entry.2.map(|json| self.deserialize_json(&json)).transpose()?,
2358                kind: self.deserialize_json(&entry.3)?,
2359                created_at,
2360            });
2361        }
2362
2363        Ok(dependent_events)
2364    }
2365
2366    async fn upsert_thread_subscriptions(
2367        &self,
2368        updates: Vec<(&RoomId, &EventId, StoredThreadSubscription)>,
2369    ) -> Result<(), Self::Error> {
2370        let values: Vec<_> = updates
2371            .into_iter()
2372            .map(|(room_id, thread_id, subscription)| {
2373                (
2374                    self.encode_key(keys::THREAD_SUBSCRIPTIONS, room_id),
2375                    self.encode_key(keys::THREAD_SUBSCRIPTIONS, thread_id),
2376                    subscription.status.as_str(),
2377                    subscription.bump_stamp,
2378                )
2379            })
2380            .collect();
2381
2382        self.write()
2383            .await?
2384            .with_transaction(move |txn| {
2385                let mut txn = txn.prepare_cached(
2386                    "INSERT INTO thread_subscriptions (room_id, event_id, status, bump_stamp)
2387                    VALUES (?, ?, ?, ?)
2388                    ON CONFLICT (room_id, event_id) DO UPDATE
2389                    SET
2390                        status =
2391                            CASE
2392                                WHEN thread_subscriptions.bump_stamp IS NULL THEN EXCLUDED.status
2393                                WHEN EXCLUDED.bump_stamp IS NULL THEN EXCLUDED.status
2394                                WHEN thread_subscriptions.bump_stamp < EXCLUDED.bump_stamp THEN EXCLUDED.status
2395                                ELSE thread_subscriptions.status
2396                            END,
2397                        bump_stamp =
2398                            CASE
2399                                WHEN thread_subscriptions.bump_stamp IS NULL THEN EXCLUDED.bump_stamp
2400                                WHEN EXCLUDED.bump_stamp IS NULL THEN thread_subscriptions.bump_stamp
2401                                WHEN thread_subscriptions.bump_stamp < EXCLUDED.bump_stamp THEN EXCLUDED.bump_stamp
2402                                ELSE thread_subscriptions.bump_stamp
2403                            END",
2404                )?;
2405
2406                for value in values {
2407                    txn.execute(value)?;
2408                }
2409
2410                Result::<_, Error>::Ok(())
2411            })
2412            .await?;
2413
2414        Ok(())
2415    }
2416
2417    async fn load_thread_subscription(
2418        &self,
2419        room_id: &RoomId,
2420        thread_id: &EventId,
2421    ) -> Result<Option<StoredThreadSubscription>, Self::Error> {
2422        let room_id = self.encode_key(keys::THREAD_SUBSCRIPTIONS, room_id);
2423        let thread_id = self.encode_key(keys::THREAD_SUBSCRIPTIONS, thread_id);
2424
2425        Ok(self
2426            .read()
2427            .await?
2428            .query_row(
2429                "SELECT status, bump_stamp FROM thread_subscriptions WHERE room_id = ? AND event_id = ?",
2430                (room_id, thread_id),
2431                |row| Ok((row.get::<_, String>(0)?, row.get::<_, Option<u64>>(1)?))
2432            )
2433            .await
2434            .optional()?
2435            .map(|(status, bump_stamp)| -> Result<_, Self::Error> {
2436                let status = ThreadSubscriptionStatus::from_str(&status).map_err(|_| {
2437                    Error::InvalidData { details: format!("Invalid thread status: {status}") }
2438                })?;
2439                Ok(StoredThreadSubscription { status, bump_stamp })
2440            })
2441            .transpose()?)
2442    }
2443
2444    async fn remove_thread_subscription(
2445        &self,
2446        room_id: &RoomId,
2447        thread_id: &EventId,
2448    ) -> Result<(), Self::Error> {
2449        let room_id = self.encode_key(keys::THREAD_SUBSCRIPTIONS, room_id);
2450        let thread_id = self.encode_key(keys::THREAD_SUBSCRIPTIONS, thread_id);
2451
2452        self.write()
2453            .await?
2454            .execute(
2455                "DELETE FROM thread_subscriptions WHERE room_id = ? AND event_id = ?",
2456                (room_id, thread_id),
2457            )
2458            .await?;
2459
2460        Ok(())
2461    }
2462
2463    async fn get_global_profile(
2464        &self,
2465        user_id: &UserId,
2466    ) -> Result<Option<UserProfile>, Self::Error> {
2467        self.read()
2468            .await?
2469            .get_global_profiles(vec![self.encode_key(keys::GLOBAL_PROFILES, user_id)])
2470            .await?
2471            .into_iter()
2472            .next()
2473            .map(|(_, data)| self.deserialize_json(&data))
2474            .transpose()
2475    }
2476
2477    async fn get_global_profiles<'a>(
2478        &self,
2479        user_ids: &'a [OwnedUserId],
2480    ) -> Result<BTreeMap<&'a UserId, UserProfile>, Self::Error> {
2481        if user_ids.is_empty() {
2482            return Ok(BTreeMap::new());
2483        }
2484
2485        let mut user_ids_map = user_ids
2486            .iter()
2487            .map(|u| (self.encode_key(keys::GLOBAL_PROFILES, u), u.as_ref()))
2488            .collect::<BTreeMap<_, _>>();
2489        let user_ids = user_ids_map.keys().cloned().collect();
2490
2491        self.read()
2492            .await?
2493            .get_global_profiles(user_ids)
2494            .await?
2495            .into_iter()
2496            .map(|(user_id, data)| {
2497                Ok((
2498                    user_ids_map
2499                        .remove(user_id.as_slice())
2500                        .expect("returned user IDs were requested"),
2501                    self.deserialize_json(&data)?,
2502                ))
2503            })
2504            .collect()
2505    }
2506
2507    async fn optimize(&self) -> Result<(), Self::Error> {
2508        Ok(self.vacuum().await?)
2509    }
2510
2511    async fn get_size(&self) -> Result<Option<usize>, Self::Error> {
2512        self.get_db_size().await
2513    }
2514
2515    async fn close(&self) -> Result<(), Self::Error> {
2516        connection::close_connections(&self.connections, "State store").await;
2517        Ok(())
2518    }
2519
2520    async fn reopen(&self) -> Result<(), Self::Error> {
2521        connection::reopen_connections(
2522            &self.connections,
2523            self.db_path.clone(),
2524            self.pool_config,
2525            self.runtime_config,
2526        )
2527        .await?;
2528        Ok(())
2529    }
2530}
2531
2532#[derive(Debug, Clone, Serialize, Deserialize)]
2533struct ReceiptData {
2534    receipt: Receipt,
2535    event_id: OwnedEventId,
2536    user_id: OwnedUserId,
2537}
2538
2539#[cfg(test)]
2540mod tests {
2541    use std::sync::{
2542        LazyLock,
2543        atomic::{AtomicU32, Ordering::SeqCst},
2544    };
2545
2546    use matrix_sdk_base::{StateStore, StoreError, statestore_integration_tests};
2547    use tempfile::{TempDir, tempdir};
2548
2549    use super::SqliteStateStore;
2550
2551    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2552    static NUM: AtomicU32 = AtomicU32::new(0);
2553
2554    async fn get_store() -> Result<impl StateStore, StoreError> {
2555        let name = NUM.fetch_add(1, SeqCst).to_string();
2556        let tmpdir_path = TMP_DIR.path().join(name);
2557
2558        tracing::info!("using store @ {}", tmpdir_path.to_str().unwrap());
2559
2560        Ok(SqliteStateStore::open(tmpdir_path.to_str().unwrap(), None).await.unwrap())
2561    }
2562
2563    statestore_integration_tests!();
2564}
2565
2566#[cfg(test)]
2567mod encrypted_tests {
2568    use std::{
2569        path::PathBuf,
2570        sync::{
2571            LazyLock,
2572            atomic::{AtomicU32, Ordering::SeqCst},
2573        },
2574    };
2575
2576    use matrix_sdk_base::{StateStore, StoreError, statestore_integration_tests};
2577    use matrix_sdk_test::async_test;
2578    use tempfile::{TempDir, tempdir};
2579
2580    use super::SqliteStateStore;
2581    use crate::{SqliteStoreConfig, utils::SqliteAsyncConnExt};
2582
2583    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2584    static NUM: AtomicU32 = AtomicU32::new(0);
2585
2586    fn new_state_store_workspace() -> PathBuf {
2587        let name = NUM.fetch_add(1, SeqCst).to_string();
2588        TMP_DIR.path().join(name)
2589    }
2590
2591    async fn get_store() -> Result<impl StateStore, StoreError> {
2592        let tmpdir_path = new_state_store_workspace();
2593
2594        tracing::info!("using store @ {}", tmpdir_path.to_str().unwrap());
2595
2596        Ok(SqliteStateStore::open(tmpdir_path.to_str().unwrap(), Some("default_test_password"))
2597            .await
2598            .unwrap())
2599    }
2600
2601    #[async_test]
2602    async fn test_pool_size() {
2603        let tmpdir_path = new_state_store_workspace();
2604        let store_open_config = SqliteStoreConfig::new(tmpdir_path).pool_max_size(42);
2605
2606        let store = SqliteStateStore::open_with_config(&store_open_config).await.unwrap();
2607
2608        let guard = store.connections.lock().await;
2609        assert_eq!(guard.as_ref().unwrap().pool.status().max_size, 42);
2610    }
2611
2612    #[async_test]
2613    async fn test_cache_size() {
2614        let tmpdir_path = new_state_store_workspace();
2615        let store_open_config = SqliteStoreConfig::new(tmpdir_path).cache_size(1500);
2616
2617        let store = SqliteStateStore::open_with_config(&store_open_config).await.unwrap();
2618
2619        let conn = store.read().await.unwrap();
2620        let cache_size =
2621            conn.query_row("PRAGMA cache_size", (), |row| row.get::<_, i32>(0)).await.unwrap();
2622
2623        // The value passed to `SqliteStoreConfig` is in bytes. Check it is
2624        // converted to kibibytes. Also, it must be a negative value because it
2625        // _is_ the size in kibibytes, not in page size.
2626        assert_eq!(cache_size, -(1500 / 1024));
2627    }
2628
2629    #[async_test]
2630    async fn test_journal_size_limit() {
2631        let tmpdir_path = new_state_store_workspace();
2632        let store_open_config = SqliteStoreConfig::new(tmpdir_path).journal_size_limit(1500);
2633
2634        let store = SqliteStateStore::open_with_config(&store_open_config).await.unwrap();
2635
2636        let conn = store.read().await.unwrap();
2637        let journal_size_limit = conn
2638            .query_row("PRAGMA journal_size_limit", (), |row| row.get::<_, u32>(0))
2639            .await
2640            .unwrap();
2641
2642        // The value passed to `SqliteStoreConfig` is in bytes. It stays in
2643        // bytes in SQLite.
2644        assert_eq!(journal_size_limit, 1500);
2645    }
2646
2647    statestore_integration_tests!();
2648}
2649
2650#[cfg(test)]
2651mod migration_tests {
2652    use std::{
2653        path::{Path, PathBuf},
2654        sync::{
2655            Arc, LazyLock,
2656            atomic::{AtomicU32, Ordering::SeqCst},
2657        },
2658    };
2659
2660    use as_variant::as_variant;
2661    use matrix_sdk_base::{
2662        RoomState, StateStore,
2663        media::{MediaFormat, MediaRequestParameters},
2664        store::{
2665            ChildTransactionId, DependentQueuedRequestKind, RoomLoadSettings,
2666            SerializableEventContent,
2667        },
2668        sync::UnreadNotificationsCount,
2669    };
2670    use matrix_sdk_test::async_test;
2671    use ruma::{
2672        EventId, MilliSecondsSinceUnixEpoch, OwnedTransactionId, RoomId, TransactionId, UserId,
2673        events::{
2674            StateEventType,
2675            room::{MediaSource, create::RoomCreateEventContent, message::RoomMessageEventContent},
2676        },
2677        room_id, server_name, user_id,
2678    };
2679    use rusqlite::Transaction;
2680    use serde::{Deserialize, Serialize};
2681    use serde_json::json;
2682    use tempfile::{TempDir, tempdir};
2683    use tokio::{fs, sync::Mutex};
2684    use zeroize::Zeroizing;
2685
2686    use super::{DATABASE_NAME, SqliteStateStore, init, keys};
2687    use crate::{
2688        OpenStoreError, Secret, SqliteStoreConfig, connection,
2689        error::{Error, Result},
2690        utils::{EncryptableStore as _, SqliteAsyncConnExt, SqliteKeyValueStoreAsyncConnExt},
2691    };
2692
2693    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2694    static NUM: AtomicU32 = AtomicU32::new(0);
2695    const SECRET: &str = "secret";
2696
2697    fn new_path() -> PathBuf {
2698        let name = NUM.fetch_add(1, SeqCst).to_string();
2699        TMP_DIR.path().join(name)
2700    }
2701
2702    async fn create_fake_db(path: &Path, version: u8) -> Result<SqliteStateStore> {
2703        let config = SqliteStoreConfig::new(path);
2704
2705        fs::create_dir_all(&config.path).await.map_err(OpenStoreError::CreateDir).unwrap();
2706
2707        let pool = config.build_pool_of_connections(DATABASE_NAME).unwrap();
2708        let db_path = pool.manager().database_path.clone();
2709        let conn = pool.get().await?;
2710
2711        init(&conn).await?;
2712
2713        let store_cipher = Some(Arc::new(
2714            conn.get_or_create_store_cipher(Secret::PassPhrase(Zeroizing::new(SECRET.to_owned())))
2715                .await
2716                .unwrap(),
2717        ));
2718        let this = SqliteStateStore {
2719            store_cipher,
2720            connections: Arc::new(Mutex::new(Some(connection::SqliteConnections {
2721                pool,
2722                write_connection: Arc::new(Mutex::new(conn)),
2723            }))),
2724            db_path,
2725            pool_config: deadpool::managed::PoolConfig::default(),
2726            runtime_config: crate::RuntimeConfig::default(),
2727        };
2728        this.run_migrations(1, Some(version)).await?;
2729
2730        Ok(this)
2731    }
2732
2733    fn room_info_v1_json(
2734        room_id: &RoomId,
2735        state: RoomState,
2736        name: Option<&str>,
2737        creator: Option<&UserId>,
2738    ) -> serde_json::Value {
2739        // Test with name set or not.
2740        let name_content = match name {
2741            Some(name) => json!({ "name": name }),
2742            None => json!({ "name": null }),
2743        };
2744        // Test with creator set or not.
2745        let create_content = match creator {
2746            Some(creator) => RoomCreateEventContent::new_v1(creator.to_owned()),
2747            None => RoomCreateEventContent::new_v11(),
2748        };
2749
2750        json!({
2751            "room_id": room_id,
2752            "room_type": state,
2753            "notification_counts": UnreadNotificationsCount::default(),
2754            "summary": {
2755                "heroes": [],
2756                "joined_member_count": 0,
2757                "invited_member_count": 0,
2758            },
2759            "members_synced": false,
2760            "base_info": {
2761                "dm_targets": [],
2762                "max_power_level": 100,
2763                "name": {
2764                    "Original": {
2765                        "content": name_content,
2766                    },
2767                },
2768                "create": {
2769                    "Original": {
2770                        "content": create_content,
2771                    }
2772                }
2773            },
2774        })
2775    }
2776
2777    #[async_test]
2778    pub async fn test_migrating_v1_to_v2() {
2779        let path = new_path();
2780        // Create and populate db.
2781        {
2782            let db = create_fake_db(&path, 1).await.unwrap();
2783            let conn = db.read().await.unwrap();
2784
2785            let this = db.clone();
2786            conn.with_transaction(move |txn| {
2787                for i in 0..5 {
2788                    let room_id = RoomId::parse(format!("!room_{i}:localhost")).unwrap();
2789                    let (state, stripped) =
2790                        if i < 3 { (RoomState::Joined, false) } else { (RoomState::Invited, true) };
2791                    let info = room_info_v1_json(&room_id, state, None, None);
2792
2793                    let room_id = this.encode_key(keys::ROOM_INFO, room_id);
2794                    let data = this.serialize_json(&info)?;
2795
2796                    txn.prepare_cached(
2797                        "INSERT INTO room_info (room_id, stripped, data)
2798                         VALUES (?, ?, ?)",
2799                    )?
2800                    .execute((room_id, stripped, data))?;
2801                }
2802
2803                Result::<_, Error>::Ok(())
2804            })
2805            .await
2806            .unwrap();
2807        }
2808
2809        // This transparently migrates to the latest version.
2810        let store = SqliteStateStore::open(path, Some(SECRET)).await.unwrap();
2811
2812        // Check all room infos are there.
2813        assert_eq!(store.get_room_infos(&RoomLoadSettings::default()).await.unwrap().len(), 5);
2814    }
2815
2816    // Add a room in version 2 format of the state store.
2817    fn add_room_v2(
2818        this: &SqliteStateStore,
2819        txn: &Transaction<'_>,
2820        room_id: &RoomId,
2821        name: Option<&str>,
2822        create_creator: Option<&UserId>,
2823        create_sender: Option<&UserId>,
2824    ) -> Result<(), Error> {
2825        let room_info_json = room_info_v1_json(room_id, RoomState::Joined, name, create_creator);
2826
2827        let encoded_room_id = this.encode_key(keys::ROOM_INFO, room_id);
2828        let encoded_state =
2829            this.encode_key(keys::ROOM_INFO, serde_json::to_string(&RoomState::Joined)?);
2830        let data = this.serialize_json(&room_info_json)?;
2831
2832        txn.prepare_cached(
2833            "INSERT INTO room_info (room_id, state, data)
2834             VALUES (?, ?, ?)",
2835        )?
2836        .execute((encoded_room_id, encoded_state, data))?;
2837
2838        // Test with or without `m.room.create` event in the room state.
2839        let Some(create_sender) = create_sender else {
2840            return Ok(());
2841        };
2842
2843        let create_content = match create_creator {
2844            Some(creator) => RoomCreateEventContent::new_v1(creator.to_owned()),
2845            None => RoomCreateEventContent::new_v11(),
2846        };
2847
2848        let event_id = EventId::new_v1(server_name!("dummy.local"));
2849        let create_event = json!({
2850            "content": create_content,
2851            "event_id": event_id,
2852            "sender": create_sender.to_owned(),
2853            "origin_server_ts": MilliSecondsSinceUnixEpoch::now(),
2854            "state_key": "",
2855            "type": "m.room.create",
2856            "unsigned": {},
2857        });
2858
2859        let encoded_room_id = this.encode_key(keys::STATE_EVENT, room_id);
2860        let encoded_event_type =
2861            this.encode_key(keys::STATE_EVENT, StateEventType::RoomCreate.to_string());
2862        let encoded_state_key = this.encode_key(keys::STATE_EVENT, "");
2863        let stripped = false;
2864        let encoded_event_id = this.encode_key(keys::STATE_EVENT, event_id);
2865        let data = this.serialize_json(&create_event)?;
2866
2867        txn.prepare_cached(
2868            "INSERT
2869             INTO state_event (room_id, event_type, state_key, stripped, event_id, data)
2870             VALUES (?, ?, ?, ?, ?, ?)",
2871        )?
2872        .execute((
2873            encoded_room_id,
2874            encoded_event_type,
2875            encoded_state_key,
2876            stripped,
2877            encoded_event_id,
2878            data,
2879        ))?;
2880
2881        Ok(())
2882    }
2883
2884    #[async_test]
2885    pub async fn test_migrating_v2_to_v3() {
2886        let path = new_path();
2887
2888        // Room A: with name, creator and sender.
2889        let room_a_id = room_id!("!room_a:dummy.local");
2890        let room_a_name = "Room A";
2891        let room_a_creator = user_id!("@creator:dummy.local");
2892        // Use a different sender to check that sender is used over creator in
2893        // migration.
2894        let room_a_create_sender = user_id!("@sender:dummy.local");
2895
2896        // Room B: without name, creator and sender.
2897        let room_b_id = room_id!("!room_b:dummy.local");
2898
2899        // Room C: only with sender.
2900        let room_c_id = room_id!("!room_c:dummy.local");
2901        let room_c_create_sender = user_id!("@creator:dummy.local");
2902
2903        // Create and populate db.
2904        {
2905            let db = create_fake_db(&path, 2).await.unwrap();
2906            let conn = db.read().await.unwrap();
2907
2908            let this = db.clone();
2909            conn.with_transaction(move |txn| {
2910                add_room_v2(
2911                    &this,
2912                    txn,
2913                    room_a_id,
2914                    Some(room_a_name),
2915                    Some(room_a_creator),
2916                    Some(room_a_create_sender),
2917                )?;
2918                add_room_v2(&this, txn, room_b_id, None, None, None)?;
2919                add_room_v2(&this, txn, room_c_id, None, None, Some(room_c_create_sender))?;
2920
2921                Result::<_, Error>::Ok(())
2922            })
2923            .await
2924            .unwrap();
2925        }
2926
2927        // This transparently migrates to the latest version.
2928        let store = SqliteStateStore::open(path, Some(SECRET)).await.unwrap();
2929
2930        // Check all room infos are there.
2931        let room_infos = store.get_room_infos(&RoomLoadSettings::default()).await.unwrap();
2932        assert_eq!(room_infos.len(), 3);
2933
2934        let room_a = room_infos.iter().find(|r| r.room_id() == room_a_id).unwrap();
2935        assert_eq!(room_a.name(), Some(room_a_name));
2936        assert_eq!(room_a.creators(), Some(vec![room_a_create_sender.to_owned()]));
2937
2938        let room_b = room_infos.iter().find(|r| r.room_id() == room_b_id).unwrap();
2939        assert_eq!(room_b.name(), None);
2940        assert_eq!(room_b.creators(), None);
2941
2942        let room_c = room_infos.iter().find(|r| r.room_id() == room_c_id).unwrap();
2943        assert_eq!(room_c.name(), None);
2944        assert_eq!(room_c.creators(), Some(vec![room_c_create_sender.to_owned()]));
2945    }
2946
2947    #[async_test]
2948    pub async fn test_migrating_v7_to_v9() {
2949        let path = new_path();
2950
2951        let room_id = room_id!("!room_a:dummy.local");
2952        let wedged_event_transaction_id = TransactionId::new();
2953        let local_event_transaction_id = TransactionId::new();
2954
2955        // Create and populate db.
2956        {
2957            let db = create_fake_db(&path, 7).await.unwrap();
2958            let conn = db.read().await.unwrap();
2959
2960            let wedge_tx = wedged_event_transaction_id.clone();
2961            let local_tx = local_event_transaction_id.clone();
2962
2963            conn.with_transaction(move |txn| {
2964                add_dependent_send_queue_event_v7(
2965                    &db,
2966                    txn,
2967                    room_id,
2968                    &local_tx,
2969                    ChildTransactionId::new(),
2970                    DependentQueuedRequestKind::RedactEvent,
2971                )?;
2972                add_send_queue_event_v7(&db, txn, &wedge_tx, room_id, true)?;
2973                add_send_queue_event_v7(&db, txn, &local_tx, room_id, false)?;
2974                Result::<_, Error>::Ok(())
2975            })
2976            .await
2977            .unwrap();
2978        }
2979
2980        // This transparently migrates to the latest version, which clears up all
2981        // requests and dependent requests.
2982        let store = SqliteStateStore::open(path, Some(SECRET)).await.unwrap();
2983
2984        let requests = store.load_send_queue_requests(room_id).await.unwrap();
2985        assert!(requests.is_empty());
2986
2987        let dependent_requests = store.load_dependent_queued_requests(room_id).await.unwrap();
2988        assert!(dependent_requests.is_empty());
2989    }
2990
2991    fn add_send_queue_event_v7(
2992        this: &SqliteStateStore,
2993        txn: &Transaction<'_>,
2994        transaction_id: &TransactionId,
2995        room_id: &RoomId,
2996        is_wedged: bool,
2997    ) -> Result<(), Error> {
2998        let content =
2999            SerializableEventContent::new(&RoomMessageEventContent::text_plain("Hello").into())?;
3000
3001        let room_id_key = this.encode_key(keys::SEND_QUEUE, room_id);
3002        let room_id_value = this.serialize_value(&room_id.to_owned())?;
3003
3004        let content = this.serialize_json(&content)?;
3005
3006        txn.prepare_cached("INSERT INTO send_queue_events (room_id, room_id_val, transaction_id, content, wedged) VALUES (?, ?, ?, ?, ?)")?
3007            .execute((room_id_key, room_id_value, transaction_id.to_string(), content, is_wedged))?;
3008
3009        Ok(())
3010    }
3011
3012    fn add_dependent_send_queue_event_v7(
3013        this: &SqliteStateStore,
3014        txn: &Transaction<'_>,
3015        room_id: &RoomId,
3016        parent_txn_id: &TransactionId,
3017        own_txn_id: ChildTransactionId,
3018        content: DependentQueuedRequestKind,
3019    ) -> Result<(), Error> {
3020        let room_id_value = this.serialize_value(&room_id.to_owned())?;
3021
3022        let parent_txn_id = parent_txn_id.to_string();
3023        let own_txn_id = own_txn_id.to_string();
3024        let content = this.serialize_json(&content)?;
3025
3026        txn.prepare_cached(
3027            "INSERT INTO dependent_send_queue_events
3028                         (room_id, parent_transaction_id, own_transaction_id, content)
3029                       VALUES (?, ?, ?, ?)",
3030        )?
3031        .execute((room_id_value, parent_txn_id, own_txn_id, content))?;
3032
3033        Ok(())
3034    }
3035
3036    #[derive(Clone, Debug, Serialize, Deserialize)]
3037    pub enum LegacyDependentQueuedRequestKind {
3038        UploadFileWithThumbnail {
3039            content_type: String,
3040            cache_key: MediaRequestParameters,
3041            related_to: OwnedTransactionId,
3042        },
3043    }
3044
3045    #[async_test]
3046    pub async fn test_dependent_queued_request_variant_renaming() {
3047        let path = new_path();
3048        let db = create_fake_db(&path, 7).await.unwrap();
3049
3050        let cache_key = MediaRequestParameters {
3051            format: MediaFormat::File,
3052            source: MediaSource::Plain("https://server.local/foobar".into()),
3053        };
3054        let related_to = TransactionId::new();
3055        let request = LegacyDependentQueuedRequestKind::UploadFileWithThumbnail {
3056            content_type: "image/png".to_owned(),
3057            cache_key,
3058            related_to: related_to.clone(),
3059        };
3060
3061        let data = db
3062            .serialize_json(&request)
3063            .expect("should be able to serialize legacy dependent request");
3064        let deserialized: DependentQueuedRequestKind = db.deserialize_json(&data).expect(
3065            "should be able to deserialize dependent request from legacy dependent request",
3066        );
3067
3068        as_variant!(deserialized, DependentQueuedRequestKind::UploadFileOrThumbnail { related_to: de_related_to, .. } => {
3069            assert_eq!(de_related_to, related_to);
3070        });
3071    }
3072}
3073
3074#[cfg(test)]
3075mod close_reopen_tests {
3076    use std::sync::{
3077        LazyLock,
3078        atomic::{AtomicU32, Ordering::SeqCst},
3079    };
3080
3081    use matrix_sdk_base::StateStore;
3082    use matrix_sdk_test::async_test;
3083    use tempfile::{TempDir, tempdir};
3084
3085    use super::SqliteStateStore;
3086
3087    static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
3088    static NUM: AtomicU32 = AtomicU32::new(0);
3089
3090    async fn new_store() -> SqliteStateStore {
3091        let name = NUM.fetch_add(1, SeqCst).to_string();
3092        let tmpdir_path = TMP_DIR.path().join(name);
3093        SqliteStateStore::open(tmpdir_path.to_str().unwrap(), None).await.unwrap()
3094    }
3095
3096    #[async_test]
3097    async fn test_close_completes_without_timeout() {
3098        let store = new_store().await;
3099
3100        // Close should complete quickly without hitting the 5s timeout.
3101        let start = std::time::Instant::now();
3102        store.close().await.unwrap();
3103        let elapsed = start.elapsed();
3104
3105        assert!(
3106            elapsed < std::time::Duration::from_secs(2),
3107            "close() took {elapsed:?}, expected < 2s (no timeout)"
3108        );
3109
3110        // Connections should be None after close.
3111        let guard = store.connections.lock().await;
3112        assert!(guard.is_none(), "connections should be None after close");
3113    }
3114
3115    #[async_test]
3116    async fn test_reopen_restores_connections() {
3117        let store = new_store().await;
3118
3119        store.close().await.unwrap();
3120
3121        // Connections should be None after close.
3122        {
3123            let guard = store.connections.lock().await;
3124            assert!(guard.is_none());
3125        }
3126
3127        store.reopen().await.unwrap();
3128
3129        // Connections should be Some after reopen.
3130        {
3131            let guard = store.connections.lock().await;
3132            assert!(guard.is_some(), "connections should be Some after reopen");
3133        }
3134    }
3135
3136    #[async_test]
3137    async fn test_close_is_idempotent() {
3138        let store = new_store().await;
3139
3140        // First close.
3141        store.close().await.unwrap();
3142        // Second close should also succeed (no-op).
3143        store.close().await.unwrap();
3144
3145        let guard = store.connections.lock().await;
3146        assert!(guard.is_none());
3147    }
3148
3149    #[async_test]
3150    async fn test_reopen_is_idempotent() {
3151        let store = new_store().await;
3152
3153        // Reopen on an active store should be a no-op.
3154        store.reopen().await.unwrap();
3155
3156        // Connections should still be Some.
3157        let guard = store.connections.lock().await;
3158        assert!(guard.is_some());
3159    }
3160
3161    #[async_test]
3162    async fn test_read_fails_when_closed() {
3163        let store = new_store().await;
3164        store.close().await.unwrap();
3165
3166        let err = store.get_custom_value(b"some_key").await;
3167        assert!(err.is_err(), "read should fail when closed");
3168
3169        let err_msg = err.unwrap_err().to_string();
3170        assert!(err_msg.contains("closed"), "error should mention 'closed', got: {err_msg}");
3171    }
3172
3173    #[async_test]
3174    async fn test_write_fails_when_closed() {
3175        let store = new_store().await;
3176        store.close().await.unwrap();
3177
3178        let err = store.set_custom_value(b"key", b"value".to_vec()).await;
3179        assert!(err.is_err(), "write should fail when closed");
3180
3181        let err_msg = err.unwrap_err().to_string();
3182        assert!(err_msg.contains("closed"), "error should mention 'closed', got: {err_msg}");
3183    }
3184
3185    #[async_test]
3186    async fn test_data_persists_across_close_reopen() {
3187        let store = new_store().await;
3188
3189        // Write some data.
3190        store.set_custom_value(b"test_key", b"test_value".to_vec()).await.unwrap();
3191
3192        // Verify it's there.
3193        let value = store.get_custom_value(b"test_key").await.unwrap();
3194        assert_eq!(value.as_deref(), Some(b"test_value".as_slice()));
3195
3196        // Close and reopen.
3197        store.close().await.unwrap();
3198        store.reopen().await.unwrap();
3199
3200        // Data should still be there after reopen.
3201        let value = store.get_custom_value(b"test_key").await.unwrap();
3202        assert_eq!(
3203            value.as_deref(),
3204            Some(b"test_value".as_slice()),
3205            "data should persist across close/reopen"
3206        );
3207    }
3208
3209    #[async_test]
3210    async fn test_multiple_close_reopen_cycles() {
3211        let store = new_store().await;
3212
3213        for i in 0..3 {
3214            let key = format!("key_{i}");
3215            let value = format!("value_{i}");
3216
3217            store.set_custom_value(key.as_bytes(), value.as_bytes().to_vec()).await.unwrap();
3218
3219            store.close().await.unwrap();
3220            store.reopen().await.unwrap();
3221
3222            // Verify all previously written data is still accessible.
3223            for j in 0..=i {
3224                let k = format!("key_{j}");
3225                let v = format!("value_{j}");
3226                let retrieved = store.get_custom_value(k.as_bytes()).await.unwrap();
3227                assert_eq!(
3228                    retrieved.as_deref(),
3229                    Some(v.as_bytes()),
3230                    "data for key_{j} should persist after cycle {i}"
3231                );
3232            }
3233        }
3234    }
3235
3236    #[async_test]
3237    async fn test_pool_is_fully_drained_after_close() {
3238        let store = new_store().await;
3239
3240        // Do a few reads to exercise the pool.
3241        let _ = store.get_custom_value(b"key1").await;
3242        let _ = store.get_custom_value(b"key2").await;
3243
3244        store.close().await.unwrap();
3245
3246        // After close, the connections field should be None (pool and write
3247        // connection have been fully torn down).
3248        let guard = store.connections.lock().await;
3249        assert!(guard.is_none(), "all connections should be released after close");
3250    }
3251
3252    #[async_test]
3253    async fn test_operations_work_immediately_after_reopen() {
3254        let store = new_store().await;
3255
3256        store.close().await.unwrap();
3257        store.reopen().await.unwrap();
3258
3259        // Write should work immediately.
3260        store.set_custom_value(b"after_reopen", b"works".to_vec()).await.unwrap();
3261
3262        // Read should work immediately.
3263        let value = store.get_custom_value(b"after_reopen").await.unwrap();
3264        assert_eq!(value.as_deref(), Some(b"works".as_slice()));
3265    }
3266
3267    #[async_test]
3268    async fn test_close_waits_for_held_read_connection_to_drain() {
3269        let store = new_store().await;
3270
3271        // Acquire a read connection and hold it, simulating an in-flight read.
3272        let held_conn = store.read().await.unwrap();
3273
3274        // Spawn close in a background task — it will close the pool and then
3275        // poll-wait for pool.status().size == 0 in the drain loop.
3276        let store_clone = store.clone();
3277        let close_handle = tokio::spawn(async move {
3278            store_clone.close().await.unwrap();
3279        });
3280
3281        // Give close() a moment to close the pool and enter the drain loop.
3282        tokio::time::sleep(std::time::Duration::from_millis(200)).await;
3283
3284        // The close task should still be running because we hold a connection.
3285        assert!(!close_handle.is_finished(), "close should be waiting for the held connection");
3286
3287        // Release the held connection — this lets pool.status().size drop to 0.
3288        drop(held_conn);
3289
3290        // Now close should complete promptly (well within the 5s timeout).
3291        let timeout = tokio::time::timeout(std::time::Duration::from_secs(3), close_handle).await;
3292        assert!(timeout.is_ok(), "close should complete after the held connection is released");
3293        timeout.unwrap().unwrap();
3294
3295        // Verify the store is fully closed.
3296        let guard = store.connections.lock().await;
3297        assert!(guard.is_none(), "connections should be None after close");
3298    }
3299}