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