Skip to main content

matrix_sdk_sqlite/
state_store.rs

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