use std::collections::{HashMap, HashSet};
use std::num::NonZeroU64;
use std::sync::{Arc, Mutex as SyncMutex, MutexGuard as SyncMutexGuard};
use anyhow::Result;
use async_lock::Mutex;
use portable_atomic::{AtomicBool, AtomicU64, Ordering};
use rand::RngExt;
use crate::libsignal::protocol::{
ProtocolAddress, SenderKeyRecord, SessionCheckoutKey, SessionCheckoutStoreResult, SessionRecord,
};
use crate::libsignal::store::sender_key_name::SenderKeyName;
use crate::store::traits::SignalStore;
type StoreIncarnation = [u8; 16];
fn new_store_incarnation() -> StoreIncarnation {
let mut incarnation = [0; 16];
rand::make_rng::<rand::rngs::StdRng>().fill(&mut incarnation);
incarnation
}
fn evict_clean_entries<V>(
cache: &mut HashMap<Arc<str>, Option<V>>,
dirty: &HashSet<Arc<str>>,
deleted: Option<&HashSet<Arc<str>>>,
max_entries: usize,
) {
if cache.len() <= high_watermark(max_entries) {
return;
}
let overflow = cache.len().saturating_sub(max_entries);
let mut negative = Vec::with_capacity(overflow);
let mut positive = Vec::with_capacity(overflow);
for (k, v) in cache.iter() {
if dirty.contains(k.as_ref()) {
continue;
}
if let Some(del) = deleted
&& del.contains(k.as_ref())
{
continue;
}
if v.is_none() {
negative.push(k.clone());
} else {
positive.push(k.clone());
}
}
for key in negative.into_iter().chain(positive).take(overflow) {
cache.remove(&key);
}
}
const DEFAULT_MAX_CACHE_ENTRIES: usize = 2_000;
const EVICTION_SLACK_DIVISOR: usize = 8;
const EVICTION_SLACK_FLOOR: usize = 16;
fn high_watermark(max_entries: usize) -> usize {
max_entries.saturating_add((max_entries / EVICTION_SLACK_DIVISOR).max(EVICTION_SLACK_FLOOR))
}
fn protocol_address_matches_user(address: &str, user: &str) -> bool {
address
.strip_prefix(user)
.is_some_and(|suffix| suffix.starts_with('@') || suffix.starts_with(':'))
}
pub struct SignalStoreCache {
sessions: Mutex<SessionStoreState>,
session_recovery_generation: AtomicU64,
has_pending_session_restores: AtomicBool,
pending_session_restores: SyncMutex<Vec<PendingSessionRestore>>,
identities: Mutex<ByteStoreState>,
sender_keys: Mutex<SenderKeyStoreState>,
has_pending_sender_key_distributions: AtomicBool,
removed_prekeys: Mutex<HashMap<u32, Arc<str>>>,
sender_key_locks: Mutex<HashMap<Arc<str>, Arc<Mutex<()>>>>,
max_entries: usize,
}
enum SessionEntry {
Present(Arc<SessionRecord>),
Absent,
CheckedOut {
had_session: bool,
token: NonZeroU64,
},
}
impl SessionEntry {
fn exists(&self) -> bool {
matches!(
self,
Self::Present(_)
| Self::CheckedOut {
had_session: true,
..
}
)
}
}
enum CachedSessionCheckout {
Missing(SessionCheckoutKey),
Absent(SessionCheckoutKey),
Busy,
Present(SessionRecord, SessionCheckoutKey),
}
struct SessionStoreState {
incarnation: StoreIncarnation,
checkout_generation: u64,
next_checkout_token: u64,
cache: HashMap<Arc<str>, SessionEntry>,
dirty: HashSet<Arc<str>>,
deleted: HashSet<Arc<str>>,
reservation_pending: HashSet<Arc<str>>,
}
impl SessionStoreState {
fn new(incarnation: StoreIncarnation) -> Self {
Self {
incarnation,
checkout_generation: 0,
next_checkout_token: 1,
cache: HashMap::new(),
dirty: HashSet::new(),
deleted: HashSet::new(),
reservation_pending: HashSet::new(),
}
}
fn key_for(&self, address: &str) -> Arc<str> {
match self.cache.get_key_value(address) {
Some((existing, _)) => existing.clone(),
None => Arc::from(address),
}
}
fn put(&mut self, address: &str, record: SessionRecord) {
let addr = self.key_for(address);
self.put_with_key(addr, record);
}
fn put_with_key(&mut self, addr: Arc<str>, mut record: SessionRecord) {
if record.has_pending_reservation() {
record.clear_pending_reservation();
self.reservation_pending.insert(addr.clone());
}
self.cache
.insert(addr.clone(), SessionEntry::Present(Arc::new(record)));
self.dirty.insert(addr.clone());
self.deleted.remove(&addr);
}
fn checkout(&mut self, address: &str) -> CachedSessionCheckout {
let token = NonZeroU64::new(self.next_checkout_token).unwrap_or(NonZeroU64::MIN);
self.next_checkout_token = self.next_checkout_token.wrapping_add(1);
if self.next_checkout_token == 0 {
self.next_checkout_token = 1;
}
let checkout = SessionCheckoutKey::new(self.checkout_generation, token);
let Some(entry) = self.cache.get_mut(address) else {
return CachedSessionCheckout::Missing(checkout);
};
match entry {
SessionEntry::Present(_) => {
let SessionEntry::Present(record) = std::mem::replace(
entry,
SessionEntry::CheckedOut {
had_session: true,
token,
},
) else {
unreachable!()
};
CachedSessionCheckout::Present(
Arc::try_unwrap(record).unwrap_or_else(|arc| (*arc).clone()),
checkout,
)
}
SessionEntry::Absent => {
*entry = SessionEntry::CheckedOut {
had_session: false,
token,
};
CachedSessionCheckout::Absent(checkout)
}
SessionEntry::CheckedOut { .. } => CachedSessionCheckout::Busy,
}
}
fn delete(&mut self, address: &str) {
let addr = self.key_for(address);
self.cache.insert(addr.clone(), SessionEntry::Absent);
self.deleted.insert(addr.clone());
self.dirty.remove(&addr);
}
fn clear(&mut self) {
self.cache.clear();
self.dirty.clear();
self.deleted.clear();
self.reservation_pending.clear();
}
fn clear_clean_entries(&mut self) {
self.cache
.retain(|_, entry| matches!(entry, SessionEntry::CheckedOut { .. }));
}
fn discard(&mut self, incarnation: StoreIncarnation, generation: u64) {
self.clear();
self.incarnation = incarnation;
self.checkout_generation = generation;
}
fn evict_if_needed(&mut self, max_entries: usize) {
if self.cache.len() <= high_watermark(max_entries) {
return;
}
let overflow = self.cache.len().saturating_sub(max_entries);
let mut negative = Vec::with_capacity(overflow);
let mut positive = Vec::with_capacity(overflow);
for (k, v) in self.cache.iter() {
if self.dirty.contains(k.as_ref()) || self.deleted.contains(k.as_ref()) {
continue;
}
match v {
SessionEntry::CheckedOut { .. } => continue, SessionEntry::Absent => negative.push(k.clone()),
SessionEntry::Present(_) => positive.push(k.clone()),
}
}
for key in negative.into_iter().chain(positive).take(overflow) {
self.cache.remove(&key);
}
}
}
struct PendingSessionRestore {
address: Arc<str>,
record: Option<SessionRecord>,
checkout: SessionCheckoutKey,
had_session: bool,
completion: Option<Arc<AtomicBool>>,
}
struct SenderKeyStoreState {
incarnation: StoreIncarnation,
cache: HashMap<Arc<str>, Option<Arc<SenderKeyRecord>>>,
dirty: HashSet<Arc<str>>,
wire_gate_pending: HashSet<Arc<str>>,
pending_distributions: HashMap<Arc<str>, Arc<[u8]>>,
}
impl SenderKeyStoreState {
fn new(incarnation: StoreIncarnation) -> Self {
Self {
incarnation,
cache: HashMap::new(),
dirty: HashSet::new(),
wire_gate_pending: HashSet::new(),
pending_distributions: HashMap::new(),
}
}
fn key_for(&self, address: &str) -> Arc<str> {
match self.cache.get_key_value(address) {
Some((existing, _)) => existing.clone(),
None => Arc::from(address),
}
}
fn put(&mut self, address: &str, mut record: SenderKeyRecord) {
let addr = self.key_for(address);
if record.is_wire_gated() {
record.clear_wire_gated();
self.wire_gate_pending.insert(addr.clone());
}
self.cache.insert(addr.clone(), Some(Arc::new(record)));
self.dirty.insert(addr.clone());
}
fn delete(&mut self, address: &str) {
let addr = self.key_for(address);
self.cache.insert(addr.clone(), None);
self.dirty.insert(addr.clone());
self.pending_distributions.remove(address);
}
fn clear(&mut self) {
self.cache.clear();
self.dirty.clear();
self.wire_gate_pending.clear();
self.pending_distributions.clear();
}
fn discard(&mut self, incarnation: StoreIncarnation) {
self.clear();
self.incarnation = incarnation;
}
fn evict_if_needed(&mut self, max_entries: usize) {
evict_clean_entries(&mut self.cache, &self.dirty, None, max_entries);
}
}
struct ByteStoreState {
cache: HashMap<Arc<str>, Option<Arc<[u8]>>>,
dirty: HashSet<Arc<str>>,
deleted: HashSet<Arc<str>>,
}
impl ByteStoreState {
fn new() -> Self {
Self {
cache: HashMap::new(),
dirty: HashSet::new(),
deleted: HashSet::new(),
}
}
fn key_for(&self, address: &str) -> Arc<str> {
match self.cache.get_key_value(address) {
Some((existing, _)) => existing.clone(),
None => Arc::from(address),
}
}
fn put_dedup(&mut self, address: &str, data: &[u8]) {
if let Some(Some(existing)) = self.cache.get(address)
&& existing.as_ref() == data
{
return;
}
self.put(address, data);
}
fn put(&mut self, address: &str, data: &[u8]) {
let addr = self.key_for(address);
self.cache.insert(addr.clone(), Some(Arc::from(data)));
self.dirty.insert(addr.clone());
self.deleted.remove(&addr);
}
fn delete(&mut self, address: &str) {
let addr = self.key_for(address);
self.cache.insert(addr.clone(), None);
self.deleted.insert(addr.clone());
self.dirty.remove(&addr);
}
fn clear(&mut self) {
self.cache.clear();
self.dirty.clear();
self.deleted.clear();
}
fn evict_if_needed(&mut self, max_entries: usize) {
evict_clean_entries(
&mut self.cache,
&self.dirty,
Some(&self.deleted),
max_entries,
);
}
}
impl Default for SignalStoreCache {
fn default() -> Self {
Self::new()
}
}
impl SignalStoreCache {
pub fn new() -> Self {
Self::with_max_entries(DEFAULT_MAX_CACHE_ENTRIES)
}
pub fn with_max_entries(max_entries: usize) -> Self {
Self::with_max_entries_and_incarnation(max_entries, new_store_incarnation())
}
fn with_max_entries_and_incarnation(max_entries: usize, incarnation: StoreIncarnation) -> Self {
Self {
sessions: Mutex::new(SessionStoreState::new(incarnation)),
session_recovery_generation: AtomicU64::new(0),
has_pending_session_restores: AtomicBool::new(false),
pending_session_restores: SyncMutex::new(Vec::new()),
identities: Mutex::new(ByteStoreState::new()),
sender_keys: Mutex::new(SenderKeyStoreState::new(incarnation)),
has_pending_sender_key_distributions: AtomicBool::new(false),
removed_prekeys: Mutex::new(HashMap::new()),
sender_key_locks: Mutex::new(HashMap::new()),
max_entries,
}
}
fn pending_session_restores(&self) -> SyncMutexGuard<'_, Vec<PendingSessionRestore>> {
self.pending_session_restores
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
fn drain_session_restores(&self, state: &mut SessionStoreState) {
if !self.has_pending_session_restores.load(Ordering::Acquire) {
return;
}
let mut pending = self.pending_session_restores();
for PendingSessionRestore {
address,
record,
checkout,
had_session,
completion,
} in pending.drain(..)
{
let key = if checkout.generation() == state.checkout_generation
&& let Some((
key,
SessionEntry::CheckedOut {
had_session: was_present,
token,
},
)) = state.cache.get_key_value(address.as_ref())
&& *was_present == had_session
&& *token == checkout.token()
{
Some(key.clone())
} else {
None
};
let restored = key.is_some();
match (key, record) {
(Some(key), Some(record)) => state.put_with_key(key, record),
(Some(key), None) => {
state.cache.insert(key, SessionEntry::Absent);
}
(None, _) => {}
}
if let Some(completion) = completion {
completion.store(restored, Ordering::Release);
}
}
self.has_pending_session_restores
.store(false, Ordering::Release);
state.evict_if_needed(self.max_entries);
}
async fn lock_sessions(&self) -> async_lock::MutexGuard<'_, SessionStoreState> {
let mut state = self.sessions.lock().await;
self.drain_session_restores(&mut state);
state
}
fn try_lock_sessions(&self) -> Option<async_lock::MutexGuard<'_, SessionStoreState>> {
let mut state = self.sessions.try_lock()?;
self.drain_session_restores(&mut state);
Some(state)
}
#[doc(hidden)]
pub fn restore_session_from_checkout(
&self,
address: &ProtocolAddress,
record: SessionRecord,
checkout: SessionCheckoutKey,
had_session: bool,
) -> SessionCheckoutStoreResult {
if checkout.generation() != self.session_recovery_generation.load(Ordering::Acquire) {
return SessionCheckoutStoreResult::Rejected;
}
if let Some(mut state) = self.try_lock_sessions() {
if checkout.generation() != self.session_recovery_generation.load(Ordering::Acquire)
|| checkout.generation() != state.checkout_generation
{
return SessionCheckoutStoreResult::Rejected;
}
let Some((
key,
SessionEntry::CheckedOut {
had_session: was_present,
token,
},
)) = state.cache.get_key_value(address.as_str())
else {
return SessionCheckoutStoreResult::Rejected;
};
if *was_present != had_session || *token != checkout.token() {
return SessionCheckoutStoreResult::Rejected;
}
let key = key.clone();
state.put_with_key(key, record);
state.evict_if_needed(self.max_entries);
return SessionCheckoutStoreResult::Stored;
}
let mut pending = self.pending_session_restores();
if checkout.generation() != self.session_recovery_generation.load(Ordering::Acquire) {
return SessionCheckoutStoreResult::Rejected;
}
let completion = Arc::new(AtomicBool::new(false));
pending.push(PendingSessionRestore {
address: Arc::from(address.as_str()),
record: Some(record),
checkout,
had_session,
completion: Some(completion.clone()),
});
self.has_pending_session_restores
.store(true, Ordering::Release);
SessionCheckoutStoreResult::Pending(completion)
}
#[doc(hidden)]
pub fn cancel_session_checkout(&self, address: &ProtocolAddress, checkout: SessionCheckoutKey) {
if checkout.generation() != self.session_recovery_generation.load(Ordering::Acquire) {
return;
}
let Some(mut state) = self.try_lock_sessions() else {
let mut pending = self.pending_session_restores();
if checkout.generation() == self.session_recovery_generation.load(Ordering::Acquire) {
pending.push(PendingSessionRestore {
address: Arc::from(address.as_str()),
record: None,
checkout,
had_session: false,
completion: None,
});
self.has_pending_session_restores
.store(true, Ordering::Release);
}
return;
};
let key = if checkout.generation()
== self.session_recovery_generation.load(Ordering::Acquire)
&& checkout.generation() == state.checkout_generation
&& let Some((
key,
SessionEntry::CheckedOut {
had_session: false,
token,
},
)) = state.cache.get_key_value(address.as_str())
&& *token == checkout.token()
{
Some(key.clone())
} else {
None
};
if let Some(key) = key {
state.cache.insert(key, SessionEntry::Absent);
state.evict_if_needed(self.max_entries);
}
}
#[doc(hidden)]
pub async fn complete_session_checkout(&self) {
drop(self.lock_sessions().await);
}
pub async fn has_state_for_user(&self, user: &str, backend: &dyn SignalStore) -> Result<bool> {
{
let state = self.lock_sessions().await;
if state
.cache
.keys()
.any(|address| protocol_address_matches_user(address, user))
{
return Ok(true);
}
}
{
let state = self.identities.lock().await;
if state
.cache
.keys()
.any(|address| protocol_address_matches_user(address, user))
{
return Ok(true);
}
}
Ok(backend.has_signal_state_for_user(user).await?)
}
pub async fn has_pending_pairwise_writes_for_user(&self, user: &str) -> bool {
{
let state = self.lock_sessions().await;
if state
.dirty
.iter()
.chain(&state.deleted)
.any(|address| protocol_address_matches_user(address, user))
{
return true;
}
}
let state = self.identities.lock().await;
state
.dirty
.iter()
.chain(&state.deleted)
.any(|address| protocol_address_matches_user(address, user))
}
fn decode_stored_session(
key: &str,
bytes: &[u8],
incarnation: &StoreIncarnation,
) -> Option<SessionRecord> {
match SessionRecord::deserialize_for_store(bytes, incarnation) {
Ok(record) => Some(record),
Err(error) => {
log::error!(
"discarding unreadable session row for addr#{:016x}: {error} — recovering with a fresh session",
wacore_binary::jid::observe_token(key)
);
crate::telemetry::session_record_quarantined();
None
}
}
}
pub async fn get_session(
&self,
address: &ProtocolAddress,
backend: &dyn SignalStore,
) -> Result<Option<SessionRecord>> {
let (record, checkout) = self.checkout_session(address, backend).await?;
if record.is_none() {
self.cancel_session_checkout(address, checkout);
}
Ok(record)
}
#[doc(hidden)]
pub async fn checkout_session(
&self,
address: &ProtocolAddress,
backend: &dyn SignalStore,
) -> Result<(Option<SessionRecord>, SessionCheckoutKey)> {
let key = address.as_str();
{
let mut state = self.lock_sessions().await;
match state.checkout(key) {
CachedSessionCheckout::Present(record, checkout) => {
return Ok((Some(record), checkout));
}
CachedSessionCheckout::Absent(checkout) => return Ok((None, checkout)),
CachedSessionCheckout::Busy => {
anyhow::bail!("session is already checked out")
}
CachedSessionCheckout::Missing(_) => {}
}
}
let backend_result = backend.get_session(key).await?;
let mut state = self.lock_sessions().await;
let checkout = match state.checkout(key) {
CachedSessionCheckout::Present(record, checkout) => {
return Ok((Some(record), checkout));
}
CachedSessionCheckout::Absent(checkout) => return Ok((None, checkout)),
CachedSessionCheckout::Busy => anyhow::bail!("session is already checked out"),
CachedSessionCheckout::Missing(checkout) => checkout,
};
match backend_result
.as_deref()
.and_then(|bytes| Self::decode_stored_session(key, bytes, &state.incarnation))
{
Some(record) => {
state.cache.insert(
Arc::from(key),
SessionEntry::CheckedOut {
had_session: true,
token: checkout.token(),
},
);
state.evict_if_needed(self.max_entries);
Ok((Some(record), checkout))
}
None => {
state.cache.insert(
Arc::from(key),
SessionEntry::CheckedOut {
had_session: false,
token: checkout.token(),
},
);
state.evict_if_needed(self.max_entries);
Ok((None, checkout))
}
}
}
#[doc(hidden)]
pub fn try_checkout_session(
&self,
address: &ProtocolAddress,
) -> Option<Result<(Option<SessionRecord>, SessionCheckoutKey)>> {
let mut state = self.try_lock_sessions()?;
match state.checkout(address.as_str()) {
CachedSessionCheckout::Present(record, checkout) => Some(Ok((Some(record), checkout))),
CachedSessionCheckout::Absent(checkout) => Some(Ok((None, checkout))),
CachedSessionCheckout::Busy => {
Some(Err(anyhow::anyhow!("session is already checked out")))
}
CachedSessionCheckout::Missing(_) => None,
}
}
pub async fn peek_session(
&self,
address: &ProtocolAddress,
backend: &dyn SignalStore,
) -> Result<Option<Arc<SessionRecord>>> {
let key = address.as_str();
{
let state = self.lock_sessions().await;
if let Some(entry) = state.cache.get(key) {
return match entry {
SessionEntry::Present(record) => Ok(Some(record.clone())),
_ => Ok(None),
};
}
}
let backend_result = backend.get_session(key).await?;
let mut state = self.lock_sessions().await;
if let Some(entry) = state.cache.get(key) {
return match entry {
SessionEntry::Present(record) => Ok(Some(record.clone())),
SessionEntry::Absent | SessionEntry::CheckedOut { .. } => Ok(None),
};
}
match backend_result
.as_deref()
.and_then(|bytes| Self::decode_stored_session(key, bytes, &state.incarnation))
{
Some(record) => {
let record = Arc::new(record);
state
.cache
.insert(Arc::from(key), SessionEntry::Present(record.clone()));
state.evict_if_needed(self.max_entries);
Ok(Some(record))
}
None => {
state.cache.insert(Arc::from(key), SessionEntry::Absent);
state.evict_if_needed(self.max_entries);
Ok(None)
}
}
}
pub async fn put_session(&self, address: &ProtocolAddress, record: SessionRecord) {
let mut state = self.lock_sessions().await;
state.put(address.as_str(), record);
state.evict_if_needed(self.max_entries);
}
#[allow(clippy::result_large_err)]
pub fn try_put_session(
&self,
address: &ProtocolAddress,
record: SessionRecord,
) -> core::result::Result<(), SessionRecord> {
match self.try_lock_sessions() {
Some(mut state) => {
state.put(address.as_str(), record);
state.evict_if_needed(self.max_entries);
Ok(())
}
None => Err(record),
}
}
pub fn try_has_session(&self, address: &ProtocolAddress) -> Option<bool> {
let state = self.try_lock_sessions()?;
state.cache.get(address.as_str()).map(SessionEntry::exists)
}
pub async fn delete_session(&self, address: &ProtocolAddress) {
let mut state = self.lock_sessions().await;
state.delete(address.as_str());
}
pub async fn has_session(
&self,
address: &ProtocolAddress,
backend: &dyn SignalStore,
) -> Result<bool> {
let key = address.as_str();
{
let state = self.lock_sessions().await;
if let Some(entry) = state.cache.get(key) {
return Ok(entry.exists());
}
}
let backend_result = backend.get_session(key).await?;
let mut state = self.lock_sessions().await;
if let Some(entry) = state.cache.get(key) {
return Ok(entry.exists());
}
let entry = match backend_result
.as_deref()
.and_then(|bytes| Self::decode_stored_session(key, bytes, &state.incarnation))
{
Some(record) => SessionEntry::Present(Arc::new(record)),
None => SessionEntry::Absent,
};
let exists = entry.exists();
state.cache.insert(Arc::from(key), entry);
state.evict_if_needed(self.max_entries);
Ok(exists)
}
pub async fn get_identity(
&self,
address: &ProtocolAddress,
backend: &dyn SignalStore,
) -> Result<Option<Arc<[u8]>>> {
let key = address.as_str();
{
let state = self.identities.lock().await;
if let Some(cached) = state.cache.get(key) {
return Ok(cached.clone());
}
}
let data = backend.load_identity(key).await?;
let arc_data = data.map(Arc::from);
let mut state = self.identities.lock().await;
if let Some(cached) = state.cache.get(key) {
return Ok(cached.clone());
}
state.cache.insert(Arc::from(key), arc_data.clone());
state.evict_if_needed(self.max_entries);
Ok(arc_data)
}
pub async fn put_identity(&self, address: &ProtocolAddress, data: &[u8]) {
let mut state = self.identities.lock().await;
state.put_dedup(address.as_str(), data);
state.evict_if_needed(self.max_entries);
}
pub fn try_get_identity(&self, address: &ProtocolAddress) -> Option<Option<Arc<[u8]>>> {
let state = self.identities.try_lock()?;
state.cache.get(address.as_str()).cloned()
}
pub fn try_put_identity(&self, address: &ProtocolAddress, data: &[u8]) -> bool {
match self.identities.try_lock() {
Some(mut state) => {
state.put_dedup(address.as_str(), data);
state.evict_if_needed(self.max_entries);
true
}
None => false,
}
}
pub async fn delete_identity(&self, address: &ProtocolAddress) {
let mut state = self.identities.lock().await;
state.delete(address.as_str());
}
pub async fn get_sender_key(
&self,
name: &SenderKeyName,
backend: &dyn SignalStore,
) -> Result<Option<Arc<SenderKeyRecord>>> {
let key = name.cache_key();
let mut state = self.sender_keys.lock().await;
if let Some(cached) = state.cache.get(key) {
return Ok(cached.clone());
}
let record = match backend.get_sender_key(key).await? {
Some(bytes) => Some(Arc::new(SenderKeyRecord::deserialize_for_store(
&bytes,
&state.incarnation,
)?)),
None => None,
};
state.cache.insert(Arc::from(key), record.clone());
state.evict_if_needed(self.max_entries);
Ok(record)
}
pub async fn put_sender_key(&self, name: &SenderKeyName, record: SenderKeyRecord) {
let mut state = self.sender_keys.lock().await;
state.put(name.cache_key(), record);
state.evict_if_needed(self.max_entries);
}
pub async fn cache_pending_sender_key_distribution(
&self,
name: &SenderKeyName,
distribution: Arc<[u8]>,
) {
let mut state = self.sender_keys.lock().await;
let key = state.key_for(name.cache_key());
state.pending_distributions.insert(key, distribution);
self.has_pending_sender_key_distributions
.store(true, Ordering::Release);
}
pub async fn pending_sender_key_distribution(&self, name: &SenderKeyName) -> Option<Arc<[u8]>> {
if !self
.has_pending_sender_key_distributions
.load(Ordering::Acquire)
{
return None;
}
self.sender_keys
.lock()
.await
.pending_distributions
.get(name.cache_key())
.cloned()
}
pub async fn clear_pending_sender_key_distribution(
&self,
name: &SenderKeyName,
expected: &[u8],
) {
let mut state = self.sender_keys.lock().await;
if state
.pending_distributions
.get(name.cache_key())
.is_some_and(|distribution| distribution.as_ref() == expected)
{
state.pending_distributions.remove(name.cache_key());
if state.pending_distributions.is_empty() {
self.has_pending_sender_key_distributions
.store(false, Ordering::Release);
}
}
}
pub async fn sender_key_lock(&self, name: &SenderKeyName) -> Arc<Mutex<()>> {
self.shared_named_lock(name.cache_key()).await
}
pub async fn session_setup_lock(&self, name: &SenderKeyName) -> Arc<Mutex<()>> {
let mut key = String::with_capacity(name.cache_key().len() + 8);
key.push_str(name.cache_key());
key.push_str("::setup");
self.shared_named_lock(&key).await
}
async fn shared_named_lock(&self, key: &str) -> Arc<Mutex<()>> {
let mut map = self.sender_key_locks.lock().await;
if let Some(lock) = map.get(key) {
return lock.clone();
}
if map.len() >= self.max_entries {
map.retain(|_, lock| Arc::strong_count(lock) > 1);
}
let lock = Arc::new(Mutex::new(()));
map.insert(Arc::from(key), lock.clone());
lock
}
pub async fn delete_sender_key(&self, cache_key: &str) {
let lock = self.shared_named_lock(cache_key).await;
let _guard = lock.lock().await;
let mut state = self.sender_keys.lock().await;
state.delete(cache_key);
if state.pending_distributions.is_empty() {
self.has_pending_sender_key_distributions
.store(false, Ordering::Release);
}
}
pub async fn delete_sender_key_durable(
&self,
name: &SenderKeyName,
backend: &dyn SignalStore,
) -> Result<()> {
let lock = self.sender_key_lock(name).await;
let _guard = lock.lock().await;
let cache_key = name.cache_key();
{
let mut state = self.sender_keys.lock().await;
state.delete(cache_key);
if state.pending_distributions.is_empty() {
self.has_pending_sender_key_distributions
.store(false, Ordering::Release);
}
}
backend.delete_sender_key(cache_key).await?;
let mut state = self.sender_keys.lock().await;
if matches!(state.cache.get(cache_key), Some(None)) {
state.dirty.remove(cache_key);
state.wire_gate_pending.remove(cache_key);
} else if state.cache.contains_key(cache_key) {
let key = state.key_for(cache_key);
state.dirty.insert(key);
}
state.evict_if_needed(self.max_entries);
Ok(())
}
pub async fn remove_prekey(&self, prekey_id: u32, session_address: &str) {
self.removed_prekeys
.lock()
.await
.insert(prekey_id, Arc::from(session_address));
}
pub async fn flush(&self, backend: &dyn SignalStore) -> Result<()> {
{
let mut state = self.lock_sessions().await;
let incarnation = state.incarnation;
let dirty_keys: Vec<_> = state.dirty.iter().cloned().collect();
let deleted_keys: Vec<_> = state.deleted.iter().cloned().collect();
let mut batch: Vec<(Arc<str>, bytes::Bytes)> = Vec::new();
for address in &dirty_keys {
if let Some(SessionEntry::Present(record)) = state.cache.get(address.as_ref()) {
let mut buf = Vec::new();
record.serialize_into_for_store(&mut buf, &incarnation);
batch.push((address.clone(), bytes::Bytes::from(buf)));
}
}
if !batch.is_empty() {
backend.put_sessions_batch(&batch).await?;
for (address, _) in &batch {
state.reservation_pending.remove(address);
}
}
for address in &deleted_keys {
backend.delete_session(address).await?;
state.reservation_pending.remove(address);
}
for key in &dirty_keys {
if !matches!(
state.cache.get(key.as_ref()),
Some(SessionEntry::CheckedOut { .. })
) {
state.dirty.remove(key);
}
}
for key in &deleted_keys {
state.deleted.remove(key);
}
state.evict_if_needed(self.max_entries);
{
let mut removed = self.removed_prekeys.lock().await;
if !removed.is_empty() {
let mut deletable: Vec<u32> = Vec::new();
for (id, addr) in removed.iter() {
let durable = match state.cache.get(addr.as_ref()) {
Some(SessionEntry::Present(_)) => Some(true),
Some(SessionEntry::CheckedOut { .. }) => Some(false),
Some(SessionEntry::Absent) | None => None,
};
let durable = match durable {
Some(d) => d,
None => backend
.get_session(addr.as_ref())
.await?
.as_deref()
.and_then(|bytes| {
Self::decode_stored_session(
addr.as_ref(),
bytes,
&state.incarnation,
)
})
.is_some(),
};
if durable {
deletable.push(*id);
}
}
for id in &deletable {
backend.remove_prekey(*id).await?;
}
for id in &deletable {
removed.remove(id);
}
}
}
}
{
let mut state = self.identities.lock().await;
let dirty_keys: Vec<_> = state.dirty.iter().cloned().collect();
let deleted_keys: Vec<_> = state.deleted.iter().cloned().collect();
let mut batch: Vec<(Arc<str>, [u8; 32])> = Vec::new();
for address in &dirty_keys {
if let Some(Some(data)) = state.cache.get(address.as_ref()) {
let key: [u8; 32] = data.as_ref().try_into().map_err(|_| {
anyhow::anyhow!(
"Corrupted identity key for {address}: expected 32 bytes, got {}",
data.len()
)
})?;
batch.push((address.clone(), key));
}
}
if !batch.is_empty() {
backend.put_identities_batch(&batch).await?;
}
for address in &deleted_keys {
backend.delete_identity(address).await?;
}
for key in &dirty_keys {
state.dirty.remove(key);
}
for key in &deleted_keys {
state.deleted.remove(key);
}
state.evict_if_needed(self.max_entries);
}
{
let mut state = self.sender_keys.lock().await;
let incarnation = state.incarnation;
let dirty_keys: Vec<_> = state.dirty.iter().cloned().collect();
let mut batch: Vec<(Arc<str>, bytes::Bytes)> = Vec::new();
for name in &dirty_keys {
match state.cache.get(name.as_ref()) {
Some(Some(record)) => {
let bytes = record
.serialize_for_store(&incarnation)
.map_err(|e| anyhow::anyhow!("sender key serialize for {name}: {e}"))?;
batch.push((name.clone(), bytes::Bytes::from(bytes)));
}
Some(None) => {
backend.delete_sender_key(name).await?;
state.wire_gate_pending.remove(name);
}
None => {}
}
}
if !batch.is_empty() {
backend.put_sender_keys_batch(&batch).await?;
for (name, _) in &batch {
state.wire_gate_pending.remove(name);
}
}
for key in &dirty_keys {
state.dirty.remove(key);
}
state.evict_if_needed(self.max_entries);
}
Ok(())
}
pub async fn needs_pre_wire_flush(&self) -> bool {
if !self.lock_sessions().await.reservation_pending.is_empty() {
return true;
}
!self.sender_keys.lock().await.wire_gate_pending.is_empty()
}
pub async fn memory_stats(
&self,
) -> (
crate::stats::CollectionStats,
crate::stats::CollectionStats,
crate::stats::CollectionStats,
) {
use crate::stats::CollectionStats;
let (session_count, session_keys_len, session_recs): (u64, usize, Vec<_>) = {
let s = self.lock_sessions().await;
let mut keys_len = 0usize;
let recs = s
.cache
.iter()
.filter_map(|(k, v)| {
keys_len += k.len();
match v {
SessionEntry::Present(rec) => Some(rec.clone()),
SessionEntry::Absent | SessionEntry::CheckedOut { .. } => None,
}
})
.collect();
(s.cache.len() as u64, keys_len, recs)
};
let session_bytes: usize = session_keys_len
+ session_recs
.iter()
.map(|r| r.estimated_size())
.sum::<usize>();
let sessions = CollectionStats::new(session_count, session_bytes as u64);
let identities = {
let i = self.identities.lock().await;
let bytes: usize = i
.cache
.iter()
.map(|(k, v)| k.len() + v.as_ref().map_or(0, |b| b.len()))
.sum();
CollectionStats::new(i.cache.len() as u64, bytes as u64)
};
let (sk_count, sk_keys_len, sk_pending_bytes, sk_recs): (u64, usize, usize, Vec<_>) = {
let sk = self.sender_keys.lock().await;
let mut keys_len = 0usize;
let recs = sk
.cache
.iter()
.filter_map(|(k, v)| {
keys_len += k.len();
v.clone()
})
.collect();
let pending_bytes = sk
.pending_distributions
.values()
.map(|distribution| distribution.len())
.sum();
let (pending_only_count, pending_only_key_bytes) = sk
.pending_distributions
.keys()
.filter(|key| !sk.cache.contains_key(key.as_ref()))
.fold((0usize, 0usize), |(count, bytes), key| {
(count + 1, bytes + key.len())
});
keys_len += pending_only_key_bytes;
(
(sk.cache.len() + pending_only_count) as u64,
keys_len,
pending_bytes,
recs,
)
};
let sk_bytes: usize = sk_keys_len
+ sk_pending_bytes
+ sk_recs.iter().map(|r| r.estimated_size()).sum::<usize>();
let sender_keys = CollectionStats::new(sk_count, sk_bytes as u64);
(sessions, identities, sender_keys)
}
pub async fn clear(&self) {
self.clear_with_incarnation(new_store_incarnation()).await;
}
async fn clear_with_incarnation(&self, incarnation: StoreIncarnation) {
self.session_recovery_generation
.fetch_add(1, Ordering::AcqRel);
{
let mut sessions = self.sessions.lock().await;
let mut pending = self.pending_session_restores();
let generation = self.session_recovery_generation.load(Ordering::Acquire);
pending.clear();
self.has_pending_session_restores
.store(false, Ordering::Release);
sessions.discard(incarnation, generation);
}
self.identities.lock().await.clear();
self.sender_keys.lock().await.discard(incarnation);
self.has_pending_sender_key_distributions
.store(false, Ordering::Release);
self.removed_prekeys.lock().await.clear();
}
#[doc(hidden)]
pub async fn clear_after_flush(&self) {
let mut sessions = self.lock_sessions().await;
if sessions.dirty.is_empty()
&& sessions.deleted.is_empty()
&& sessions.reservation_pending.is_empty()
{
sessions.clear_clean_entries();
if sessions.cache.is_empty() {
self.removed_prekeys.lock().await.clear();
}
}
drop(sessions);
let mut identities = self.identities.lock().await;
if identities.dirty.is_empty() && identities.deleted.is_empty() {
identities.clear();
}
drop(identities);
let mut sender_keys = self.sender_keys.lock().await;
if sender_keys.dirty.is_empty()
&& sender_keys.wire_gate_pending.is_empty()
&& sender_keys.pending_distributions.is_empty()
{
sender_keys.clear();
self.has_pending_sender_key_distributions
.store(false, Ordering::Release);
}
}
}
#[cfg(test)]
mod sender_key_lock_tests {
use super::*;
use crate::libsignal::store::sender_key_name::SenderKeyName;
use crate::store::error::Result as StoreResult;
use bytes::Bytes;
struct BlockingSessionLookup {
started: async_lock::Barrier,
release: async_lock::Barrier,
}
impl BlockingSessionLookup {
fn new() -> Self {
Self {
started: async_lock::Barrier::new(2),
release: async_lock::Barrier::new(2),
}
}
}
#[async_trait::async_trait]
impl SignalStore for BlockingSessionLookup {
async fn put_identity(&self, _: &str, _: [u8; 32]) -> StoreResult<()> {
unreachable!()
}
async fn load_identity(&self, _: &str) -> StoreResult<Option<[u8; 32]>> {
unreachable!()
}
async fn delete_identity(&self, _: &str) -> StoreResult<()> {
unreachable!()
}
async fn get_session(&self, _: &str) -> StoreResult<Option<Bytes>> {
self.started.wait().await;
self.release.wait().await;
Ok(None)
}
async fn has_session(&self, _: &str) -> StoreResult<bool> {
self.started.wait().await;
self.release.wait().await;
Ok(false)
}
async fn put_session(&self, _: &str, _: &[u8]) -> StoreResult<()> {
unreachable!()
}
async fn delete_session(&self, _: &str) -> StoreResult<()> {
unreachable!()
}
async fn store_prekey(&self, _: u32, _: &[u8], _: bool) -> StoreResult<()> {
unreachable!()
}
async fn load_prekey(&self, _: u32) -> StoreResult<Option<Bytes>> {
unreachable!()
}
async fn mark_prekeys_uploaded(&self, _: &[u32]) -> StoreResult<()> {
unreachable!()
}
async fn remove_prekey(&self, _: u32) -> StoreResult<()> {
unreachable!()
}
async fn get_max_prekey_id(&self) -> StoreResult<u32> {
unreachable!()
}
async fn store_signed_prekey(&self, _: u32, _: &[u8]) -> StoreResult<()> {
unreachable!()
}
async fn load_signed_prekey(&self, _: u32) -> StoreResult<Option<Vec<u8>>> {
unreachable!()
}
async fn load_all_signed_prekeys(&self) -> StoreResult<Vec<(u32, Vec<u8>)>> {
unreachable!()
}
async fn remove_signed_prekey(&self, _: u32) -> StoreResult<()> {
unreachable!()
}
async fn put_sender_key(&self, _: &str, _: &[u8]) -> StoreResult<()> {
unreachable!()
}
async fn get_sender_key(&self, _: &str) -> StoreResult<Option<Vec<u8>>> {
unreachable!()
}
async fn delete_sender_key(&self, _: &str) -> StoreResult<()> {
unreachable!()
}
}
async fn wait_for_lock_waiter(lock: &Arc<Mutex<()>>, baseline: usize) {
for _ in 0..10_000 {
if Arc::strong_count(lock) > baseline {
return;
}
tokio::task::yield_now().await;
}
panic!("task did not reach the contested lock");
}
#[tokio::test]
async fn same_name_shares_one_lock() {
let cache = SignalStoreCache::new();
let a = SenderKeyName::from_parts("g1@g.us", "u1@s.whatsapp.net:0");
let b = SenderKeyName::from_parts("g2@g.us", "u1@s.whatsapp.net:0");
let l1 = cache.sender_key_lock(&a).await;
let l2 = cache.sender_key_lock(&a).await;
let l3 = cache.sender_key_lock(&b).await;
assert!(Arc::ptr_eq(&l1, &l2), "same name must share one lock");
assert!(!Arc::ptr_eq(&l1, &l3), "different names must not share");
}
#[tokio::test]
async fn same_name_lock_is_mutually_exclusive() {
let cache = SignalStoreCache::new();
let name = SenderKeyName::from_parts("g@g.us", "u@s.whatsapp.net:0");
let lock = cache.sender_key_lock(&name).await;
let guard = lock.lock().await;
assert!(
lock.try_lock().is_none(),
"held lock must block a second acquire"
);
drop(guard);
assert!(lock.try_lock().is_some(), "released lock must reacquire");
}
#[tokio::test]
async fn delete_waits_for_the_chain_lock() {
let cache = Arc::new(SignalStoreCache::new());
let backend = crate::store::in_memory::InMemoryBackend::new();
let name = SenderKeyName::from_parts("g@g.us", "u@s.whatsapp.net:0");
cache
.put_sender_key(&name, SenderKeyRecord::new_empty())
.await;
let lock = cache.sender_key_lock(&name).await;
let held = lock.lock().await;
let lock_refs = Arc::strong_count(&lock);
let started = Arc::new(async_lock::Barrier::new(2));
let task = tokio::spawn({
let cache = cache.clone();
let started = started.clone();
let cache_key = name.cache_key().to_string();
async move {
started.wait().await;
cache.delete_sender_key(&cache_key).await;
}
});
started.wait().await;
wait_for_lock_waiter(&lock, lock_refs).await;
assert!(
cache
.get_sender_key(&name, &backend)
.await
.unwrap()
.is_some(),
"delete must wait for the in-flight chain mutation"
);
drop(held);
task.await.expect("delete task");
assert!(
cache
.get_sender_key(&name, &backend)
.await
.unwrap()
.is_none(),
"delete must run after the mutation releases the chain"
);
}
#[tokio::test]
async fn warm_sender_key_hit_shares_arc_not_deep_clone() {
let cache = SignalStoreCache::new();
let backend = crate::store::in_memory::InMemoryBackend::new();
let name = SenderKeyName::from_parts("g@g.us", "u@s.whatsapp.net:0");
cache
.put_sender_key(&name, SenderKeyRecord::new_empty())
.await;
let a = cache
.get_sender_key(&name, &backend)
.await
.unwrap()
.expect("warm hit");
let b = cache
.get_sender_key(&name, &backend)
.await
.unwrap()
.expect("warm hit");
assert!(Arc::ptr_eq(&a, &b));
}
#[tokio::test]
async fn try_put_session_marks_dirty_and_flushes() {
let cache = SignalStoreCache::new();
let backend = crate::store::in_memory::InMemoryBackend::new();
let addr = ProtocolAddress::new("15550009999", 1.into());
assert!(
cache
.try_put_session(&addr, SessionRecord::new_fresh())
.is_ok(),
"uncontended try_put_session must succeed"
);
assert_eq!(cache.try_has_session(&addr), Some(true));
cache.flush(&backend).await.unwrap();
assert!(
SignalStore::get_session(&backend, addr.as_str())
.await
.unwrap()
.is_some(),
"flush must persist a session stored via the fast path"
);
}
#[tokio::test]
async fn try_session_paths_fall_back_under_contention() {
let cache = SignalStoreCache::new();
let addr = ProtocolAddress::new("15550009999", 1.into());
let guard = cache.sessions.lock().await;
assert!(
cache
.try_put_session(&addr, SessionRecord::new_fresh())
.is_err(),
"held sessions lock must reject try_put_session"
);
assert_eq!(
cache.try_has_session(&addr),
None,
"held sessions lock must reject try_has_session"
);
assert!(
cache.try_checkout_session(&addr).is_none(),
"held sessions lock must defer checkout"
);
drop(guard);
assert_eq!(
cache.try_has_session(&addr),
None,
"unknown entry must defer to the async path"
);
assert!(cache.try_checkout_session(&addr).is_none());
assert!(
cache
.try_put_session(&addr, SessionRecord::new_fresh())
.is_ok(),
"released lock must accept try_put_session"
);
assert_eq!(cache.try_has_session(&addr), Some(true));
}
#[tokio::test]
async fn cancelled_checkout_queues_under_contention_and_remains_flushable() {
let cache = SignalStoreCache::new();
let backend = crate::store::in_memory::InMemoryBackend::new();
let addr = ProtocolAddress::new("15550008888", 1.into());
cache.put_session(&addr, SessionRecord::new_fresh()).await;
let (record, generation) = cache.checkout_session(&addr, &backend).await.unwrap();
let sessions = cache.sessions.lock().await;
let SessionCheckoutStoreResult::Pending(completion) = cache.restore_session_from_checkout(
&addr,
record.expect("checked-out record"),
generation,
true,
) else {
panic!("contended restore must be queued")
};
drop(completion);
assert_eq!(cache.pending_session_restores().len(), 1);
drop(sessions);
cache.flush(&backend).await.unwrap();
assert!(
SignalStore::get_session(&backend, addr.as_str())
.await
.unwrap()
.is_some(),
"a queued cancellation restore must not strand dirty state"
);
}
#[tokio::test]
async fn lossy_clear_rejects_an_older_checkout_generation() {
let cache = SignalStoreCache::new();
let backend = crate::store::in_memory::InMemoryBackend::new();
let addr = ProtocolAddress::new("15550007777", 1.into());
cache.put_session(&addr, SessionRecord::new_fresh()).await;
let (record, generation) = cache.checkout_session(&addr, &backend).await.unwrap();
cache.clear().await;
assert!(matches!(
cache.restore_session_from_checkout(
&addr,
record.expect("checked-out record"),
generation,
true,
),
SessionCheckoutStoreResult::Rejected
));
assert!(cache.peek_session(&addr, &backend).await.unwrap().is_none());
}
#[tokio::test]
async fn lossy_clear_invalidates_checkouts_before_waiting_for_the_cache() {
let cache = Arc::new(SignalStoreCache::new());
let backend = crate::store::in_memory::InMemoryBackend::new();
let addr = ProtocolAddress::new("15550007776", 1.into());
cache.put_session(&addr, SessionRecord::new_fresh()).await;
let (record, checkout) = cache.checkout_session(&addr, &backend).await.unwrap();
let sessions = cache.sessions.lock().await;
let clear = tokio::spawn({
let cache = cache.clone();
async move { cache.clear().await }
});
for _ in 0..10_000 {
if cache.session_recovery_generation.load(Ordering::Acquire) != checkout.generation() {
break;
}
tokio::task::yield_now().await;
}
assert_ne!(
cache.session_recovery_generation.load(Ordering::Acquire),
checkout.generation(),
"clear must invalidate owners before waiting"
);
assert!(matches!(
cache.restore_session_from_checkout(
&addr,
record.expect("checked-out record"),
checkout,
true,
),
SessionCheckoutStoreResult::Rejected
));
drop(sessions);
clear.await.unwrap();
}
#[tokio::test]
async fn stale_checkout_cannot_overwrite_a_new_owner() {
let cache = SignalStoreCache::new();
let addr = ProtocolAddress::new("15550007775", 1.into());
cache.put_session(&addr, SessionRecord::new_fresh()).await;
let (old_record, old_checkout) = cache
.try_checkout_session(&addr)
.expect("warm checkout")
.expect("old owner");
cache.put_session(&addr, SessionRecord::new_fresh()).await;
let (new_record, new_checkout) = cache
.try_checkout_session(&addr)
.expect("warm checkout")
.expect("new owner");
assert_ne!(old_checkout, new_checkout);
assert!(matches!(
cache.restore_session_from_checkout(
&addr,
old_record.expect("old record"),
old_checkout,
true,
),
SessionCheckoutStoreResult::Rejected
));
assert!(matches!(
cache.restore_session_from_checkout(
&addr,
new_record.expect("new record"),
new_checkout,
true,
),
SessionCheckoutStoreResult::Stored
));
}
#[tokio::test]
async fn checkout_rejects_a_competing_owner() {
let cache = SignalStoreCache::new();
let addr = ProtocolAddress::new("15550007770", 1.into());
cache.put_session(&addr, SessionRecord::new_fresh()).await;
let (record, generation) = cache
.try_checkout_session(&addr)
.expect("warm checkout")
.expect("first owner");
let error = match cache
.try_checkout_session(&addr)
.expect("checked-out slots are known")
{
Ok(_) => panic!("a second owner must be rejected"),
Err(error) => error,
};
assert!(error.to_string().contains("already checked out"));
assert!(matches!(
cache.restore_session_from_checkout(
&addr,
record.expect("first owner"),
generation,
true,
),
SessionCheckoutStoreResult::Stored
));
}
#[tokio::test]
async fn restore_does_not_resurrect_a_deleted_slot() {
let cache = SignalStoreCache::new();
let backend = crate::store::in_memory::InMemoryBackend::new();
let addr = ProtocolAddress::new("15550007771", 1.into());
cache.put_session(&addr, SessionRecord::new_fresh()).await;
let (record, generation) = cache.checkout_session(&addr, &backend).await.unwrap();
cache.delete_session(&addr).await;
assert!(matches!(
cache.restore_session_from_checkout(
&addr,
record.expect("checked-out record"),
generation,
true,
),
SessionCheckoutStoreResult::Rejected
));
assert!(cache.peek_session(&addr, &backend).await.unwrap().is_none());
}
#[tokio::test]
async fn queued_restore_does_not_overwrite_a_delete() {
let cache = SignalStoreCache::new();
let backend = crate::store::in_memory::InMemoryBackend::new();
let addr = ProtocolAddress::new("15550007772", 1.into());
cache.put_session(&addr, SessionRecord::new_fresh()).await;
let (record, generation) = cache.checkout_session(&addr, &backend).await.unwrap();
let mut sessions = cache.sessions.lock().await;
sessions.delete(addr.as_str());
let SessionCheckoutStoreResult::Pending(completion) = cache.restore_session_from_checkout(
&addr,
record.expect("checked-out record"),
generation,
true,
) else {
panic!("contended restore must be queued")
};
drop(sessions);
cache.complete_session_checkout().await;
assert!(!completion.load(Ordering::Acquire));
assert!(cache.peek_session(&addr, &backend).await.unwrap().is_none());
}
#[tokio::test]
async fn empty_checkout_reserves_and_releases_its_slot() {
let cache = SignalStoreCache::new();
let backend = crate::store::in_memory::InMemoryBackend::new();
let addr = ProtocolAddress::new("15550007773", 1.into());
let (record, generation) = cache.checkout_session(&addr, &backend).await.unwrap();
assert!(record.is_none());
assert_eq!(cache.try_has_session(&addr), Some(false));
assert!(!cache.has_session(&addr, &backend).await.unwrap());
assert!(cache.checkout_session(&addr, &backend).await.is_err());
cache.cancel_session_checkout(&addr, generation);
let (record, generation) = cache.checkout_session(&addr, &backend).await.unwrap();
assert!(record.is_none());
let sessions = cache.sessions.lock().await;
cache.cancel_session_checkout(&addr, generation);
assert_eq!(cache.pending_session_restores().len(), 1);
drop(sessions);
cache.complete_session_checkout().await;
assert_eq!(cache.try_has_session(&addr), Some(false));
}
#[tokio::test]
async fn peek_prefers_a_cache_write_that_wins_the_backend_race() {
let cache = Arc::new(SignalStoreCache::new());
let backend = Arc::new(BlockingSessionLookup::new());
let addr = ProtocolAddress::new("15550007774", 1.into());
let peek = tokio::spawn({
let cache = cache.clone();
let backend = backend.clone();
let addr = addr.clone();
async move { cache.peek_session(&addr, backend.as_ref()).await }
});
backend.started.wait().await;
cache.put_session(&addr, SessionRecord::new_fresh()).await;
backend.release.wait().await;
assert!(peek.await.unwrap().unwrap().is_some());
}
#[tokio::test]
async fn existence_prefers_a_cache_write_that_wins_the_backend_race() {
let cache = Arc::new(SignalStoreCache::new());
let backend = Arc::new(BlockingSessionLookup::new());
let addr = ProtocolAddress::new("15550007772", 2.into());
let exists = tokio::spawn({
let cache = cache.clone();
let backend = backend.clone();
let addr = addr.clone();
async move { cache.has_session(&addr, backend.as_ref()).await }
});
backend.started.wait().await;
cache.put_session(&addr, SessionRecord::new_fresh()).await;
backend.release.wait().await;
assert!(exists.await.unwrap().unwrap());
}
#[tokio::test]
async fn try_has_session_reports_known_absent() {
let cache = SignalStoreCache::new();
let addr = ProtocolAddress::new("15550009999", 1.into());
cache.delete_session(&addr).await;
assert_eq!(
cache.try_has_session(&addr),
Some(false),
"negative-cached entry must answer synchronously"
);
}
#[tokio::test]
async fn try_identity_paths_cover_hit_miss_and_contention() {
let cache = SignalStoreCache::new();
let addr = ProtocolAddress::new("15550009999", 1.into());
let key_bytes = [7u8; 32];
assert_eq!(
cache.try_get_identity(&addr),
None,
"unknown entry must defer to the async path"
);
assert!(cache.try_put_identity(&addr, &key_bytes));
match cache.try_get_identity(&addr) {
Some(Some(bytes)) => assert_eq!(bytes.as_ref(), &key_bytes),
other => panic!("expected cached identity, got {other:?}"),
}
let guard = cache.identities.lock().await;
assert_eq!(cache.try_get_identity(&addr), None);
assert!(!cache.try_put_identity(&addr, &key_bytes));
drop(guard);
cache.delete_identity(&addr).await;
assert_eq!(
cache.try_get_identity(&addr),
Some(None),
"known-absent identity must answer synchronously"
);
}
}
#[cfg(test)]
mod consumed_prekey_atomicity_tests {
use super::*;
use crate::store::in_memory::InMemoryBackend;
use crate::store::traits::SignalStore;
const PREKEY_ID: u32 = 4242;
async fn seed(backend: &InMemoryBackend) -> ProtocolAddress {
backend
.store_prekey(PREKEY_ID, b"durable-prekey", false)
.await
.unwrap();
ProtocolAddress::new("bob", 1.into())
}
#[tokio::test]
async fn consumed_prekey_stays_durable_until_session_flush() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
let addr = seed(&backend).await;
cache.put_session(&addr, SessionRecord::new_fresh()).await;
cache.remove_prekey(PREKEY_ID, addr.as_str()).await;
assert!(
backend.load_prekey(PREKEY_ID).await.unwrap().is_some(),
"consumed prekey must remain in the backend until the session flush"
);
assert!(
backend.get_session(addr.as_str()).await.unwrap().is_none(),
"session is only volatile before flush"
);
cache.flush(&backend).await.unwrap();
assert!(
backend.get_session(addr.as_str()).await.unwrap().is_some(),
"session must be durable after flush"
);
assert!(
backend.load_prekey(PREKEY_ID).await.unwrap().is_none(),
"prekey must be deleted once the session it produced is durable"
);
}
#[tokio::test]
async fn checked_out_session_defers_prekey_delete_until_durable() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
let addr = seed(&backend).await;
cache.put_session(&addr, SessionRecord::new_fresh()).await;
cache.remove_prekey(PREKEY_ID, addr.as_str()).await;
let taken = cache.get_session(&addr, &backend).await.unwrap();
assert!(taken.is_some(), "the promoted session should be readable");
cache.flush(&backend).await.unwrap();
assert!(
backend.get_session(addr.as_str()).await.unwrap().is_none(),
"a checked-out session is not persisted by this flush"
);
assert!(
backend.load_prekey(PREKEY_ID).await.unwrap().is_some(),
"prekey must not be deleted while its session is checked out (still volatile)"
);
cache.put_session(&addr, taken.unwrap()).await;
cache.flush(&backend).await.unwrap();
assert!(
backend.get_session(addr.as_str()).await.unwrap().is_some(),
"session is durable after the reader returned it"
);
assert!(
backend.load_prekey(PREKEY_ID).await.unwrap().is_none(),
"the deferred prekey deletion commits once the session is durable"
);
}
#[tokio::test]
async fn one_flush_drains_persisted_session_prekey_and_defers_checked_out_one() {
const PREKEY_A: u32 = 5101;
const PREKEY_B: u32 = 5102;
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
backend.store_prekey(PREKEY_A, b"a", false).await.unwrap();
backend.store_prekey(PREKEY_B, b"b", false).await.unwrap();
let addr_a = ProtocolAddress::new("alice", 1.into());
let addr_b = ProtocolAddress::new("bob", 1.into());
cache.put_session(&addr_a, SessionRecord::new_fresh()).await;
cache.remove_prekey(PREKEY_A, addr_a.as_str()).await;
cache.put_session(&addr_b, SessionRecord::new_fresh()).await;
cache.remove_prekey(PREKEY_B, addr_b.as_str()).await;
let taken_b = cache.get_session(&addr_b, &backend).await.unwrap();
assert!(taken_b.is_some(), "B's promoted session should be readable");
cache.flush(&backend).await.unwrap();
assert!(
backend
.get_session(addr_a.as_str())
.await
.unwrap()
.is_some(),
"A's session must be durable after the flush"
);
assert!(
backend.load_prekey(PREKEY_A).await.unwrap().is_none(),
"A's prekey must be deleted: its session was persisted this flush"
);
assert!(
backend.load_prekey(PREKEY_B).await.unwrap().is_some(),
"B's prekey must be deferred while B's session is checked out"
);
assert!(
cache.removed_prekeys.lock().await.contains_key(&PREKEY_B),
"B's prekey stays buffered for a later flush"
);
assert!(
!cache.removed_prekeys.lock().await.contains_key(&PREKEY_A),
"A's prekey must be drained from the buffer, not left to leak"
);
cache.put_session(&addr_b, taken_b.unwrap()).await;
cache.flush(&backend).await.unwrap();
assert!(
backend.load_prekey(PREKEY_B).await.unwrap().is_none(),
"B's prekey is deleted once B's session is durable"
);
}
#[tokio::test]
async fn clear_before_flush_keeps_prekey_so_pkmsg_can_rebuild() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
let addr = seed(&backend).await;
cache.put_session(&addr, SessionRecord::new_fresh()).await;
cache.remove_prekey(PREKEY_ID, addr.as_str()).await;
cache.clear().await;
assert!(
backend.get_session(addr.as_str()).await.unwrap().is_none(),
"volatile session is dropped on clear"
);
assert!(
backend.load_prekey(PREKEY_ID).await.unwrap().is_some(),
"prekey must survive a clear that discarded its unflushed session"
);
cache.flush(&backend).await.unwrap();
assert!(
backend.load_prekey(PREKEY_ID).await.unwrap().is_some(),
"cleared buffer must not delete the prekey on a later flush"
);
}
#[tokio::test]
async fn prekey_behind_an_unreadable_session_row_survives_flush() {
use super::lease_reload_tests::leased_session;
use crate::libsignal::protocol::consts::MAX_RESERVATION_FAST_FORWARD;
let backend = InMemoryBackend::new();
let addr = seed(&backend).await;
let writer = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xA1; 16],
);
let mut stranded = leased_session();
stranded.reserve_sender_chain_counters(MAX_RESERVATION_FAST_FORWARD);
writer.put_session(&addr, stranded).await;
writer.flush(&backend).await.unwrap();
assert!(
backend.get_session(addr.as_str()).await.unwrap().is_some(),
"the row is there; what follows is about whether it decodes"
);
let restarted = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xB2; 16],
);
restarted.remove_prekey(PREKEY_ID, addr.as_str()).await;
restarted.flush(&backend).await.unwrap();
assert!(
backend.load_prekey(PREKEY_ID).await.unwrap().is_some(),
"a prekey behind a row that does not decode must survive the flush"
);
}
#[tokio::test]
async fn prekey_without_a_persisted_session_survives_flush() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
let addr = seed(&backend).await;
cache.remove_prekey(PREKEY_ID, addr.as_str()).await;
cache.flush(&backend).await.unwrap();
assert!(
backend.load_prekey(PREKEY_ID).await.unwrap().is_some(),
"a prekey with no durable session must survive the flush"
);
assert!(
cache.removed_prekeys.lock().await.contains_key(&PREKEY_ID),
"it stays buffered; a later clear() drops it, keeping the prekey durable"
);
}
#[tokio::test]
async fn prekey_buffered_after_session_already_durable_is_deleted() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
let addr = seed(&backend).await;
cache.put_session(&addr, SessionRecord::new_fresh()).await;
cache.flush(&backend).await.unwrap();
assert!(backend.get_session(addr.as_str()).await.unwrap().is_some());
cache.remove_prekey(PREKEY_ID, addr.as_str()).await;
cache.flush(&backend).await.unwrap();
assert!(
backend.load_prekey(PREKEY_ID).await.unwrap().is_none(),
"prekey of an already-durable session must be deleted on the next flush"
);
}
#[tokio::test]
async fn failed_session_flush_does_not_delete_prekey() {
struct FailingSessions(InMemoryBackend);
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
impl SignalStore for FailingSessions {
async fn put_sessions_batch(
&self,
_sessions: &[(Arc<str>, bytes::Bytes)],
) -> crate::store::error::Result<()> {
Err(crate::store::error::StoreError::Validation(
"simulated session write failure".to_string(),
))
}
async fn put_identity(
&self,
address: &str,
key: [u8; 32],
) -> crate::store::error::Result<()> {
self.0.put_identity(address, key).await
}
async fn load_identity(
&self,
address: &str,
) -> crate::store::error::Result<Option<[u8; 32]>> {
self.0.load_identity(address).await
}
async fn delete_identity(&self, address: &str) -> crate::store::error::Result<()> {
self.0.delete_identity(address).await
}
async fn get_session(
&self,
address: &str,
) -> crate::store::error::Result<Option<bytes::Bytes>> {
self.0.get_session(address).await
}
async fn put_session(
&self,
address: &str,
session: &[u8],
) -> crate::store::error::Result<()> {
self.0.put_session(address, session).await
}
async fn delete_session(&self, address: &str) -> crate::store::error::Result<()> {
self.0.delete_session(address).await
}
async fn store_prekey(
&self,
id: u32,
record: &[u8],
uploaded: bool,
) -> crate::store::error::Result<()> {
self.0.store_prekey(id, record, uploaded).await
}
async fn load_prekey(
&self,
id: u32,
) -> crate::store::error::Result<Option<bytes::Bytes>> {
self.0.load_prekey(id).await
}
async fn remove_prekey(&self, id: u32) -> crate::store::error::Result<()> {
self.0.remove_prekey(id).await
}
async fn mark_prekeys_uploaded(&self, ids: &[u32]) -> crate::store::error::Result<()> {
self.0.mark_prekeys_uploaded(ids).await
}
async fn get_max_prekey_id(&self) -> crate::store::error::Result<u32> {
self.0.get_max_prekey_id().await
}
async fn store_signed_prekey(
&self,
id: u32,
record: &[u8],
) -> crate::store::error::Result<()> {
self.0.store_signed_prekey(id, record).await
}
async fn load_signed_prekey(
&self,
id: u32,
) -> crate::store::error::Result<Option<Vec<u8>>> {
self.0.load_signed_prekey(id).await
}
async fn load_all_signed_prekeys(
&self,
) -> crate::store::error::Result<Vec<(u32, Vec<u8>)>> {
self.0.load_all_signed_prekeys().await
}
async fn remove_signed_prekey(&self, id: u32) -> crate::store::error::Result<()> {
self.0.remove_signed_prekey(id).await
}
async fn put_sender_key(
&self,
address: &str,
record: &[u8],
) -> crate::store::error::Result<()> {
self.0.put_sender_key(address, record).await
}
async fn get_sender_key(
&self,
address: &str,
) -> crate::store::error::Result<Option<Vec<u8>>> {
self.0.get_sender_key(address).await
}
async fn delete_sender_key(&self, address: &str) -> crate::store::error::Result<()> {
self.0.delete_sender_key(address).await
}
}
let inner = InMemoryBackend::new();
let addr = seed(&inner).await;
let backend = FailingSessions(inner);
let cache = SignalStoreCache::new();
cache.put_session(&addr, SessionRecord::new_fresh()).await;
cache.remove_prekey(PREKEY_ID, addr.as_str()).await;
assert!(cache.flush(&backend).await.is_err());
assert!(
backend.load_prekey(PREKEY_ID).await.unwrap().is_some(),
"prekey must not be deleted when the session write fails"
);
assert!(
cache.removed_prekeys.lock().await.contains_key(&PREKEY_ID),
"buffered prekey removal must persist across a failed flush"
);
}
#[tokio::test]
async fn concurrent_decrypt_does_not_lose_prekey_during_flush() {
use std::sync::Arc as StdArc;
use std::sync::atomic::{AtomicBool, Ordering};
const PREKEY_A: u32 = 1001;
const PREKEY_B: u32 = 1002;
struct GatedBackend {
inner: InMemoryBackend,
cache: StdArc<SignalStoreCache>,
drained_without_sessions_lock: StdArc<AtomicBool>,
violation: StdArc<AtomicBool>,
addr_b: String,
}
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
impl SignalStore for GatedBackend {
async fn put_sessions_batch(
&self,
sessions: &[(Arc<str>, bytes::Bytes)],
) -> crate::store::error::Result<()> {
for _ in 0..64 {
tokio::task::yield_now().await;
}
self.inner.put_sessions_batch(sessions).await
}
async fn mark_prekeys_uploaded(&self, ids: &[u32]) -> crate::store::error::Result<()> {
self.inner.mark_prekeys_uploaded(ids).await
}
async fn remove_prekey(&self, id: u32) -> crate::store::error::Result<()> {
if self.cache.sessions.try_lock().is_some() {
self.drained_without_sessions_lock
.store(true, Ordering::SeqCst);
}
if id == PREKEY_B
&& self
.inner
.get_session(&self.addr_b)
.await
.unwrap()
.is_none()
{
self.violation.store(true, Ordering::SeqCst);
}
self.inner.remove_prekey(id).await
}
async fn put_identity(
&self,
address: &str,
key: [u8; 32],
) -> crate::store::error::Result<()> {
self.inner.put_identity(address, key).await
}
async fn load_identity(
&self,
address: &str,
) -> crate::store::error::Result<Option<[u8; 32]>> {
self.inner.load_identity(address).await
}
async fn delete_identity(&self, address: &str) -> crate::store::error::Result<()> {
self.inner.delete_identity(address).await
}
async fn get_session(
&self,
address: &str,
) -> crate::store::error::Result<Option<bytes::Bytes>> {
self.inner.get_session(address).await
}
async fn put_session(
&self,
address: &str,
session: &[u8],
) -> crate::store::error::Result<()> {
self.inner.put_session(address, session).await
}
async fn delete_session(&self, address: &str) -> crate::store::error::Result<()> {
self.inner.delete_session(address).await
}
async fn store_prekey(
&self,
id: u32,
record: &[u8],
uploaded: bool,
) -> crate::store::error::Result<()> {
self.inner.store_prekey(id, record, uploaded).await
}
async fn load_prekey(
&self,
id: u32,
) -> crate::store::error::Result<Option<bytes::Bytes>> {
self.inner.load_prekey(id).await
}
async fn get_max_prekey_id(&self) -> crate::store::error::Result<u32> {
self.inner.get_max_prekey_id().await
}
async fn store_signed_prekey(
&self,
id: u32,
record: &[u8],
) -> crate::store::error::Result<()> {
self.inner.store_signed_prekey(id, record).await
}
async fn load_signed_prekey(
&self,
id: u32,
) -> crate::store::error::Result<Option<Vec<u8>>> {
self.inner.load_signed_prekey(id).await
}
async fn load_all_signed_prekeys(
&self,
) -> crate::store::error::Result<Vec<(u32, Vec<u8>)>> {
self.inner.load_all_signed_prekeys().await
}
async fn remove_signed_prekey(&self, id: u32) -> crate::store::error::Result<()> {
self.inner.remove_signed_prekey(id).await
}
async fn put_sender_key(
&self,
address: &str,
record: &[u8],
) -> crate::store::error::Result<()> {
self.inner.put_sender_key(address, record).await
}
async fn get_sender_key(
&self,
address: &str,
) -> crate::store::error::Result<Option<Vec<u8>>> {
self.inner.get_sender_key(address).await
}
async fn delete_sender_key(&self, address: &str) -> crate::store::error::Result<()> {
self.inner.delete_sender_key(address).await
}
}
let inner = InMemoryBackend::new();
inner
.store_prekey(PREKEY_A, b"prekey-a", false)
.await
.unwrap();
inner
.store_prekey(PREKEY_B, b"prekey-b", false)
.await
.unwrap();
let addr_a = ProtocolAddress::new("alice", 1.into());
let addr_b = ProtocolAddress::new("bob", 1.into());
let cache = StdArc::new(SignalStoreCache::new());
let violation = StdArc::new(AtomicBool::new(false));
let drained_without_sessions_lock = StdArc::new(AtomicBool::new(false));
let backend = StdArc::new(GatedBackend {
inner,
cache: cache.clone(),
drained_without_sessions_lock: drained_without_sessions_lock.clone(),
violation: violation.clone(),
addr_b: addr_b.as_str().to_string(),
});
cache.put_session(&addr_a, SessionRecord::new_fresh()).await;
cache.remove_prekey(PREKEY_A, addr_a.as_str()).await;
let b_cache = cache.clone();
let addr_b_task = addr_b.clone();
let b_task = tokio::spawn(async move {
b_cache
.put_session(&addr_b_task, SessionRecord::new_fresh())
.await;
b_cache.remove_prekey(PREKEY_B, addr_b_task.as_str()).await;
});
cache.flush(backend.as_ref()).await.unwrap();
b_task.await.unwrap();
assert!(
!drained_without_sessions_lock.load(Ordering::SeqCst),
"prekey was drained without holding the sessions lock (regression)"
);
assert!(
!violation.load(Ordering::SeqCst),
"flush deleted B's prekey while B's session was still volatile"
);
assert!(
backend
.get_session(addr_a.as_str())
.await
.unwrap()
.is_some(),
"sender A's session must be durable after its flush"
);
assert!(
backend.load_prekey(PREKEY_A).await.unwrap().is_none(),
"sender A's consumed prekey must be deleted with its session"
);
assert!(
backend.load_prekey(PREKEY_B).await.unwrap().is_some(),
"B's prekey must survive a concurrent flush that did not persist B's session"
);
assert!(
cache.removed_prekeys.lock().await.contains_key(&PREKEY_B),
"B's prekey removal stays buffered for B's own flush"
);
cache.flush(backend.as_ref()).await.unwrap();
assert!(
backend
.get_session(addr_b.as_str())
.await
.unwrap()
.is_some(),
"B's session must be durable after B's flush"
);
assert!(
backend.load_prekey(PREKEY_B).await.unwrap().is_none(),
"B's prekey is deleted only once B's session is durable"
);
assert!(
!violation.load(Ordering::SeqCst),
"B's prekey delete must coincide with B's durable session"
);
}
}
#[cfg(test)]
mod eviction_tests {
use super::*;
use crate::libsignal::protocol::{DeviceId, ProtocolAddress};
use crate::store::in_memory::InMemoryBackend;
fn addr(i: usize) -> ProtocolAddress {
ProtocolAddress::new(&format!("user{i}@s.whatsapp.net"), DeviceId::new(0))
}
#[test]
fn high_watermark_is_above_max_and_amortizes() {
assert!(high_watermark(2_000) > 2_000);
assert_eq!(
high_watermark(2_000),
2_000 + 2_000 / EVICTION_SLACK_DIVISOR
);
assert_eq!(high_watermark(4), 4 + EVICTION_SLACK_FLOOR);
}
#[tokio::test]
async fn eviction_bounds_cache_over_many_inserts() {
let max = 64usize;
let cache = SignalStoreCache::with_max_entries(max);
let backend = InMemoryBackend::new();
for i in 0..(max * 4) {
cache.put_identity(&addr(i), &[0u8; 32]).await;
cache.flush(&backend).await.unwrap();
}
let len = cache.identities.lock().await.cache.len();
assert!(
len <= high_watermark(max),
"cache grew past the high watermark: len={len} watermark={}",
high_watermark(max)
);
assert!(
len >= max,
"eviction was too aggressive: len={len} max={max}"
);
}
#[tokio::test]
async fn read_over_capacity_stays_bounded() {
let max = 64usize;
let cache = SignalStoreCache::with_max_entries(max);
let backend = InMemoryBackend::new();
let watermark = high_watermark(max);
for i in 0..watermark {
cache.put_identity(&addr(i), &[0u8; 32]).await;
cache.flush(&backend).await.unwrap();
}
let before = cache.identities.lock().await.cache.len();
assert_eq!(before, watermark, "setup should fill exactly to watermark");
let missing = addr(watermark + 1);
let got = cache.get_identity(&missing, &backend).await.unwrap();
assert!(got.is_none());
let after = cache.identities.lock().await.cache.len();
assert!(
after <= watermark,
"a read over capacity must stay bounded: after={after} watermark={watermark}"
);
}
#[tokio::test]
async fn read_flood_of_unique_keys_stays_bounded() {
let max = 64usize;
let cache = SignalStoreCache::with_max_entries(max);
let backend = InMemoryBackend::new();
for i in 0..(max * 8) {
assert!(
cache
.get_identity(&addr(i), &backend)
.await
.unwrap()
.is_none()
);
}
let len = cache.identities.lock().await.cache.len();
assert!(
len <= high_watermark(max),
"unique-read flood must stay bounded: len={len} watermark={}",
high_watermark(max)
);
}
#[tokio::test]
async fn dirty_entries_are_never_evicted() {
let max = 64usize;
let cache = SignalStoreCache::with_max_entries(max);
let total = high_watermark(max) * 2;
for i in 0..total {
cache.put_identity(&addr(i), &[0u8; 32]).await;
}
let len = cache.identities.lock().await.cache.len();
assert_eq!(
len, total,
"dirty (unflushed) entries must never be evicted"
);
}
#[tokio::test]
async fn checked_out_sessions_are_never_evicted() {
let max = 64usize;
let cache = SignalStoreCache::with_max_entries(max);
let backend = InMemoryBackend::new();
let pinned = addr(0);
cache.put_session(&pinned, SessionRecord::new_fresh()).await;
cache.flush(&backend).await.unwrap();
let taken = cache.get_session(&pinned, &backend).await.unwrap();
assert!(taken.is_some(), "session should be present before checkout");
let watermark = high_watermark(max);
for i in 1..(watermark + 8) {
assert!(!cache.has_session(&addr(i), &backend).await.unwrap());
}
cache
.put_session(&addr(99_999), SessionRecord::new_fresh())
.await;
{
let state = cache.sessions.lock().await;
let entry = state.cache.get(pinned.as_str());
assert!(
matches!(entry, Some(SessionEntry::CheckedOut { .. })),
"checked-out session must survive eviction"
);
assert!(
state.cache.len() <= high_watermark(max) + 1,
"eviction must bound the session cache: len={}",
state.cache.len()
);
}
}
}
#[cfg(test)]
mod lease_reload_tests {
use super::*;
use crate::libsignal::protocol::{
ChainKey, IdentityKey, KeyPair, RootKey, SenderKeyStore, SessionState,
create_sender_key_distribution_message, group_decrypt, group_encrypt,
process_sender_key_distribution_message,
};
use crate::store::in_memory::InMemoryBackend;
struct CachedSenderKeyStore<'a> {
cache: &'a SignalStoreCache,
backend: &'a InMemoryBackend,
}
#[async_trait::async_trait]
impl SenderKeyStore for CachedSenderKeyStore<'_> {
async fn store_sender_key(
&mut self,
name: &SenderKeyName,
record: SenderKeyRecord,
) -> crate::libsignal::protocol::error::Result<()> {
self.cache.put_sender_key(name, record).await;
Ok(())
}
async fn load_sender_key(
&self,
name: &SenderKeyName,
) -> crate::libsignal::protocol::error::Result<Option<SenderKeyRecord>> {
Ok(self
.cache
.get_sender_key(name, self.backend)
.await
.expect("test backend")
.map(|record| (*record).clone()))
}
}
fn sender_key_name() -> SenderKeyName {
SenderKeyName::from_parts("group@g.us", "15550001000@s.whatsapp.net:0")
}
pub(super) fn leased_session() -> SessionRecord {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let local = IdentityKey::new(KeyPair::generate(&mut rng).public_key);
let remote = IdentityKey::new(KeyPair::generate(&mut rng).public_key);
let base_key = KeyPair::generate(&mut rng).public_key;
let mut state = SessionState::new(3, &local, &remote, &RootKey::new([0; 32]), &base_key);
state.set_sender_chain(&KeyPair::generate(&mut rng), &ChainKey::new([1; 32], 0));
let mut record = SessionRecord::new(state);
record.reserve_sender_chain_counters(0);
record
}
fn session_chain_index(record: &SessionRecord) -> u32 {
record
.session_state()
.expect("session")
.get_sender_chain_key()
.expect("sender chain")
.index()
}
#[tokio::test]
async fn post_flush_clear_preserves_only_live_checkouts() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
let active = ProtocolAddress::new("15550001007", 1.into());
let idle = ProtocolAddress::new("15550001008", 1.into());
cache.put_session(&active, leased_session()).await;
cache.put_session(&idle, leased_session()).await;
cache.flush(&backend).await.expect("flush");
let (record, checkout) = cache.checkout_session(&active, &backend).await.unwrap();
cache.remove_prekey(7, active.as_str()).await;
cache.clear_after_flush().await;
{
let state = cache.sessions.lock().await;
assert!(matches!(
state.cache.get(active.as_str()),
Some(SessionEntry::CheckedOut { .. })
));
assert!(!state.cache.contains_key(idle.as_str()));
}
assert!(cache.removed_prekeys.lock().await.contains_key(&7));
assert!(matches!(
cache.restore_session_from_checkout(
&active,
record.expect("checked-out record"),
checkout,
true,
),
SessionCheckoutStoreResult::Stored
));
}
#[tokio::test]
async fn dm_clean_reload_is_exact_but_new_cache_burns_the_lease() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xA1; 16],
);
let address = ProtocolAddress::new("15550001001", 1.into());
cache.put_session(&address, leased_session()).await;
cache.flush(&backend).await.expect("flush");
cache.clear_after_flush().await;
let clean = cache
.get_session(&address, &backend)
.await
.expect("cache load")
.expect("session");
assert_eq!(session_chain_index(&clean), 0);
let replacement = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xB2; 16],
);
let recovered = replacement
.get_session(&address, &backend)
.await
.expect("recovery load")
.expect("session");
assert_eq!(
session_chain_index(&recovered),
crate::libsignal::protocol::consts::SENDER_CHAIN_RESERVATION_BATCH
);
}
#[tokio::test]
async fn an_unreadable_session_row_is_reported_absent_so_recovery_can_replace_it() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xA1; 16],
);
let address = ProtocolAddress::new("15550001009", 1.into());
let mut stranded = leased_session();
stranded.reserve_sender_chain_counters(
crate::libsignal::protocol::consts::MAX_RESERVATION_FAST_FORWARD,
);
assert_eq!(session_chain_index(&stranded), 0);
cache.put_session(&address, stranded).await;
cache.flush(&backend).await.expect("flush");
cache.clear_after_flush().await;
assert!(
cache
.get_session(&address, &backend)
.await
.expect("live reload")
.is_some()
);
let restarted = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xB2; 16],
);
assert!(
restarted
.get_session(&address, &backend)
.await
.expect("an unreadable row must not fail the load")
.is_none()
);
assert!(
!restarted
.has_session(&address, &backend)
.await
.expect("has_session"),
"the quarantined address must look session-less so ensure_e2e_sessions rebuilds it"
);
}
#[tokio::test]
async fn a_cold_existence_probe_does_not_report_a_quarantined_row_as_present() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xA1; 16],
);
let address = ProtocolAddress::new("15550001010", 1.into());
let mut stranded = leased_session();
stranded.reserve_sender_chain_counters(
crate::libsignal::protocol::consts::MAX_RESERVATION_FAST_FORWARD,
);
cache.put_session(&address, stranded).await;
cache.flush(&backend).await.expect("flush");
let restarted = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xB2; 16],
);
assert!(
!restarted
.has_session(&address, &backend)
.await
.expect("a quarantined row must not fail the probe"),
"a row the next checkout would discard must not be reported present"
);
assert!(
restarted
.get_session(&address, &backend)
.await
.expect("checkout")
.is_none()
);
}
#[tokio::test]
async fn incomplete_session_flush_retains_newer_state_and_fails_closed_on_recovery() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xA1; 16],
);
let address = ProtocolAddress::new("15550001002", 1.into());
cache.put_session(&address, leased_session()).await;
cache.flush(&backend).await.expect("initial flush");
let mut advanced = cache
.get_session(&address, &backend)
.await
.expect("cache load")
.expect("session");
let next = advanced
.session_state()
.expect("session")
.get_sender_chain_key()
.expect("sender chain")
.next_chain_key()
.expect("chain advance");
advanced
.session_state_mut()
.expect("session")
.set_sender_chain_key(&next)
.expect("chain update");
cache.put_session(&address, advanced).await;
let checked_out = cache
.get_session(&address, &backend)
.await
.expect("cache checkout")
.expect("session");
cache.flush(&backend).await.expect("skipped flush");
cache.clear_after_flush().await;
{
let state = cache.sessions.lock().await;
assert_eq!(state.incarnation, [0xA1; 16]);
assert!(state.dirty.contains(address.as_str()));
assert!(matches!(
state.cache.get(address.as_str()),
Some(SessionEntry::CheckedOut { .. })
));
}
let replacement = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xB2; 16],
);
let recovered = replacement
.get_session(&address, &backend)
.await
.expect("recovery load")
.expect("session");
assert_eq!(
session_chain_index(&recovered),
crate::libsignal::protocol::consts::SENDER_CHAIN_RESERVATION_BATCH
);
cache.put_session(&address, checked_out).await;
cache.flush(&backend).await.expect("retry flush");
cache.clear_after_flush().await;
let exact = cache
.get_session(&address, &backend)
.await
.expect("exact reload")
.expect("session");
assert_eq!(session_chain_index(&exact), 1);
}
#[tokio::test]
async fn repeated_clean_reloads_keep_group_messages_within_forward_jump_limit() {
let sender_backend = InMemoryBackend::new();
let sender_cache = SignalStoreCache::new();
let mut sender = CachedSenderKeyStore {
cache: &sender_cache,
backend: &sender_backend,
};
let receiver_backend = InMemoryBackend::new();
let receiver_cache = SignalStoreCache::new();
let mut receiver = CachedSenderKeyStore {
cache: &receiver_cache,
backend: &receiver_backend,
};
let name = sender_key_name();
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let skdm = create_sender_key_distribution_message(&name, &mut sender, &mut rng)
.await
.expect("sender setup");
process_sender_key_distribution_message(&name, &skdm, &mut receiver)
.await
.expect("receiver setup");
let mut last = None;
for expected_iteration in 0..=32 {
let message = group_encrypt(&mut sender, &name, b"payload", &mut rng)
.await
.expect("group encrypt");
assert_eq!(message.iteration(), expected_iteration);
last = Some(message);
sender_cache.flush(&sender_backend).await.expect("flush");
sender_cache.clear_after_flush().await;
}
let plaintext = group_decrypt(last.expect("message").serialized(), &mut receiver, &name)
.await
.expect("a peer may miss every preceding message");
assert_eq!(plaintext, b"payload");
}
#[tokio::test]
async fn clean_sender_key_eviction_does_not_burn_a_lease() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
let mut store = CachedSenderKeyStore {
cache: &cache,
backend: &backend,
};
let name = sender_key_name();
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
create_sender_key_distribution_message(&name, &mut store, &mut rng)
.await
.expect("sender setup");
let first = group_encrypt(&mut store, &name, b"first", &mut rng)
.await
.expect("first send");
assert_eq!(first.iteration(), 0);
cache.flush(&backend).await.expect("flush");
assert!(
cache
.sender_keys
.lock()
.await
.cache
.remove(name.cache_key())
.is_some()
);
let second = group_encrypt(&mut store, &name, b"second", &mut rng)
.await
.expect("send after eviction");
assert_eq!(second.iteration(), 1);
}
#[tokio::test]
async fn dirty_sender_key_stays_resident_while_recovery_fails_closed() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xA1; 16],
);
let mut store = CachedSenderKeyStore {
cache: &cache,
backend: &backend,
};
let name = sender_key_name();
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
create_sender_key_distribution_message(&name, &mut store, &mut rng)
.await
.expect("sender setup");
let first = group_encrypt(&mut store, &name, b"first", &mut rng)
.await
.expect("first send");
assert_eq!(first.iteration(), 0);
cache.flush(&backend).await.expect("flush");
let unflushed = group_encrypt(&mut store, &name, b"unflushed", &mut rng)
.await
.expect("unflushed send");
assert_eq!(unflushed.iteration(), 1);
cache.clear_after_flush().await;
{
let state = cache.sender_keys.lock().await;
assert_eq!(state.incarnation, [0xA1; 16]);
assert!(state.dirty.contains(name.cache_key()));
assert!(state.cache.contains_key(name.cache_key()));
}
let resumed = group_encrypt(&mut store, &name, b"resumed", &mut rng)
.await
.expect("resident send");
assert_eq!(resumed.iteration(), 2);
let replacement = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xB2; 16],
);
let mut recovered_store = CachedSenderKeyStore {
cache: &replacement,
backend: &backend,
};
let recovered = group_encrypt(&mut recovered_store, &name, b"recovered", &mut rng)
.await
.expect("recovery send");
assert_eq!(
recovered.iteration(),
crate::libsignal::protocol::consts::SENDER_CHAIN_RESERVATION_BATCH
);
cache.flush(&backend).await.expect("retry flush");
cache.clear_after_flush().await;
let exact = group_encrypt(&mut store, &name, b"exact", &mut rng)
.await
.expect("exact reload");
assert_eq!(exact.iteration(), 3);
}
}
#[cfg(test)]
mod pre_wire_gate_tests {
use super::*;
use std::sync::atomic::{AtomicBool, Ordering};
use crate::libsignal::store::sender_key_name::SenderKeyName;
use crate::store::in_memory::InMemoryBackend;
use async_lock::Barrier;
fn addr(user: &str) -> ProtocolAddress {
ProtocolAddress::new(user, 1.into())
}
fn leased_record() -> SessionRecord {
let mut record = SessionRecord::new_fresh();
record.reserve_sender_chain_counters(0);
record
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum DeleteTarget {
Session,
SenderKey,
}
struct DeleteBarrierBackend {
inner: InMemoryBackend,
target: DeleteTarget,
entered: Barrier,
release: Barrier,
fail_delete: AtomicBool,
}
impl DeleteBarrierBackend {
fn new(target: DeleteTarget) -> Self {
Self {
inner: InMemoryBackend::new(),
target,
entered: Barrier::new(2),
release: Barrier::new(2),
fail_delete: AtomicBool::new(true),
}
}
async fn gate_delete(&self, target: DeleteTarget) -> crate::store::error::Result<()> {
if self.target != target {
return Ok(());
}
self.entered.wait().await;
self.release.wait().await;
if self.fail_delete.load(Ordering::Acquire) {
return Err(crate::store::error::StoreError::Validation(
"simulated delete failure".to_string(),
));
}
Ok(())
}
}
#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
impl SignalStore for DeleteBarrierBackend {
async fn put_identity(
&self,
address: &str,
key: [u8; 32],
) -> crate::store::error::Result<()> {
self.inner.put_identity(address, key).await
}
async fn load_identity(
&self,
address: &str,
) -> crate::store::error::Result<Option<[u8; 32]>> {
self.inner.load_identity(address).await
}
async fn delete_identity(&self, address: &str) -> crate::store::error::Result<()> {
self.inner.delete_identity(address).await
}
async fn get_session(
&self,
address: &str,
) -> crate::store::error::Result<Option<bytes::Bytes>> {
self.inner.get_session(address).await
}
async fn put_session(
&self,
address: &str,
session: &[u8],
) -> crate::store::error::Result<()> {
self.inner.put_session(address, session).await
}
async fn delete_session(&self, address: &str) -> crate::store::error::Result<()> {
self.gate_delete(DeleteTarget::Session).await?;
self.inner.delete_session(address).await
}
async fn store_prekey(
&self,
id: u32,
record: &[u8],
uploaded: bool,
) -> crate::store::error::Result<()> {
self.inner.store_prekey(id, record, uploaded).await
}
async fn load_prekey(&self, id: u32) -> crate::store::error::Result<Option<bytes::Bytes>> {
self.inner.load_prekey(id).await
}
async fn mark_prekeys_uploaded(&self, ids: &[u32]) -> crate::store::error::Result<()> {
self.inner.mark_prekeys_uploaded(ids).await
}
async fn remove_prekey(&self, id: u32) -> crate::store::error::Result<()> {
self.inner.remove_prekey(id).await
}
async fn get_max_prekey_id(&self) -> crate::store::error::Result<u32> {
self.inner.get_max_prekey_id().await
}
async fn store_signed_prekey(
&self,
id: u32,
record: &[u8],
) -> crate::store::error::Result<()> {
self.inner.store_signed_prekey(id, record).await
}
async fn load_signed_prekey(
&self,
id: u32,
) -> crate::store::error::Result<Option<Vec<u8>>> {
self.inner.load_signed_prekey(id).await
}
async fn load_all_signed_prekeys(
&self,
) -> crate::store::error::Result<Vec<(u32, Vec<u8>)>> {
self.inner.load_all_signed_prekeys().await
}
async fn remove_signed_prekey(&self, id: u32) -> crate::store::error::Result<()> {
self.inner.remove_signed_prekey(id).await
}
async fn put_sender_key(
&self,
address: &str,
record: &[u8],
) -> crate::store::error::Result<()> {
self.inner.put_sender_key(address, record).await
}
async fn get_sender_key(
&self,
address: &str,
) -> crate::store::error::Result<Option<Vec<u8>>> {
self.inner.get_sender_key(address).await
}
async fn delete_sender_key(&self, address: &str) -> crate::store::error::Result<()> {
self.gate_delete(DeleteTarget::SenderKey).await?;
self.inner.delete_sender_key(address).await
}
}
async fn run_gated_flush(
cache: Arc<SignalStoreCache>,
backend: Arc<DeleteBarrierBackend>,
) -> Result<()> {
let flush_cache = cache.clone();
let flush_backend = backend.clone();
let task = tokio::spawn(async move { flush_cache.flush(flush_backend.as_ref()).await });
backend.entered.wait().await;
backend.release.wait().await;
task.await.expect("flush task")
}
#[tokio::test]
async fn session_lease_gates_until_a_successful_flush() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
cache
.put_session(&addr("15550000001"), SessionRecord::new_fresh())
.await;
assert!(
!cache.needs_pre_wire_flush().await,
"a dirty session without a raised lease must not gate the wire"
);
cache
.put_session(&addr("15550000002"), leased_record())
.await;
assert!(cache.needs_pre_wire_flush().await);
cache.flush(&backend).await.unwrap();
assert!(
!cache.needs_pre_wire_flush().await,
"a persisted lease releases the gate"
);
}
#[tokio::test]
async fn failed_flush_keeps_the_gate_closed() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
cache
.put_session(&addr("15550000003"), leased_record())
.await;
backend.set_fail_session_writes(true);
assert!(cache.flush(&backend).await.is_err());
assert!(
cache.needs_pre_wire_flush().await,
"an unpersisted lease must keep gating the wire"
);
backend.set_fail_session_writes(false);
cache.flush(&backend).await.unwrap();
assert!(!cache.needs_pre_wire_flush().await);
}
#[tokio::test]
async fn checked_out_session_keeps_its_lease_pending_across_a_flush() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
let a = addr("15550000004");
cache.put_session(&a, leased_record()).await;
let taken = cache.get_session(&a, &backend).await.unwrap().unwrap();
cache.flush(&backend).await.unwrap();
assert!(
cache.needs_pre_wire_flush().await,
"a checked-out lease was not persisted and must keep the gate closed"
);
cache.put_session(&a, taken).await;
cache.flush(&backend).await.unwrap();
assert!(!cache.needs_pre_wire_flush().await);
}
#[tokio::test]
async fn only_encrypt_marked_sender_keys_gate_the_wire() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
let name = SenderKeyName::from_parts("g@g.us", "u@s.whatsapp.net:0");
cache
.put_sender_key(&name, SenderKeyRecord::new_empty())
.await;
assert!(
!cache.needs_pre_wire_flush().await,
"a decrypt-side sender-key write must not gate the wire"
);
let mut outbound = SenderKeyRecord::new_empty();
outbound.mark_wire_gated();
cache.put_sender_key(&name, outbound).await;
assert!(cache.needs_pre_wire_flush().await);
cache.flush(&backend).await.unwrap();
assert!(!cache.needs_pre_wire_flush().await);
}
#[tokio::test]
async fn failed_flush_keeps_the_sender_key_gate_closed() {
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::new();
let name = SenderKeyName::from_parts("g@g.us", "u@s.whatsapp.net:0");
let mut outbound = SenderKeyRecord::new_empty();
outbound.mark_wire_gated();
cache.put_sender_key(&name, outbound).await;
backend.set_fail_sender_key_writes(true);
assert!(cache.flush(&backend).await.is_err());
assert!(
cache.needs_pre_wire_flush().await,
"an unpersisted sender-key advance must keep gating the wire"
);
backend.set_fail_sender_key_writes(false);
cache.flush(&backend).await.unwrap();
assert!(!cache.needs_pre_wire_flush().await);
}
#[tokio::test]
async fn session_tombstone_keeps_gate_until_delete_is_durable() {
let cache = Arc::new(SignalStoreCache::new());
let backend = Arc::new(DeleteBarrierBackend::new(DeleteTarget::Session));
let address = addr("15550000007");
backend
.inner
.put_session(address.as_str(), b"durable session")
.await
.unwrap();
cache.put_session(&address, leased_record()).await;
cache.delete_session(&address).await;
assert!(cache.needs_pre_wire_flush().await);
assert!(
run_gated_flush(cache.clone(), backend.clone())
.await
.is_err()
);
assert!(cache.needs_pre_wire_flush().await);
assert!(
backend
.inner
.get_session(address.as_str())
.await
.unwrap()
.is_some()
);
backend.fail_delete.store(false, Ordering::Release);
run_gated_flush(cache.clone(), backend.clone())
.await
.unwrap();
assert!(!cache.needs_pre_wire_flush().await);
assert!(
backend
.inner
.get_session(address.as_str())
.await
.unwrap()
.is_none()
);
}
#[tokio::test]
async fn sender_key_tombstone_keeps_gate_until_delete_is_durable() {
let cache = Arc::new(SignalStoreCache::new());
let backend = Arc::new(DeleteBarrierBackend::new(DeleteTarget::SenderKey));
let name = SenderKeyName::from_parts("g@g.us", "u@s.whatsapp.net:0");
backend
.inner
.put_sender_key(name.cache_key(), b"durable sender key")
.await
.unwrap();
let mut outbound = SenderKeyRecord::new_empty();
outbound.mark_wire_gated();
cache.put_sender_key(&name, outbound).await;
cache.delete_sender_key(name.cache_key()).await;
assert!(cache.needs_pre_wire_flush().await);
assert!(
run_gated_flush(cache.clone(), backend.clone())
.await
.is_err()
);
assert!(cache.needs_pre_wire_flush().await);
assert!(
backend
.inner
.get_sender_key(name.cache_key())
.await
.unwrap()
.is_some()
);
backend.fail_delete.store(false, Ordering::Release);
run_gated_flush(cache.clone(), backend.clone())
.await
.unwrap();
assert!(!cache.needs_pre_wire_flush().await);
assert!(
backend
.inner
.get_sender_key(name.cache_key())
.await
.unwrap()
.is_none()
);
}
#[tokio::test]
async fn durable_sender_key_delete_does_not_block_unrelated_chains() {
let cache = Arc::new(SignalStoreCache::new());
let backend = Arc::new(DeleteBarrierBackend::new(DeleteTarget::SenderKey));
backend.fail_delete.store(false, Ordering::Release);
let target = SenderKeyName::from_parts("g1@g.us", "u@s.whatsapp.net:0");
let unrelated = SenderKeyName::from_parts("g2@g.us", "u@s.whatsapp.net:0");
cache
.put_sender_key(&target, SenderKeyRecord::new_empty())
.await;
let target_lock = cache.sender_key_lock(&target).await;
let deletion = tokio::spawn({
let cache = cache.clone();
let backend = backend.clone();
async move {
cache
.delete_sender_key_durable(&target, backend.as_ref())
.await
}
});
backend.entered.wait().await;
assert!(
target_lock.try_lock().is_none(),
"the target chain must remain serialized during backend deletion"
);
tokio::time::timeout(
std::time::Duration::from_secs(1),
cache.put_sender_key(&unrelated, SenderKeyRecord::new_empty()),
)
.await
.expect("backend latency for one chain must not hold the global cache lock");
backend.release.wait().await;
deletion
.await
.expect("delete task")
.expect("durable delete");
assert!(
cache
.get_sender_key(&unrelated, backend.as_ref())
.await
.unwrap()
.is_some(),
"unrelated state must remain available"
);
}
#[tokio::test]
async fn clear_after_flush_retains_every_post_flush_write_and_wire_gate() {
const PREKEY_ID: u32 = 7001;
let backend = InMemoryBackend::new();
let cache = SignalStoreCache::with_max_entries_and_incarnation(
DEFAULT_MAX_CACHE_ENTRIES,
[0xA1; 16],
);
let address = addr("15550000005");
let name = SenderKeyName::from_parts("g@g.us", "u@s.whatsapp.net:0");
cache.flush(&backend).await.unwrap();
cache.put_session(&address, leased_record()).await;
cache.put_identity(&address, &[7; 32]).await;
backend
.store_prekey(PREKEY_ID, b"prekey", false)
.await
.unwrap();
cache.remove_prekey(PREKEY_ID, address.as_str()).await;
let mut outbound = SenderKeyRecord::new_empty();
outbound.mark_wire_gated();
cache.put_sender_key(&name, outbound).await;
cache.clear_after_flush().await;
assert!(cache.needs_pre_wire_flush().await);
{
let sessions = cache.sessions.lock().await;
assert_eq!(sessions.incarnation, [0xA1; 16]);
assert!(sessions.dirty.contains(address.as_str()));
assert!(sessions.reservation_pending.contains(address.as_str()));
}
{
let identities = cache.identities.lock().await;
assert!(identities.dirty.contains(address.as_str()));
}
{
let sender_keys = cache.sender_keys.lock().await;
assert_eq!(sender_keys.incarnation, [0xA1; 16]);
assert!(sender_keys.dirty.contains(name.cache_key()));
assert!(sender_keys.wire_gate_pending.contains(name.cache_key()));
}
assert!(cache.removed_prekeys.lock().await.contains_key(&PREKEY_ID));
cache.flush(&backend).await.unwrap();
assert!(!cache.needs_pre_wire_flush().await);
assert!(
backend
.get_session(address.as_str())
.await
.unwrap()
.is_some()
);
assert_eq!(
backend.load_identity(address.as_str()).await.unwrap(),
Some([7; 32])
);
assert!(
backend
.get_sender_key(name.cache_key())
.await
.unwrap()
.is_some()
);
assert!(backend.load_prekey(PREKEY_ID).await.unwrap().is_none());
}
#[tokio::test]
async fn clear_drops_a_pending_tombstone_gate() {
let cache = SignalStoreCache::new();
let a = addr("15550000006");
cache.put_session(&a, leased_record()).await;
cache.delete_session(&a).await;
assert!(cache.needs_pre_wire_flush().await);
cache.clear().await;
assert!(!cache.needs_pre_wire_flush().await);
}
}
#[cfg(test)]
#[path = "signal_cache_durability_chaos.rs"]
mod durability_chaos_tests;