diff --git a/src/client.rs b/src/client.rs index 710ccf9fa..4e3ed97e4 100644 --- a/src/client.rs +++ b/src/client.rs @@ -28,9 +28,9 @@ use log::{debug, error, info, trace, warn}; use rand::RngCore; use scopeguard; use std::collections::{HashMap, HashSet}; +use std::sync::{Arc, OnceLock, Weak}; use wacore_binary::jid::Jid; -use std::sync::Arc; use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicUsize, Ordering}; use thiserror::Error; @@ -40,7 +40,10 @@ use wacore::appstate::patch_decode::WAPatchName; use wacore::client::context::GroupInfo; use waproto::whatsapp as wa; -use crate::socket::{NoiseSocket, SocketError, error::EncryptSendError}; +use crate::socket::{ + NoiseSocket, SocketError, + error::{EncryptSendError, EncryptSendErrorKind}, +}; use crate::sync_task::MajorSyncTask; /// Type alias for chatstate event handler functions. @@ -89,11 +92,13 @@ pub(crate) struct RetryMetrics { pub struct Client { pub(crate) core: wacore::client::CoreClient, + self_weak: OnceLock>, pub(crate) persistence_manager: Arc, pub(crate) media_conn: Arc>>, pub(crate) is_logged_in: Arc, + connection_alive: Arc, pub(crate) is_connecting: Arc, pub(crate) is_running: Arc, pub(crate) shutdown_notifier: Arc, @@ -232,6 +237,10 @@ pub struct Client { } impl Client { + pub(crate) fn shared(&self) -> Option> { + self.self_weak.get().and_then(Weak::upgrade) + } + /// Enable or disable skipping of history sync notifications at runtime. /// /// When enabled, the client will acknowledge incoming history sync @@ -261,9 +270,11 @@ impl Client { let this = Self { core, + self_weak: OnceLock::new(), persistence_manager: persistence_manager.clone(), media_conn: Arc::new(RwLock::new(None)), is_logged_in: Arc::new(AtomicBool::new(false)), + connection_alive: Arc::new(AtomicBool::new(false)), is_connecting: Arc::new(AtomicBool::new(false)), is_running: Arc::new(AtomicBool::new(false)), shutdown_notifier: Arc::new(Notify::new()), @@ -371,6 +382,9 @@ impl Client { }; let arc = Arc::new(this); + if arc.self_weak.set(Arc::downgrade(&arc)).is_err() { + warn!("client self reference was already set"); + } // Warm up the LID-PN cache from persistent storage let warm_up_arc = arc.clone(); @@ -559,6 +573,7 @@ impl Client { // handle_success will properly process the stanza even if // a previous connection's post-login task bailed out early. self.is_logged_in.store(false, Ordering::Relaxed); + self.connection_alive.store(false, Ordering::Release); self.offline_sync_completed.store(false, Ordering::Relaxed); let version_future = crate::version::resolve_and_update_version( @@ -585,6 +600,7 @@ impl Client { *self.transport.lock().await = Some(transport); *self.transport_events.lock().await = Some(transport_events); *self.noise_socket.lock().await = Some(noise_socket); + self.connection_alive.store(true, Ordering::Release); // Notify waiters that socket is ready (before login) self.socket_ready_notifier.notify_waiters(); @@ -598,6 +614,7 @@ impl Client { pub async fn disconnect(self: &Arc) { info!("Disconnecting client intentionally."); self.expected_disconnect.store(true, Ordering::Relaxed); + self.connection_alive.store(false, Ordering::Release); self.is_running.store(false, Ordering::Relaxed); self.shutdown_notifier.notify_waiters(); @@ -609,6 +626,7 @@ impl Client { async fn cleanup_connection_state(&self) { self.is_logged_in.store(false, Ordering::Relaxed); + self.connection_alive.store(false, Ordering::Release); *self.transport.lock().await = None; *self.transport_events.lock().await = None; *self.noise_socket.lock().await = None; @@ -898,20 +916,44 @@ impl Client { /// Uses Arc to avoid cloning when spawning the async task. async fn maybe_deferred_ack(self: &Arc, node: Arc) { if self.synchronous_ack { - if let Err(e) = self.send_ack_for(&node).await { + if let Err(e) = self.send_ack_for(&node).await + && !Self::is_connection_closed_error(&e) + { warn!("Failed to send ack: {e:?}"); } } else { let this = self.clone(); // Node is already in Arc - just clone the Arc (cheap), not the Node tokio::spawn(async move { - if let Err(e) = this.send_ack_for(&node).await { + if let Err(e) = this.send_ack_for(&node).await + && !Self::is_connection_closed_error(&e) + { warn!("Failed to send ack: {e:?}"); } }); } } + fn is_terminal_transport_send_error(e: &EncryptSendError) -> bool { + match e.kind { + EncryptSendErrorKind::ChannelClosed | EncryptSendErrorKind::Join => true, + EncryptSendErrorKind::Transport => e + .source + .downcast_ref::() + .is_some_and(|err| *err == wacore::net::TransportSendError::ConnectionClosed), + _ => false, + } + } + + fn is_connection_closed_error(e: &ClientError) -> bool { + match e { + ClientError::NotConnected => true, + ClientError::Socket(SocketError::SocketClosed) => true, + ClientError::EncryptSend(send_err) => Self::is_terminal_transport_send_error(send_err), + _ => false, + } + } + /// Build and send an node corresponding to the given stanza. async fn send_ack_for(&self, node: &Node) -> Result<(), ClientError> { if !self.is_connected() || self.expected_disconnect.load(Ordering::Relaxed) { @@ -1872,9 +1914,11 @@ impl Client { } pub fn is_connected(&self) -> bool { - self.noise_socket - .try_lock() - .is_ok_and(|guard| guard.is_some()) + self.connection_alive.load(Ordering::Acquire) + && self + .noise_socket + .try_lock() + .is_ok_and(|guard| guard.is_some()) } pub fn is_logged_in(&self) -> bool { @@ -2012,6 +2056,12 @@ impl Client { } pub async fn send_node(&self, node: Node) -> Result<(), ClientError> { + if !self.connection_alive.load(Ordering::Acquire) + || self.expected_disconnect.load(Ordering::Relaxed) + { + return Err(ClientError::NotConnected); + } + let noise_socket_arc = { self.noise_socket.lock().await.clone() }; let noise_socket = match noise_socket_arc { Some(socket) => socket, @@ -2044,6 +2094,9 @@ impl Client { { Ok(bufs) => bufs, Err(mut e) => { + if Self::is_terminal_transport_send_error(&e) { + self.connection_alive.store(false, Ordering::Release); + } let p_buf = std::mem::take(&mut e.plaintext_buf); let mut pool = self.plaintext_buffer_pool.lock().await; if p_buf.capacity() <= MAX_POOLED_BUFFER_CAP { @@ -2364,6 +2417,39 @@ mod tests { ); } + #[test] + fn test_terminal_transport_send_error_detection() { + let err = EncryptSendError::transport( + wacore::net::TransportSendError::ConnectionClosed, + Vec::new(), + Vec::new(), + ); + + assert!( + Client::is_terminal_transport_send_error(&err), + "closed websocket send errors must be treated as terminal" + ); + let client_err: ClientError = err.into(); + assert!( + Client::is_connection_closed_error(&client_err), + "terminal send errors should be classified as closed connection errors" + ); + } + + #[test] + fn test_non_terminal_transport_send_error_detection() { + let err = EncryptSendError::transport( + anyhow::anyhow!("temporary websocket write timeout"), + Vec::new(), + Vec::new(), + ); + + assert!( + !Client::is_terminal_transport_send_error(&err), + "temporary transport errors should not always be classified as terminal" + ); + } + /// Test that the lid_pn_cache correctly stores and retrieves LID mappings. /// /// This is critical for the LID-PN session mismatch fix. When we receive a message diff --git a/src/main.rs b/src/main.rs index 0e79deb5c..995b54f10 100644 --- a/src/main.rs +++ b/src/main.rs @@ -214,12 +214,6 @@ fn main() { Event::Connected(_) => { info!("✅ Bot connected successfully!"); } - Event::Receipt(receipt) => { - info!( - "Got receipt for message(s) {:?}, type: {:?}", - receipt.message_ids, receipt.r#type - ); - } Event::LoggedOut(_) => { error!("❌ Bot was logged out!"); } diff --git a/src/message.rs b/src/message.rs index 511d625c4..c6961fd8f 100644 --- a/src/message.rs +++ b/src/message.rs @@ -1,5 +1,5 @@ use crate::client::Client; -use crate::store::signal_adapter::SignalProtocolStoreAdapter; +use crate::store::signal_adapter::BatchedSignalProtocolStoreAdapter; use crate::types::events::Event; use crate::types::message::MessageInfo; use chrono::DateTime; @@ -34,6 +34,46 @@ const MAX_DECRYPT_RETRIES: u8 = 5; /// WhatsApp Web logs metrics when retry count exceeds this value. const HIGH_RETRY_COUNT_THRESHOLD: u8 = 3; +/// Runs Signal message decryption on the blocking thread pool. +/// +/// With the batched cache adapter, store operations (load/store session, identity) +/// resolve from in-memory HashMaps, making the work effectively CPU-bound (crypto). +/// On rare cache misses the inner store access will block the thread via `block_on`, +/// which is acceptable since the blocking pool has many threads. +async fn decrypt_session_message_with_blocking( + parsed_message: wacore::libsignal::protocol::CiphertextMessage, + signal_address: wacore::libsignal::protocol::ProtocolAddress, + adapter: BatchedSignalProtocolStoreAdapter, +) -> Result< + ( + BatchedSignalProtocolStoreAdapter, + wacore::libsignal::protocol::CiphertextMessage, + Result, SignalProtocolError>, + ), + tokio::task::JoinError, +> { + let runtime = tokio::runtime::Handle::current(); + tokio::task::spawn_blocking(move || { + let mut adapter = adapter; + let mut rng = rand::rngs::OsRng.unwrap_err(); + let decrypt_res = runtime.block_on(async { + message_decrypt( + &parsed_message, + &signal_address, + &mut adapter.session_store, + &mut adapter.identity_store, + &mut adapter.pre_key_store, + &adapter.signed_pre_key_store, + &mut rng, + UsePQRatchet::No, + ) + .await + }); + (adapter, parsed_message, decrypt_res) + }) + .await +} + /// Retry reason codes matching WhatsApp Web's RetryReason enum. /// These are included in the retry receipt to help the sender understand /// why the message couldn't be decrypted. @@ -650,8 +690,7 @@ impl Client { let _session_guard = session_mutex.lock().await; let mut adapter = - SignalProtocolStoreAdapter::new(self.persistence_manager.get_device_arc().await); - let rng = rand::rngs::OsRng; + BatchedSignalProtocolStoreAdapter::new(self.persistence_manager.get_device_arc().await); let mut any_success = false; let mut any_duplicate = false; let mut dispatched_undecryptable = false; @@ -722,17 +761,26 @@ impl Client { } } - let decrypt_res = message_decrypt( - &parsed_message, - &signal_address, - &mut adapter.session_store, - &mut adapter.identity_store, - &mut adapter.pre_key_store, - &adapter.signed_pre_key_store, - &mut rng.unwrap_err(), - UsePQRatchet::No, - ) - .await; + let (next_adapter, parsed_message, decrypt_res) = + match decrypt_session_message_with_blocking( + parsed_message, + signal_address.clone(), + adapter, + ) + .await + { + Ok(result) => result, + Err(e) => { + log::error!( + "spawn_blocking failed while decrypting message {} from {}: {}", + info.id, + info.source.sender, + e + ); + return (any_success, any_duplicate, dispatched_undecryptable); + } + }; + adapter = next_adapter; match decrypt_res { Ok(padded_plaintext) => { @@ -794,6 +842,7 @@ impl Client { log::warn!("Failed to delete old identity for {}: {:?}", address, err); } else { log::info!("Successfully cleared old identity for {}", address); + adapter.invalidate_identity(address).await; } // Re-attempt decryption with the new identity @@ -803,17 +852,26 @@ impl Client { address ); - let retry_decrypt_res = message_decrypt( - &parsed_message, - &signal_address, - &mut adapter.session_store, - &mut adapter.identity_store, - &mut adapter.pre_key_store, - &adapter.signed_pre_key_store, - &mut rng.unwrap_err(), - UsePQRatchet::No, - ) - .await; + let (next_adapter, _parsed_message, retry_decrypt_res) = + match decrypt_session_message_with_blocking( + parsed_message, + signal_address.clone(), + adapter, + ) + .await + { + Ok(result) => result, + Err(e) => { + log::error!( + "spawn_blocking failed while retrying decrypt for message {} from {}: {}", + info.id, + info.source.sender, + e + ); + return (any_success, any_duplicate, dispatched_undecryptable); + } + }; + adapter = next_adapter; match retry_decrypt_res { Ok(padded_plaintext) => { @@ -914,11 +972,8 @@ impl Client { info.source.sender ); - // Delete the stale session - let device_arc = self.persistence_manager.get_device_arc().await; - let device_guard = device_arc.write().await; - let address_str = signal_address.to_string(); - if let Err(err) = device_guard.backend.delete_session(&address_str).await { + // Delete the stale session (also invalidates the adapter cache) + if let Err(err) = adapter.delete_session(&signal_address).await { log::warn!( "Failed to delete stale session for {}: {:?}", signal_address, @@ -930,7 +985,6 @@ impl Client { signal_address ); } - drop(device_guard); // Send retry receipt so the sender resends with a PreKeySignalMessage dispatched_undecryptable = @@ -976,6 +1030,9 @@ impl Client { } } } + if let Err(flush_err) = adapter.flush().await { + log::warn!("Failed to flush cached Signal stores after decrypt batch: {flush_err}"); + } (any_success, any_duplicate, dispatched_undecryptable) } diff --git a/src/receipt.rs b/src/receipt.rs index b7be47c95..f36bf993a 100644 --- a/src/receipt.rs +++ b/src/receipt.rs @@ -1,7 +1,7 @@ use crate::client::Client; use crate::types::events::{Event, Receipt}; use crate::types::presence::ReceiptType; -use log::info; +use log::{debug, info}; use std::collections::HashMap; use std::sync::Arc; use wacore_binary::builder::NodeBuilder; @@ -25,7 +25,7 @@ impl Client { let receipt_type = ReceiptType::from(receipt_type_str.to_string()); - info!("Received receipt type '{receipt_type:?}' for message {id} from {from}"); + debug!("Received receipt type '{receipt_type:?}' for message {id} from {from}"); let from_clone = from.clone(); let sender = if from.is_group() { diff --git a/src/send.rs b/src/send.rs index 171291965..4f7f0fe47 100644 --- a/src/send.rs +++ b/src/send.rs @@ -1,5 +1,5 @@ use crate::client::Client; -use crate::store::signal_adapter::SignalProtocolStoreAdapter; +use crate::store::signal_adapter::BatchedSignalProtocolStoreAdapter; use anyhow::anyhow; use wacore::client::context::SendContextResolver; use wacore::libsignal::protocol::SignalProtocolError; @@ -174,17 +174,32 @@ impl Client { let _session_guard = session_mutex.lock().await; let device_store_arc = self.persistence_manager.get_device_arc().await; - let mut store_adapter = SignalProtocolStoreAdapter::new(device_store_arc); - - wacore::send::prepare_peer_stanza( - &mut store_adapter.session_store, - &mut store_adapter.identity_store, - to, - encryption_jid, - message, - request_id, - ) - .await? + let message_for_encrypt = message.clone(); + let to_for_encrypt = to.clone(); + let encryption_jid_for_encrypt = encryption_jid.clone(); + // Peer encryption is a single-session operation. With the batched + // cache, store ops resolve from memory so the work is CPU-bound. + let (mut store_adapter, stanza_result) = tokio::task::spawn_blocking(move || { + let runtime = tokio::runtime::Handle::current(); + let mut store_adapter = BatchedSignalProtocolStoreAdapter::new(device_store_arc); + let stanza_result = runtime.block_on(async { + wacore::send::prepare_peer_stanza( + &mut store_adapter.session_store, + &mut store_adapter.identity_store, + to_for_encrypt, + encryption_jid_for_encrypt, + &message_for_encrypt, + request_id, + ) + .await + }); + (store_adapter, stanza_result) + }) + .await + .map_err(|e| anyhow!("spawn_blocking failed during peer encryption: {e}"))?; + let stanza = stanza_result?; + store_adapter.flush().await?; + stanza } else if to.is_group() { // Group messages: No client-level lock needed. // Each participant device is encrypted separately with its own per-device lock @@ -239,7 +254,8 @@ impl Client { force_key_distribution || !key_exists }; - let mut store_adapter = SignalProtocolStoreAdapter::new(device_store_arc.clone()); + let mut store_adapter = + BatchedSignalProtocolStoreAdapter::new(device_store_arc.clone()); let mut stores = wacore::send::SignalStores { session_store: &mut store_adapter.session_store, @@ -332,7 +348,7 @@ impl Client { let is_full_distribution = force_skdm || skdm_target_devices.is_none(); let devices_receiving_skdm: Vec = skdm_target_devices.clone().unwrap_or_default(); - match wacore::send::prepare_group_stanza( + let stanza = match wacore::send::prepare_group_stanza( &mut stores, self, &mut group_info, @@ -393,7 +409,7 @@ impl Client { } let mut store_adapter_retry = - SignalProtocolStoreAdapter::new(device_store_arc.clone()); + BatchedSignalProtocolStoreAdapter::new(device_store_arc.clone()); let mut stores_retry = wacore::send::SignalStores { session_store: &mut store_adapter_retry.session_store, identity_store: &mut store_adapter_retry.identity_store, @@ -402,7 +418,7 @@ impl Client { sender_key_store: &mut store_adapter_retry.sender_key_store, }; - wacore::send::prepare_group_stanza( + let stanza = wacore::send::prepare_group_stanza( &mut stores_retry, self, &mut group_info, @@ -417,12 +433,21 @@ impl Client { edit.clone(), extra_stanza_nodes.clone(), ) - .await? + .await?; + store_adapter_retry.flush().await?; + stanza } else { + if let Err(flush_err) = store_adapter.flush().await { + log::warn!( + "Failed to flush cached Signal stores after group prepare error: {flush_err}" + ); + } return Err(e); } } - } + }; + store_adapter.flush().await?; + stanza } else { // Direct message: Acquire lock only during encryption @@ -464,28 +489,49 @@ impl Client { let _session_guard = session_mutex.lock().await; let device_store_arc = self.persistence_manager.get_device_arc().await; - let mut store_adapter = SignalProtocolStoreAdapter::new(device_store_arc); - - let mut stores = wacore::send::SignalStores { - session_store: &mut store_adapter.session_store, - identity_store: &mut store_adapter.identity_store, - prekey_store: &mut store_adapter.pre_key_store, - signed_prekey_store: &store_adapter.signed_pre_key_store, - sender_key_store: &mut store_adapter.sender_key_store, - }; - - wacore::send::prepare_dm_stanza( - &mut stores, - self, - &own_jid, - account_info.as_ref(), - to, - message, - request_id, - edit, - extra_stanza_nodes, - ) - .await? + let client = self + .shared() + .ok_or_else(|| anyhow!("client is shutting down"))?; + let to_for_encrypt = to.clone(); + let own_jid_for_encrypt = own_jid.clone(); + let account_info_for_encrypt = account_info.clone(); + let message_for_encrypt = message.clone(); + let edit_for_encrypt = edit.clone(); + let extra_nodes_for_encrypt = extra_stanza_nodes.clone(); + + // DM encryption involves multiple devices. With the batched cache, + // store ops resolve from memory so the work is CPU-bound. + let (mut store_adapter, stanza_result) = tokio::task::spawn_blocking(move || { + let runtime = tokio::runtime::Handle::current(); + let mut store_adapter = BatchedSignalProtocolStoreAdapter::new(device_store_arc); + let mut stores = wacore::send::SignalStores { + session_store: &mut store_adapter.session_store, + identity_store: &mut store_adapter.identity_store, + prekey_store: &mut store_adapter.pre_key_store, + signed_prekey_store: &store_adapter.signed_pre_key_store, + sender_key_store: &mut store_adapter.sender_key_store, + }; + let stanza_result = runtime.block_on(async { + wacore::send::prepare_dm_stanza( + &mut stores, + client.as_ref(), + &own_jid_for_encrypt, + account_info_for_encrypt.as_ref(), + to_for_encrypt, + &message_for_encrypt, + request_id, + edit_for_encrypt, + extra_nodes_for_encrypt, + ) + .await + }); + (store_adapter, stanza_result) + }) + .await + .map_err(|e| anyhow!("spawn_blocking failed during dm encryption: {e}"))?; + let stanza = stanza_result?; + store_adapter.flush().await?; + stanza }; self.send_node(stanza_to_send).await.map_err(|e| e.into()) diff --git a/src/store/signal.rs b/src/store/signal.rs index f04142d1c..a97b10e60 100644 --- a/src/store/signal.rs +++ b/src/store/signal.rs @@ -204,26 +204,35 @@ impl IdentityKeyStore for Device { identity_key: &IdentityKey, ) -> SignalResult { let address_str = address.to_string(); - let key_bytes = identity_key.public_key().public_key_bytes(); - let existing_identity_opt = self.get_identity(address).await?; + let key_bytes: [u8; 32] = identity_key + .public_key() + .public_key_bytes() + .try_into() + .map_err(|_| SignalProtocolError::InvalidArgument("Invalid key length".into()))?; + let existing_identity_bytes = self + .backend + .load_identity(&address_str) + .await + .map_err(|e| SignalProtocolError::InvalidState("backend get_identity", e.to_string()))? + .filter(|b| !b.is_empty()); + + if existing_identity_bytes + .as_deref() + .is_some_and(|existing| existing == key_bytes.as_slice()) + { + return Ok(IdentityChange::NewOrUnchanged); + } self.backend - .put_identity( - &address_str, - key_bytes.try_into().map_err(|_| { - SignalProtocolError::InvalidArgument("Invalid key length".into()) - })?, - ) + .put_identity(&address_str, key_bytes) .await .map_err(|e| { SignalProtocolError::InvalidState("backend put_identity", e.to_string()) })?; - match existing_identity_opt { - None => Ok(IdentityChange::NewOrUnchanged), - Some(existing) if &existing == identity_key => Ok(IdentityChange::NewOrUnchanged), - Some(_) => Ok(IdentityChange::ReplacedExisting), - } + Ok(IdentityChange::from_changed( + existing_identity_bytes.is_some(), + )) } async fn is_trusted_identity( diff --git a/src/store/signal_adapter.rs b/src/store/signal_adapter.rs index 2a1e6a086..478f6b264 100644 --- a/src/store/signal_adapter.rs +++ b/src/store/signal_adapter.rs @@ -1,12 +1,14 @@ use crate::store::Device; use async_trait::async_trait; +use std::collections::{HashMap, HashSet}; use std::sync::Arc; -use tokio::sync::RwLock; +use tokio::sync::{Mutex, RwLock}; use wacore::libsignal::protocol::{ Direction, IdentityChange, IdentityKey, IdentityKeyPair, IdentityKeyStore, PreKeyId, PreKeyRecord, PreKeyStore, ProtocolAddress, SessionRecord, SessionStore, SignalProtocolError, SignedPreKeyId, SignedPreKeyRecord, SignedPreKeyStore, }; +use wacore::libsignal::protocol::{SenderKeyRecord, SenderKeyStore}; use wacore::libsignal::store::record_helpers as wacore_record; use wacore::libsignal::store::sender_key_name::SenderKeyName; @@ -53,6 +55,207 @@ impl SignalProtocolStoreAdapter { } } +#[derive(Default)] +struct BatchWriteCache { + sessions: HashMap>, + dirty_sessions: HashSet, + identities: HashMap>, + dirty_identities: HashSet, + sender_keys: HashMap>, + dirty_sender_keys: HashSet, +} + +#[derive(Clone)] +pub struct CachedSessionAdapter { + inner: SessionAdapter, + cache: Arc>, +} + +impl CachedSessionAdapter { + fn new(inner: SessionAdapter, cache: Arc>) -> Self { + Self { inner, cache } + } + + async fn flush(&mut self) -> Result<(), SignalProtocolError> { + let pending_writes: Vec<_> = { + let cache = self.cache.lock().await; + cache + .dirty_sessions + .iter() + .filter_map(|address| { + cache + .sessions + .get(address) + .and_then(|opt| opt.as_ref()) + .map(|record| (address.clone(), record.clone())) + }) + .collect() + }; + + for (address, record) in pending_writes { + SessionStore::store_session(&mut self.inner, &address, &record).await?; + self.cache.lock().await.dirty_sessions.remove(&address); + } + + Ok(()) + } + + pub async fn delete_session( + &mut self, + address: &ProtocolAddress, + ) -> Result<(), SignalProtocolError> { + let addr_str = address.to_string(); + let device = self.inner.0.device.write().await; + device + .backend + .delete_session(&addr_str) + .await + .map_err(|e| SignalProtocolError::InvalidState("delete_session", e.to_string()))?; + drop(device); + self.invalidate(address).await; + Ok(()) + } + + async fn invalidate(&mut self, address: &ProtocolAddress) { + let mut cache = self.cache.lock().await; + cache.sessions.remove(address); + cache.dirty_sessions.remove(address); + } +} + +#[derive(Clone)] +pub struct CachedIdentityAdapter { + inner: IdentityAdapter, + cache: Arc>, +} + +impl CachedIdentityAdapter { + fn new(inner: IdentityAdapter, cache: Arc>) -> Self { + Self { inner, cache } + } + + async fn flush(&mut self) -> Result<(), SignalProtocolError> { + let pending_writes: Vec<_> = { + let cache = self.cache.lock().await; + cache + .dirty_identities + .iter() + .filter_map(|address| { + cache + .identities + .get(address) + .and_then(|opt| opt.as_ref()) + .map(|identity| (address.clone(), *identity)) + }) + .collect() + }; + + for (address, identity) in pending_writes { + IdentityKeyStore::save_identity(&mut self.inner, &address, &identity).await?; + self.cache.lock().await.dirty_identities.remove(&address); + } + + Ok(()) + } + + async fn invalidate(&mut self, address: &ProtocolAddress) { + let mut cache = self.cache.lock().await; + cache.identities.remove(address); + cache.dirty_identities.remove(address); + } +} + +#[derive(Clone)] +pub struct CachedSenderKeyAdapter { + inner: SenderKeyAdapter, + cache: Arc>, +} + +impl CachedSenderKeyAdapter { + fn new(inner: SenderKeyAdapter, cache: Arc>) -> Self { + Self { inner, cache } + } + + async fn flush(&mut self) -> Result<(), SignalProtocolError> { + let pending_writes: Vec<_> = { + let cache = self.cache.lock().await; + cache + .dirty_sender_keys + .iter() + .filter_map(|name| { + cache + .sender_keys + .get(name) + .and_then(|opt| opt.as_ref()) + .map(|record| (name.clone(), record.clone())) + }) + .collect() + }; + + for (sender_key_name, record) in pending_writes { + SenderKeyStore::store_sender_key(&mut self.inner, &sender_key_name, &record).await?; + self.cache + .lock() + .await + .dirty_sender_keys + .remove(&sender_key_name); + } + + Ok(()) + } + + async fn invalidate(&mut self, sender_key_name: &SenderKeyName) { + let mut cache = self.cache.lock().await; + cache.sender_keys.remove(sender_key_name); + cache.dirty_sender_keys.remove(sender_key_name); + } +} + +pub struct BatchedSignalProtocolStoreAdapter { + pub session_store: CachedSessionAdapter, + pub identity_store: CachedIdentityAdapter, + pub pre_key_store: PreKeyAdapter, + pub signed_pre_key_store: SignedPreKeyAdapter, + pub sender_key_store: CachedSenderKeyAdapter, +} + +impl BatchedSignalProtocolStoreAdapter { + pub fn new(device: Arc>) -> Self { + let base = SignalProtocolStoreAdapter::new(device); + let cache = Arc::new(Mutex::new(BatchWriteCache::default())); + Self { + session_store: CachedSessionAdapter::new(base.session_store, Arc::clone(&cache)), + identity_store: CachedIdentityAdapter::new(base.identity_store, Arc::clone(&cache)), + pre_key_store: base.pre_key_store, + signed_pre_key_store: base.signed_pre_key_store, + sender_key_store: CachedSenderKeyAdapter::new(base.sender_key_store, cache), + } + } + + pub async fn flush(&mut self) -> Result<(), SignalProtocolError> { + // Preserve original operation order: identity first, then session. + self.identity_store.flush().await?; + self.session_store.flush().await?; + self.sender_key_store.flush().await?; + Ok(()) + } + + pub async fn delete_session( + &mut self, + address: &ProtocolAddress, + ) -> Result<(), SignalProtocolError> { + self.session_store.delete_session(address).await + } + + pub async fn invalidate_identity(&mut self, address: &ProtocolAddress) { + self.identity_store.invalidate(address).await; + } + + pub async fn invalidate_sender_key(&mut self, sender_key_name: &SenderKeyName) { + self.sender_key_store.invalidate(sender_key_name).await; + } +} + #[async_trait] impl SessionStore for SessionAdapter { async fn load_session( @@ -81,50 +284,6 @@ impl SessionStore for SessionAdapter { let addr_str = address.to_string(); let device = self.0.device.read().await; - let existing_session = device - .backend - .get_session(&addr_str) - .await - .ok() - .flatten() - .and_then(|data| SessionRecord::deserialize(&data).ok()); - - if let (Some(existing), Some(new_state)) = (&existing_session, record.session_state()) { - if let Some(existing_state) = existing.session_state() { - let old_base_key = existing_state.alice_base_key(); - let new_base_key = new_state.alice_base_key(); - - if old_base_key != new_base_key { - let backtrace = std::backtrace::Backtrace::force_capture(); - log::warn!( - target: "signal_session_store", - "⚠️ SESSION BASE KEY CHANGED for {}!\n\ - Old base_key: {}\n\ - New base_key: {}\n\ - Old version: {:?}, New version: {:?}\n\ - Old prev_sessions: {}, New prev_sessions: {}\n\ - This will cause MAC verification failures on future messages!\n\ - Backtrace:\n{}", - addr_str, - hex::encode(old_base_key), - hex::encode(new_base_key), - existing_state.session_version(), - new_state.session_version(), - existing.previous_session_count(), - record.previous_session_count(), - backtrace - ); - } - } - } else if let (None, Some(state)) = (&existing_session, record.session_state()) { - log::debug!( - target: "signal_session_store", - "Creating new session for {}: base_key={}", - addr_str, - hex::encode(&state.alice_base_key()[..8.min(state.alice_base_key().len())]) - ); - } - let record_bytes = record.serialize()?; device .backend @@ -159,18 +318,10 @@ impl IdentityKeyStore for IdentityAdapter { address: &ProtocolAddress, identity: &IdentityKey, ) -> Result { - let existing_identity = self.get_identity(address).await?; - let mut device = self.0.device.write().await; IdentityKeyStore::save_identity(&mut *device, address, identity) .await - .map_err(|e| SignalProtocolError::InvalidState("save_identity", e.to_string()))?; - - match existing_identity { - None => Ok(IdentityChange::NewOrUnchanged), - Some(existing) if &existing == identity => Ok(IdentityChange::NewOrUnchanged), - Some(_) => Ok(IdentityChange::ReplacedExisting), - } + .map_err(|e| SignalProtocolError::InvalidState("save_identity", e.to_string())) } async fn is_trusted_identity( @@ -196,6 +347,99 @@ impl IdentityKeyStore for IdentityAdapter { } } +#[async_trait] +impl SessionStore for CachedSessionAdapter { + async fn load_session( + &self, + address: &ProtocolAddress, + ) -> Result, SignalProtocolError> { + { + let cache = self.cache.lock().await; + if let Some(cached) = cache.sessions.get(address) { + return Ok(cached.clone()); + } + } + + let loaded = SessionStore::load_session(&self.inner, address).await?; + let mut cache = self.cache.lock().await; + let cached = cache + .sessions + .entry(address.clone()) + .or_insert_with(|| loaded.clone()); + Ok(cached.clone()) + } + + async fn store_session( + &mut self, + address: &ProtocolAddress, + record: &SessionRecord, + ) -> Result<(), SignalProtocolError> { + let mut cache = self.cache.lock().await; + cache.sessions.insert(address.clone(), Some(record.clone())); + cache.dirty_sessions.insert(address.clone()); + Ok(()) + } +} + +#[async_trait] +impl IdentityKeyStore for CachedIdentityAdapter { + async fn get_identity_key_pair(&self) -> Result { + IdentityKeyStore::get_identity_key_pair(&self.inner).await + } + + async fn get_local_registration_id(&self) -> Result { + IdentityKeyStore::get_local_registration_id(&self.inner).await + } + + async fn save_identity( + &mut self, + address: &ProtocolAddress, + identity: &IdentityKey, + ) -> Result { + let existing = self.get_identity(address).await?; + let changed = existing.is_some_and(|current| current != *identity); + if existing.is_none() || changed { + let mut cache = self.cache.lock().await; + cache.identities.insert(address.clone(), Some(*identity)); + cache.dirty_identities.insert(address.clone()); + } + Ok(IdentityChange::from_changed(changed)) + } + + async fn is_trusted_identity( + &self, + address: &ProtocolAddress, + identity: &IdentityKey, + _direction: Direction, + ) -> Result { + let existing = self.get_identity(address).await?; + Ok(match existing { + None => true, + Some(stored) => stored == *identity, + }) + } + + async fn get_identity( + &self, + address: &ProtocolAddress, + ) -> Result, SignalProtocolError> { + { + let cache = self.cache.lock().await; + if let Some(cached) = cache.identities.get(address) { + return Ok(*cached); + } + } + + let loaded = IdentityKeyStore::get_identity(&self.inner, address).await?; + let mut cache = self.cache.lock().await; + let cached = cache + .identities + .entry(address.clone()) + .or_insert_with(|| loaded); + Ok(*cached) + } +} + #[async_trait] impl PreKeyStore for PreKeyAdapter { async fn get_pre_key(&self, prekey_id: PreKeyId) -> Result { @@ -274,3 +518,39 @@ impl wacore::libsignal::protocol::SenderKeyStore for SenderKeyAdapter { .await } } + +#[async_trait] +impl SenderKeyStore for CachedSenderKeyAdapter { + async fn store_sender_key( + &mut self, + sender_key_name: &SenderKeyName, + record: &SenderKeyRecord, + ) -> wacore::libsignal::protocol::error::Result<()> { + let mut cache = self.cache.lock().await; + cache + .sender_keys + .insert(sender_key_name.clone(), Some(record.clone())); + cache.dirty_sender_keys.insert(sender_key_name.clone()); + Ok(()) + } + + async fn load_sender_key( + &mut self, + sender_key_name: &SenderKeyName, + ) -> wacore::libsignal::protocol::error::Result> { + { + let cache = self.cache.lock().await; + if let Some(cached) = cache.sender_keys.get(sender_key_name) { + return Ok(cached.clone()); + } + } + + let loaded = SenderKeyStore::load_sender_key(&mut self.inner, sender_key_name).await?; + let mut cache = self.cache.lock().await; + let cached = cache + .sender_keys + .entry(sender_key_name.clone()) + .or_insert_with(|| loaded.clone()); + Ok(cached.clone()) + } +} diff --git a/transports/tokio-transport/src/lib.rs b/transports/tokio-transport/src/lib.rs index 437cc1408..ff70e41c5 100644 --- a/transports/tokio-transport/src/lib.rs +++ b/transports/tokio-transport/src/lib.rs @@ -12,7 +12,9 @@ use std::sync::{Arc, Once}; use tokio::net::TcpStream; use tokio::sync::Mutex; use tokio_websockets::{ClientBuilder, Connector, MaybeTlsStream, Message, WebSocketStream}; -use wacore::net::{Transport, TransportEvent, TransportFactory, WHATSAPP_WEB_WS_URL}; +use wacore::net::{ + Transport, TransportEvent, TransportFactory, TransportSendError, WHATSAPP_WEB_WS_URL, +}; /// Ensures the rustls crypto provider is only installed once static CRYPTO_PROVIDER_INIT: Once = Once::new(); @@ -136,14 +138,22 @@ impl Transport for TokioWebSocketTransport { /// The caller is responsible for any framing. async fn send(&self, data: Vec) -> Result<(), anyhow::Error> { let mut sink_guard = self.ws_sink.lock().await; - let sink = sink_guard - .as_mut() - .ok_or_else(|| anyhow::anyhow!("Socket is closed"))?; + let Some(sink) = sink_guard.as_mut() else { + let mut is_connected_guard = self.is_connected.lock().await; + *is_connected_guard = false; + return Err(TransportSendError::ConnectionClosed.into()); + }; debug!("--> Sending {} bytes", data.len()); - sink.send(Message::binary(data)) - .await - .map_err(|e| anyhow::anyhow!("WebSocket send error: {}", e))?; + if let Err(e) = sink.send(Message::binary(data)).await { + // Ensure future sends fail fast after a terminal websocket send failure. + *sink_guard = None; + let mut is_connected_guard = self.is_connected.lock().await; + *is_connected_guard = false; + warn!("WebSocket send error; marking transport as closed: {e}"); + return Err(TransportSendError::ConnectionClosed.into()); + } + Ok(()) } diff --git a/wacore/libsignal/benches/libsignal_benchmark.rs b/wacore/libsignal/benches/libsignal_benchmark.rs index 48fdb2e58..4efdf6b6a 100644 --- a/wacore/libsignal/benches/libsignal_benchmark.rs +++ b/wacore/libsignal/benches/libsignal_benchmark.rs @@ -1,4 +1,5 @@ use std::collections::HashMap; +use std::sync::OnceLock; use async_trait::async_trait; use iai_callgrind::{ @@ -17,18 +18,48 @@ use wacore_libsignal::protocol::{ }; use wacore_libsignal::store::sender_key_name::SenderKeyName; +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum BenchmarkStoreMode { + InMemory, + BackendSerialized, +} + +fn benchmark_store_mode() -> BenchmarkStoreMode { + static MODE: OnceLock = OnceLock::new(); + // Toggle with `WACORE_LIBSIGNAL_BENCH_MODE=backend` to include + // serialize/deserialize overhead on each store access. + *MODE.get_or_init(|| match std::env::var("WACORE_LIBSIGNAL_BENCH_MODE") { + Ok(raw) => { + let normalized = raw.trim().to_ascii_lowercase(); + if matches!( + normalized.as_str(), + "backend" | "serialized" | "backend-serialized" + ) { + BenchmarkStoreMode::BackendSerialized + } else { + BenchmarkStoreMode::InMemory + } + } + Err(_) => BenchmarkStoreMode::InMemory, + }) +} + struct InMemoryIdentityKeyStore { + mode: BenchmarkStoreMode, identity_key_pair: IdentityKeyPair, registration_id: u32, identities: HashMap, + serialized_identities: HashMap, } impl InMemoryIdentityKeyStore { fn new(identity_key_pair: IdentityKeyPair, registration_id: u32) -> Self { Self { + mode: benchmark_store_mode(), identity_key_pair, registration_id, identities: HashMap::new(), + serialized_identities: HashMap::new(), } } } @@ -51,10 +82,20 @@ impl IdentityKeyStore for InMemoryIdentityKeyStore { identity: &IdentityKey, ) -> wacore_libsignal::protocol::error::Result { let changed = self - .identities - .get(address) - .is_some_and(|existing| existing != identity); - self.identities.insert(address.clone(), *identity); + .get_identity(address) + .await? + .is_some_and(|existing| existing != *identity); + + match self.mode { + BenchmarkStoreMode::InMemory => { + self.identities.insert(address.clone(), *identity); + } + BenchmarkStoreMode::BackendSerialized => { + self.serialized_identities + .insert(address.clone(), identity.serialize()); + } + } + Ok(IdentityChange::from_changed(changed)) } @@ -71,18 +112,29 @@ impl IdentityKeyStore for InMemoryIdentityKeyStore { &self, address: &ProtocolAddress, ) -> wacore_libsignal::protocol::error::Result> { - Ok(self.identities.get(address).cloned()) + match self.mode { + BenchmarkStoreMode::InMemory => Ok(self.identities.get(address).cloned()), + BenchmarkStoreMode::BackendSerialized => self + .serialized_identities + .get(address) + .map(|bytes| IdentityKey::try_from(bytes.as_slice())) + .transpose(), + } } } struct InMemoryPreKeyStore { + mode: BenchmarkStoreMode, prekeys: HashMap, + serialized_prekeys: HashMap>, } impl InMemoryPreKeyStore { fn new() -> Self { Self { + mode: benchmark_store_mode(), prekeys: HashMap::new(), + serialized_prekeys: HashMap::new(), } } } @@ -93,10 +145,18 @@ impl PreKeyStore for InMemoryPreKeyStore { &self, prekey_id: PreKeyId, ) -> wacore_libsignal::protocol::error::Result { - self.prekeys - .get(&prekey_id) - .cloned() - .ok_or(wacore_libsignal::protocol::SignalProtocolError::InvalidPreKeyId) + match self.mode { + BenchmarkStoreMode::InMemory => self + .prekeys + .get(&prekey_id) + .cloned() + .ok_or(wacore_libsignal::protocol::SignalProtocolError::InvalidPreKeyId), + BenchmarkStoreMode::BackendSerialized => self + .serialized_prekeys + .get(&prekey_id) + .ok_or(wacore_libsignal::protocol::SignalProtocolError::InvalidPreKeyId) + .and_then(|bytes| PreKeyRecord::deserialize(bytes)), + } } async fn save_pre_key( @@ -104,7 +164,15 @@ impl PreKeyStore for InMemoryPreKeyStore { prekey_id: PreKeyId, record: &PreKeyRecord, ) -> wacore_libsignal::protocol::error::Result<()> { - self.prekeys.insert(prekey_id, record.clone()); + match self.mode { + BenchmarkStoreMode::InMemory => { + self.prekeys.insert(prekey_id, record.clone()); + } + BenchmarkStoreMode::BackendSerialized => { + self.serialized_prekeys + .insert(prekey_id, record.serialize()?); + } + } Ok(()) } @@ -112,19 +180,30 @@ impl PreKeyStore for InMemoryPreKeyStore { &mut self, prekey_id: PreKeyId, ) -> wacore_libsignal::protocol::error::Result<()> { - self.prekeys.remove(&prekey_id); + match self.mode { + BenchmarkStoreMode::InMemory => { + self.prekeys.remove(&prekey_id); + } + BenchmarkStoreMode::BackendSerialized => { + self.serialized_prekeys.remove(&prekey_id); + } + } Ok(()) } } struct InMemorySignedPreKeyStore { + mode: BenchmarkStoreMode, signed_prekeys: HashMap, + serialized_signed_prekeys: HashMap>, } impl InMemorySignedPreKeyStore { fn new() -> Self { Self { + mode: benchmark_store_mode(), signed_prekeys: HashMap::new(), + serialized_signed_prekeys: HashMap::new(), } } } @@ -135,10 +214,18 @@ impl SignedPreKeyStore for InMemorySignedPreKeyStore { &self, signed_prekey_id: SignedPreKeyId, ) -> wacore_libsignal::protocol::error::Result { - self.signed_prekeys - .get(&signed_prekey_id) - .cloned() - .ok_or(wacore_libsignal::protocol::SignalProtocolError::InvalidSignedPreKeyId) + match self.mode { + BenchmarkStoreMode::InMemory => self + .signed_prekeys + .get(&signed_prekey_id) + .cloned() + .ok_or(wacore_libsignal::protocol::SignalProtocolError::InvalidSignedPreKeyId), + BenchmarkStoreMode::BackendSerialized => self + .serialized_signed_prekeys + .get(&signed_prekey_id) + .ok_or(wacore_libsignal::protocol::SignalProtocolError::InvalidSignedPreKeyId) + .and_then(|bytes| SignedPreKeyRecord::deserialize(bytes)), + } } async fn save_signed_pre_key( @@ -146,19 +233,31 @@ impl SignedPreKeyStore for InMemorySignedPreKeyStore { signed_prekey_id: SignedPreKeyId, record: &SignedPreKeyRecord, ) -> wacore_libsignal::protocol::error::Result<()> { - self.signed_prekeys.insert(signed_prekey_id, record.clone()); + match self.mode { + BenchmarkStoreMode::InMemory => { + self.signed_prekeys.insert(signed_prekey_id, record.clone()); + } + BenchmarkStoreMode::BackendSerialized => { + self.serialized_signed_prekeys + .insert(signed_prekey_id, record.serialize()?); + } + } Ok(()) } } struct InMemorySessionStore { + mode: BenchmarkStoreMode, sessions: HashMap, + serialized_sessions: HashMap>, } impl InMemorySessionStore { fn new() -> Self { Self { + mode: benchmark_store_mode(), sessions: HashMap::new(), + serialized_sessions: HashMap::new(), } } } @@ -169,7 +268,14 @@ impl SessionStore for InMemorySessionStore { &self, address: &ProtocolAddress, ) -> wacore_libsignal::protocol::error::Result> { - Ok(self.sessions.get(address).cloned()) + match self.mode { + BenchmarkStoreMode::InMemory => Ok(self.sessions.get(address).cloned()), + BenchmarkStoreMode::BackendSerialized => self + .serialized_sessions + .get(address) + .map(|bytes| SessionRecord::deserialize(bytes)) + .transpose(), + } } async fn store_session( @@ -177,19 +283,31 @@ impl SessionStore for InMemorySessionStore { address: &ProtocolAddress, record: &SessionRecord, ) -> wacore_libsignal::protocol::error::Result<()> { - self.sessions.insert(address.clone(), record.clone()); + match self.mode { + BenchmarkStoreMode::InMemory => { + self.sessions.insert(address.clone(), record.clone()); + } + BenchmarkStoreMode::BackendSerialized => { + self.serialized_sessions + .insert(address.clone(), record.serialize()?); + } + } Ok(()) } } struct InMemorySenderKeyStore { + mode: BenchmarkStoreMode, sender_keys: HashMap, + serialized_sender_keys: HashMap>, } impl InMemorySenderKeyStore { fn new() -> Self { Self { + mode: benchmark_store_mode(), sender_keys: HashMap::new(), + serialized_sender_keys: HashMap::new(), } } } @@ -201,8 +319,16 @@ impl SenderKeyStore for InMemorySenderKeyStore { sender_key_name: &SenderKeyName, record: &SenderKeyRecord, ) -> wacore_libsignal::protocol::error::Result<()> { - self.sender_keys - .insert(sender_key_name.clone(), record.clone()); + match self.mode { + BenchmarkStoreMode::InMemory => { + self.sender_keys + .insert(sender_key_name.clone(), record.clone()); + } + BenchmarkStoreMode::BackendSerialized => { + self.serialized_sender_keys + .insert(sender_key_name.clone(), record.serialize()?); + } + } Ok(()) } @@ -210,7 +336,14 @@ impl SenderKeyStore for InMemorySenderKeyStore { &mut self, sender_key_name: &SenderKeyName, ) -> wacore_libsignal::protocol::error::Result> { - Ok(self.sender_keys.get(sender_key_name).cloned()) + match self.mode { + BenchmarkStoreMode::InMemory => Ok(self.sender_keys.get(sender_key_name).cloned()), + BenchmarkStoreMode::BackendSerialized => self + .serialized_sender_keys + .get(sender_key_name) + .map(|bytes| SenderKeyRecord::deserialize(bytes)) + .transpose(), + } } } diff --git a/wacore/src/net.rs b/wacore/src/net.rs index c9709cc60..6b4f669ff 100644 --- a/wacore/src/net.rs +++ b/wacore/src/net.rs @@ -3,6 +3,7 @@ use async_trait::async_trait; use bytes::Bytes; use std::collections::HashMap; use std::sync::Arc; +use thiserror::Error; /// Default WhatsApp Web websocket endpoint. pub const WHATSAPP_WEB_WS_URL: &str = "wss://web.whatsapp.com/ws/chat"; @@ -38,6 +39,13 @@ pub trait TransportFactory: Send + Sync { ) -> Result<(Arc, async_channel::Receiver), anyhow::Error>; } +/// Typed transport send failures that callers can downcast from `anyhow::Error`. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)] +pub enum TransportSendError { + #[error("transport connection is closed")] + ConnectionClosed, +} + /// A simple structure to represent an HTTP request #[derive(Debug, Clone)] pub struct HttpRequest { diff --git a/wacore/src/send.rs b/wacore/src/send.rs index 044725cf2..5b7f008a7 100644 --- a/wacore/src/send.rs +++ b/wacore/src/send.rs @@ -1,7 +1,7 @@ use crate::client::context::{GroupInfo, SendContextResolver}; use crate::libsignal::protocol::{ CiphertextMessage, SENDERKEY_MESSAGE_CURRENT_VERSION, SenderKeyDistributionMessage, - SenderKeyMessage, SenderKeyRecord, SenderKeyStore, SignalProtocolError, UsePQRatchet, + SenderKeyRecord, SenderKeyStore, SignalProtocolError, UsePQRatchet, group_encrypt, message_encrypt, process_prekey_bundle, }; use crate::libsignal::store::sender_key_name::SenderKeyName; @@ -12,86 +12,14 @@ use crate::reporting_token::{ use crate::types::jid::JidExt; use anyhow::{Result, anyhow}; use prost::Message as ProtoMessage; -use rand::{CryptoRng, Rng, TryRngCore as _}; +use rand::{Rng, TryRngCore as _}; use std::collections::HashSet; use wacore_binary::builder::NodeBuilder; use wacore_binary::jid::{Jid, JidExt as _}; use wacore_binary::node::{Attrs, Node}; -use wacore_libsignal::crypto::aes_256_cbc_encrypt_into; use waproto::whatsapp as wa; use waproto::whatsapp::message::DeviceSentMessage; -pub async fn encrypt_group_message( - sender_key_store: &mut S, - group_jid: &Jid, - sender_jid: &Jid, - plaintext: &[u8], - csprng: &mut R, -) -> Result -where - S: SenderKeyStore + ?Sized, - R: Rng + CryptoRng, -{ - let sender_address = sender_jid.to_protocol_address(); - let sender_key_name = SenderKeyName::new(group_jid.to_string(), sender_address.to_string()); - log::debug!( - "Attempting to load sender key for group {} sender {}", - sender_key_name.group_id(), - sender_key_name.sender_id() - ); - - let mut record = sender_key_store - .load_sender_key(&sender_key_name) - .await? - .ok_or_else(|| { - SignalProtocolError::NoSenderKeyState(format!( - "no sender key record for group {} sender {}", - sender_key_name.group_id(), - sender_key_name.sender_id() - )) - })?; - - let sender_key_state = record - .sender_key_state_mut() - .map_err(|e| anyhow!("Invalid SenderKey session: {:?}", e))?; - - let sender_chain_key = sender_key_state - .sender_chain_key() - .ok_or_else(|| anyhow!("Invalid SenderKey session: missing chain key"))?; - - let message_keys = sender_chain_key.sender_message_key(); - - let mut ciphertext = Vec::new(); - aes_256_cbc_encrypt_into( - plaintext, - message_keys.cipher_key(), - message_keys.iv(), - &mut ciphertext, - ) - .map_err(|_| anyhow!("AES encryption failed"))?; - - let signing_key = sender_key_state - .signing_key_private() - .map_err(|e| anyhow!("Invalid SenderKey session: missing signing key: {:?}", e))?; - - let skm = SenderKeyMessage::new( - SENDERKEY_MESSAGE_CURRENT_VERSION, - sender_key_state.chain_id(), - message_keys.iteration(), - ciphertext.into_boxed_slice(), - csprng, - &signing_key, - )?; - - sender_key_state.set_sender_chain_key(sender_chain_key.next()?); - - sender_key_store - .store_sender_key(&sender_key_name, &record) - .await?; - - Ok(skm) -} - pub struct SignalStores<'a, S, I, P, SP> { pub sender_key_store: &'a mut (dyn crate::libsignal::protocol::SenderKeyStore + Send + Sync), pub session_store: &'a mut S, @@ -100,6 +28,37 @@ pub struct SignalStores<'a, S, I, P, SP> { pub signed_prekey_store: &'a SP, } +fn build_participant_enc_node( + device_jid: &Jid, + encrypted_payload: CiphertextMessage, + enc_extra_attrs: &Attrs, +) -> Option<(Node, bool)> { + let (enc_type, serialized_bytes, includes_prekey_message) = match encrypted_payload { + CiphertextMessage::PreKeySignalMessage(msg) => ("pkmsg", msg.serialized().to_vec(), true), + CiphertextMessage::SignalMessage(msg) => ("msg", msg.serialized().to_vec(), false), + _ => return None, + }; + + let mut enc_attrs = Attrs::new(); + enc_attrs.insert("v".to_string(), "2".to_string()); + enc_attrs.insert("type".to_string(), enc_type.to_string()); + for (k, v) in enc_extra_attrs.iter() { + enc_attrs.insert(k.clone(), v.clone()); + } + + let enc_node = NodeBuilder::new("enc") + .attrs(enc_attrs) + .bytes(serialized_bytes) + .build(); + + let participant_node = NodeBuilder::new("to") + .attr("jid", device_jid.to_string()) + .children([enc_node]) + .build(); + + Some((participant_node, includes_prekey_message)) +} + async fn encrypt_for_devices<'a, S, I, P, SP>( stores: &mut SignalStores<'a, S, I, P, SP>, resolver: &dyn SendContextResolver, @@ -113,67 +72,19 @@ where P: crate::libsignal::protocol::PreKeyStore + Send + Sync, SP: crate::libsignal::protocol::SignedPreKeyStore + Send + Sync, { - // Build a map of device JIDs to their effective encryption JIDs. - // For phone number JIDs, check if we have an existing session under the corresponding LID. - // This handles the case where a session was established via a message with sender_lid, - // and now we're sending a reply using the phone number address. + let mut participant_nodes = Vec::new(); + let mut includes_prekey_message = false; let mut jid_to_encryption_jid: std::collections::HashMap = std::collections::HashMap::new(); let mut jids_needing_prekeys = Vec::new(); + let mut prekey_set = HashSet::new(); + // Build a deterministic PN->LID preferred mapping up-front, then attempt + // encryption directly. This avoids an extra pre-encrypt session load. for device_jid in devices { - // WhatsApp Web's SignalAddress.toString() normalizes PN → LID before - // creating signal addresses. We do the same: check LID session FIRST. - // This prevents using stale PN sessions when a newer LID session exists. - if device_jid.is_pn() - && let Some(lid_user) = resolver.get_lid_for_phone(&device_jid.user).await - { - // Construct the LID JID with the same device ID - let lid_jid = Jid::lid_device(lid_user, device_jid.device); - let lid_address = lid_jid.to_protocol_address(); - - if stores - .session_store - .load_session(&lid_address) - .await? - .is_some() - { - // Found existing session under LID address - use it! - log::debug!( - "Using LID session {} for PN {} (LID-first lookup)", - lid_jid, - device_jid - ); - jid_to_encryption_jid.insert(device_jid.clone(), lid_jid); - continue; - } - } - - // Fall back to direct address lookup (for LID JIDs or PN without LID mapping) - let signal_address = device_jid.to_protocol_address(); - if stores - .session_store - .load_session(&signal_address) - .await? - .is_some() - { - // Session exists under direct address, use it - jid_to_encryption_jid.insert(device_jid.clone(), device_jid.clone()); - continue; - } - - // No session found - need to fetch prekeys and create session. - // Keep device_jid for prekey fetch (server returns bundles keyed by this), - // but normalize to LID for the actual session creation. let encryption_jid = if device_jid.is_pn() { if let Some(lid_user) = resolver.get_lid_for_phone(&device_jid.user).await { - let lid_jid = Jid::lid_device(lid_user, device_jid.device); - log::debug!( - "Will create LID session {} for PN {} (no existing session)", - lid_jid, - device_jid - ); - lid_jid + Jid::lid_device(lid_user, device_jid.device) } else { device_jid.clone() } @@ -181,8 +92,85 @@ where device_jid.clone() }; jid_to_encryption_jid.insert(device_jid.clone(), encryption_jid); - // Use original device_jid for prekey fetch (HashMap key match) - jids_needing_prekeys.push(device_jid.clone()); + } + + // First pass: try encrypting with current mapping; collect only SessionNotFound + // devices for pre-key fetch and session creation. + for device_jid in devices { + let encryption_jid = jid_to_encryption_jid + .get(device_jid) + .unwrap_or(device_jid) + .clone(); + let signal_address = encryption_jid.to_protocol_address(); + + match message_encrypt( + plaintext_to_encrypt, + &signal_address, + stores.session_store, + stores.identity_store, + ) + .await + { + Ok(encrypted_payload) => { + if let Some((participant_node, includes_prekey)) = + build_participant_enc_node(device_jid, encrypted_payload, enc_extra_attrs) + { + includes_prekey_message |= includes_prekey; + participant_nodes.push(participant_node); + } + } + Err(SignalProtocolError::SessionNotFound(_)) if encryption_jid != *device_jid => { + // LID-first fallback: if LID session is missing, try the direct PN address + // before fetching pre-keys. + let fallback_address = device_jid.to_protocol_address(); + match message_encrypt( + plaintext_to_encrypt, + &fallback_address, + stores.session_store, + stores.identity_store, + ) + .await + { + Ok(encrypted_payload) => { + jid_to_encryption_jid.insert(device_jid.clone(), device_jid.clone()); + if let Some((participant_node, includes_prekey)) = + build_participant_enc_node( + device_jid, + encrypted_payload, + enc_extra_attrs, + ) + { + includes_prekey_message |= includes_prekey; + participant_nodes.push(participant_node); + } + } + Err(SignalProtocolError::SessionNotFound(_)) => { + if prekey_set.insert(device_jid.clone()) { + jids_needing_prekeys.push(device_jid.clone()); + } + } + Err(e) => { + log::warn!( + "Failed to encrypt message for device {} (direct fallback): {}. Skipping this device.", + &fallback_address, + e + ); + } + } + } + Err(SignalProtocolError::SessionNotFound(_)) => { + if prekey_set.insert(device_jid.clone()) { + jids_needing_prekeys.push(device_jid.clone()); + } + } + Err(e) => { + log::warn!( + "Failed to encrypt message for device {}: {}. Skipping this device.", + &signal_address, + e + ); + } + } } if !jids_needing_prekeys.is_empty() { @@ -195,21 +183,18 @@ where .await?; for device_jid in &jids_needing_prekeys { - // Use the LID-normalized encryption JID for session creation + // Keep LID-preferred session establishment when available. let mut encryption_jid = jid_to_encryption_jid .get(device_jid) .unwrap_or(device_jid) .clone(); - // Normalize agent to 0 for LID JIDs to match how pre-key bundles are stored. - // The JID parsing logic in `prekeys.rs` forces agent=0 for LID, so we must match that here. + // LID bundles are keyed with agent=0. if encryption_jid.is_lid() { encryption_jid.agent = 0; } let signal_address = encryption_jid.to_protocol_address(); - // Fix: Use the normalized device_jid to lookup the bundle - // Use centralized normalization logic to avoid mismatches let lookup_jid = device_jid.normalize_for_prekey_bundle(); match prekey_bundles.get(&lookup_jid) { Some(bundle) => { @@ -223,20 +208,13 @@ where ) .await { - Ok(_) => { - // Session established successfully - } + Ok(_) => {} Err(SignalProtocolError::UntrustedIdentity(ref addr)) => { - // The stored identity doesn't match the server's identity. - // This typically happens when a user reinstalls WhatsApp. - // We trust the server's identity and update our local store, - // then retry establishing the session. log::info!( "Untrusted identity for device {}. Updating identity and retrying session establishment.", addr ); - // Get the new identity from the prekey bundle and save it let new_identity = match bundle.identity_key() { Ok(key) => key, Err(e) => { @@ -249,7 +227,6 @@ where } }; - // Save the new identity (this replaces the old one) if let Err(e) = stores .identity_store .save_identity(&signal_address, new_identity) @@ -263,12 +240,6 @@ where continue; } - log::debug!( - "Identity updated for {}. Retrying session establishment.", - addr - ); - - // Retry processing the prekey bundle with the updated identity match process_prekey_bundle( &signal_address, stores.session_store, @@ -296,7 +267,6 @@ where } } Err(e) => { - // Propagate other unexpected errors return Err(anyhow::anyhow!( "Failed to process pre-key bundle for {}: {:?}", signal_address, @@ -313,62 +283,68 @@ where } } } - } - let mut participant_nodes = Vec::new(); - let mut includes_prekey_message = false; - - for device_jid in devices { - // Use the effective encryption JID (may be LID if we found an existing LID session) - let encryption_jid = jid_to_encryption_jid.get(device_jid).unwrap_or(device_jid); - let signal_address = encryption_jid.to_protocol_address(); - - // Try to encrypt for this device. If it fails (e.g., no session established), - // log a warning and skip this device instead of failing the entire operation. - match message_encrypt( - plaintext_to_encrypt, - &signal_address, - stores.session_store, - stores.identity_store, - ) - .await - { - Ok(encrypted_payload) => { - let (enc_type, serialized_bytes) = match encrypted_payload { - CiphertextMessage::PreKeySignalMessage(msg) => { - includes_prekey_message = true; - ("pkmsg", msg.serialized().to_vec()) + // Second pass: retry encryption only for devices that needed pre-keys. + for device_jid in &jids_needing_prekeys { + let encryption_jid = jid_to_encryption_jid + .get(device_jid) + .unwrap_or(device_jid) + .clone(); + let signal_address = encryption_jid.to_protocol_address(); + match message_encrypt( + plaintext_to_encrypt, + &signal_address, + stores.session_store, + stores.identity_store, + ) + .await + { + Ok(encrypted_payload) => { + if let Some((participant_node, includes_prekey)) = + build_participant_enc_node(device_jid, encrypted_payload, enc_extra_attrs) + { + includes_prekey_message |= includes_prekey; + participant_nodes.push(participant_node); } - CiphertextMessage::SignalMessage(msg) => ("msg", msg.serialized().to_vec()), - _ => continue, - }; - - let mut enc_attrs = Attrs::new(); - enc_attrs.insert("v".to_string(), "2".to_string()); - enc_attrs.insert("type".to_string(), enc_type.to_string()); - for (k, v) in enc_extra_attrs.iter() { - enc_attrs.insert(k.clone(), v.clone()); } - - let enc_node = NodeBuilder::new("enc") - .attrs(enc_attrs) - .bytes(serialized_bytes) - .build(); - // Use the original device_jid for the `to` attribute (what the server expects), - // but we encrypted using the encryption_jid's session - participant_nodes.push( - NodeBuilder::new("to") - .attr("jid", device_jid.to_string()) - .children([enc_node]) - .build(), - ); - } - Err(e) => { - log::warn!( - "Failed to encrypt message for device {}: {}. Skipping this device.", - &signal_address, - e - ); + Err(SignalProtocolError::SessionNotFound(_)) if encryption_jid != *device_jid => { + let fallback_address = device_jid.to_protocol_address(); + match message_encrypt( + plaintext_to_encrypt, + &fallback_address, + stores.session_store, + stores.identity_store, + ) + .await + { + Ok(encrypted_payload) => { + if let Some((participant_node, includes_prekey)) = + build_participant_enc_node( + device_jid, + encrypted_payload, + enc_extra_attrs, + ) + { + includes_prekey_message |= includes_prekey; + participant_nodes.push(participant_node); + } + } + Err(e) => { + log::warn!( + "Failed to encrypt message for device {} after pre-key setup: {}. Skipping this device.", + &fallback_address, + e + ); + } + } + } + Err(e) => { + log::warn!( + "Failed to encrypt message for device {} after pre-key setup: {}. Skipping this device.", + &signal_address, + e + ); + } } } } @@ -774,10 +750,11 @@ pub async fn prepare_group_stanza< } let plaintext = MessageUtils::pad_message_v2(message_for_encryption.encode_to_vec()); - let skmsg = encrypt_group_message( + let sender_address = own_sending_jid.to_protocol_address(); + let sender_key_name = SenderKeyName::new(to_jid.to_string(), sender_address.to_string()); + let skmsg = group_encrypt( stores.sender_key_store, - &to_jid, - &own_sending_jid, + &sender_key_name, &plaintext, &mut rand::rngs::OsRng.unwrap_err(), )