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