use crate::schema::*;
use async_trait::async_trait;
use bytes::Bytes;
use diesel::prelude::*;
use diesel::r2d2::{ConnectionManager, Pool};
use diesel::result::{DatabaseErrorKind, Error as DieselError};
use diesel::sqlite::SqliteConnection;
use diesel::upsert::excluded;
use diesel_migrations::{EmbeddedMigrations, MigrationHarness, embed_migrations};
use log::warn;
use std::sync::Arc;
use std::time::Duration;
use wacore::appstate::hash::HashState;
use wacore::appstate::processor::AppStateMutationMAC;
use wacore::libsignal::protocol::{KeyPair, PrivateKey, PublicKey};
use wacore::store::Device as CoreDevice;
use wacore::store::error::{Result, StoreError};
use wacore::store::traits::*;
enum DieselOrStore {
Diesel(DieselError),
Store(StoreError),
}
impl From<DieselOrStore> for StoreError {
fn from(e: DieselOrStore) -> Self {
match e {
DieselOrStore::Diesel(e) => StoreError::Database(Box::new(e)),
DieselOrStore::Store(e) => e,
}
}
}
fn is_retriable_sqlite_error(error: &DieselError) -> bool {
match error {
DieselError::DatabaseError(DatabaseErrorKind::Unknown, info) => {
let msg = info.message();
msg.contains("locked") || msg.contains("busy")
}
_ => false,
}
}
const MIGRATIONS: EmbeddedMigrations = embed_migrations!("migrations");
pub(crate) type SqlitePool = Pool<ConnectionManager<SqliteConnection>>;
#[derive(Queryable, Selectable)]
#[diesel(table_name = device)]
#[allow(dead_code)]
struct DeviceRow {
id: i32,
lid: String,
pn: String,
registration_id: i32,
noise_key: Vec<u8>,
identity_key: Vec<u8>,
signed_pre_key: Vec<u8>,
signed_pre_key_id: i32,
signed_pre_key_signature: Vec<u8>,
adv_secret_key: Vec<u8>,
account: Option<Vec<u8>>,
push_name: String,
app_version_primary: i32,
app_version_secondary: i32,
app_version_tertiary: i64,
app_version_last_fetched_ms: i64,
edge_routing_info: Option<Vec<u8>>,
props_hash: Option<String>,
next_pre_key_id: i32,
nct_salt: Option<Vec<u8>>,
server_has_prekeys: bool,
server_cert_chain: Option<Vec<u8>>,
login_counter: i32,
first_unupload_pre_key_id: i32,
lid_migrated: bool,
last_signed_pre_key_rotation_ms: i64,
read_receipts_disabled: bool,
}
const ID_PARAM_CHUNK: usize = 900;
const MSG_SECRET_INSERT_CHUNK_SIZE: usize = 100;
#[derive(Clone)]
pub(crate) struct ReadPool {
pub(crate) pool: SqlitePool,
pub(crate) semaphore: Arc<tokio::sync::Semaphore>,
}
#[derive(Clone)]
pub struct SqliteStore {
pub(crate) pool: SqlitePool,
pub(crate) db_semaphore: Arc<tokio::sync::Semaphore>,
pub(crate) reads: Option<ReadPool>,
pub(crate) database_path: String,
device_id: i32,
}
#[derive(Debug, Clone, Copy)]
pub enum Synchronous {
Off,
Normal,
Full,
}
impl Synchronous {
fn as_pragma(self) -> &'static str {
match self {
Synchronous::Off => "OFF",
Synchronous::Normal => "NORMAL",
Synchronous::Full => "FULL",
}
}
}
pub type ConnectionInitHook = Arc<
dyn Fn(
&mut SqliteConnection,
) -> std::result::Result<(), Box<dyn std::error::Error + Send + Sync>>
+ Send
+ Sync,
>;
#[derive(Clone)]
pub struct SqliteStoreConfig {
pub pool_size: u32,
pub read_pool_size: u32,
pub cache_size_kib: u32,
pub mmap_size: Option<u64>,
pub busy_timeout: Duration,
pub synchronous: Synchronous,
pub thread_pool: Option<Arc<scheduled_thread_pool::ScheduledThreadPool>>,
pub connection_init: Option<ConnectionInitHook>,
}
impl Default for SqliteStoreConfig {
fn default() -> Self {
Self {
pool_size: 1,
read_pool_size: 0,
cache_size_kib: 512,
mmap_size: None,
busy_timeout: Duration::from_secs(30),
synchronous: Synchronous::Normal,
thread_pool: None,
connection_init: None,
}
}
}
impl SqliteStoreConfig {
pub fn with_read_pool_size(mut self, n: u32) -> Self {
self.read_pool_size = n;
self
}
pub fn with_mmap_size(mut self, bytes: u64) -> Self {
self.mmap_size = Some(bytes);
self
}
pub fn with_connection_init<F>(mut self, hook: F) -> Self
where
F: Fn(
&mut SqliteConnection,
) -> std::result::Result<(), Box<dyn std::error::Error + Send + Sync>>
+ Send
+ Sync
+ 'static,
{
self.connection_init = Some(Arc::new(hook));
self
}
}
#[derive(Clone)]
struct ConnectionOptions {
cache_size_kib: u32,
mmap_size: Option<u64>,
busy_timeout_ms: u64,
synchronous: Synchronous,
connection_init: Option<ConnectionInitHook>,
query_only: bool,
}
impl std::fmt::Debug for ConnectionOptions {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ConnectionOptions")
.field("cache_size_kib", &self.cache_size_kib)
.field("mmap_size", &self.mmap_size)
.field("busy_timeout_ms", &self.busy_timeout_ms)
.field("synchronous", &self.synchronous)
.field(
"connection_init",
&self.connection_init.as_ref().map(|_| ()),
)
.finish()
}
}
impl diesel::r2d2::CustomizeConnection<SqliteConnection, diesel::r2d2::Error>
for ConnectionOptions
{
fn on_acquire(
&self,
conn: &mut SqliteConnection,
) -> std::result::Result<(), diesel::r2d2::Error> {
if let Some(init) = &self.connection_init {
init(conn).map_err(|e| {
diesel::r2d2::Error::QueryError(diesel::result::Error::QueryBuilderError(e))
})?;
}
let mut pragmas = vec![
format!("PRAGMA busy_timeout = {};", self.busy_timeout_ms),
format!("PRAGMA synchronous = {};", self.synchronous.as_pragma()),
format!("PRAGMA cache_size = -{};", self.cache_size_kib),
"PRAGMA temp_store = memory;".to_string(),
"PRAGMA foreign_keys = ON;".to_string(),
];
if let Some(mmap_size) = self.mmap_size.filter(|&n| n > 0) {
pragmas.push(format!("PRAGMA mmap_size = {mmap_size};"));
}
if self.query_only {
pragmas.push("PRAGMA query_only = 1;".to_string());
}
for pragma in pragmas {
diesel::sql_query(pragma)
.execute(conn)
.map_err(diesel::r2d2::Error::QueryError)?;
}
Ok(())
}
}
fn parse_database_path(database_url: &str) -> Result<String> {
if database_url == ":memory:" {
return Err(StoreError::InvalidConfig(
"Snapshot not supported for in-memory databases".to_string(),
));
}
let path = database_url
.split(['?', '#'])
.next()
.unwrap_or(database_url);
let path = path.trim_start_matches("sqlite://");
if path == ":memory:" || path.starts_with(":memory:?") {
return Err(StoreError::InvalidConfig(
"Snapshot not supported for in-memory databases".to_string(),
));
}
Ok(path.to_string())
}
fn is_shared_cache(database_url: &str) -> bool {
let Some((_, query)) = database_url.split_once('?') else {
return false;
};
if !database_url.starts_with("file:") {
return false;
}
query
.split('#')
.next()
.unwrap_or(query)
.split('&')
.filter_map(|param| param.split_once('='))
.find(|(key, _)| *key == "cache")
.is_some_and(|(_, value)| value.eq_ignore_ascii_case("shared"))
}
fn shared_r2d2_thread_pool() -> Arc<scheduled_thread_pool::ScheduledThreadPool> {
static POOL: std::sync::OnceLock<Arc<scheduled_thread_pool::ScheduledThreadPool>> =
std::sync::OnceLock::new();
POOL.get_or_init(|| {
Arc::new(
scheduled_thread_pool::ScheduledThreadPool::builder()
.num_threads(2)
.thread_name_pattern("r2d2-shared-{}")
.build(),
)
})
.clone()
}
impl SqliteStore {
pub async fn new(database_url: &str) -> std::result::Result<Self, StoreError> {
Self::build(database_url, 1, SqliteStoreConfig::default()).await
}
pub async fn with_config(
database_url: &str,
config: SqliteStoreConfig,
) -> std::result::Result<Self, StoreError> {
Self::build(database_url, 1, config).await
}
pub async fn new_for_device(
database_url: &str,
device_id: i32,
) -> std::result::Result<Self, StoreError> {
Self::build(database_url, device_id, SqliteStoreConfig::default()).await
}
pub async fn with_config_for_device(
database_url: &str,
device_id: i32,
config: SqliteStoreConfig,
) -> std::result::Result<Self, StoreError> {
Self::build(database_url, device_id, config).await
}
async fn build(
database_url: &str,
device_id: i32,
config: SqliteStoreConfig,
) -> std::result::Result<Self, StoreError> {
let manager = ConnectionManager::<SqliteConnection>::new(database_url);
let pool_size = config.pool_size.max(1);
let read_pool_size = config.read_pool_size;
let thread_pool = config.thread_pool.unwrap_or_else(shared_r2d2_thread_pool);
let read_thread_pool = Arc::clone(&thread_pool);
let options = ConnectionOptions {
cache_size_kib: config.cache_size_kib,
mmap_size: config.mmap_size,
busy_timeout_ms: if config.busy_timeout.is_zero() {
0
} else {
config.busy_timeout.as_millis().clamp(1, i32::MAX as u128) as u64
},
synchronous: config.synchronous,
connection_init: config.connection_init,
query_only: false,
};
let read_options = ConnectionOptions {
query_only: true,
..options.clone()
};
let db_url = database_url.to_string();
let (pool, journal_mode) = tokio::task::spawn_blocking(
move || -> std::result::Result<(SqlitePool, String), StoreError> {
let pool = Pool::builder()
.max_size(pool_size)
.test_on_check_out(false)
.thread_pool(thread_pool)
.connection_customizer(Box::new(options))
.build(manager)
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
#[derive(diesel::QueryableByName)]
struct JournalMode {
#[diesel(sql_type = diesel::sql_types::Text)]
journal_mode: String,
}
let journal_mode = diesel::sql_query("PRAGMA journal_mode = WAL;")
.get_result::<JournalMode>(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?
.journal_mode;
conn.run_pending_migrations(MIGRATIONS)
.map_err(StoreError::Migration)?;
Ok((pool, journal_mode))
},
)
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
let wal = journal_mode.eq_ignore_ascii_case("wal");
let shared_cache = is_shared_cache(&db_url);
let declined = if !wal {
Some(format!("journal_mode is '{journal_mode}', not WAL"))
} else if shared_cache {
Some("the URI opts into shared cache, whose table locks block the writer".to_string())
} else {
None
};
if read_pool_size > 0
&& let Some(reason) = &declined
{
log::warn!("sqlite-storage: read_pool_size={read_pool_size} ignored, {reason}");
}
let reads = if read_pool_size > 0 && declined.is_none() {
let manager = ConnectionManager::<SqliteConnection>::new(&db_url);
let pool = tokio::task::spawn_blocking(
move || -> std::result::Result<SqlitePool, StoreError> {
Pool::builder()
.max_size(read_pool_size)
.test_on_check_out(false)
.thread_pool(read_thread_pool)
.connection_customizer(Box::new(read_options))
.build(manager)
.map_err(|e| StoreError::Connection(Box::new(e)))
},
)
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Some(ReadPool {
pool,
semaphore: Arc::new(tokio::sync::Semaphore::new(read_pool_size as usize)),
})
} else {
None
};
let database_path = parse_database_path(database_url)?;
Ok(Self {
pool,
db_semaphore: Arc::new(tokio::sync::Semaphore::new(pool_size as usize)),
reads,
database_path,
device_id,
})
}
pub fn device_id(&self) -> i32 {
self.device_id
}
async fn with_semaphore<F, T>(&self, f: F) -> Result<T>
where
F: FnOnce() -> Result<T> + Send + 'static,
T: Send + 'static,
{
let permit = self
.db_semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| StoreError::Database(Box::new(e)))?;
let result = tokio::task::spawn_blocking(move || {
let res = f();
drop(permit);
res
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(result)
}
async fn with_retry<F, T>(&self, op_name: &str, make_op: F) -> Result<T>
where
F: Fn() -> Box<
dyn FnOnce(&mut SqliteConnection) -> std::result::Result<T, DieselError> + Send,
>,
T: Send + 'static,
{
const MAX_RETRIES: u32 = 5;
for attempt in 0..=MAX_RETRIES {
let permit = self
.db_semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| StoreError::Database(Box::new(e)))?;
let pool = self.pool.clone();
let op = make_op();
let result =
tokio::task::spawn_blocking(move || -> std::result::Result<T, DieselOrStore> {
let _permit = permit;
let mut conn = pool
.get()
.map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
op(&mut conn).map_err(DieselOrStore::Diesel)
})
.await;
match result {
Ok(Ok(val)) => return Ok(val),
Ok(Err(DieselOrStore::Diesel(ref e)))
if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
{
let delay_ms = 10u64 * (1u64 << attempt.min(4));
if attempt >= 1 {
warn!(
"{op_name} busy/locked, retry {}/{} in {delay_ms}ms: {e}",
attempt + 1,
MAX_RETRIES + 1
);
}
tokio::time::sleep(Duration::from_millis(delay_ms)).await;
}
Ok(Err(e)) => return Err(e.into()),
Err(e) => return Err(StoreError::Database(Box::new(e))),
}
}
Err(StoreError::RetriesExhausted {
op: op_name.to_string(),
})
}
fn serialize_keypair(&self, key_pair: &KeyPair) -> Result<Vec<u8>> {
let mut bytes = Vec::with_capacity(64);
bytes.extend_from_slice(key_pair.private_key.serialize());
bytes.extend_from_slice(key_pair.public_key.public_key_bytes());
Ok(bytes)
}
fn deserialize_keypair(&self, bytes: &[u8]) -> Result<KeyPair> {
if bytes.len() != 64 {
return Err(StoreError::Validation(format!(
"Invalid KeyPair length: {}",
bytes.len()
)));
}
let private_key = PrivateKey::deserialize(&bytes[0..32])
.map_err(|e| StoreError::Serialization(Box::new(e)))?;
let public_key = PublicKey::from_djb_public_key_bytes(&bytes[32..64])
.map_err(|e| StoreError::Serialization(Box::new(e)))?;
Ok(KeyPair::new(public_key, private_key))
}
pub async fn save_device_data_for_device(
&self,
device_id: i32,
device_data: &CoreDevice,
) -> Result<()> {
let noise_key_data: Arc<[u8]> = self.serialize_keypair(&device_data.noise_key)?.into();
let identity_key_data: Arc<[u8]> =
self.serialize_keypair(&device_data.identity_key)?.into();
let signed_pre_key_data: Arc<[u8]> =
self.serialize_keypair(&device_data.signed_pre_key)?.into();
let account_data: Option<Arc<[u8]>> = device_data
.account
.as_ref()
.map(|a| Arc::from(wacore::store::device::account_serde::to_bytes(a)));
let registration_id = device_data.registration_id as i32;
let signed_pre_key_id = device_data.signed_pre_key_id as i32;
let signed_pre_key_signature: Arc<[u8]> =
Arc::from(&device_data.signed_pre_key_signature[..]);
let adv_secret_key: Arc<[u8]> = Arc::from(&device_data.adv_secret_key[..]);
let push_name: Arc<str> = Arc::from(device_data.push_name.as_str());
let app_version_primary = device_data.app_version_primary as i32;
let app_version_secondary = device_data.app_version_secondary as i32;
let app_version_tertiary = device_data.app_version_tertiary as i64;
let app_version_last_fetched_ms = device_data.app_version_last_fetched_ms;
let edge_routing_info: Option<Arc<[u8]>> =
device_data.edge_routing_info.as_deref().map(Arc::from);
let props_hash: Option<Arc<str>> = device_data.props_hash.as_deref().map(Arc::from);
let next_pre_key_id = device_data.next_pre_key_id as i32;
let first_unupload_pre_key_id = device_data.first_unupload_pre_key_id as i32;
let server_has_prekeys = device_data.server_has_prekeys;
let nct_salt: Option<Arc<[u8]>> = device_data.nct_salt.as_deref().map(Arc::from);
let server_cert_chain: Option<Arc<[u8]>> = device_data
.server_cert_chain
.as_ref()
.map(|chain| Arc::from(crate::wire::encode_server_cert_chain(chain)));
let login_counter = device_data.login_counter;
let lid_migrated = device_data.lid_migrated;
let last_signed_pre_key_rotation_ms = device_data.last_signed_pre_key_rotation_ms;
let read_receipts_disabled = device_data.read_receipts_disabled;
let new_lid: Arc<str> = Arc::from(
device_data
.lid
.as_ref()
.map(|j| j.to_string())
.unwrap_or_default()
.as_str(),
);
let new_pn: Arc<str> = Arc::from(
device_data
.pn
.as_ref()
.map(|j| j.to_string())
.unwrap_or_default()
.as_str(),
);
self.with_retry("save_device_data", || {
let noise_key_data = Arc::clone(&noise_key_data);
let identity_key_data = Arc::clone(&identity_key_data);
let signed_pre_key_data = Arc::clone(&signed_pre_key_data);
let account_data = account_data.clone();
let signed_pre_key_signature = Arc::clone(&signed_pre_key_signature);
let adv_secret_key = Arc::clone(&adv_secret_key);
let push_name = Arc::clone(&push_name);
let edge_routing_info = edge_routing_info.clone();
let props_hash = props_hash.clone();
let nct_salt = nct_salt.clone();
let server_cert_chain = server_cert_chain.clone();
let new_lid = Arc::clone(&new_lid);
let new_pn = Arc::clone(&new_pn);
Box::new(move |conn: &mut SqliteConnection| {
diesel::insert_into(device::table)
.values((
device::id.eq(device_id),
device::lid.eq(&*new_lid),
device::pn.eq(&*new_pn),
device::registration_id.eq(registration_id),
device::noise_key.eq(&*noise_key_data),
device::identity_key.eq(&*identity_key_data),
device::signed_pre_key.eq(&*signed_pre_key_data),
device::signed_pre_key_id.eq(signed_pre_key_id),
device::signed_pre_key_signature.eq(&*signed_pre_key_signature),
device::adv_secret_key.eq(&*adv_secret_key),
device::account.eq(account_data.as_deref()),
device::push_name.eq(&*push_name),
device::app_version_primary.eq(app_version_primary),
device::app_version_secondary.eq(app_version_secondary),
device::app_version_tertiary.eq(app_version_tertiary),
device::app_version_last_fetched_ms.eq(app_version_last_fetched_ms),
device::edge_routing_info.eq(edge_routing_info.as_deref()),
device::props_hash.eq(props_hash.as_deref()),
device::next_pre_key_id.eq(next_pre_key_id),
device::first_unupload_pre_key_id.eq(first_unupload_pre_key_id),
device::server_has_prekeys.eq(server_has_prekeys),
device::nct_salt.eq(nct_salt.as_deref()),
device::server_cert_chain.eq(server_cert_chain.as_deref()),
device::login_counter.eq(login_counter),
device::lid_migrated.eq(lid_migrated),
device::last_signed_pre_key_rotation_ms.eq(last_signed_pre_key_rotation_ms),
device::read_receipts_disabled.eq(read_receipts_disabled),
))
.on_conflict(device::id)
.do_update()
.set((
device::lid.eq(excluded(device::lid)),
device::pn.eq(excluded(device::pn)),
device::registration_id.eq(excluded(device::registration_id)),
device::noise_key.eq(excluded(device::noise_key)),
device::identity_key.eq(excluded(device::identity_key)),
device::signed_pre_key.eq(excluded(device::signed_pre_key)),
device::signed_pre_key_id.eq(excluded(device::signed_pre_key_id)),
device::signed_pre_key_signature
.eq(excluded(device::signed_pre_key_signature)),
device::adv_secret_key.eq(excluded(device::adv_secret_key)),
device::account.eq(excluded(device::account)),
device::push_name.eq(excluded(device::push_name)),
device::app_version_primary.eq(excluded(device::app_version_primary)),
device::app_version_secondary.eq(excluded(device::app_version_secondary)),
device::app_version_tertiary.eq(excluded(device::app_version_tertiary)),
device::app_version_last_fetched_ms
.eq(excluded(device::app_version_last_fetched_ms)),
device::edge_routing_info.eq(excluded(device::edge_routing_info)),
device::props_hash.eq(excluded(device::props_hash)),
device::next_pre_key_id.eq(excluded(device::next_pre_key_id)),
device::first_unupload_pre_key_id
.eq(excluded(device::first_unupload_pre_key_id)),
device::server_has_prekeys.eq(excluded(device::server_has_prekeys)),
device::nct_salt.eq(excluded(device::nct_salt)),
device::server_cert_chain.eq(excluded(device::server_cert_chain)),
device::login_counter.eq(excluded(device::login_counter)),
device::lid_migrated.eq(excluded(device::lid_migrated)),
device::last_signed_pre_key_rotation_ms
.eq(excluded(device::last_signed_pre_key_rotation_ms)),
device::read_receipts_disabled.eq(excluded(device::read_receipts_disabled)),
))
.execute(conn)
.map(|_| ())
})
})
.await
}
pub async fn create_new_device(&self) -> Result<i32> {
let device_id = self.device_id;
let new_device = wacore::store::Device::new();
let noise_key_data: Arc<[u8]> = self.serialize_keypair(&new_device.noise_key)?.into();
let identity_key_data: Arc<[u8]> = self.serialize_keypair(&new_device.identity_key)?.into();
let signed_pre_key_data: Arc<[u8]> =
self.serialize_keypair(&new_device.signed_pre_key)?.into();
let registration_id = new_device.registration_id as i32;
let signed_pre_key_id = new_device.signed_pre_key_id as i32;
let signed_pre_key_signature: Arc<[u8]> =
Arc::from(&new_device.signed_pre_key_signature[..]);
let adv_secret_key: Arc<[u8]> = Arc::from(&new_device.adv_secret_key[..]);
let push_name: Arc<str> = Arc::from(new_device.push_name.as_str());
let app_version_primary = new_device.app_version_primary as i32;
let app_version_secondary = new_device.app_version_secondary as i32;
let app_version_tertiary = new_device.app_version_tertiary as i64;
let app_version_last_fetched_ms = new_device.app_version_last_fetched_ms;
let next_pre_key_id = new_device.next_pre_key_id as i32;
let first_unupload_pre_key_id = new_device.first_unupload_pre_key_id as i32;
let server_has_prekeys = new_device.server_has_prekeys;
let last_signed_pre_key_rotation_ms = new_device.last_signed_pre_key_rotation_ms;
self.with_retry("create_new_device", || {
let noise_key_data = Arc::clone(&noise_key_data);
let identity_key_data = Arc::clone(&identity_key_data);
let signed_pre_key_data = Arc::clone(&signed_pre_key_data);
let signed_pre_key_signature = Arc::clone(&signed_pre_key_signature);
let adv_secret_key = Arc::clone(&adv_secret_key);
let push_name = Arc::clone(&push_name);
Box::new(move |conn: &mut SqliteConnection| {
diesel::insert_into(device::table)
.values((
device::id.eq(device_id),
device::lid.eq(""),
device::pn.eq(""),
device::registration_id.eq(registration_id),
device::noise_key.eq(&*noise_key_data),
device::identity_key.eq(&*identity_key_data),
device::signed_pre_key.eq(&*signed_pre_key_data),
device::signed_pre_key_id.eq(signed_pre_key_id),
device::signed_pre_key_signature.eq(&*signed_pre_key_signature),
device::adv_secret_key.eq(&*adv_secret_key),
device::account.eq(None::<&[u8]>),
device::push_name.eq(&*push_name),
device::app_version_primary.eq(app_version_primary),
device::app_version_secondary.eq(app_version_secondary),
device::app_version_tertiary.eq(app_version_tertiary),
device::app_version_last_fetched_ms.eq(app_version_last_fetched_ms),
device::edge_routing_info.eq(None::<&[u8]>),
device::props_hash.eq(None::<&str>),
device::next_pre_key_id.eq(next_pre_key_id),
device::first_unupload_pre_key_id.eq(first_unupload_pre_key_id),
device::server_has_prekeys.eq(server_has_prekeys),
device::nct_salt.eq(None::<&[u8]>),
device::server_cert_chain.eq(None::<&[u8]>),
device::login_counter.eq(0i32),
device::lid_migrated.eq(false),
device::last_signed_pre_key_rotation_ms.eq(last_signed_pre_key_rotation_ms),
device::read_receipts_disabled.eq(false),
))
.execute(conn)
.map(|_| device_id)
})
})
.await
}
pub async fn device_exists(&self, device_id: i32) -> Result<bool> {
use crate::schema::device;
let pool = self.pool.clone();
tokio::task::spawn_blocking(move || -> Result<bool> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let count: i64 = device::table
.filter(device::id.eq(device_id))
.count()
.get_result(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(count > 0)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
pub async fn load_device_data_for_device(&self, device_id: i32) -> Result<Option<CoreDevice>> {
use crate::schema::device;
let pool = self.pool.clone();
let row = tokio::task::spawn_blocking(move || -> Result<Option<DeviceRow>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let result = device::table
.filter(device::id.eq(device_id))
.first::<DeviceRow>(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(result)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
if let Some(row) = row {
let pn = if !row.pn.is_empty() {
row.pn.parse().ok()
} else {
None
};
let lid = if !row.lid.is_empty() {
row.lid.parse().ok()
} else {
None
};
let noise_key = self.deserialize_keypair(&row.noise_key)?;
let identity_key = self.deserialize_keypair(&row.identity_key)?;
let signed_pre_key = self.deserialize_keypair(&row.signed_pre_key)?;
let signed_pre_key_signature: [u8; 64] =
row.signed_pre_key_signature.try_into().map_err(|_| {
StoreError::Validation("Invalid signed_pre_key_signature length".to_string())
})?;
let adv_secret_key: [u8; 32] = row
.adv_secret_key
.try_into()
.map_err(|_| StoreError::Validation("Invalid adv_secret_key length".to_string()))?;
let account = row
.account
.map(|data| {
wacore::store::device::account_serde::from_bytes(&data)
.map_err(|e| StoreError::Serialization(Box::new(e)))
})
.transpose()?;
Ok(Some(CoreDevice {
pn,
lid,
registration_id: row.registration_id as u32,
noise_key,
identity_key,
signed_pre_key,
signed_pre_key_id: row.signed_pre_key_id as u32,
signed_pre_key_signature,
adv_secret_key,
account: account.map(Arc::new),
push_name: row.push_name,
app_version_primary: row.app_version_primary as u32,
app_version_secondary: row.app_version_secondary as u32,
app_version_tertiary: row.app_version_tertiary.try_into().unwrap_or(0u32),
app_version_last_fetched_ms: row.app_version_last_fetched_ms,
device_props: Arc::new(wacore::store::device::DEVICE_PROPS.clone()),
client_profile: wacore::client_profile::ClientProfile::web(),
edge_routing_info: row.edge_routing_info,
props_hash: row.props_hash,
next_pre_key_id: row.next_pre_key_id as u32,
first_unupload_pre_key_id: row.first_unupload_pre_key_id as u32,
server_has_prekeys: row.server_has_prekeys,
nct_salt: row.nct_salt,
nct_salt_sync_seen: false,
server_cert_chain: row
.server_cert_chain
.as_deref()
.and_then(|bytes| {
match crate::wire::decode_server_cert_chain(bytes) {
Ok(chain) => Some(chain),
Err(e) => {
log::warn!(
"device {} server_cert_chain blob ({} bytes) failed to decode: {e}; \
dropping cache, next connect will use XX",
self.device_id,
bytes.len(),
);
None
}
}
}),
login_counter: row.login_counter,
lid_migrated: row.lid_migrated,
last_signed_pre_key_rotation_ms: row.last_signed_pre_key_rotation_ms,
read_receipts_disabled: row.read_receipts_disabled,
}))
} else {
Ok(None)
}
}
pub async fn put_identity_for_device(
&self,
address: &str,
key: [u8; 32],
device_id: i32,
) -> Result<()> {
let pool = self.pool.clone();
let db_semaphore = self.db_semaphore.clone();
let address_owned = address.to_string();
let key_vec = key.to_vec();
const MAX_RETRIES: u32 = 5;
for attempt in 0..=MAX_RETRIES {
let permit = db_semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| StoreError::Database(Box::new(e)))?;
let pool_clone = pool.clone();
let address_clone = address_owned.clone();
let key_clone = key_vec.clone();
let result =
tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
let mut conn = pool_clone
.get()
.map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
diesel::insert_into(identities::table)
.values((
identities::address.eq(address_clone),
identities::key.eq(&key_clone[..]),
identities::device_id.eq(device_id),
))
.on_conflict((identities::address, identities::device_id))
.do_update()
.set(identities::key.eq(&key_clone[..]))
.execute(&mut conn)
.map_err(DieselOrStore::Diesel)?;
Ok(())
})
.await;
drop(permit);
match result {
Ok(Ok(())) => return Ok(()),
Ok(Err(DieselOrStore::Diesel(ref e)))
if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
{
let delay_ms = 10 * 2u64.pow(attempt);
warn!(
"Identity write failed (attempt {}/{}): {e}. Retrying in {delay_ms}ms...",
attempt + 1,
MAX_RETRIES + 1,
);
tokio::time::sleep(Duration::from_millis(delay_ms)).await;
continue;
}
Ok(Err(e)) => return Err(e.into()),
Err(e) => return Err(StoreError::Database(Box::new(e))),
}
}
Err(StoreError::RetriesExhausted {
op: format!("identity_write (after {} attempts)", MAX_RETRIES + 1),
})
}
pub async fn delete_identity_for_device(&self, address: &str, device_id: i32) -> Result<()> {
let pool = self.pool.clone();
let address_owned = address.to_string();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::delete(
identities::table
.filter(identities::address.eq(address_owned))
.filter(identities::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
pub async fn load_identity_for_device(
&self,
address: &str,
device_id: i32,
) -> Result<Option<Vec<u8>>> {
let pool = self.pool.clone();
let address = address.to_string();
let result = self
.with_semaphore(move || -> Result<Option<Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let res: Option<Vec<u8>> = identities::table
.select(identities::key)
.filter(identities::address.eq(address))
.filter(identities::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(res)
})
.await?;
Ok(result)
}
pub async fn get_session_for_device(
&self,
address: &str,
device_id: i32,
) -> Result<Option<Vec<u8>>> {
let pool = self.pool.clone();
let address_for_query = address.to_string();
let result = self
.with_semaphore(move || -> Result<Option<Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let res: Option<Vec<u8>> = sessions::table
.select(sessions::record)
.filter(sessions::address.eq(address_for_query.clone()))
.filter(sessions::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(res)
})
.await?;
Ok(result)
}
pub async fn put_session_for_device(
&self,
address: &str,
session: &[u8],
device_id: i32,
) -> Result<()> {
let pool = self.pool.clone();
let db_semaphore = self.db_semaphore.clone();
let address_owned = address.to_string();
let session_vec = session.to_vec();
const MAX_RETRIES: u32 = 5;
for attempt in 0..=MAX_RETRIES {
let permit = db_semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| StoreError::Database(Box::new(e)))?;
let pool_clone = pool.clone();
let address_clone = address_owned.clone();
let session_clone = session_vec.clone();
let result =
tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
let mut conn = pool_clone
.get()
.map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
diesel::insert_into(sessions::table)
.values((
sessions::address.eq(address_clone),
sessions::record.eq(&session_clone),
sessions::device_id.eq(device_id),
))
.on_conflict((sessions::address, sessions::device_id))
.do_update()
.set(sessions::record.eq(&session_clone))
.execute(&mut conn)
.map_err(DieselOrStore::Diesel)?;
Ok(())
})
.await;
drop(permit);
match result {
Ok(Ok(())) => return Ok(()),
Ok(Err(DieselOrStore::Diesel(ref e)))
if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
{
let delay_ms = 10 * 2u64.pow(attempt);
warn!(
"Session write failed (attempt {}/{}): {e}. Retrying in {delay_ms}ms...",
attempt + 1,
MAX_RETRIES + 1,
);
tokio::time::sleep(Duration::from_millis(delay_ms)).await;
continue;
}
Ok(Err(e)) => return Err(e.into()),
Err(e) => return Err(StoreError::Database(Box::new(e))),
}
}
Err(StoreError::RetriesExhausted {
op: format!("session_write (after {} attempts)", MAX_RETRIES + 1),
})
}
pub async fn delete_session_for_device(&self, address: &str, device_id: i32) -> Result<()> {
let pool = self.pool.clone();
let address_owned = address.to_string();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::delete(
sessions::table
.filter(sessions::address.eq(address_owned))
.filter(sessions::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
pub async fn put_sender_key_for_device(
&self,
address: &str,
record: &[u8],
device_id: i32,
) -> Result<()> {
let pool = self.pool.clone();
let address = address.to_string();
let record_vec = record.to_vec();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::insert_into(sender_keys::table)
.values((
sender_keys::address.eq(address),
sender_keys::record.eq(&record_vec),
sender_keys::device_id.eq(device_id),
))
.on_conflict((sender_keys::address, sender_keys::device_id))
.do_update()
.set(sender_keys::record.eq(&record_vec))
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
pub async fn get_sender_key_for_device(
&self,
address: &str,
device_id: i32,
) -> Result<Option<Vec<u8>>> {
let pool = self.pool.clone();
let address = address.to_string();
tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let res: Option<Vec<u8>> = sender_keys::table
.select(sender_keys::record)
.filter(sender_keys::address.eq(address))
.filter(sender_keys::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(res)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
pub async fn delete_sender_key_for_device(&self, address: &str, device_id: i32) -> Result<()> {
let pool = self.pool.clone();
let address = address.to_string();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::delete(
sender_keys::table
.filter(sender_keys::address.eq(address))
.filter(sender_keys::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
pub async fn get_app_state_sync_key_for_device(
&self,
key_id: &[u8],
device_id: i32,
) -> Result<Option<AppStateSyncKey>> {
let pool = self.pool.clone();
let key_id = key_id.to_vec();
let res: Option<Vec<u8>> =
tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let res: Option<Vec<u8>> = app_state_keys::table
.select(app_state_keys::key_data)
.filter(app_state_keys::key_id.eq(&key_id))
.filter(app_state_keys::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(res)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
if let Some(data) = res {
match crate::wire::decode_app_state_sync_key(&data) {
Ok(key) => Ok(Some(key)),
Err(e) => {
warn!(
"app_state_sync_key blob ({} bytes) failed to decode: {e}; \
treating as absent, key will be re-requested",
data.len()
);
Ok(None)
}
}
} else {
Ok(None)
}
}
pub async fn set_app_state_sync_key_for_device(
&self,
key_id: &[u8],
key: AppStateSyncKey,
device_id: i32,
) -> Result<()> {
let pool = self.pool.clone();
let key_id = key_id.to_vec();
let data = crate::wire::encode_app_state_sync_key(&key);
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::insert_into(app_state_keys::table)
.values((
app_state_keys::key_id.eq(&key_id),
app_state_keys::key_data.eq(&data),
app_state_keys::device_id.eq(device_id),
))
.on_conflict((app_state_keys::key_id, app_state_keys::device_id))
.do_update()
.set(app_state_keys::key_data.eq(&data))
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
pub async fn get_latest_app_state_sync_key_id_for_device(
&self,
device_id: i32,
) -> Result<Option<Vec<u8>>> {
let pool = self.pool.clone();
let res: Option<Vec<u8>> =
tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let candidates: Vec<(Vec<u8>, Vec<u8>)> = app_state_keys::table
.select((app_state_keys::key_id, app_state_keys::key_data))
.filter(app_state_keys::device_id.eq(device_id))
.order(app_state_keys::key_id.desc())
.load(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
let res = candidates
.into_iter()
.find(|(_, data)| crate::wire::decode_app_state_sync_key(data).is_ok())
.map(|(key_id, _)| key_id);
Ok(res)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(res)
}
pub async fn get_app_state_version_for_device(
&self,
name: &str,
device_id: i32,
) -> Result<HashState> {
let pool = self.pool.clone();
let name = name.to_string();
let res: Option<Vec<u8>> =
tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let res: Option<Vec<u8>> = app_state_versions::table
.select(app_state_versions::state_data)
.filter(app_state_versions::name.eq(name))
.filter(app_state_versions::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(res)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
if let Some(data) = res {
match crate::wire::decode_hash_state(&data) {
Ok(state) => Ok(state),
Err(e) => {
warn!(
"app_state_version blob ({} bytes) failed to decode: {e}; \
resetting to default, collection will re-sync from 0",
data.len()
);
Ok(HashState::default())
}
}
} else {
Ok(HashState::default())
}
}
pub async fn set_app_state_version_for_device(
&self,
name: &str,
state: HashState,
device_id: i32,
) -> Result<()> {
let name = name.to_string();
let data = crate::wire::encode_hash_state(&state);
self.with_retry("set_app_state_version", || {
let name = name.clone();
let data = data.clone();
Box::new(move |conn: &mut SqliteConnection| {
diesel::insert_into(app_state_versions::table)
.values((
app_state_versions::name.eq(&name),
app_state_versions::state_data.eq(&data),
app_state_versions::device_id.eq(device_id),
))
.on_conflict((app_state_versions::name, app_state_versions::device_id))
.do_update()
.set(app_state_versions::state_data.eq(&data))
.execute(conn)?;
Ok(())
})
})
.await
}
pub async fn put_app_state_mutation_macs_for_device(
&self,
name: &str,
version: u64,
mutations: &[AppStateMutationMAC],
device_id: i32,
) -> Result<()> {
if mutations.is_empty() {
return Ok(());
}
let name = name.to_string();
let mutations: Vec<AppStateMutationMAC> = mutations.to_vec();
self.with_retry("put_app_state_mutation_macs", || {
let name = name.clone();
let mutations = mutations.clone();
Box::new(move |conn: &mut SqliteConnection| {
let records: Vec<_> = mutations
.iter()
.map(|m| {
(
app_state_mutation_macs::name.eq(&name),
app_state_mutation_macs::version.eq(version as i64),
app_state_mutation_macs::index_mac.eq(&m.index_mac),
app_state_mutation_macs::value_mac.eq(&m.value_mac),
app_state_mutation_macs::device_id.eq(device_id),
)
})
.collect();
const CHUNK_SIZE: usize = 100;
for chunk in records.chunks(CHUNK_SIZE) {
diesel::insert_into(app_state_mutation_macs::table)
.values(chunk)
.on_conflict((
app_state_mutation_macs::name,
app_state_mutation_macs::index_mac,
app_state_mutation_macs::device_id,
))
.do_update()
.set((
app_state_mutation_macs::version
.eq(excluded(app_state_mutation_macs::version)),
app_state_mutation_macs::value_mac
.eq(excluded(app_state_mutation_macs::value_mac)),
))
.execute(conn)?;
}
Ok(())
})
})
.await
}
pub async fn delete_app_state_mutation_macs_for_device(
&self,
name: &str,
index_macs: &[Vec<u8>],
device_id: i32,
) -> Result<()> {
if index_macs.is_empty() {
return Ok(());
}
let name = name.to_string();
let index_macs: Vec<Vec<u8>> = index_macs.to_vec();
self.with_retry("delete_app_state_mutation_macs", || {
let name = name.clone();
let index_macs = index_macs.clone();
Box::new(move |conn: &mut SqliteConnection| {
const CHUNK_SIZE: usize = 500;
for chunk in index_macs.chunks(CHUNK_SIZE) {
diesel::delete(
app_state_mutation_macs::table.filter(
app_state_mutation_macs::name
.eq(&name)
.and(app_state_mutation_macs::index_mac.eq_any(chunk))
.and(app_state_mutation_macs::device_id.eq(device_id)),
),
)
.execute(conn)?;
}
Ok(())
})
})
.await
}
pub async fn get_app_state_mutation_mac_for_device(
&self,
name: &str,
index_mac: &[u8],
device_id: i32,
) -> Result<Option<Vec<u8>>> {
let pool = self.pool.clone();
let name = name.to_string();
let index_mac = index_mac.to_vec();
tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let res: Option<Vec<u8>> = app_state_mutation_macs::table
.select(app_state_mutation_macs::value_mac)
.filter(app_state_mutation_macs::name.eq(&name))
.filter(app_state_mutation_macs::index_mac.eq(&index_mac))
.filter(app_state_mutation_macs::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(res)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
pub async fn get_app_state_mutation_macs_batch_for_device(
&self,
name: &str,
index_macs: &[[u8; 32]],
device_id: i32,
) -> Result<std::collections::HashMap<[u8; 32], Vec<u8>>> {
if index_macs.is_empty() {
return Ok(std::collections::HashMap::new());
}
let pool = self.pool.clone();
let name = name.to_string();
let index_macs: Vec<[u8; 32]> = index_macs.to_vec();
tokio::task::spawn_blocking(
move || -> Result<std::collections::HashMap<[u8; 32], Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let mut out = std::collections::HashMap::with_capacity(index_macs.len());
const CHUNK_SIZE: usize = 500;
for chunk in index_macs.chunks(CHUNK_SIZE) {
let chunk_slices: Vec<&[u8]> = chunk.iter().map(|m| m.as_slice()).collect();
let rows: Vec<(Vec<u8>, Vec<u8>)> = app_state_mutation_macs::table
.select((
app_state_mutation_macs::index_mac,
app_state_mutation_macs::value_mac,
))
.filter(app_state_mutation_macs::name.eq(&name))
.filter(app_state_mutation_macs::index_mac.eq_any(chunk_slices))
.filter(app_state_mutation_macs::device_id.eq(device_id))
.load(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
out.extend(rows.into_iter().filter_map(|(k, v)| {
<[u8; 32]>::try_from(k.as_slice()).ok().map(|k| (k, v))
}));
}
Ok(out)
},
)
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl SignalStore for SqliteStore {
async fn put_identity(&self, address: &str, key: [u8; 32]) -> Result<()> {
self.put_identity_for_device(address, key, self.device_id)
.await
}
async fn put_identities_batch(&self, identities: &[(Arc<str>, [u8; 32])]) -> Result<()> {
if identities.is_empty() {
return Ok(());
}
let device_id = self.device_id;
let batch = Arc::new(identities.to_vec());
self.with_retry("put_identities_batch", || {
let batch = batch.clone();
Box::new(move |conn: &mut SqliteConnection| {
conn.transaction(|conn| {
for (address, key) in batch.iter() {
diesel::insert_into(identities::table)
.values((
identities::address.eq(address.as_ref()),
identities::key.eq(&key[..]),
identities::device_id.eq(device_id),
))
.on_conflict((identities::address, identities::device_id))
.do_update()
.set(identities::key.eq(&key[..]))
.execute(conn)?;
}
Ok(())
})
})
})
.await
}
async fn load_identity(&self, address: &str) -> Result<Option<[u8; 32]>> {
let blob = self
.load_identity_for_device(address, self.device_id)
.await?;
match blob {
None => Ok(None),
Some(v) => Ok(Some(v.try_into().map_err(|v: Vec<u8>| {
StoreError::Validation(format!(
"identity key for '{}' has invalid length {} (expected 32)",
address,
v.len()
))
})?)),
}
}
async fn delete_identity(&self, address: &str) -> Result<()> {
self.delete_identity_for_device(address, self.device_id)
.await
}
async fn get_session(&self, address: &str) -> Result<Option<Bytes>> {
Ok(self
.get_session_for_device(address, self.device_id)
.await?
.map(Bytes::from))
}
async fn has_session(&self, address: &str) -> Result<bool> {
let pool = self.pool.clone();
let device_id = self.device_id;
let address_owned = address.to_string();
self.with_semaphore(move || -> Result<bool> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let exists = diesel::select(diesel::dsl::exists(
sessions::table
.filter(sessions::address.eq(&address_owned))
.filter(sessions::device_id.eq(device_id)),
))
.get_result(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(exists)
})
.await
}
async fn has_signal_state_for_user(&self, user: &str) -> Result<bool> {
let pool = self.pool.clone();
let device_id = self.device_id;
let pat_at = format!("{user}@%");
let pat_dev = format!("{user}:%");
self.with_semaphore(move || -> Result<bool> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let has_session = diesel::select(diesel::dsl::exists(
sessions::table
.filter(sessions::device_id.eq(device_id))
.filter(
sessions::address
.like(&pat_at)
.or(sessions::address.like(&pat_dev)),
),
))
.get_result::<bool>(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
if has_session {
return Ok(true);
}
let has_identity = diesel::select(diesel::dsl::exists(
identities::table
.filter(identities::device_id.eq(device_id))
.filter(
identities::address
.like(&pat_at)
.or(identities::address.like(&pat_dev)),
),
))
.get_result::<bool>(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(has_identity)
})
.await
}
async fn put_session(&self, address: &str, session: &[u8]) -> Result<()> {
self.put_session_for_device(address, session, self.device_id)
.await
}
async fn put_sessions_batch(&self, sessions: &[(Arc<str>, Bytes)]) -> Result<()> {
if sessions.is_empty() {
return Ok(());
}
let device_id = self.device_id;
let batch = Arc::new(sessions.to_vec());
self.with_retry("put_sessions_batch", || {
let batch = batch.clone();
Box::new(move |conn: &mut SqliteConnection| {
conn.transaction(|conn| {
for (address, record) in batch.iter() {
diesel::insert_into(sessions::table)
.values((
sessions::address.eq(address.as_ref()),
sessions::record.eq(record.as_ref()),
sessions::device_id.eq(device_id),
))
.on_conflict((sessions::address, sessions::device_id))
.do_update()
.set(sessions::record.eq(record.as_ref()))
.execute(conn)?;
}
Ok(())
})
})
})
.await
}
async fn delete_session(&self, address: &str) -> Result<()> {
self.delete_session_for_device(address, self.device_id)
.await
}
async fn store_prekey(&self, id: u32, record: &[u8], uploaded: bool) -> Result<()> {
let pool = self.pool.clone();
let db_semaphore = self.db_semaphore.clone();
let device_id = self.device_id;
let record = record.to_vec();
const MAX_RETRIES: u32 = 5;
for attempt in 0..=MAX_RETRIES {
let permit = db_semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| StoreError::Database(Box::new(e)))?;
let pool_clone = pool.clone();
let record_clone = record.clone();
let result =
tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
let mut conn = pool_clone
.get()
.map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
diesel::insert_into(prekeys::table)
.values((
prekeys::id.eq(id as i32),
prekeys::key.eq(&record_clone),
prekeys::uploaded.eq(uploaded),
prekeys::device_id.eq(device_id),
))
.on_conflict((prekeys::id, prekeys::device_id))
.do_update()
.set((
prekeys::key.eq(&record_clone),
prekeys::uploaded.eq(uploaded),
))
.execute(&mut conn)
.map_err(DieselOrStore::Diesel)?;
Ok(())
})
.await;
drop(permit);
match result {
Ok(Ok(())) => return Ok(()),
Ok(Err(DieselOrStore::Diesel(ref e)))
if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
{
let delay_ms = 10u64 * (1u64 << attempt.min(4));
tokio::time::sleep(Duration::from_millis(delay_ms)).await;
}
Ok(Err(e)) => return Err(e.into()),
Err(e) => return Err(StoreError::Database(Box::new(e))),
}
}
Err(StoreError::RetriesExhausted {
op: "store_prekey".to_string(),
})
}
async fn store_prekeys_batch(&self, keys: &[(u32, Bytes)], uploaded: bool) -> Result<()> {
if keys.is_empty() {
return Ok(());
}
let pool = self.pool.clone();
let db_semaphore = self.db_semaphore.clone();
let device_id = self.device_id;
let keys: Vec<(u32, Bytes)> = keys.to_vec();
const MAX_RETRIES: u32 = 5;
for attempt in 0..=MAX_RETRIES {
let permit = db_semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| StoreError::Database(Box::new(e)))?;
let pool_clone = pool.clone();
let keys_clone = keys.clone();
let result =
tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
let mut conn = pool_clone
.get()
.map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
conn.transaction(|conn| {
for (id, record) in &keys_clone {
diesel::insert_into(prekeys::table)
.values((
prekeys::id.eq(*id as i32),
prekeys::key.eq(record.as_ref()),
prekeys::uploaded.eq(uploaded),
prekeys::device_id.eq(device_id),
))
.on_conflict((prekeys::id, prekeys::device_id))
.do_update()
.set((
prekeys::key.eq(record.as_ref()),
prekeys::uploaded.eq(uploaded),
))
.execute(conn)?;
}
Ok::<(), diesel::result::Error>(())
})
.map_err(DieselOrStore::Diesel)
})
.await;
drop(permit);
match result {
Ok(Ok(())) => return Ok(()),
Ok(Err(DieselOrStore::Diesel(ref e)))
if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
{
let delay_ms = 10u64 * (1u64 << attempt.min(4));
tokio::time::sleep(Duration::from_millis(delay_ms)).await;
}
Ok(Err(e)) => return Err(e.into()),
Err(e) => return Err(StoreError::Database(Box::new(e))),
}
}
Err(StoreError::RetriesExhausted {
op: "store_prekeys_batch".to_string(),
})
}
async fn load_prekey(&self, id: u32) -> Result<Option<Bytes>> {
let pool = self.pool.clone();
let device_id = self.device_id;
tokio::task::spawn_blocking(move || -> Result<Option<Bytes>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let res: Option<Vec<u8>> = prekeys::table
.select(prekeys::key)
.filter(prekeys::id.eq(id as i32))
.filter(prekeys::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(res.map(Bytes::from))
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn load_prekeys_batch(&self, ids: &[u32]) -> Result<Vec<(u32, Bytes)>> {
if ids.is_empty() {
return Ok(Vec::new());
}
let pool = self.pool.clone();
let device_id = self.device_id;
let ids: Vec<i32> = ids.iter().map(|&id| id as i32).collect();
self.with_semaphore(move || -> Result<Vec<(u32, Bytes)>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let mut out = Vec::with_capacity(ids.len());
for chunk in ids.chunks(ID_PARAM_CHUNK) {
let rows: Vec<(i32, Vec<u8>)> = prekeys::table
.select((prekeys::id, prekeys::key))
.filter(prekeys::id.eq_any(chunk))
.filter(prekeys::device_id.eq(device_id))
.load(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
out.extend(
rows.into_iter()
.map(|(id, key)| (id as u32, Bytes::from(key))),
);
}
Ok(out)
})
.await
}
async fn remove_prekey(&self, id: u32) -> Result<()> {
let pool = self.pool.clone();
let db_semaphore = self.db_semaphore.clone();
let device_id = self.device_id;
const MAX_RETRIES: u32 = 5;
for attempt in 0..=MAX_RETRIES {
let permit = db_semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| StoreError::Database(Box::new(e)))?;
let pool_clone = pool.clone();
let result =
tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
let mut conn = pool_clone
.get()
.map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
diesel::delete(
prekeys::table
.filter(prekeys::id.eq(id as i32))
.filter(prekeys::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(DieselOrStore::Diesel)?;
Ok(())
})
.await;
drop(permit);
match result {
Ok(Ok(())) => return Ok(()),
Ok(Err(DieselOrStore::Diesel(ref e)))
if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
{
let delay_ms = 10u64 * (1u64 << attempt.min(4));
tokio::time::sleep(Duration::from_millis(delay_ms)).await;
}
Ok(Err(e)) => return Err(e.into()),
Err(e) => return Err(StoreError::Database(Box::new(e))),
}
}
Err(StoreError::RetriesExhausted {
op: "remove_prekey".to_string(),
})
}
async fn mark_prekeys_uploaded(&self, ids: &[u32]) -> Result<()> {
if ids.is_empty() {
return Ok(());
}
let device_id = self.device_id;
let ids: Vec<i32> = ids.iter().map(|&id| id as i32).collect();
self.with_retry("mark_prekeys_uploaded", move || {
let ids = ids.clone();
Box::new(move |conn: &mut SqliteConnection| {
for chunk in ids.chunks(ID_PARAM_CHUNK) {
diesel::update(
prekeys::table
.filter(prekeys::id.eq_any(chunk.to_vec()))
.filter(prekeys::device_id.eq(device_id)),
)
.set(prekeys::uploaded.eq(true))
.execute(conn)?;
}
Ok(())
})
})
.await
}
async fn get_max_prekey_id(&self) -> Result<u32> {
let pool = self.pool.clone();
let device_id = self.device_id;
let db_semaphore = self.db_semaphore.clone();
let _permit = db_semaphore
.acquire()
.await
.map_err(|e| StoreError::Database(Box::new(e)))?;
tokio::task::spawn_blocking(move || -> Result<u32> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
use diesel::dsl::max;
let result: Option<i32> = prekeys::table
.filter(prekeys::device_id.eq(device_id))
.select(max(prekeys::id))
.first(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(result.unwrap_or(0) as u32)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn store_signed_prekey(&self, id: u32, record: &[u8]) -> Result<()> {
let pool = self.pool.clone();
let db_semaphore = self.db_semaphore.clone();
let device_id = self.device_id;
let record = record.to_vec();
const MAX_RETRIES: u32 = 5;
for attempt in 0..=MAX_RETRIES {
let permit = db_semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| StoreError::Database(Box::new(e)))?;
let pool_clone = pool.clone();
let record_clone = record.clone();
let result =
tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
let mut conn = pool_clone
.get()
.map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
diesel::insert_into(signed_prekeys::table)
.values((
signed_prekeys::id.eq(id as i32),
signed_prekeys::record.eq(&record_clone),
signed_prekeys::device_id.eq(device_id),
))
.on_conflict((signed_prekeys::id, signed_prekeys::device_id))
.do_update()
.set(signed_prekeys::record.eq(&record_clone))
.execute(&mut conn)
.map_err(DieselOrStore::Diesel)?;
Ok(())
})
.await;
drop(permit);
match result {
Ok(Ok(())) => return Ok(()),
Ok(Err(DieselOrStore::Diesel(ref e)))
if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
{
let delay_ms = 10u64 * (1u64 << attempt.min(4));
tokio::time::sleep(Duration::from_millis(delay_ms)).await;
}
Ok(Err(e)) => return Err(e.into()),
Err(e) => return Err(StoreError::Database(Box::new(e))),
}
}
Err(StoreError::RetriesExhausted {
op: "store_signed_prekey".to_string(),
})
}
async fn load_signed_prekey(&self, id: u32) -> Result<Option<Vec<u8>>> {
let pool = self.pool.clone();
let device_id = self.device_id;
tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let res: Option<Vec<u8>> = signed_prekeys::table
.select(signed_prekeys::record)
.filter(signed_prekeys::id.eq(id as i32))
.filter(signed_prekeys::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(res)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn load_all_signed_prekeys(&self) -> Result<Vec<(u32, Vec<u8>)>> {
let pool = self.pool.clone();
let device_id = self.device_id;
tokio::task::spawn_blocking(move || -> Result<Vec<(u32, Vec<u8>)>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let results: Vec<(i32, Vec<u8>)> = signed_prekeys::table
.select((signed_prekeys::id, signed_prekeys::record))
.filter(signed_prekeys::device_id.eq(device_id))
.load(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(results
.into_iter()
.map(|(id, record)| (id as u32, record))
.collect())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn remove_signed_prekey(&self, id: u32) -> Result<()> {
let pool = self.pool.clone();
let db_semaphore = self.db_semaphore.clone();
let device_id = self.device_id;
const MAX_RETRIES: u32 = 5;
for attempt in 0..=MAX_RETRIES {
let permit = db_semaphore
.clone()
.acquire_owned()
.await
.map_err(|e| StoreError::Database(Box::new(e)))?;
let pool_clone = pool.clone();
let result =
tokio::task::spawn_blocking(move || -> std::result::Result<(), DieselOrStore> {
let mut conn = pool_clone
.get()
.map_err(|e| DieselOrStore::Store(StoreError::Connection(Box::new(e))))?;
diesel::delete(
signed_prekeys::table
.filter(signed_prekeys::id.eq(id as i32))
.filter(signed_prekeys::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(DieselOrStore::Diesel)?;
Ok(())
})
.await;
drop(permit);
match result {
Ok(Ok(())) => return Ok(()),
Ok(Err(DieselOrStore::Diesel(ref e)))
if is_retriable_sqlite_error(e) && attempt < MAX_RETRIES =>
{
let delay_ms = 10u64 * (1u64 << attempt.min(4));
tokio::time::sleep(Duration::from_millis(delay_ms)).await;
}
Ok(Err(e)) => return Err(e.into()),
Err(e) => return Err(StoreError::Database(Box::new(e))),
}
}
Err(StoreError::RetriesExhausted {
op: "remove_signed_prekey".to_string(),
})
}
async fn put_sender_key(&self, address: &str, record: &[u8]) -> Result<()> {
self.put_sender_key_for_device(address, record, self.device_id)
.await
}
async fn put_sender_keys_batch(&self, sender_keys: &[(Arc<str>, Bytes)]) -> Result<()> {
if sender_keys.is_empty() {
return Ok(());
}
let device_id = self.device_id;
let batch = Arc::new(sender_keys.to_vec());
self.with_retry("put_sender_keys_batch", || {
let batch = batch.clone();
Box::new(move |conn: &mut SqliteConnection| {
conn.transaction(|conn| {
for (address, record) in batch.iter() {
diesel::insert_into(sender_keys::table)
.values((
sender_keys::address.eq(address.as_ref()),
sender_keys::record.eq(record.as_ref()),
sender_keys::device_id.eq(device_id),
))
.on_conflict((sender_keys::address, sender_keys::device_id))
.do_update()
.set(sender_keys::record.eq(record.as_ref()))
.execute(conn)?;
}
Ok(())
})
})
})
.await
}
async fn get_sender_key(&self, address: &str) -> Result<Option<Vec<u8>>> {
self.get_sender_key_for_device(address, self.device_id)
.await
}
async fn delete_sender_key(&self, address: &str) -> Result<()> {
self.delete_sender_key_for_device(address, self.device_id)
.await
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl AppSyncStore for SqliteStore {
async fn get_sync_key(&self, key_id: &[u8]) -> Result<Option<AppStateSyncKey>> {
self.get_app_state_sync_key_for_device(key_id, self.device_id)
.await
}
async fn set_sync_key(&self, key_id: &[u8], key: AppStateSyncKey) -> Result<()> {
self.set_app_state_sync_key_for_device(key_id, key, self.device_id)
.await
}
async fn get_version(&self, name: &str) -> Result<HashState> {
self.get_app_state_version_for_device(name, self.device_id)
.await
}
async fn set_version(&self, name: &str, state: HashState) -> Result<()> {
self.set_app_state_version_for_device(name, state, self.device_id)
.await
}
async fn put_mutation_macs(
&self,
name: &str,
version: u64,
mutations: &[AppStateMutationMAC],
) -> Result<()> {
self.put_app_state_mutation_macs_for_device(name, version, mutations, self.device_id)
.await
}
async fn get_mutation_mac(&self, name: &str, index_mac: &[u8]) -> Result<Option<Vec<u8>>> {
self.get_app_state_mutation_mac_for_device(name, index_mac, self.device_id)
.await
}
async fn get_mutation_macs(
&self,
name: &str,
index_macs: &[[u8; 32]],
) -> Result<std::collections::HashMap<[u8; 32], Vec<u8>>> {
self.get_app_state_mutation_macs_batch_for_device(name, index_macs, self.device_id)
.await
}
async fn delete_mutation_macs(&self, name: &str, index_macs: &[Vec<u8>]) -> Result<()> {
self.delete_app_state_mutation_macs_for_device(name, index_macs, self.device_id)
.await
}
async fn clear_mutation_macs(&self, name: &str) -> Result<()> {
let device_id = self.device_id;
let name = name.to_string();
self.with_retry("clear_mutation_macs", || {
let name = name.clone();
Box::new(move |conn: &mut SqliteConnection| {
diesel::delete(
app_state_mutation_macs::table
.filter(app_state_mutation_macs::name.eq(&name))
.filter(app_state_mutation_macs::device_id.eq(device_id)),
)
.execute(conn)?;
Ok(())
})
})
.await
}
async fn get_latest_sync_key_id(&self) -> Result<Option<Vec<u8>>> {
self.get_latest_app_state_sync_key_id_for_device(self.device_id)
.await
}
}
fn insert_pending_inbound_row(
conn: &mut SqliteConnection,
device_id: i32,
chat: &str,
sender: &str,
id: &str,
message: &[u8],
) -> QueryResult<usize> {
diesel::replace_into(pending_inbound_messages::table)
.values((
pending_inbound_messages::chat.eq(chat),
pending_inbound_messages::sender.eq(sender),
pending_inbound_messages::id.eq(id),
pending_inbound_messages::message.eq(message),
pending_inbound_messages::device_id.eq(device_id),
))
.execute(conn)
}
fn delete_pending_inbound_row(
conn: &mut SqliteConnection,
device_id: i32,
chat: &str,
sender: &str,
id: &str,
) -> QueryResult<usize> {
diesel::delete(
pending_inbound_messages::table
.filter(pending_inbound_messages::chat.eq(chat))
.filter(pending_inbound_messages::sender.eq(sender))
.filter(pending_inbound_messages::id.eq(id))
.filter(pending_inbound_messages::device_id.eq(device_id)),
)
.execute(conn)
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl ProtocolStore for SqliteStore {
async fn get_sender_key_devices(&self, group_jid: &str) -> Result<Vec<(String, bool)>> {
let pool = self.pool.clone();
let device_id = self.device_id;
let group_jid = group_jid.to_string();
tokio::task::spawn_blocking(move || -> Result<Vec<(String, bool)>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let rows: Vec<(String, i32)> = sender_key_devices::table
.select((sender_key_devices::device_jid, sender_key_devices::has_key))
.filter(sender_key_devices::group_jid.eq(&group_jid))
.filter(sender_key_devices::device_id.eq(device_id))
.load(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(rows
.into_iter()
.map(|(jid, has_key)| (jid, has_key != 0))
.collect())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn set_sender_key_status(&self, group_jid: &str, entries: &[(&str, bool)]) -> Result<()> {
if entries.is_empty() {
return Ok(());
}
let device_id = self.device_id;
let group_jid = group_jid.to_string();
let owned_entries: Arc<Vec<(String, bool)>> = Arc::new(
entries
.iter()
.map(|(jid, has_key)| (jid.to_string(), *has_key))
.collect(),
);
let now = wacore::time::now_secs();
self.with_retry("set_sender_key_status", || {
let group_jid = group_jid.clone();
let owned_entries = Arc::clone(&owned_entries);
Box::new(move |conn: &mut SqliteConnection| {
let values: Vec<_> = owned_entries
.iter()
.map(|(device_jid, has_key)| {
(
sender_key_devices::group_jid.eq(&group_jid),
sender_key_devices::device_jid.eq(device_jid),
sender_key_devices::has_key.eq(i32::from(*has_key)),
sender_key_devices::device_id.eq(device_id),
sender_key_devices::updated_at.eq(now),
)
})
.collect();
const CHUNK_SIZE: usize = 190;
for chunk in values.chunks(CHUNK_SIZE) {
diesel::insert_into(sender_key_devices::table)
.values(chunk)
.on_conflict((
sender_key_devices::group_jid,
sender_key_devices::device_jid,
sender_key_devices::device_id,
))
.do_update()
.set((
sender_key_devices::has_key.eq(excluded(sender_key_devices::has_key)),
sender_key_devices::updated_at.eq(now),
))
.execute(conn)?;
}
Ok(())
})
})
.await
}
async fn clear_sender_key_devices(&self, group_jid: &str) -> Result<()> {
let device_id = self.device_id;
let group_jid = group_jid.to_string();
self.with_retry("clear_sender_key_devices", || {
let group_jid = group_jid.clone();
Box::new(move |conn: &mut SqliteConnection| {
diesel::delete(
sender_key_devices::table
.filter(sender_key_devices::group_jid.eq(&group_jid))
.filter(sender_key_devices::device_id.eq(device_id)),
)
.execute(conn)?;
Ok(())
})
})
.await
}
async fn clear_all_sender_key_devices(&self) -> Result<()> {
let device_id = self.device_id;
self.with_retry("clear_all_sender_key_devices", || {
Box::new(move |conn: &mut SqliteConnection| {
diesel::delete(
sender_key_devices::table.filter(sender_key_devices::device_id.eq(device_id)),
)
.execute(conn)?;
Ok(())
})
})
.await
}
async fn delete_sender_key_device_rows(&self, device_jids: &[&str]) -> Result<()> {
if device_jids.is_empty() {
return Ok(());
}
let device_id = self.device_id;
let owned: Arc<Vec<String>> = Arc::new(device_jids.iter().map(|s| s.to_string()).collect());
self.with_retry("delete_sender_key_device_rows", || {
let owned = Arc::clone(&owned);
Box::new(move |conn: &mut SqliteConnection| {
const CHUNK: usize = 190;
for chunk in owned.chunks(CHUNK) {
diesel::delete(
sender_key_devices::table
.filter(sender_key_devices::device_jid.eq_any(chunk))
.filter(sender_key_devices::device_id.eq(device_id)),
)
.execute(conn)?;
}
Ok(())
})
})
.await
}
async fn get_lid_mapping(&self, lid: &str) -> Result<Option<LidPnMappingEntry>> {
let pool = self.pool.clone();
let device_id = self.device_id;
let lid = lid.to_string();
tokio::task::spawn_blocking(move || -> Result<Option<LidPnMappingEntry>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let row: Option<(String, String, i64, String, i64)> = lid_pn_mapping::table
.select((
lid_pn_mapping::lid,
lid_pn_mapping::phone_number,
lid_pn_mapping::created_at,
lid_pn_mapping::learning_source,
lid_pn_mapping::updated_at,
))
.filter(lid_pn_mapping::lid.eq(&lid))
.filter(lid_pn_mapping::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(row.map(
|(lid, phone_number, created_at, learning_source, updated_at)| LidPnMappingEntry {
lid,
phone_number,
created_at,
updated_at,
learning_source,
},
))
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn get_pn_mapping(&self, phone: &str) -> Result<Option<LidPnMappingEntry>> {
let pool = self.pool.clone();
let device_id = self.device_id;
let phone = phone.to_string();
tokio::task::spawn_blocking(move || -> Result<Option<LidPnMappingEntry>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let row: Option<(String, String, i64, String, i64)> = lid_pn_mapping::table
.select((
lid_pn_mapping::lid,
lid_pn_mapping::phone_number,
lid_pn_mapping::created_at,
lid_pn_mapping::learning_source,
lid_pn_mapping::updated_at,
))
.filter(lid_pn_mapping::phone_number.eq(&phone))
.filter(lid_pn_mapping::device_id.eq(device_id))
.order(lid_pn_mapping::updated_at.desc())
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(row.map(
|(lid, phone_number, created_at, learning_source, updated_at)| LidPnMappingEntry {
lid,
phone_number,
created_at,
updated_at,
learning_source,
},
))
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn put_lid_mapping(&self, entry: &LidPnMappingEntry) -> Result<()> {
self.put_lid_mappings(std::slice::from_ref(entry)).await
}
async fn put_lid_mappings(&self, entries: &[LidPnMappingEntry]) -> Result<()> {
if entries.is_empty() {
return Ok(());
}
let device_id = self.device_id;
let entries: Arc<Vec<LidPnMappingEntry>> = Arc::new(entries.to_vec());
self.with_retry("put_lid_mappings", move || {
let entries = Arc::clone(&entries);
Box::new(move |conn: &mut SqliteConnection| {
conn.transaction::<_, DieselError, _>(|conn| {
for entry in entries.iter() {
diesel::insert_into(lid_pn_mapping::table)
.values((
lid_pn_mapping::lid.eq(&entry.lid),
lid_pn_mapping::phone_number.eq(&entry.phone_number),
lid_pn_mapping::created_at.eq(entry.created_at),
lid_pn_mapping::learning_source.eq(&entry.learning_source),
lid_pn_mapping::updated_at.eq(entry.updated_at),
lid_pn_mapping::device_id.eq(device_id),
))
.on_conflict((lid_pn_mapping::lid, lid_pn_mapping::device_id))
.do_update()
.set((
lid_pn_mapping::phone_number.eq(&entry.phone_number),
lid_pn_mapping::learning_source.eq(&entry.learning_source),
lid_pn_mapping::updated_at.eq(entry.updated_at),
))
.execute(conn)?;
}
Ok(())
})
})
})
.await
}
async fn get_all_lid_mappings(&self) -> Result<Vec<LidPnMappingEntry>> {
let pool = self.pool.clone();
let device_id = self.device_id;
tokio::task::spawn_blocking(move || -> Result<Vec<LidPnMappingEntry>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let rows: Vec<(String, String, i64, String, i64)> = lid_pn_mapping::table
.select((
lid_pn_mapping::lid,
lid_pn_mapping::phone_number,
lid_pn_mapping::created_at,
lid_pn_mapping::learning_source,
lid_pn_mapping::updated_at,
))
.filter(lid_pn_mapping::device_id.eq(device_id))
.load(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(rows
.into_iter()
.map(
|(lid, phone_number, created_at, learning_source, updated_at)| {
LidPnMappingEntry {
lid,
phone_number,
created_at,
updated_at,
learning_source,
}
},
)
.collect())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn save_base_key(&self, address: &str, message_id: &str, base_key: &[u8]) -> Result<()> {
let pool = self.pool.clone();
let device_id = self.device_id;
let address = address.to_string();
let message_id = message_id.to_string();
let base_key = base_key.to_vec();
let now = wacore::time::now_secs() as i32;
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::insert_into(base_keys::table)
.values((
base_keys::address.eq(&address),
base_keys::message_id.eq(&message_id),
base_keys::base_key.eq(&base_key),
base_keys::device_id.eq(device_id),
base_keys::created_at.eq(now),
))
.on_conflict((
base_keys::address,
base_keys::message_id,
base_keys::device_id,
))
.do_update()
.set(base_keys::base_key.eq(&base_key))
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
async fn has_same_base_key(
&self,
address: &str,
message_id: &str,
current_base_key: &[u8],
) -> Result<bool> {
let pool = self.pool.clone();
let device_id = self.device_id;
let address = address.to_string();
let message_id = message_id.to_string();
let current_base_key = current_base_key.to_vec();
tokio::task::spawn_blocking(move || -> Result<bool> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let stored_key: Option<Vec<u8>> = base_keys::table
.select(base_keys::base_key)
.filter(base_keys::address.eq(&address))
.filter(base_keys::message_id.eq(&message_id))
.filter(base_keys::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(stored_key.as_ref() == Some(¤t_base_key))
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn delete_base_key(&self, address: &str, message_id: &str) -> Result<()> {
let pool = self.pool.clone();
let device_id = self.device_id;
let address = address.to_string();
let message_id = message_id.to_string();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::delete(
base_keys::table
.filter(base_keys::address.eq(&address))
.filter(base_keys::message_id.eq(&message_id))
.filter(base_keys::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
async fn update_device_list(&self, record: DeviceListRecord) -> Result<()> {
let pool = self.pool.clone();
let device_id = self.device_id;
let devices_json = serde_json::to_string(&record.devices)
.map_err(|e| StoreError::Serialization(Box::new(e)))?;
let now = wacore::time::now_secs() as i32;
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let raw_id_i32 = record.raw_id.map(|r| r as i32);
diesel::insert_into(device_registry::table)
.values((
device_registry::user_id.eq(&record.user),
device_registry::devices_json.eq(&devices_json),
device_registry::timestamp.eq(record.timestamp as i32),
device_registry::phash.eq(&record.phash),
device_registry::device_id.eq(device_id),
device_registry::updated_at.eq(now),
device_registry::raw_id.eq(raw_id_i32),
))
.on_conflict((device_registry::user_id, device_registry::device_id))
.do_update()
.set((
device_registry::devices_json.eq(&devices_json),
device_registry::timestamp.eq(record.timestamp as i32),
device_registry::phash.eq(&record.phash),
device_registry::updated_at.eq(now),
device_registry::raw_id.eq(raw_id_i32),
))
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
async fn update_device_lists(&self, records: Vec<DeviceListRecord>) -> Result<()> {
if records.is_empty() {
return Ok(());
}
let device_id = self.device_id;
let now = wacore::time::now_secs() as i32;
struct PreparedRow {
user: String,
devices_json: String,
timestamp: i32,
phash: Option<String>,
raw_id: Option<i32>,
}
let prepared: Vec<PreparedRow> = records
.into_iter()
.map(|r| {
let devices_json = serde_json::to_string(&r.devices)
.map_err(|e| StoreError::Serialization(Box::new(e)))?;
Ok(PreparedRow {
user: r.user,
devices_json,
timestamp: r.timestamp as i32,
phash: r.phash,
raw_id: r.raw_id.map(|v| v as i32),
})
})
.collect::<Result<Vec<_>>>()?;
let prepared = Arc::new(prepared);
self.with_retry("update_device_lists", move || {
let prepared = Arc::clone(&prepared);
Box::new(move |conn: &mut SqliteConnection| {
conn.transaction::<_, DieselError, _>(|conn| {
for row in prepared.iter() {
diesel::insert_into(device_registry::table)
.values((
device_registry::user_id.eq(&row.user),
device_registry::devices_json.eq(&row.devices_json),
device_registry::timestamp.eq(row.timestamp),
device_registry::phash.eq(&row.phash),
device_registry::device_id.eq(device_id),
device_registry::updated_at.eq(now),
device_registry::raw_id.eq(row.raw_id),
))
.on_conflict((device_registry::user_id, device_registry::device_id))
.do_update()
.set((
device_registry::devices_json.eq(&row.devices_json),
device_registry::timestamp.eq(row.timestamp),
device_registry::phash.eq(&row.phash),
device_registry::updated_at.eq(now),
device_registry::raw_id.eq(row.raw_id),
))
.execute(conn)?;
}
Ok(())
})
})
})
.await
}
async fn get_devices(&self, user: &str) -> Result<Option<DeviceListRecord>> {
let pool = self.pool.clone();
let device_id = self.device_id;
let user = user.to_string();
tokio::task::spawn_blocking(move || -> Result<Option<DeviceListRecord>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let row: Option<(String, String, i32, Option<String>, Option<i32>)> =
device_registry::table
.select((
device_registry::user_id,
device_registry::devices_json,
device_registry::timestamp,
device_registry::phash,
device_registry::raw_id,
))
.filter(device_registry::user_id.eq(&user))
.filter(device_registry::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
match row {
Some((user, devices_json, timestamp, phash, raw_id)) => {
let devices: Vec<DeviceInfo> = serde_json::from_str(&devices_json)
.map_err(|e| StoreError::Serialization(Box::new(e)))?;
Ok(Some(DeviceListRecord {
user,
devices,
timestamp: timestamp as i64,
phash,
raw_id: raw_id.map(|r| r as u32),
}))
}
None => Ok(None),
}
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn delete_devices(&self, user: &str) -> Result<()> {
let pool = self.pool.clone();
let device_id = self.device_id;
let user = user.to_string();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::delete(
device_registry::table
.filter(device_registry::user_id.eq(&user))
.filter(device_registry::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
async fn get_group_metadata(&self, group_jid: &str) -> Result<Option<Vec<u8>>> {
let pool = self.pool.clone();
let device_id = self.device_id;
let group_jid = group_jid.to_string();
tokio::task::spawn_blocking(move || -> Result<Option<Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let row: Option<Vec<u8>> = group_metadata::table
.select(group_metadata::info)
.filter(group_metadata::group_jid.eq(&group_jid))
.filter(group_metadata::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(row)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn put_group_metadata(&self, group_jid: &str, blob: &[u8]) -> Result<()> {
let pool = self.pool.clone();
let device_id = self.device_id;
let group_jid = group_jid.to_string();
let blob = blob.to_vec();
let now = wacore::time::now_secs();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::insert_into(group_metadata::table)
.values((
group_metadata::group_jid.eq(&group_jid),
group_metadata::info.eq(&blob),
group_metadata::device_id.eq(device_id),
group_metadata::updated_at.eq(now),
))
.on_conflict((group_metadata::group_jid, group_metadata::device_id))
.do_update()
.set((
group_metadata::info.eq(&blob),
group_metadata::updated_at.eq(now),
))
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
async fn delete_group_metadata(&self, group_jid: &str) -> Result<()> {
let pool = self.pool.clone();
let device_id = self.device_id;
let group_jid = group_jid.to_string();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::delete(
group_metadata::table
.filter(group_metadata::group_jid.eq(&group_jid))
.filter(group_metadata::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
async fn get_tc_token(&self, jid: &str) -> Result<Option<TcTokenEntry>> {
let pool = self.pool.clone();
let device_id = self.device_id;
let jid = jid.to_string();
tokio::task::spawn_blocking(move || -> Result<Option<TcTokenEntry>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let row: Option<(Vec<u8>, i64, Option<i64>)> = tc_tokens::table
.select((
tc_tokens::token,
tc_tokens::token_timestamp,
tc_tokens::sender_timestamp,
))
.filter(tc_tokens::jid.eq(&jid))
.filter(tc_tokens::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(
row.map(|(token, token_timestamp, sender_timestamp)| TcTokenEntry {
token,
token_timestamp,
sender_timestamp,
}),
)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn put_tc_token(&self, jid: &str, entry: &TcTokenEntry) -> Result<()> {
let pool = self.pool.clone();
let device_id = self.device_id;
let jid = jid.to_string();
let entry = entry.clone();
let now = wacore::time::now_secs();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::insert_into(tc_tokens::table)
.values((
tc_tokens::jid.eq(&jid),
tc_tokens::token.eq(&entry.token),
tc_tokens::token_timestamp.eq(entry.token_timestamp),
tc_tokens::sender_timestamp.eq(entry.sender_timestamp),
tc_tokens::device_id.eq(device_id),
tc_tokens::updated_at.eq(now),
))
.on_conflict((tc_tokens::jid, tc_tokens::device_id))
.do_update()
.set((
tc_tokens::token.eq(&entry.token),
tc_tokens::token_timestamp.eq(entry.token_timestamp),
tc_tokens::sender_timestamp.eq(entry.sender_timestamp),
tc_tokens::updated_at.eq(now),
))
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
async fn delete_tc_token(&self, jid: &str) -> Result<()> {
let pool = self.pool.clone();
let device_id = self.device_id;
let jid = jid.to_string();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::delete(
tc_tokens::table
.filter(tc_tokens::jid.eq(&jid))
.filter(tc_tokens::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
async fn get_all_tc_token_jids(&self) -> Result<Vec<String>> {
let pool = self.pool.clone();
let device_id = self.device_id;
tokio::task::spawn_blocking(move || -> Result<Vec<String>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let jids: Vec<String> = tc_tokens::table
.select(tc_tokens::jid)
.filter(tc_tokens::device_id.eq(device_id))
.load(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(jids)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn delete_expired_tc_tokens(&self, token_cutoff: i64, sender_cutoff: i64) -> Result<u32> {
let pool = self.pool.clone();
let device_id = self.device_id;
tokio::task::spawn_blocking(move || -> Result<u32> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let deleted = diesel::delete(
tc_tokens::table
.filter(
tc_tokens::token
.eq(Vec::<u8>::new())
.or(tc_tokens::token_timestamp.lt(token_cutoff)),
)
.filter(
tc_tokens::sender_timestamp
.is_null()
.or(tc_tokens::sender_timestamp.lt(sender_cutoff)),
)
.filter(tc_tokens::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(deleted as u32)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn store_received_tc_token(
&self,
jid: &str,
token: &[u8],
token_timestamp: i64,
) -> Result<()> {
let device_id = self.device_id;
let jid = jid.to_string();
let token = token.to_vec();
let now = wacore::time::now_secs();
self.with_retry("store_received_tc_token", || {
let jid = jid.clone();
let token = token.clone();
Box::new(move |conn: &mut SqliteConnection| {
conn.immediate_transaction(|conn| -> QueryResult<()> {
let existing: Option<(Vec<u8>, i64)> = tc_tokens::table
.filter(tc_tokens::jid.eq(&jid))
.filter(tc_tokens::device_id.eq(device_id))
.select((tc_tokens::token, tc_tokens::token_timestamp))
.first(conn)
.optional()?;
let write = match &existing {
Some((existing_token, existing_ts)) => {
existing_token.is_empty() || token_timestamp >= *existing_ts
}
None => true,
};
if write {
diesel::insert_into(tc_tokens::table)
.values((
tc_tokens::jid.eq(&jid),
tc_tokens::token.eq(&token),
tc_tokens::token_timestamp.eq(token_timestamp),
tc_tokens::sender_timestamp.eq(None::<i64>),
tc_tokens::device_id.eq(device_id),
tc_tokens::updated_at.eq(now),
))
.on_conflict((tc_tokens::jid, tc_tokens::device_id))
.do_update()
.set((
tc_tokens::token.eq(&token),
tc_tokens::token_timestamp.eq(token_timestamp),
tc_tokens::updated_at.eq(now),
))
.execute(conn)?;
}
Ok(())
})
})
})
.await
}
async fn touch_tc_token_sender_timestamp(
&self,
jid: &str,
sender_timestamp: i64,
) -> Result<()> {
let pool = self.pool.clone();
let device_id = self.device_id;
let jid = jid.to_string();
let now = wacore::time::now_secs();
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
diesel::insert_into(tc_tokens::table)
.values((
tc_tokens::jid.eq(&jid),
tc_tokens::token.eq(Vec::<u8>::new()),
tc_tokens::token_timestamp.eq(sender_timestamp),
tc_tokens::sender_timestamp.eq(Some(sender_timestamp)),
tc_tokens::device_id.eq(device_id),
tc_tokens::updated_at.eq(now),
))
.on_conflict((tc_tokens::jid, tc_tokens::device_id))
.do_update()
.set((
tc_tokens::sender_timestamp.eq(diesel::dsl::sql::<
diesel::sql_types::Nullable<diesel::sql_types::BigInt>,
>(
"MAX(COALESCE(sender_timestamp, "
)
.bind::<diesel::sql_types::BigInt, _>(sender_timestamp)
.sql("), ")
.bind::<diesel::sql_types::BigInt, _>(sender_timestamp)
.sql(")")),
tc_tokens::updated_at.eq(now),
))
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
async fn store_sent_message(
&self,
chat_jid: &str,
message_id: &str,
payload: &[u8],
) -> Result<()> {
let chat_jid = chat_jid.to_string();
let message_id = message_id.to_string();
let payload: Arc<Vec<u8>> = Arc::new(payload.to_vec());
let device_id = self.device_id;
self.with_retry("store_sent_message", || {
let chat_jid = chat_jid.clone();
let message_id = message_id.clone();
let payload = Arc::clone(&payload);
Box::new(move |conn: &mut SqliteConnection| {
diesel::replace_into(sent_messages::table)
.values((
sent_messages::chat_jid.eq(&chat_jid),
sent_messages::message_id.eq(&message_id),
sent_messages::payload.eq(payload.as_slice()),
sent_messages::device_id.eq(device_id),
))
.execute(conn)?;
Ok(())
})
})
.await
}
async fn take_sent_message(&self, chat_jid: &str, message_id: &str) -> Result<Option<Vec<u8>>> {
let chat_jid = chat_jid.to_string();
let message_id = message_id.to_string();
let device_id = self.device_id;
self.with_retry("take_sent_message", || {
let chat_jid = chat_jid.clone();
let message_id = message_id.clone();
Box::new(move |conn: &mut SqliteConnection| {
conn.immediate_transaction(|conn| {
let row: Option<Vec<u8>> = sent_messages::table
.select(sent_messages::payload)
.filter(sent_messages::chat_jid.eq(&chat_jid))
.filter(sent_messages::message_id.eq(&message_id))
.filter(sent_messages::device_id.eq(device_id))
.first(conn)
.optional()?;
if row.is_some() {
diesel::delete(
sent_messages::table
.filter(sent_messages::chat_jid.eq(&chat_jid))
.filter(sent_messages::message_id.eq(&message_id))
.filter(sent_messages::device_id.eq(device_id)),
)
.execute(conn)?;
}
Ok(row)
})
})
})
.await
}
async fn delete_expired_sent_messages(&self, cutoff_timestamp: i64) -> Result<u32> {
let pool = self.pool.clone();
let device_id = self.device_id;
tokio::task::spawn_blocking(move || -> Result<u32> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let deleted = diesel::delete(
sent_messages::table
.filter(sent_messages::created_at.lt(cutoff_timestamp))
.filter(sent_messages::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(deleted as u32)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn store_pending_inbound(
&self,
chat: &str,
sender: &str,
id: &str,
message: &[u8],
) -> Result<()> {
let chat = chat.to_string();
let sender = sender.to_string();
let id = id.to_string();
let message: Arc<Vec<u8>> = Arc::new(message.to_vec());
let device_id = self.device_id;
self.with_retry("store_pending_inbound", || {
let chat = chat.clone();
let sender = sender.clone();
let id = id.clone();
let message = Arc::clone(&message);
Box::new(move |conn: &mut SqliteConnection| {
insert_pending_inbound_row(conn, device_id, &chat, &sender, &id, &message)?;
Ok(())
})
})
.await
}
async fn get_pending_inbound(
&self,
chat: &str,
sender: &str,
id: &str,
) -> Result<Option<Vec<u8>>> {
let chat = chat.to_string();
let sender = sender.to_string();
let id = id.to_string();
let device_id = self.device_id;
self.with_retry("get_pending_inbound", || {
let chat = chat.clone();
let sender = sender.clone();
let id = id.clone();
Box::new(move |conn: &mut SqliteConnection| {
let row: Option<Vec<u8>> = pending_inbound_messages::table
.select(pending_inbound_messages::message)
.filter(pending_inbound_messages::chat.eq(&chat))
.filter(pending_inbound_messages::sender.eq(&sender))
.filter(pending_inbound_messages::id.eq(&id))
.filter(pending_inbound_messages::device_id.eq(device_id))
.first(conn)
.optional()?;
Ok(row)
})
})
.await
}
async fn delete_pending_inbound(&self, chat: &str, sender: &str, id: &str) -> Result<()> {
let chat = chat.to_string();
let sender = sender.to_string();
let id = id.to_string();
let device_id = self.device_id;
self.with_retry("delete_pending_inbound", || {
let chat = chat.clone();
let sender = sender.clone();
let id = id.clone();
Box::new(move |conn: &mut SqliteConnection| {
delete_pending_inbound_row(conn, device_id, &chat, &sender, &id)?;
Ok(())
})
})
.await
}
async fn delete_expired_pending_inbound(&self, cutoff_timestamp: i64) -> Result<u32> {
let pool = self.pool.clone();
let device_id = self.device_id;
tokio::task::spawn_blocking(move || -> Result<u32> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let deleted = diesel::delete(
pending_inbound_messages::table
.filter(pending_inbound_messages::inserted_at.lt(cutoff_timestamp))
.filter(pending_inbound_messages::device_id.eq(device_id)),
)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(deleted as u32)
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))?
}
async fn store_pending_inbound_batch(&self, rows: &[PendingInboundRow<'_>]) -> Result<()> {
if rows.is_empty() {
return Ok(());
}
let rows: Arc<Vec<(String, String, String, Vec<u8>)>> = Arc::new(
rows.iter()
.map(|r| {
(
r.chat.to_string(),
r.sender.to_string(),
r.id.to_string(),
r.message.to_vec(),
)
})
.collect(),
);
let device_id = self.device_id;
self.with_retry("store_pending_inbound_batch", || {
let rows = Arc::clone(&rows);
Box::new(move |conn: &mut SqliteConnection| {
conn.transaction(|conn| {
for (chat, sender, id, message) in rows.iter() {
insert_pending_inbound_row(conn, device_id, chat, sender, id, message)?;
}
Ok(())
})
})
})
.await
}
async fn delete_pending_inbound_batch(&self, keys: &[PendingInboundKey<'_>]) -> Result<()> {
if keys.is_empty() {
return Ok(());
}
let keys: Arc<Vec<(String, String, String)>> = Arc::new(
keys.iter()
.map(|k| (k.chat.to_string(), k.sender.to_string(), k.id.to_string()))
.collect(),
);
let device_id = self.device_id;
self.with_retry("delete_pending_inbound_batch", || {
let keys = Arc::clone(&keys);
Box::new(move |conn: &mut SqliteConnection| {
conn.transaction(|conn| {
for (chat, sender, id) in keys.iter() {
delete_pending_inbound_row(conn, device_id, chat, sender, id)?;
}
Ok(())
})
})
})
.await
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl MsgSecretStore for SqliteStore {
async fn put_msg_secrets(&self, entries: Vec<MsgSecretEntry>) -> Result<usize> {
if entries.is_empty() {
return Ok(0);
}
let device_id = self.device_id;
let entries = Arc::new(entries);
let now = wacore::time::now_secs();
self.with_retry("put_msg_secrets", || {
let entries = Arc::clone(&entries);
Box::new(move |conn: &mut SqliteConnection| {
conn.immediate_transaction(|conn| {
let mut stored = 0usize;
for chunk in entries.chunks(MSG_SECRET_INSERT_CHUNK_SIZE) {
let records: Vec<_> = chunk
.iter()
.map(|entry| {
(
msg_secrets::chat.eq(entry.chat.as_ref()),
msg_secrets::sender.eq(entry.sender.as_ref()),
msg_secrets::msg_id.eq(entry.msg_id.as_ref()),
msg_secrets::secret.eq(entry.secret.as_ref()),
msg_secrets::device_id.eq(device_id),
msg_secrets::created_at.eq(now),
msg_secrets::expires_at.eq(entry.expires_at),
msg_secrets::message_ts.eq(entry.message_ts),
)
})
.collect();
stored += diesel::insert_into(msg_secrets::table)
.values(&records)
.on_conflict((
msg_secrets::chat,
msg_secrets::sender,
msg_secrets::msg_id,
msg_secrets::device_id,
))
.do_update()
.set((
msg_secrets::secret.eq(excluded(msg_secrets::secret)),
msg_secrets::created_at.eq(now),
msg_secrets::expires_at.eq(diesel::dsl::sql::<
diesel::sql_types::BigInt,
>(
"CASE WHEN msg_secrets.expires_at = 0 \
OR excluded.expires_at = 0 THEN 0 \
ELSE MAX(msg_secrets.expires_at, excluded.expires_at) END",
)),
msg_secrets::message_ts.eq(diesel::dsl::sql::<
diesel::sql_types::BigInt,
>(
"MAX(msg_secrets.message_ts, excluded.message_ts)",
)),
))
.execute(conn)?;
}
Ok(stored)
})
})
})
.await
}
async fn get_msg_secret(
&self,
chat: &str,
sender: &str,
msg_id: &str,
) -> Result<Option<Vec<u8>>> {
let pool = self.pool.clone();
let device_id = self.device_id;
let chat = chat.to_string();
let sender = sender.to_string();
let msg_id = msg_id.to_string();
self.with_semaphore(move || -> Result<Option<Vec<u8>>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let row: Option<Vec<u8>> = msg_secrets::table
.select(msg_secrets::secret)
.filter(msg_secrets::chat.eq(&chat))
.filter(msg_secrets::sender.eq(&sender))
.filter(msg_secrets::msg_id.eq(&msg_id))
.filter(msg_secrets::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(row)
})
.await
}
async fn get_msg_secret_with_ts(
&self,
chat: &str,
sender: &str,
msg_id: &str,
) -> Result<Option<(Vec<u8>, i64)>> {
let pool = self.pool.clone();
let device_id = self.device_id;
let chat = chat.to_string();
let sender = sender.to_string();
let msg_id = msg_id.to_string();
self.with_semaphore(move || -> Result<Option<(Vec<u8>, i64)>> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let row: Option<(Vec<u8>, i64)> = msg_secrets::table
.select((msg_secrets::secret, msg_secrets::message_ts))
.filter(msg_secrets::chat.eq(&chat))
.filter(msg_secrets::sender.eq(&sender))
.filter(msg_secrets::msg_id.eq(&msg_id))
.filter(msg_secrets::device_id.eq(device_id))
.first(&mut conn)
.optional()
.map_err(|e| StoreError::Database(Box::new(e)))?;
Ok(row)
})
.await
}
async fn delete_expired_msg_secrets(&self, cutoff_timestamp: i64) -> Result<u32> {
let device_id = self.device_id;
self.with_retry("delete_expired_msg_secrets", || {
Box::new(move |conn: &mut SqliteConnection| {
let deleted = diesel::delete(
msg_secrets::table
.filter(msg_secrets::expires_at.ne(0))
.filter(msg_secrets::expires_at.le(cutoff_timestamp))
.filter(msg_secrets::device_id.eq(device_id)),
)
.execute(conn)?;
Ok(deleted as u32)
})
})
.await
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl DeviceStore for SqliteStore {
async fn save(&self, device: &CoreDevice) -> Result<()> {
SqliteStore::save_device_data_for_device(self, self.device_id, device).await
}
async fn load(&self) -> Result<Option<CoreDevice>> {
SqliteStore::load_device_data_for_device(self, self.device_id).await
}
async fn exists(&self) -> Result<bool> {
SqliteStore::device_exists(self, self.device_id).await
}
async fn create(&self) -> Result<i32> {
SqliteStore::create_new_device(self).await
}
async fn snapshot_db(&self, name: &str, extra_content: Option<&[u8]>) -> Result<()> {
fn sanitize_snapshot_name(name: &str) -> Result<String> {
const MAX_LENGTH: usize = 100;
let sanitized: String = name
.chars()
.map(|c| {
if c.is_ascii_alphanumeric() || c == '_' || c == '-' || c == '.' {
c
} else {
'_'
}
})
.collect();
let sanitized = sanitized
.split('.')
.filter(|part| !part.is_empty() && *part != "..")
.collect::<Vec<_>>()
.join(".");
let sanitized = sanitized.trim_matches(['/', '\\', '.']);
if sanitized.is_empty() {
return Err(StoreError::InvalidConfig(
"Snapshot name cannot be empty after sanitization".to_string(),
));
}
if sanitized.len() > MAX_LENGTH {
return Err(StoreError::InvalidConfig(format!(
"Snapshot name exceeds maximum length of {} characters",
MAX_LENGTH
)));
}
Ok(sanitized.to_string())
}
let sanitized_name = sanitize_snapshot_name(name)?;
let pool = self.pool.clone();
let db_path = self.database_path.clone();
let extra_data = extra_content.map(|b| b.to_vec());
tokio::task::spawn_blocking(move || -> Result<()> {
let mut conn = pool
.get()
.map_err(|e| StoreError::Connection(Box::new(e)))?;
let timestamp = wacore::time::now_secs();
let target_path = format!("{}.snapshot-{}-{}", db_path, timestamp, sanitized_name);
let query = format!("VACUUM INTO '{}'", target_path.replace("'", "''"));
diesel::sql_query(query)
.execute(&mut conn)
.map_err(|e| StoreError::Database(Box::new(e)))?;
if let Some(data) = extra_data {
let extra_path = format!("{}.json", target_path);
std::fs::write(&extra_path, data)?;
}
Ok(())
})
.await
.map_err(|e| StoreError::Database(Box::new(e)))??;
Ok(())
}
async fn resource_report(&self) -> wacore::stats::StorageResourceReport {
let pool = self.pool.clone();
let read_pool = self.reads.as_ref().map(|reads| reads.pool.clone());
tokio::task::spawn_blocking(move || {
let Some(mut conn) = pool.try_get() else {
return wacore::stats::StorageResourceReport::default();
};
let (Some(page_size), Some(page_count), Some(cache_size)) = (
pragma_i64(&mut conn, "page_size"),
pragma_i64(&mut conn, "page_count"),
pragma_i64(&mut conn, "cache_size"),
) else {
return wacore::stats::StorageResourceReport::default();
};
let page_size = page_size.max(0) as u64;
let page_count = page_count.max(0) as u64;
let cache_cap_bytes = if cache_size < 0 {
cache_size.unsigned_abs().saturating_mul(1024)
} else {
(cache_size as u64).saturating_mul(page_size)
};
let db_bytes = page_count.saturating_mul(page_size);
let per_conn_cache = cache_cap_bytes.min(db_bytes);
let open_connections = pool.state().connections.max(1) as u64
+ read_pool.map_or(0, |reads| reads.state().connections as u64);
wacore::stats::StorageResourceReport {
memory_bytes: Some(per_conn_cache.saturating_mul(open_connections)),
pages: Some(page_count),
..Default::default()
}
})
.await
.unwrap_or_default()
}
}
fn pragma_i64(conn: &mut SqliteConnection, pragma: &str) -> Option<i64> {
if pragma.is_empty()
|| !pragma
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'_')
{
debug_assert!(false, "pragma_i64 requires an identifier, got {pragma:?}");
return None;
}
#[derive(diesel::QueryableByName)]
struct Row {
#[diesel(sql_type = diesel::sql_types::BigInt)]
value: i64,
}
let sql = format!("SELECT {pragma} AS value FROM pragma_{pragma}()");
diesel::sql_query(sql)
.get_result::<Row>(conn)
.ok()
.map(|r| r.value)
}
#[cfg(test)]
mod tests {
use super::*;
async fn create_test_store() -> SqliteStore {
use portable_atomic::AtomicU64;
use std::sync::atomic::Ordering;
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let db_name = format!(
"file:memdb_test_{}_{}?mode=memory&cache=shared",
std::process::id(),
id
);
SqliteStore::new(&db_name)
.await
.expect("Failed to create test store")
}
#[tokio::test]
async fn with_config_custom_tuning_builds_and_operates() {
use portable_atomic::AtomicU64;
use std::sync::atomic::Ordering;
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let db_name = format!(
"file:memdb_cfg_{}_{}?mode=memory&cache=shared",
std::process::id(),
id
);
let def = SqliteStoreConfig::default();
assert_eq!(def.pool_size, 1);
assert_eq!(def.cache_size_kib, 512);
let config = SqliteStoreConfig {
pool_size: 2,
read_pool_size: 0,
cache_size_kib: 4096,
mmap_size: None,
busy_timeout: Duration::from_secs(7),
synchronous: Synchronous::Full,
thread_pool: Some(Arc::new(
scheduled_thread_pool::ScheduledThreadPool::builder()
.num_threads(1)
.build(),
)),
connection_init: None,
};
let store = SqliteStore::with_config(&db_name, config)
.await
.expect("custom-config store");
let mac = AppStateMutationMAC {
index_mac: vec![1u8; 32],
value_mac: vec![2u8; 32],
};
store
.put_app_state_mutation_macs_for_device("c", 1, std::slice::from_ref(&mac), 1)
.await
.unwrap();
let got = store
.get_app_state_mutation_mac_for_device("c", &mac.index_mac, 1)
.await
.unwrap();
assert_eq!(got, Some(mac.value_mac));
#[derive(diesel::QueryableByName)]
struct CacheSync {
#[diesel(sql_type = diesel::sql_types::BigInt)]
cache: i64,
#[diesel(sql_type = diesel::sql_types::BigInt)]
sync: i64,
}
#[derive(diesel::QueryableByName)]
struct Busy {
#[diesel(sql_type = diesel::sql_types::BigInt)]
timeout: i64,
}
let mut conn = store.pool.get().unwrap();
let cs: CacheSync = diesel::sql_query(
"SELECT cs.cache_size AS cache, sy.synchronous AS sync \
FROM pragma_cache_size cs, pragma_synchronous sy",
)
.get_result(&mut conn)
.unwrap();
let busy: Busy = diesel::sql_query("PRAGMA busy_timeout")
.get_result(&mut conn)
.unwrap();
assert_eq!(cs.cache, -4096, "cache_size_kib applied as negative KiB");
assert_eq!(cs.sync, 2, "synchronous = FULL");
assert_eq!(busy.timeout, 7000, "busy_timeout = 7s");
}
#[tokio::test]
async fn connection_init_runs_before_pragmas_and_migrations() {
use portable_atomic::AtomicU64;
use std::sync::atomic::{AtomicBool, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let db_name = format!(
"file:memdb_init_{}_{}?mode=memory&cache=shared",
std::process::id(),
id
);
let calls = Arc::new(AtomicU64::new(0));
let saw_migrations_table = Arc::new(AtomicBool::new(false));
let saw_store_pragmas = Arc::new(AtomicBool::new(false));
#[derive(diesel::QueryableByName)]
struct Count {
#[diesel(sql_type = diesel::sql_types::BigInt)]
n: i64,
}
let config = {
let calls = calls.clone();
let saw_migrations_table = saw_migrations_table.clone();
let saw_store_pragmas = saw_store_pragmas.clone();
SqliteStoreConfig::default().with_connection_init(move |conn| {
calls.fetch_add(1, Ordering::Relaxed);
let migrated: Count = diesel::sql_query(
"SELECT count(*) AS n FROM sqlite_master \
WHERE name = '__diesel_schema_migrations'",
)
.get_result(conn)?;
if migrated.n > 0 {
saw_migrations_table.store(true, Ordering::Relaxed);
}
let busy: Count = diesel::sql_query("SELECT timeout AS n FROM pragma_busy_timeout")
.get_result(conn)?;
if busy.n != 0 {
saw_store_pragmas.store(true, Ordering::Relaxed);
}
Ok(())
})
};
let store = SqliteStore::with_config(&db_name, config)
.await
.expect("store with connection_init");
assert!(calls.load(Ordering::Relaxed) >= 1, "hook ran");
assert!(
!saw_migrations_table.load(Ordering::Relaxed),
"hook ran before migrations on the first connection"
);
assert!(
!saw_store_pragmas.load(Ordering::Relaxed),
"hook ran before the store's own pragmas"
);
let mac = AppStateMutationMAC {
index_mac: vec![3u8; 32],
value_mac: vec![4u8; 32],
};
store
.put_app_state_mutation_macs_for_device("ci", 1, std::slice::from_ref(&mac), 1)
.await
.unwrap();
assert_eq!(
store
.get_app_state_mutation_mac_for_device("ci", &mac.index_mac, 1)
.await
.unwrap(),
Some(mac.value_mac)
);
}
#[test]
fn connection_init_error_rejects_connection_before_pragmas() {
let mut conn = SqliteConnection::establish(":memory:").expect("raw connection");
let options = ConnectionOptions {
cache_size_kib: 512,
mmap_size: None,
busy_timeout_ms: 30_000,
synchronous: Synchronous::Normal,
connection_init: Some(Arc::new(|_conn: &mut SqliteConnection| {
Err("wrong key".into())
})),
query_only: false,
};
use diesel::r2d2::CustomizeConnection;
let err = options
.on_acquire(&mut conn)
.expect_err("hook error surfaces");
assert!(err.to_string().contains("wrong key"));
#[derive(diesel::QueryableByName)]
struct Busy {
#[diesel(sql_type = diesel::sql_types::BigInt)]
timeout: i64,
}
let busy: Busy = diesel::sql_query("PRAGMA busy_timeout")
.get_result(&mut conn)
.unwrap();
assert_eq!(busy.timeout, 0);
}
#[tokio::test]
async fn batch_mutation_macs_matches_per_item() {
let store = create_test_store().await;
let name = "regular";
let device_id = 1;
let macs: Vec<AppStateMutationMAC> = (0..25u8)
.map(|i| {
let mut index_mac = vec![0u8; 32];
index_mac[0] = i;
AppStateMutationMAC {
index_mac,
value_mac: vec![i; 32],
}
})
.collect();
store
.put_app_state_mutation_macs_for_device(name, 1, &macs, device_id)
.await
.unwrap();
let mut index_macs: Vec<[u8; 32]> = macs
.iter()
.map(|m| m.index_mac.as_slice().try_into().unwrap())
.collect();
index_macs.push([0xFF; 32]);
let batch = store
.get_app_state_mutation_macs_batch_for_device(name, &index_macs, device_id)
.await
.unwrap();
assert_eq!(batch.len(), macs.len());
assert!(!batch.contains_key(&[0xFF; 32]));
for m in &macs {
let key: [u8; 32] = m.index_mac.as_slice().try_into().unwrap();
let per_item = store
.get_app_state_mutation_mac_for_device(name, &m.index_mac, device_id)
.await
.unwrap();
assert_eq!(per_item.as_ref(), batch.get(&key));
assert_eq!(batch.get(&key), Some(&m.value_mac));
}
let empty = store
.get_app_state_mutation_macs_batch_for_device(name, &[], device_id)
.await
.unwrap();
assert!(empty.is_empty());
}
#[tokio::test]
async fn clear_mutation_macs_wipes_only_named_collection() {
let store = create_test_store().await;
let mac = |i: u8| AppStateMutationMAC {
index_mac: vec![i; 32],
value_mac: vec![i; 32],
};
store
.put_mutation_macs("regular", 1, &[mac(1)])
.await
.unwrap();
store
.put_mutation_macs("critical", 1, &[mac(2)])
.await
.unwrap();
store.clear_mutation_macs("regular").await.unwrap();
assert!(
store
.get_mutation_mac("regular", &[1; 32])
.await
.unwrap()
.is_none()
);
assert!(
store
.get_mutation_mac("critical", &[2; 32])
.await
.unwrap()
.is_some()
);
}
#[tokio::test]
async fn put_signal_batches_persist_and_upsert() {
use std::sync::Arc;
let store = create_test_store().await;
let sessions: Vec<(Arc<str>, Bytes)> = (0..5u8)
.map(|i| {
(
Arc::from(format!("user{i}@s.whatsapp.net").as_str()),
Bytes::from(vec![i; 8]),
)
})
.collect();
store.put_sessions_batch(&sessions).await.unwrap();
for (addr, bytes) in &sessions {
assert_eq!(
store.get_session(addr).await.unwrap().as_deref(),
Some(bytes.as_ref())
);
}
let identities: Vec<(Arc<str>, [u8; 32])> = (0..5u8)
.map(|i| {
(
Arc::from(format!("user{i}@s.whatsapp.net").as_str()),
[i; 32],
)
})
.collect();
store.put_identities_batch(&identities).await.unwrap();
for (addr, key) in &identities {
assert_eq!(store.load_identity(addr).await.unwrap(), Some(*key));
}
let sender_keys: Vec<(Arc<str>, Bytes)> = (0..5u8)
.map(|i| {
(
Arc::from(format!("g@g.us::user{i}").as_str()),
Bytes::from(vec![i; 16]),
)
})
.collect();
store.put_sender_keys_batch(&sender_keys).await.unwrap();
for (addr, bytes) in &sender_keys {
assert_eq!(
store.get_sender_key(addr).await.unwrap().as_deref(),
Some(bytes.as_ref())
);
}
let updated: Vec<(Arc<str>, Bytes)> = sessions
.iter()
.map(|(addr, _)| (addr.clone(), Bytes::from(vec![0xAA; 8])))
.collect();
store.put_sessions_batch(&updated).await.unwrap();
for (addr, _) in &sessions {
assert_eq!(
store.get_session(addr).await.unwrap().as_deref(),
Some([0xAA; 8].as_slice())
);
}
let dup: Arc<str> = Arc::from("dup@s.whatsapp.net");
store
.put_sessions_batch(&[
(dup.clone(), Bytes::from(vec![1u8; 4])),
(dup.clone(), Bytes::from(vec![2u8; 4])),
])
.await
.unwrap();
assert_eq!(
store.get_session(&dup).await.unwrap().as_deref(),
Some([2u8; 4].as_slice())
);
store.put_sessions_batch(&[]).await.unwrap();
store.put_identities_batch(&[]).await.unwrap();
store.put_sender_keys_batch(&[]).await.unwrap();
}
#[test]
fn test_parse_database_path_regular_path() {
let path = "/var/lib/whatsapp/database.db";
let result = parse_database_path(path).unwrap();
assert_eq!(result, "/var/lib/whatsapp/database.db");
}
#[test]
fn test_parse_database_path_with_sqlite_prefix() {
let path = "sqlite:///var/lib/whatsapp/database.db";
let result = parse_database_path(path).unwrap();
assert_eq!(result, "/var/lib/whatsapp/database.db");
}
#[test]
fn test_parse_database_path_with_query_params() {
let path = "file:database.db?mode=memory&cache=shared";
let result = parse_database_path(path).unwrap();
assert_eq!(result, "file:database.db");
}
#[test]
fn test_parse_database_path_with_fragment() {
let path = "file:database.db#fragment";
let result = parse_database_path(path).unwrap();
assert_eq!(result, "file:database.db");
}
#[test]
fn test_parse_database_path_with_both_query_and_fragment() {
let path = "sqlite:///var/lib/database.db?mode=ro#backup";
let result = parse_database_path(path).unwrap();
assert_eq!(result, "/var/lib/database.db");
}
#[test]
fn test_parse_database_path_in_memory_rejected() {
let result = parse_database_path(":memory:");
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("not supported"));
}
#[test]
fn test_parse_database_path_in_memory_with_query_rejected() {
let result = parse_database_path(":memory:?cache=shared");
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("not supported"));
}
#[tokio::test]
async fn test_device_registry_save_and_get() {
let store = create_test_store().await;
let record = DeviceListRecord {
user: "1234567890".to_string(),
devices: vec![DeviceInfo::new(0, None), DeviceInfo::new(1, Some(42))],
timestamp: 1234567890,
phash: Some("2:abcdef".to_string()),
raw_id: None,
};
store.update_device_list(record).await.expect("save failed");
let loaded = store
.get_devices("1234567890")
.await
.expect("get failed")
.expect("record should exist");
assert_eq!(loaded.user, "1234567890");
assert_eq!(loaded.devices.len(), 2);
assert_eq!(loaded.devices[0].device_id, 0);
assert_eq!(loaded.devices[1].device_id, 1);
assert_eq!(loaded.devices[1].key_index, Some(42));
assert_eq!(loaded.phash, Some("2:abcdef".to_string()));
}
#[tokio::test]
async fn test_device_registry_update_existing() {
let store = create_test_store().await;
let record1 = DeviceListRecord {
user: "1234567890".to_string(),
devices: vec![DeviceInfo::new(0, None)],
timestamp: 1000,
phash: Some("2:old".to_string()),
raw_id: None,
};
store
.update_device_list(record1)
.await
.expect("save1 failed");
let record2 = DeviceListRecord {
user: "1234567890".to_string(),
devices: vec![DeviceInfo::new(0, None), DeviceInfo::new(2, None)],
timestamp: 2000,
phash: Some("2:new".to_string()),
raw_id: None,
};
store
.update_device_list(record2)
.await
.expect("save2 failed");
let loaded = store
.get_devices("1234567890")
.await
.expect("get failed")
.expect("record should exist");
assert_eq!(loaded.devices.len(), 2);
assert_eq!(loaded.phash, Some("2:new".to_string()));
}
#[tokio::test]
async fn test_device_registry_get_nonexistent() {
let store = create_test_store().await;
let result = store.get_devices("nonexistent").await.expect("get failed");
assert!(result.is_none());
}
#[tokio::test]
async fn test_sender_key_devices_set_and_get() {
let store = create_test_store().await;
let group = "group123@g.us";
store
.set_sender_key_status(group, &[("user1:5@lid", true), ("user2:3@lid", false)])
.await
.expect("set failed");
let devices = store
.get_sender_key_devices(group)
.await
.expect("get failed");
assert_eq!(devices.len(), 2);
assert!(devices.contains(&("user1:5@lid".to_string(), true)));
assert!(devices.contains(&("user2:3@lid".to_string(), false)));
}
#[tokio::test]
async fn test_sender_key_devices_upsert_overwrites() {
let store = create_test_store().await;
let group = "group123@g.us";
store
.set_sender_key_status(group, &[("user1:5@lid", false)])
.await
.expect("set failed");
store
.set_sender_key_status(group, &[("user1:5@lid", true)])
.await
.expect("set failed");
let devices = store
.get_sender_key_devices(group)
.await
.expect("get failed");
assert_eq!(devices.len(), 1);
assert_eq!(devices[0], ("user1:5@lid".to_string(), true));
}
#[tokio::test]
async fn test_sender_key_devices_clear() {
let store = create_test_store().await;
let group = "group123@g.us";
store
.set_sender_key_status(group, &[("user1:5@lid", true), ("user2:3@lid", true)])
.await
.expect("set failed");
store
.clear_sender_key_devices(group)
.await
.expect("clear failed");
let devices = store
.get_sender_key_devices(group)
.await
.expect("get failed");
assert!(devices.is_empty());
}
#[tokio::test]
async fn test_tc_token_put_and_get() {
let store = create_test_store().await;
let entry = TcTokenEntry {
token: vec![1, 2, 3, 4, 5],
token_timestamp: 1707000000,
sender_timestamp: Some(1707000100),
};
store
.put_tc_token("user@lid", &entry)
.await
.expect("put failed");
let loaded = store
.get_tc_token("user@lid")
.await
.expect("get failed")
.expect("should exist");
assert_eq!(loaded.token, vec![1, 2, 3, 4, 5]);
assert_eq!(loaded.token_timestamp, 1707000000);
assert_eq!(loaded.sender_timestamp, Some(1707000100));
}
#[tokio::test]
async fn test_tc_token_upsert() {
let store = create_test_store().await;
let entry1 = TcTokenEntry {
token: vec![1, 2, 3],
token_timestamp: 1000,
sender_timestamp: None,
};
store.put_tc_token("user@lid", &entry1).await.unwrap();
let entry2 = TcTokenEntry {
token: vec![4, 5, 6],
token_timestamp: 2000,
sender_timestamp: Some(1500),
};
store.put_tc_token("user@lid", &entry2).await.unwrap();
let loaded = store.get_tc_token("user@lid").await.unwrap().unwrap();
assert_eq!(loaded.token, vec![4, 5, 6]);
assert_eq!(loaded.token_timestamp, 2000);
assert_eq!(loaded.sender_timestamp, Some(1500));
}
#[tokio::test]
async fn test_tc_token_delete() {
let store = create_test_store().await;
let entry = TcTokenEntry {
token: vec![1, 2, 3],
token_timestamp: 1000,
sender_timestamp: None,
};
store.put_tc_token("user@lid", &entry).await.unwrap();
store.delete_tc_token("user@lid").await.unwrap();
let result = store.get_tc_token("user@lid").await.unwrap();
assert!(result.is_none());
}
#[tokio::test]
async fn test_touch_and_store_received_preserve_each_others_field() {
let store = create_test_store().await;
store
.touch_tc_token_sender_timestamp("user@lid", 5000)
.await
.unwrap();
store
.store_received_tc_token("user@lid", &[7, 8, 9], 4000)
.await
.unwrap();
let a = store.get_tc_token("user@lid").await.unwrap().unwrap();
assert_eq!(a.token, vec![7, 8, 9]);
assert_eq!(a.token_timestamp, 4000);
assert_eq!(a.sender_timestamp, Some(5000));
store
.touch_tc_token_sender_timestamp("user@lid", 6000)
.await
.unwrap();
let b = store.get_tc_token("user@lid").await.unwrap().unwrap();
assert_eq!(b.token, vec![7, 8, 9], "touch must keep the real token");
assert_eq!(b.sender_timestamp, Some(6000));
store
.touch_tc_token_sender_timestamp("user@lid", 1000)
.await
.unwrap();
let c = store.get_tc_token("user@lid").await.unwrap().unwrap();
assert_eq!(c.sender_timestamp, Some(6000), "touch is advance-only");
}
#[tokio::test]
async fn store_received_tc_token_is_newer_wins() {
let store = create_test_store().await;
store
.store_received_tc_token("c@lid", &[1, 1, 1], 5000)
.await
.unwrap();
store
.store_received_tc_token("c@lid", &[2, 2, 2], 3000)
.await
.unwrap();
let e = store.get_tc_token("c@lid").await.unwrap().unwrap();
assert_eq!(e.token, vec![1, 1, 1], "older write must not overwrite");
assert_eq!(e.token_timestamp, 5000);
store
.store_received_tc_token("c@lid", &[3, 3, 3], 7000)
.await
.unwrap();
let e = store.get_tc_token("c@lid").await.unwrap().unwrap();
assert_eq!(e.token, vec![3, 3, 3]);
assert_eq!(e.token_timestamp, 7000);
store
.touch_tc_token_sender_timestamp("p@lid", 9000)
.await
.unwrap();
store
.store_received_tc_token("p@lid", &[4, 4, 4], 6000)
.await
.unwrap();
let e = store.get_tc_token("p@lid").await.unwrap().unwrap();
assert_eq!(e.token, vec![4, 4, 4], "placeholder must accept real token");
assert_eq!(e.token_timestamp, 6000);
assert_eq!(e.sender_timestamp, Some(9000), "sender bucket preserved");
}
#[tokio::test]
async fn test_delete_expired_two_window_pruning() {
let store = create_test_store().await;
store
.touch_tc_token_sender_timestamp("recent_ph@lid", 2500)
.await
.unwrap();
store
.touch_tc_token_sender_timestamp("stale_ph@lid", 100)
.await
.unwrap();
store
.put_tc_token(
"expired_live_sender@lid",
&TcTokenEntry {
token: vec![1],
token_timestamp: 1,
sender_timestamp: Some(2500),
},
)
.await
.unwrap();
store
.put_tc_token(
"orphan_expired@lid",
&TcTokenEntry {
token: vec![2],
token_timestamp: 1,
sender_timestamp: None,
},
)
.await
.unwrap();
let removed = store.delete_expired_tc_tokens(1000, 2000).await.unwrap();
assert_eq!(removed, 2);
assert!(store.get_tc_token("recent_ph@lid").await.unwrap().is_some());
assert!(store.get_tc_token("stale_ph@lid").await.unwrap().is_none());
assert!(
store
.get_tc_token("expired_live_sender@lid")
.await
.unwrap()
.is_some()
);
assert!(
store
.get_tc_token("orphan_expired@lid")
.await
.unwrap()
.is_none()
);
}
#[tokio::test]
async fn test_tc_token_get_all_jids() {
let store = create_test_store().await;
let entry = TcTokenEntry {
token: vec![1],
token_timestamp: 1000,
sender_timestamp: None,
};
store.put_tc_token("user1@lid", &entry).await.unwrap();
store.put_tc_token("user2@lid", &entry).await.unwrap();
store.put_tc_token("user3@lid", &entry).await.unwrap();
let mut jids = store.get_all_tc_token_jids().await.unwrap();
jids.sort();
assert_eq!(jids, vec!["user1@lid", "user2@lid", "user3@lid"]);
}
#[tokio::test]
async fn test_tc_token_delete_expired() {
let store = create_test_store().await;
let old = TcTokenEntry {
token: vec![1],
token_timestamp: 1000,
sender_timestamp: None,
};
let recent = TcTokenEntry {
token: vec![2],
token_timestamp: 5000,
sender_timestamp: None,
};
store.put_tc_token("old@lid", &old).await.unwrap();
store.put_tc_token("recent@lid", &recent).await.unwrap();
let deleted = store.delete_expired_tc_tokens(3000, 3000).await.unwrap();
assert_eq!(deleted, 1);
assert!(store.get_tc_token("old@lid").await.unwrap().is_none());
assert!(store.get_tc_token("recent@lid").await.unwrap().is_some());
}
#[tokio::test]
async fn test_tc_token_get_nonexistent() {
let store = create_test_store().await;
let result = store.get_tc_token("nonexistent@lid").await.unwrap();
assert!(result.is_none());
}
#[tokio::test]
async fn test_sender_key_devices_different_groups() {
let store = create_test_store().await;
let group1 = "group1@g.us";
let group2 = "group2@g.us";
store
.set_sender_key_status(group1, &[("user:5@lid", true)])
.await
.expect("set failed");
let g1 = store.get_sender_key_devices(group1).await.unwrap();
assert_eq!(g1.len(), 1);
let g2 = store.get_sender_key_devices(group2).await.unwrap();
assert!(g2.is_empty());
}
#[tokio::test]
async fn test_create_new_device_uses_configured_device_id() {
use portable_atomic::AtomicU64;
use std::sync::atomic::Ordering;
static COUNTER: AtomicU64 = AtomicU64::new(100);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let db_name = format!(
"file:memdb_devid_{}_{}?mode=memory&cache=shared",
std::process::id(),
id
);
let device_id = 42;
let store = SqliteStore::new_for_device(&db_name, device_id)
.await
.expect("Failed to create test store");
assert!(!store.device_exists(device_id).await.unwrap());
let returned_id = store.create_new_device().await.unwrap();
assert_eq!(returned_id, device_id);
assert!(store.device_exists(device_id).await.unwrap());
if device_id != 1 {
assert!(!store.device_exists(1).await.unwrap());
}
let loaded = store.load_device_data_for_device(device_id).await.unwrap();
assert!(
loaded.is_some(),
"device data should be loadable by configured id"
);
}
#[tokio::test]
async fn mark_prekeys_uploaded_never_resurrects_deleted_rows() {
let store = create_test_store().await;
store
.store_prekey(1, b"record-1", false)
.await
.expect("store");
store
.store_prekey(2, b"record-2", false)
.await
.expect("store");
store.remove_prekey(1).await.expect("consume");
store
.mark_prekeys_uploaded(&[1, 2])
.await
.expect("mark uploaded");
let gone = store.load_prekey(1).await.expect("load");
assert!(gone.is_none(), "consumed key must not be resurrected");
let live = store.load_prekey(2).await.expect("load");
assert!(live.is_some(), "live key still present");
}
#[tokio::test]
async fn test_prekey_watermarks_survive_save_load_roundtrip() {
use portable_atomic::AtomicU64;
use std::sync::atomic::Ordering;
static COUNTER: AtomicU64 = AtomicU64::new(300);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let db_name = format!(
"file:memdb_pkwatermark_{}_{}?mode=memory&cache=shared",
std::process::id(),
id
);
let device_id = 9;
let _writer = SqliteStore::new_for_device(&db_name, device_id)
.await
.expect("create store");
_writer.create_new_device().await.expect("create device");
let mut device = _writer
.load_device_data_for_device(device_id)
.await
.expect("load")
.expect("device should exist after create");
assert_eq!(
device.first_unupload_pre_key_id, 0,
"fresh device starts with the watermark unset"
);
device.next_pre_key_id = 913;
device.first_unupload_pre_key_id = 101;
_writer
.save_device_data_for_device(device_id, &device)
.await
.expect("save with watermarks");
let store = SqliteStore::new_for_device(&db_name, device_id)
.await
.expect("reopen store");
let loaded = store
.load_device_data_for_device(device_id)
.await
.expect("load")
.expect("device should exist after reopen");
assert_eq!(loaded.next_pre_key_id, 913);
assert_eq!(
loaded.first_unupload_pre_key_id, 101,
"first_unupload_pre_key_id must survive a save/load roundtrip"
);
}
#[tokio::test]
async fn test_server_cert_chain_survives_save_load_roundtrip() {
use portable_atomic::AtomicU64;
use std::sync::atomic::Ordering;
use wacore::store::device::{CachedNoiseCert, CachedServerCertChain};
static COUNTER: AtomicU64 = AtomicU64::new(200);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let db_name = format!(
"file:memdb_certchain_{}_{}?mode=memory&cache=shared",
std::process::id(),
id
);
let device_id = 7;
let chain = CachedServerCertChain {
intermediate: CachedNoiseCert {
key: [0xAB; 32],
not_before: 1_700_000_000,
not_after: 1_900_000_000,
},
leaf: CachedNoiseCert {
key: [0xCD; 32],
not_before: 1_700_000_500,
not_after: 1_899_999_500,
},
};
let _writer = SqliteStore::new_for_device(&db_name, device_id)
.await
.expect("create store");
_writer.create_new_device().await.expect("create device");
let mut device = _writer
.load_device_data_for_device(device_id)
.await
.expect("load")
.expect("device should exist after create");
device.server_cert_chain = Some(chain.clone());
_writer
.save_device_data_for_device(device_id, &device)
.await
.expect("save with cert chain");
let store = SqliteStore::new_for_device(&db_name, device_id)
.await
.expect("reopen store");
let loaded = store
.load_device_data_for_device(device_id)
.await
.expect("load")
.expect("device should exist after reopen");
assert_eq!(
loaded.server_cert_chain.as_ref(),
Some(&chain),
"server_cert_chain must survive a save/load roundtrip"
);
let mut device = loaded;
device.server_cert_chain = None;
store
.save_device_data_for_device(device_id, &device)
.await
.expect("save with cleared cert chain");
let reloaded = store
.load_device_data_for_device(device_id)
.await
.expect("reload")
.expect("device should exist");
assert!(
reloaded.server_cert_chain.is_none(),
"cleared chain must round-trip as None"
);
}
#[tokio::test]
async fn legacy_bincode_blobs_self_heal_then_overwrite() {
use diesel::{ExpressionMethods, RunQueryDsl, sql_query};
use wacore::appstate::hash::HashState;
use wacore::store::traits::AppStateSyncKey;
let legacy_sync_key = {
let mut v = vec![0x20u8]; v.extend([0x11u8; 32]);
v.extend([0x04, 0xaa, 0xbb, 0xcc, 0xdd, 0xfc, 0x00, 0xe2, 0xa7, 0xca]);
v
};
let legacy_hash_state = {
let mut v = vec![0x07u8]; v.push(0xde);
v.push(0xad);
v.extend([0u8; 125]);
v.push(0xbe);
v.push(0x00); v
};
let store = create_test_store().await;
let device_id = store.device_id;
let key_id = b"legacy-key".to_vec();
{
let kid = key_id.clone();
let blob = legacy_sync_key.clone();
store
.with_retry("insert_legacy_key", move || {
let kid = kid.clone();
let blob = blob.clone();
Box::new(move |conn| {
diesel::insert_into(app_state_keys::table)
.values((
app_state_keys::key_id.eq(kid),
app_state_keys::key_data.eq(blob),
app_state_keys::device_id.eq(device_id),
))
.execute(conn)
.map(|_| ())
})
})
.await
.expect("insert legacy key row");
}
let name = "critical_block";
{
let blob = legacy_hash_state.clone();
store
.with_retry("insert_legacy_version", move || {
let blob = blob.clone();
Box::new(move |conn| {
diesel::insert_into(app_state_versions::table)
.values((
app_state_versions::name.eq(name),
app_state_versions::state_data.eq(blob),
app_state_versions::device_id.eq(device_id),
))
.execute(conn)
.map(|_| ())
})
})
.await
.expect("insert legacy version row");
}
assert!(
store
.get_app_state_sync_key_for_device(&key_id, device_id)
.await
.expect("legacy sync-key blob must not surface a decode error")
.is_none(),
"a legacy bincode sync-key row must read back as absent"
);
assert_eq!(
store
.get_app_state_version_for_device(name, device_id)
.await
.expect("legacy version blob must not surface a decode error")
.version,
0,
"a legacy bincode version row must reset to default (re-sync from 0)"
);
store
.set_app_state_sync_key_for_device(
&key_id,
AppStateSyncKey {
key_data: vec![7u8; 32],
fingerprint: vec![1, 2, 3],
timestamp: 99,
},
device_id,
)
.await
.expect("overwrite key");
let healed_key = store
.get_app_state_sync_key_for_device(&key_id, device_id)
.await
.expect("get key")
.expect("re-shared key must persist over the legacy row");
assert_eq!(healed_key.key_data, vec![7u8; 32]);
assert_eq!(healed_key.timestamp, 99);
store
.set_app_state_version_for_device(
name,
HashState {
version: 5,
..HashState::default()
},
device_id,
)
.await
.expect("overwrite version");
assert_eq!(
store
.get_app_state_version_for_device(name, device_id)
.await
.expect("get version")
.version,
5,
"a re-synced version must persist over the legacy row"
);
store
.with_retry("corrupt_key", || {
Box::new(|conn| {
sql_query("UPDATE app_state_keys SET key_data = X'00ff00ff'")
.execute(conn)
.map(|_| ())
})
})
.await
.expect("corrupt key blob");
assert!(
store
.get_app_state_sync_key_for_device(&key_id, device_id)
.await
.expect("corrupt key blob must not error")
.is_none(),
"an arbitrarily corrupt sync-key blob must also read back as absent"
);
}
#[tokio::test]
async fn latest_sync_key_skips_undecodable_rows() {
use diesel::{ExpressionMethods, RunQueryDsl};
use wacore::store::traits::AppStateSyncKey;
let legacy_blob = {
let mut v = vec![0x20u8];
v.extend([0x11u8; 32]);
v.extend([0x04, 0xaa, 0xbb, 0xcc, 0xdd, 0xfc, 0x00, 0xe2, 0xa7, 0xca]);
v
};
let store = create_test_store().await;
let device_id = store.device_id;
let good_id = b"key-aaa".to_vec();
store
.set_app_state_sync_key_for_device(
&good_id,
AppStateSyncKey {
key_data: vec![7u8; 32],
fingerprint: vec![1],
timestamp: 1,
},
device_id,
)
.await
.unwrap();
let bad_id = b"key-zzz".to_vec();
{
let bid = bad_id.clone();
let blob = legacy_blob.clone();
store
.with_retry("insert_stale_key", move || {
let bid = bid.clone();
let blob = blob.clone();
Box::new(move |conn| {
diesel::insert_into(app_state_keys::table)
.values((
app_state_keys::key_id.eq(bid),
app_state_keys::key_data.eq(blob),
app_state_keys::device_id.eq(device_id),
))
.execute(conn)
.map(|_| ())
})
})
.await
.unwrap();
}
assert_eq!(
store
.get_latest_app_state_sync_key_id_for_device(device_id)
.await
.unwrap(),
Some(good_id),
"latest-key selection must skip undecodable bincode rows"
);
}
#[tokio::test]
async fn group_metadata_round_trip_sqlite() {
use wacore::store::traits::ProtocolStore;
let store = create_test_store().await;
let jid = "120363000000000001@g.us";
assert!(store.get_group_metadata(jid).await.unwrap().is_none());
store.put_group_metadata(jid, b"blob-v1").await.unwrap();
assert_eq!(
store.get_group_metadata(jid).await.unwrap().as_deref(),
Some(&b"blob-v1"[..])
);
store.put_group_metadata(jid, b"blob-v2").await.unwrap();
assert_eq!(
store.get_group_metadata(jid).await.unwrap().as_deref(),
Some(&b"blob-v2"[..])
);
store.delete_group_metadata(jid).await.unwrap();
assert!(store.get_group_metadata(jid).await.unwrap().is_none());
}
#[tokio::test]
async fn msg_secret_round_trip_sqlite() {
let store = create_test_store().await;
let secret = [0xABu8; 32];
store
.put_msg_secret("12345@s.whatsapp.net", "9999@lid", "MID1", &secret)
.await
.expect("put");
let got = store
.get_msg_secret("12345@s.whatsapp.net", "9999@lid", "MID1")
.await
.expect("get")
.expect("must exist");
assert_eq!(got, secret.to_vec());
}
#[tokio::test]
async fn msg_secret_miss_returns_none_sqlite() {
let store = create_test_store().await;
assert!(
store
.get_msg_secret("any@s.whatsapp.net", "any@lid", "NOPE")
.await
.expect("get")
.is_none()
);
}
#[tokio::test]
async fn msg_secret_upsert_replaces_secret() {
let store = create_test_store().await;
store
.put_msg_secret("c", "s", "M", &[1u8; 32])
.await
.expect("put 1");
store
.put_msg_secret("c", "s", "M", &[9u8; 32])
.await
.expect("put 2");
let got = store.get_msg_secret("c", "s", "M").await.unwrap().unwrap();
assert_eq!(got, vec![9u8; 32], "ON CONFLICT must overwrite");
}
#[tokio::test]
async fn msg_secret_scoped_by_three_columns() {
let store = create_test_store().await;
store
.put_msg_secret("c1", "s1", "M1", &[1u8; 32])
.await
.unwrap();
store
.put_msg_secret("c1", "s1", "M2", &[2u8; 32])
.await
.unwrap();
store
.put_msg_secret("c1", "s2", "M1", &[3u8; 32])
.await
.unwrap();
store
.put_msg_secret("c2", "s1", "M1", &[4u8; 32])
.await
.unwrap();
for (chat, sender, msg_id, expected) in [
("c1", "s1", "M1", 1u8),
("c1", "s1", "M2", 2),
("c1", "s2", "M1", 3),
("c2", "s1", "M1", 4),
] {
let got = store
.get_msg_secret(chat, sender, msg_id)
.await
.unwrap()
.unwrap_or_else(|| panic!("missing ({chat},{sender},{msg_id})"));
assert_eq!(got, vec![expected; 32]);
}
}
#[tokio::test]
async fn msg_secret_batch_upserts_in_one_call() {
const ORIGINAL_SECRET_BYTE: u8 = 0x5a;
const UPDATED_SECRET_BYTE: u8 = 0xa5;
let store = create_test_store().await;
let mut entries: Vec<_> = (0..=MSG_SECRET_INSERT_CHUNK_SIZE)
.map(|index| MsgSecretEntry {
chat: "c".into(),
sender: "s".into(),
msg_id: format!("M{index}").into(),
secret: [ORIGINAL_SECRET_BYTE; wacore::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: 0,
message_ts: 0,
})
.collect();
entries.push(MsgSecretEntry {
chat: "c".into(),
sender: "s".into(),
msg_id: "M0".into(),
secret: [UPDATED_SECRET_BYTE; wacore::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: 0,
message_ts: 0,
});
let expected_stored = entries.len();
let stored = store.put_msg_secrets(entries).await.unwrap();
assert_eq!(stored, expected_stored);
assert_eq!(
store.get_msg_secret("c", "s", "M0").await.unwrap().unwrap(),
vec![UPDATED_SECRET_BYTE; wacore::reporting_token::MESSAGE_SECRET_SIZE]
);
assert_eq!(
store
.get_msg_secret("c", "s", &format!("M{MSG_SECRET_INSERT_CHUNK_SIZE}"))
.await
.unwrap()
.unwrap(),
vec![ORIGINAL_SECRET_BYTE; wacore::reporting_token::MESSAGE_SECRET_SIZE]
);
}
#[tokio::test]
async fn delete_expired_msg_secrets_deletes_only_passed_deadlines() {
let store = create_test_store().await;
let now = wacore::time::now_secs();
store
.put_msg_secrets(vec![
MsgSecretEntry {
chat: "c".into(),
sender: "s".into(),
msg_id: "NEVER".into(),
secret: [1u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: 0,
message_ts: 0,
},
MsgSecretEntry {
chat: "c".into(),
sender: "s".into(),
msg_id: "FUTURE".into(),
secret: [2u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: now + 86_400,
message_ts: 0,
},
MsgSecretEntry {
chat: "c".into(),
sender: "s".into(),
msg_id: "PAST".into(),
secret: [3u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: now - 86_400,
message_ts: 0,
},
])
.await
.unwrap();
let removed = store.delete_expired_msg_secrets(now).await.unwrap();
assert_eq!(
removed, 1,
"only the row whose deadline has passed is deleted"
);
assert!(
store
.get_msg_secret("c", "s", "NEVER")
.await
.unwrap()
.is_some(),
"expires_at = 0 never expires"
);
assert!(
store
.get_msg_secret("c", "s", "FUTURE")
.await
.unwrap()
.is_some(),
"a future deadline survives"
);
assert!(
store
.get_msg_secret("c", "s", "PAST")
.await
.unwrap()
.is_none(),
"a passed deadline is pruned"
);
}
#[tokio::test]
async fn put_msg_secrets_keeps_later_deadline_on_conflict() {
let store = create_test_store().await;
let now = wacore::time::now_secs();
store
.put_msg_secrets(vec![MsgSecretEntry {
chat: "c".into(),
sender: "s".into(),
msg_id: "M".into(),
secret: [1u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: now + 90 * 86_400,
message_ts: 0,
}])
.await
.unwrap();
store
.put_msg_secrets(vec![MsgSecretEntry {
chat: "c".into(),
sender: "s".into(),
msg_id: "M".into(),
secret: [1u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: now + 30 * 86_400,
message_ts: 0,
}])
.await
.unwrap();
let removed = store
.delete_expired_msg_secrets(now + 60 * 86_400)
.await
.unwrap();
assert_eq!(removed, 0, "conflict must keep the later (90d) deadline");
store
.put_msg_secret("c", "s", "M", &[1u8; 32])
.await
.unwrap();
let removed = store
.delete_expired_msg_secrets(now + 200 * 86_400)
.await
.unwrap();
assert_eq!(removed, 0, "a 0 (never) deadline wins over any finite one");
}
#[tokio::test]
async fn get_msg_secret_with_ts_round_trips_and_keeps_parent_ts() {
let store = create_test_store().await;
let parent_ts = 1_700_000_000i64;
store
.put_msg_secrets(vec![MsgSecretEntry {
chat: "c".into(),
sender: "s".into(),
msg_id: "M".into(),
secret: [5u8; wacore::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: 0,
message_ts: parent_ts,
}])
.await
.unwrap();
assert_eq!(
store.get_msg_secret_with_ts("c", "s", "M").await.unwrap(),
Some((vec![5u8; 32], parent_ts))
);
store
.put_msg_secret("c", "s", "M", &[5u8; 32])
.await
.unwrap();
assert_eq!(
store.get_msg_secret_with_ts("c", "s", "M").await.unwrap(),
Some((vec![5u8; 32], parent_ts)),
"message_ts (immutable parent time) must survive a 0-ts redelivery"
);
assert_eq!(
store
.get_msg_secret_with_ts("c", "s", "MISSING")
.await
.unwrap(),
None
);
}
#[tokio::test]
async fn msg_secret_isolated_per_device_id() {
use portable_atomic::AtomicU64;
use std::sync::atomic::Ordering;
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let shared_url = format!(
"file:memdb_msgsecret_iso_{}_{}?mode=memory&cache=shared",
std::process::id(),
id
);
let store_a = SqliteStore::new_for_device(&shared_url, 1)
.await
.expect("store_a");
let store_b = SqliteStore::new_for_device(&shared_url, 2)
.await
.expect("store_b");
store_a
.put_msg_secret("c", "s", "M", &[7u8; 32])
.await
.unwrap();
assert!(
store_b
.get_msg_secret("c", "s", "M")
.await
.unwrap()
.is_none(),
"same DB, different device_id must not see each other's secrets"
);
assert_eq!(
store_a
.get_msg_secret("c", "s", "M")
.await
.unwrap()
.unwrap(),
vec![7u8; 32],
"device_a still sees its own write"
);
}
#[tokio::test]
async fn resource_report_bounds_cache_by_db_size_and_cap() {
let store = create_test_store().await; let device_id = 1;
let macs: Vec<AppStateMutationMAC> = (0..500u32)
.map(|i| {
let mut index_mac = vec![0u8; 32];
index_mac[..4].copy_from_slice(&i.to_le_bytes());
AppStateMutationMAC {
index_mac,
value_mac: vec![(i % 251) as u8; 32],
}
})
.collect();
store
.put_app_state_mutation_macs_for_device("coll", 1, &macs, device_id)
.await
.unwrap();
let report = store.resource_report().await;
let pages = report.pages.expect("SQLite reports a page count");
assert!(pages > 0, "a migrated + seeded DB has pages");
let mem = report
.memory_bytes
.expect("SQLite reports a cache estimate");
assert!(mem > 0, "cache-in-use estimate is non-zero for a seeded DB");
assert!(
mem <= 512 * 1024,
"estimate never exceeds the configured 512 KiB cap, got {mem}"
);
assert_eq!(report.total_bytes(), mem, "total_bytes == memory_bytes");
assert_eq!(report.io_read_bytes, None);
assert_eq!(report.io_write_bytes, None);
}
#[test]
fn mmap_size_config_is_opt_in() {
assert_eq!(
SqliteStoreConfig::default().mmap_size,
None,
"default leaves mmap off (current behavior)"
);
assert_eq!(
SqliteStoreConfig::default()
.with_mmap_size(64 * 1024 * 1024)
.mmap_size,
Some(64 * 1024 * 1024),
"builder sets the field"
);
}
#[tokio::test]
async fn mmap_size_applies_pragma_and_store_operates() {
use portable_atomic::AtomicU64;
use std::sync::atomic::Ordering;
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = COUNTER.fetch_add(1, Ordering::Relaxed);
let path =
std::env::temp_dir().join(format!("wa_mmap_test_{}_{}.db", std::process::id(), id));
let url = path.to_str().unwrap().to_string();
let read_mmap = |store: &SqliteStore| -> i64 {
#[derive(diesel::QueryableByName)]
struct M {
#[diesel(sql_type = diesel::sql_types::BigInt)]
mmap_size: i64,
}
let mut conn = store.pool.get().unwrap();
diesel::sql_query("PRAGMA mmap_size")
.get_result::<M>(&mut conn)
.map(|m| m.mmap_size)
.unwrap_or(-1)
};
let def_store = SqliteStore::new(&url).await.expect("default store");
assert_eq!(read_mmap(&def_store), 0, "default keeps mmap off");
drop(def_store);
const MMAP: u64 = 64 * 1024 * 1024;
let cfg = SqliteStoreConfig::default().with_mmap_size(MMAP);
let store = SqliteStore::with_config(&url, cfg)
.await
.expect("mmap store builds");
store
.put_identity("559980000001@s.whatsapp.net", [9u8; 32])
.await
.expect("store operates with mmap set");
let applied = read_mmap(&store);
assert!(
applied == MMAP as i64 || applied == 0,
"mmap_size is applied when the VFS supports it, got {applied}"
);
drop(store);
for suffix in ["", "-wal", "-shm"] {
let _ = std::fs::remove_file(format!("{url}{suffix}"));
}
}
}