use hashbrown::hash_map::Entry;
use hashbrown::{Equivalent, HashMap as HbHashMap};
use std::collections::HashMap;
use std::collections::hash_map::RandomState;
use std::hash::Hash;
use std::sync::Arc;
#[cfg(any(test, feature = "test-util"))]
use std::sync::atomic::{AtomicBool, AtomicU32};
use std::sync::atomic::{AtomicI32, Ordering};
use crate::appstate::hash::HashState;
use crate::store::Device;
use crate::store::error::Result;
use crate::store::traits::*;
use async_lock::Mutex;
use async_trait::async_trait;
use bytes::Bytes;
use wacore_appstate::processor::AppStateMutationMAC;
type SentMessageKey = (String, String);
struct SentMessageEntry {
payload: Vec<u8>,
timestamp: i64,
}
struct PreKeyEntry {
record: Bytes,
}
type BaseKeyKey = (String, String);
type MsgSecretRow = (MessageSecret, i64, i64);
#[derive(Eq, Hash, PartialEq)]
struct MsgSecretKey {
chat: Arc<str>,
sender: Arc<str>,
msg_id: Arc<str>,
}
#[derive(Hash)]
struct MsgSecretKeyRef<'a> {
chat: &'a str,
sender: &'a str,
msg_id: &'a str,
}
impl Equivalent<MsgSecretKey> for MsgSecretKeyRef<'_> {
fn equivalent(&self, key: &MsgSecretKey) -> bool {
self.chat == key.chat.as_ref()
&& self.sender == key.sender.as_ref()
&& self.msg_id == key.msg_id.as_ref()
}
}
type MsgSecretMap = HbHashMap<MsgSecretKey, MsgSecretRow, RandomState>;
#[derive(Default)]
struct InMemoryState {
identities: HashMap<String, [u8; 32]>,
sessions: HashMap<String, Bytes>,
prekeys: HashMap<u32, PreKeyEntry>,
signed_prekeys: HashMap<u32, Vec<u8>>,
sender_keys: HashMap<String, Vec<u8>>,
sync_keys: HashMap<Vec<u8>, AppStateSyncKey>,
latest_sync_key_id: Option<Vec<u8>>,
versions: HashMap<String, HashState>,
mutation_macs: HashMap<(String, Vec<u8>), Vec<u8>>,
sender_key_devices: HashMap<String, HashMap<String, bool>>,
lid_mappings: HashMap<String, LidPnMappingEntry>,
pn_to_lid: HashMap<String, String>,
base_keys: HashMap<BaseKeyKey, Vec<u8>>,
device_lists: HashMap<String, DeviceListRecord>,
group_metadata: HashMap<String, Vec<u8>>,
tc_tokens: HashMap<String, TcTokenEntry>,
sent_messages: HashMap<SentMessageKey, SentMessageEntry>,
pending_inbound: HashMap<(String, String, String), (Vec<u8>, i64)>,
msg_secrets: MsgSecretMap,
device: Option<Device>,
}
const MAX_SENT_MESSAGES: usize = 4096;
pub struct InMemoryBackend {
state: Mutex<InMemoryState>,
next_device_id: AtomicI32,
#[cfg(any(test, feature = "test-util"))]
session_batch_writes: AtomicU32,
#[cfg(any(test, feature = "test-util"))]
sender_key_batch_writes: AtomicU32,
#[cfg(any(test, feature = "test-util"))]
fail_session_writes: AtomicBool,
#[cfg(any(test, feature = "test-util"))]
fail_sender_key_writes: AtomicBool,
}
impl InMemoryBackend {
pub fn new() -> Self {
Self {
state: Mutex::new(InMemoryState::default()),
next_device_id: AtomicI32::new(1),
#[cfg(any(test, feature = "test-util"))]
session_batch_writes: AtomicU32::new(0),
#[cfg(any(test, feature = "test-util"))]
sender_key_batch_writes: AtomicU32::new(0),
#[cfg(any(test, feature = "test-util"))]
fail_session_writes: AtomicBool::new(false),
#[cfg(any(test, feature = "test-util"))]
fail_sender_key_writes: AtomicBool::new(false),
}
}
#[cfg(any(test, feature = "test-util"))]
pub fn session_batch_write_count(&self) -> u32 {
self.session_batch_writes.load(Ordering::Relaxed)
}
#[cfg(any(test, feature = "test-util"))]
pub fn sender_key_batch_write_count(&self) -> u32 {
self.sender_key_batch_writes.load(Ordering::Relaxed)
}
#[cfg(any(test, feature = "test-util"))]
pub fn set_fail_session_writes(&self, fail: bool) {
self.fail_session_writes.store(fail, Ordering::Relaxed);
}
#[cfg(any(test, feature = "test-util"))]
pub fn set_fail_sender_key_writes(&self, fail: bool) {
self.fail_sender_key_writes.store(fail, Ordering::Relaxed);
}
#[cfg(any(test, feature = "test-util"))]
pub async fn remove_sync_key_for_test(&self, key_id: &[u8]) -> bool {
self.state.lock().await.sync_keys.remove(key_id).is_some()
}
#[cfg(any(test, feature = "test-util"))]
pub async fn sync_key_count_for_test(&self) -> usize {
self.state.lock().await.sync_keys.len()
}
}
impl Default for InMemoryBackend {
fn default() -> Self {
Self::new()
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl SignalStore for InMemoryBackend {
async fn put_identity(&self, address: &str, key: [u8; 32]) -> Result<()> {
self.state
.lock()
.await
.identities
.insert(address.to_string(), key);
Ok(())
}
async fn load_identity(&self, address: &str) -> Result<Option<[u8; 32]>> {
Ok(self.state.lock().await.identities.get(address).copied())
}
async fn delete_identity(&self, address: &str) -> Result<()> {
self.state.lock().await.identities.remove(address);
Ok(())
}
async fn get_session(&self, address: &str) -> Result<Option<Bytes>> {
Ok(self.state.lock().await.sessions.get(address).cloned())
}
async fn put_session(&self, address: &str, session: &[u8]) -> Result<()> {
self.state
.lock()
.await
.sessions
.insert(address.to_string(), Bytes::copy_from_slice(session));
Ok(())
}
async fn put_sessions_batch(&self, sessions: &[(Arc<str>, Bytes)]) -> Result<()> {
#[cfg(any(test, feature = "test-util"))]
{
self.session_batch_writes.fetch_add(1, Ordering::Relaxed);
if self.fail_session_writes.load(Ordering::Relaxed) {
return Err(crate::store::error::StoreError::Io(std::io::Error::other(
"put_sessions_batch failing (test hook)",
)));
}
}
let mut state = self.state.lock().await;
state.sessions.reserve(sessions.len());
for (address, session) in sessions {
if let Some(stored) = state.sessions.get_mut(address.as_ref()) {
*stored = session.clone();
} else {
state.sessions.insert(address.to_string(), session.clone());
}
}
Ok(())
}
async fn has_session(&self, address: &str) -> Result<bool> {
Ok(self.state.lock().await.sessions.contains_key(address))
}
async fn has_signal_state_for_user(&self, user: &str) -> Result<bool> {
fn matches(addr: &str, user: &str) -> bool {
addr.strip_prefix(user)
.is_some_and(|rest| rest.starts_with('@') || rest.starts_with(':'))
}
let state = self.state.lock().await;
Ok(state.sessions.keys().any(|k| matches(k, user))
|| state.identities.keys().any(|k| matches(k, user)))
}
async fn delete_session(&self, address: &str) -> Result<()> {
self.state.lock().await.sessions.remove(address);
Ok(())
}
async fn store_prekey(&self, id: u32, record: &[u8], _uploaded: bool) -> Result<()> {
self.state.lock().await.prekeys.insert(
id,
PreKeyEntry {
record: Bytes::copy_from_slice(record),
},
);
Ok(())
}
async fn mark_prekeys_uploaded(&self, _ids: &[u32]) -> Result<()> {
Ok(())
}
async fn store_prekeys_batch(&self, keys: &[(u32, Bytes)], _uploaded: bool) -> Result<()> {
let mut state = self.state.lock().await;
for (id, record) in keys {
state.prekeys.insert(
*id,
PreKeyEntry {
record: record.clone(),
},
);
}
Ok(())
}
async fn load_prekey(&self, id: u32) -> Result<Option<Bytes>> {
Ok(self
.state
.lock()
.await
.prekeys
.get(&id)
.map(|e| e.record.clone()))
}
async fn load_prekeys_batch(&self, ids: &[u32]) -> Result<Vec<(u32, Bytes)>> {
let state = self.state.lock().await;
let mut result = Vec::with_capacity(ids.len());
for &id in ids {
if let Some(entry) = state.prekeys.get(&id) {
result.push((id, entry.record.clone()));
}
}
Ok(result)
}
async fn remove_prekey(&self, id: u32) -> Result<()> {
self.state.lock().await.prekeys.remove(&id);
Ok(())
}
async fn get_max_prekey_id(&self) -> Result<u32> {
Ok(self
.state
.lock()
.await
.prekeys
.keys()
.copied()
.max()
.unwrap_or(0))
}
async fn store_signed_prekey(&self, id: u32, record: &[u8]) -> Result<()> {
self.state
.lock()
.await
.signed_prekeys
.insert(id, record.to_vec());
Ok(())
}
async fn load_signed_prekey(&self, id: u32) -> Result<Option<Vec<u8>>> {
Ok(self.state.lock().await.signed_prekeys.get(&id).cloned())
}
async fn load_all_signed_prekeys(&self) -> Result<Vec<(u32, Vec<u8>)>> {
Ok(self
.state
.lock()
.await
.signed_prekeys
.iter()
.map(|(id, rec)| (*id, rec.clone()))
.collect())
}
async fn remove_signed_prekey(&self, id: u32) -> Result<()> {
self.state.lock().await.signed_prekeys.remove(&id);
Ok(())
}
async fn put_sender_key(&self, address: &str, record: &[u8]) -> Result<()> {
#[cfg(any(test, feature = "test-util"))]
if self.fail_sender_key_writes.load(Ordering::Relaxed) {
return Err(crate::store::error::StoreError::Io(std::io::Error::other(
"put_sender_key failing (test hook)",
)));
}
self.state
.lock()
.await
.sender_keys
.insert(address.to_string(), record.to_vec());
Ok(())
}
async fn put_sender_keys_batch(&self, sender_keys: &[(Arc<str>, Bytes)]) -> Result<()> {
#[cfg(any(test, feature = "test-util"))]
{
self.sender_key_batch_writes.fetch_add(1, Ordering::Relaxed);
if self.fail_sender_key_writes.load(Ordering::Relaxed) {
return Err(crate::store::error::StoreError::Io(std::io::Error::other(
"put_sender_keys_batch failing (test hook)",
)));
}
}
let mut state = self.state.lock().await;
state.sender_keys.reserve(sender_keys.len());
for (address, record) in sender_keys {
if let Some(stored) = state.sender_keys.get_mut(address.as_ref()) {
stored.clear();
stored.extend_from_slice(record);
} else {
state
.sender_keys
.insert(address.to_string(), record.to_vec());
}
}
Ok(())
}
async fn get_sender_key(&self, address: &str) -> Result<Option<Vec<u8>>> {
Ok(self.state.lock().await.sender_keys.get(address).cloned())
}
async fn delete_sender_key(&self, address: &str) -> Result<()> {
self.state.lock().await.sender_keys.remove(address);
Ok(())
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl AppSyncStore for InMemoryBackend {
async fn get_sync_key(&self, key_id: &[u8]) -> Result<Option<AppStateSyncKey>> {
Ok(self.state.lock().await.sync_keys.get(key_id).cloned())
}
async fn set_sync_key(&self, key_id: &[u8], key: AppStateSyncKey) -> Result<()> {
let mut s = self.state.lock().await;
s.sync_keys.insert(key_id.to_vec(), key);
s.latest_sync_key_id = Some(key_id.to_vec());
Ok(())
}
async fn get_version(&self, name: &str) -> Result<HashState> {
Ok(self
.state
.lock()
.await
.versions
.get(name)
.cloned()
.unwrap_or_default())
}
async fn set_version(&self, name: &str, state: HashState) -> Result<()> {
self.state
.lock()
.await
.versions
.insert(name.to_string(), state);
Ok(())
}
async fn put_mutation_macs(
&self,
name: &str,
_version: u64,
mutations: &[AppStateMutationMAC],
) -> Result<()> {
let mut s = self.state.lock().await;
for m in mutations {
s.mutation_macs
.insert((name.to_string(), m.index_mac.clone()), m.value_mac.clone());
}
Ok(())
}
async fn get_mutation_mac(&self, name: &str, index_mac: &[u8]) -> Result<Option<Vec<u8>>> {
Ok(self
.state
.lock()
.await
.mutation_macs
.get(&(name.to_string(), index_mac.to_vec()))
.cloned())
}
async fn delete_mutation_macs(&self, name: &str, index_macs: &[Vec<u8>]) -> Result<()> {
let mut s = self.state.lock().await;
for im in index_macs {
s.mutation_macs.remove(&(name.to_string(), im.clone()));
}
Ok(())
}
async fn clear_mutation_macs(&self, name: &str) -> Result<()> {
self.state
.lock()
.await
.mutation_macs
.retain(|(n, _), _| n != name);
Ok(())
}
async fn get_latest_sync_key_id(&self) -> Result<Option<Vec<u8>>> {
Ok(self.state.lock().await.latest_sync_key_id.clone())
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl ProtocolStore for InMemoryBackend {
async fn get_sender_key_devices(&self, group_jid: &str) -> Result<Vec<(String, bool)>> {
Ok(self
.state
.lock()
.await
.sender_key_devices
.get(group_jid)
.map(|map| map.iter().map(|(k, v)| (k.clone(), *v)).collect())
.unwrap_or_default())
}
async fn set_sender_key_status(&self, group_jid: &str, entries: &[(&str, bool)]) -> Result<()> {
let mut s = self.state.lock().await;
let map = s
.sender_key_devices
.entry(group_jid.to_string())
.or_default();
for (device_jid, has_key) in entries {
map.insert(device_jid.to_string(), *has_key);
}
Ok(())
}
async fn clear_sender_key_devices(&self, group_jid: &str) -> Result<()> {
self.state.lock().await.sender_key_devices.remove(group_jid);
Ok(())
}
async fn clear_all_sender_key_devices(&self) -> Result<()> {
self.state.lock().await.sender_key_devices.clear();
Ok(())
}
async fn delete_sender_key_device_rows(&self, device_jids: &[&str]) -> Result<()> {
if device_jids.is_empty() {
return Ok(());
}
let mut state = self.state.lock().await;
for group_map in state.sender_key_devices.values_mut() {
group_map.retain(|jid, _| !device_jids.contains(&jid.as_str()));
}
Ok(())
}
async fn get_lid_mapping(&self, lid: &str) -> Result<Option<LidPnMappingEntry>> {
Ok(self.state.lock().await.lid_mappings.get(lid).cloned())
}
async fn get_pn_mapping(&self, phone: &str) -> Result<Option<LidPnMappingEntry>> {
let s = self.state.lock().await;
let entry = s
.pn_to_lid
.get(phone)
.and_then(|lid| s.lid_mappings.get(lid))
.cloned();
Ok(entry)
}
async fn put_lid_mapping(&self, entry: &LidPnMappingEntry) -> Result<()> {
let mut s = self.state.lock().await;
if let Some(old_phone) = s
.lid_mappings
.get(&entry.lid)
.filter(|old| old.phone_number != entry.phone_number)
.map(|old| old.phone_number.clone())
{
s.pn_to_lid.remove(&old_phone);
}
s.pn_to_lid
.insert(entry.phone_number.clone(), entry.lid.clone());
s.lid_mappings.insert(entry.lid.clone(), entry.clone());
Ok(())
}
async fn get_all_lid_mappings(&self) -> Result<Vec<LidPnMappingEntry>> {
Ok(self
.state
.lock()
.await
.lid_mappings
.values()
.cloned()
.collect())
}
async fn save_base_key(&self, address: &str, message_id: &str, base_key: &[u8]) -> Result<()> {
self.state.lock().await.base_keys.insert(
(address.to_string(), message_id.to_string()),
base_key.to_vec(),
);
Ok(())
}
async fn has_same_base_key(
&self,
address: &str,
message_id: &str,
current_base_key: &[u8],
) -> Result<bool> {
let s = self.state.lock().await;
let same = s
.base_keys
.get(&(address.to_string(), message_id.to_string()))
.is_some_and(|stored| stored == current_base_key);
Ok(same)
}
async fn delete_base_key(&self, address: &str, message_id: &str) -> Result<()> {
self.state
.lock()
.await
.base_keys
.remove(&(address.to_string(), message_id.to_string()));
Ok(())
}
async fn update_device_list(&self, record: DeviceListRecord) -> Result<()> {
self.state
.lock()
.await
.device_lists
.insert(record.user.clone(), record);
Ok(())
}
async fn get_devices(&self, user: &str) -> Result<Option<DeviceListRecord>> {
Ok(self.state.lock().await.device_lists.get(user).cloned())
}
async fn delete_devices(&self, user: &str) -> Result<()> {
self.state.lock().await.device_lists.remove(user);
Ok(())
}
async fn get_group_metadata(&self, group_jid: &str) -> Result<Option<Vec<u8>>> {
Ok(self
.state
.lock()
.await
.group_metadata
.get(group_jid)
.cloned())
}
async fn put_group_metadata(&self, group_jid: &str, blob: &[u8]) -> Result<()> {
self.state
.lock()
.await
.group_metadata
.insert(group_jid.to_string(), blob.to_vec());
Ok(())
}
async fn delete_group_metadata(&self, group_jid: &str) -> Result<()> {
self.state.lock().await.group_metadata.remove(group_jid);
Ok(())
}
async fn get_tc_token(&self, jid: &str) -> Result<Option<TcTokenEntry>> {
Ok(self.state.lock().await.tc_tokens.get(jid).cloned())
}
async fn put_tc_token(&self, jid: &str, entry: &TcTokenEntry) -> Result<()> {
self.state
.lock()
.await
.tc_tokens
.insert(jid.to_string(), entry.clone());
Ok(())
}
async fn delete_tc_token(&self, jid: &str) -> Result<()> {
self.state.lock().await.tc_tokens.remove(jid);
Ok(())
}
async fn get_all_tc_token_jids(&self) -> Result<Vec<String>> {
Ok(self.state.lock().await.tc_tokens.keys().cloned().collect())
}
async fn delete_expired_tc_tokens(&self, token_cutoff: i64, sender_cutoff: i64) -> Result<u32> {
let mut s = self.state.lock().await;
let before = s.tc_tokens.len();
s.tc_tokens.retain(|_, entry| {
let token_live = !entry.token.is_empty() && entry.token_timestamp >= token_cutoff;
let sender_live = entry.sender_timestamp.is_some_and(|ts| ts >= sender_cutoff);
token_live || sender_live
});
Ok((before - s.tc_tokens.len()) as u32)
}
async fn touch_tc_token_sender_timestamp(
&self,
jid: &str,
sender_timestamp: i64,
) -> Result<()> {
let mut s = self.state.lock().await;
match s.tc_tokens.get_mut(jid) {
Some(entry) => {
entry.sender_timestamp = Some(
entry
.sender_timestamp
.map_or(sender_timestamp, |e| e.max(sender_timestamp)),
);
}
None => {
s.tc_tokens.insert(
jid.to_string(),
TcTokenEntry {
token: Vec::new(),
token_timestamp: sender_timestamp,
sender_timestamp: Some(sender_timestamp),
},
);
}
}
Ok(())
}
async fn store_received_tc_token(
&self,
jid: &str,
token: &[u8],
token_timestamp: i64,
) -> Result<()> {
let mut s = self.state.lock().await;
match s.tc_tokens.get_mut(jid) {
Some(entry) => {
if entry.token.is_empty() || token_timestamp >= entry.token_timestamp {
entry.token = token.to_vec();
entry.token_timestamp = token_timestamp;
}
}
None => {
s.tc_tokens.insert(
jid.to_string(),
TcTokenEntry {
token: token.to_vec(),
token_timestamp,
sender_timestamp: None,
},
);
}
}
Ok(())
}
async fn store_sent_message(
&self,
chat_jid: &str,
message_id: &str,
payload: &[u8],
) -> Result<()> {
let now = crate::time::now_secs();
let mut s = self.state.lock().await;
if s.sent_messages.len() >= MAX_SENT_MESSAGES {
let target = MAX_SENT_MESSAGES * 3 / 4;
let drop_count = s.sent_messages.len().saturating_sub(target);
if drop_count > 0 {
let mut ages: Vec<i64> = s.sent_messages.values().map(|e| e.timestamp).collect();
let (_, &mut cutoff, _) = ages.select_nth_unstable(drop_count - 1);
let mut removed = 0usize;
s.sent_messages.retain(|_, e| {
if e.timestamp < cutoff {
removed += 1;
false
} else {
true
}
});
let mut remaining = drop_count.saturating_sub(removed);
if remaining > 0 {
s.sent_messages.retain(|_, e| {
if remaining > 0 && e.timestamp == cutoff {
remaining -= 1;
false
} else {
true
}
});
}
}
}
s.sent_messages.insert(
(chat_jid.to_string(), message_id.to_string()),
SentMessageEntry {
payload: payload.to_vec(),
timestamp: now,
},
);
Ok(())
}
async fn take_sent_message(&self, chat_jid: &str, message_id: &str) -> Result<Option<Vec<u8>>> {
Ok(self
.state
.lock()
.await
.sent_messages
.remove(&(chat_jid.to_string(), message_id.to_string()))
.map(|e| e.payload))
}
async fn delete_expired_sent_messages(&self, cutoff_timestamp: i64) -> Result<u32> {
let mut s = self.state.lock().await;
let before = s.sent_messages.len();
s.sent_messages
.retain(|_, entry| entry.timestamp >= cutoff_timestamp);
Ok((before - s.sent_messages.len()) as u32)
}
async fn store_pending_inbound(
&self,
chat: &str,
sender: &str,
id: &str,
message: &[u8],
) -> Result<()> {
let now = crate::time::now_secs();
self.state.lock().await.pending_inbound.insert(
(chat.to_string(), sender.to_string(), id.to_string()),
(message.to_vec(), now),
);
Ok(())
}
async fn get_pending_inbound(
&self,
chat: &str,
sender: &str,
id: &str,
) -> Result<Option<Vec<u8>>> {
let key = (chat.to_string(), sender.to_string(), id.to_string());
Ok(self
.state
.lock()
.await
.pending_inbound
.get(&key)
.map(|(bytes, _)| bytes.clone()))
}
async fn delete_pending_inbound(&self, chat: &str, sender: &str, id: &str) -> Result<()> {
let key = (chat.to_string(), sender.to_string(), id.to_string());
self.state.lock().await.pending_inbound.remove(&key);
Ok(())
}
async fn delete_expired_pending_inbound(&self, cutoff_timestamp: i64) -> Result<u32> {
let mut s = self.state.lock().await;
let before = s.pending_inbound.len();
s.pending_inbound
.retain(|_, (_, inserted_at)| *inserted_at >= cutoff_timestamp);
Ok((before - s.pending_inbound.len()) as u32)
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl MsgSecretStore for InMemoryBackend {
async fn put_msg_secrets(&self, entries: Vec<MsgSecretEntry>) -> Result<usize> {
use crate::store::traits::{merge_msg_secret_expiry, merge_msg_secret_message_ts};
let stored = entries.len();
let mut state = self.state.lock().await;
if state.msg_secrets.is_empty() {
state.msg_secrets.reserve(stored);
}
for entry in entries {
let key = MsgSecretKey {
chat: entry.chat,
sender: entry.sender,
msg_id: entry.msg_id,
};
match state.msg_secrets.entry(key) {
Entry::Occupied(mut occupied) => {
let (secret, expires_at, message_ts) = occupied.get_mut();
*secret = entry.secret;
*expires_at = merge_msg_secret_expiry(*expires_at, entry.expires_at);
*message_ts = merge_msg_secret_message_ts(*message_ts, entry.message_ts);
}
Entry::Vacant(vacant) => {
vacant.insert((entry.secret, entry.expires_at, entry.message_ts));
}
}
}
Ok(stored)
}
async fn get_msg_secret(
&self,
chat: &str,
sender: &str,
msg_id: &str,
) -> Result<Option<Vec<u8>>> {
Ok(self
.get_msg_secret_with_ts(chat, sender, msg_id)
.await?
.map(|(secret, _)| secret))
}
async fn get_msg_secret_with_ts(
&self,
chat: &str,
sender: &str,
msg_id: &str,
) -> Result<Option<(Vec<u8>, i64)>> {
Ok(self
.state
.lock()
.await
.msg_secrets
.get(&MsgSecretKeyRef {
chat,
sender,
msg_id,
})
.map(|(secret, _, message_ts)| (secret.to_vec(), *message_ts)))
}
async fn delete_expired_msg_secrets(&self, cutoff_timestamp: i64) -> Result<u32> {
let mut state = self.state.lock().await;
let before = state.msg_secrets.len();
state
.msg_secrets
.retain(|_, (_, expires_at, _)| *expires_at == 0 || *expires_at > cutoff_timestamp);
Ok((before - state.msg_secrets.len()) as u32)
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl DeviceStore for InMemoryBackend {
async fn save(&self, device: &Device) -> Result<()> {
self.state.lock().await.device = Some(device.clone());
Ok(())
}
async fn load(&self) -> Result<Option<Device>> {
Ok(self.state.lock().await.device.clone())
}
async fn exists(&self) -> Result<bool> {
Ok(self.state.lock().await.device.is_some())
}
async fn create(&self) -> Result<i32> {
let id = self.next_device_id.fetch_add(1, Ordering::Relaxed);
let mut state = self.state.lock().await;
if state.device.is_none() {
state.device = Some(Device::new());
}
Ok(id)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn is_backend<T: Backend>() {}
#[test]
fn in_memory_backend_implements_backend() {
is_backend::<InMemoryBackend>();
}
#[tokio::test]
async fn put_sessions_batch_inserts_and_updates() {
let backend = InMemoryBackend::new();
let first: Arc<str> = "15550000001:1@s.whatsapp.net".into();
let second: Arc<str> = "15550000002:2@s.whatsapp.net".into();
backend
.put_sessions_batch(&[
(first.clone(), Bytes::from_static(b"first")),
(second.clone(), Bytes::from_static(b"second")),
])
.await
.unwrap();
backend
.put_sessions_batch(&[(first.clone(), Bytes::from_static(b"updated"))])
.await
.unwrap();
assert_eq!(
backend.get_session(&first).await.unwrap().unwrap(),
Bytes::from_static(b"updated")
);
assert_eq!(
backend.get_session(&second).await.unwrap().unwrap(),
Bytes::from_static(b"second")
);
}
#[tokio::test]
async fn group_metadata_round_trip() {
use crate::store::traits::ProtocolStore;
let backend = InMemoryBackend::new();
let jid = "120363000000000001@g.us";
assert!(backend.get_group_metadata(jid).await.unwrap().is_none());
backend.put_group_metadata(jid, b"blob-v1").await.unwrap();
assert_eq!(
backend.get_group_metadata(jid).await.unwrap().as_deref(),
Some(&b"blob-v1"[..])
);
backend.put_group_metadata(jid, b"blob-v2").await.unwrap();
assert_eq!(
backend.get_group_metadata(jid).await.unwrap().as_deref(),
Some(&b"blob-v2"[..])
);
backend.delete_group_metadata(jid).await.unwrap();
assert!(backend.get_group_metadata(jid).await.unwrap().is_none());
}
#[tokio::test]
async fn clear_mutation_macs_wipes_only_named_collection() {
use crate::store::traits::AppSyncStore;
let backend = InMemoryBackend::new();
let mac = |i: u8, v: u8| AppStateMutationMAC {
index_mac: vec![i],
value_mac: vec![v],
};
backend
.put_mutation_macs("regular", 1, &[mac(1, 10)])
.await
.unwrap();
backend
.put_mutation_macs("critical", 1, &[mac(2, 20)])
.await
.unwrap();
backend.clear_mutation_macs("regular").await.unwrap();
assert!(
backend
.get_mutation_mac("regular", &[1])
.await
.unwrap()
.is_none()
);
assert_eq!(
backend.get_mutation_mac("critical", &[2]).await.unwrap(),
Some(vec![20])
);
}
#[tokio::test]
async fn has_signal_state_for_user_matches_by_user_prefix() {
let backend = InMemoryBackend::new();
let user = "5511999990000";
assert!(!backend.has_signal_state_for_user(user).await.unwrap());
backend
.put_session("5511999990000@s.whatsapp.net", b"sess")
.await
.unwrap();
assert!(backend.has_signal_state_for_user(user).await.unwrap());
let other = InMemoryBackend::new();
other
.put_session("55119999900001@s.whatsapp.net", b"sess")
.await
.unwrap();
assert!(!other.has_signal_state_for_user(user).await.unwrap());
let dev = InMemoryBackend::new();
dev.put_identity("5511999990000:5@s.whatsapp.net", [7u8; 32])
.await
.unwrap();
assert!(dev.has_signal_state_for_user(user).await.unwrap());
}
#[tokio::test]
async fn store_sent_message_is_memory_bounded() {
let backend = InMemoryBackend::new();
for i in 0..(MAX_SENT_MESSAGES + 500) {
backend
.store_sent_message("chat@g.us", &format!("m{i}"), b"payload")
.await
.unwrap();
}
let len = backend.state.lock().await.sent_messages.len();
assert!(
len <= MAX_SENT_MESSAGES,
"sent_messages must stay within the hard cap, got {len}"
);
let last = format!("m{}", MAX_SENT_MESSAGES + 500 - 1);
assert!(
backend
.take_sent_message("chat@g.us", &last)
.await
.unwrap()
.is_some(),
"the newest message must survive count-cap eviction"
);
}
#[tokio::test]
async fn store_sent_message_eviction_trims_when_all_timestamps_tie() {
let backend = InMemoryBackend::new();
for i in 0..MAX_SENT_MESSAGES {
backend
.store_sent_message("chat@g.us", &format!("m{i}"), b"payload")
.await
.unwrap();
}
{
let mut s = backend.state.lock().await;
for entry in s.sent_messages.values_mut() {
entry.timestamp = 1_000;
}
}
backend
.store_sent_message("chat@g.us", "trigger", b"payload")
.await
.unwrap();
let target = MAX_SENT_MESSAGES * 3 / 4;
let s = backend.state.lock().await;
assert_eq!(
s.sent_messages.len(),
target + 1,
"eviction must trim to 3/4 of the cap plus the insert that triggered it"
);
assert!(
s.sent_messages
.contains_key(&("chat@g.us".to_string(), "trigger".to_string())),
"the insert that triggered eviction must survive it"
);
}
#[tokio::test]
async fn store_sent_message_eviction_drops_the_oldest_first() {
let backend = InMemoryBackend::new();
for i in 0..MAX_SENT_MESSAGES {
backend
.store_sent_message("chat@g.us", &format!("m{i}"), b"payload")
.await
.unwrap();
}
let old_ids: Vec<String> = (0..16).map(|i| format!("m{i}")).collect();
{
let mut s = backend.state.lock().await;
for (key, entry) in s.sent_messages.iter_mut() {
entry.timestamp = if old_ids.contains(&key.1) { 500 } else { 1_000 };
}
}
backend
.store_sent_message("chat@g.us", "trigger", b"payload")
.await
.unwrap();
let s = backend.state.lock().await;
for id in &old_ids {
assert!(
!s.sent_messages
.contains_key(&("chat@g.us".to_string(), id.clone())),
"entry {id} is older than the cutoff and must have been evicted"
);
}
}
#[tokio::test]
async fn msg_secret_round_trip() {
let backend = InMemoryBackend::new();
let secret = [7u8; 32];
backend
.put_msg_secret("12345@s.whatsapp.net", "9999@lid", "MID1", &secret)
.await
.unwrap();
let got = backend
.get_msg_secret("12345@s.whatsapp.net", "9999@lid", "MID1")
.await
.unwrap();
assert_eq!(got.as_deref(), Some(&secret[..]));
}
#[tokio::test]
async fn msg_secret_miss_returns_none() {
let backend = InMemoryBackend::new();
assert!(
backend
.get_msg_secret("12345@s.whatsapp.net", "9999@lid", "MID1")
.await
.unwrap()
.is_none(),
"absent secret must return None"
);
}
#[tokio::test]
async fn msg_secret_keyed_by_all_three_columns() {
let backend = InMemoryBackend::new();
backend
.put_msg_secret("chatA", "senderX", "M1", &[1u8; 32])
.await
.unwrap();
backend
.put_msg_secret("chatA", "senderX", "M2", &[2u8; 32])
.await
.unwrap();
backend
.put_msg_secret("chatA", "senderY", "M1", &[3u8; 32])
.await
.unwrap();
backend
.put_msg_secret("chatB", "senderX", "M1", &[4u8; 32])
.await
.unwrap();
assert_eq!(
backend
.get_msg_secret("chatA", "senderX", "M1")
.await
.unwrap()
.unwrap(),
vec![1u8; 32]
);
assert_eq!(
backend
.get_msg_secret("chatA", "senderX", "M2")
.await
.unwrap()
.unwrap(),
vec![2u8; 32]
);
assert_eq!(
backend
.get_msg_secret("chatA", "senderY", "M1")
.await
.unwrap()
.unwrap(),
vec![3u8; 32]
);
assert_eq!(
backend
.get_msg_secret("chatB", "senderX", "M1")
.await
.unwrap()
.unwrap(),
vec![4u8; 32]
);
}
#[tokio::test]
async fn msg_secret_batch_round_trip_and_overwrite() {
let backend = InMemoryBackend::new();
let stored = backend
.put_msg_secrets(vec![
MsgSecretEntry {
chat: "chat".into(),
sender: "sender".into(),
msg_id: "M1".into(),
secret: [1u8; crate::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: 0,
message_ts: 0,
},
MsgSecretEntry {
chat: "chat".into(),
sender: "sender".into(),
msg_id: "M2".into(),
secret: [2u8; crate::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: 0,
message_ts: 0,
},
MsgSecretEntry {
chat: "chat".into(),
sender: "sender".into(),
msg_id: "M1".into(),
secret: [9u8; crate::reporting_token::MESSAGE_SECRET_SIZE],
expires_at: 0,
message_ts: 0,
},
])
.await
.unwrap();
assert_eq!(stored, 3);
assert_eq!(
backend
.get_msg_secret("chat", "sender", "M1")
.await
.unwrap()
.unwrap(),
vec![9u8; 32]
);
assert_eq!(
backend
.get_msg_secret("chat", "sender", "M2")
.await
.unwrap()
.unwrap(),
vec![2u8; 32]
);
}
#[tokio::test]
async fn delete_expired_msg_secrets_removes_only_old_rows() {
let backend = InMemoryBackend::new();
backend
.put_msg_secret("c", "s", "OLD", &[1u8; 32])
.await
.unwrap();
{
let mut state = backend.state.lock().await;
let entry = state
.msg_secrets
.get_mut(&MsgSecretKeyRef {
chat: "c",
sender: "s",
msg_id: "OLD",
})
.unwrap();
entry.1 = crate::time::now_secs() - 86_400 * 30;
}
backend
.put_msg_secret("c", "s", "NEW", &[2u8; 32])
.await
.unwrap();
let cutoff = crate::time::now_secs() - 86_400 * 14;
let removed = backend.delete_expired_msg_secrets(cutoff).await.unwrap();
assert_eq!(removed, 1);
assert!(
backend
.get_msg_secret("c", "s", "OLD")
.await
.unwrap()
.is_none()
);
assert!(
backend
.get_msg_secret("c", "s", "NEW")
.await
.unwrap()
.is_some()
);
}
#[tokio::test]
async fn msg_secret_overwrite_on_same_key() {
let backend = InMemoryBackend::new();
backend
.put_msg_secret("chat", "sender", "M", &[1u8; 32])
.await
.unwrap();
backend
.put_msg_secret("chat", "sender", "M", &[9u8; 32])
.await
.unwrap();
assert_eq!(
backend
.get_msg_secret("chat", "sender", "M")
.await
.unwrap()
.unwrap(),
vec![9u8; 32],
"last write wins for the same composite key"
);
}
#[tokio::test]
async fn touch_tc_token_creates_placeholder_then_preserves_real_token() {
let backend = InMemoryBackend::new();
backend
.touch_tc_token_sender_timestamp("u1", 1000)
.await
.unwrap();
let placeholder = backend.get_tc_token("u1").await.unwrap().unwrap();
assert!(placeholder.token.is_empty());
assert_eq!(placeholder.sender_timestamp, Some(1000));
backend
.put_tc_token(
"u1",
&TcTokenEntry {
token: vec![7, 8, 9],
token_timestamp: 2000,
sender_timestamp: None,
},
)
.await
.unwrap();
backend
.touch_tc_token_sender_timestamp("u1", 3000)
.await
.unwrap();
let merged = backend.get_tc_token("u1").await.unwrap().unwrap();
assert_eq!(
merged.token,
vec![7, 8, 9],
"touch must not clobber the real token"
);
assert_eq!(merged.token_timestamp, 2000);
assert_eq!(merged.sender_timestamp, Some(3000));
}
#[tokio::test]
async fn touch_sender_timestamp_only_advances() {
let backend = InMemoryBackend::new();
backend
.touch_tc_token_sender_timestamp("uadv", 5000)
.await
.unwrap();
backend
.touch_tc_token_sender_timestamp("uadv", 3000)
.await
.unwrap();
assert_eq!(
backend
.get_tc_token("uadv")
.await
.unwrap()
.unwrap()
.sender_timestamp,
Some(5000)
);
}
#[tokio::test]
async fn store_received_tc_token_preserves_sender_timestamp() {
let backend = InMemoryBackend::new();
backend
.touch_tc_token_sender_timestamp("u2", 5000)
.await
.unwrap();
backend
.store_received_tc_token("u2", &[1, 2, 3], 4000)
.await
.unwrap();
let entry = backend.get_tc_token("u2").await.unwrap().unwrap();
assert_eq!(entry.token, vec![1, 2, 3]);
assert_eq!(entry.token_timestamp, 4000);
assert_eq!(
entry.sender_timestamp,
Some(5000),
"store_received_tc_token must not drop the sender bucket"
);
backend
.store_received_tc_token("u3", &[9], 4000)
.await
.unwrap();
let fresh = backend.get_tc_token("u3").await.unwrap().unwrap();
assert_eq!(fresh.sender_timestamp, None);
}
#[tokio::test]
async fn store_received_tc_token_is_newer_wins() {
let backend = InMemoryBackend::new();
backend
.store_received_tc_token("c", &[1, 1, 1], 5000)
.await
.unwrap();
backend
.store_received_tc_token("c", &[2, 2, 2], 3000)
.await
.unwrap();
let e = backend.get_tc_token("c").await.unwrap().unwrap();
assert_eq!(e.token, vec![1, 1, 1], "older write must not overwrite");
assert_eq!(e.token_timestamp, 5000);
backend
.store_received_tc_token("c", &[3, 3, 3], 7000)
.await
.unwrap();
let e = backend.get_tc_token("c").await.unwrap().unwrap();
assert_eq!(e.token, vec![3, 3, 3]);
assert_eq!(e.token_timestamp, 7000);
backend
.touch_tc_token_sender_timestamp("p", 9000)
.await
.unwrap();
backend
.store_received_tc_token("p", &[4, 4, 4], 6000)
.await
.unwrap();
let e = backend.get_tc_token("p").await.unwrap().unwrap();
assert_eq!(
e.token,
vec![4, 4, 4],
"placeholder must accept first real token"
);
assert_eq!(e.token_timestamp, 6000);
assert_eq!(e.sender_timestamp, Some(9000), "sender bucket preserved");
}
#[tokio::test]
async fn prune_respects_sender_and_token_windows() {
let backend = InMemoryBackend::new();
backend
.touch_tc_token_sender_timestamp("recent_ph", 2500)
.await
.unwrap();
backend
.touch_tc_token_sender_timestamp("stale_ph", 100)
.await
.unwrap();
backend
.put_tc_token(
"expired_tok_live_sender",
&TcTokenEntry {
token: vec![1],
token_timestamp: 1,
sender_timestamp: Some(2500),
},
)
.await
.unwrap();
backend
.put_tc_token(
"orphan_expired",
&TcTokenEntry {
token: vec![2],
token_timestamp: 1,
sender_timestamp: None,
},
)
.await
.unwrap();
backend
.put_tc_token(
"fresh_tok",
&TcTokenEntry {
token: vec![3],
token_timestamp: 5000,
sender_timestamp: None,
},
)
.await
.unwrap();
let removed = backend.delete_expired_tc_tokens(1000, 2000).await.unwrap();
assert_eq!(removed, 2, "only fully-stale rows are pruned");
assert!(backend.get_tc_token("recent_ph").await.unwrap().is_some());
assert!(backend.get_tc_token("stale_ph").await.unwrap().is_none());
assert!(
backend
.get_tc_token("expired_tok_live_sender")
.await
.unwrap()
.is_some()
);
assert!(
backend
.get_tc_token("orphan_expired")
.await
.unwrap()
.is_none()
);
assert!(backend.get_tc_token("fresh_tok").await.unwrap().is_some());
}
}