1use std::{
18 collections::HashMap,
19 fmt,
20 iter::once,
21 ops::{Deref, Not},
22 path::{Path, PathBuf},
23 sync::Arc,
24};
25
26use async_trait::async_trait;
27use deadpool::managed::PoolConfig;
28use matrix_sdk_base::{
29 cross_process_lock::CrossProcessLockGeneration,
30 deserialized_responses::TimelineEvent,
31 event_cache::{
32 Event, Gap,
33 store::{EventCacheStore, extract_event_relation},
34 thread::ThreadInfo,
35 },
36 linked_chunk::{
37 ChunkContent, ChunkIdentifier, ChunkIdentifierGenerator, ChunkMetadata, LinkedChunkId,
38 Position, RawChunk, Update,
39 },
40 timer,
41};
42use matrix_sdk_store_encryption::StoreCipher;
43use ruma::{
44 EventId, MilliSecondsSinceUnixEpoch, OwnedEventId, RoomId, events::relation::RelationType,
45};
46use rusqlite::{
47 OptionalExtension, ToSql, Transaction, TransactionBehavior, params, params_from_iter,
48};
49use tokio::{
50 fs,
51 sync::{Mutex, OwnedMutexGuard},
52};
53use tracing::{debug, error, instrument, trace};
54
55use crate::{
56 OpenStoreError, RuntimeConfig, Secret, SqliteStoreConfig,
57 connection::{self, Connection as SqliteAsyncConn, Pool as SqlitePool, SqliteConnections},
58 error::{Error, Result},
59 utils::{
60 EncryptableStore, Key, SqliteAsyncConnExt, SqliteKeyValueStoreAsyncConnExt,
61 SqliteKeyValueStoreConnExt, SqliteTransactionExt, host_parameters,
62 },
63};
64
65mod keys {
66 pub const LINKED_CHUNKS: &str = "linked_chunks";
68 pub const EVENTS: &str = "events";
69}
70
71const DATABASE_NAME: &str = "matrix-sdk-event-cache.sqlite3";
73
74const CHUNK_TYPE_EVENT_TYPE_STRING: &str = "E";
77const CHUNK_TYPE_GAP_TYPE_STRING: &str = "G";
80
81struct Encryption {
86 cipher: Option<StoreCipher>,
87}
88
89impl Encryption {
90 fn encode_event(&self, event: &TimelineEvent) -> Result<EncodedEvent> {
91 let serialized = serde_json::to_vec(event)?;
92
93 let raw_event = event.raw();
95 let (relates_to, rel_type) = extract_event_relation(raw_event).unzip();
96
97 let content = self.encode_value(serialized)?;
99
100 Ok(EncodedEvent {
101 content,
102 rel_type,
103 relates_to: relates_to
104 .map(|relates_to| self.encode_event_id(keys::EVENTS, &relates_to)),
105 })
106 }
107
108 fn decode_event(&self, raw_encoded_event: &[u8]) -> Result<Event> {
109 Ok(serde_json::from_slice(&self.decode_value(raw_encoded_event)?)?)
110 }
111
112 fn encode_event_id(&self, table_name: &str, event_id: &EventId) -> Key {
115 self.encode_key(table_name, event_id)
116 }
117
118 fn encode_room_id(&self, table_name: &str, room_id: &RoomId) -> Key {
120 self.encode_key(table_name, room_id)
121 }
122
123 fn encode_linked_chunk(&self, table_name: &str, linked_chunk_id: &LinkedChunkId<'_>) -> Key {
125 self.encode_key(table_name, linked_chunk_id.storage_key())
126 }
127
128 fn encode_thread_id(&self, thread_id: &EventId) -> Result<Vec<u8>> {
130 self.encode_value(String::from(thread_id.as_str()))
131 }
132
133 #[cfg(test)]
135 fn decode_thread_id(&self, encoded_thread_id: &[u8]) -> Result<OwnedEventId> {
136 let as_slice = self.decode_value(encoded_thread_id)?;
137 let as_str = str::from_utf8(as_slice.as_ref())?;
138
139 Ok(EventId::parse(as_str)?)
140 }
141
142 fn encode_thread_info(&self, thread_info: &ThreadInfo) -> Result<Vec<u8>> {
144 self.encode_value(serde_json::to_vec(thread_info)?)
145 }
146
147 fn decode_thread_info(&self, encoded_thread_info: &[u8]) -> Result<ThreadInfo> {
149 Ok(serde_json::from_slice(&self.decode_value(encoded_thread_info)?)?)
150 }
151}
152
153impl EncryptableStore for Encryption {
154 fn get_cypher(&self) -> Option<&StoreCipher> {
155 self.cipher.as_ref()
156 }
157}
158
159#[derive(Clone)]
161pub struct SqliteEventCacheStore {
162 encryption: Arc<Encryption>,
164
165 connections: Arc<Mutex<Option<SqliteConnections>>>,
167
168 db_path: PathBuf,
170
171 pool_config: PoolConfig,
173
174 runtime_config: RuntimeConfig,
176}
177
178#[cfg(not(tarpaulin_include))]
179impl fmt::Debug for SqliteEventCacheStore {
180 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
181 f.debug_struct("SqliteEventCacheStore").finish_non_exhaustive()
182 }
183}
184
185impl SqliteEventCacheStore {
186 pub async fn open(
189 path: impl AsRef<Path>,
190 passphrase: Option<&str>,
191 ) -> Result<Self, OpenStoreError> {
192 Self::open_with_config(&SqliteStoreConfig::new(path).passphrase(passphrase)).await
193 }
194
195 pub async fn open_with_key(
198 path: impl AsRef<Path>,
199 key: Option<&[u8]>,
200 ) -> Result<Self, OpenStoreError> {
201 Self::open_with_config(&SqliteStoreConfig::new(path).key(key)).await
202 }
203
204 #[instrument(skip(config), fields(path = ?config.path))]
206 pub async fn open_with_config(config: &SqliteStoreConfig) -> Result<Self, OpenStoreError> {
207 debug!(?config);
208
209 let _timer = timer!("open_with_config");
210
211 fs::create_dir_all(&config.path).await.map_err(OpenStoreError::CreateDir)?;
212
213 let db_path = config.path.join(DATABASE_NAME);
214 let pool_config = config.pool_config();
215 let runtime_config = config.runtime_config();
216
217 let pool = config.build_pool_of_connections(DATABASE_NAME)?;
218
219 let this =
220 Self::open_with_pool(pool, db_path, pool_config, runtime_config, config.secret.clone())
221 .await?;
222
223 this.write().await?.apply_runtime_config(runtime_config).await?;
225
226 Ok(this)
227 }
228
229 async fn open_with_pool(
232 pool: SqlitePool,
233 db_path: PathBuf,
234 pool_config: PoolConfig,
235 runtime_config: RuntimeConfig,
236 secret: Option<Secret>,
237 ) -> Result<Self, OpenStoreError> {
238 let conn = pool.get().await?;
239
240 let version = conn.db_version().await?;
241
242 run_migrations(&conn, version).await?;
243
244 conn.wal_checkpoint().await;
245
246 let cipher = match secret {
247 Some(s) => Some(conn.get_or_create_store_cipher(s).await?),
248 None => None,
249 };
250
251 let connections = SqliteConnections {
252 pool,
253 write_connection: Arc::new(Mutex::new(conn)),
255 };
256
257 Ok(Self {
258 encryption: Arc::new(Encryption { cipher }),
259 connections: Arc::new(Mutex::new(Some(connections))),
260 db_path,
261 pool_config,
262 runtime_config,
263 })
264 }
265
266 #[instrument(skip_all)]
268 async fn read(&self) -> Result<SqliteAsyncConn> {
269 let pool = {
270 let guard = self.connections.lock().await;
271 let conns = guard.as_ref().ok_or(Error::StoreClosed)?;
272 conns.pool.clone()
273 };
274
275 let connection = pool.get().await?;
276
277 connection.execute_batch("PRAGMA foreign_keys = ON;").await?;
282
283 Ok(connection)
284 }
285
286 #[instrument(skip_all)]
288 async fn write(&self) -> Result<OwnedMutexGuard<SqliteAsyncConn>> {
289 let write_connection = {
290 let guard = self.connections.lock().await;
291 let conns = guard.as_ref().ok_or(Error::StoreClosed)?;
292 conns.write_connection.clone()
293 };
294
295 let connection = write_connection.lock_owned().await;
296
297 connection.execute_batch("PRAGMA foreign_keys = ON;").await?;
302
303 Ok(connection)
304 }
305
306 fn map_row_to_chunk(
307 row: &rusqlite::Row<'_>,
308 ) -> Result<(u64, Option<u64>, Option<u64>, String), rusqlite::Error> {
309 Ok((
310 row.get::<_, u64>(0)?,
311 row.get::<_, Option<u64>>(1)?,
312 row.get::<_, Option<u64>>(2)?,
313 row.get::<_, String>(3)?,
314 ))
315 }
316
317 pub async fn vacuum(&self) -> Result<()> {
318 let write_connection = {
319 let guard = self.connections.lock().await;
320 let conns = guard.as_ref().ok_or(Error::StoreClosed)?;
321 conns.write_connection.clone()
322 };
323 write_connection.lock().await.vacuum().await
324 }
325
326 async fn get_db_size(&self) -> Result<Option<usize>> {
327 let pool = {
328 let guard = self.connections.lock().await;
329 let conns = guard.as_ref().ok_or(Error::StoreClosed)?;
330 conns.pool.clone()
331 };
332 Ok(Some(pool.get().await?.get_db_size().await?))
333 }
334
335 pub async fn close(&self) -> Result<()> {
336 connection::close_connections(&self.connections, "Event cache store").await;
337 Ok(())
338 }
339
340 pub async fn reopen(&self) -> Result<()> {
341 connection::reopen_connections(
342 &self.connections,
343 self.db_path.clone(),
344 self.pool_config,
345 self.runtime_config,
346 )
347 .await?;
348 Ok(())
349 }
350
351 #[cfg(test)]
353 async fn pool_max_size(&self) -> Option<usize> {
354 let guard = self.connections.lock().await;
355 guard.as_ref().map(|conns| conns.pool.status().max_size)
356 }
357}
358
359struct EncodedEvent {
360 content: Vec<u8>,
361 rel_type: Option<String>,
362 relates_to: Option<Key>,
363}
364
365trait TransactionExtForLinkedChunks {
366 fn rebuild_chunk(
367 &self,
368 encryption: &Encryption,
369 linked_chunk_id: &Key,
370 previous: Option<u64>,
371 index: u64,
372 next: Option<u64>,
373 chunk_type: &str,
374 ) -> Result<RawChunk<Event, Gap>>;
375
376 fn load_gap_content(
377 &self,
378 encryption: &Encryption,
379 linked_chunk_id: &Key,
380 chunk_id: ChunkIdentifier,
381 ) -> Result<Gap>;
382
383 fn load_events_content(
384 &self,
385 encryption: &Encryption,
386 linked_chunk_id: &Key,
387 chunk_id: ChunkIdentifier,
388 ) -> Result<Vec<Event>>;
389}
390
391impl TransactionExtForLinkedChunks for Transaction<'_> {
392 fn rebuild_chunk(
393 &self,
394 encryption: &Encryption,
395 linked_chunk_id: &Key,
396 previous: Option<u64>,
397 id: u64,
398 next: Option<u64>,
399 chunk_type: &str,
400 ) -> Result<RawChunk<Event, Gap>> {
401 let previous = previous.map(ChunkIdentifier::new);
402 let next = next.map(ChunkIdentifier::new);
403 let id = ChunkIdentifier::new(id);
404
405 match chunk_type {
406 CHUNK_TYPE_GAP_TYPE_STRING => {
407 let gap = self.load_gap_content(encryption, linked_chunk_id, id)?;
409 Ok(RawChunk { content: ChunkContent::Gap(gap), previous, identifier: id, next })
410 }
411
412 CHUNK_TYPE_EVENT_TYPE_STRING => {
413 let events = self.load_events_content(encryption, linked_chunk_id, id)?;
415 Ok(RawChunk {
416 content: ChunkContent::Items(events),
417 previous,
418 identifier: id,
419 next,
420 })
421 }
422
423 other => {
424 Err(Error::InvalidData {
426 details: format!("a linked chunk has an unknown type {other}"),
427 })
428 }
429 }
430 }
431
432 fn load_gap_content(
433 &self,
434 encryption: &Encryption,
435 linked_chunk_id: &Key,
436 chunk_id: ChunkIdentifier,
437 ) -> Result<Gap> {
438 let encoded_prev_token: Vec<u8> = self.query_one(
441 "SELECT prev_token FROM gap_chunks WHERE chunk_id = ? AND linked_chunk_id = ?",
442 (chunk_id.index(), &linked_chunk_id),
443 |row| row.get(0),
444 )?;
445 let prev_token_bytes = encryption.decode_value(&encoded_prev_token)?;
446 let prev_token = String::from_utf8(prev_token_bytes.into_owned())?;
447 Ok(Gap { token: prev_token })
448 }
449
450 fn load_events_content(
451 &self,
452 encryption: &Encryption,
453 linked_chunk_id: &Key,
454 chunk_id: ChunkIdentifier,
455 ) -> Result<Vec<Event>> {
456 let mut events = Vec::new();
458
459 for event_data in self
460 .prepare(
461 "SELECT events.content \
462 FROM event_chunks ec, events \
463 WHERE events.event_id = ec.event_id AND ec.chunk_id = ? AND ec.linked_chunk_id = ? \
464 ORDER BY ec.position ASC",
465 )?
466 .query_map((chunk_id.index(), &linked_chunk_id), |row| row.get::<_, Vec<u8>>(0))?
467 {
468 events.push(encryption.decode_event(&event_data?)?);
469 }
470
471 Ok(events)
472 }
473}
474
475async fn run_migrations(conn: &SqliteAsyncConn, version: u8) -> Result<()> {
477 conn.execute_batch("PRAGMA foreign_keys = ON;").await?;
479
480 if version < 1 {
481 debug!("Creating database");
482 conn.execute_batch("PRAGMA journal_mode = wal;").await?;
486 conn.with_transaction(|txn| {
487 txn.execute_batch(include_str!("../migrations/event_cache_store/001_init.sql"))?;
488 txn.set_db_version(1)
489 })
490 .await?;
491 }
492
493 if version < 2 {
494 debug!("Upgrading database to version 2");
495 conn.with_transaction(|txn| {
496 txn.execute_batch(include_str!("../migrations/event_cache_store/002_lease_locks.sql"))?;
497 txn.set_db_version(2)
498 })
499 .await?;
500 }
501
502 if version < 3 {
503 debug!("Upgrading database to version 3");
504 conn.with_transaction(|txn| {
505 txn.execute_batch(include_str!("../migrations/event_cache_store/003_events.sql"))?;
506 txn.set_db_version(3)
507 })
508 .await?;
509 }
510
511 if version < 4 {
512 debug!("Upgrading database to version 4");
513 conn.with_transaction(|txn| {
514 txn.execute_batch(include_str!(
515 "../migrations/event_cache_store/004_ignore_policy.sql"
516 ))?;
517 txn.set_db_version(4)
518 })
519 .await?;
520 }
521
522 if version < 5 {
523 debug!("Upgrading database to version 5");
524 conn.with_transaction(|txn| {
525 txn.execute_batch(include_str!(
526 "../migrations/event_cache_store/005_events_index_on_event_id.sql"
527 ))?;
528 txn.set_db_version(5)
529 })
530 .await?;
531 }
532
533 if version < 6 {
534 debug!("Upgrading database to version 6");
535 conn.with_transaction(|txn| {
536 txn.execute_batch(include_str!("../migrations/event_cache_store/006_events.sql"))?;
537 txn.set_db_version(6)
538 })
539 .await?;
540 }
541
542 if version < 7 {
543 debug!("Upgrading database to version 7");
544 conn.with_transaction(|txn| {
545 txn.execute_batch(include_str!(
546 "../migrations/event_cache_store/007_event_chunks.sql"
547 ))?;
548 txn.set_db_version(7)
549 })
550 .await?;
551 }
552
553 if version < 8 {
554 debug!("Upgrading database to version 8");
555 conn.with_transaction(|txn| {
556 txn.execute_batch(include_str!(
557 "../migrations/event_cache_store/008_linked_chunk_id.sql"
558 ))?;
559 txn.set_db_version(8)
560 })
561 .await?;
562 }
563
564 if version < 9 {
565 debug!("Upgrading database to version 9");
566 conn.with_transaction(|txn| {
567 txn.execute_batch(include_str!(
568 "../migrations/event_cache_store/009_related_event_index.sql"
569 ))?;
570 txn.set_db_version(9)
571 })
572 .await?;
573 }
574
575 if version < 10 {
576 debug!("Upgrading database to version 10");
577 conn.with_transaction(|txn| {
578 txn.execute_batch(include_str!("../migrations/event_cache_store/010_drop_media.sql"))?;
579 txn.set_db_version(10)
580 })
581 .await?;
582
583 if version >= 1 {
584 conn.vacuum().await?;
587 }
588 }
589
590 if version < 11 {
591 debug!("Upgrading database to version 11");
592 conn.with_transaction(|txn| {
593 txn.execute_batch(include_str!(
594 "../migrations/event_cache_store/011_empty_event_cache.sql"
595 ))?;
596 txn.set_db_version(11)
597 })
598 .await?;
599 }
600
601 if version < 12 {
602 debug!("Upgrading database to version 12");
603 conn.with_transaction(|txn| {
604 txn.execute_batch(include_str!(
605 "../migrations/event_cache_store/012_store_event_type.sql"
606 ))?;
607 txn.set_db_version(12)
608 })
609 .await?;
610 }
611
612 if version < 13 {
613 debug!("Upgrading database to version 13");
614 conn.with_transaction(|txn| {
615 txn.execute_batch(include_str!(
616 "../migrations/event_cache_store/013_lease_locks_with_generation.sql"
617 ))?;
618 txn.set_db_version(13)
619 })
620 .await?;
621 }
622
623 if version < 14 {
624 debug!("Upgrading database to version 14");
625 conn.with_transaction(|txn| {
626 txn.execute_batch(include_str!(
627 "../migrations/event_cache_store/014_event_chunks_event_id_index.sql"
628 ))?;
629 txn.set_db_version(14)
630 })
631 .await?;
632 }
633
634 if version < 15 {
635 debug!("Upgrading database to version 15");
636 conn.with_transaction(|txn| {
637 txn.execute_batch(include_str!(
638 "../migrations/event_cache_store/015_event_ids_are_encoded.sql"
639 ))?;
640 txn.set_db_version(15)
641 })
642 .await?;
643 }
644
645 if version < 16 {
646 debug!("Upgrading database to version 16");
647 conn.with_transaction(|txn| {
648 txn.execute_batch(include_str!("../migrations/event_cache_store/016_threads.sql"))?;
649 txn.set_db_version(16)
650 })
651 .await?;
652 }
653
654 if version < 17 {
655 debug!("Upgrading database to version 17");
656 conn.with_transaction(|txn| {
657 txn.execute_batch(include_str!(
658 "../migrations/event_cache_store/017_threads_with_thread_infos.sql"
659 ))?;
660 txn.set_db_version(17)
661 })
662 .await?;
663 }
664
665 if version < 18 {
666 debug!("Upgrading database to version 18");
667 conn.with_transaction(|txn| {
668 txn.execute_batch(include_str!("../migrations/event_cache_store/018_event_chunks_unique_linked_chunk_id_event_id.sql"))?;
669 txn.set_db_version(18)
670 })
671 .await?;
672 }
673
674 Ok(())
675}
676
677#[async_trait]
678impl EventCacheStore for SqliteEventCacheStore {
679 type Error = Error;
680
681 #[instrument(skip(self))]
682 async fn try_take_leased_lock(
683 &self,
684 lease_duration_ms: u32,
685 key: &str,
686 holder: &str,
687 ) -> Result<Option<CrossProcessLockGeneration>> {
688 let key = key.to_owned();
689 let holder = holder.to_owned();
690
691 let now: u64 = MilliSecondsSinceUnixEpoch::now().get().into();
692 let expiration = now + lease_duration_ms as u64;
693
694 let generation = self
696 .write()
697 .await?
698 .with_transaction(move |txn| {
699 txn.query_one(
700 "INSERT INTO lease_locks (key, holder, expiration) \
701 VALUES (?1, ?2, ?3) \
702 ON CONFLICT (key) \
703 DO \
704 UPDATE SET \
705 holder = excluded.holder, \
706 expiration = excluded.expiration, \
707 generation = \
708 CASE holder \
709 WHEN excluded.holder THEN generation \
710 ELSE generation + 1 \
711 END \
712 WHERE \
713 holder = excluded.holder \
714 OR expiration < ?4 \
715 RETURNING generation",
716 (key, holder, expiration, now),
717 |row| row.get(0),
718 )
719 .optional()
720 })
721 .await?;
722
723 Ok(generation)
724 }
725
726 #[instrument(skip(self, updates))]
727 async fn handle_linked_chunk_updates(
728 &self,
729 linked_chunk_id: LinkedChunkId<'_>,
730 updates: Vec<Update<Event, Gap>>,
731 ) -> Result<(), Self::Error> {
732 let _timer = timer!("method");
733
734 let hashed_linked_chunk_id =
735 self.encryption.encode_linked_chunk(keys::LINKED_CHUNKS, &linked_chunk_id);
736 let hashed_room_id =
737 self.encryption.encode_room_id(keys::EVENTS, linked_chunk_id.room_id());
738 let encryption = self.encryption.clone();
739
740 with_immediate_transaction(self, move |txn| {
743 for update in updates {
744 match update {
745 Update::NewItemsChunk { previous, new, next } => {
746 let previous = previous.as_ref().map(ChunkIdentifier::index);
747 let new = new.index();
748 let next = next.as_ref().map(ChunkIdentifier::index);
749
750 trace!("new events chunk (prev={previous:?}, i={new}, next={next:?})");
751
752 insert_chunk(
753 txn,
754 &hashed_linked_chunk_id,
755 previous,
756 new,
757 next,
758 CHUNK_TYPE_EVENT_TYPE_STRING,
759 )?;
760 }
761
762 Update::NewGapChunk { previous, new, next, gap } => {
763 let hashed_prev_token = encryption.encode_value(gap.token)?;
764
765 let previous = previous.as_ref().map(ChunkIdentifier::index);
766 let new = new.index();
767 let next = next.as_ref().map(ChunkIdentifier::index);
768
769 trace!("new gap chunk (prev={previous:?}, i={new}, next={next:?})");
770
771 insert_chunk(
773 txn,
774 &hashed_linked_chunk_id,
775 previous,
776 new,
777 next,
778 CHUNK_TYPE_GAP_TYPE_STRING,
779 )?;
780
781 txn.execute(
783 r#"
784 INSERT INTO gap_chunks(chunk_id, linked_chunk_id, prev_token)
785 VALUES (?, ?, ?)
786 "#,
787 (new, &hashed_linked_chunk_id, hashed_prev_token),
788 )?;
789 }
790
791 Update::RemoveChunk(chunk_identifier) => {
792 let chunk_id = chunk_identifier.index();
793
794 trace!("removing chunk @ {chunk_id}");
795
796 let (previous, next): (Option<usize>, Option<usize>) = txn.query_one(
798 "SELECT previous, next FROM linked_chunks WHERE id = ? AND linked_chunk_id = ?",
799 (chunk_id, &hashed_linked_chunk_id),
800 |row| Ok((row.get(0)?, row.get(1)?))
801 )?;
802
803 if let Some(previous) = previous {
805 txn.execute("UPDATE linked_chunks SET next = ? WHERE id = ? AND linked_chunk_id = ?", (next, previous, &hashed_linked_chunk_id))?;
806 }
807
808 if let Some(next) = next {
810 txn.execute("UPDATE linked_chunks SET previous = ? WHERE id = ? AND linked_chunk_id = ?", (previous, next, &hashed_linked_chunk_id))?;
811 }
812
813 txn.execute("DELETE FROM linked_chunks WHERE id = ? AND linked_chunk_id = ?", (chunk_id, &hashed_linked_chunk_id))?;
816 }
817
818 Update::PushItems { at, items } => {
819 if items.is_empty() {
820 continue;
822 }
823
824 let chunk_id = at.chunk_identifier().index();
825
826 trace!("pushing {} items @ {chunk_id}", items.len());
827
828 let mut chunk_statement = txn.prepare(
829 "INSERT INTO event_chunks(chunk_id, linked_chunk_id, event_id, position) VALUES (?, ?, ?, ?)"
830 )?;
831
832 let mut content_statement = txn.prepare(
839 "INSERT OR REPLACE INTO events(room_id, event_id, event_type, session_id, content, relates_to, rel_type) VALUES (?, ?, ?, ?, ?, ?, ?)"
840 )?;
841
842 let invalid_event = |event: TimelineEvent| {
843 let Some(event_id) = event.event_id() else {
844 error!("Trying to push an event with no ID");
845 return None;
846 };
847
848 let Some(event_type) = event.kind.event_type() else {
849 error!(%event_id, "Trying to save an event with no event type");
850 return None;
851 };
852
853 Some((event_id.to_owned(), event_type, event))
854 };
855
856 for (i, (event_id, event_type, event)) in items.into_iter().filter_map(invalid_event).enumerate() {
857 let hashed_event_id = encryption.encode_event_id(
858 keys::EVENTS,
864 &event_id,
865 );
866
867 {
869 let index = at.index() + i;
870
871 chunk_statement.execute((chunk_id, &hashed_linked_chunk_id, &hashed_event_id, index))?;
872 }
873
874 {
876 let hashed_session_id = event.kind.session_id().map(|s| encryption.encode_key(keys::EVENTS, s));
877 let hashed_event_type = encryption.encode_key(keys::EVENTS, event_type);
878 let encoded_event = encryption.encode_event(&event)?;
879
880 content_statement.execute((
881 &hashed_room_id,
882 &hashed_event_id,
883 hashed_event_type,
884 hashed_session_id,
885 encoded_event.content,
886 encoded_event.relates_to,
887 encoded_event.rel_type
888 ))?;
889 }
890 }
891 }
892
893 Update::ReplaceItem { at, item: event } => {
894 let chunk_id = at.chunk_identifier().index();
895 let index = at.index();
896
897 trace!("replacing item @ {chunk_id}:{index}");
898
899 let Some(event_id) = event.event_id().map(|event_id| event_id.to_owned()) else {
901 error!("Trying to replace an event with a new one that has no ID");
902 continue;
903 };
904
905 let Some(event_type) = event.kind.event_type() else {
906 error!(%event_id, "Trying to save an event with no event type");
907 continue;
908 };
909
910 let hashed_event_id = encryption.encode_event_id(keys::EVENTS, &event_id);
911
912 {
920 let hashed_session_id = event.kind.session_id().map(|s| encryption.encode_key(keys::EVENTS, s));
921 let hashed_event_type = encryption.encode_key(keys::EVENTS, event_type);
922 let encoded_event = encryption.encode_event(&event)?;
923
924 txn.execute(
925 "INSERT OR REPLACE INTO events(room_id, event_id, event_type, session_id, content, relates_to, rel_type) VALUES (?, ?, ?, ?, ?, ?, ?)",
926 (
927 &hashed_room_id,
928 &hashed_event_id,
929 hashed_event_type,
930 hashed_session_id,
931 encoded_event.content,
932 encoded_event.relates_to,
933 encoded_event.rel_type
934 ),
935 )?;
936 }
937
938 {
940 txn.execute(
942 r#"UPDATE event_chunks SET event_id = ? WHERE linked_chunk_id = ? AND chunk_id = ? AND position = ?"#,
943 (&hashed_event_id, &hashed_linked_chunk_id, chunk_id, index)
944 )?;
945 }
946 }
947
948 Update::RemoveItem { at } => {
949 let chunk_id = at.chunk_identifier().index();
950 let index = at.index();
951
952 trace!("removing item @ {chunk_id}:{index}");
953
954 txn.execute("DELETE FROM event_chunks WHERE linked_chunk_id = ? AND chunk_id = ? AND position = ?", (&hashed_linked_chunk_id, chunk_id, index))?;
956
957 txn.execute(
1058 r#"
1059 UPDATE event_chunks
1060 SET position = -(position - 1)
1061 WHERE linked_chunk_id = ? AND chunk_id = ? AND position > ?
1062 "#,
1063 (&hashed_linked_chunk_id, chunk_id, index)
1064 )?;
1065 txn.execute(
1066 r#"
1067 UPDATE event_chunks
1068 SET position = -position
1069 WHERE position < 0 AND linked_chunk_id = ? AND chunk_id = ?
1070 "#,
1071 (&hashed_linked_chunk_id, chunk_id)
1072 )?;
1073
1074 }
1077
1078 Update::DetachLastItems { at } => {
1079 let chunk_id = at.chunk_identifier().index();
1080 let index = at.index();
1081
1082 trace!("truncating items >= {chunk_id}:{index}");
1083
1084 txn.execute("DELETE FROM event_chunks WHERE linked_chunk_id = ? AND chunk_id = ? AND position >= ?", (&hashed_linked_chunk_id, chunk_id, index))?;
1086
1087 }
1090
1091 Update::Clear => {
1092 trace!("clearing items");
1093
1094 txn.execute(
1096 "DELETE FROM linked_chunks WHERE linked_chunk_id = ?",
1097 (&hashed_linked_chunk_id,),
1098 )?;
1099
1100 }
1103
1104 Update::StartReattachItems | Update::EndReattachItems => {
1105 }
1107 }
1108 }
1109
1110 Ok(())
1111 })
1112 .await?;
1113
1114 Ok(())
1115 }
1116
1117 #[instrument(skip(self))]
1118 async fn load_all_chunks(
1119 &self,
1120 linked_chunk_id: LinkedChunkId<'_>,
1121 ) -> Result<Vec<RawChunk<Event, Gap>>, Self::Error> {
1122 let _timer = timer!("method");
1123
1124 let hashed_linked_chunk_id =
1125 self.encryption.encode_linked_chunk(keys::LINKED_CHUNKS, &linked_chunk_id);
1126 let encryption = self.encryption.clone();
1127
1128 let result = self
1129 .read()
1130 .await?
1131 .with_transaction(move |txn| -> Result<_> {
1132 let mut items = Vec::new();
1133
1134 for data in txn
1136 .prepare(
1137 "SELECT id, previous, next, type FROM linked_chunks WHERE linked_chunk_id = ? ORDER BY id",
1138 )?
1139 .query_map((&hashed_linked_chunk_id,), Self::map_row_to_chunk)?
1140 {
1141 let (id, previous, next, chunk_type) = data?;
1142 let new = txn.rebuild_chunk(
1143 &encryption,
1144 &hashed_linked_chunk_id,
1145 previous,
1146 id,
1147 next,
1148 chunk_type.as_str(),
1149 )?;
1150 items.push(new);
1151 }
1152
1153 Ok(items)
1154 })
1155 .await?;
1156
1157 Ok(result)
1158 }
1159
1160 #[instrument(skip(self))]
1161 async fn load_all_chunks_metadata(
1162 &self,
1163 linked_chunk_id: LinkedChunkId<'_>,
1164 ) -> Result<Vec<ChunkMetadata>, Self::Error> {
1165 let _timer = timer!("method");
1166
1167 let hashed_linked_chunk_id =
1168 self.encryption.encode_linked_chunk(keys::LINKED_CHUNKS, &linked_chunk_id);
1169
1170 self.read()
1171 .await?
1172 .with_transaction(move |txn| -> Result<_> {
1173 let num_events_by_chunk_ids = txn
1206 .prepare(
1207 "SELECT ec.chunk_id, COUNT(ec.event_id) \
1208 FROM event_chunks as ec \
1209 WHERE ec.linked_chunk_id = ? \
1210 GROUP BY ec.chunk_id",
1211 )?
1212 .query_map((&hashed_linked_chunk_id,), |row| {
1213 Ok((row.get::<_, u64>(0)?, row.get::<_, usize>(1)?))
1214 })?
1215 .collect::<Result<HashMap<_, _>, _>>()?;
1216
1217 txn.prepare(
1218 "SELECT \
1219 lc.id, \
1220 lc.previous, \
1221 lc.next, \
1222 lc.type \
1223 FROM linked_chunks as lc \
1224 WHERE lc.linked_chunk_id = ? \
1225 ORDER BY lc.id",
1226 )?
1227 .query_map((&hashed_linked_chunk_id,), |row| {
1228 Ok((
1229 row.get::<_, u64>(0)?,
1230 row.get::<_, Option<u64>>(1)?,
1231 row.get::<_, Option<u64>>(2)?,
1232 row.get::<_, String>(3)?,
1233 ))
1234 })?
1235 .map(|data| -> Result<_> {
1236 let (id, previous, next, chunk_type) = data?;
1237
1238 let num_items = if chunk_type == CHUNK_TYPE_GAP_TYPE_STRING {
1245 0
1246 } else {
1247 num_events_by_chunk_ids.get(&id).copied().unwrap_or(0)
1248 };
1249
1250 Ok(ChunkMetadata {
1251 identifier: ChunkIdentifier::new(id),
1252 previous: previous.map(ChunkIdentifier::new),
1253 next: next.map(ChunkIdentifier::new),
1254 num_items,
1255 })
1256 })
1257 .collect::<Result<Vec<_>, _>>()
1258 })
1259 .await
1260 }
1261
1262 #[instrument(skip(self))]
1263 async fn load_last_chunk(
1264 &self,
1265 linked_chunk_id: LinkedChunkId<'_>,
1266 ) -> Result<(Option<RawChunk<Event, Gap>>, ChunkIdentifierGenerator), Self::Error> {
1267 let _timer = timer!("method");
1268
1269 let hashed_linked_chunk_id =
1270 self.encryption.encode_linked_chunk(keys::LINKED_CHUNKS, &linked_chunk_id);
1271 let encryption = self.encryption.clone();
1272
1273 self
1274 .read()
1275 .await?
1276 .with_transaction(move |txn| -> Result<_> {
1277 let (observed_max_identifier, number_of_chunks) = txn
1279 .prepare(
1280 "SELECT MAX(id), COUNT(*) FROM linked_chunks WHERE linked_chunk_id = ?"
1281 )?
1282 .query_one(
1283 (&hashed_linked_chunk_id,),
1284 |row| {
1285 Ok((
1286 row.get::<_, Option<u64>>(0)?,
1291 row.get::<_, u64>(1)?,
1292 ))
1293 }
1294 )?;
1295
1296 let chunk_identifier_generator = match observed_max_identifier {
1297 Some(max_observed_identifier) => {
1298 ChunkIdentifierGenerator::new_from_previous_chunk_identifier(
1299 ChunkIdentifier::new(max_observed_identifier)
1300 )
1301 },
1302 None => ChunkIdentifierGenerator::new_from_scratch(),
1303 };
1304
1305 let Some((chunk_identifier, previous_chunk, chunk_type)) = txn
1307 .prepare(
1308 "SELECT id, previous, type FROM linked_chunks WHERE linked_chunk_id = ? AND next IS NULL"
1309 )?
1310 .query_one(
1311 (&hashed_linked_chunk_id,),
1312 |row| {
1313 Ok((
1314 row.get::<_, u64>(0)?,
1315 row.get::<_, Option<u64>>(1)?,
1316 row.get::<_, String>(2)?,
1317 ))
1318 }
1319 )
1320 .optional()?
1321 else {
1322 if number_of_chunks == 0 {
1325 return Ok((None, chunk_identifier_generator));
1326 }
1327 else {
1334 return Err(Error::InvalidData {
1335 details:
1336 "last chunk is not found but chunks exist: the linked chunk contains a cycle"
1337 .to_owned()
1338 }
1339 )
1340 }
1341 };
1342
1343 let last_chunk = txn.rebuild_chunk(
1345 &encryption,
1346 &hashed_linked_chunk_id,
1347 previous_chunk,
1348 chunk_identifier,
1349 None,
1350 &chunk_type
1351 )?;
1352
1353 Ok((Some(last_chunk), chunk_identifier_generator))
1354 })
1355 .await
1356 }
1357
1358 #[instrument(skip(self))]
1359 async fn load_previous_chunk(
1360 &self,
1361 linked_chunk_id: LinkedChunkId<'_>,
1362 before_chunk_identifier: ChunkIdentifier,
1363 ) -> Result<Option<RawChunk<Event, Gap>>, Self::Error> {
1364 let _timer = timer!("method");
1365
1366 let hashed_linked_chunk_id =
1367 self.encryption.encode_linked_chunk(keys::LINKED_CHUNKS, &linked_chunk_id);
1368 let encryption = self.encryption.clone();
1369
1370 self
1371 .read()
1372 .await?
1373 .with_transaction(move |txn| -> Result<_> {
1374 let Some((chunk_identifier, previous_chunk, next_chunk, chunk_type)) = txn
1376 .prepare(
1377 "SELECT id, previous, next, type FROM linked_chunks WHERE linked_chunk_id = ? AND next = ?"
1378 )?
1379 .query_one(
1380 (&hashed_linked_chunk_id, before_chunk_identifier.index()),
1381 |row| {
1382 Ok((
1383 row.get::<_, u64>(0)?,
1384 row.get::<_, Option<u64>>(1)?,
1385 row.get::<_, Option<u64>>(2)?,
1386 row.get::<_, String>(3)?,
1387 ))
1388 }
1389 )
1390 .optional()?
1391 else {
1392 return Ok(None);
1394 };
1395
1396 let last_chunk = txn.rebuild_chunk(
1398 &encryption,
1399 &hashed_linked_chunk_id,
1400 previous_chunk,
1401 chunk_identifier,
1402 next_chunk,
1403 &chunk_type
1404 )?;
1405
1406 Ok(Some(last_chunk))
1407 })
1408 .await
1409 }
1410
1411 async fn load_thread_info(
1412 &self,
1413 room_id: &RoomId,
1414 thread_id: &EventId,
1415 insert_default_if_missing: bool,
1416 ) -> Result<Option<ThreadInfo>, Self::Error> {
1417 let linked_chunk_id = LinkedChunkId::Thread(room_id, thread_id);
1418 let hashed_linked_chunk_id =
1419 self.encryption.encode_linked_chunk(keys::LINKED_CHUNKS, &linked_chunk_id);
1420 let encryption = self.encryption.clone();
1421
1422 let maybe_thread_info = self
1427 .read()
1428 .await?
1429 .with_transaction(move |txn| {
1430 let maybe_encoded_thread_info = txn
1431 .query_one(
1432 "SELECT info FROM threads WHERE linked_chunk_id = ?",
1433 (hashed_linked_chunk_id,),
1434 |row| row.get::<_, Vec<u8>>(0),
1435 )
1436 .optional()?;
1437
1438 maybe_encoded_thread_info
1439 .map(|encoded_thread_info| {
1440 encryption.decode_thread_info(encoded_thread_info.as_slice())
1441 })
1442 .transpose()
1443 })
1444 .await?;
1445
1446 if let Some(thread_info) = maybe_thread_info {
1447 return Ok(Some(thread_info));
1448 } else if !insert_default_if_missing {
1449 return Ok(None);
1451 }
1452
1453 let thread_info = ThreadInfo::default();
1456 let hashed_linked_chunk_id =
1457 self.encryption.encode_linked_chunk(keys::LINKED_CHUNKS, &linked_chunk_id);
1458 let hashed_room_id = self.encryption.encode_room_id(keys::EVENTS, room_id);
1459 let hashed_thread_id = self.encryption.encode_thread_id(thread_id)?;
1460 let encoded_thread_id = self.encryption.encode_thread_info(&thread_info)?;
1461
1462 self.write()
1463 .await?
1464 .with_transaction(move |txn| {
1465 txn.execute(
1466 "INSERT INTO threads VALUES (?, ?, ?, ?)",
1467 (hashed_linked_chunk_id, hashed_room_id, hashed_thread_id, encoded_thread_id),
1468 )?;
1469
1470 Ok::<(), Self::Error>(())
1471 })
1472 .await?;
1473
1474 Ok(Some(thread_info))
1475 }
1476
1477 async fn update_thread_info(
1478 &self,
1479 room_id: &RoomId,
1480 thread_id: &EventId,
1481 thread_info: &ThreadInfo,
1482 ) -> Result<(), Self::Error> {
1483 let hashed_linked_chunk_id = self
1484 .encryption
1485 .encode_linked_chunk(keys::LINKED_CHUNKS, &LinkedChunkId::Thread(room_id, thread_id));
1486 let encoded_thread_info = self.encryption.encode_thread_info(thread_info)?;
1487
1488 self.write()
1489 .await?
1490 .with_transaction(move |txn| {
1491 txn.execute(
1492 "UPDATE threads SET info = ? WHERE linked_chunk_id = ?",
1493 (encoded_thread_info, hashed_linked_chunk_id),
1494 )?;
1495
1496 Ok(())
1497 })
1498 .await
1499 }
1500
1501 #[instrument(skip(self))]
1502 async fn clear_all_events(&self, room_id: Option<&RoomId>) -> Result<(), Self::Error> {
1503 let _timer = timer!("method");
1504
1505 match room_id {
1506 None => {
1508 self.write()
1509 .await?
1510 .with_transaction(move |txn| {
1511 txn.execute("DELETE FROM linked_chunks", ())?;
1513
1514 txn.execute("DELETE FROM events", ())?;
1517
1518 Ok(())
1519 })
1520 .await
1521 }
1522
1523 Some(room_id) => {
1525 let encryption = self.encryption.clone();
1526 let room_id = room_id.to_owned();
1527
1528 self.write()
1529 .await?
1530 .with_transaction(move |txn| {
1531 {
1533 let mut delete =
1534 txn.prepare("DELETE FROM linked_chunks WHERE linked_chunk_id = ?")?;
1535
1536 for linked_chunk_id in
1537 [LinkedChunkId::Room(&room_id), LinkedChunkId::PinnedEvents(&room_id)]
1538 {
1539 let linked_chunk_id = encryption.encode_linked_chunk(keys::LINKED_CHUNKS, &linked_chunk_id);
1540
1541 delete.execute((&linked_chunk_id,))?;
1545 }
1546 }
1547
1548 let encoded_room_id = encryption.encode_room_id(keys::EVENTS, &room_id);
1549
1550 {
1552 txn.execute(
1553 "DELETE FROM linked_chunks WHERE linked_chunk_id IN (SELECT linked_chunk_id FROM threads WHERE room_id = ?)",
1554 (&encoded_room_id,),
1555 )?;
1556 }
1557
1558 txn.execute(
1560 "DELETE FROM events WHERE room_id = ?",
1561 (encoded_room_id,),
1562 )?;
1563
1564 Ok(())
1565 })
1566 .await
1567 }
1568 }
1569 }
1570
1571 #[instrument(skip(self, event_ids))]
1572 async fn filter_duplicated_events(
1573 &self,
1574 linked_chunk_id: LinkedChunkId<'_>,
1575 event_ids: Vec<OwnedEventId>,
1576 ) -> Result<Vec<(OwnedEventId, Position)>, Self::Error> {
1577 let _timer = timer!("method");
1578
1579 if event_ids.is_empty() {
1583 return Ok(Vec::new());
1584 }
1585
1586 let hashed_linked_chunk_id =
1588 self.encryption.encode_linked_chunk(keys::LINKED_CHUNKS, &linked_chunk_id);
1589 let event_ids_and_hashed_event_ids = event_ids
1590 .into_iter()
1591 .map(|event_id| {
1592 let hashed_event_id = self.encryption.encode_event_id(
1593 keys::EVENTS,
1598 &event_id,
1599 );
1600
1601 (event_id, hashed_event_id)
1602 })
1603 .collect::<Vec<_>>();
1604
1605 self.read()
1606 .await?
1607 .with_transaction(move |txn| -> Result<_> {
1608 txn.chunk_large_query_over(
1609 event_ids_and_hashed_event_ids,
1610 None,
1611 move |txn, event_ids_and_hashed_event_ids| {
1612 let query = format!(
1613 "SELECT event_id, chunk_id, position \
1614 FROM event_chunks \
1615 WHERE linked_chunk_id = ? AND event_id IN ({}) \
1616 ORDER BY chunk_id ASC, position ASC",
1617 event_ids_and_hashed_event_ids.host_parameters(),
1618 );
1619
1620 let parameters = params_from_iter(
1621 once(
1623 hashed_linked_chunk_id
1624 .to_sql()
1625 .unwrap(),
1627 )
1628 .chain(
1630 event_ids_and_hashed_event_ids.iter().map(
1631 |(_event_id, hashed_event_id)| {
1632 hashed_event_id
1633 .to_sql()
1634 .unwrap()
1637 },
1638 ),
1639 ),
1640 );
1641
1642 let mut duplicated_events = Vec::new();
1643
1644 for duplicated_event in
1645 txn.prepare(&query)?.query_map(parameters, |row| {
1646 Ok((
1647 row.get::<_, Vec<u8>>(0)?,
1648 row.get::<_, u64>(1)?,
1649 row.get::<_, usize>(2)?,
1650 ))
1651 })?
1652 {
1653 let (duplicated_hashed_event_id, chunk_identifier, index) =
1654 duplicated_event?;
1655
1656 let Some(duplicated_event_id) = event_ids_and_hashed_event_ids
1661 .iter()
1662 .find_map(|(event_id, hashed_event_id)| {
1663 (hashed_event_id.deref() == duplicated_hashed_event_id)
1664 .then_some(event_id.clone())
1665 })
1666 else {
1667 error!(
1668 "Unreachable: found a duplicated event that was not requested"
1669 );
1670 continue;
1671 };
1672
1673 duplicated_events.push((
1674 duplicated_event_id,
1675 Position::new(ChunkIdentifier::new(chunk_identifier), index),
1676 ));
1677 }
1678
1679 Ok(duplicated_events)
1680 },
1681 )
1682 })
1683 .await
1684 }
1685
1686 #[instrument(skip(self, event_id))]
1687 async fn find_event(
1688 &self,
1689 room_id: &RoomId,
1690 event_id: &EventId,
1691 ) -> Result<Option<Event>, Self::Error> {
1692 let _timer = timer!("method");
1693
1694 let encryption = self.encryption.clone();
1695
1696 let hashed_room_id = self.encryption.encode_room_id(keys::EVENTS, room_id);
1697 let hashed_event_id = self.encryption.encode_event_id(keys::EVENTS, event_id);
1698
1699 self.read()
1700 .await?
1701 .with_transaction(move |txn| -> Result<_> {
1702 let Some(event) = txn
1703 .prepare("SELECT content FROM events WHERE event_id = ? AND room_id = ?")?
1704 .query_one((hashed_event_id, hashed_room_id), |row| row.get::<_, Vec<u8>>(0))
1705 .optional()?
1706 else {
1707 return Ok(None);
1709 };
1710
1711 Ok(Some(encryption.decode_event(&event)?))
1712 })
1713 .await
1714 }
1715
1716 #[instrument(skip(self, event_id, filters))]
1717 async fn find_event_relations(
1718 &self,
1719 room_id: &RoomId,
1720 event_id: &EventId,
1721 filters: Option<&[RelationType]>,
1722 ) -> Result<Vec<(Event, Option<Position>)>, Self::Error> {
1723 let _timer = timer!("method");
1724
1725 let hashed_room_id = self.encryption.encode_room_id(keys::EVENTS, room_id);
1726 let hashed_linked_chunk_id =
1727 self.encryption.encode_linked_chunk(keys::LINKED_CHUNKS, &LinkedChunkId::Room(room_id));
1728 let hashed_event_id = self.encryption.encode_event_id(keys::EVENTS, event_id);
1729
1730 let filters = filters.map(ToOwned::to_owned);
1731 let encryption = self.encryption.clone();
1732
1733 self.read()
1734 .await?
1735 .with_transaction(move |txn| -> Result<_> {
1736 find_event_relations_transaction(
1737 &encryption,
1738 hashed_room_id,
1739 hashed_linked_chunk_id,
1740 hashed_event_id,
1741 filters,
1742 txn,
1743 )
1744 })
1745 .await
1746 }
1747
1748 #[instrument(skip(self))]
1749 async fn get_room_events(
1750 &self,
1751 room_id: &RoomId,
1752 event_type: Option<&str>,
1753 session_id: Option<&str>,
1754 ) -> Result<Vec<Event>, Self::Error> {
1755 let _timer = timer!("method");
1756
1757 let encryption = self.encryption.clone();
1758
1759 let hashed_room_id = self.encryption.encode_room_id(keys::EVENTS, room_id);
1760 let hashed_event_type = event_type.map(|e| self.encryption.encode_key(keys::EVENTS, e));
1761 let hashed_session_id = session_id.map(|s| self.encryption.encode_key(keys::EVENTS, s));
1762
1763 self.read()
1764 .await?
1765 .with_transaction(move |txn| -> Result<_> {
1766 #[allow(clippy::redundant_clone)]
1771 let (query, keys) = match (hashed_event_type, hashed_session_id) {
1772 (None, None) => {
1773 ("SELECT content FROM events WHERE room_id = ?", params![hashed_room_id])
1774 }
1775 (None, Some(session_id)) => (
1776 "SELECT content FROM events WHERE room_id = ?1 AND session_id = ?2",
1777 params![hashed_room_id, session_id.to_owned()],
1778 ),
1779 (Some(event_type), None) => (
1780 "SELECT content FROM events WHERE room_id = ? AND event_type = ?",
1781 params![hashed_room_id, event_type.to_owned()]
1782 ),
1783 (Some(event_type), Some(session_id)) => (
1784 "SELECT content FROM events WHERE room_id = ?1 AND event_type = ?2 AND session_id = ?3",
1785 params![hashed_room_id, event_type.to_owned(), session_id.to_owned()],
1786 ),
1787 };
1788
1789 let mut statement = txn.prepare(query)?;
1790
1791 statement
1792 .query_map(keys, |row| row.get::<_, Vec<u8>>(0))?
1793 .map(|maybe_encoded_event| {
1794 encryption.decode_event(&maybe_encoded_event?)
1795 })
1796 .collect::<Result<Vec<_>>>()
1797 })
1798 .await
1799 }
1800
1801 #[instrument(skip(self, event))]
1802 async fn save_event(&self, room_id: &RoomId, event: Event) -> Result<(), Self::Error> {
1803 let _timer = timer!("method");
1804
1805 let Some(event_id) = event.event_id() else {
1806 error!("Trying to save an event with no ID");
1807 return Ok(());
1808 };
1809
1810 let Some(event_type) = event.kind.event_type() else {
1811 error!(%event_id, "Trying to save an event with no event type");
1812 return Ok(());
1813 };
1814
1815 let hashed_event_type = self.encryption.encode_key(keys::EVENTS, event_type);
1816 let hashed_session_id =
1817 event.kind.session_id().map(|s| self.encryption.encode_key(keys::EVENTS, s));
1818
1819 let hashed_room_id = self.encryption.encode_room_id(keys::EVENTS, room_id);
1820 let hashed_event_id = self.encryption.encode_event_id(keys::EVENTS, event_id);
1821 let encoded_event = self.encryption.encode_event(&event)?;
1822
1823 self.write()
1824 .await?
1825 .with_transaction(move |txn| -> Result<_> {
1826 txn.execute(
1827 "INSERT OR REPLACE INTO events(room_id, event_id, event_type, session_id, content, relates_to, rel_type) VALUES (?, ?, ?, ?, ?, ?, ?)",
1828 (
1829 &hashed_room_id,
1830 hashed_event_id,
1831 hashed_event_type,
1832 hashed_session_id,
1833 encoded_event.content,
1834 encoded_event.relates_to,
1835 encoded_event.rel_type
1836 )
1837 )?;
1838
1839 Ok(())
1840 })
1841 .await
1842 }
1843
1844 async fn close(&self) -> Result<(), Self::Error> {
1845 SqliteEventCacheStore::close(self).await
1846 }
1847
1848 async fn reopen(&self) -> Result<(), Self::Error> {
1849 SqliteEventCacheStore::reopen(self).await
1850 }
1851
1852 async fn optimize(&self) -> Result<(), Self::Error> {
1853 Ok(self.vacuum().await?)
1854 }
1855
1856 async fn get_size(&self) -> Result<Option<usize>, Self::Error> {
1857 self.get_db_size().await
1858 }
1859}
1860
1861fn find_event_relations_transaction(
1862 encryption: &Encryption,
1863 hashed_room_id: Key,
1864 hashed_linked_chunk_id: Key,
1865 hashed_event_id: Key,
1866 filters: Option<Vec<RelationType>>,
1867 txn: &Transaction<'_>,
1868) -> Result<Vec<(Event, Option<Position>)>> {
1869 let get_rows = |row: &rusqlite::Row<'_>| {
1870 Ok((
1871 row.get::<_, Vec<u8>>(0)?,
1872 row.get::<_, Option<u64>>(1)?,
1873 row.get::<_, Option<usize>>(2)?,
1874 ))
1875 };
1876
1877 let collect_results = |transaction| {
1879 let mut related = Vec::new();
1880
1881 for result in transaction {
1882 let (event, chunk_id, index): (Vec<u8>, Option<u64>, _) = result?;
1883 let event = encryption.decode_event(&event)?;
1884
1885 let pos = chunk_id
1889 .zip(index)
1890 .map(|(chunk_id, index)| Position::new(ChunkIdentifier::new(chunk_id), index));
1891
1892 related.push((event, pos));
1893 }
1894
1895 Ok(related)
1896 };
1897
1898 if let Some(filters) = filters
1899 && filters.is_empty().not()
1900 {
1901 let query = format!(
1902 "SELECT events.content, event_chunks.chunk_id, event_chunks.position \
1903 FROM events \
1904 LEFT JOIN event_chunks ON events.event_id = event_chunks.event_id AND event_chunks.linked_chunk_id = ? \
1905 WHERE events.relates_to = ? AND events.room_id = ? AND events.rel_type IN ({})",
1906 host_parameters(filters.len())
1907 );
1908
1909 let filter_strings: Vec<_> = filters.iter().map(|f| f.to_string()).collect();
1914 let filters_params: Vec<_> = filter_strings
1915 .iter()
1916 .map(|f| f.to_sql().expect("converting a string to SQL should work"))
1917 .collect();
1918
1919 let parameters = params_from_iter(
1920 [
1921 hashed_linked_chunk_id.to_sql().expect(
1922 "We should be able to convert a hashed linked chunk ID to a SQLite value",
1923 ),
1924 hashed_event_id
1925 .to_sql()
1926 .expect("We should be able to convert an event ID to a SQLite value"),
1927 hashed_room_id
1928 .to_sql()
1929 .expect("We should be able to convert a room ID to a SQLite value"),
1930 ]
1931 .into_iter()
1932 .chain(filters_params),
1933 );
1934
1935 let mut transaction = txn.prepare(&query)?;
1936 let transaction = transaction.query_map(parameters, get_rows)?;
1937
1938 collect_results(transaction)
1939 } else {
1940 let query = "SELECT events.content, event_chunks.chunk_id, event_chunks.position \
1941 FROM events \
1942 LEFT JOIN event_chunks ON events.event_id = event_chunks.event_id AND event_chunks.linked_chunk_id = ? \
1943 WHERE events.relates_to = ? AND events.room_id = ?";
1944 let parameters = (hashed_linked_chunk_id, hashed_event_id, hashed_room_id);
1945
1946 let mut transaction = txn.prepare(query)?;
1947 let transaction = transaction.query_map(parameters, get_rows)?;
1948
1949 collect_results(transaction)
1950 }
1951}
1952
1953async fn with_immediate_transaction<
1958 T: Send + 'static,
1959 F: FnOnce(&Transaction<'_>) -> Result<T, Error> + Send + 'static,
1960>(
1961 this: &SqliteEventCacheStore,
1962 f: F,
1963) -> Result<T, Error> {
1964 this.write()
1965 .await?
1966 .interact(move |conn| -> Result<T, Error> {
1967 conn.set_transaction_behavior(TransactionBehavior::Immediate);
1972
1973 let code = || -> Result<T, Error> {
1974 let txn = conn.transaction()?;
1975 let res = f(&txn)?;
1976 txn.commit()?;
1977 Ok(res)
1978 };
1979
1980 let res = code();
1981
1982 conn.set_transaction_behavior(TransactionBehavior::Deferred);
1985
1986 res
1987 })
1988 .await
1989 .unwrap()
1991}
1992
1993fn insert_chunk(
1994 txn: &Transaction<'_>,
1995 linked_chunk_id: &Key,
1996 previous: Option<u64>,
1997 new: u64,
1998 next: Option<u64>,
1999 type_str: &str,
2000) -> rusqlite::Result<()> {
2001 txn.execute(
2003 r#"
2004 INSERT INTO linked_chunks(id, linked_chunk_id, previous, next, type)
2005 VALUES (?, ?, ?, ?, ?)
2006 "#,
2007 (new, linked_chunk_id, previous, next, type_str),
2008 )?;
2009
2010 if let Some(previous) = previous {
2012 let updated = txn.execute(
2013 r#"
2014 UPDATE linked_chunks
2015 SET next = ?
2016 WHERE id = ? AND linked_chunk_id = ?
2017 "#,
2018 (new, previous, linked_chunk_id),
2019 )?;
2020 if updated < 1 {
2021 return Err(rusqlite::Error::QueryReturnedNoRows);
2022 }
2023 if updated > 1 {
2024 return Err(rusqlite::Error::QueryReturnedMoreThanOneRow);
2025 }
2026 }
2027
2028 if let Some(next) = next {
2030 let updated = txn.execute(
2031 r#"
2032 UPDATE linked_chunks
2033 SET previous = ?
2034 WHERE id = ? AND linked_chunk_id = ?
2035 "#,
2036 (new, next, linked_chunk_id),
2037 )?;
2038 if updated < 1 {
2039 return Err(rusqlite::Error::QueryReturnedNoRows);
2040 }
2041 if updated > 1 {
2042 return Err(rusqlite::Error::QueryReturnedMoreThanOneRow);
2043 }
2044 }
2045
2046 Ok(())
2047}
2048
2049#[cfg(test)]
2050mod tests {
2051 use std::{
2052 path::PathBuf,
2053 sync::{
2054 LazyLock,
2055 atomic::{AtomicU32, Ordering::SeqCst},
2056 },
2057 };
2058
2059 use assert_matches::assert_matches;
2060 use matrix_sdk_base::{
2061 event_cache::store::{
2062 EventCacheStore, EventCacheStoreError, IntoEventCacheStore,
2063 integration_tests::EventCacheStoreIntegrationTests,
2064 },
2065 event_cache_store_integration_tests, event_cache_store_integration_tests_time,
2066 linked_chunk::{ChunkIdentifier, LinkedChunkId, Update},
2067 };
2068 use matrix_sdk_test::{DEFAULT_TEST_ROOM_ID, async_test};
2069 use ruma::{OwnedEventId, event_id};
2070 use tempfile::{TempDir, tempdir};
2071
2072 use super::{SqliteEventCacheStore, keys};
2073 use crate::{SqliteStoreConfig, utils::SqliteAsyncConnExt};
2074
2075 static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2076 static NUM: AtomicU32 = AtomicU32::new(0);
2077
2078 fn new_event_cache_store_workspace() -> PathBuf {
2079 let name = NUM.fetch_add(1, SeqCst).to_string();
2080 TMP_DIR.path().join(name)
2081 }
2082
2083 async fn get_event_cache_store() -> Result<SqliteEventCacheStore, EventCacheStoreError> {
2084 let tmpdir_path = new_event_cache_store_workspace();
2085
2086 tracing::info!("using event cache store @ {}", tmpdir_path.to_str().unwrap());
2087
2088 Ok(SqliteEventCacheStore::open(tmpdir_path.to_str().unwrap(), None).await.unwrap())
2089 }
2090
2091 event_cache_store_integration_tests!();
2092 event_cache_store_integration_tests_time!();
2093
2094 #[async_test]
2095 async fn test_encryption_encode_decode_thread_id_roundtrip() {
2096 let store = get_event_cache_store().await.expect("creating cache store failed");
2097 let thread_id: OwnedEventId = event_id!("$event").to_owned();
2098
2099 let encoded_thread_id = store.encryption.encode_thread_id(&thread_id).unwrap();
2100 let decoded_thread_id: OwnedEventId =
2101 store.encryption.decode_thread_id(&encoded_thread_id).unwrap();
2102
2103 assert_eq!(thread_id, decoded_thread_id);
2104 }
2105
2106 #[async_test]
2107 async fn test_pool_size() {
2108 let tmpdir_path = new_event_cache_store_workspace();
2109 let store_open_config = SqliteStoreConfig::new(tmpdir_path).pool_max_size(42);
2110
2111 let store = SqliteEventCacheStore::open_with_config(&store_open_config).await.unwrap();
2112
2113 assert_eq!(store.pool_max_size().await, Some(42));
2114 }
2115
2116 #[async_test]
2117 async fn test_linked_chunk_remove_chunk() {
2118 let store = get_event_cache_store().await.expect("creating cache store failed");
2119
2120 store.clone().into_event_cache_store().test_linked_chunk_remove_chunk().await;
2122
2123 let gaps = store
2125 .read()
2126 .await
2127 .unwrap()
2128 .with_transaction(|txn| -> rusqlite::Result<_> {
2129 let mut gaps = Vec::new();
2130 for data in txn
2131 .prepare("SELECT chunk_id FROM gap_chunks ORDER BY chunk_id")?
2132 .query_map((), |row| row.get::<_, u64>(0))?
2133 {
2134 gaps.push(data?);
2135 }
2136 Ok(gaps)
2137 })
2138 .await
2139 .unwrap();
2140
2141 assert_eq!(gaps, vec![42, 44]);
2144 }
2145
2146 #[async_test]
2147 async fn test_linked_chunk_remove_item() {
2148 let store = get_event_cache_store().await.expect("creating cache store failed");
2149
2150 store.clone().into_event_cache_store().test_linked_chunk_remove_item().await;
2152
2153 let room_id = *DEFAULT_TEST_ROOM_ID;
2154 let hashed_linked_chunk_id = store
2155 .encryption
2156 .encode_linked_chunk(keys::LINKED_CHUNKS, &LinkedChunkId::Room(room_id));
2157
2158 let num_rows: u64 = store
2160 .read()
2161 .await
2162 .unwrap()
2163 .with_transaction(move |txn| {
2164 txn.query_one(
2165 "SELECT COUNT(*) FROM event_chunks WHERE chunk_id = 42 AND linked_chunk_id = ? AND position IN (2, 3, 4)",
2166 (hashed_linked_chunk_id,),
2167 |row| row.get(0),
2168 )
2169 })
2170 .await
2171 .unwrap();
2172 assert_eq!(num_rows, 3);
2173 }
2174
2175 #[async_test]
2176 async fn test_linked_chunk_clear() {
2177 let store = get_event_cache_store().await.expect("creating cache store failed");
2178
2179 store.clone().into_event_cache_store().test_linked_chunk_clear().await;
2181
2182 store
2184 .read()
2185 .await
2186 .unwrap()
2187 .with_transaction(|txn| -> rusqlite::Result<_> {
2188 let num_gaps = txn
2189 .prepare("SELECT COUNT(chunk_id) FROM gap_chunks ORDER BY chunk_id")?
2190 .query_one((), |row| row.get::<_, u64>(0))?;
2191 assert_eq!(num_gaps, 0);
2192
2193 let num_events = txn
2194 .prepare("SELECT COUNT(event_id) FROM event_chunks ORDER BY chunk_id")?
2195 .query_one((), |row| row.get::<_, u64>(0))?;
2196 assert_eq!(num_events, 0);
2197
2198 Ok(())
2199 })
2200 .await
2201 .unwrap();
2202 }
2203
2204 #[async_test]
2205 async fn test_linked_chunk_update_is_a_transaction() {
2206 let store = get_event_cache_store().await.expect("creating cache store failed");
2207
2208 let room_id = *DEFAULT_TEST_ROOM_ID;
2209 let linked_chunk_id = LinkedChunkId::Room(room_id);
2210
2211 let err = store
2214 .handle_linked_chunk_updates(
2215 linked_chunk_id,
2216 vec![
2217 Update::NewItemsChunk {
2218 previous: None,
2219 new: ChunkIdentifier::new(42),
2220 next: None,
2221 },
2222 Update::NewItemsChunk {
2223 previous: None,
2224 new: ChunkIdentifier::new(42),
2225 next: None,
2226 },
2227 ],
2228 )
2229 .await
2230 .unwrap_err();
2231
2232 assert_matches!(err, crate::error::Error::Sqlite(err) => {
2234 assert_matches!(err.sqlite_error_code(), Some(rusqlite::ErrorCode::ConstraintViolation));
2235 });
2236
2237 let chunks = store.load_all_chunks(linked_chunk_id).await.unwrap();
2241 assert!(chunks.is_empty());
2242 }
2243}
2244
2245#[cfg(test)]
2246mod encrypted_tests {
2247 use std::sync::{
2248 LazyLock,
2249 atomic::{AtomicU32, Ordering::SeqCst},
2250 };
2251
2252 use matrix_sdk_base::{
2253 event_cache::store::{EventCacheStore, EventCacheStoreError},
2254 event_cache_store_integration_tests, event_cache_store_integration_tests_time,
2255 };
2256 use matrix_sdk_test::{async_test, event_factory::EventFactory};
2257 use ruma::{
2258 event_id,
2259 events::{relation::RelationType, room::message::RoomMessageEventContentWithoutRelation},
2260 room_id, user_id,
2261 };
2262 use tempfile::{TempDir, tempdir};
2263
2264 use super::SqliteEventCacheStore;
2265
2266 static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2267 static NUM: AtomicU32 = AtomicU32::new(0);
2268
2269 async fn get_event_cache_store() -> Result<SqliteEventCacheStore, EventCacheStoreError> {
2270 let name = NUM.fetch_add(1, SeqCst).to_string();
2271 let tmpdir_path = TMP_DIR.path().join(name);
2272
2273 tracing::info!("using event cache store @ {}", tmpdir_path.to_str().unwrap());
2274
2275 Ok(SqliteEventCacheStore::open(
2276 tmpdir_path.to_str().unwrap(),
2277 Some("default_test_password"),
2278 )
2279 .await
2280 .unwrap())
2281 }
2282
2283 event_cache_store_integration_tests!();
2284 event_cache_store_integration_tests_time!();
2285
2286 #[async_test]
2287 async fn test_no_sqlite_injection_in_find_event_relations() {
2288 let room_id = room_id!("!test:localhost");
2289 let another_room_id = room_id!("!r1:matrix.org");
2290 let sender = user_id!("@alice:localhost");
2291
2292 let store = get_event_cache_store()
2293 .await
2294 .expect("We should be able to create a new, empty, event cache store");
2295
2296 let f = EventFactory::new().room(room_id).sender(sender);
2297
2298 let event_id = event_id!("$DO_NOT_FIND_ME:matrix.org");
2300 let event = f.text_msg("DO NOT FIND").event_id(event_id).into_event();
2301
2302 let edit_id = event_id!("$find_me:matrix.org");
2304 let edit = f
2305 .text_msg("Find me")
2306 .event_id(edit_id)
2307 .edit(event_id, RoomMessageEventContentWithoutRelation::text_plain("jebote"))
2308 .into_event();
2309
2310 let f = f.room(another_room_id);
2312
2313 let another_event_id = event_id!("$DO_NOT_FIND_ME_EITHER:matrix.org");
2314 let another_event =
2315 f.text_msg("DO NOT FIND ME EITHER").event_id(another_event_id).into_event();
2316
2317 store.save_event(room_id, event).await.unwrap();
2319 store.save_event(room_id, edit).await.unwrap();
2320 store.save_event(another_room_id, another_event).await.unwrap();
2321
2322 let filter = Some(vec![RelationType::Replacement, "x\") OR 1=1; --".into()]);
2326
2327 let results = store
2329 .find_event_relations(room_id, event_id, filter.as_deref())
2330 .await
2331 .expect("We should be able to attempt to find event relations");
2332
2333 similar_asserts::assert_eq!(
2336 results.len(),
2337 1,
2338 "We should only have loaded events for the first room {results:#?}"
2339 );
2340
2341 let (found_event, _) = &results[0];
2343 assert_eq!(
2344 found_event.event_id(),
2345 Some(edit_id),
2346 "The single event we found should be the edit event"
2347 );
2348 }
2349}
2350
2351#[cfg(test)]
2352mod close_reopen_tests {
2353 use std::sync::{
2354 LazyLock,
2355 atomic::{AtomicU32, Ordering::SeqCst},
2356 };
2357
2358 use matrix_sdk_base::{event_cache::store::EventCacheStore, linked_chunk::LinkedChunkId};
2359 use matrix_sdk_test::{DEFAULT_TEST_ROOM_ID, async_test};
2360 use tempfile::{TempDir, tempdir};
2361
2362 use super::SqliteEventCacheStore;
2363
2364 static TMP_DIR: LazyLock<TempDir> = LazyLock::new(|| tempdir().unwrap());
2365 static NUM: AtomicU32 = AtomicU32::new(0);
2366
2367 async fn new_store() -> SqliteEventCacheStore {
2368 let name = NUM.fetch_add(1, SeqCst).to_string();
2369 let tmpdir_path = TMP_DIR.path().join(name);
2370 SqliteEventCacheStore::open(tmpdir_path, None).await.unwrap()
2371 }
2372
2373 #[async_test]
2374 async fn test_close_completes_without_timeout() {
2375 let store = new_store().await;
2376
2377 let start = std::time::Instant::now();
2379 store.close().await.unwrap();
2380 let elapsed = start.elapsed();
2381
2382 assert!(
2383 elapsed < std::time::Duration::from_secs(2),
2384 "close() took {elapsed:?}, expected < 2s (no timeout)"
2385 );
2386
2387 let guard = store.connections.lock().await;
2389 assert!(guard.is_none(), "connections should be None after close");
2390 }
2391
2392 #[async_test]
2393 async fn test_reopen_restores_connections() {
2394 let store = new_store().await;
2395
2396 store.close().await.unwrap();
2397
2398 {
2399 let guard = store.connections.lock().await;
2400 assert!(guard.is_none());
2401 }
2402
2403 store.reopen().await.unwrap();
2404
2405 {
2406 let guard = store.connections.lock().await;
2407 assert!(guard.is_some(), "connections should be Some after reopen");
2408 }
2409 }
2410
2411 #[async_test]
2412 async fn test_close_is_idempotent() {
2413 let store = new_store().await;
2414
2415 store.close().await.unwrap();
2416 store.close().await.unwrap();
2418
2419 let guard = store.connections.lock().await;
2420 assert!(guard.is_none());
2421 }
2422
2423 #[async_test]
2424 async fn test_reopen_is_idempotent() {
2425 let store = new_store().await;
2426
2427 store.reopen().await.unwrap();
2429
2430 let guard = store.connections.lock().await;
2431 assert!(guard.is_some());
2432 }
2433
2434 #[async_test]
2435 async fn test_read_fails_when_closed() {
2436 let store = new_store().await;
2437 store.close().await.unwrap();
2438
2439 let err = store.load_all_chunks(LinkedChunkId::Room(*DEFAULT_TEST_ROOM_ID)).await;
2440 assert!(err.is_err(), "read should fail when closed");
2441
2442 let err_msg = err.unwrap_err().to_string();
2443 assert!(err_msg.contains("closed"), "error should mention 'closed', got: {err_msg}");
2444 }
2445
2446 #[async_test]
2447 async fn test_write_fails_when_closed() {
2448 let store = new_store().await;
2449 store.close().await.unwrap();
2450
2451 let err = store.try_take_leased_lock(1000, "test_lock", "holder").await;
2452 assert!(err.is_err(), "write should fail when closed");
2453
2454 let err_msg = err.unwrap_err().to_string();
2455 assert!(err_msg.contains("closed"), "error should mention 'closed', got: {err_msg}");
2456 }
2457
2458 #[async_test]
2459 async fn test_data_persists_across_close_reopen() {
2460 let store = new_store().await;
2461
2462 let result = store.try_take_leased_lock(60_000, "test_lock", "holder").await.unwrap();
2464 assert!(result.is_some(), "should have acquired the lock");
2465
2466 store.close().await.unwrap();
2468 store.reopen().await.unwrap();
2469
2470 let result = store.try_take_leased_lock(60_000, "test_lock", "other_holder").await.unwrap();
2472 assert!(result.is_none(), "lock should still be held by the original holder after reopen");
2473 }
2474
2475 #[async_test]
2476 async fn test_multiple_close_reopen_cycles() {
2477 let store = new_store().await;
2478
2479 for _ in 0..5 {
2480 store.close().await.unwrap();
2481 store.reopen().await.unwrap();
2482
2483 let result = store.load_all_chunks(LinkedChunkId::Room(*DEFAULT_TEST_ROOM_ID)).await;
2485 assert!(result.is_ok(), "store should work after close/reopen cycle");
2486 }
2487 }
2488
2489 #[async_test]
2490 async fn test_pool_is_fully_drained_after_close() {
2491 let store = new_store().await;
2492
2493 let _ = store.load_all_chunks(LinkedChunkId::Room(*DEFAULT_TEST_ROOM_ID)).await;
2495 let _ = store.load_all_chunks(LinkedChunkId::Room(*DEFAULT_TEST_ROOM_ID)).await;
2496
2497 store.close().await.unwrap();
2498
2499 let guard = store.connections.lock().await;
2502 assert!(guard.is_none(), "all connections should be released after close");
2503 }
2504
2505 #[async_test]
2506 async fn test_operations_work_immediately_after_reopen() {
2507 let store = new_store().await;
2508
2509 store.close().await.unwrap();
2510 store.reopen().await.unwrap();
2511
2512 let result = store.load_all_chunks(LinkedChunkId::Room(*DEFAULT_TEST_ROOM_ID)).await;
2514 assert!(result.is_ok(), "read should succeed immediately after reopen");
2515
2516 let result = store.try_take_leased_lock(1000, "test_lock", "holder").await;
2518 assert!(result.is_ok(), "write should succeed immediately after reopen");
2519 }
2520
2521 #[async_test]
2522 async fn test_close_waits_for_held_read_connection_to_drain() {
2523 let store = new_store().await;
2524
2525 let held_conn = store.read().await.unwrap();
2527
2528 let store_clone = store.clone();
2531 let close_handle = tokio::spawn(async move {
2532 store_clone.close().await.unwrap();
2533 });
2534
2535 tokio::time::sleep(std::time::Duration::from_millis(200)).await;
2537
2538 assert!(!close_handle.is_finished(), "close should be waiting for the held connection");
2540
2541 drop(held_conn);
2543
2544 let timeout = tokio::time::timeout(std::time::Duration::from_secs(3), close_handle).await;
2546 assert!(timeout.is_ok(), "close should complete after the held connection is released");
2547 timeout.unwrap().unwrap();
2548
2549 let guard = store.connections.lock().await;
2551 assert!(guard.is_none(), "connections should be None after close");
2552 }
2553}