diff --git a/http_clients/ureq-client/src/lib.rs b/http_clients/ureq-client/src/lib.rs index 4df9dfe48..43fd4dbf7 100644 --- a/http_clients/ureq-client/src/lib.rs +++ b/http_clients/ureq-client/src/lib.rs @@ -27,20 +27,24 @@ impl Default for UreqHttpClient { } fn build_agent() -> ureq::Agent { + use ureq::config::Config; + + #[allow(unused_mut)] + let mut builder = Config::builder() + // 16 KB per buffer instead of the 128 KB default. + // WA API payloads are small JSON; media uses streaming I/O. + .input_buffer_size(16 * 1024) + .output_buffer_size(16 * 1024) + .max_idle_connections(3) + .max_idle_connections_per_host(2); + #[cfg(feature = "danger-skip-tls-verify")] { - use ureq::config::Config; use ureq::tls::TlsConfig; - Config::builder() - .tls_config(TlsConfig::builder().disable_verification(true).build()) - .build() - .into() + builder = builder.tls_config(TlsConfig::builder().disable_verification(true).build()); } - #[cfg(not(feature = "danger-skip-tls-verify"))] - { - ureq::Agent::new_with_defaults() - } + builder.build().into() } #[async_trait] diff --git a/src/prekeys.rs b/src/prekeys.rs index 71633e8e9..357cb793d 100644 --- a/src/prekeys.rs +++ b/src/prekeys.rs @@ -311,42 +311,51 @@ impl Client { guard.backend.clone() }; - // Load each prekey referenced by the server digest and extract its public key + // Batch-load all prekeys referenced by the server digest + let loaded = match backend.load_prekeys_batch(&response.prekey_ids).await { + Ok(v) => v, + Err(e) => { + log::warn!("digestKey: failed to batch-load prekeys: {:?}, skipping", e); + return Ok(()); + } + }; + + // Build a lookup so we preserve the server-requested order. + // Dedupe the expected count since the server may send duplicate IDs. + let loaded_map: std::collections::HashMap> = loaded.into_iter().collect(); + let unique_requested: std::collections::HashSet<&u32> = + response.prekey_ids.iter().collect(); + + if loaded_map.len() < unique_requested.len() { + log::warn!( + "digestKey: missing {} local prekeys, skipping", + unique_requested.len() - loaded_map.len() + ); + return Ok(()); + } + let mut prekey_pubkeys = Vec::with_capacity(response.prekey_ids.len()); for prekey_id in &response.prekey_ids { - match backend.load_prekey(*prekey_id).await { - Ok(Some(record_bytes)) => { - use prost::Message; - match waproto::whatsapp::PreKeyRecordStructure::decode(record_bytes.as_slice()) - { - Ok(record) => { - if let Some(pk) = record.public_key { - prekey_pubkeys.push(pk); - } else { - log::warn!( - "digestKey: prekey {} has no public key, skipping", - prekey_id - ); - return Ok(()); - } - } - Err(e) => { - log::warn!( - "digestKey: failed to decode prekey {}: {}, skipping", - prekey_id, - e - ); - return Ok(()); - } + let Some(record_bytes) = loaded_map.get(prekey_id) else { + log::warn!("digestKey: missing local prekey {}, skipping", prekey_id); + return Ok(()); + }; + use prost::Message; + match waproto::whatsapp::PreKeyRecordStructure::decode(record_bytes.as_slice()) { + Ok(record) => { + if let Some(pk) = record.public_key { + prekey_pubkeys.push(pk); + } else { + log::warn!( + "digestKey: prekey {} has no public key, skipping", + prekey_id + ); + return Ok(()); } } - Ok(None) => { - log::warn!("digestKey: missing local prekey {}, skipping", prekey_id); - return Ok(()); - } Err(e) => { log::warn!( - "digestKey: failed to load prekey {}: {:?}, skipping", + "digestKey: failed to decode prekey {}: {}, skipping", prekey_id, e ); diff --git a/src/request.rs b/src/request.rs index 15c337533..9260be0df 100644 --- a/src/request.rs +++ b/src/request.rs @@ -135,11 +135,12 @@ impl Client { return Err(IqError::NotConnected); } + let default_timeout = Duration::from_secs(75); + let iq_timeout = query.timeout.unwrap_or(default_timeout); let req_id = query .id .clone() .unwrap_or_else(|| self.generate_request_id()); - let default_timeout = Duration::from_secs(75); let (tx, rx) = futures::channel::oneshot::channel(); self.response_waiters @@ -148,7 +149,7 @@ impl Client { .insert(req_id.clone(), tx); let request_utils = self.get_request_utils(); - let node = request_utils.build_iq_node(&query, Some(req_id.clone())); + let node = request_utils.build_iq_node(query, Some(req_id.clone())); // Register the shutdown listener BEFORE sending to avoid a window where // a shutdown fires between send_node() completing and listen() being called. @@ -172,7 +173,6 @@ impl Client { // Race the IQ response against shutdown so we fail fast on disconnect // instead of waiting the full timeout. - let iq_timeout = query.timeout.unwrap_or(default_timeout); futures::select! { result = rt_timeout(&*self.runtime, iq_timeout, rx).fuse() => { diff --git a/storages/sqlite-storage/src/sqlite_store.rs b/storages/sqlite-storage/src/sqlite_store.rs index 9f8b1cf1a..d79bfae9e 100644 --- a/storages/sqlite-storage/src/sqlite_store.rs +++ b/storages/sqlite-storage/src/sqlite_store.rs @@ -1188,6 +1188,26 @@ impl SignalStore for SqliteStore { self.get_session_for_device(address, self.device_id).await } + async fn has_session(&self, address: &str) -> Result { + let pool = self.pool.clone(); + let device_id = self.device_id; + let address_owned = address.to_string(); + self.with_semaphore(move || -> Result { + let mut conn = pool + .get() + .map_err(|e| StoreError::Connection(e.to_string()))?; + let exists = diesel::select(diesel::dsl::exists( + sessions::table + .filter(sessions::address.eq(&address_owned)) + .filter(sessions::device_id.eq(device_id)), + )) + .get_result(&mut conn) + .map_err(|e| StoreError::Database(e.to_string()))?; + Ok(exists) + }) + .await + } + async fn put_session(&self, address: &str, session: &[u8]) -> Result<()> { self.put_session_for_device(address, session, self.device_id) .await @@ -1346,6 +1366,28 @@ impl SignalStore for SqliteStore { .map_err(|e| StoreError::Database(e.to_string()))? } + async fn load_prekeys_batch(&self, ids: &[u32]) -> Result)>> { + if ids.is_empty() { + return Ok(Vec::new()); + } + let pool = self.pool.clone(); + let device_id = self.device_id; + let ids: Vec = ids.iter().map(|&id| id as i32).collect(); + self.with_semaphore(move || -> Result)>> { + let mut conn = pool + .get() + .map_err(|e| StoreError::Connection(e.to_string()))?; + let rows: Vec<(i32, Vec)> = prekeys::table + .select((prekeys::id, prekeys::key)) + .filter(prekeys::id.eq_any(&ids)) + .filter(prekeys::device_id.eq(device_id)) + .load(&mut conn) + .map_err(|e| StoreError::Database(e.to_string()))?; + Ok(rows.into_iter().map(|(id, key)| (id as u32, key)).collect()) + }) + .await + } + async fn remove_prekey(&self, id: u32) -> Result<()> { let pool = self.pool.clone(); let db_semaphore = self.db_semaphore.clone(); diff --git a/wacore/src/iq/prekeys.rs b/wacore/src/iq/prekeys.rs index f31bda2d5..801d2b070 100644 --- a/wacore/src/iq/prekeys.rs +++ b/wacore/src/iq/prekeys.rs @@ -365,20 +365,15 @@ impl IqSpec for PreKeyUploadSpec { type Response = (); fn build_iq(&self) -> InfoQuery<'static> { - // Convert PublicKeys to 32-byte raw values for the wire - let pre_keys_bytes: Vec<(u32, Vec)> = self - .pre_keys - .iter() - .map(|(id, pk)| (*id, pk.public_key_bytes().to_vec())) - .collect(); - let content = PreKeyUtils::build_upload_prekeys_request( self.registration_id, self.identity_key.public_key_bytes().to_vec(), self.signed_pre_key_id, self.signed_pre_key_public.public_key_bytes().to_vec(), self.signed_pre_key_signature.clone(), - &pre_keys_bytes, + self.pre_keys + .iter() + .map(|(id, pk)| (*id, pk.public_key_bytes().to_vec())), ); InfoQuery::set( diff --git a/wacore/src/iq/props.rs b/wacore/src/iq/props.rs index 808d78d95..28181449a 100644 --- a/wacore/src/iq/props.rs +++ b/wacore/src/iq/props.rs @@ -106,7 +106,7 @@ impl crate::protocol::ProtocolNode for AbProp { } let config_value = optional_attr(node, "config_value") .ok_or_else(|| anyhow::anyhow!("missing config_value in prop"))? - .to_string(); + .into_owned(); let config_expo_key = optional_attr(node, "config_expo_key").and_then(|s| s.parse().ok()); Ok(Self { @@ -189,31 +189,31 @@ impl crate::protocol::ProtocolNode for AbPropConfig { } fn try_from_node_ref(node: &NodeRef<'_>) -> Result { + use crate::iq::node::optional_attr; + if node.tag != "prop" { return Err(anyhow::anyhow!("expected , got <{}>", node.tag)); } - let experiment = AbProp::try_from_node_ref(node); - if let Ok(prop) = experiment { - return Ok(Self::Experiment(prop)); - } - - let sampling = SamplingProp::try_from_node_ref(node); - if let Ok(prop) = sampling { - return Ok(Self::Sampling(prop)); + // Check discriminating attribute to avoid double-parse allocations + let has_config = optional_attr(node, "config_code").is_some(); + let has_event = optional_attr(node, "event_code").is_some(); + + if has_config && has_event { + Err(anyhow::anyhow!( + "prop has both config_code and event_code (attrs: {:?})", + node.attrs + )) + } else if has_config { + Ok(Self::Experiment(AbProp::try_from_node_ref(node)?)) + } else if has_event { + Ok(Self::Sampling(SamplingProp::try_from_node_ref(node)?)) + } else { + Err(anyhow::anyhow!( + "prop has neither config_code nor event_code (attrs: {:?})", + node.attrs + )) } - - let experiment_err = experiment - .err() - .unwrap_or_else(|| anyhow::anyhow!("unknown error")); - let sampling_err = sampling - .err() - .unwrap_or_else(|| anyhow::anyhow!("unknown error")); - Err(anyhow::anyhow!( - "prop did not match experiment or sampling config: experiment_err={}; sampling_err={}", - experiment_err, - sampling_err - )) } } diff --git a/wacore/src/prekeys.rs b/wacore/src/prekeys.rs index 6d21e2921..a909ed00a 100644 --- a/wacore/src/prekeys.rs +++ b/wacore/src/prekeys.rs @@ -47,17 +47,17 @@ impl PreKeyUtils { signed_pre_key_id: u32, signed_pre_key_public_bytes: Vec, signed_pre_key_signature: Vec, - pre_keys: &[(u32, Vec)], + pre_keys: impl IntoIterator)>, ) -> Vec { - let mut pre_key_nodes = Vec::new(); + let pre_keys = pre_keys.into_iter(); + let (lower, upper) = pre_keys.size_hint(); + let mut pre_key_nodes = Vec::with_capacity(upper.unwrap_or(lower)); for (pre_key_id, public_bytes) in pre_keys { let id_bytes = pre_key_id.to_be_bytes()[1..].to_vec(); let node = NodeBuilder::new("key") .children([ NodeBuilder::new("id").bytes(id_bytes).build(), - NodeBuilder::new("value") - .bytes(public_bytes.clone()) - .build(), + NodeBuilder::new("value").bytes(public_bytes).build(), ]) .build(); pre_key_nodes.push(node); diff --git a/wacore/src/request.rs b/wacore/src/request.rs index d1a383f1c..75558cc6a 100644 --- a/wacore/src/request.rs +++ b/wacore/src/request.rs @@ -178,30 +178,22 @@ impl RequestUtils { id } - pub fn build_iq_node(&self, query: &InfoQuery<'_>, req_id: Option) -> Node { + pub fn build_iq_node(&self, query: InfoQuery<'_>, req_id: Option) -> Node { let id = req_id.unwrap_or_else(|| self.generate_request_id()); let mut builder = NodeBuilder::new("iq") .attr("id", id) .attr("xmlns", query.namespace) .attr("type", query.query_type.as_str()) - .attr("to", &query.to); + .attr("to", query.to); - if let Some(target) = &query.target + if let Some(target) = query.target && !target.is_empty() { builder = builder.attr("target", target); } - if let Some(content) = &query.content { - match content { - NodeContent::Bytes(b) => builder = builder.bytes(b.clone()), - NodeContent::String(s) => builder = builder.string_content(s.clone()), - NodeContent::Nodes(n) => builder = builder.children(n.clone()), - } - } - - builder.build() + builder.apply_content(query.content).build() } pub fn parse_iq_response(&self, response_node: &NodeRef<'_>) -> Result<(), IqError> { diff --git a/wacore/src/store/in_memory.rs b/wacore/src/store/in_memory.rs index 059659897..a5f4a2b68 100644 --- a/wacore/src/store/in_memory.rs +++ b/wacore/src/store/in_memory.rs @@ -155,6 +155,10 @@ impl SignalStore for InMemoryBackend { Ok(()) } + async fn has_session(&self, address: &str) -> Result { + Ok(self.state.lock().await.sessions.contains_key(address)) + } + async fn delete_session(&self, address: &str) -> Result<()> { self.state.lock().await.sessions.remove(address); Ok(()) @@ -170,6 +174,19 @@ impl SignalStore for InMemoryBackend { Ok(()) } + async fn store_prekeys_batch(&self, keys: &[(u32, Vec)], _uploaded: bool) -> Result<()> { + let mut state = self.state.lock().await; + for (id, record) in keys { + state.prekeys.insert( + *id, + PreKeyEntry { + record: record.clone(), + }, + ); + } + Ok(()) + } + async fn load_prekey(&self, id: u32) -> Result>> { Ok(self .state @@ -180,6 +197,17 @@ impl SignalStore for InMemoryBackend { .map(|e| e.record.clone())) } + async fn load_prekeys_batch(&self, ids: &[u32]) -> Result)>> { + let state = self.state.lock().await; + let mut result = Vec::with_capacity(ids.len()); + for &id in ids { + if let Some(entry) = state.prekeys.get(&id) { + result.push((id, entry.record.clone())); + } + } + Ok(result) + } + async fn remove_prekey(&self, id: u32) -> Result<()> { self.state.lock().await.prekeys.remove(&id); Ok(()) diff --git a/wacore/src/store/traits.rs b/wacore/src/store/traits.rs index 5f82cbf8f..875e1b6fe 100644 --- a/wacore/src/store/traits.rs +++ b/wacore/src/store/traits.rs @@ -126,6 +126,18 @@ pub trait SignalStore: Send + Sync { /// Load a pre-key by ID. async fn load_prekey(&self, id: u32) -> Result>>; + /// Load multiple pre-keys by ID in a single batch operation. + /// Returns only the keys that exist. + async fn load_prekeys_batch(&self, ids: &[u32]) -> Result)>> { + let mut result = Vec::with_capacity(ids.len()); + for &id in ids { + if let Some(record) = self.load_prekey(id).await? { + result.push((id, record)); + } + } + Ok(result) + } + /// Remove a pre-key. async fn remove_prekey(&self, id: u32) -> Result<()>;