1use crate::schema::*;
2use async_trait::async_trait;
3use bytes::Bytes;
4use diesel::prelude::*;
5use diesel::r2d2::{ConnectionManager, Pool};
6use diesel::result::{DatabaseErrorKind, Error as DieselError};
7use diesel::sqlite::SqliteConnection;
8use diesel::upsert::excluded;
9use diesel_migrations::{EmbeddedMigrations, MigrationHarness, embed_migrations};
10use log::warn;
11use std::sync::Arc;
12use std::time::Duration;
13use wacore::appstate::hash::HashState;
14use wacore::appstate::processor::AppStateMutationMAC;
15use wacore::libsignal::protocol::{KeyPair, PrivateKey, PublicKey};
16use wacore::store::Device as CoreDevice;
17use wacore::store::error::{Result, StoreError};
18use wacore::store::traits::*;
19
20enum DieselOrStore {
24 Diesel(DieselError),
25 Store(StoreError),
26}
27
28impl From<DieselOrStore> for StoreError {
29 fn from(e: DieselOrStore) -> Self {
30 match e {
31 DieselOrStore::Diesel(e) => StoreError::Database(Box::new(e)),
32 DieselOrStore::Store(e) => e,
33 }
34 }
35}
36
37fn is_retriable_sqlite_error(error: &DieselError) -> bool {
43 match error {
44 DieselError::DatabaseError(DatabaseErrorKind::Unknown, info) => {
45 let msg = info.message();
46 msg.contains("locked") || msg.contains("busy")
47 }
48 _ => false,
49 }
50}
51
52const MIGRATIONS: EmbeddedMigrations = embed_migrations!("migrations");
53
54pub(crate) type SqlitePool = Pool<ConnectionManager<SqliteConnection>>;
55
56#[derive(Queryable, Selectable)]
62#[diesel(table_name = device)]
63#[allow(dead_code)]
64struct DeviceRow {
65 id: i32,
66 lid: String,
67 pn: String,
68 registration_id: i32,
69 noise_key: Vec<u8>,
70 identity_key: Vec<u8>,
71 signed_pre_key: Vec<u8>,
72 signed_pre_key_id: i32,
73 signed_pre_key_signature: Vec<u8>,
74 adv_secret_key: Vec<u8>,
75 account: Option<Vec<u8>>,
76 push_name: String,
77 app_version_primary: i32,
78 app_version_secondary: i32,
79 app_version_tertiary: i64,
80 app_version_last_fetched_ms: i64,
81 edge_routing_info: Option<Vec<u8>>,
82 props_hash: Option<String>,
83 next_pre_key_id: i32,
84 nct_salt: Option<Vec<u8>>,
85 server_has_prekeys: bool,
86 server_cert_chain: Option<Vec<u8>>,
87 login_counter: i32,
88 first_unupload_pre_key_id: i32,
89 lid_migrated: bool,
90 last_signed_pre_key_rotation_ms: i64,
91 read_receipts_disabled: bool,
92}
93
94const ID_PARAM_CHUNK: usize = 900;
96const MSG_SECRET_INSERT_CHUNK_SIZE: usize = 100;
99
100#[derive(Clone)]
102pub(crate) struct ReadPool {
103 pub(crate) pool: SqlitePool,
104 pub(crate) semaphore: Arc<tokio::sync::Semaphore>,
107}
108
109#[derive(Clone)]
110pub struct SqliteStore {
111 pub(crate) pool: SqlitePool,
112 pub(crate) db_semaphore: Arc<tokio::sync::Semaphore>,
113 pub(crate) reads: Option<ReadPool>,
126 pub(crate) database_path: String,
127 device_id: i32,
128}
129
130#[derive(Debug, Clone, Copy)]
132pub enum Synchronous {
133 Off,
134 Normal,
135 Full,
136}
137
138impl Synchronous {
139 fn as_pragma(self) -> &'static str {
140 match self {
141 Synchronous::Off => "OFF",
142 Synchronous::Normal => "NORMAL",
143 Synchronous::Full => "FULL",
144 }
145 }
146}
147
148pub type ConnectionInitHook = Arc<
159 dyn Fn(
160 &mut SqliteConnection,
161 ) -> std::result::Result<(), Box<dyn std::error::Error + Send + Sync>>
162 + Send
163 + Sync,
164>;
165
166#[derive(Clone)]
174pub struct SqliteStoreConfig {
175 pub pool_size: u32,
185 pub read_pool_size: u32,
199 pub cache_size_kib: u32,
201 pub mmap_size: Option<u64>,
211 pub busy_timeout: Duration,
213 pub synchronous: Synchronous,
215 pub thread_pool: Option<Arc<scheduled_thread_pool::ScheduledThreadPool>>,
218 pub connection_init: Option<ConnectionInitHook>,
222}
223
224impl Default for SqliteStoreConfig {
225 fn default() -> Self {
226 Self {
227 pool_size: 1,
228 read_pool_size: 0,
229 cache_size_kib: 512,
230 mmap_size: None,
231 busy_timeout: Duration::from_secs(30),
232 synchronous: Synchronous::Normal,
233 thread_pool: None,
234 connection_init: None,
235 }
236 }
237}
238
239impl SqliteStoreConfig {
240 pub fn with_read_pool_size(mut self, n: u32) -> Self {
244 self.read_pool_size = n;
245 self
246 }
247
248 pub fn with_mmap_size(mut self, bytes: u64) -> Self {
252 self.mmap_size = Some(bytes);
253 self
254 }
255
256 pub fn with_connection_init<F>(mut self, hook: F) -> Self
278 where
279 F: Fn(
280 &mut SqliteConnection,
281 ) -> std::result::Result<(), Box<dyn std::error::Error + Send + Sync>>
282 + Send
283 + Sync
284 + 'static,
285 {
286 self.connection_init = Some(Arc::new(hook));
287 self
288 }
289}
290
291#[derive(Clone)]
292struct ConnectionOptions {
293 cache_size_kib: u32,
294 mmap_size: Option<u64>,
295 busy_timeout_ms: u64,
296 synchronous: Synchronous,
297 connection_init: Option<ConnectionInitHook>,
298 query_only: bool,
302}
303
304impl std::fmt::Debug for ConnectionOptions {
305 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
306 f.debug_struct("ConnectionOptions")
307 .field("cache_size_kib", &self.cache_size_kib)
308 .field("mmap_size", &self.mmap_size)
309 .field("busy_timeout_ms", &self.busy_timeout_ms)
310 .field("synchronous", &self.synchronous)
311 .field(
312 "connection_init",
313 &self.connection_init.as_ref().map(|_| ()),
314 )
315 .finish()
316 }
317}
318
319impl diesel::r2d2::CustomizeConnection<SqliteConnection, diesel::r2d2::Error>
320 for ConnectionOptions
321{
322 fn on_acquire(
323 &self,
324 conn: &mut SqliteConnection,
325 ) -> std::result::Result<(), diesel::r2d2::Error> {
326 if let Some(init) = &self.connection_init {
329 init(conn).map_err(|e| {
330 diesel::r2d2::Error::QueryError(diesel::result::Error::QueryBuilderError(e))
331 })?;
332 }
333 let mut pragmas = vec![
336 format!("PRAGMA busy_timeout = {};", self.busy_timeout_ms),
337 format!("PRAGMA synchronous = {};", self.synchronous.as_pragma()),
338 format!("PRAGMA cache_size = -{};", self.cache_size_kib),
339 "PRAGMA temp_store = memory;".to_string(),
340 "PRAGMA foreign_keys = ON;".to_string(),
341 ];
342 if let Some(mmap_size) = self.mmap_size.filter(|&n| n > 0) {
345 pragmas.push(format!("PRAGMA mmap_size = {mmap_size};"));
346 }
347 if self.query_only {
350 pragmas.push("PRAGMA query_only = 1;".to_string());
351 }
352 for pragma in pragmas {
353 diesel::sql_query(pragma)
354 .execute(conn)
355 .map_err(diesel::r2d2::Error::QueryError)?;
356 }
357 Ok(())
358 }
359}
360
361fn parse_database_path(database_url: &str) -> Result<String> {
362 if database_url == ":memory:" {
364 return Err(StoreError::InvalidConfig(
365 "Snapshot not supported for in-memory databases".to_string(),
366 ));
367 }
368
369 let path = database_url
371 .split(['?', '#'])
372 .next()
373 .unwrap_or(database_url);
374
375 let path = path.trim_start_matches("sqlite://");
377
378 if path == ":memory:" || path.starts_with(":memory:?") {
380 return Err(StoreError::InvalidConfig(
381 "Snapshot not supported for in-memory databases".to_string(),
382 ));
383 }
384
385 Ok(path.to_string())
386}
387
388fn is_shared_cache(database_url: &str) -> bool {
394 let Some((_, query)) = database_url.split_once('?') else {
395 return false;
396 };
397 if !database_url.starts_with("file:") {
398 return false;
399 }
400 query
401 .split('#')
402 .next()
403 .unwrap_or(query)
404 .split('&')
405 .filter_map(|param| param.split_once('='))
406 .find(|(key, _)| *key == "cache")
407 .is_some_and(|(_, value)| value.eq_ignore_ascii_case("shared"))
408}
409
410fn shared_r2d2_thread_pool() -> Arc<scheduled_thread_pool::ScheduledThreadPool> {
416 static POOL: std::sync::OnceLock<Arc<scheduled_thread_pool::ScheduledThreadPool>> =
417 std::sync::OnceLock::new();
418 POOL.get_or_init(|| {
419 Arc::new(
420 scheduled_thread_pool::ScheduledThreadPool::builder()
421 .num_threads(2)
422 .thread_name_pattern("r2d2-shared-{}")
423 .build(),
424 )
425 })
426 .clone()
427}
428
429impl SqliteStore {
430 pub async fn new(database_url: &str) -> std::result::Result<Self, StoreError> {
432 Self::build(database_url, 1, SqliteStoreConfig::default()).await
433 }
434
435 pub async fn with_config(
438 database_url: &str,
439 config: SqliteStoreConfig,
440 ) -> std::result::Result<Self, StoreError> {
441 Self::build(database_url, 1, config).await
442 }
443
444 pub async fn new_for_device(
445 database_url: &str,
446 device_id: i32,
447 ) -> std::result::Result<Self, StoreError> {
448 Self::build(database_url, device_id, SqliteStoreConfig::default()).await
449 }
450
451 pub async fn with_config_for_device(
453 database_url: &str,
454 device_id: i32,
455 config: SqliteStoreConfig,
456 ) -> std::result::Result<Self, StoreError> {
457 Self::build(database_url, device_id, config).await
458 }
459
460 async fn build(
461 database_url: &str,
462 device_id: i32,
463 config: SqliteStoreConfig,
464 ) -> std::result::Result<Self, StoreError> {
465 let manager = ConnectionManager::<SqliteConnection>::new(database_url);
466 let pool_size = config.pool_size.max(1);
470 let read_pool_size = config.read_pool_size;
471 let thread_pool = config.thread_pool.unwrap_or_else(shared_r2d2_thread_pool);
472 let read_thread_pool = Arc::clone(&thread_pool);
473
474 let options = ConnectionOptions {
475 cache_size_kib: config.cache_size_kib,
476 mmap_size: config.mmap_size,
477 busy_timeout_ms: if config.busy_timeout.is_zero() {
481 0
482 } else {
483 config.busy_timeout.as_millis().clamp(1, i32::MAX as u128) as u64
484 },
485 synchronous: config.synchronous,
486 connection_init: config.connection_init,
487 query_only: false,
488 };
489 let read_options = ConnectionOptions {
490 query_only: true,
491 ..options.clone()
492 };
493
494 let db_url = database_url.to_string();
498 let (pool, journal_mode) = tokio::task::spawn_blocking(
499 move || -> std::result::Result<(SqlitePool, String), StoreError> {
500 let pool = Pool::builder()
505 .max_size(pool_size)
506 .test_on_check_out(false)
507 .thread_pool(thread_pool)
508 .connection_customizer(Box::new(options))
509 .build(manager)
510 .map_err(|e| StoreError::Connection(Box::new(e)))?;
511
512 let mut conn = pool
513 .get()
514 .map_err(|e| StoreError::Connection(Box::new(e)))?;
515 #[derive(diesel::QueryableByName)]
519 struct JournalMode {
520 #[diesel(sql_type = diesel::sql_types::Text)]
521 journal_mode: String,
522 }
523 let journal_mode = diesel::sql_query("PRAGMA journal_mode = WAL;")
524 .get_result::<JournalMode>(&mut conn)
525 .map_err(|e| StoreError::Database(Box::new(e)))?
526 .journal_mode;
527 conn.run_pending_migrations(MIGRATIONS)
528 .map_err(StoreError::Migration)?;
529
530 Ok((pool, journal_mode))
531 },
532 )
533 .await
534 .map_err(|e| StoreError::Database(Box::new(e)))??;
535
536 let wal = journal_mode.eq_ignore_ascii_case("wal");
541 let shared_cache = is_shared_cache(&db_url);
546 let declined = if !wal {
547 Some(format!("journal_mode is '{journal_mode}', not WAL"))
548 } else if shared_cache {
549 Some("the URI opts into shared cache, whose table locks block the writer".to_string())
550 } else {
551 None
552 };
553 if read_pool_size > 0
554 && let Some(reason) = &declined
555 {
556 log::warn!("sqlite-storage: read_pool_size={read_pool_size} ignored, {reason}");
557 }
558 let reads = if read_pool_size > 0 && declined.is_none() {
559 let manager = ConnectionManager::<SqliteConnection>::new(&db_url);
560 let pool = tokio::task::spawn_blocking(
561 move || -> std::result::Result<SqlitePool, StoreError> {
562 Pool::builder()
563 .max_size(read_pool_size)
564 .test_on_check_out(false)
565 .thread_pool(read_thread_pool)
566 .connection_customizer(Box::new(read_options))
567 .build(manager)
568 .map_err(|e| StoreError::Connection(Box::new(e)))
569 },
570 )
571 .await
572 .map_err(|e| StoreError::Database(Box::new(e)))??;
573 Some(ReadPool {
574 pool,
575 semaphore: Arc::new(tokio::sync::Semaphore::new(read_pool_size as usize)),
576 })
577 } else {
578 None
579 };
580
581 let database_path = parse_database_path(database_url)?;
582
583 Ok(Self {
584 pool,
585 db_semaphore: Arc::new(tokio::sync::Semaphore::new(pool_size as usize)),
586 reads,
587 database_path,
588 device_id,
589 })
590 }
591
592 pub fn device_id(&self) -> i32 {
593 self.device_id
594 }
595
596 async fn with_semaphore<F, T>(&self, f: F) -> Result<T>
597 where
598 F: FnOnce() -> Result<T> + Send + 'static,
599 T: Send + 'static,
600 {
601 let permit = self
602 .db_semaphore
603 .clone()
604 .acquire_owned()
605 .await
606 .map_err(|e| StoreError::Database(Box::new(e)))?;
607 let result = tokio::task::spawn_blocking(move || {
608 let res = f();
609 drop(permit);
610 res
611 })
612 .await
613 .map_err(|e| StoreError::Database(Box::new(e)))??;
614 Ok(result)
615 }
616
617 async fn with_retry<F, T>(&self, op_name: &str, make_op: F) -> Result<T>
621 where
622 F: Fn() -> Box<
623 dyn FnOnce(&mut SqliteConnection) -> std::result::Result<T, DieselError> + Send,
624 >,
625 T: Send + 'static,
626 {
627 const MAX_RETRIES: u32 = 5;
628
629 for attempt in 0..=MAX_RETRIES {
630 let permit = self
631 .db_semaphore
632 .clone()
633 .acquire_owned()
634 .await
635 .map_err(|e| StoreError::Database(Box::new(e)))?;
636
637 let pool = self.pool.clone();
638 let op = make_op();
639
640 let result =
641 tokio::task::spawn_blocking(move || -> std::result::Result<T, DieselOrStore> {
642 let _permit = permit;
643 let mut conn = pool
644 .get()
645 .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
646 op(&mut conn).map_err(DieselOrStore::Diesel)
647 })
648 .await;
649
650 match result {
651 Ok(Ok(val)) => return Ok(val),
652 Ok(Err(DieselOrStore::Diesel(ref e)))
653 if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
654 {
655 let delay_ms = 10u64 * (1u64 << attempt.min(4));
656 if attempt >= 1 {
659 warn!(
660 "{op_name} busy/locked, retry {}/{} in {delay_ms}ms: {e}",
661 attempt + 1,
662 MAX_RETRIES + 1
663 );
664 }
665 tokio::time::sleep(Duration::from_millis(delay_ms)).await;
666 }
667 Ok(Err(e)) => return Err(e.into()),
668 Err(e) => return Err(StoreError::Database(Box::new(e))),
669 }
670 }
671
672 Err(StoreError::RetriesExhausted {
673 op: op_name.to_string(),
674 })
675 }
676
677 fn serialize_keypair(&self, key_pair: &KeyPair) -> Result<Vec<u8>> {
678 let mut bytes = Vec::with_capacity(64);
679 bytes.extend_from_slice(key_pair.private_key.serialize());
680 bytes.extend_from_slice(key_pair.public_key.public_key_bytes());
681 Ok(bytes)
682 }
683
684 fn deserialize_keypair(&self, bytes: &[u8]) -> Result<KeyPair> {
685 if bytes.len() != 64 {
686 return Err(StoreError::Validation(format!(
687 "Invalid KeyPair length: {}",
688 bytes.len()
689 )));
690 }
691
692 let private_key = PrivateKey::deserialize(&bytes[0..32])
693 .map_err(|e| StoreError::Serialization(Box::new(e)))?;
694 let public_key = PublicKey::from_djb_public_key_bytes(&bytes[32..64])
695 .map_err(|e| StoreError::Serialization(Box::new(e)))?;
696
697 Ok(KeyPair::new(public_key, private_key))
698 }
699
700 pub async fn save_device_data_for_device(
701 &self,
702 device_id: i32,
703 device_data: &CoreDevice,
704 ) -> Result<()> {
705 let noise_key_data: Arc<[u8]> = self.serialize_keypair(&device_data.noise_key)?.into();
707 let identity_key_data: Arc<[u8]> =
708 self.serialize_keypair(&device_data.identity_key)?.into();
709 let signed_pre_key_data: Arc<[u8]> =
710 self.serialize_keypair(&device_data.signed_pre_key)?.into();
711 let account_data: Option<Arc<[u8]>> = device_data
712 .account
713 .as_ref()
714 .map(|a| Arc::from(wacore::store::device::account_serde::to_bytes(a)));
715 let registration_id = device_data.registration_id as i32;
716 let signed_pre_key_id = device_data.signed_pre_key_id as i32;
717 let signed_pre_key_signature: Arc<[u8]> =
718 Arc::from(&device_data.signed_pre_key_signature[..]);
719 let adv_secret_key: Arc<[u8]> = Arc::from(&device_data.adv_secret_key[..]);
720 let push_name: Arc<str> = Arc::from(device_data.push_name.as_str());
721 let app_version_primary = device_data.app_version_primary as i32;
722 let app_version_secondary = device_data.app_version_secondary as i32;
723 let app_version_tertiary = device_data.app_version_tertiary as i64;
724 let app_version_last_fetched_ms = device_data.app_version_last_fetched_ms;
725 let edge_routing_info: Option<Arc<[u8]>> =
726 device_data.edge_routing_info.as_deref().map(Arc::from);
727 let props_hash: Option<Arc<str>> = device_data.props_hash.as_deref().map(Arc::from);
728 let next_pre_key_id = device_data.next_pre_key_id as i32;
729 let first_unupload_pre_key_id = device_data.first_unupload_pre_key_id as i32;
730 let server_has_prekeys = device_data.server_has_prekeys;
731 let nct_salt: Option<Arc<[u8]>> = device_data.nct_salt.as_deref().map(Arc::from);
732 let server_cert_chain: Option<Arc<[u8]>> = device_data
733 .server_cert_chain
734 .as_ref()
735 .map(|chain| Arc::from(crate::wire::encode_server_cert_chain(chain)));
736 let login_counter = device_data.login_counter;
737 let lid_migrated = device_data.lid_migrated;
738 let last_signed_pre_key_rotation_ms = device_data.last_signed_pre_key_rotation_ms;
739 let read_receipts_disabled = device_data.read_receipts_disabled;
740 let new_lid: Arc<str> = Arc::from(
741 device_data
742 .lid
743 .as_ref()
744 .map(|j| j.to_string())
745 .unwrap_or_default()
746 .as_str(),
747 );
748 let new_pn: Arc<str> = Arc::from(
749 device_data
750 .pn
751 .as_ref()
752 .map(|j| j.to_string())
753 .unwrap_or_default()
754 .as_str(),
755 );
756
757 self.with_retry("save_device_data", || {
758 let noise_key_data = Arc::clone(&noise_key_data);
759 let identity_key_data = Arc::clone(&identity_key_data);
760 let signed_pre_key_data = Arc::clone(&signed_pre_key_data);
761 let account_data = account_data.clone();
762 let signed_pre_key_signature = Arc::clone(&signed_pre_key_signature);
763 let adv_secret_key = Arc::clone(&adv_secret_key);
764 let push_name = Arc::clone(&push_name);
765 let edge_routing_info = edge_routing_info.clone();
766 let props_hash = props_hash.clone();
767 let nct_salt = nct_salt.clone();
768 let server_cert_chain = server_cert_chain.clone();
769 let new_lid = Arc::clone(&new_lid);
770 let new_pn = Arc::clone(&new_pn);
771
772 Box::new(move |conn: &mut SqliteConnection| {
773 diesel::insert_into(device::table)
774 .values((
775 device::id.eq(device_id),
776 device::lid.eq(&*new_lid),
777 device::pn.eq(&*new_pn),
778 device::registration_id.eq(registration_id),
779 device::noise_key.eq(&*noise_key_data),
780 device::identity_key.eq(&*identity_key_data),
781 device::signed_pre_key.eq(&*signed_pre_key_data),
782 device::signed_pre_key_id.eq(signed_pre_key_id),
783 device::signed_pre_key_signature.eq(&*signed_pre_key_signature),
784 device::adv_secret_key.eq(&*adv_secret_key),
785 device::account.eq(account_data.as_deref()),
786 device::push_name.eq(&*push_name),
787 device::app_version_primary.eq(app_version_primary),
788 device::app_version_secondary.eq(app_version_secondary),
789 device::app_version_tertiary.eq(app_version_tertiary),
790 device::app_version_last_fetched_ms.eq(app_version_last_fetched_ms),
791 device::edge_routing_info.eq(edge_routing_info.as_deref()),
792 device::props_hash.eq(props_hash.as_deref()),
793 device::next_pre_key_id.eq(next_pre_key_id),
794 device::first_unupload_pre_key_id.eq(first_unupload_pre_key_id),
795 device::server_has_prekeys.eq(server_has_prekeys),
796 device::nct_salt.eq(nct_salt.as_deref()),
797 device::server_cert_chain.eq(server_cert_chain.as_deref()),
798 device::login_counter.eq(login_counter),
799 device::lid_migrated.eq(lid_migrated),
800 device::last_signed_pre_key_rotation_ms.eq(last_signed_pre_key_rotation_ms),
801 device::read_receipts_disabled.eq(read_receipts_disabled),
802 ))
803 .on_conflict(device::id)
804 .do_update()
805 .set((
806 device::lid.eq(excluded(device::lid)),
807 device::pn.eq(excluded(device::pn)),
808 device::registration_id.eq(excluded(device::registration_id)),
809 device::noise_key.eq(excluded(device::noise_key)),
810 device::identity_key.eq(excluded(device::identity_key)),
811 device::signed_pre_key.eq(excluded(device::signed_pre_key)),
812 device::signed_pre_key_id.eq(excluded(device::signed_pre_key_id)),
813 device::signed_pre_key_signature
814 .eq(excluded(device::signed_pre_key_signature)),
815 device::adv_secret_key.eq(excluded(device::adv_secret_key)),
816 device::account.eq(excluded(device::account)),
817 device::push_name.eq(excluded(device::push_name)),
818 device::app_version_primary.eq(excluded(device::app_version_primary)),
819 device::app_version_secondary.eq(excluded(device::app_version_secondary)),
820 device::app_version_tertiary.eq(excluded(device::app_version_tertiary)),
821 device::app_version_last_fetched_ms
822 .eq(excluded(device::app_version_last_fetched_ms)),
823 device::edge_routing_info.eq(excluded(device::edge_routing_info)),
824 device::props_hash.eq(excluded(device::props_hash)),
825 device::next_pre_key_id.eq(excluded(device::next_pre_key_id)),
826 device::first_unupload_pre_key_id
827 .eq(excluded(device::first_unupload_pre_key_id)),
828 device::server_has_prekeys.eq(excluded(device::server_has_prekeys)),
829 device::nct_salt.eq(excluded(device::nct_salt)),
830 device::server_cert_chain.eq(excluded(device::server_cert_chain)),
831 device::login_counter.eq(excluded(device::login_counter)),
832 device::lid_migrated.eq(excluded(device::lid_migrated)),
833 device::last_signed_pre_key_rotation_ms
834 .eq(excluded(device::last_signed_pre_key_rotation_ms)),
835 device::read_receipts_disabled.eq(excluded(device::read_receipts_disabled)),
836 ))
837 .execute(conn)
838 .map(|_| ())
839 })
840 })
841 .await
842 }
843
844 pub async fn create_new_device(&self) -> Result<i32> {
845 let device_id = self.device_id;
846 let new_device = wacore::store::Device::new();
847
848 let noise_key_data: Arc<[u8]> = self.serialize_keypair(&new_device.noise_key)?.into();
849 let identity_key_data: Arc<[u8]> = self.serialize_keypair(&new_device.identity_key)?.into();
850 let signed_pre_key_data: Arc<[u8]> =
851 self.serialize_keypair(&new_device.signed_pre_key)?.into();
852 let registration_id = new_device.registration_id as i32;
853 let signed_pre_key_id = new_device.signed_pre_key_id as i32;
854 let signed_pre_key_signature: Arc<[u8]> =
855 Arc::from(&new_device.signed_pre_key_signature[..]);
856 let adv_secret_key: Arc<[u8]> = Arc::from(&new_device.adv_secret_key[..]);
857 let push_name: Arc<str> = Arc::from(new_device.push_name.as_str());
858 let app_version_primary = new_device.app_version_primary as i32;
859 let app_version_secondary = new_device.app_version_secondary as i32;
860 let app_version_tertiary = new_device.app_version_tertiary as i64;
861 let app_version_last_fetched_ms = new_device.app_version_last_fetched_ms;
862 let next_pre_key_id = new_device.next_pre_key_id as i32;
863 let first_unupload_pre_key_id = new_device.first_unupload_pre_key_id as i32;
864 let server_has_prekeys = new_device.server_has_prekeys;
865 let last_signed_pre_key_rotation_ms = new_device.last_signed_pre_key_rotation_ms;
866
867 self.with_retry("create_new_device", || {
868 let noise_key_data = Arc::clone(&noise_key_data);
869 let identity_key_data = Arc::clone(&identity_key_data);
870 let signed_pre_key_data = Arc::clone(&signed_pre_key_data);
871 let signed_pre_key_signature = Arc::clone(&signed_pre_key_signature);
872 let adv_secret_key = Arc::clone(&adv_secret_key);
873 let push_name = Arc::clone(&push_name);
874
875 Box::new(move |conn: &mut SqliteConnection| {
876 diesel::insert_into(device::table)
877 .values((
878 device::id.eq(device_id),
879 device::lid.eq(""),
880 device::pn.eq(""),
881 device::registration_id.eq(registration_id),
882 device::noise_key.eq(&*noise_key_data),
883 device::identity_key.eq(&*identity_key_data),
884 device::signed_pre_key.eq(&*signed_pre_key_data),
885 device::signed_pre_key_id.eq(signed_pre_key_id),
886 device::signed_pre_key_signature.eq(&*signed_pre_key_signature),
887 device::adv_secret_key.eq(&*adv_secret_key),
888 device::account.eq(None::<&[u8]>),
889 device::push_name.eq(&*push_name),
890 device::app_version_primary.eq(app_version_primary),
891 device::app_version_secondary.eq(app_version_secondary),
892 device::app_version_tertiary.eq(app_version_tertiary),
893 device::app_version_last_fetched_ms.eq(app_version_last_fetched_ms),
894 device::edge_routing_info.eq(None::<&[u8]>),
895 device::props_hash.eq(None::<&str>),
896 device::next_pre_key_id.eq(next_pre_key_id),
897 device::first_unupload_pre_key_id.eq(first_unupload_pre_key_id),
898 device::server_has_prekeys.eq(server_has_prekeys),
899 device::nct_salt.eq(None::<&[u8]>),
900 device::server_cert_chain.eq(None::<&[u8]>),
901 device::login_counter.eq(0i32),
902 device::lid_migrated.eq(false),
903 device::last_signed_pre_key_rotation_ms.eq(last_signed_pre_key_rotation_ms),
904 device::read_receipts_disabled.eq(false),
905 ))
906 .execute(conn)
907 .map(|_| device_id)
908 })
909 })
910 .await
911 }
912
913 pub async fn device_exists(&self, device_id: i32) -> Result<bool> {
914 use crate::schema::device;
915
916 let pool = self.pool.clone();
917 tokio::task::spawn_blocking(move || -> Result<bool> {
918 let mut conn = pool
919 .get()
920 .map_err(|e| StoreError::Connection(Box::new(e)))?;
921
922 let count: i64 = device::table
923 .filter(device::id.eq(device_id))
924 .count()
925 .get_result(&mut conn)
926 .map_err(|e| StoreError::Database(Box::new(e)))?;
927
928 Ok(count > 0)
929 })
930 .await
931 .map_err(|e| StoreError::Database(Box::new(e)))?
932 }
933
934 pub async fn load_device_data_for_device(&self, device_id: i32) -> Result<Option<CoreDevice>> {
935 use crate::schema::device;
936
937 let pool = self.pool.clone();
938 let row = tokio::task::spawn_blocking(move || -> Result<Option<DeviceRow>> {
939 let mut conn = pool
940 .get()
941 .map_err(|e| StoreError::Connection(Box::new(e)))?;
942 let result = device::table
943 .filter(device::id.eq(device_id))
944 .first::<DeviceRow>(&mut conn)
945 .optional()
946 .map_err(|e| StoreError::Database(Box::new(e)))?;
947 Ok(result)
948 })
949 .await
950 .map_err(|e| StoreError::Database(Box::new(e)))??;
951
952 if let Some(row) = row {
953 let pn = if !row.pn.is_empty() {
954 row.pn.parse().ok()
955 } else {
956 None
957 };
958 let lid = if !row.lid.is_empty() {
959 row.lid.parse().ok()
960 } else {
961 None
962 };
963
964 let noise_key = self.deserialize_keypair(&row.noise_key)?;
965 let identity_key = self.deserialize_keypair(&row.identity_key)?;
966 let signed_pre_key = self.deserialize_keypair(&row.signed_pre_key)?;
967
968 let signed_pre_key_signature: [u8; 64] =
969 row.signed_pre_key_signature.try_into().map_err(|_| {
970 StoreError::Validation("Invalid signed_pre_key_signature length".to_string())
971 })?;
972
973 let adv_secret_key: [u8; 32] = row
974 .adv_secret_key
975 .try_into()
976 .map_err(|_| StoreError::Validation("Invalid adv_secret_key length".to_string()))?;
977
978 let account = row
979 .account
980 .map(|data| {
981 wacore::store::device::account_serde::from_bytes(&data)
982 .map_err(|e| StoreError::Serialization(Box::new(e)))
983 })
984 .transpose()?;
985
986 Ok(Some(CoreDevice {
987 pn,
988 lid,
989 registration_id: row.registration_id as u32,
990 noise_key,
991 identity_key,
992 signed_pre_key,
993 signed_pre_key_id: row.signed_pre_key_id as u32,
994 signed_pre_key_signature,
995 adv_secret_key,
996 account: account.map(Arc::new),
997 push_name: row.push_name,
998 app_version_primary: row.app_version_primary as u32,
999 app_version_secondary: row.app_version_secondary as u32,
1000 app_version_tertiary: row.app_version_tertiary.try_into().unwrap_or(0u32),
1001 app_version_last_fetched_ms: row.app_version_last_fetched_ms,
1002 device_props: Arc::new(wacore::store::device::DEVICE_PROPS.clone()),
1003 client_profile: wacore::client_profile::ClientProfile::web(),
1004 edge_routing_info: row.edge_routing_info,
1005 props_hash: row.props_hash,
1006 next_pre_key_id: row.next_pre_key_id as u32,
1007 first_unupload_pre_key_id: row.first_unupload_pre_key_id as u32,
1008 server_has_prekeys: row.server_has_prekeys,
1009 nct_salt: row.nct_salt,
1010 nct_salt_sync_seen: false,
1011 server_cert_chain: row
1012 .server_cert_chain
1013 .as_deref()
1014 .and_then(|bytes| {
1015 match crate::wire::decode_server_cert_chain(bytes) {
1021 Ok(chain) => Some(chain),
1022 Err(e) => {
1023 log::warn!(
1024 "device {} server_cert_chain blob ({} bytes) failed to decode: {e}; \
1025 dropping cache, next connect will use XX",
1026 self.device_id,
1027 bytes.len(),
1028 );
1029 None
1030 }
1031 }
1032 }),
1033 login_counter: row.login_counter,
1034 lid_migrated: row.lid_migrated,
1035 last_signed_pre_key_rotation_ms: row.last_signed_pre_key_rotation_ms,
1036 read_receipts_disabled: row.read_receipts_disabled,
1037 }))
1038 } else {
1039 Ok(None)
1040 }
1041 }
1042
1043 pub async fn put_identity_for_device(
1044 &self,
1045 address: &str,
1046 key: [u8; 32],
1047 device_id: i32,
1048 ) -> Result<()> {
1049 let pool = self.pool.clone();
1050 let db_semaphore = self.db_semaphore.clone();
1051 let address_owned = address.to_string();
1052 let key_vec = key.to_vec();
1053
1054 const MAX_RETRIES: u32 = 5;
1055
1056 for attempt in 0..=MAX_RETRIES {
1057 let permit = db_semaphore
1058 .clone()
1059 .acquire_owned()
1060 .await
1061 .map_err(|e| StoreError::Database(Box::new(e)))?;
1062
1063 let pool_clone = pool.clone();
1064 let address_clone = address_owned.clone();
1065 let key_clone = key_vec.clone();
1066
1067 let result =
1068 tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
1069 let mut conn = pool_clone
1070 .get()
1071 .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
1072 diesel::insert_into(identities::table)
1073 .values((
1074 identities::address.eq(address_clone),
1075 identities::key.eq(&key_clone[..]),
1076 identities::device_id.eq(device_id),
1077 ))
1078 .on_conflict((identities::address, identities::device_id))
1079 .do_update()
1080 .set(identities::key.eq(&key_clone[..]))
1081 .execute(&mut conn)
1082 .map_err(DieselOrStore::Diesel)?;
1083 Ok(())
1084 })
1085 .await;
1086
1087 drop(permit);
1088
1089 match result {
1090 Ok(Ok(())) => return Ok(()),
1091 Ok(Err(DieselOrStore::Diesel(ref e)))
1092 if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
1093 {
1094 let delay_ms = 10 * 2u64.pow(attempt);
1095 warn!(
1096 "Identity write failed (attempt {}/{}): {e}. Retrying in {delay_ms}ms...",
1097 attempt + 1,
1098 MAX_RETRIES + 1,
1099 );
1100 tokio::time::sleep(Duration::from_millis(delay_ms)).await;
1101 continue;
1102 }
1103 Ok(Err(e)) => return Err(e.into()),
1104 Err(e) => return Err(StoreError::Database(Box::new(e))),
1105 }
1106 }
1107
1108 Err(StoreError::RetriesExhausted {
1109 op: format!("identity_write (after {} attempts)", MAX_RETRIES + 1),
1110 })
1111 }
1112
1113 pub async fn delete_identity_for_device(&self, address: &str, device_id: i32) -> Result<()> {
1114 let pool = self.pool.clone();
1115 let address_owned = address.to_string();
1116
1117 tokio::task::spawn_blocking(move || -> Result<()> {
1118 let mut conn = pool
1119 .get()
1120 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1121 diesel::delete(
1122 identities::table
1123 .filter(identities::address.eq(address_owned))
1124 .filter(identities::device_id.eq(device_id)),
1125 )
1126 .execute(&mut conn)
1127 .map_err(|e| StoreError::Database(Box::new(e)))?;
1128 Ok(())
1129 })
1130 .await
1131 .map_err(|e| StoreError::Database(Box::new(e)))??;
1132
1133 Ok(())
1134 }
1135
1136 pub async fn load_identity_for_device(
1137 &self,
1138 address: &str,
1139 device_id: i32,
1140 ) -> Result<Option<Vec<u8>>> {
1141 let pool = self.pool.clone();
1142 let address = address.to_string();
1143 let result = self
1144 .with_semaphore(move || -> Result<Option<Vec<u8>>> {
1145 let mut conn = pool
1146 .get()
1147 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1148 let res: Option<Vec<u8>> = identities::table
1149 .select(identities::key)
1150 .filter(identities::address.eq(address))
1151 .filter(identities::device_id.eq(device_id))
1152 .first(&mut conn)
1153 .optional()
1154 .map_err(|e| StoreError::Database(Box::new(e)))?;
1155 Ok(res)
1156 })
1157 .await?;
1158
1159 Ok(result)
1160 }
1161
1162 pub async fn get_session_for_device(
1163 &self,
1164 address: &str,
1165 device_id: i32,
1166 ) -> Result<Option<Vec<u8>>> {
1167 let pool = self.pool.clone();
1168 let address_for_query = address.to_string();
1169 let result = self
1170 .with_semaphore(move || -> Result<Option<Vec<u8>>> {
1171 let mut conn = pool
1172 .get()
1173 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1174 let res: Option<Vec<u8>> = sessions::table
1175 .select(sessions::record)
1176 .filter(sessions::address.eq(address_for_query.clone()))
1177 .filter(sessions::device_id.eq(device_id))
1178 .first(&mut conn)
1179 .optional()
1180 .map_err(|e| StoreError::Database(Box::new(e)))?;
1181
1182 Ok(res)
1183 })
1184 .await?;
1185
1186 Ok(result)
1187 }
1188
1189 pub async fn put_session_for_device(
1190 &self,
1191 address: &str,
1192 session: &[u8],
1193 device_id: i32,
1194 ) -> Result<()> {
1195 let pool = self.pool.clone();
1196 let db_semaphore = self.db_semaphore.clone();
1197 let address_owned = address.to_string();
1198 let session_vec = session.to_vec();
1199
1200 const MAX_RETRIES: u32 = 5;
1201
1202 for attempt in 0..=MAX_RETRIES {
1203 let permit = db_semaphore
1204 .clone()
1205 .acquire_owned()
1206 .await
1207 .map_err(|e| StoreError::Database(Box::new(e)))?;
1208
1209 let pool_clone = pool.clone();
1210 let address_clone = address_owned.clone();
1211 let session_clone = session_vec.clone();
1212
1213 let result =
1214 tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
1215 let mut conn = pool_clone
1216 .get()
1217 .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
1218 diesel::insert_into(sessions::table)
1219 .values((
1220 sessions::address.eq(address_clone),
1221 sessions::record.eq(&session_clone),
1222 sessions::device_id.eq(device_id),
1223 ))
1224 .on_conflict((sessions::address, sessions::device_id))
1225 .do_update()
1226 .set(sessions::record.eq(&session_clone))
1227 .execute(&mut conn)
1228 .map_err(DieselOrStore::Diesel)?;
1229 Ok(())
1230 })
1231 .await;
1232
1233 drop(permit);
1234
1235 match result {
1236 Ok(Ok(())) => return Ok(()),
1237 Ok(Err(DieselOrStore::Diesel(ref e)))
1238 if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
1239 {
1240 let delay_ms = 10 * 2u64.pow(attempt);
1241 warn!(
1242 "Session write failed (attempt {}/{}): {e}. Retrying in {delay_ms}ms...",
1243 attempt + 1,
1244 MAX_RETRIES + 1,
1245 );
1246 tokio::time::sleep(Duration::from_millis(delay_ms)).await;
1247 continue;
1248 }
1249 Ok(Err(e)) => return Err(e.into()),
1250 Err(e) => return Err(StoreError::Database(Box::new(e))),
1251 }
1252 }
1253
1254 Err(StoreError::RetriesExhausted {
1255 op: format!("session_write (after {} attempts)", MAX_RETRIES + 1),
1256 })
1257 }
1258
1259 pub async fn delete_session_for_device(&self, address: &str, device_id: i32) -> Result<()> {
1260 let pool = self.pool.clone();
1261 let address_owned = address.to_string();
1262
1263 tokio::task::spawn_blocking(move || -> Result<()> {
1264 let mut conn = pool
1265 .get()
1266 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1267 diesel::delete(
1268 sessions::table
1269 .filter(sessions::address.eq(address_owned))
1270 .filter(sessions::device_id.eq(device_id)),
1271 )
1272 .execute(&mut conn)
1273 .map_err(|e| StoreError::Database(Box::new(e)))?;
1274 Ok(())
1275 })
1276 .await
1277 .map_err(|e| StoreError::Database(Box::new(e)))??;
1278
1279 Ok(())
1280 }
1281
1282 pub async fn put_sender_key_for_device(
1283 &self,
1284 address: &str,
1285 record: &[u8],
1286 device_id: i32,
1287 ) -> Result<()> {
1288 let pool = self.pool.clone();
1289 let address = address.to_string();
1290 let record_vec = record.to_vec();
1291 tokio::task::spawn_blocking(move || -> Result<()> {
1292 let mut conn = pool
1293 .get()
1294 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1295 diesel::insert_into(sender_keys::table)
1296 .values((
1297 sender_keys::address.eq(address),
1298 sender_keys::record.eq(&record_vec),
1299 sender_keys::device_id.eq(device_id),
1300 ))
1301 .on_conflict((sender_keys::address, sender_keys::device_id))
1302 .do_update()
1303 .set(sender_keys::record.eq(&record_vec))
1304 .execute(&mut conn)
1305 .map_err(|e| StoreError::Database(Box::new(e)))?;
1306 Ok(())
1307 })
1308 .await
1309 .map_err(|e| StoreError::Database(Box::new(e)))??;
1310 Ok(())
1311 }
1312
1313 pub async fn get_sender_key_for_device(
1314 &self,
1315 address: &str,
1316 device_id: i32,
1317 ) -> Result<Option<Vec<u8>>> {
1318 let pool = self.pool.clone();
1319 let address = address.to_string();
1320 tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
1321 let mut conn = pool
1322 .get()
1323 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1324 let res: Option<Vec<u8>> = sender_keys::table
1325 .select(sender_keys::record)
1326 .filter(sender_keys::address.eq(address))
1327 .filter(sender_keys::device_id.eq(device_id))
1328 .first(&mut conn)
1329 .optional()
1330 .map_err(|e| StoreError::Database(Box::new(e)))?;
1331 Ok(res)
1332 })
1333 .await
1334 .map_err(|e| StoreError::Database(Box::new(e)))?
1335 }
1336
1337 pub async fn delete_sender_key_for_device(&self, address: &str, device_id: i32) -> Result<()> {
1338 let pool = self.pool.clone();
1339 let address = address.to_string();
1340 tokio::task::spawn_blocking(move || -> Result<()> {
1341 let mut conn = pool
1342 .get()
1343 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1344 diesel::delete(
1345 sender_keys::table
1346 .filter(sender_keys::address.eq(address))
1347 .filter(sender_keys::device_id.eq(device_id)),
1348 )
1349 .execute(&mut conn)
1350 .map_err(|e| StoreError::Database(Box::new(e)))?;
1351 Ok(())
1352 })
1353 .await
1354 .map_err(|e| StoreError::Database(Box::new(e)))??;
1355 Ok(())
1356 }
1357
1358 pub async fn get_app_state_sync_key_for_device(
1359 &self,
1360 key_id: &[u8],
1361 device_id: i32,
1362 ) -> Result<Option<AppStateSyncKey>> {
1363 let pool = self.pool.clone();
1364 let key_id = key_id.to_vec();
1365 let res: Option<Vec<u8>> =
1366 tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
1367 let mut conn = pool
1368 .get()
1369 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1370 let res: Option<Vec<u8>> = app_state_keys::table
1371 .select(app_state_keys::key_data)
1372 .filter(app_state_keys::key_id.eq(&key_id))
1373 .filter(app_state_keys::device_id.eq(device_id))
1374 .first(&mut conn)
1375 .optional()
1376 .map_err(|e| StoreError::Database(Box::new(e)))?;
1377 Ok(res)
1378 })
1379 .await
1380 .map_err(|e| StoreError::Database(Box::new(e)))??;
1381
1382 if let Some(data) = res {
1383 match crate::wire::decode_app_state_sync_key(&data) {
1387 Ok(key) => Ok(Some(key)),
1388 Err(e) => {
1389 warn!(
1390 "app_state_sync_key blob ({} bytes) failed to decode: {e}; \
1391 treating as absent, key will be re-requested",
1392 data.len()
1393 );
1394 Ok(None)
1395 }
1396 }
1397 } else {
1398 Ok(None)
1399 }
1400 }
1401
1402 pub async fn set_app_state_sync_key_for_device(
1403 &self,
1404 key_id: &[u8],
1405 key: AppStateSyncKey,
1406 device_id: i32,
1407 ) -> Result<()> {
1408 let pool = self.pool.clone();
1409 let key_id = key_id.to_vec();
1410 let data = crate::wire::encode_app_state_sync_key(&key);
1411 tokio::task::spawn_blocking(move || -> Result<()> {
1412 let mut conn = pool
1413 .get()
1414 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1415 diesel::insert_into(app_state_keys::table)
1416 .values((
1417 app_state_keys::key_id.eq(&key_id),
1418 app_state_keys::key_data.eq(&data),
1419 app_state_keys::device_id.eq(device_id),
1420 ))
1421 .on_conflict((app_state_keys::key_id, app_state_keys::device_id))
1422 .do_update()
1423 .set(app_state_keys::key_data.eq(&data))
1424 .execute(&mut conn)
1425 .map_err(|e| StoreError::Database(Box::new(e)))?;
1426 Ok(())
1427 })
1428 .await
1429 .map_err(|e| StoreError::Database(Box::new(e)))??;
1430 Ok(())
1431 }
1432
1433 pub async fn get_latest_app_state_sync_key_id_for_device(
1434 &self,
1435 device_id: i32,
1436 ) -> Result<Option<Vec<u8>>> {
1437 let pool = self.pool.clone();
1438 let res: Option<Vec<u8>> =
1439 tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
1440 let mut conn = pool
1441 .get()
1442 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1443 let candidates: Vec<(Vec<u8>, Vec<u8>)> = app_state_keys::table
1450 .select((app_state_keys::key_id, app_state_keys::key_data))
1451 .filter(app_state_keys::device_id.eq(device_id))
1452 .order(app_state_keys::key_id.desc())
1453 .load(&mut conn)
1454 .map_err(|e| StoreError::Database(Box::new(e)))?;
1455 let res = candidates
1456 .into_iter()
1457 .find(|(_, data)| crate::wire::decode_app_state_sync_key(data).is_ok())
1458 .map(|(key_id, _)| key_id);
1459 Ok(res)
1460 })
1461 .await
1462 .map_err(|e| StoreError::Database(Box::new(e)))??;
1463 Ok(res)
1464 }
1465
1466 pub async fn get_app_state_version_for_device(
1467 &self,
1468 name: &str,
1469 device_id: i32,
1470 ) -> Result<HashState> {
1471 let pool = self.pool.clone();
1472 let name = name.to_string();
1473 let res: Option<Vec<u8>> =
1474 tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
1475 let mut conn = pool
1476 .get()
1477 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1478 let res: Option<Vec<u8>> = app_state_versions::table
1479 .select(app_state_versions::state_data)
1480 .filter(app_state_versions::name.eq(name))
1481 .filter(app_state_versions::device_id.eq(device_id))
1482 .first(&mut conn)
1483 .optional()
1484 .map_err(|e| StoreError::Database(Box::new(e)))?;
1485 Ok(res)
1486 })
1487 .await
1488 .map_err(|e| StoreError::Database(Box::new(e)))??;
1489
1490 if let Some(data) = res {
1491 match crate::wire::decode_hash_state(&data) {
1494 Ok(state) => Ok(state),
1495 Err(e) => {
1496 warn!(
1497 "app_state_version blob ({} bytes) failed to decode: {e}; \
1498 resetting to default, collection will re-sync from 0",
1499 data.len()
1500 );
1501 Ok(HashState::default())
1502 }
1503 }
1504 } else {
1505 Ok(HashState::default())
1506 }
1507 }
1508
1509 pub async fn set_app_state_version_for_device(
1510 &self,
1511 name: &str,
1512 state: HashState,
1513 device_id: i32,
1514 ) -> Result<()> {
1515 let name = name.to_string();
1516 let data = crate::wire::encode_hash_state(&state);
1517 self.with_retry("set_app_state_version", || {
1518 let name = name.clone();
1519 let data = data.clone();
1520 Box::new(move |conn: &mut SqliteConnection| {
1521 diesel::insert_into(app_state_versions::table)
1522 .values((
1523 app_state_versions::name.eq(&name),
1524 app_state_versions::state_data.eq(&data),
1525 app_state_versions::device_id.eq(device_id),
1526 ))
1527 .on_conflict((app_state_versions::name, app_state_versions::device_id))
1528 .do_update()
1529 .set(app_state_versions::state_data.eq(&data))
1530 .execute(conn)?;
1531 Ok(())
1532 })
1533 })
1534 .await
1535 }
1536
1537 pub async fn put_app_state_mutation_macs_for_device(
1538 &self,
1539 name: &str,
1540 version: u64,
1541 mutations: &[AppStateMutationMAC],
1542 device_id: i32,
1543 ) -> Result<()> {
1544 if mutations.is_empty() {
1545 return Ok(());
1546 }
1547 let name = name.to_string();
1548 let mutations: Vec<AppStateMutationMAC> = mutations.to_vec();
1549 self.with_retry("put_app_state_mutation_macs", || {
1550 let name = name.clone();
1551 let mutations = mutations.clone();
1552 Box::new(move |conn: &mut SqliteConnection| {
1553 let records: Vec<_> = mutations
1554 .iter()
1555 .map(|m| {
1556 (
1557 app_state_mutation_macs::name.eq(&name),
1558 app_state_mutation_macs::version.eq(version as i64),
1559 app_state_mutation_macs::index_mac.eq(&m.index_mac),
1560 app_state_mutation_macs::value_mac.eq(&m.value_mac),
1561 app_state_mutation_macs::device_id.eq(device_id),
1562 )
1563 })
1564 .collect();
1565
1566 const CHUNK_SIZE: usize = 100;
1569
1570 for chunk in records.chunks(CHUNK_SIZE) {
1571 diesel::insert_into(app_state_mutation_macs::table)
1572 .values(chunk)
1573 .on_conflict((
1574 app_state_mutation_macs::name,
1575 app_state_mutation_macs::index_mac,
1576 app_state_mutation_macs::device_id,
1577 ))
1578 .do_update()
1579 .set((
1580 app_state_mutation_macs::version
1581 .eq(excluded(app_state_mutation_macs::version)),
1582 app_state_mutation_macs::value_mac
1583 .eq(excluded(app_state_mutation_macs::value_mac)),
1584 ))
1585 .execute(conn)?;
1586 }
1587 Ok(())
1588 })
1589 })
1590 .await
1591 }
1592
1593 pub async fn delete_app_state_mutation_macs_for_device(
1594 &self,
1595 name: &str,
1596 index_macs: &[Vec<u8>],
1597 device_id: i32,
1598 ) -> Result<()> {
1599 if index_macs.is_empty() {
1600 return Ok(());
1601 }
1602 let name = name.to_string();
1603 let index_macs: Vec<Vec<u8>> = index_macs.to_vec();
1604 self.with_retry("delete_app_state_mutation_macs", || {
1605 let name = name.clone();
1606 let index_macs = index_macs.clone();
1607 Box::new(move |conn: &mut SqliteConnection| {
1608 const CHUNK_SIZE: usize = 500;
1611
1612 for chunk in index_macs.chunks(CHUNK_SIZE) {
1613 diesel::delete(
1614 app_state_mutation_macs::table.filter(
1615 app_state_mutation_macs::name
1616 .eq(&name)
1617 .and(app_state_mutation_macs::index_mac.eq_any(chunk))
1618 .and(app_state_mutation_macs::device_id.eq(device_id)),
1619 ),
1620 )
1621 .execute(conn)?;
1622 }
1623 Ok(())
1624 })
1625 })
1626 .await
1627 }
1628
1629 pub async fn get_app_state_mutation_mac_for_device(
1630 &self,
1631 name: &str,
1632 index_mac: &[u8],
1633 device_id: i32,
1634 ) -> Result<Option<Vec<u8>>> {
1635 let pool = self.pool.clone();
1636 let name = name.to_string();
1637 let index_mac = index_mac.to_vec();
1638 tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
1639 let mut conn = pool
1640 .get()
1641 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1642 let res: Option<Vec<u8>> = app_state_mutation_macs::table
1643 .select(app_state_mutation_macs::value_mac)
1644 .filter(app_state_mutation_macs::name.eq(&name))
1645 .filter(app_state_mutation_macs::index_mac.eq(&index_mac))
1646 .filter(app_state_mutation_macs::device_id.eq(device_id))
1647 .first(&mut conn)
1648 .optional()
1649 .map_err(|e| StoreError::Database(Box::new(e)))?;
1650 Ok(res)
1651 })
1652 .await
1653 .map_err(|e| StoreError::Database(Box::new(e)))?
1654 }
1655
1656 pub async fn get_app_state_mutation_macs_batch_for_device(
1660 &self,
1661 name: &str,
1662 index_macs: &[[u8; 32]],
1663 device_id: i32,
1664 ) -> Result<std::collections::HashMap<[u8; 32], Vec<u8>>> {
1665 if index_macs.is_empty() {
1666 return Ok(std::collections::HashMap::new());
1667 }
1668 let pool = self.pool.clone();
1669 let name = name.to_string();
1670 let index_macs: Vec<[u8; 32]> = index_macs.to_vec();
1671 tokio::task::spawn_blocking(
1672 move || -> Result<std::collections::HashMap<[u8; 32], Vec<u8>>> {
1673 let mut conn = pool
1674 .get()
1675 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1676 let mut out = std::collections::HashMap::with_capacity(index_macs.len());
1677 const CHUNK_SIZE: usize = 500;
1678 for chunk in index_macs.chunks(CHUNK_SIZE) {
1679 let chunk_slices: Vec<&[u8]> = chunk.iter().map(|m| m.as_slice()).collect();
1680 let rows: Vec<(Vec<u8>, Vec<u8>)> = app_state_mutation_macs::table
1681 .select((
1682 app_state_mutation_macs::index_mac,
1683 app_state_mutation_macs::value_mac,
1684 ))
1685 .filter(app_state_mutation_macs::name.eq(&name))
1686 .filter(app_state_mutation_macs::index_mac.eq_any(chunk_slices))
1687 .filter(app_state_mutation_macs::device_id.eq(device_id))
1688 .load(&mut conn)
1689 .map_err(|e| StoreError::Database(Box::new(e)))?;
1690 out.extend(rows.into_iter().filter_map(|(k, v)| {
1693 <[u8; 32]>::try_from(k.as_slice()).ok().map(|k| (k, v))
1694 }));
1695 }
1696 Ok(out)
1697 },
1698 )
1699 .await
1700 .map_err(|e| StoreError::Database(Box::new(e)))?
1701 }
1702}
1703
1704#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
1705#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
1706impl SignalStore for SqliteStore {
1707 async fn put_identity(&self, address: &str, key: [u8; 32]) -> Result<()> {
1708 self.put_identity_for_device(address, key, self.device_id)
1709 .await
1710 }
1711
1712 async fn put_identities_batch(&self, identities: &[(Arc<str>, [u8; 32])]) -> Result<()> {
1713 if identities.is_empty() {
1714 return Ok(());
1715 }
1716
1717 let device_id = self.device_id;
1718 let batch = Arc::new(identities.to_vec());
1721 self.with_retry("put_identities_batch", || {
1722 let batch = batch.clone();
1723 Box::new(move |conn: &mut SqliteConnection| {
1724 conn.transaction(|conn| {
1725 for (address, key) in batch.iter() {
1726 diesel::insert_into(identities::table)
1727 .values((
1728 identities::address.eq(address.as_ref()),
1729 identities::key.eq(&key[..]),
1730 identities::device_id.eq(device_id),
1731 ))
1732 .on_conflict((identities::address, identities::device_id))
1733 .do_update()
1734 .set(identities::key.eq(&key[..]))
1735 .execute(conn)?;
1736 }
1737 Ok(())
1738 })
1739 })
1740 })
1741 .await
1742 }
1743
1744 async fn load_identity(&self, address: &str) -> Result<Option<[u8; 32]>> {
1745 let blob = self
1746 .load_identity_for_device(address, self.device_id)
1747 .await?;
1748 match blob {
1749 None => Ok(None),
1750 Some(v) => Ok(Some(v.try_into().map_err(|v: Vec<u8>| {
1751 StoreError::Validation(format!(
1752 "identity key for '{}' has invalid length {} (expected 32)",
1753 address,
1754 v.len()
1755 ))
1756 })?)),
1757 }
1758 }
1759
1760 async fn delete_identity(&self, address: &str) -> Result<()> {
1761 self.delete_identity_for_device(address, self.device_id)
1762 .await
1763 }
1764
1765 async fn get_session(&self, address: &str) -> Result<Option<Bytes>> {
1766 Ok(self
1767 .get_session_for_device(address, self.device_id)
1768 .await?
1769 .map(Bytes::from))
1770 }
1771
1772 async fn has_session(&self, address: &str) -> Result<bool> {
1773 let pool = self.pool.clone();
1774 let device_id = self.device_id;
1775 let address_owned = address.to_string();
1776 self.with_semaphore(move || -> Result<bool> {
1777 let mut conn = pool
1778 .get()
1779 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1780 let exists = diesel::select(diesel::dsl::exists(
1781 sessions::table
1782 .filter(sessions::address.eq(&address_owned))
1783 .filter(sessions::device_id.eq(device_id)),
1784 ))
1785 .get_result(&mut conn)
1786 .map_err(|e| StoreError::Database(Box::new(e)))?;
1787 Ok(exists)
1788 })
1789 .await
1790 }
1791
1792 async fn has_signal_state_for_user(&self, user: &str) -> Result<bool> {
1793 let pool = self.pool.clone();
1794 let device_id = self.device_id;
1795 let pat_at = format!("{user}@%");
1798 let pat_dev = format!("{user}:%");
1799 self.with_semaphore(move || -> Result<bool> {
1800 let mut conn = pool
1801 .get()
1802 .map_err(|e| StoreError::Connection(Box::new(e)))?;
1803 let has_session = diesel::select(diesel::dsl::exists(
1804 sessions::table
1805 .filter(sessions::device_id.eq(device_id))
1806 .filter(
1807 sessions::address
1808 .like(&pat_at)
1809 .or(sessions::address.like(&pat_dev)),
1810 ),
1811 ))
1812 .get_result::<bool>(&mut conn)
1813 .map_err(|e| StoreError::Database(Box::new(e)))?;
1814 if has_session {
1815 return Ok(true);
1816 }
1817 let has_identity = diesel::select(diesel::dsl::exists(
1818 identities::table
1819 .filter(identities::device_id.eq(device_id))
1820 .filter(
1821 identities::address
1822 .like(&pat_at)
1823 .or(identities::address.like(&pat_dev)),
1824 ),
1825 ))
1826 .get_result::<bool>(&mut conn)
1827 .map_err(|e| StoreError::Database(Box::new(e)))?;
1828 Ok(has_identity)
1829 })
1830 .await
1831 }
1832
1833 async fn put_session(&self, address: &str, session: &[u8]) -> Result<()> {
1834 self.put_session_for_device(address, session, self.device_id)
1835 .await
1836 }
1837
1838 async fn put_sessions_batch(&self, sessions: &[(Arc<str>, Bytes)]) -> Result<()> {
1839 if sessions.is_empty() {
1840 return Ok(());
1841 }
1842
1843 let device_id = self.device_id;
1844 let batch = Arc::new(sessions.to_vec());
1845 self.with_retry("put_sessions_batch", || {
1846 let batch = batch.clone();
1847 Box::new(move |conn: &mut SqliteConnection| {
1848 conn.transaction(|conn| {
1849 for (address, record) in batch.iter() {
1850 diesel::insert_into(sessions::table)
1851 .values((
1852 sessions::address.eq(address.as_ref()),
1853 sessions::record.eq(record.as_ref()),
1854 sessions::device_id.eq(device_id),
1855 ))
1856 .on_conflict((sessions::address, sessions::device_id))
1857 .do_update()
1858 .set(sessions::record.eq(record.as_ref()))
1859 .execute(conn)?;
1860 }
1861 Ok(())
1862 })
1863 })
1864 })
1865 .await
1866 }
1867
1868 async fn delete_session(&self, address: &str) -> Result<()> {
1869 self.delete_session_for_device(address, self.device_id)
1870 .await
1871 }
1872
1873 async fn store_prekey(&self, id: u32, record: &[u8], uploaded: bool) -> Result<()> {
1874 let pool = self.pool.clone();
1875 let db_semaphore = self.db_semaphore.clone();
1876 let device_id = self.device_id;
1877 let record = record.to_vec();
1878
1879 const MAX_RETRIES: u32 = 5;
1880
1881 for attempt in 0..=MAX_RETRIES {
1882 let permit = db_semaphore
1883 .clone()
1884 .acquire_owned()
1885 .await
1886 .map_err(|e| StoreError::Database(Box::new(e)))?;
1887
1888 let pool_clone = pool.clone();
1889 let record_clone = record.clone();
1890
1891 let result =
1892 tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
1893 let mut conn = pool_clone
1894 .get()
1895 .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
1896 diesel::insert_into(prekeys::table)
1897 .values((
1898 prekeys::id.eq(id as i32),
1899 prekeys::key.eq(&record_clone),
1900 prekeys::uploaded.eq(uploaded),
1901 prekeys::device_id.eq(device_id),
1902 ))
1903 .on_conflict((prekeys::id, prekeys::device_id))
1904 .do_update()
1905 .set((
1906 prekeys::key.eq(&record_clone),
1907 prekeys::uploaded.eq(uploaded),
1908 ))
1909 .execute(&mut conn)
1910 .map_err(DieselOrStore::Diesel)?;
1911 Ok(())
1912 })
1913 .await;
1914
1915 drop(permit);
1916
1917 match result {
1918 Ok(Ok(())) => return Ok(()),
1919 Ok(Err(DieselOrStore::Diesel(ref e)))
1920 if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
1921 {
1922 let delay_ms = 10u64 * (1u64 << attempt.min(4));
1923 tokio::time::sleep(Duration::from_millis(delay_ms)).await;
1924 }
1925 Ok(Err(e)) => return Err(e.into()),
1926 Err(e) => return Err(StoreError::Database(Box::new(e))),
1927 }
1928 }
1929
1930 Err(StoreError::RetriesExhausted {
1931 op: "store_prekey".to_string(),
1932 })
1933 }
1934
1935 async fn store_prekeys_batch(&self, keys: &[(u32, Bytes)], uploaded: bool) -> Result<()> {
1936 if keys.is_empty() {
1937 return Ok(());
1938 }
1939
1940 let pool = self.pool.clone();
1941 let db_semaphore = self.db_semaphore.clone();
1942 let device_id = self.device_id;
1943 let keys: Vec<(u32, Bytes)> = keys.to_vec();
1944
1945 const MAX_RETRIES: u32 = 5;
1946
1947 for attempt in 0..=MAX_RETRIES {
1948 let permit = db_semaphore
1949 .clone()
1950 .acquire_owned()
1951 .await
1952 .map_err(|e| StoreError::Database(Box::new(e)))?;
1953
1954 let pool_clone = pool.clone();
1955 let keys_clone = keys.clone();
1956
1957 let result =
1958 tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
1959 let mut conn = pool_clone
1960 .get()
1961 .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
1962
1963 conn.transaction(|conn| {
1964 for (id, record) in &keys_clone {
1965 diesel::insert_into(prekeys::table)
1966 .values((
1967 prekeys::id.eq(*id as i32),
1968 prekeys::key.eq(record.as_ref()),
1969 prekeys::uploaded.eq(uploaded),
1970 prekeys::device_id.eq(device_id),
1971 ))
1972 .on_conflict((prekeys::id, prekeys::device_id))
1973 .do_update()
1974 .set((
1975 prekeys::key.eq(record.as_ref()),
1976 prekeys::uploaded.eq(uploaded),
1977 ))
1978 .execute(conn)?;
1979 }
1980 Ok::<(), diesel::result::Error>(())
1981 })
1982 .map_err(DieselOrStore::Diesel)
1983 })
1984 .await;
1985
1986 drop(permit);
1987
1988 match result {
1989 Ok(Ok(())) => return Ok(()),
1990 Ok(Err(DieselOrStore::Diesel(ref e)))
1991 if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
1992 {
1993 let delay_ms = 10u64 * (1u64 << attempt.min(4));
1994 tokio::time::sleep(Duration::from_millis(delay_ms)).await;
1995 }
1996 Ok(Err(e)) => return Err(e.into()),
1997 Err(e) => return Err(StoreError::Database(Box::new(e))),
1998 }
1999 }
2000
2001 Err(StoreError::RetriesExhausted {
2002 op: "store_prekeys_batch".to_string(),
2003 })
2004 }
2005
2006 async fn load_prekey(&self, id: u32) -> Result<Option<Bytes>> {
2007 let pool = self.pool.clone();
2008 let device_id = self.device_id;
2009 tokio::task::spawn_blocking(move || -> Result<Option<Bytes>> {
2010 let mut conn = pool
2011 .get()
2012 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2013 let res: Option<Vec<u8>> = prekeys::table
2014 .select(prekeys::key)
2015 .filter(prekeys::id.eq(id as i32))
2016 .filter(prekeys::device_id.eq(device_id))
2017 .first(&mut conn)
2018 .optional()
2019 .map_err(|e| StoreError::Database(Box::new(e)))?;
2020 Ok(res.map(Bytes::from))
2021 })
2022 .await
2023 .map_err(|e| StoreError::Database(Box::new(e)))?
2024 }
2025
2026 async fn load_prekeys_batch(&self, ids: &[u32]) -> Result<Vec<(u32, Bytes)>> {
2027 if ids.is_empty() {
2028 return Ok(Vec::new());
2029 }
2030 let pool = self.pool.clone();
2031 let device_id = self.device_id;
2032 let ids: Vec<i32> = ids.iter().map(|&id| id as i32).collect();
2033 self.with_semaphore(move || -> Result<Vec<(u32, Bytes)>> {
2034 let mut conn = pool
2035 .get()
2036 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2037 let mut out = Vec::with_capacity(ids.len());
2040 for chunk in ids.chunks(ID_PARAM_CHUNK) {
2041 let rows: Vec<(i32, Vec<u8>)> = prekeys::table
2042 .select((prekeys::id, prekeys::key))
2043 .filter(prekeys::id.eq_any(chunk))
2044 .filter(prekeys::device_id.eq(device_id))
2045 .load(&mut conn)
2046 .map_err(|e| StoreError::Database(Box::new(e)))?;
2047 out.extend(
2048 rows.into_iter()
2049 .map(|(id, key)| (id as u32, Bytes::from(key))),
2050 );
2051 }
2052 Ok(out)
2053 })
2054 .await
2055 }
2056
2057 async fn remove_prekey(&self, id: u32) -> Result<()> {
2058 let pool = self.pool.clone();
2059 let db_semaphore = self.db_semaphore.clone();
2060 let device_id = self.device_id;
2061
2062 const MAX_RETRIES: u32 = 5;
2063
2064 for attempt in 0..=MAX_RETRIES {
2065 let permit = db_semaphore
2066 .clone()
2067 .acquire_owned()
2068 .await
2069 .map_err(|e| StoreError::Database(Box::new(e)))?;
2070
2071 let pool_clone = pool.clone();
2072
2073 let result =
2074 tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
2075 let mut conn = pool_clone
2076 .get()
2077 .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
2078 diesel::delete(
2079 prekeys::table
2080 .filter(prekeys::id.eq(id as i32))
2081 .filter(prekeys::device_id.eq(device_id)),
2082 )
2083 .execute(&mut conn)
2084 .map_err(DieselOrStore::Diesel)?;
2085 Ok(())
2086 })
2087 .await;
2088
2089 drop(permit);
2090
2091 match result {
2092 Ok(Ok(())) => return Ok(()),
2093 Ok(Err(DieselOrStore::Diesel(ref e)))
2094 if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
2095 {
2096 let delay_ms = 10u64 * (1u64 << attempt.min(4));
2097 tokio::time::sleep(Duration::from_millis(delay_ms)).await;
2098 }
2099 Ok(Err(e)) => return Err(e.into()),
2100 Err(e) => return Err(StoreError::Database(Box::new(e))),
2101 }
2102 }
2103
2104 Err(StoreError::RetriesExhausted {
2105 op: "remove_prekey".to_string(),
2106 })
2107 }
2108
2109 async fn mark_prekeys_uploaded(&self, ids: &[u32]) -> Result<()> {
2110 if ids.is_empty() {
2111 return Ok(());
2112 }
2113 let device_id = self.device_id;
2114 let ids: Vec<i32> = ids.iter().map(|&id| id as i32).collect();
2115 self.with_retry("mark_prekeys_uploaded", move || {
2116 let ids = ids.clone();
2117 Box::new(move |conn: &mut SqliteConnection| {
2118 for chunk in ids.chunks(ID_PARAM_CHUNK) {
2121 diesel::update(
2122 prekeys::table
2123 .filter(prekeys::id.eq_any(chunk.to_vec()))
2124 .filter(prekeys::device_id.eq(device_id)),
2125 )
2126 .set(prekeys::uploaded.eq(true))
2127 .execute(conn)?;
2128 }
2129 Ok(())
2130 })
2131 })
2132 .await
2133 }
2134
2135 async fn get_max_prekey_id(&self) -> Result<u32> {
2136 let pool = self.pool.clone();
2137 let device_id = self.device_id;
2138 let db_semaphore = self.db_semaphore.clone();
2139 let _permit = db_semaphore
2140 .acquire()
2141 .await
2142 .map_err(|e| StoreError::Database(Box::new(e)))?;
2143
2144 tokio::task::spawn_blocking(move || -> Result<u32> {
2145 let mut conn = pool
2146 .get()
2147 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2148 use diesel::dsl::max;
2149 let result: Option<i32> = prekeys::table
2150 .filter(prekeys::device_id.eq(device_id))
2151 .select(max(prekeys::id))
2152 .first(&mut conn)
2153 .map_err(|e| StoreError::Database(Box::new(e)))?;
2154 Ok(result.unwrap_or(0) as u32)
2155 })
2156 .await
2157 .map_err(|e| StoreError::Database(Box::new(e)))?
2158 }
2159
2160 async fn store_signed_prekey(&self, id: u32, record: &[u8]) -> Result<()> {
2161 let pool = self.pool.clone();
2162 let db_semaphore = self.db_semaphore.clone();
2163 let device_id = self.device_id;
2164 let record = record.to_vec();
2165
2166 const MAX_RETRIES: u32 = 5;
2167
2168 for attempt in 0..=MAX_RETRIES {
2169 let permit = db_semaphore
2170 .clone()
2171 .acquire_owned()
2172 .await
2173 .map_err(|e| StoreError::Database(Box::new(e)))?;
2174
2175 let pool_clone = pool.clone();
2176 let record_clone = record.clone();
2177
2178 let result =
2179 tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
2180 let mut conn = pool_clone
2181 .get()
2182 .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
2183 diesel::insert_into(signed_prekeys::table)
2184 .values((
2185 signed_prekeys::id.eq(id as i32),
2186 signed_prekeys::record.eq(&record_clone),
2187 signed_prekeys::device_id.eq(device_id),
2188 ))
2189 .on_conflict((signed_prekeys::id, signed_prekeys::device_id))
2190 .do_update()
2191 .set(signed_prekeys::record.eq(&record_clone))
2192 .execute(&mut conn)
2193 .map_err(DieselOrStore::Diesel)?;
2194 Ok(())
2195 })
2196 .await;
2197
2198 drop(permit);
2199
2200 match result {
2201 Ok(Ok(())) => return Ok(()),
2202 Ok(Err(DieselOrStore::Diesel(ref e)))
2203 if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
2204 {
2205 let delay_ms = 10u64 * (1u64 << attempt.min(4));
2206 tokio::time::sleep(Duration::from_millis(delay_ms)).await;
2207 }
2208 Ok(Err(e)) => return Err(e.into()),
2209 Err(e) => return Err(StoreError::Database(Box::new(e))),
2210 }
2211 }
2212
2213 Err(StoreError::RetriesExhausted {
2214 op: "store_signed_prekey".to_string(),
2215 })
2216 }
2217
2218 async fn load_signed_prekey(&self, id: u32) -> Result<Option<Vec<u8>>> {
2219 let pool = self.pool.clone();
2220 let device_id = self.device_id;
2221 tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
2222 let mut conn = pool
2223 .get()
2224 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2225 let res: Option<Vec<u8>> = signed_prekeys::table
2226 .select(signed_prekeys::record)
2227 .filter(signed_prekeys::id.eq(id as i32))
2228 .filter(signed_prekeys::device_id.eq(device_id))
2229 .first(&mut conn)
2230 .optional()
2231 .map_err(|e| StoreError::Database(Box::new(e)))?;
2232 Ok(res)
2233 })
2234 .await
2235 .map_err(|e| StoreError::Database(Box::new(e)))?
2236 }
2237
2238 async fn load_all_signed_prekeys(&self) -> Result<Vec<(u32, Vec<u8>)>> {
2239 let pool = self.pool.clone();
2240 let device_id = self.device_id;
2241 tokio::task::spawn_blocking(move || -> Result<Vec<(u32, Vec<u8>)>> {
2242 let mut conn = pool
2243 .get()
2244 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2245 let results: Vec<(i32, Vec<u8>)> = signed_prekeys::table
2246 .select((signed_prekeys::id, signed_prekeys::record))
2247 .filter(signed_prekeys::device_id.eq(device_id))
2248 .load(&mut conn)
2249 .map_err(|e| StoreError::Database(Box::new(e)))?;
2250 Ok(results
2251 .into_iter()
2252 .map(|(id, record)| (id as u32, record))
2253 .collect())
2254 })
2255 .await
2256 .map_err(|e| StoreError::Database(Box::new(e)))?
2257 }
2258
2259 async fn remove_signed_prekey(&self, id: u32) -> Result<()> {
2260 let pool = self.pool.clone();
2261 let db_semaphore = self.db_semaphore.clone();
2262 let device_id = self.device_id;
2263
2264 const MAX_RETRIES: u32 = 5;
2265
2266 for attempt in 0..=MAX_RETRIES {
2267 let permit = db_semaphore
2268 .clone()
2269 .acquire_owned()
2270 .await
2271 .map_err(|e| StoreError::Database(Box::new(e)))?;
2272
2273 let pool_clone = pool.clone();
2274
2275 let result =
2276 tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
2277 let mut conn = pool_clone
2278 .get()
2279 .map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
2280 diesel::delete(
2281 signed_prekeys::table
2282 .filter(signed_prekeys::id.eq(id as i32))
2283 .filter(signed_prekeys::device_id.eq(device_id)),
2284 )
2285 .execute(&mut conn)
2286 .map_err(DieselOrStore::Diesel)?;
2287 Ok(())
2288 })
2289 .await;
2290
2291 drop(permit);
2292
2293 match result {
2294 Ok(Ok(())) => return Ok(()),
2295 Ok(Err(DieselOrStore::Diesel(ref e)))
2296 if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
2297 {
2298 let delay_ms = 10u64 * (1u64 << attempt.min(4));
2299 tokio::time::sleep(Duration::from_millis(delay_ms)).await;
2300 }
2301 Ok(Err(e)) => return Err(e.into()),
2302 Err(e) => return Err(StoreError::Database(Box::new(e))),
2303 }
2304 }
2305
2306 Err(StoreError::RetriesExhausted {
2307 op: "remove_signed_prekey".to_string(),
2308 })
2309 }
2310
2311 async fn put_sender_key(&self, address: &str, record: &[u8]) -> Result<()> {
2312 self.put_sender_key_for_device(address, record, self.device_id)
2313 .await
2314 }
2315
2316 async fn put_sender_keys_batch(&self, sender_keys: &[(Arc<str>, Bytes)]) -> Result<()> {
2317 if sender_keys.is_empty() {
2318 return Ok(());
2319 }
2320
2321 let device_id = self.device_id;
2322 let batch = Arc::new(sender_keys.to_vec());
2323 self.with_retry("put_sender_keys_batch", || {
2324 let batch = batch.clone();
2325 Box::new(move |conn: &mut SqliteConnection| {
2326 conn.transaction(|conn| {
2327 for (address, record) in batch.iter() {
2328 diesel::insert_into(sender_keys::table)
2329 .values((
2330 sender_keys::address.eq(address.as_ref()),
2331 sender_keys::record.eq(record.as_ref()),
2332 sender_keys::device_id.eq(device_id),
2333 ))
2334 .on_conflict((sender_keys::address, sender_keys::device_id))
2335 .do_update()
2336 .set(sender_keys::record.eq(record.as_ref()))
2337 .execute(conn)?;
2338 }
2339 Ok(())
2340 })
2341 })
2342 })
2343 .await
2344 }
2345
2346 async fn get_sender_key(&self, address: &str) -> Result<Option<Vec<u8>>> {
2347 self.get_sender_key_for_device(address, self.device_id)
2348 .await
2349 }
2350
2351 async fn delete_sender_key(&self, address: &str) -> Result<()> {
2352 self.delete_sender_key_for_device(address, self.device_id)
2353 .await
2354 }
2355}
2356
2357#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
2358#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
2359impl AppSyncStore for SqliteStore {
2360 async fn get_sync_key(&self, key_id: &[u8]) -> Result<Option<AppStateSyncKey>> {
2361 self.get_app_state_sync_key_for_device(key_id, self.device_id)
2362 .await
2363 }
2364
2365 async fn set_sync_key(&self, key_id: &[u8], key: AppStateSyncKey) -> Result<()> {
2366 self.set_app_state_sync_key_for_device(key_id, key, self.device_id)
2367 .await
2368 }
2369
2370 async fn get_version(&self, name: &str) -> Result<HashState> {
2371 self.get_app_state_version_for_device(name, self.device_id)
2372 .await
2373 }
2374
2375 async fn set_version(&self, name: &str, state: HashState) -> Result<()> {
2376 self.set_app_state_version_for_device(name, state, self.device_id)
2377 .await
2378 }
2379
2380 async fn put_mutation_macs(
2381 &self,
2382 name: &str,
2383 version: u64,
2384 mutations: &[AppStateMutationMAC],
2385 ) -> Result<()> {
2386 self.put_app_state_mutation_macs_for_device(name, version, mutations, self.device_id)
2387 .await
2388 }
2389
2390 async fn get_mutation_mac(&self, name: &str, index_mac: &[u8]) -> Result<Option<Vec<u8>>> {
2391 self.get_app_state_mutation_mac_for_device(name, index_mac, self.device_id)
2392 .await
2393 }
2394
2395 async fn get_mutation_macs(
2396 &self,
2397 name: &str,
2398 index_macs: &[[u8; 32]],
2399 ) -> Result<std::collections::HashMap<[u8; 32], Vec<u8>>> {
2400 self.get_app_state_mutation_macs_batch_for_device(name, index_macs, self.device_id)
2401 .await
2402 }
2403
2404 async fn delete_mutation_macs(&self, name: &str, index_macs: &[Vec<u8>]) -> Result<()> {
2405 self.delete_app_state_mutation_macs_for_device(name, index_macs, self.device_id)
2406 .await
2407 }
2408
2409 async fn clear_mutation_macs(&self, name: &str) -> Result<()> {
2410 let device_id = self.device_id;
2411 let name = name.to_string();
2412 self.with_retry("clear_mutation_macs", || {
2413 let name = name.clone();
2414 Box::new(move |conn: &mut SqliteConnection| {
2415 diesel::delete(
2416 app_state_mutation_macs::table
2417 .filter(app_state_mutation_macs::name.eq(&name))
2418 .filter(app_state_mutation_macs::device_id.eq(device_id)),
2419 )
2420 .execute(conn)?;
2421 Ok(())
2422 })
2423 })
2424 .await
2425 }
2426
2427 async fn get_latest_sync_key_id(&self) -> Result<Option<Vec<u8>>> {
2428 self.get_latest_app_state_sync_key_id_for_device(self.device_id)
2429 .await
2430 }
2431}
2432
2433fn insert_pending_inbound_row(
2437 conn: &mut SqliteConnection,
2438 device_id: i32,
2439 chat: &str,
2440 sender: &str,
2441 id: &str,
2442 message: &[u8],
2443) -> QueryResult<usize> {
2444 diesel::replace_into(pending_inbound_messages::table)
2445 .values((
2446 pending_inbound_messages::chat.eq(chat),
2447 pending_inbound_messages::sender.eq(sender),
2448 pending_inbound_messages::id.eq(id),
2449 pending_inbound_messages::message.eq(message),
2450 pending_inbound_messages::device_id.eq(device_id),
2451 ))
2452 .execute(conn)
2453}
2454
2455fn delete_pending_inbound_row(
2457 conn: &mut SqliteConnection,
2458 device_id: i32,
2459 chat: &str,
2460 sender: &str,
2461 id: &str,
2462) -> QueryResult<usize> {
2463 diesel::delete(
2464 pending_inbound_messages::table
2465 .filter(pending_inbound_messages::chat.eq(chat))
2466 .filter(pending_inbound_messages::sender.eq(sender))
2467 .filter(pending_inbound_messages::id.eq(id))
2468 .filter(pending_inbound_messages::device_id.eq(device_id)),
2469 )
2470 .execute(conn)
2471}
2472
2473#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
2474#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
2475impl ProtocolStore for SqliteStore {
2476 async fn get_sender_key_devices(&self, group_jid: &str) -> Result<Vec<(String, bool)>> {
2477 let pool = self.pool.clone();
2478 let device_id = self.device_id;
2479 let group_jid = group_jid.to_string();
2480 tokio::task::spawn_blocking(move || -> Result<Vec<(String, bool)>> {
2481 let mut conn = pool
2482 .get()
2483 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2484 let rows: Vec<(String, i32)> = sender_key_devices::table
2485 .select((sender_key_devices::device_jid, sender_key_devices::has_key))
2486 .filter(sender_key_devices::group_jid.eq(&group_jid))
2487 .filter(sender_key_devices::device_id.eq(device_id))
2488 .load(&mut conn)
2489 .map_err(|e| StoreError::Database(Box::new(e)))?;
2490 Ok(rows
2491 .into_iter()
2492 .map(|(jid, has_key)| (jid, has_key != 0))
2493 .collect())
2494 })
2495 .await
2496 .map_err(|e| StoreError::Database(Box::new(e)))?
2497 }
2498
2499 async fn set_sender_key_status(&self, group_jid: &str, entries: &[(&str, bool)]) -> Result<()> {
2500 if entries.is_empty() {
2501 return Ok(());
2502 }
2503 let device_id = self.device_id;
2504 let group_jid = group_jid.to_string();
2505 let owned_entries: Arc<Vec<(String, bool)>> = Arc::new(
2506 entries
2507 .iter()
2508 .map(|(jid, has_key)| (jid.to_string(), *has_key))
2509 .collect(),
2510 );
2511 let now = wacore::time::now_secs();
2512 self.with_retry("set_sender_key_status", || {
2513 let group_jid = group_jid.clone();
2514 let owned_entries = Arc::clone(&owned_entries);
2515 Box::new(move |conn: &mut SqliteConnection| {
2516 let values: Vec<_> = owned_entries
2517 .iter()
2518 .map(|(device_jid, has_key)| {
2519 (
2520 sender_key_devices::group_jid.eq(&group_jid),
2521 sender_key_devices::device_jid.eq(device_jid),
2522 sender_key_devices::has_key.eq(i32::from(*has_key)),
2523 sender_key_devices::device_id.eq(device_id),
2524 sender_key_devices::updated_at.eq(now),
2525 )
2526 })
2527 .collect();
2528
2529 const CHUNK_SIZE: usize = 190;
2530
2531 for chunk in values.chunks(CHUNK_SIZE) {
2532 diesel::insert_into(sender_key_devices::table)
2533 .values(chunk)
2534 .on_conflict((
2535 sender_key_devices::group_jid,
2536 sender_key_devices::device_jid,
2537 sender_key_devices::device_id,
2538 ))
2539 .do_update()
2540 .set((
2541 sender_key_devices::has_key.eq(excluded(sender_key_devices::has_key)),
2542 sender_key_devices::updated_at.eq(now),
2543 ))
2544 .execute(conn)?;
2545 }
2546 Ok(())
2547 })
2548 })
2549 .await
2550 }
2551
2552 async fn clear_sender_key_devices(&self, group_jid: &str) -> Result<()> {
2553 let device_id = self.device_id;
2554 let group_jid = group_jid.to_string();
2555 self.with_retry("clear_sender_key_devices", || {
2556 let group_jid = group_jid.clone();
2557 Box::new(move |conn: &mut SqliteConnection| {
2558 diesel::delete(
2559 sender_key_devices::table
2560 .filter(sender_key_devices::group_jid.eq(&group_jid))
2561 .filter(sender_key_devices::device_id.eq(device_id)),
2562 )
2563 .execute(conn)?;
2564 Ok(())
2565 })
2566 })
2567 .await
2568 }
2569
2570 async fn clear_all_sender_key_devices(&self) -> Result<()> {
2571 let device_id = self.device_id;
2572 self.with_retry("clear_all_sender_key_devices", || {
2573 Box::new(move |conn: &mut SqliteConnection| {
2574 diesel::delete(
2575 sender_key_devices::table.filter(sender_key_devices::device_id.eq(device_id)),
2576 )
2577 .execute(conn)?;
2578 Ok(())
2579 })
2580 })
2581 .await
2582 }
2583
2584 async fn delete_sender_key_device_rows(&self, device_jids: &[&str]) -> Result<()> {
2585 if device_jids.is_empty() {
2586 return Ok(());
2587 }
2588 let device_id = self.device_id;
2589 let owned: Arc<Vec<String>> = Arc::new(device_jids.iter().map(|s| s.to_string()).collect());
2590 self.with_retry("delete_sender_key_device_rows", || {
2591 let owned = Arc::clone(&owned);
2592 Box::new(move |conn: &mut SqliteConnection| {
2593 const CHUNK: usize = 190;
2594 for chunk in owned.chunks(CHUNK) {
2595 diesel::delete(
2596 sender_key_devices::table
2597 .filter(sender_key_devices::device_jid.eq_any(chunk))
2598 .filter(sender_key_devices::device_id.eq(device_id)),
2599 )
2600 .execute(conn)?;
2601 }
2602 Ok(())
2603 })
2604 })
2605 .await
2606 }
2607
2608 async fn get_lid_mapping(&self, lid: &str) -> Result<Option<LidPnMappingEntry>> {
2609 let pool = self.pool.clone();
2610 let device_id = self.device_id;
2611 let lid = lid.to_string();
2612 tokio::task::spawn_blocking(move || -> Result<Option<LidPnMappingEntry>> {
2613 let mut conn = pool
2614 .get()
2615 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2616 let row: Option<(String, String, i64, String, i64)> = lid_pn_mapping::table
2617 .select((
2618 lid_pn_mapping::lid,
2619 lid_pn_mapping::phone_number,
2620 lid_pn_mapping::created_at,
2621 lid_pn_mapping::learning_source,
2622 lid_pn_mapping::updated_at,
2623 ))
2624 .filter(lid_pn_mapping::lid.eq(&lid))
2625 .filter(lid_pn_mapping::device_id.eq(device_id))
2626 .first(&mut conn)
2627 .optional()
2628 .map_err(|e| StoreError::Database(Box::new(e)))?;
2629 Ok(row.map(
2630 |(lid, phone_number, created_at, learning_source, updated_at)| LidPnMappingEntry {
2631 lid,
2632 phone_number,
2633 created_at,
2634 updated_at,
2635 learning_source,
2636 },
2637 ))
2638 })
2639 .await
2640 .map_err(|e| StoreError::Database(Box::new(e)))?
2641 }
2642
2643 async fn get_pn_mapping(&self, phone: &str) -> Result<Option<LidPnMappingEntry>> {
2644 let pool = self.pool.clone();
2645 let device_id = self.device_id;
2646 let phone = phone.to_string();
2647 tokio::task::spawn_blocking(move || -> Result<Option<LidPnMappingEntry>> {
2648 let mut conn = pool
2649 .get()
2650 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2651 let row: Option<(String, String, i64, String, i64)> = lid_pn_mapping::table
2652 .select((
2653 lid_pn_mapping::lid,
2654 lid_pn_mapping::phone_number,
2655 lid_pn_mapping::created_at,
2656 lid_pn_mapping::learning_source,
2657 lid_pn_mapping::updated_at,
2658 ))
2659 .filter(lid_pn_mapping::phone_number.eq(&phone))
2660 .filter(lid_pn_mapping::device_id.eq(device_id))
2661 .order(lid_pn_mapping::updated_at.desc())
2662 .first(&mut conn)
2663 .optional()
2664 .map_err(|e| StoreError::Database(Box::new(e)))?;
2665 Ok(row.map(
2666 |(lid, phone_number, created_at, learning_source, updated_at)| LidPnMappingEntry {
2667 lid,
2668 phone_number,
2669 created_at,
2670 updated_at,
2671 learning_source,
2672 },
2673 ))
2674 })
2675 .await
2676 .map_err(|e| StoreError::Database(Box::new(e)))?
2677 }
2678
2679 async fn put_lid_mapping(&self, entry: &LidPnMappingEntry) -> Result<()> {
2680 self.put_lid_mappings(std::slice::from_ref(entry)).await
2681 }
2682
2683 async fn put_lid_mappings(&self, entries: &[LidPnMappingEntry]) -> Result<()> {
2684 if entries.is_empty() {
2685 return Ok(());
2686 }
2687 let device_id = self.device_id;
2688 let entries: Arc<Vec<LidPnMappingEntry>> = Arc::new(entries.to_vec());
2692 self.with_retry("put_lid_mappings", move || {
2693 let entries = Arc::clone(&entries);
2694 Box::new(move |conn: &mut SqliteConnection| {
2695 conn.transaction::<_, DieselError, _>(|conn| {
2696 for entry in entries.iter() {
2697 diesel::insert_into(lid_pn_mapping::table)
2698 .values((
2699 lid_pn_mapping::lid.eq(&entry.lid),
2700 lid_pn_mapping::phone_number.eq(&entry.phone_number),
2701 lid_pn_mapping::created_at.eq(entry.created_at),
2702 lid_pn_mapping::learning_source.eq(&entry.learning_source),
2703 lid_pn_mapping::updated_at.eq(entry.updated_at),
2704 lid_pn_mapping::device_id.eq(device_id),
2705 ))
2706 .on_conflict((lid_pn_mapping::lid, lid_pn_mapping::device_id))
2707 .do_update()
2708 .set((
2709 lid_pn_mapping::phone_number.eq(&entry.phone_number),
2710 lid_pn_mapping::learning_source.eq(&entry.learning_source),
2711 lid_pn_mapping::updated_at.eq(entry.updated_at),
2712 ))
2713 .execute(conn)?;
2714 }
2715 Ok(())
2716 })
2717 })
2718 })
2719 .await
2720 }
2721
2722 async fn get_all_lid_mappings(&self) -> Result<Vec<LidPnMappingEntry>> {
2723 let pool = self.pool.clone();
2724 let device_id = self.device_id;
2725 tokio::task::spawn_blocking(move || -> Result<Vec<LidPnMappingEntry>> {
2726 let mut conn = pool
2727 .get()
2728 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2729 let rows: Vec<(String, String, i64, String, i64)> = lid_pn_mapping::table
2730 .select((
2731 lid_pn_mapping::lid,
2732 lid_pn_mapping::phone_number,
2733 lid_pn_mapping::created_at,
2734 lid_pn_mapping::learning_source,
2735 lid_pn_mapping::updated_at,
2736 ))
2737 .filter(lid_pn_mapping::device_id.eq(device_id))
2738 .load(&mut conn)
2739 .map_err(|e| StoreError::Database(Box::new(e)))?;
2740 Ok(rows
2741 .into_iter()
2742 .map(
2743 |(lid, phone_number, created_at, learning_source, updated_at)| {
2744 LidPnMappingEntry {
2745 lid,
2746 phone_number,
2747 created_at,
2748 updated_at,
2749 learning_source,
2750 }
2751 },
2752 )
2753 .collect())
2754 })
2755 .await
2756 .map_err(|e| StoreError::Database(Box::new(e)))?
2757 }
2758
2759 async fn save_base_key(&self, address: &str, message_id: &str, base_key: &[u8]) -> Result<()> {
2760 let pool = self.pool.clone();
2761 let device_id = self.device_id;
2762 let address = address.to_string();
2763 let message_id = message_id.to_string();
2764 let base_key = base_key.to_vec();
2765 let now = wacore::time::now_secs() as i32;
2766 tokio::task::spawn_blocking(move || -> Result<()> {
2767 let mut conn = pool
2768 .get()
2769 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2770 diesel::insert_into(base_keys::table)
2771 .values((
2772 base_keys::address.eq(&address),
2773 base_keys::message_id.eq(&message_id),
2774 base_keys::base_key.eq(&base_key),
2775 base_keys::device_id.eq(device_id),
2776 base_keys::created_at.eq(now),
2777 ))
2778 .on_conflict((
2779 base_keys::address,
2780 base_keys::message_id,
2781 base_keys::device_id,
2782 ))
2783 .do_update()
2784 .set(base_keys::base_key.eq(&base_key))
2785 .execute(&mut conn)
2786 .map_err(|e| StoreError::Database(Box::new(e)))?;
2787 Ok(())
2788 })
2789 .await
2790 .map_err(|e| StoreError::Database(Box::new(e)))??;
2791 Ok(())
2792 }
2793
2794 async fn has_same_base_key(
2795 &self,
2796 address: &str,
2797 message_id: &str,
2798 current_base_key: &[u8],
2799 ) -> Result<bool> {
2800 let pool = self.pool.clone();
2801 let device_id = self.device_id;
2802 let address = address.to_string();
2803 let message_id = message_id.to_string();
2804 let current_base_key = current_base_key.to_vec();
2805 tokio::task::spawn_blocking(move || -> Result<bool> {
2806 let mut conn = pool
2807 .get()
2808 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2809 let stored_key: Option<Vec<u8>> = base_keys::table
2810 .select(base_keys::base_key)
2811 .filter(base_keys::address.eq(&address))
2812 .filter(base_keys::message_id.eq(&message_id))
2813 .filter(base_keys::device_id.eq(device_id))
2814 .first(&mut conn)
2815 .optional()
2816 .map_err(|e| StoreError::Database(Box::new(e)))?;
2817 Ok(stored_key.as_ref() == Some(¤t_base_key))
2818 })
2819 .await
2820 .map_err(|e| StoreError::Database(Box::new(e)))?
2821 }
2822
2823 async fn delete_base_key(&self, address: &str, message_id: &str) -> Result<()> {
2824 let pool = self.pool.clone();
2825 let device_id = self.device_id;
2826 let address = address.to_string();
2827 let message_id = message_id.to_string();
2828 tokio::task::spawn_blocking(move || -> Result<()> {
2829 let mut conn = pool
2830 .get()
2831 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2832 diesel::delete(
2833 base_keys::table
2834 .filter(base_keys::address.eq(&address))
2835 .filter(base_keys::message_id.eq(&message_id))
2836 .filter(base_keys::device_id.eq(device_id)),
2837 )
2838 .execute(&mut conn)
2839 .map_err(|e| StoreError::Database(Box::new(e)))?;
2840 Ok(())
2841 })
2842 .await
2843 .map_err(|e| StoreError::Database(Box::new(e)))??;
2844 Ok(())
2845 }
2846
2847 async fn update_device_list(&self, record: DeviceListRecord) -> Result<()> {
2848 let pool = self.pool.clone();
2849 let device_id = self.device_id;
2850 let devices_json = serde_json::to_string(&record.devices)
2851 .map_err(|e| StoreError::Serialization(Box::new(e)))?;
2852 let now = wacore::time::now_secs() as i32;
2853 tokio::task::spawn_blocking(move || -> Result<()> {
2854 let mut conn = pool
2855 .get()
2856 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2857 let raw_id_i32 = record.raw_id.map(|r| r as i32);
2858 diesel::insert_into(device_registry::table)
2859 .values((
2860 device_registry::user_id.eq(&record.user),
2861 device_registry::devices_json.eq(&devices_json),
2862 device_registry::timestamp.eq(record.timestamp as i32),
2863 device_registry::phash.eq(&record.phash),
2864 device_registry::device_id.eq(device_id),
2865 device_registry::updated_at.eq(now),
2866 device_registry::raw_id.eq(raw_id_i32),
2867 ))
2868 .on_conflict((device_registry::user_id, device_registry::device_id))
2869 .do_update()
2870 .set((
2871 device_registry::devices_json.eq(&devices_json),
2872 device_registry::timestamp.eq(record.timestamp as i32),
2873 device_registry::phash.eq(&record.phash),
2874 device_registry::updated_at.eq(now),
2875 device_registry::raw_id.eq(raw_id_i32),
2876 ))
2877 .execute(&mut conn)
2878 .map_err(|e| StoreError::Database(Box::new(e)))?;
2879 Ok(())
2880 })
2881 .await
2882 .map_err(|e| StoreError::Database(Box::new(e)))??;
2883 Ok(())
2884 }
2885
2886 async fn update_device_lists(&self, records: Vec<DeviceListRecord>) -> Result<()> {
2887 if records.is_empty() {
2888 return Ok(());
2889 }
2890 let device_id = self.device_id;
2891 let now = wacore::time::now_secs() as i32;
2892
2893 struct PreparedRow {
2897 user: String,
2898 devices_json: String,
2899 timestamp: i32,
2900 phash: Option<String>,
2901 raw_id: Option<i32>,
2902 }
2903
2904 let prepared: Vec<PreparedRow> = records
2905 .into_iter()
2906 .map(|r| {
2907 let devices_json = serde_json::to_string(&r.devices)
2908 .map_err(|e| StoreError::Serialization(Box::new(e)))?;
2909 Ok(PreparedRow {
2910 user: r.user,
2911 devices_json,
2912 timestamp: r.timestamp as i32,
2913 phash: r.phash,
2914 raw_id: r.raw_id.map(|v| v as i32),
2915 })
2916 })
2917 .collect::<Result<Vec<_>>>()?;
2918 let prepared = Arc::new(prepared);
2919
2920 self.with_retry("update_device_lists", move || {
2921 let prepared = Arc::clone(&prepared);
2922 Box::new(move |conn: &mut SqliteConnection| {
2923 conn.transaction::<_, DieselError, _>(|conn| {
2924 for row in prepared.iter() {
2925 diesel::insert_into(device_registry::table)
2926 .values((
2927 device_registry::user_id.eq(&row.user),
2928 device_registry::devices_json.eq(&row.devices_json),
2929 device_registry::timestamp.eq(row.timestamp),
2930 device_registry::phash.eq(&row.phash),
2931 device_registry::device_id.eq(device_id),
2932 device_registry::updated_at.eq(now),
2933 device_registry::raw_id.eq(row.raw_id),
2934 ))
2935 .on_conflict((device_registry::user_id, device_registry::device_id))
2936 .do_update()
2937 .set((
2938 device_registry::devices_json.eq(&row.devices_json),
2939 device_registry::timestamp.eq(row.timestamp),
2940 device_registry::phash.eq(&row.phash),
2941 device_registry::updated_at.eq(now),
2942 device_registry::raw_id.eq(row.raw_id),
2943 ))
2944 .execute(conn)?;
2945 }
2946 Ok(())
2947 })
2948 })
2949 })
2950 .await
2951 }
2952
2953 async fn get_devices(&self, user: &str) -> Result<Option<DeviceListRecord>> {
2954 let pool = self.pool.clone();
2955 let device_id = self.device_id;
2956 let user = user.to_string();
2957 tokio::task::spawn_blocking(move || -> Result<Option<DeviceListRecord>> {
2958 let mut conn = pool
2959 .get()
2960 .map_err(|e| StoreError::Connection(Box::new(e)))?;
2961 let row: Option<(String, String, i32, Option<String>, Option<i32>)> =
2962 device_registry::table
2963 .select((
2964 device_registry::user_id,
2965 device_registry::devices_json,
2966 device_registry::timestamp,
2967 device_registry::phash,
2968 device_registry::raw_id,
2969 ))
2970 .filter(device_registry::user_id.eq(&user))
2971 .filter(device_registry::device_id.eq(device_id))
2972 .first(&mut conn)
2973 .optional()
2974 .map_err(|e| StoreError::Database(Box::new(e)))?;
2975 match row {
2976 Some((user, devices_json, timestamp, phash, raw_id)) => {
2977 let devices: Vec<DeviceInfo> = serde_json::from_str(&devices_json)
2978 .map_err(|e| StoreError::Serialization(Box::new(e)))?;
2979 Ok(Some(DeviceListRecord {
2980 user,
2981 devices,
2982 timestamp: timestamp as i64,
2983 phash,
2984 raw_id: raw_id.map(|r| r as u32),
2985 }))
2986 }
2987 None => Ok(None),
2988 }
2989 })
2990 .await
2991 .map_err(|e| StoreError::Database(Box::new(e)))?
2992 }
2993
2994 async fn delete_devices(&self, user: &str) -> Result<()> {
2995 let pool = self.pool.clone();
2996 let device_id = self.device_id;
2997 let user = user.to_string();
2998 tokio::task::spawn_blocking(move || -> Result<()> {
2999 let mut conn = pool
3000 .get()
3001 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3002 diesel::delete(
3003 device_registry::table
3004 .filter(device_registry::user_id.eq(&user))
3005 .filter(device_registry::device_id.eq(device_id)),
3006 )
3007 .execute(&mut conn)
3008 .map_err(|e| StoreError::Database(Box::new(e)))?;
3009 Ok(())
3010 })
3011 .await
3012 .map_err(|e| StoreError::Database(Box::new(e)))??;
3013 Ok(())
3014 }
3015
3016 async fn get_group_metadata(&self, group_jid: &str) -> Result<Option<Vec<u8>>> {
3017 let pool = self.pool.clone();
3018 let device_id = self.device_id;
3019 let group_jid = group_jid.to_string();
3020 tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
3021 let mut conn = pool
3022 .get()
3023 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3024 let row: Option<Vec<u8>> = group_metadata::table
3025 .select(group_metadata::info)
3026 .filter(group_metadata::group_jid.eq(&group_jid))
3027 .filter(group_metadata::device_id.eq(device_id))
3028 .first(&mut conn)
3029 .optional()
3030 .map_err(|e| StoreError::Database(Box::new(e)))?;
3031 Ok(row)
3032 })
3033 .await
3034 .map_err(|e| StoreError::Database(Box::new(e)))?
3035 }
3036
3037 async fn put_group_metadata(&self, group_jid: &str, blob: &[u8]) -> Result<()> {
3038 let pool = self.pool.clone();
3039 let device_id = self.device_id;
3040 let group_jid = group_jid.to_string();
3041 let blob = blob.to_vec();
3042 let now = wacore::time::now_secs();
3043 tokio::task::spawn_blocking(move || -> Result<()> {
3044 let mut conn = pool
3045 .get()
3046 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3047 diesel::insert_into(group_metadata::table)
3048 .values((
3049 group_metadata::group_jid.eq(&group_jid),
3050 group_metadata::info.eq(&blob),
3051 group_metadata::device_id.eq(device_id),
3052 group_metadata::updated_at.eq(now),
3053 ))
3054 .on_conflict((group_metadata::group_jid, group_metadata::device_id))
3055 .do_update()
3056 .set((
3057 group_metadata::info.eq(&blob),
3058 group_metadata::updated_at.eq(now),
3059 ))
3060 .execute(&mut conn)
3061 .map_err(|e| StoreError::Database(Box::new(e)))?;
3062 Ok(())
3063 })
3064 .await
3065 .map_err(|e| StoreError::Database(Box::new(e)))??;
3066 Ok(())
3067 }
3068
3069 async fn delete_group_metadata(&self, group_jid: &str) -> Result<()> {
3070 let pool = self.pool.clone();
3071 let device_id = self.device_id;
3072 let group_jid = group_jid.to_string();
3073 tokio::task::spawn_blocking(move || -> Result<()> {
3074 let mut conn = pool
3075 .get()
3076 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3077 diesel::delete(
3078 group_metadata::table
3079 .filter(group_metadata::group_jid.eq(&group_jid))
3080 .filter(group_metadata::device_id.eq(device_id)),
3081 )
3082 .execute(&mut conn)
3083 .map_err(|e| StoreError::Database(Box::new(e)))?;
3084 Ok(())
3085 })
3086 .await
3087 .map_err(|e| StoreError::Database(Box::new(e)))??;
3088 Ok(())
3089 }
3090
3091 async fn get_tc_token(&self, jid: &str) -> Result<Option<TcTokenEntry>> {
3092 let pool = self.pool.clone();
3093 let device_id = self.device_id;
3094 let jid = jid.to_string();
3095 tokio::task::spawn_blocking(move || -> Result<Option<TcTokenEntry>> {
3096 let mut conn = pool
3097 .get()
3098 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3099 let row: Option<(Vec<u8>, i64, Option<i64>)> = tc_tokens::table
3100 .select((
3101 tc_tokens::token,
3102 tc_tokens::token_timestamp,
3103 tc_tokens::sender_timestamp,
3104 ))
3105 .filter(tc_tokens::jid.eq(&jid))
3106 .filter(tc_tokens::device_id.eq(device_id))
3107 .first(&mut conn)
3108 .optional()
3109 .map_err(|e| StoreError::Database(Box::new(e)))?;
3110 Ok(
3111 row.map(|(token, token_timestamp, sender_timestamp)| TcTokenEntry {
3112 token,
3113 token_timestamp,
3114 sender_timestamp,
3115 }),
3116 )
3117 })
3118 .await
3119 .map_err(|e| StoreError::Database(Box::new(e)))?
3120 }
3121
3122 async fn put_tc_token(&self, jid: &str, entry: &TcTokenEntry) -> Result<()> {
3123 let pool = self.pool.clone();
3124 let device_id = self.device_id;
3125 let jid = jid.to_string();
3126 let entry = entry.clone();
3127 let now = wacore::time::now_secs();
3128 tokio::task::spawn_blocking(move || -> Result<()> {
3129 let mut conn = pool
3130 .get()
3131 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3132 diesel::insert_into(tc_tokens::table)
3133 .values((
3134 tc_tokens::jid.eq(&jid),
3135 tc_tokens::token.eq(&entry.token),
3136 tc_tokens::token_timestamp.eq(entry.token_timestamp),
3137 tc_tokens::sender_timestamp.eq(entry.sender_timestamp),
3138 tc_tokens::device_id.eq(device_id),
3139 tc_tokens::updated_at.eq(now),
3140 ))
3141 .on_conflict((tc_tokens::jid, tc_tokens::device_id))
3142 .do_update()
3143 .set((
3144 tc_tokens::token.eq(&entry.token),
3145 tc_tokens::token_timestamp.eq(entry.token_timestamp),
3146 tc_tokens::sender_timestamp.eq(entry.sender_timestamp),
3147 tc_tokens::updated_at.eq(now),
3148 ))
3149 .execute(&mut conn)
3150 .map_err(|e| StoreError::Database(Box::new(e)))?;
3151 Ok(())
3152 })
3153 .await
3154 .map_err(|e| StoreError::Database(Box::new(e)))??;
3155 Ok(())
3156 }
3157
3158 async fn delete_tc_token(&self, jid: &str) -> Result<()> {
3159 let pool = self.pool.clone();
3160 let device_id = self.device_id;
3161 let jid = jid.to_string();
3162 tokio::task::spawn_blocking(move || -> Result<()> {
3163 let mut conn = pool
3164 .get()
3165 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3166 diesel::delete(
3167 tc_tokens::table
3168 .filter(tc_tokens::jid.eq(&jid))
3169 .filter(tc_tokens::device_id.eq(device_id)),
3170 )
3171 .execute(&mut conn)
3172 .map_err(|e| StoreError::Database(Box::new(e)))?;
3173 Ok(())
3174 })
3175 .await
3176 .map_err(|e| StoreError::Database(Box::new(e)))??;
3177 Ok(())
3178 }
3179
3180 async fn get_all_tc_token_jids(&self) -> Result<Vec<String>> {
3181 let pool = self.pool.clone();
3182 let device_id = self.device_id;
3183 tokio::task::spawn_blocking(move || -> Result<Vec<String>> {
3184 let mut conn = pool
3185 .get()
3186 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3187 let jids: Vec<String> = tc_tokens::table
3188 .select(tc_tokens::jid)
3189 .filter(tc_tokens::device_id.eq(device_id))
3190 .load(&mut conn)
3191 .map_err(|e| StoreError::Database(Box::new(e)))?;
3192 Ok(jids)
3193 })
3194 .await
3195 .map_err(|e| StoreError::Database(Box::new(e)))?
3196 }
3197
3198 async fn delete_expired_tc_tokens(&self, token_cutoff: i64, sender_cutoff: i64) -> Result<u32> {
3199 let pool = self.pool.clone();
3200 let device_id = self.device_id;
3201 tokio::task::spawn_blocking(move || -> Result<u32> {
3202 let mut conn = pool
3203 .get()
3204 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3205 let deleted = diesel::delete(
3210 tc_tokens::table
3211 .filter(
3212 tc_tokens::token
3213 .eq(Vec::<u8>::new())
3214 .or(tc_tokens::token_timestamp.lt(token_cutoff)),
3215 )
3216 .filter(
3217 tc_tokens::sender_timestamp
3218 .is_null()
3219 .or(tc_tokens::sender_timestamp.lt(sender_cutoff)),
3220 )
3221 .filter(tc_tokens::device_id.eq(device_id)),
3222 )
3223 .execute(&mut conn)
3224 .map_err(|e| StoreError::Database(Box::new(e)))?;
3225 Ok(deleted as u32)
3226 })
3227 .await
3228 .map_err(|e| StoreError::Database(Box::new(e)))?
3229 }
3230
3231 async fn store_received_tc_token(
3232 &self,
3233 jid: &str,
3234 token: &[u8],
3235 token_timestamp: i64,
3236 ) -> Result<()> {
3237 let device_id = self.device_id;
3238 let jid = jid.to_string();
3239 let token = token.to_vec();
3240 let now = wacore::time::now_secs();
3241 self.with_retry("store_received_tc_token", || {
3246 let jid = jid.clone();
3247 let token = token.clone();
3248 Box::new(move |conn: &mut SqliteConnection| {
3249 conn.immediate_transaction(|conn| -> QueryResult<()> {
3250 let existing: Option<(Vec<u8>, i64)> = tc_tokens::table
3251 .filter(tc_tokens::jid.eq(&jid))
3252 .filter(tc_tokens::device_id.eq(device_id))
3253 .select((tc_tokens::token, tc_tokens::token_timestamp))
3254 .first(conn)
3255 .optional()?;
3256 let write = match &existing {
3257 Some((existing_token, existing_ts)) => {
3258 existing_token.is_empty() || token_timestamp >= *existing_ts
3259 }
3260 None => true,
3261 };
3262 if write {
3263 diesel::insert_into(tc_tokens::table)
3264 .values((
3265 tc_tokens::jid.eq(&jid),
3266 tc_tokens::token.eq(&token),
3267 tc_tokens::token_timestamp.eq(token_timestamp),
3268 tc_tokens::sender_timestamp.eq(None::<i64>),
3269 tc_tokens::device_id.eq(device_id),
3270 tc_tokens::updated_at.eq(now),
3271 ))
3272 .on_conflict((tc_tokens::jid, tc_tokens::device_id))
3273 .do_update()
3274 .set((
3275 tc_tokens::token.eq(&token),
3276 tc_tokens::token_timestamp.eq(token_timestamp),
3277 tc_tokens::updated_at.eq(now),
3278 ))
3279 .execute(conn)?;
3280 }
3281 Ok(())
3282 })
3283 })
3284 })
3285 .await
3286 }
3287
3288 async fn touch_tc_token_sender_timestamp(
3289 &self,
3290 jid: &str,
3291 sender_timestamp: i64,
3292 ) -> Result<()> {
3293 let pool = self.pool.clone();
3294 let device_id = self.device_id;
3295 let jid = jid.to_string();
3296 let now = wacore::time::now_secs();
3297 tokio::task::spawn_blocking(move || -> Result<()> {
3298 let mut conn = pool
3299 .get()
3300 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3301 diesel::insert_into(tc_tokens::table)
3305 .values((
3306 tc_tokens::jid.eq(&jid),
3307 tc_tokens::token.eq(Vec::<u8>::new()),
3308 tc_tokens::token_timestamp.eq(sender_timestamp),
3309 tc_tokens::sender_timestamp.eq(Some(sender_timestamp)),
3310 tc_tokens::device_id.eq(device_id),
3311 tc_tokens::updated_at.eq(now),
3312 ))
3313 .on_conflict((tc_tokens::jid, tc_tokens::device_id))
3314 .do_update()
3315 .set((
3316 tc_tokens::sender_timestamp.eq(diesel::dsl::sql::<
3320 diesel::sql_types::Nullable<diesel::sql_types::BigInt>,
3321 >(
3322 "MAX(COALESCE(sender_timestamp, "
3323 )
3324 .bind::<diesel::sql_types::BigInt, _>(sender_timestamp)
3325 .sql("), ")
3326 .bind::<diesel::sql_types::BigInt, _>(sender_timestamp)
3327 .sql(")")),
3328 tc_tokens::updated_at.eq(now),
3329 ))
3330 .execute(&mut conn)
3331 .map_err(|e| StoreError::Database(Box::new(e)))?;
3332 Ok(())
3333 })
3334 .await
3335 .map_err(|e| StoreError::Database(Box::new(e)))??;
3336 Ok(())
3337 }
3338
3339 async fn store_sent_message(
3340 &self,
3341 chat_jid: &str,
3342 message_id: &str,
3343 payload: &[u8],
3344 ) -> Result<()> {
3345 let chat_jid = chat_jid.to_string();
3346 let message_id = message_id.to_string();
3347 let payload: Arc<Vec<u8>> = Arc::new(payload.to_vec());
3349 let device_id = self.device_id;
3350 self.with_retry("store_sent_message", || {
3351 let chat_jid = chat_jid.clone();
3352 let message_id = message_id.clone();
3353 let payload = Arc::clone(&payload);
3354 Box::new(move |conn: &mut SqliteConnection| {
3355 diesel::replace_into(sent_messages::table)
3356 .values((
3357 sent_messages::chat_jid.eq(&chat_jid),
3358 sent_messages::message_id.eq(&message_id),
3359 sent_messages::payload.eq(payload.as_slice()),
3360 sent_messages::device_id.eq(device_id),
3361 ))
3362 .execute(conn)?;
3363 Ok(())
3364 })
3365 })
3366 .await
3367 }
3368
3369 async fn take_sent_message(&self, chat_jid: &str, message_id: &str) -> Result<Option<Vec<u8>>> {
3370 let chat_jid = chat_jid.to_string();
3371 let message_id = message_id.to_string();
3372 let device_id = self.device_id;
3373 self.with_retry("take_sent_message", || {
3375 let chat_jid = chat_jid.clone();
3376 let message_id = message_id.clone();
3377 Box::new(move |conn: &mut SqliteConnection| {
3378 conn.immediate_transaction(|conn| {
3379 let row: Option<Vec<u8>> = sent_messages::table
3380 .select(sent_messages::payload)
3381 .filter(sent_messages::chat_jid.eq(&chat_jid))
3382 .filter(sent_messages::message_id.eq(&message_id))
3383 .filter(sent_messages::device_id.eq(device_id))
3384 .first(conn)
3385 .optional()?;
3386 if row.is_some() {
3387 diesel::delete(
3388 sent_messages::table
3389 .filter(sent_messages::chat_jid.eq(&chat_jid))
3390 .filter(sent_messages::message_id.eq(&message_id))
3391 .filter(sent_messages::device_id.eq(device_id)),
3392 )
3393 .execute(conn)?;
3394 }
3395 Ok(row)
3396 })
3397 })
3398 })
3399 .await
3400 }
3401
3402 async fn delete_expired_sent_messages(&self, cutoff_timestamp: i64) -> Result<u32> {
3403 let pool = self.pool.clone();
3404 let device_id = self.device_id;
3405 tokio::task::spawn_blocking(move || -> Result<u32> {
3406 let mut conn = pool
3407 .get()
3408 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3409 let deleted = diesel::delete(
3410 sent_messages::table
3411 .filter(sent_messages::created_at.lt(cutoff_timestamp))
3412 .filter(sent_messages::device_id.eq(device_id)),
3413 )
3414 .execute(&mut conn)
3415 .map_err(|e| StoreError::Database(Box::new(e)))?;
3416 Ok(deleted as u32)
3417 })
3418 .await
3419 .map_err(|e| StoreError::Database(Box::new(e)))?
3420 }
3421
3422 async fn store_pending_inbound(
3423 &self,
3424 chat: &str,
3425 sender: &str,
3426 id: &str,
3427 message: &[u8],
3428 ) -> Result<()> {
3429 let chat = chat.to_string();
3432 let sender = sender.to_string();
3433 let id = id.to_string();
3434 let message: Arc<Vec<u8>> = Arc::new(message.to_vec());
3436 let device_id = self.device_id;
3437 self.with_retry("store_pending_inbound", || {
3438 let chat = chat.clone();
3439 let sender = sender.clone();
3440 let id = id.clone();
3441 let message = Arc::clone(&message);
3442 Box::new(move |conn: &mut SqliteConnection| {
3443 insert_pending_inbound_row(conn, device_id, &chat, &sender, &id, &message)?;
3444 Ok(())
3445 })
3446 })
3447 .await
3448 }
3449
3450 async fn get_pending_inbound(
3451 &self,
3452 chat: &str,
3453 sender: &str,
3454 id: &str,
3455 ) -> Result<Option<Vec<u8>>> {
3456 let chat = chat.to_string();
3457 let sender = sender.to_string();
3458 let id = id.to_string();
3459 let device_id = self.device_id;
3460 self.with_retry("get_pending_inbound", || {
3463 let chat = chat.clone();
3464 let sender = sender.clone();
3465 let id = id.clone();
3466 Box::new(move |conn: &mut SqliteConnection| {
3467 let row: Option<Vec<u8>> = pending_inbound_messages::table
3468 .select(pending_inbound_messages::message)
3469 .filter(pending_inbound_messages::chat.eq(&chat))
3470 .filter(pending_inbound_messages::sender.eq(&sender))
3471 .filter(pending_inbound_messages::id.eq(&id))
3472 .filter(pending_inbound_messages::device_id.eq(device_id))
3473 .first(conn)
3474 .optional()?;
3475 Ok(row)
3476 })
3477 })
3478 .await
3479 }
3480
3481 async fn delete_pending_inbound(&self, chat: &str, sender: &str, id: &str) -> Result<()> {
3482 let chat = chat.to_string();
3483 let sender = sender.to_string();
3484 let id = id.to_string();
3485 let device_id = self.device_id;
3486 self.with_retry("delete_pending_inbound", || {
3487 let chat = chat.clone();
3488 let sender = sender.clone();
3489 let id = id.clone();
3490 Box::new(move |conn: &mut SqliteConnection| {
3491 delete_pending_inbound_row(conn, device_id, &chat, &sender, &id)?;
3492 Ok(())
3493 })
3494 })
3495 .await
3496 }
3497
3498 async fn delete_expired_pending_inbound(&self, cutoff_timestamp: i64) -> Result<u32> {
3499 let pool = self.pool.clone();
3500 let device_id = self.device_id;
3501 tokio::task::spawn_blocking(move || -> Result<u32> {
3502 let mut conn = pool
3503 .get()
3504 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3505 let deleted = diesel::delete(
3506 pending_inbound_messages::table
3507 .filter(pending_inbound_messages::inserted_at.lt(cutoff_timestamp))
3508 .filter(pending_inbound_messages::device_id.eq(device_id)),
3509 )
3510 .execute(&mut conn)
3511 .map_err(|e| StoreError::Database(Box::new(e)))?;
3512 Ok(deleted as u32)
3513 })
3514 .await
3515 .map_err(|e| StoreError::Database(Box::new(e)))?
3516 }
3517
3518 async fn store_pending_inbound_batch(&self, rows: &[PendingInboundRow<'_>]) -> Result<()> {
3519 if rows.is_empty() {
3520 return Ok(());
3521 }
3522 let rows: Arc<Vec<(String, String, String, Vec<u8>)>> = Arc::new(
3525 rows.iter()
3526 .map(|r| {
3527 (
3528 r.chat.to_string(),
3529 r.sender.to_string(),
3530 r.id.to_string(),
3531 r.message.to_vec(),
3532 )
3533 })
3534 .collect(),
3535 );
3536 let device_id = self.device_id;
3537 self.with_retry("store_pending_inbound_batch", || {
3538 let rows = Arc::clone(&rows);
3539 Box::new(move |conn: &mut SqliteConnection| {
3540 conn.transaction(|conn| {
3546 for (chat, sender, id, message) in rows.iter() {
3547 insert_pending_inbound_row(conn, device_id, chat, sender, id, message)?;
3548 }
3549 Ok(())
3550 })
3551 })
3552 })
3553 .await
3554 }
3555
3556 async fn delete_pending_inbound_batch(&self, keys: &[PendingInboundKey<'_>]) -> Result<()> {
3557 if keys.is_empty() {
3558 return Ok(());
3559 }
3560 let keys: Arc<Vec<(String, String, String)>> = Arc::new(
3561 keys.iter()
3562 .map(|k| (k.chat.to_string(), k.sender.to_string(), k.id.to_string()))
3563 .collect(),
3564 );
3565 let device_id = self.device_id;
3566 self.with_retry("delete_pending_inbound_batch", || {
3567 let keys = Arc::clone(&keys);
3568 Box::new(move |conn: &mut SqliteConnection| {
3569 conn.transaction(|conn| {
3573 for (chat, sender, id) in keys.iter() {
3574 delete_pending_inbound_row(conn, device_id, chat, sender, id)?;
3575 }
3576 Ok(())
3577 })
3578 })
3579 })
3580 .await
3581 }
3582}
3583
3584#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
3585#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
3586impl MsgSecretStore for SqliteStore {
3587 async fn put_msg_secrets(&self, entries: Vec<MsgSecretEntry>) -> Result<usize> {
3588 if entries.is_empty() {
3589 return Ok(0);
3590 }
3591
3592 let device_id = self.device_id;
3593 let entries = Arc::new(entries);
3597 let now = wacore::time::now_secs();
3598 self.with_retry("put_msg_secrets", || {
3599 let entries = Arc::clone(&entries);
3600 Box::new(move |conn: &mut SqliteConnection| {
3601 conn.immediate_transaction(|conn| {
3602 let mut stored = 0usize;
3603 for chunk in entries.chunks(MSG_SECRET_INSERT_CHUNK_SIZE) {
3604 let records: Vec<_> = chunk
3608 .iter()
3609 .map(|entry| {
3610 (
3611 msg_secrets::chat.eq(entry.chat.as_ref()),
3612 msg_secrets::sender.eq(entry.sender.as_ref()),
3613 msg_secrets::msg_id.eq(entry.msg_id.as_ref()),
3614 msg_secrets::secret.eq(entry.secret.as_ref()),
3615 msg_secrets::device_id.eq(device_id),
3616 msg_secrets::created_at.eq(now),
3617 msg_secrets::expires_at.eq(entry.expires_at),
3618 msg_secrets::message_ts.eq(entry.message_ts),
3619 )
3620 })
3621 .collect();
3622 stored += diesel::insert_into(msg_secrets::table)
3623 .values(&records)
3624 .on_conflict((
3625 msg_secrets::chat,
3626 msg_secrets::sender,
3627 msg_secrets::msg_id,
3628 msg_secrets::device_id,
3629 ))
3630 .do_update()
3631 .set((
3632 msg_secrets::secret.eq(excluded(msg_secrets::secret)),
3633 msg_secrets::created_at.eq(now),
3634 msg_secrets::expires_at.eq(diesel::dsl::sql::<
3638 diesel::sql_types::BigInt,
3639 >(
3640 "CASE WHEN msg_secrets.expires_at = 0 \
3641 OR excluded.expires_at = 0 THEN 0 \
3642 ELSE MAX(msg_secrets.expires_at, excluded.expires_at) END",
3643 )),
3644 msg_secrets::message_ts.eq(diesel::dsl::sql::<
3647 diesel::sql_types::BigInt,
3648 >(
3649 "MAX(msg_secrets.message_ts, excluded.message_ts)",
3650 )),
3651 ))
3652 .execute(conn)?;
3653 }
3654 Ok(stored)
3655 })
3656 })
3657 })
3658 .await
3659 }
3660
3661 async fn get_msg_secret(
3662 &self,
3663 chat: &str,
3664 sender: &str,
3665 msg_id: &str,
3666 ) -> Result<Option<Vec<u8>>> {
3667 let pool = self.pool.clone();
3671 let device_id = self.device_id;
3672 let chat = chat.to_string();
3673 let sender = sender.to_string();
3674 let msg_id = msg_id.to_string();
3675 self.with_semaphore(move || -> Result<Option<Vec<u8>>> {
3676 let mut conn = pool
3677 .get()
3678 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3679 let row: Option<Vec<u8>> = msg_secrets::table
3680 .select(msg_secrets::secret)
3681 .filter(msg_secrets::chat.eq(&chat))
3682 .filter(msg_secrets::sender.eq(&sender))
3683 .filter(msg_secrets::msg_id.eq(&msg_id))
3684 .filter(msg_secrets::device_id.eq(device_id))
3685 .first(&mut conn)
3686 .optional()
3687 .map_err(|e| StoreError::Database(Box::new(e)))?;
3688 Ok(row)
3689 })
3690 .await
3691 }
3692
3693 async fn get_msg_secret_with_ts(
3694 &self,
3695 chat: &str,
3696 sender: &str,
3697 msg_id: &str,
3698 ) -> Result<Option<(Vec<u8>, i64)>> {
3699 let pool = self.pool.clone();
3704 let device_id = self.device_id;
3705 let chat = chat.to_string();
3706 let sender = sender.to_string();
3707 let msg_id = msg_id.to_string();
3708 self.with_semaphore(move || -> Result<Option<(Vec<u8>, i64)>> {
3709 let mut conn = pool
3710 .get()
3711 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3712 let row: Option<(Vec<u8>, i64)> = msg_secrets::table
3713 .select((msg_secrets::secret, msg_secrets::message_ts))
3714 .filter(msg_secrets::chat.eq(&chat))
3715 .filter(msg_secrets::sender.eq(&sender))
3716 .filter(msg_secrets::msg_id.eq(&msg_id))
3717 .filter(msg_secrets::device_id.eq(device_id))
3718 .first(&mut conn)
3719 .optional()
3720 .map_err(|e| StoreError::Database(Box::new(e)))?;
3721 Ok(row)
3722 })
3723 .await
3724 }
3725
3726 async fn delete_expired_msg_secrets(&self, cutoff_timestamp: i64) -> Result<u32> {
3727 let device_id = self.device_id;
3728 self.with_retry("delete_expired_msg_secrets", || {
3729 Box::new(move |conn: &mut SqliteConnection| {
3730 let deleted = diesel::delete(
3732 msg_secrets::table
3733 .filter(msg_secrets::expires_at.ne(0))
3734 .filter(msg_secrets::expires_at.le(cutoff_timestamp))
3735 .filter(msg_secrets::device_id.eq(device_id)),
3736 )
3737 .execute(conn)?;
3738 Ok(deleted as u32)
3739 })
3740 })
3741 .await
3742 }
3743}
3744
3745#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
3746#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
3747impl DeviceStore for SqliteStore {
3748 async fn save(&self, device: &CoreDevice) -> Result<()> {
3749 SqliteStore::save_device_data_for_device(self, self.device_id, device).await
3750 }
3751
3752 async fn load(&self) -> Result<Option<CoreDevice>> {
3753 SqliteStore::load_device_data_for_device(self, self.device_id).await
3754 }
3755
3756 async fn exists(&self) -> Result<bool> {
3757 SqliteStore::device_exists(self, self.device_id).await
3758 }
3759
3760 async fn create(&self) -> Result<i32> {
3761 SqliteStore::create_new_device(self).await
3762 }
3763
3764 async fn snapshot_db(&self, name: &str, extra_content: Option<&[u8]>) -> Result<()> {
3765 fn sanitize_snapshot_name(name: &str) -> Result<String> {
3766 const MAX_LENGTH: usize = 100;
3767
3768 let sanitized: String = name
3769 .chars()
3770 .map(|c| {
3771 if c.is_ascii_alphanumeric() || c == '_' || c == '-' || c == '.' {
3772 c
3773 } else {
3774 '_'
3775 }
3776 })
3777 .collect();
3778
3779 let sanitized = sanitized
3780 .split('.')
3781 .filter(|part| !part.is_empty() && *part != "..")
3782 .collect::<Vec<_>>()
3783 .join(".");
3784
3785 let sanitized = sanitized.trim_matches(['/', '\\', '.']);
3786
3787 if sanitized.is_empty() {
3788 return Err(StoreError::InvalidConfig(
3789 "Snapshot name cannot be empty after sanitization".to_string(),
3790 ));
3791 }
3792
3793 if sanitized.len() > MAX_LENGTH {
3794 return Err(StoreError::InvalidConfig(format!(
3795 "Snapshot name exceeds maximum length of {} characters",
3796 MAX_LENGTH
3797 )));
3798 }
3799
3800 Ok(sanitized.to_string())
3801 }
3802
3803 let sanitized_name = sanitize_snapshot_name(name)?;
3804
3805 let pool = self.pool.clone();
3806 let db_path = self.database_path.clone();
3807 let extra_data = extra_content.map(|b| b.to_vec());
3808
3809 tokio::task::spawn_blocking(move || -> Result<()> {
3810 let mut conn = pool
3811 .get()
3812 .map_err(|e| StoreError::Connection(Box::new(e)))?;
3813
3814 let timestamp = wacore::time::now_secs();
3815
3816 let target_path = format!("{}.snapshot-{}-{}", db_path, timestamp, sanitized_name);
3818
3819 let query = format!("VACUUM INTO '{}'", target_path.replace("'", "''"));
3822
3823 diesel::sql_query(query)
3824 .execute(&mut conn)
3825 .map_err(|e| StoreError::Database(Box::new(e)))?;
3826
3827 if let Some(data) = extra_data {
3829 let extra_path = format!("{}.json", target_path);
3830 std::fs::write(&extra_path, data)?;
3831 }
3832
3833 Ok(())
3834 })
3835 .await
3836 .map_err(|e| StoreError::Database(Box::new(e)))??;
3837
3838 Ok(())
3839 }
3840
3841 async fn resource_report(&self) -> wacore::stats::StorageResourceReport {
3860 let pool = self.pool.clone();
3861 let read_pool = self.reads.as_ref().map(|reads| reads.pool.clone());
3865 tokio::task::spawn_blocking(move || {
3866 let Some(mut conn) = pool.try_get() else {
3871 return wacore::stats::StorageResourceReport::default();
3872 };
3873 let (Some(page_size), Some(page_count), Some(cache_size)) = (
3877 pragma_i64(&mut conn, "page_size"),
3878 pragma_i64(&mut conn, "page_count"),
3879 pragma_i64(&mut conn, "cache_size"),
3880 ) else {
3881 return wacore::stats::StorageResourceReport::default();
3882 };
3883 let page_size = page_size.max(0) as u64;
3884 let page_count = page_count.max(0) as u64;
3885 let cache_cap_bytes = if cache_size < 0 {
3887 cache_size.unsigned_abs().saturating_mul(1024)
3888 } else {
3889 (cache_size as u64).saturating_mul(page_size)
3890 };
3891 let db_bytes = page_count.saturating_mul(page_size);
3892 let per_conn_cache = cache_cap_bytes.min(db_bytes);
3893 let open_connections = pool.state().connections.max(1) as u64
3897 + read_pool.map_or(0, |reads| reads.state().connections as u64);
3898 wacore::stats::StorageResourceReport {
3899 memory_bytes: Some(per_conn_cache.saturating_mul(open_connections)),
3900 pages: Some(page_count),
3901 ..Default::default()
3902 }
3903 })
3904 .await
3905 .unwrap_or_default()
3906 }
3907}
3908
3909fn pragma_i64(conn: &mut SqliteConnection, pragma: &str) -> Option<i64> {
3913 if pragma.is_empty()
3917 || !pragma
3918 .bytes()
3919 .all(|b| b.is_ascii_alphanumeric() || b == b'_')
3920 {
3921 debug_assert!(false, "pragma_i64 requires an identifier, got {pragma:?}");
3922 return None;
3923 }
3924 #[derive(diesel::QueryableByName)]
3925 struct Row {
3926 #[diesel(sql_type = diesel::sql_types::BigInt)]
3927 value: i64,
3928 }
3929 let sql = format!("SELECT {pragma} AS value FROM pragma_{pragma}()");
3932 diesel::sql_query(sql)
3933 .get_result::<Row>(conn)
3934 .ok()
3935 .map(|r| r.value)
3936}
3937
3938#[cfg(test)]
3939mod tests {
3940 use super::*;
3941
3942 async fn create_test_store() -> SqliteStore {
3943 use portable_atomic::AtomicU64;
3944 use std::sync::atomic::Ordering;
3945 static COUNTER: AtomicU64 = AtomicU64::new(0);
3946 let id = COUNTER.fetch_add(1, Ordering::Relaxed);
3947 let db_name = format!(
3948 "file:memdb_test_{}_{}?mode=memory&cache=shared",
3949 std::process::id(),
3950 id
3951 );
3952 SqliteStore::new(&db_name)
3953 .await
3954 .expect("Failed to create test store")
3955 }
3956
3957 #[tokio::test]
3958 async fn with_config_custom_tuning_builds_and_operates() {
3959 use portable_atomic::AtomicU64;
3960 use std::sync::atomic::Ordering;
3961 static COUNTER: AtomicU64 = AtomicU64::new(0);
3962 let id = COUNTER.fetch_add(1, Ordering::Relaxed);
3963 let db_name = format!(
3964 "file:memdb_cfg_{}_{}?mode=memory&cache=shared",
3965 std::process::id(),
3966 id
3967 );
3968
3969 let def = SqliteStoreConfig::default();
3971 assert_eq!(def.pool_size, 1);
3972 assert_eq!(def.cache_size_kib, 512);
3973
3974 let config = SqliteStoreConfig {
3977 pool_size: 2,
3978 read_pool_size: 0,
3979 cache_size_kib: 4096,
3980 mmap_size: None,
3981 busy_timeout: Duration::from_secs(7),
3982 synchronous: Synchronous::Full,
3983 thread_pool: Some(Arc::new(
3984 scheduled_thread_pool::ScheduledThreadPool::builder()
3985 .num_threads(1)
3986 .build(),
3987 )),
3988 connection_init: None,
3989 };
3990 let store = SqliteStore::with_config(&db_name, config)
3991 .await
3992 .expect("custom-config store");
3993
3994 let mac = AppStateMutationMAC {
3995 index_mac: vec![1u8; 32],
3996 value_mac: vec![2u8; 32],
3997 };
3998 store
3999 .put_app_state_mutation_macs_for_device("c", 1, std::slice::from_ref(&mac), 1)
4000 .await
4001 .unwrap();
4002 let got = store
4003 .get_app_state_mutation_mac_for_device("c", &mac.index_mac, 1)
4004 .await
4005 .unwrap();
4006 assert_eq!(got, Some(mac.value_mac));
4007
4008 #[derive(diesel::QueryableByName)]
4011 struct CacheSync {
4012 #[diesel(sql_type = diesel::sql_types::BigInt)]
4013 cache: i64,
4014 #[diesel(sql_type = diesel::sql_types::BigInt)]
4015 sync: i64,
4016 }
4017 #[derive(diesel::QueryableByName)]
4018 struct Busy {
4019 #[diesel(sql_type = diesel::sql_types::BigInt)]
4020 timeout: i64,
4021 }
4022 let mut conn = store.pool.get().unwrap();
4023 let cs: CacheSync = diesel::sql_query(
4024 "SELECT cs.cache_size AS cache, sy.synchronous AS sync \
4025 FROM pragma_cache_size cs, pragma_synchronous sy",
4026 )
4027 .get_result(&mut conn)
4028 .unwrap();
4029 let busy: Busy = diesel::sql_query("PRAGMA busy_timeout")
4030 .get_result(&mut conn)
4031 .unwrap();
4032 assert_eq!(cs.cache, -4096, "cache_size_kib applied as negative KiB");
4033 assert_eq!(cs.sync, 2, "synchronous = FULL");
4034 assert_eq!(busy.timeout, 7000, "busy_timeout = 7s");
4035 }
4036
4037 #[tokio::test]
4038 async fn connection_init_runs_before_pragmas_and_migrations() {
4039 use portable_atomic::AtomicU64;
4040 use std::sync::atomic::{AtomicBool, Ordering};
4041 static COUNTER: AtomicU64 = AtomicU64::new(0);
4042 let id = COUNTER.fetch_add(1, Ordering::Relaxed);
4043 let db_name = format!(
4044 "file:memdb_init_{}_{}?mode=memory&cache=shared",
4045 std::process::id(),
4046 id
4047 );
4048
4049 let calls = Arc::new(AtomicU64::new(0));
4050 let saw_migrations_table = Arc::new(AtomicBool::new(false));
4051 let saw_store_pragmas = Arc::new(AtomicBool::new(false));
4052
4053 #[derive(diesel::QueryableByName)]
4054 struct Count {
4055 #[diesel(sql_type = diesel::sql_types::BigInt)]
4056 n: i64,
4057 }
4058 let config = {
4059 let calls = calls.clone();
4060 let saw_migrations_table = saw_migrations_table.clone();
4061 let saw_store_pragmas = saw_store_pragmas.clone();
4062 SqliteStoreConfig::default().with_connection_init(move |conn| {
4063 calls.fetch_add(1, Ordering::Relaxed);
4064 let migrated: Count = diesel::sql_query(
4065 "SELECT count(*) AS n FROM sqlite_master \
4066 WHERE name = '__diesel_schema_migrations'",
4067 )
4068 .get_result(conn)?;
4069 if migrated.n > 0 {
4070 saw_migrations_table.store(true, Ordering::Relaxed);
4071 }
4072 let busy: Count = diesel::sql_query("SELECT timeout AS n FROM pragma_busy_timeout")
4075 .get_result(conn)?;
4076 if busy.n != 0 {
4077 saw_store_pragmas.store(true, Ordering::Relaxed);
4078 }
4079 Ok(())
4080 })
4081 };
4082
4083 let store = SqliteStore::with_config(&db_name, config)
4084 .await
4085 .expect("store with connection_init");
4086
4087 assert!(calls.load(Ordering::Relaxed) >= 1, "hook ran");
4088 assert!(
4089 !saw_migrations_table.load(Ordering::Relaxed),
4090 "hook ran before migrations on the first connection"
4091 );
4092 assert!(
4093 !saw_store_pragmas.load(Ordering::Relaxed),
4094 "hook ran before the store's own pragmas"
4095 );
4096
4097 let mac = AppStateMutationMAC {
4099 index_mac: vec![3u8; 32],
4100 value_mac: vec![4u8; 32],
4101 };
4102 store
4103 .put_app_state_mutation_macs_for_device("ci", 1, std::slice::from_ref(&mac), 1)
4104 .await
4105 .unwrap();
4106 assert_eq!(
4107 store
4108 .get_app_state_mutation_mac_for_device("ci", &mac.index_mac, 1)
4109 .await
4110 .unwrap(),
4111 Some(mac.value_mac)
4112 );
4113 }
4114
4115 #[test]
4116 fn connection_init_error_rejects_connection_before_pragmas() {
4117 let mut conn = SqliteConnection::establish(":memory:").expect("raw connection");
4118 let options = ConnectionOptions {
4119 cache_size_kib: 512,
4120 mmap_size: None,
4121 busy_timeout_ms: 30_000,
4122 synchronous: Synchronous::Normal,
4123 connection_init: Some(Arc::new(|_conn: &mut SqliteConnection| {
4124 Err("wrong key".into())
4125 })),
4126 query_only: false,
4127 };
4128
4129 use diesel::r2d2::CustomizeConnection;
4130 let err = options
4131 .on_acquire(&mut conn)
4132 .expect_err("hook error surfaces");
4133 assert!(err.to_string().contains("wrong key"));
4134
4135 #[derive(diesel::QueryableByName)]
4137 struct Busy {
4138 #[diesel(sql_type = diesel::sql_types::BigInt)]
4139 timeout: i64,
4140 }
4141 let busy: Busy = diesel::sql_query("PRAGMA busy_timeout")
4142 .get_result(&mut conn)
4143 .unwrap();
4144 assert_eq!(busy.timeout, 0);
4145 }
4146
4147 #[tokio::test]
4148 async fn batch_mutation_macs_matches_per_item() {
4149 let store = create_test_store().await;
4150 let name = "regular";
4151 let device_id = 1;
4152
4153 let macs: Vec<AppStateMutationMAC> = (0..25u8)
4154 .map(|i| {
4155 let mut index_mac = vec![0u8; 32];
4156 index_mac[0] = i;
4157 AppStateMutationMAC {
4158 index_mac,
4159 value_mac: vec![i; 32],
4160 }
4161 })
4162 .collect();
4163 store
4164 .put_app_state_mutation_macs_for_device(name, 1, &macs, device_id)
4165 .await
4166 .unwrap();
4167
4168 let mut index_macs: Vec<[u8; 32]> = macs
4169 .iter()
4170 .map(|m| m.index_mac.as_slice().try_into().unwrap())
4171 .collect();
4172 index_macs.push([0xFF; 32]);
4174
4175 let batch = store
4176 .get_app_state_mutation_macs_batch_for_device(name, &index_macs, device_id)
4177 .await
4178 .unwrap();
4179
4180 assert_eq!(batch.len(), macs.len());
4181 assert!(!batch.contains_key(&[0xFF; 32]));
4182 for m in &macs {
4183 let key: [u8; 32] = m.index_mac.as_slice().try_into().unwrap();
4184 let per_item = store
4186 .get_app_state_mutation_mac_for_device(name, &m.index_mac, device_id)
4187 .await
4188 .unwrap();
4189 assert_eq!(per_item.as_ref(), batch.get(&key));
4190 assert_eq!(batch.get(&key), Some(&m.value_mac));
4191 }
4192
4193 let empty = store
4195 .get_app_state_mutation_macs_batch_for_device(name, &[], device_id)
4196 .await
4197 .unwrap();
4198 assert!(empty.is_empty());
4199 }
4200
4201 #[tokio::test]
4202 async fn clear_mutation_macs_wipes_only_named_collection() {
4203 let store = create_test_store().await;
4204 let mac = |i: u8| AppStateMutationMAC {
4205 index_mac: vec![i; 32],
4206 value_mac: vec![i; 32],
4207 };
4208 store
4209 .put_mutation_macs("regular", 1, &[mac(1)])
4210 .await
4211 .unwrap();
4212 store
4213 .put_mutation_macs("critical", 1, &[mac(2)])
4214 .await
4215 .unwrap();
4216
4217 store.clear_mutation_macs("regular").await.unwrap();
4218
4219 assert!(
4220 store
4221 .get_mutation_mac("regular", &[1; 32])
4222 .await
4223 .unwrap()
4224 .is_none()
4225 );
4226 assert!(
4227 store
4228 .get_mutation_mac("critical", &[2; 32])
4229 .await
4230 .unwrap()
4231 .is_some()
4232 );
4233 }
4234
4235 #[tokio::test]
4236 async fn put_signal_batches_persist_and_upsert() {
4237 use std::sync::Arc;
4238 let store = create_test_store().await;
4239
4240 let sessions: Vec<(Arc<str>, Bytes)> = (0..5u8)
4241 .map(|i| {
4242 (
4243 Arc::from(format!("user{i}@s.whatsapp.net").as_str()),
4244 Bytes::from(vec![i; 8]),
4245 )
4246 })
4247 .collect();
4248 store.put_sessions_batch(&sessions).await.unwrap();
4249 for (addr, bytes) in &sessions {
4250 assert_eq!(
4251 store.get_session(addr).await.unwrap().as_deref(),
4252 Some(bytes.as_ref())
4253 );
4254 }
4255
4256 let identities: Vec<(Arc<str>, [u8; 32])> = (0..5u8)
4257 .map(|i| {
4258 (
4259 Arc::from(format!("user{i}@s.whatsapp.net").as_str()),
4260 [i; 32],
4261 )
4262 })
4263 .collect();
4264 store.put_identities_batch(&identities).await.unwrap();
4265 for (addr, key) in &identities {
4266 assert_eq!(store.load_identity(addr).await.unwrap(), Some(*key));
4267 }
4268
4269 let sender_keys: Vec<(Arc<str>, Bytes)> = (0..5u8)
4270 .map(|i| {
4271 (
4272 Arc::from(format!("g@g.us::user{i}").as_str()),
4273 Bytes::from(vec![i; 16]),
4274 )
4275 })
4276 .collect();
4277 store.put_sender_keys_batch(&sender_keys).await.unwrap();
4278 for (addr, bytes) in &sender_keys {
4279 assert_eq!(
4280 store.get_sender_key(addr).await.unwrap().as_deref(),
4281 Some(bytes.as_ref())
4282 );
4283 }
4284
4285 let updated: Vec<(Arc<str>, Bytes)> = sessions
4287 .iter()
4288 .map(|(addr, _)| (addr.clone(), Bytes::from(vec![0xAA; 8])))
4289 .collect();
4290 store.put_sessions_batch(&updated).await.unwrap();
4291 for (addr, _) in &sessions {
4292 assert_eq!(
4293 store.get_session(addr).await.unwrap().as_deref(),
4294 Some([0xAA; 8].as_slice())
4295 );
4296 }
4297
4298 let dup: Arc<str> = Arc::from("dup@s.whatsapp.net");
4301 store
4302 .put_sessions_batch(&[
4303 (dup.clone(), Bytes::from(vec![1u8; 4])),
4304 (dup.clone(), Bytes::from(vec![2u8; 4])),
4305 ])
4306 .await
4307 .unwrap();
4308 assert_eq!(
4309 store.get_session(&dup).await.unwrap().as_deref(),
4310 Some([2u8; 4].as_slice())
4311 );
4312
4313 store.put_sessions_batch(&[]).await.unwrap();
4315 store.put_identities_batch(&[]).await.unwrap();
4316 store.put_sender_keys_batch(&[]).await.unwrap();
4317 }
4318
4319 #[test]
4320 fn test_parse_database_path_regular_path() {
4321 let path = "/var/lib/whatsapp/database.db";
4322 let result = parse_database_path(path).unwrap();
4323 assert_eq!(result, "/var/lib/whatsapp/database.db");
4324 }
4325
4326 #[test]
4327 fn test_parse_database_path_with_sqlite_prefix() {
4328 let path = "sqlite:///var/lib/whatsapp/database.db";
4329 let result = parse_database_path(path).unwrap();
4330 assert_eq!(result, "/var/lib/whatsapp/database.db");
4331 }
4332
4333 #[test]
4334 fn test_parse_database_path_with_query_params() {
4335 let path = "file:database.db?mode=memory&cache=shared";
4336 let result = parse_database_path(path).unwrap();
4337 assert_eq!(result, "file:database.db");
4338 }
4339
4340 #[test]
4341 fn test_parse_database_path_with_fragment() {
4342 let path = "file:database.db#fragment";
4343 let result = parse_database_path(path).unwrap();
4344 assert_eq!(result, "file:database.db");
4345 }
4346
4347 #[test]
4348 fn test_parse_database_path_with_both_query_and_fragment() {
4349 let path = "sqlite:///var/lib/database.db?mode=ro#backup";
4350 let result = parse_database_path(path).unwrap();
4351 assert_eq!(result, "/var/lib/database.db");
4352 }
4353
4354 #[test]
4355 fn test_parse_database_path_in_memory_rejected() {
4356 let result = parse_database_path(":memory:");
4357 assert!(result.is_err());
4358 assert!(result.unwrap_err().to_string().contains("not supported"));
4359 }
4360
4361 #[test]
4362 fn test_parse_database_path_in_memory_with_query_rejected() {
4363 let result = parse_database_path(":memory:?cache=shared");
4364 assert!(result.is_err());
4365 assert!(result.unwrap_err().to_string().contains("not supported"));
4366 }
4367
4368 #[tokio::test]
4369 async fn test_device_registry_save_and_get() {
4370 let store = create_test_store().await;
4371
4372 let record = DeviceListRecord {
4373 user: "1234567890".to_string(),
4374 devices: vec![DeviceInfo::new(0, None), DeviceInfo::new(1, Some(42))],
4375 timestamp: 1234567890,
4376 phash: Some("2:abcdef".to_string()),
4377 raw_id: None,
4378 };
4379
4380 store.update_device_list(record).await.expect("save failed");
4381 let loaded = store
4382 .get_devices("1234567890")
4383 .await
4384 .expect("get failed")
4385 .expect("record should exist");
4386
4387 assert_eq!(loaded.user, "1234567890");
4388 assert_eq!(loaded.devices.len(), 2);
4389 assert_eq!(loaded.devices[0].device_id, 0);
4390 assert_eq!(loaded.devices[1].device_id, 1);
4391 assert_eq!(loaded.devices[1].key_index, Some(42));
4392 assert_eq!(loaded.phash, Some("2:abcdef".to_string()));
4393 }
4394
4395 #[tokio::test]
4396 async fn test_device_registry_update_existing() {
4397 let store = create_test_store().await;
4398
4399 let record1 = DeviceListRecord {
4400 user: "1234567890".to_string(),
4401 devices: vec![DeviceInfo::new(0, None)],
4402 timestamp: 1000,
4403 phash: Some("2:old".to_string()),
4404 raw_id: None,
4405 };
4406 store
4407 .update_device_list(record1)
4408 .await
4409 .expect("save1 failed");
4410
4411 let record2 = DeviceListRecord {
4412 user: "1234567890".to_string(),
4413 devices: vec![DeviceInfo::new(0, None), DeviceInfo::new(2, None)],
4414 timestamp: 2000,
4415 phash: Some("2:new".to_string()),
4416 raw_id: None,
4417 };
4418 store
4419 .update_device_list(record2)
4420 .await
4421 .expect("save2 failed");
4422
4423 let loaded = store
4424 .get_devices("1234567890")
4425 .await
4426 .expect("get failed")
4427 .expect("record should exist");
4428
4429 assert_eq!(loaded.devices.len(), 2);
4430 assert_eq!(loaded.phash, Some("2:new".to_string()));
4431 }
4432
4433 #[tokio::test]
4434 async fn test_device_registry_get_nonexistent() {
4435 let store = create_test_store().await;
4436 let result = store.get_devices("nonexistent").await.expect("get failed");
4437 assert!(result.is_none());
4438 }
4439
4440 #[tokio::test]
4441 async fn test_sender_key_devices_set_and_get() {
4442 let store = create_test_store().await;
4443
4444 let group = "group123@g.us";
4445
4446 store
4448 .set_sender_key_status(group, &[("user1:5@lid", true), ("user2:3@lid", false)])
4449 .await
4450 .expect("set failed");
4451
4452 let devices = store
4453 .get_sender_key_devices(group)
4454 .await
4455 .expect("get failed");
4456 assert_eq!(devices.len(), 2);
4457 assert!(devices.contains(&("user1:5@lid".to_string(), true)));
4458 assert!(devices.contains(&("user2:3@lid".to_string(), false)));
4459 }
4460
4461 #[tokio::test]
4462 async fn test_sender_key_devices_upsert_overwrites() {
4463 let store = create_test_store().await;
4464
4465 let group = "group123@g.us";
4466
4467 store
4469 .set_sender_key_status(group, &[("user1:5@lid", false)])
4470 .await
4471 .expect("set failed");
4472
4473 store
4475 .set_sender_key_status(group, &[("user1:5@lid", true)])
4476 .await
4477 .expect("set failed");
4478
4479 let devices = store
4480 .get_sender_key_devices(group)
4481 .await
4482 .expect("get failed");
4483 assert_eq!(devices.len(), 1);
4484 assert_eq!(devices[0], ("user1:5@lid".to_string(), true));
4485 }
4486
4487 #[tokio::test]
4488 async fn test_sender_key_devices_clear() {
4489 let store = create_test_store().await;
4490
4491 let group = "group123@g.us";
4492
4493 store
4494 .set_sender_key_status(group, &[("user1:5@lid", true), ("user2:3@lid", true)])
4495 .await
4496 .expect("set failed");
4497
4498 store
4499 .clear_sender_key_devices(group)
4500 .await
4501 .expect("clear failed");
4502
4503 let devices = store
4504 .get_sender_key_devices(group)
4505 .await
4506 .expect("get failed");
4507 assert!(devices.is_empty());
4508 }
4509
4510 #[tokio::test]
4511 async fn test_tc_token_put_and_get() {
4512 let store = create_test_store().await;
4513
4514 let entry = TcTokenEntry {
4515 token: vec![1, 2, 3, 4, 5],
4516 token_timestamp: 1707000000,
4517 sender_timestamp: Some(1707000100),
4518 };
4519
4520 store
4521 .put_tc_token("user@lid", &entry)
4522 .await
4523 .expect("put failed");
4524
4525 let loaded = store
4526 .get_tc_token("user@lid")
4527 .await
4528 .expect("get failed")
4529 .expect("should exist");
4530
4531 assert_eq!(loaded.token, vec![1, 2, 3, 4, 5]);
4532 assert_eq!(loaded.token_timestamp, 1707000000);
4533 assert_eq!(loaded.sender_timestamp, Some(1707000100));
4534 }
4535
4536 #[tokio::test]
4537 async fn test_tc_token_upsert() {
4538 let store = create_test_store().await;
4539
4540 let entry1 = TcTokenEntry {
4541 token: vec![1, 2, 3],
4542 token_timestamp: 1000,
4543 sender_timestamp: None,
4544 };
4545 store.put_tc_token("user@lid", &entry1).await.unwrap();
4546
4547 let entry2 = TcTokenEntry {
4548 token: vec![4, 5, 6],
4549 token_timestamp: 2000,
4550 sender_timestamp: Some(1500),
4551 };
4552 store.put_tc_token("user@lid", &entry2).await.unwrap();
4553
4554 let loaded = store.get_tc_token("user@lid").await.unwrap().unwrap();
4555 assert_eq!(loaded.token, vec![4, 5, 6]);
4556 assert_eq!(loaded.token_timestamp, 2000);
4557 assert_eq!(loaded.sender_timestamp, Some(1500));
4558 }
4559
4560 #[tokio::test]
4561 async fn test_tc_token_delete() {
4562 let store = create_test_store().await;
4563
4564 let entry = TcTokenEntry {
4565 token: vec![1, 2, 3],
4566 token_timestamp: 1000,
4567 sender_timestamp: None,
4568 };
4569 store.put_tc_token("user@lid", &entry).await.unwrap();
4570 store.delete_tc_token("user@lid").await.unwrap();
4571
4572 let result = store.get_tc_token("user@lid").await.unwrap();
4573 assert!(result.is_none());
4574 }
4575
4576 #[tokio::test]
4577 async fn test_touch_and_store_received_preserve_each_others_field() {
4578 let store = create_test_store().await;
4579
4580 store
4583 .touch_tc_token_sender_timestamp("user@lid", 5000)
4584 .await
4585 .unwrap();
4586 store
4587 .store_received_tc_token("user@lid", &[7, 8, 9], 4000)
4588 .await
4589 .unwrap();
4590 let a = store.get_tc_token("user@lid").await.unwrap().unwrap();
4591 assert_eq!(a.token, vec![7, 8, 9]);
4592 assert_eq!(a.token_timestamp, 4000);
4593 assert_eq!(a.sender_timestamp, Some(5000));
4594
4595 store
4597 .touch_tc_token_sender_timestamp("user@lid", 6000)
4598 .await
4599 .unwrap();
4600 let b = store.get_tc_token("user@lid").await.unwrap().unwrap();
4601 assert_eq!(b.token, vec![7, 8, 9], "touch must keep the real token");
4602 assert_eq!(b.sender_timestamp, Some(6000));
4603
4604 store
4606 .touch_tc_token_sender_timestamp("user@lid", 1000)
4607 .await
4608 .unwrap();
4609 let c = store.get_tc_token("user@lid").await.unwrap().unwrap();
4610 assert_eq!(c.sender_timestamp, Some(6000), "touch is advance-only");
4611 }
4612
4613 #[tokio::test]
4614 async fn store_received_tc_token_is_newer_wins() {
4615 let store = create_test_store().await;
4616
4617 store
4619 .store_received_tc_token("c@lid", &[1, 1, 1], 5000)
4620 .await
4621 .unwrap();
4622
4623 store
4626 .store_received_tc_token("c@lid", &[2, 2, 2], 3000)
4627 .await
4628 .unwrap();
4629 let e = store.get_tc_token("c@lid").await.unwrap().unwrap();
4630 assert_eq!(e.token, vec![1, 1, 1], "older write must not overwrite");
4631 assert_eq!(e.token_timestamp, 5000);
4632
4633 store
4635 .store_received_tc_token("c@lid", &[3, 3, 3], 7000)
4636 .await
4637 .unwrap();
4638 let e = store.get_tc_token("c@lid").await.unwrap().unwrap();
4639 assert_eq!(e.token, vec![3, 3, 3]);
4640 assert_eq!(e.token_timestamp, 7000);
4641
4642 store
4645 .touch_tc_token_sender_timestamp("p@lid", 9000)
4646 .await
4647 .unwrap();
4648 store
4649 .store_received_tc_token("p@lid", &[4, 4, 4], 6000)
4650 .await
4651 .unwrap();
4652 let e = store.get_tc_token("p@lid").await.unwrap().unwrap();
4653 assert_eq!(e.token, vec![4, 4, 4], "placeholder must accept real token");
4654 assert_eq!(e.token_timestamp, 6000);
4655 assert_eq!(e.sender_timestamp, Some(9000), "sender bucket preserved");
4656 }
4657
4658 #[tokio::test]
4659 async fn test_delete_expired_two_window_pruning() {
4660 let store = create_test_store().await;
4661 store
4665 .touch_tc_token_sender_timestamp("recent_ph@lid", 2500)
4666 .await
4667 .unwrap();
4668 store
4670 .touch_tc_token_sender_timestamp("stale_ph@lid", 100)
4671 .await
4672 .unwrap();
4673 store
4675 .put_tc_token(
4676 "expired_live_sender@lid",
4677 &TcTokenEntry {
4678 token: vec![1],
4679 token_timestamp: 1,
4680 sender_timestamp: Some(2500),
4681 },
4682 )
4683 .await
4684 .unwrap();
4685 store
4687 .put_tc_token(
4688 "orphan_expired@lid",
4689 &TcTokenEntry {
4690 token: vec![2],
4691 token_timestamp: 1,
4692 sender_timestamp: None,
4693 },
4694 )
4695 .await
4696 .unwrap();
4697
4698 let removed = store.delete_expired_tc_tokens(1000, 2000).await.unwrap();
4699 assert_eq!(removed, 2);
4700 assert!(store.get_tc_token("recent_ph@lid").await.unwrap().is_some());
4701 assert!(store.get_tc_token("stale_ph@lid").await.unwrap().is_none());
4702 assert!(
4703 store
4704 .get_tc_token("expired_live_sender@lid")
4705 .await
4706 .unwrap()
4707 .is_some()
4708 );
4709 assert!(
4710 store
4711 .get_tc_token("orphan_expired@lid")
4712 .await
4713 .unwrap()
4714 .is_none()
4715 );
4716 }
4717
4718 #[tokio::test]
4719 async fn test_tc_token_get_all_jids() {
4720 let store = create_test_store().await;
4721
4722 let entry = TcTokenEntry {
4723 token: vec![1],
4724 token_timestamp: 1000,
4725 sender_timestamp: None,
4726 };
4727 store.put_tc_token("user1@lid", &entry).await.unwrap();
4728 store.put_tc_token("user2@lid", &entry).await.unwrap();
4729 store.put_tc_token("user3@lid", &entry).await.unwrap();
4730
4731 let mut jids = store.get_all_tc_token_jids().await.unwrap();
4732 jids.sort();
4733 assert_eq!(jids, vec!["user1@lid", "user2@lid", "user3@lid"]);
4734 }
4735
4736 #[tokio::test]
4737 async fn test_tc_token_delete_expired() {
4738 let store = create_test_store().await;
4739
4740 let old = TcTokenEntry {
4741 token: vec![1],
4742 token_timestamp: 1000,
4743 sender_timestamp: None,
4744 };
4745 let recent = TcTokenEntry {
4746 token: vec![2],
4747 token_timestamp: 5000,
4748 sender_timestamp: None,
4749 };
4750 store.put_tc_token("old@lid", &old).await.unwrap();
4751 store.put_tc_token("recent@lid", &recent).await.unwrap();
4752
4753 let deleted = store.delete_expired_tc_tokens(3000, 3000).await.unwrap();
4755 assert_eq!(deleted, 1);
4756
4757 assert!(store.get_tc_token("old@lid").await.unwrap().is_none());
4758 assert!(store.get_tc_token("recent@lid").await.unwrap().is_some());
4759 }
4760
4761 #[tokio::test]
4762 async fn test_tc_token_get_nonexistent() {
4763 let store = create_test_store().await;
4764 let result = store.get_tc_token("nonexistent@lid").await.unwrap();
4765 assert!(result.is_none());
4766 }
4767
4768 #[tokio::test]
4769 async fn test_sender_key_devices_different_groups() {
4770 let store = create_test_store().await;
4771
4772 let group1 = "group1@g.us";
4773 let group2 = "group2@g.us";
4774
4775 store
4776 .set_sender_key_status(group1, &[("user:5@lid", true)])
4777 .await
4778 .expect("set failed");
4779
4780 let g1 = store.get_sender_key_devices(group1).await.unwrap();
4781 assert_eq!(g1.len(), 1);
4782
4783 let g2 = store.get_sender_key_devices(group2).await.unwrap();
4784 assert!(g2.is_empty());
4785 }
4786
4787 #[tokio::test]
4788 async fn test_create_new_device_uses_configured_device_id() {
4789 use portable_atomic::AtomicU64;
4790 use std::sync::atomic::Ordering;
4791 static COUNTER: AtomicU64 = AtomicU64::new(100);
4792 let id = COUNTER.fetch_add(1, Ordering::Relaxed);
4793 let db_name = format!(
4794 "file:memdb_devid_{}_{}?mode=memory&cache=shared",
4795 std::process::id(),
4796 id
4797 );
4798
4799 let device_id = 42;
4800 let store = SqliteStore::new_for_device(&db_name, device_id)
4801 .await
4802 .expect("Failed to create test store");
4803
4804 assert!(!store.device_exists(device_id).await.unwrap());
4805 let returned_id = store.create_new_device().await.unwrap();
4806 assert_eq!(returned_id, device_id);
4807 assert!(store.device_exists(device_id).await.unwrap());
4808
4809 if device_id != 1 {
4811 assert!(!store.device_exists(1).await.unwrap());
4812 }
4813
4814 let loaded = store.load_device_data_for_device(device_id).await.unwrap();
4815 assert!(
4816 loaded.is_some(),
4817 "device data should be loadable by configured id"
4818 );
4819 }
4820
4821 #[tokio::test]
4824 async fn mark_prekeys_uploaded_never_resurrects_deleted_rows() {
4825 let store = create_test_store().await;
4826 store
4827 .store_prekey(1, b"record-1", false)
4828 .await
4829 .expect("store");
4830 store
4831 .store_prekey(2, b"record-2", false)
4832 .await
4833 .expect("store");
4834 store.remove_prekey(1).await.expect("consume");
4835
4836 store
4837 .mark_prekeys_uploaded(&[1, 2])
4838 .await
4839 .expect("mark uploaded");
4840
4841 let gone = store.load_prekey(1).await.expect("load");
4842 assert!(gone.is_none(), "consumed key must not be resurrected");
4843 let live = store.load_prekey(2).await.expect("load");
4844 assert!(live.is_some(), "live key still present");
4845 }
4846
4847 #[tokio::test]
4852 async fn test_prekey_watermarks_survive_save_load_roundtrip() {
4853 use portable_atomic::AtomicU64;
4854 use std::sync::atomic::Ordering;
4855
4856 static COUNTER: AtomicU64 = AtomicU64::new(300);
4857 let id = COUNTER.fetch_add(1, Ordering::Relaxed);
4858 let db_name = format!(
4859 "file:memdb_pkwatermark_{}_{}?mode=memory&cache=shared",
4860 std::process::id(),
4861 id
4862 );
4863
4864 let device_id = 9;
4865 let _writer = SqliteStore::new_for_device(&db_name, device_id)
4866 .await
4867 .expect("create store");
4868 _writer.create_new_device().await.expect("create device");
4869
4870 let mut device = _writer
4871 .load_device_data_for_device(device_id)
4872 .await
4873 .expect("load")
4874 .expect("device should exist after create");
4875 assert_eq!(
4876 device.first_unupload_pre_key_id, 0,
4877 "fresh device starts with the watermark unset"
4878 );
4879 device.next_pre_key_id = 913;
4880 device.first_unupload_pre_key_id = 101;
4881 _writer
4882 .save_device_data_for_device(device_id, &device)
4883 .await
4884 .expect("save with watermarks");
4885
4886 let store = SqliteStore::new_for_device(&db_name, device_id)
4887 .await
4888 .expect("reopen store");
4889 let loaded = store
4890 .load_device_data_for_device(device_id)
4891 .await
4892 .expect("load")
4893 .expect("device should exist after reopen");
4894 assert_eq!(loaded.next_pre_key_id, 913);
4895 assert_eq!(
4896 loaded.first_unupload_pre_key_id, 101,
4897 "first_unupload_pre_key_id must survive a save/load roundtrip"
4898 );
4899 }
4900
4901 #[tokio::test]
4908 async fn test_server_cert_chain_survives_save_load_roundtrip() {
4909 use portable_atomic::AtomicU64;
4910 use std::sync::atomic::Ordering;
4911 use wacore::store::device::{CachedNoiseCert, CachedServerCertChain};
4912
4913 static COUNTER: AtomicU64 = AtomicU64::new(200);
4914 let id = COUNTER.fetch_add(1, Ordering::Relaxed);
4915 let db_name = format!(
4919 "file:memdb_certchain_{}_{}?mode=memory&cache=shared",
4920 std::process::id(),
4921 id
4922 );
4923
4924 let device_id = 7;
4925 let chain = CachedServerCertChain {
4926 intermediate: CachedNoiseCert {
4927 key: [0xAB; 32],
4928 not_before: 1_700_000_000,
4929 not_after: 1_900_000_000,
4930 },
4931 leaf: CachedNoiseCert {
4932 key: [0xCD; 32],
4933 not_before: 1_700_000_500,
4934 not_after: 1_899_999_500,
4935 },
4936 };
4937
4938 let _writer = SqliteStore::new_for_device(&db_name, device_id)
4944 .await
4945 .expect("create store");
4946 _writer.create_new_device().await.expect("create device");
4947
4948 let mut device = _writer
4949 .load_device_data_for_device(device_id)
4950 .await
4951 .expect("load")
4952 .expect("device should exist after create");
4953 device.server_cert_chain = Some(chain.clone());
4954 _writer
4955 .save_device_data_for_device(device_id, &device)
4956 .await
4957 .expect("save with cert chain");
4958
4959 let store = SqliteStore::new_for_device(&db_name, device_id)
4964 .await
4965 .expect("reopen store");
4966 let loaded = store
4967 .load_device_data_for_device(device_id)
4968 .await
4969 .expect("load")
4970 .expect("device should exist after reopen");
4971 assert_eq!(
4972 loaded.server_cert_chain.as_ref(),
4973 Some(&chain),
4974 "server_cert_chain must survive a save/load roundtrip"
4975 );
4976
4977 let mut device = loaded;
4980 device.server_cert_chain = None;
4981 store
4982 .save_device_data_for_device(device_id, &device)
4983 .await
4984 .expect("save with cleared cert chain");
4985
4986 let reloaded = store
4987 .load_device_data_for_device(device_id)
4988 .await
4989 .expect("reload")
4990 .expect("device should exist");
4991 assert!(
4992 reloaded.server_cert_chain.is_none(),
4993 "cleared chain must round-trip as None"
4994 );
4995 }
4996
4997 #[tokio::test]
5002 async fn legacy_bincode_blobs_self_heal_then_overwrite() {
5003 use diesel::{ExpressionMethods, RunQueryDsl, sql_query};
5004 use wacore::appstate::hash::HashState;
5005 use wacore::store::traits::AppStateSyncKey;
5006
5007 let legacy_sync_key = {
5012 let mut v = vec![0x20u8]; v.extend([0x11u8; 32]);
5014 v.extend([0x04, 0xaa, 0xbb, 0xcc, 0xdd, 0xfc, 0x00, 0xe2, 0xa7, 0xca]);
5015 v
5016 };
5017 let legacy_hash_state = {
5019 let mut v = vec![0x07u8]; v.push(0xde);
5021 v.push(0xad);
5022 v.extend([0u8; 125]);
5023 v.push(0xbe);
5024 v.push(0x00); v
5026 };
5027
5028 let store = create_test_store().await;
5029 let device_id = store.device_id;
5030
5031 let key_id = b"legacy-key".to_vec();
5034 {
5035 let kid = key_id.clone();
5036 let blob = legacy_sync_key.clone();
5037 store
5038 .with_retry("insert_legacy_key", move || {
5039 let kid = kid.clone();
5040 let blob = blob.clone();
5041 Box::new(move |conn| {
5042 diesel::insert_into(app_state_keys::table)
5043 .values((
5044 app_state_keys::key_id.eq(kid),
5045 app_state_keys::key_data.eq(blob),
5046 app_state_keys::device_id.eq(device_id),
5047 ))
5048 .execute(conn)
5049 .map(|_| ())
5050 })
5051 })
5052 .await
5053 .expect("insert legacy key row");
5054 }
5055 let name = "critical_block";
5056 {
5057 let blob = legacy_hash_state.clone();
5058 store
5059 .with_retry("insert_legacy_version", move || {
5060 let blob = blob.clone();
5061 Box::new(move |conn| {
5062 diesel::insert_into(app_state_versions::table)
5063 .values((
5064 app_state_versions::name.eq(name),
5065 app_state_versions::state_data.eq(blob),
5066 app_state_versions::device_id.eq(device_id),
5067 ))
5068 .execute(conn)
5069 .map(|_| ())
5070 })
5071 })
5072 .await
5073 .expect("insert legacy version row");
5074 }
5075
5076 assert!(
5079 store
5080 .get_app_state_sync_key_for_device(&key_id, device_id)
5081 .await
5082 .expect("legacy sync-key blob must not surface a decode error")
5083 .is_none(),
5084 "a legacy bincode sync-key row must read back as absent"
5085 );
5086 assert_eq!(
5087 store
5088 .get_app_state_version_for_device(name, device_id)
5089 .await
5090 .expect("legacy version blob must not surface a decode error")
5091 .version,
5092 0,
5093 "a legacy bincode version row must reset to default (re-sync from 0)"
5094 );
5095
5096 store
5099 .set_app_state_sync_key_for_device(
5100 &key_id,
5101 AppStateSyncKey {
5102 key_data: vec![7u8; 32],
5103 fingerprint: vec![1, 2, 3],
5104 timestamp: 99,
5105 },
5106 device_id,
5107 )
5108 .await
5109 .expect("overwrite key");
5110 let healed_key = store
5111 .get_app_state_sync_key_for_device(&key_id, device_id)
5112 .await
5113 .expect("get key")
5114 .expect("re-shared key must persist over the legacy row");
5115 assert_eq!(healed_key.key_data, vec![7u8; 32]);
5116 assert_eq!(healed_key.timestamp, 99);
5117
5118 store
5119 .set_app_state_version_for_device(
5120 name,
5121 HashState {
5122 version: 5,
5123 ..HashState::default()
5124 },
5125 device_id,
5126 )
5127 .await
5128 .expect("overwrite version");
5129 assert_eq!(
5130 store
5131 .get_app_state_version_for_device(name, device_id)
5132 .await
5133 .expect("get version")
5134 .version,
5135 5,
5136 "a re-synced version must persist over the legacy row"
5137 );
5138
5139 store
5141 .with_retry("corrupt_key", || {
5142 Box::new(|conn| {
5143 sql_query("UPDATE app_state_keys SET key_data = X'00ff00ff'")
5144 .execute(conn)
5145 .map(|_| ())
5146 })
5147 })
5148 .await
5149 .expect("corrupt key blob");
5150 assert!(
5151 store
5152 .get_app_state_sync_key_for_device(&key_id, device_id)
5153 .await
5154 .expect("corrupt key blob must not error")
5155 .is_none(),
5156 "an arbitrarily corrupt sync-key blob must also read back as absent"
5157 );
5158 }
5159
5160 #[tokio::test]
5164 async fn latest_sync_key_skips_undecodable_rows() {
5165 use diesel::{ExpressionMethods, RunQueryDsl};
5166 use wacore::store::traits::AppStateSyncKey;
5167
5168 let legacy_blob = {
5170 let mut v = vec![0x20u8];
5171 v.extend([0x11u8; 32]);
5172 v.extend([0x04, 0xaa, 0xbb, 0xcc, 0xdd, 0xfc, 0x00, 0xe2, 0xa7, 0xca]);
5173 v
5174 };
5175
5176 let store = create_test_store().await;
5177 let device_id = store.device_id;
5178
5179 let good_id = b"key-aaa".to_vec();
5181 store
5182 .set_app_state_sync_key_for_device(
5183 &good_id,
5184 AppStateSyncKey {
5185 key_data: vec![7u8; 32],
5186 fingerprint: vec![1],
5187 timestamp: 1,
5188 },
5189 device_id,
5190 )
5191 .await
5192 .unwrap();
5193
5194 let bad_id = b"key-zzz".to_vec();
5196 {
5197 let bid = bad_id.clone();
5198 let blob = legacy_blob.clone();
5199 store
5200 .with_retry("insert_stale_key", move || {
5201 let bid = bid.clone();
5202 let blob = blob.clone();
5203 Box::new(move |conn| {
5204 diesel::insert_into(app_state_keys::table)
5205 .values((
5206 app_state_keys::key_id.eq(bid),
5207 app_state_keys::key_data.eq(blob),
5208 app_state_keys::device_id.eq(device_id),
5209 ))
5210 .execute(conn)
5211 .map(|_| ())
5212 })
5213 })
5214 .await
5215 .unwrap();
5216 }
5217
5218 assert_eq!(
5220 store
5221 .get_latest_app_state_sync_key_id_for_device(device_id)
5222 .await
5223 .unwrap(),
5224 Some(good_id),
5225 "latest-key selection must skip undecodable bincode rows"
5226 );
5227 }
5228
5229 #[tokio::test]
5230 async fn group_metadata_round_trip_sqlite() {
5231 use wacore::store::traits::ProtocolStore;
5232 let store = create_test_store().await;
5233 let jid = "120363000000000001@g.us";
5234
5235 assert!(store.get_group_metadata(jid).await.unwrap().is_none());
5236
5237 store.put_group_metadata(jid, b"blob-v1").await.unwrap();
5238 assert_eq!(
5239 store.get_group_metadata(jid).await.unwrap().as_deref(),
5240 Some(&b"blob-v1"[..])
5241 );
5242
5243 store.put_group_metadata(jid, b"blob-v2").await.unwrap();
5245 assert_eq!(
5246 store.get_group_metadata(jid).await.unwrap().as_deref(),
5247 Some(&b"blob-v2"[..])
5248 );
5249
5250 store.delete_group_metadata(jid).await.unwrap();
5252 assert!(store.get_group_metadata(jid).await.unwrap().is_none());
5253 }
5254
5255 #[tokio::test]
5256 async fn msg_secret_round_trip_sqlite() {
5257 let store = create_test_store().await;
5258 let secret = [0xABu8; 32];
5259 store
5260 .put_msg_secret("12345@s.whatsapp.net", "9999@lid", "MID1", &secret)
5261 .await
5262 .expect("put");
5263 let got = store
5264 .get_msg_secret("12345@s.whatsapp.net", "9999@lid", "MID1")
5265 .await
5266 .expect("get")
5267 .expect("must exist");
5268 assert_eq!(got, secret.to_vec());
5269 }
5270
5271 #[tokio::test]
5272 async fn msg_secret_miss_returns_none_sqlite() {
5273 let store = create_test_store().await;
5274 assert!(
5275 store
5276 .get_msg_secret("any@s.whatsapp.net", "any@lid", "NOPE")
5277 .await
5278 .expect("get")
5279 .is_none()
5280 );
5281 }
5282
5283 #[tokio::test]
5284 async fn msg_secret_upsert_replaces_secret() {
5285 let store = create_test_store().await;
5286 store
5287 .put_msg_secret("c", "s", "M", &[1u8; 32])
5288 .await
5289 .expect("put 1");
5290 store
5291 .put_msg_secret("c", "s", "M", &[9u8; 32])
5292 .await
5293 .expect("put 2");
5294 let got = store.get_msg_secret("c", "s", "M").await.unwrap().unwrap();
5295 assert_eq!(got, vec![9u8; 32], "ON CONFLICT must overwrite");
5296 }
5297
5298 #[tokio::test]
5299 async fn msg_secret_scoped_by_three_columns() {
5300 let store = create_test_store().await;
5301 store
5302 .put_msg_secret("c1", "s1", "M1", &[1u8; 32])
5303 .await
5304 .unwrap();
5305 store
5306 .put_msg_secret("c1", "s1", "M2", &[2u8; 32])
5307 .await
5308 .unwrap();
5309 store
5310 .put_msg_secret("c1", "s2", "M1", &[3u8; 32])
5311 .await
5312 .unwrap();
5313 store
5314 .put_msg_secret("c2", "s1", "M1", &[4u8; 32])
5315 .await
5316 .unwrap();
5317
5318 for (chat, sender, msg_id, expected) in [
5319 ("c1", "s1", "M1", 1u8),
5320 ("c1", "s1", "M2", 2),
5321 ("c1", "s2", "M1", 3),
5322 ("c2", "s1", "M1", 4),
5323 ] {
5324 let got = store
5325 .get_msg_secret(chat, sender, msg_id)
5326 .await
5327 .unwrap()
5328 .unwrap_or_else(|| panic!("missing ({chat},{sender},{msg_id})"));
5329 assert_eq!(got, vec![expected; 32]);
5330 }
5331 }
5332
5333 #[tokio::test]
5334 async fn msg_secret_batch_upserts_in_one_call() {
5335 const ORIGINAL_SECRET_BYTE: u8 = 0x5a;
5336 const UPDATED_SECRET_BYTE: u8 = 0xa5;
5337
5338 let store = create_test_store().await;
5339 let mut entries: Vec<_> = (0..=MSG_SECRET_INSERT_CHUNK_SIZE)
5340 .map(|index| MsgSecretEntry {
5341 chat: "c".into(),
5342 sender: "s".into(),
5343 msg_id: format!("M{index}").into(),
5344 secret: [ORIGINAL_SECRET_BYTE; wacore::reporting_token::MESSAGE_SECRET_SIZE],
5345 expires_at: 0,
5346 message_ts: 0,
5347 })
5348 .collect();
5349 entries.push(MsgSecretEntry {
5352 chat: "c".into(),
5353 sender: "s".into(),
5354 msg_id: "M0".into(),
5355 secret: [UPDATED_SECRET_BYTE; wacore::reporting_token::MESSAGE_SECRET_SIZE],
5356 expires_at: 0,
5357 message_ts: 0,
5358 });
5359 let expected_stored = entries.len();
5360 let stored = store.put_msg_secrets(entries).await.unwrap();
5361
5362 assert_eq!(stored, expected_stored);
5363 assert_eq!(
5364 store.get_msg_secret("c", "s", "M0").await.unwrap().unwrap(),
5365 vec![UPDATED_SECRET_BYTE; wacore::reporting_token::MESSAGE_SECRET_SIZE]
5366 );
5367 assert_eq!(
5368 store
5369 .get_msg_secret("c", "s", &format!("M{MSG_SECRET_INSERT_CHUNK_SIZE}"))
5370 .await
5371 .unwrap()
5372 .unwrap(),
5373 vec![ORIGINAL_SECRET_BYTE; wacore::reporting_token::MESSAGE_SECRET_SIZE]
5374 );
5375 }
5376
5377 #[tokio::test]
5378 async fn delete_expired_msg_secrets_deletes_only_passed_deadlines() {
5379 let store = create_test_store().await;
5380 let now = wacore::time::now_secs();
5381 store
5382 .put_msg_secrets(vec![
5383 MsgSecretEntry {
5384 chat: "c".into(),
5385 sender: "s".into(),
5386 msg_id: "NEVER".into(),
5387 secret: [1u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
5388 expires_at: 0,
5389 message_ts: 0,
5390 },
5391 MsgSecretEntry {
5392 chat: "c".into(),
5393 sender: "s".into(),
5394 msg_id: "FUTURE".into(),
5395 secret: [2u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
5396 expires_at: now + 86_400,
5397 message_ts: 0,
5398 },
5399 MsgSecretEntry {
5400 chat: "c".into(),
5401 sender: "s".into(),
5402 msg_id: "PAST".into(),
5403 secret: [3u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
5404 expires_at: now - 86_400,
5405 message_ts: 0,
5406 },
5407 ])
5408 .await
5409 .unwrap();
5410
5411 let removed = store.delete_expired_msg_secrets(now).await.unwrap();
5412 assert_eq!(
5413 removed, 1,
5414 "only the row whose deadline has passed is deleted"
5415 );
5416 assert!(
5417 store
5418 .get_msg_secret("c", "s", "NEVER")
5419 .await
5420 .unwrap()
5421 .is_some(),
5422 "expires_at = 0 never expires"
5423 );
5424 assert!(
5425 store
5426 .get_msg_secret("c", "s", "FUTURE")
5427 .await
5428 .unwrap()
5429 .is_some(),
5430 "a future deadline survives"
5431 );
5432 assert!(
5433 store
5434 .get_msg_secret("c", "s", "PAST")
5435 .await
5436 .unwrap()
5437 .is_none(),
5438 "a passed deadline is pruned"
5439 );
5440 }
5441
5442 #[tokio::test]
5443 async fn put_msg_secrets_keeps_later_deadline_on_conflict() {
5444 let store = create_test_store().await;
5445 let now = wacore::time::now_secs();
5446 store
5449 .put_msg_secrets(vec![MsgSecretEntry {
5450 chat: "c".into(),
5451 sender: "s".into(),
5452 msg_id: "M".into(),
5453 secret: [1u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
5454 expires_at: now + 90 * 86_400,
5455 message_ts: 0,
5456 }])
5457 .await
5458 .unwrap();
5459 store
5460 .put_msg_secrets(vec![MsgSecretEntry {
5461 chat: "c".into(),
5462 sender: "s".into(),
5463 msg_id: "M".into(),
5464 secret: [1u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
5465 expires_at: now + 30 * 86_400,
5466 message_ts: 0,
5467 }])
5468 .await
5469 .unwrap();
5470 let removed = store
5472 .delete_expired_msg_secrets(now + 60 * 86_400)
5473 .await
5474 .unwrap();
5475 assert_eq!(removed, 0, "conflict must keep the later (90d) deadline");
5476
5477 store
5479 .put_msg_secret("c", "s", "M", &[1u8; 32])
5480 .await
5481 .unwrap();
5482 let removed = store
5483 .delete_expired_msg_secrets(now + 200 * 86_400)
5484 .await
5485 .unwrap();
5486 assert_eq!(removed, 0, "a 0 (never) deadline wins over any finite one");
5487 }
5488
5489 #[tokio::test]
5490 async fn get_msg_secret_with_ts_round_trips_and_keeps_parent_ts() {
5491 let store = create_test_store().await;
5492 let parent_ts = 1_700_000_000i64;
5493 store
5494 .put_msg_secrets(vec![MsgSecretEntry {
5495 chat: "c".into(),
5496 sender: "s".into(),
5497 msg_id: "M".into(),
5498 secret: [5u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
5499 expires_at: 0,
5500 message_ts: parent_ts,
5501 }])
5502 .await
5503 .unwrap();
5504 assert_eq!(
5505 store.get_msg_secret_with_ts("c", "s", "M").await.unwrap(),
5506 Some((vec![5u8; 32], parent_ts))
5507 );
5508
5509 store
5511 .put_msg_secret("c", "s", "M", &[5u8; 32])
5512 .await
5513 .unwrap();
5514 assert_eq!(
5515 store.get_msg_secret_with_ts("c", "s", "M").await.unwrap(),
5516 Some((vec![5u8; 32], parent_ts)),
5517 "message_ts (immutable parent time) must survive a 0-ts redelivery"
5518 );
5519
5520 assert_eq!(
5522 store
5523 .get_msg_secret_with_ts("c", "s", "MISSING")
5524 .await
5525 .unwrap(),
5526 None
5527 );
5528 }
5529
5530 #[tokio::test]
5533 async fn msg_secret_isolated_per_device_id() {
5534 use portable_atomic::AtomicU64;
5535 use std::sync::atomic::Ordering;
5536 static COUNTER: AtomicU64 = AtomicU64::new(0);
5537 let id = COUNTER.fetch_add(1, Ordering::Relaxed);
5538 let shared_url = format!(
5539 "file:memdb_msgsecret_iso_{}_{}?mode=memory&cache=shared",
5540 std::process::id(),
5541 id
5542 );
5543 let store_a = SqliteStore::new_for_device(&shared_url, 1)
5544 .await
5545 .expect("store_a");
5546 let store_b = SqliteStore::new_for_device(&shared_url, 2)
5547 .await
5548 .expect("store_b");
5549
5550 store_a
5551 .put_msg_secret("c", "s", "M", &[7u8; 32])
5552 .await
5553 .unwrap();
5554 assert!(
5555 store_b
5556 .get_msg_secret("c", "s", "M")
5557 .await
5558 .unwrap()
5559 .is_none(),
5560 "same DB, different device_id must not see each other's secrets"
5561 );
5562 assert_eq!(
5563 store_a
5564 .get_msg_secret("c", "s", "M")
5565 .await
5566 .unwrap()
5567 .unwrap(),
5568 vec![7u8; 32],
5569 "device_a still sees its own write"
5570 );
5571 }
5572
5573 #[tokio::test]
5576 async fn resource_report_bounds_cache_by_db_size_and_cap() {
5577 let store = create_test_store().await; let device_id = 1;
5579
5580 let macs: Vec<AppStateMutationMAC> = (0..500u32)
5582 .map(|i| {
5583 let mut index_mac = vec![0u8; 32];
5584 index_mac[..4].copy_from_slice(&i.to_le_bytes());
5585 AppStateMutationMAC {
5586 index_mac,
5587 value_mac: vec![(i % 251) as u8; 32],
5588 }
5589 })
5590 .collect();
5591 store
5592 .put_app_state_mutation_macs_for_device("coll", 1, &macs, device_id)
5593 .await
5594 .unwrap();
5595
5596 let report = store.resource_report().await;
5597
5598 let pages = report.pages.expect("SQLite reports a page count");
5599 assert!(pages > 0, "a migrated + seeded DB has pages");
5600
5601 let mem = report
5602 .memory_bytes
5603 .expect("SQLite reports a cache estimate");
5604 assert!(mem > 0, "cache-in-use estimate is non-zero for a seeded DB");
5605 assert!(
5608 mem <= 512 * 1024,
5609 "estimate never exceeds the configured 512 KiB cap, got {mem}"
5610 );
5611 assert_eq!(report.total_bytes(), mem, "total_bytes == memory_bytes");
5612 assert_eq!(report.io_read_bytes, None);
5614 assert_eq!(report.io_write_bytes, None);
5615 }
5616
5617 #[test]
5621 fn mmap_size_config_is_opt_in() {
5622 assert_eq!(
5623 SqliteStoreConfig::default().mmap_size,
5624 None,
5625 "default leaves mmap off (current behavior)"
5626 );
5627 assert_eq!(
5628 SqliteStoreConfig::default()
5629 .with_mmap_size(64 * 1024 * 1024)
5630 .mmap_size,
5631 Some(64 * 1024 * 1024),
5632 "builder sets the field"
5633 );
5634 }
5635
5636 #[tokio::test]
5637 async fn mmap_size_applies_pragma_and_store_operates() {
5638 use portable_atomic::AtomicU64;
5639 use std::sync::atomic::Ordering;
5640 static COUNTER: AtomicU64 = AtomicU64::new(0);
5641 let id = COUNTER.fetch_add(1, Ordering::Relaxed);
5642 let path =
5644 std::env::temp_dir().join(format!("wa_mmap_test_{}_{}.db", std::process::id(), id));
5645 let url = path.to_str().unwrap().to_string();
5646
5647 let read_mmap = |store: &SqliteStore| -> i64 {
5651 #[derive(diesel::QueryableByName)]
5652 struct M {
5653 #[diesel(sql_type = diesel::sql_types::BigInt)]
5654 mmap_size: i64,
5655 }
5656 let mut conn = store.pool.get().unwrap();
5657 diesel::sql_query("PRAGMA mmap_size")
5658 .get_result::<M>(&mut conn)
5659 .map(|m| m.mmap_size)
5660 .unwrap_or(-1)
5661 };
5662
5663 let def_store = SqliteStore::new(&url).await.expect("default store");
5666 assert_eq!(read_mmap(&def_store), 0, "default keeps mmap off");
5667 drop(def_store);
5668
5669 const MMAP: u64 = 64 * 1024 * 1024;
5672 let cfg = SqliteStoreConfig::default().with_mmap_size(MMAP);
5673 let store = SqliteStore::with_config(&url, cfg)
5674 .await
5675 .expect("mmap store builds");
5676 store
5677 .put_identity("559980000001@s.whatsapp.net", [9u8; 32])
5678 .await
5679 .expect("store operates with mmap set");
5680 let applied = read_mmap(&store);
5683 assert!(
5684 applied == MMAP as i64 || applied == 0,
5685 "mmap_size is applied when the VFS supports it, got {applied}"
5686 );
5687 drop(store);
5688
5689 for suffix in ["", "-wal", "-shm"] {
5691 let _ = std::fs::remove_file(format!("{url}{suffix}"));
5692 }
5693 }
5694}