use super::*;
use anyhow::Context;
#[cfg_attr(
feature = "tracing",
tracing::instrument(name = "wa.send.encrypt_group", level = "debug", skip_all, err(Debug))
)]
pub async fn encrypt_group_message<S, R>(
sender_key_store: &mut S,
sender_key_name: &SenderKeyName,
plaintext: &[u8],
csprng: &mut R,
) -> Result<SenderKeyMessage>
where
S: SenderKeyStore + ?Sized,
R: Rng + CryptoRng,
{
crate::libsignal::protocol::group_encrypt(sender_key_store, sender_key_name, plaintext, csprng)
.await
.context("group encrypt failed")
}
pub trait CloneableSessionStore: crate::libsignal::protocol::SessionStore {
fn clone_box(&self) -> Box<dyn CloneableSessionStore + Send + Sync>;
}
impl<T> CloneableSessionStore for T
where
T: crate::libsignal::protocol::SessionStore + Clone + Send + Sync + 'static,
{
fn clone_box(&self) -> Box<dyn CloneableSessionStore + Send + Sync> {
Box::new(self.clone())
}
}
pub trait CloneableIdentityStore: crate::libsignal::protocol::IdentityKeyStore {
fn clone_box(&self) -> Box<dyn CloneableIdentityStore + Send + Sync>;
}
impl<T> CloneableIdentityStore for T
where
T: crate::libsignal::protocol::IdentityKeyStore + Clone + Send + Sync + 'static,
{
fn clone_box(&self) -> Box<dyn CloneableIdentityStore + Send + Sync> {
Box::new(self.clone())
}
}
pub struct SignalStores<'a> {
pub sender_key_store: &'a mut (dyn SenderKeyStore + Send + Sync),
pub session_store: &'a mut (dyn CloneableSessionStore + Send + Sync),
pub identity_store: &'a mut (dyn CloneableIdentityStore + Send + Sync),
pub prekey_store: &'a mut (dyn crate::libsignal::protocol::PreKeyStore + Send + Sync),
pub signed_prekey_store: &'a (dyn crate::libsignal::protocol::SignedPreKeyStore + Send + Sync),
}
pub(crate) const UNREGISTERED_DEVICE_CODE: u16 = 406;
pub(crate) fn is_device_unregistered_error(err: &anyhow::Error) -> bool {
crate::request::ServerErrorCode::from_anyhow(err)
.is_some_and(|e| e.code == UNREGISTERED_DEVICE_CODE)
}
pub struct EncryptResult {
pub participant_nodes: Vec<Node>,
pub includes_prekey_message: bool,
pub encrypted_devices: Vec<Jid>,
pub had_unregistered_device: bool,
pub rejected_devices: Vec<Jid>,
}
pub(crate) struct EncryptAttempt {
pub result: EncryptResult,
pub first_error: Option<anyhow::Error>,
}
pub struct EncryptedDevice {
pub device_jid: Jid,
pub enc_type: &'static str,
pub is_prekey: bool,
pub ciphertext: Vec<u8>,
}
pub struct EncryptForDevicesRaw {
pub devices: Vec<EncryptedDevice>,
pub includes_prekey_message: bool,
pub had_unregistered_device: bool,
pub rejected_devices: Vec<Jid>,
}
struct RawEncryptAttempt {
result: EncryptForDevicesRaw,
first_error: Option<anyhow::Error>,
}
pub fn needs_device_identity(
includes_prekey: bool,
account: Option<&wa::ADVSignedDeviceIdentity>,
) -> Result<Option<Vec<u8>>> {
if !includes_prekey {
return Ok(None);
}
let acc = account
.ok_or_else(|| anyhow!("pkmsg requires <device-identity> but no ADV account is present"))?;
Ok(Some(waproto::codec::adv_signed_device_identity_to_vec(acc)))
}
const ENCRYPT_FANOUT_CONCURRENCY: usize = 16;
struct EncryptOneResult {
enc_type: &'static str,
is_prekey: bool,
ciphertext: Vec<u8>,
}
#[derive(Debug, thiserror::Error)]
#[error("spawned task did not produce a result (panic or runtime shutdown)")]
struct SpawnCanceled;
struct Spawned<T> {
rx: futures::channel::oneshot::Receiver<T>,
abort: Option<AbortHandle>,
}
impl<T> Future for Spawned<T> {
type Output = std::result::Result<T, SpawnCanceled>;
fn poll(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Self::Output> {
match std::pin::Pin::new(&mut self.rx).poll(cx) {
std::task::Poll::Ready(Ok(value)) => {
if let Some(handle) = self.abort.take() {
handle.detach();
}
std::task::Poll::Ready(Ok(value))
}
std::task::Poll::Ready(Err(_)) => {
if let Some(handle) = self.abort.take() {
handle.detach();
}
std::task::Poll::Ready(Err(SpawnCanceled))
}
std::task::Poll::Pending => std::task::Poll::Pending,
}
}
}
impl<T> Drop for Spawned<T> {
fn drop(&mut self) {
if let Some(handle) = self.abort.take() {
handle.abort();
}
}
}
#[cfg(not(target_arch = "wasm32"))]
fn spawn_oneshot<F, T>(
rt: &dyn Runtime,
fut: F,
) -> impl Future<Output = std::result::Result<T, SpawnCanceled>> + Send + 'static
where
F: Future<Output = T> + Send + 'static,
T: Send + 'static,
{
let (tx, rx) = futures::channel::oneshot::channel();
let abort = rt.spawn(Box::pin(async move {
let _ = tx.send(fut.await);
}));
Spawned {
rx,
abort: Some(abort),
}
}
#[cfg(target_arch = "wasm32")]
fn spawn_oneshot<F, T>(
rt: &dyn Runtime,
fut: F,
) -> impl Future<Output = std::result::Result<T, SpawnCanceled>> + 'static
where
F: Future<Output = T> + 'static,
T: 'static,
{
let (tx, rx) = futures::channel::oneshot::channel();
let abort = rt.spawn(Box::pin(async move {
let _ = tx.send(fut.await);
}));
Spawned {
rx,
abort: Some(abort),
}
}
async fn encrypt_one_device(
plaintext: &[u8],
addr: &ProtocolAddress,
session_store: &mut dyn crate::libsignal::protocol::SessionStore,
identity_store: &mut dyn crate::libsignal::protocol::IdentityKeyStore,
device_jid: Jid,
) -> (Jid, Result<Option<EncryptOneResult>>) {
match message_encrypt(plaintext, addr, session_store, identity_store).await {
Ok(encrypted_payload) => {
let Some((enc_type, is_prekey, serialized_bytes)) =
extract_ciphertext(encrypted_payload)
else {
return (device_jid, Ok(None));
};
(
device_jid,
Ok(Some(EncryptOneResult {
enc_type,
is_prekey,
ciphertext: serialized_bytes.into(),
})),
)
}
Err(error) => (
device_jid,
Err(anyhow::Error::new(error).context(format!("failed to encrypt for {addr}"))),
),
}
}
fn push_raw_result(
(device_jid, res): (Jid, Result<Option<EncryptOneResult>>),
devices: &mut Vec<EncryptedDevice>,
includes_prekey_message: &mut bool,
first_error: &mut Option<anyhow::Error>,
) {
match res {
Ok(Some(one)) => {
*includes_prekey_message |= one.is_prekey;
devices.push(EncryptedDevice {
device_jid,
enc_type: one.enc_type,
is_prekey: one.is_prekey,
ciphertext: one.ciphertext,
});
}
Ok(None) => {}
Err(error) => {
log::warn!("Failed to encrypt for device: {error:#}. Skipping.");
if first_error.is_none() {
*first_error = Some(error);
}
}
}
}
fn encrypted_device_to_participant_node(
one: EncryptedDevice,
mediatype: Option<&str>,
hide_decrypt_fail: bool,
) -> Node {
let mut enc_builder = NodeBuilder::new("enc")
.attr("v", stanza::ENC_VERSION)
.attr("type", one.enc_type);
if let Some(mt) = mediatype {
enc_builder = enc_builder.attr("mediatype", mt);
}
if hide_decrypt_fail {
enc_builder = enc_builder.attr("decrypt-fail", "hide");
}
let enc_node = enc_builder.bytes(one.ciphertext).build();
NodeBuilder::new("to")
.attr("jid", one.device_jid)
.children([enc_node])
.build()
}
pub async fn encrypt_for_devices(
runtime: &dyn Runtime,
stores: &mut SignalStores<'_>,
resolver: &dyn SendContextResolver,
devices: &[Jid],
plaintext_to_encrypt: &[u8],
hide_decrypt_fail: bool,
mediatype: Option<&str>,
) -> Result<EncryptResult> {
let plan = ensure_sessions_for_devices(runtime, stores, resolver, devices).await?;
encrypt_for_devices_with_sessions(
runtime,
stores,
devices,
plaintext_to_encrypt,
hide_decrypt_fail,
mediatype,
plan,
)
.await
}
pub struct EncryptFanoutSummary {
pub includes_prekey_message: bool,
pub had_unregistered_device: bool,
}
#[allow(clippy::too_many_arguments)]
pub async fn encrypt_for_devices_into(
runtime: &dyn Runtime,
stores: &mut SignalStores<'_>,
resolver: &dyn SendContextResolver,
devices: &[Jid],
plaintext_to_encrypt: &[u8],
hide_decrypt_fail: bool,
mediatype: Option<&str>,
participant_nodes: &mut Vec<Node>,
) -> Result<EncryptFanoutSummary> {
let plan = ensure_sessions_for_devices(runtime, stores, resolver, devices).await?;
let RawEncryptAttempt { result: raw, .. } = encrypt_for_devices_with_sessions_raw_detailed(
runtime,
stores,
devices,
plaintext_to_encrypt,
plan,
)
.await?;
participant_nodes.reserve(raw.devices.len());
for one in raw.devices {
participant_nodes.push(encrypted_device_to_participant_node(
one,
mediatype,
hide_decrypt_fail,
));
}
Ok(EncryptFanoutSummary {
includes_prekey_message: raw.includes_prekey_message,
had_unregistered_device: raw.had_unregistered_device,
})
}
pub struct SessionPlan {
device_count: usize,
encryption_overrides: Vec<Option<Jid>>,
pub had_unregistered_device: bool,
pub rejected_devices: Vec<Jid>,
first_error: Option<anyhow::Error>,
}
impl SessionPlan {
pub fn assume_ready(device_count: usize) -> Self {
Self {
device_count,
encryption_overrides: Vec::new(),
had_unregistered_device: false,
rejected_devices: Vec::new(),
first_error: None,
}
}
}
fn encryption_override_at(overrides: &[Option<Jid>], index: usize) -> Option<&Jid> {
overrides.get(index).and_then(Option::as_ref)
}
fn record_encryption_override(
overrides: &mut Vec<Option<Jid>>,
device_count: usize,
index: usize,
jid: Jid,
) {
if overrides.is_empty() {
overrides.resize(device_count, None);
}
overrides[index] = Some(jid);
}
#[cfg_attr(feature = "tracing", tracing::instrument(name = "wa.send.ensure_sessions", level = "debug", skip_all, fields(count = devices.len()), err(Debug)))]
pub async fn ensure_sessions_for_devices(
runtime: &dyn Runtime,
stores: &mut SignalStores<'_>,
resolver: &dyn SendContextResolver,
devices: &[Jid],
) -> Result<SessionPlan> {
let mut encryption_overrides: Vec<Option<Jid>> = Vec::new();
let mut indices_needing_prekeys: Vec<usize> = Vec::new();
let mut had_406 = false;
let mut rejected_devices: Vec<Jid> = Vec::new();
let mut first_error = None;
let mut reusable_addr = crate::types::jid::make_reusable_protocol_address();
for (idx, device_jid) in devices.iter().enumerate() {
if device_jid.is_pn()
&& let Some(lid_user) = resolver.get_lid_for_phone(&device_jid.user).await
{
let lid_jid = Jid::lid_device(lid_user, device_jid.device);
lid_jid.reset_protocol_address(&mut reusable_addr);
if wacore_libsignal::protocol::has_session(stores.session_store, &reusable_addr).await?
{
log::debug!(
"Using LID session {} for PN {} (LID-first lookup)",
lid_jid.observe(),
device_jid.observe()
);
record_encryption_override(&mut encryption_overrides, devices.len(), idx, lid_jid);
continue;
}
}
device_jid.reset_protocol_address(&mut reusable_addr);
if wacore_libsignal::protocol::has_session(stores.session_store, &reusable_addr).await? {
continue;
}
if device_jid.is_pn()
&& 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.observe(),
device_jid.observe()
);
record_encryption_override(&mut encryption_overrides, devices.len(), idx, lid_jid);
}
indices_needing_prekeys.push(idx);
}
if !indices_needing_prekeys.is_empty() {
log::debug!(
"Fetching prekeys for {} devices without sessions",
indices_needing_prekeys.len()
);
let jids_for_fetch: Vec<Jid> = indices_needing_prekeys
.iter()
.map(|&i| devices[i].clone())
.collect();
let prekey_bundles = match resolver
.fetch_prekeys_for_identity_check(&jids_for_fetch)
.await
{
Ok(outcome) => {
rejected_devices.extend(
outcome
.rejected
.iter()
.filter(|device| device.code == UNREGISTERED_DEVICE_CODE)
.map(|device| device.jid.clone()),
);
if !rejected_devices.is_empty() {
log::debug!(
"prekey fetch rejected {} of {} device(s) by name",
rejected_devices.len(),
jids_for_fetch.len()
);
had_406 = true;
}
outcome.bundles
}
Err(e) if is_device_unregistered_error(&e) => {
log::debug!(
"Prekey fetch returned 406 for {} device(s); skipping them this round",
jids_for_fetch.len()
);
had_406 = true;
first_error = Some(e);
std::collections::HashMap::new()
}
Err(e) => return Err(e),
};
let prekey_bundles = std::sync::Arc::new(prekey_bundles);
let total = indices_needing_prekeys.len();
let mut next_spawn = 0usize;
let make_session_task = |spawn_idx: usize| {
let idx = indices_needing_prekeys[spawn_idx];
let lookup_jid = devices[idx].clone();
let encryption_jid = encryption_override_at(&encryption_overrides, idx)
.cloned()
.unwrap_or_else(|| lookup_jid.clone());
let bundles = prekey_bundles.clone();
let mut session_store = stores.session_store.clone_box();
let mut identity_store = stores.identity_store.clone_box();
spawn_oneshot(runtime, async move {
let mut addr = crate::types::jid::make_reusable_protocol_address();
encryption_jid.reset_protocol_address(&mut addr);
let Some(bundle) = bundles.get(&lookup_jid) else {
log::debug!(
"No pre-key bundle returned for device {}. This device will be skipped for encryption.",
addr
);
return Ok::<Option<Jid>, anyhow::Error>(None);
};
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
match process_prekey_bundle(
&addr,
&mut *session_store,
&mut *identity_store,
bundle,
&mut rng,
UsePQRatchet::No,
)
.await
{
Ok(IdentityChange::ReplacedExisting) => Ok(Some(encryption_jid)),
Ok(IdentityChange::NewOrUnchanged) => Ok(None),
Err(error) => Err(anyhow::Error::new(error)
.context(format!("failed to process pre-key bundle for {addr}"))),
}
})
};
let mut in_flight: FuturesUnordered<_> = FuturesUnordered::new();
while next_spawn < total && in_flight.len() < ENCRYPT_FANOUT_CONCURRENCY {
in_flight.push(make_session_task(next_spawn));
next_spawn += 1;
}
while let Some(spawn_result) = in_flight.next().await {
match spawn_result {
Ok(Ok(Some(changed_jid))) => resolver.on_local_identity_change(&changed_jid),
Ok(Ok(None)) => {}
Ok(Err(e)) => {
log::warn!("Group session setup failed for a device, skipping it: {e}");
if first_error.is_none() {
first_error = Some(e);
}
}
Err(error) => {
log::warn!(
"Session-establishment task did not deliver a result; skipping device."
);
if first_error.is_none() {
first_error = Some(anyhow::Error::new(error));
}
}
}
if next_spawn < total {
in_flight.push(make_session_task(next_spawn));
next_spawn += 1;
}
}
}
Ok(SessionPlan {
device_count: devices.len(),
encryption_overrides,
had_unregistered_device: had_406,
rejected_devices,
first_error,
})
}
pub async fn encrypt_for_devices_with_sessions(
runtime: &dyn Runtime,
stores: &mut SignalStores<'_>,
devices: &[Jid],
plaintext_to_encrypt: &[u8],
hide_decrypt_fail: bool,
mediatype: Option<&str>,
plan: SessionPlan,
) -> Result<EncryptResult> {
Ok(encrypt_for_devices_with_sessions_detailed(
runtime,
stores,
devices,
plaintext_to_encrypt,
hide_decrypt_fail,
mediatype,
plan,
)
.await?
.result)
}
#[cfg_attr(feature = "tracing", tracing::instrument(name = "wa.send.encrypt_fanout", level = "debug", skip_all, fields(count = devices.len()), err(Debug)))]
pub(crate) async fn encrypt_for_devices_with_sessions_detailed(
runtime: &dyn Runtime,
stores: &mut SignalStores<'_>,
devices: &[Jid],
plaintext_to_encrypt: &[u8],
hide_decrypt_fail: bool,
mediatype: Option<&str>,
plan: SessionPlan,
) -> Result<EncryptAttempt> {
let RawEncryptAttempt {
result: raw,
first_error,
} = encrypt_for_devices_with_sessions_raw_detailed(
runtime,
stores,
devices,
plaintext_to_encrypt,
plan,
)
.await?;
let mut participant_nodes = Vec::with_capacity(raw.devices.len());
let mut encrypted_devices = Vec::with_capacity(raw.devices.len());
for one in raw.devices {
encrypted_devices.push(one.device_jid.clone());
participant_nodes.push(encrypted_device_to_participant_node(
one,
mediatype,
hide_decrypt_fail,
));
}
Ok(EncryptAttempt {
result: EncryptResult {
participant_nodes,
includes_prekey_message: raw.includes_prekey_message,
encrypted_devices,
had_unregistered_device: raw.had_unregistered_device,
rejected_devices: raw.rejected_devices.clone(),
},
first_error,
})
}
pub async fn encrypt_for_devices_with_sessions_raw(
runtime: &dyn Runtime,
stores: &mut SignalStores<'_>,
devices: &[Jid],
plaintext_to_encrypt: &[u8],
plan: SessionPlan,
) -> Result<EncryptForDevicesRaw> {
Ok(encrypt_for_devices_with_sessions_raw_detailed(
runtime,
stores,
devices,
plaintext_to_encrypt,
plan,
)
.await?
.result)
}
#[cfg_attr(feature = "tracing", tracing::instrument(name = "wa.send.encrypt_fanout_raw", level = "debug", skip_all, fields(count = devices.len()), err(Debug)))]
async fn encrypt_for_devices_with_sessions_raw_detailed(
runtime: &dyn Runtime,
stores: &mut SignalStores<'_>,
devices: &[Jid],
plaintext_to_encrypt: &[u8],
plan: SessionPlan,
) -> Result<RawEncryptAttempt> {
debug_assert_eq!(
plan.device_count,
devices.len(),
"SessionPlan built for a different device list"
);
let SessionPlan {
device_count: _,
encryption_overrides,
had_unregistered_device,
rejected_devices,
mut first_error,
} = plan;
let mut encrypted = Vec::with_capacity(devices.len());
let mut includes_prekey_message = false;
if devices.len() == 1 {
let device_jid = devices[0].clone();
let addr = encryption_override_at(&encryption_overrides, 0)
.unwrap_or(&devices[0])
.to_protocol_address();
let res = encrypt_one_device(
plaintext_to_encrypt,
&addr,
&mut *stores.session_store,
&mut *stores.identity_store,
device_jid,
)
.await;
push_raw_result(
res,
&mut encrypted,
&mut includes_prekey_message,
&mut first_error,
);
} else {
let plaintext_arc: std::sync::Arc<[u8]> = std::sync::Arc::from(plaintext_to_encrypt);
let total = devices.len();
let num_chunks = ENCRYPT_FANOUT_CONCURRENCY.min(total);
let mut in_flight: FuturesUnordered<_> = FuturesUnordered::new();
for chunk_idx in 0..num_chunks {
let chunk_start = chunk_idx * total / num_chunks;
let chunk_end = (chunk_idx + 1) * total / num_chunks;
let jobs: Vec<(ProtocolAddress, Jid)> = (chunk_start..chunk_end)
.map(|idx| {
let addr = encryption_override_at(&encryption_overrides, idx)
.unwrap_or(&devices[idx])
.to_protocol_address();
(addr, devices[idx].clone())
})
.collect();
let plaintext = plaintext_arc.clone();
let mut session_store = stores.session_store.clone_box();
let mut identity_store = stores.identity_store.clone_box();
in_flight.push(spawn_oneshot(runtime, async move {
let mut out = Vec::with_capacity(jobs.len());
for (addr, device_jid) in jobs {
out.push(
encrypt_one_device(
&plaintext,
&addr,
&mut *session_store,
&mut *identity_store,
device_jid,
)
.await,
);
}
out
}));
}
while let Some(spawn_result) = in_flight.next().await {
match spawn_result {
Ok(results) => {
for res in results {
push_raw_result(
res,
&mut encrypted,
&mut includes_prekey_message,
&mut first_error,
);
}
}
Err(error) => {
log::warn!(
"Encrypt chunk did not deliver a result; up to ~{} device(s) skipped this send.",
total.div_ceil(num_chunks)
);
if first_error.is_none() {
first_error = Some(anyhow::Error::new(error));
}
}
}
}
}
Ok(RawEncryptAttempt {
result: EncryptForDevicesRaw {
devices: encrypted,
includes_prekey_message,
had_unregistered_device,
rejected_devices,
},
first_error,
})
}
#[cfg(test)]
mod encryption_override_tests {
use super::{SessionPlan, encryption_override_at, record_encryption_override};
use wacore_binary::Jid;
fn lid(user: &str, device: u16) -> Jid {
Jid::lid_device(user.to_owned(), device)
}
#[test]
fn an_empty_map_answers_every_index_without_allocating() {
let overrides: Vec<Option<Jid>> = Vec::new();
assert_eq!(overrides.capacity(), 0, "no override must mean no buffer");
for index in [0, 1, 7, usize::MAX] {
assert!(encryption_override_at(&overrides, index).is_none());
}
let plan = SessionPlan::assume_ready(4);
assert!(
plan.encryption_overrides.is_empty(),
"a plan that overrides nothing must carry no override buffer"
);
assert_eq!(plan.device_count, 4, "the slice length is still recorded");
}
#[test]
fn recording_materializes_the_whole_map_once() {
let mut overrides: Vec<Option<Jid>> = Vec::new();
record_encryption_override(&mut overrides, 3, 2, lid("100000000000001", 5));
assert_eq!(overrides.len(), 3);
assert!(encryption_override_at(&overrides, 0).is_none());
assert!(encryption_override_at(&overrides, 1).is_none());
assert_eq!(
encryption_override_at(&overrides, 2),
Some(&lid("100000000000001", 5))
);
record_encryption_override(&mut overrides, 3, 0, lid("100000000000002", 0));
assert_eq!(overrides.len(), 3);
assert_eq!(
encryption_override_at(&overrides, 0),
Some(&lid("100000000000002", 0))
);
assert_eq!(
encryption_override_at(&overrides, 2),
Some(&lid("100000000000001", 5))
);
record_encryption_override(&mut overrides, 3, 2, lid("100000000000003", 1));
assert_eq!(overrides.len(), 3);
assert_eq!(
encryption_override_at(&overrides, 2),
Some(&lid("100000000000003", 1))
);
}
#[test]
fn a_single_device_plan_records_and_reads_index_zero() {
let mut overrides: Vec<Option<Jid>> = Vec::new();
assert!(encryption_override_at(&overrides, 0).is_none());
record_encryption_override(&mut overrides, 1, 0, lid("100000000000009", 33));
assert_eq!(
encryption_override_at(&overrides, 0),
Some(&lid("100000000000009", 33))
);
assert!(encryption_override_at(&overrides, 1).is_none());
}
}