use super::*;
pub fn retain_skdm_distribution_targets(devices: &mut Vec<Jid>, own_sending_jid: &Jid) {
devices.retain(|device| {
!(device.user == own_sending_jid.user && device.device == own_sending_jid.device)
&& !device.is_hosted()
});
}
pub struct PreparedGroupStanza {
pub node: Node,
pub skdm_devices: Vec<Jid>,
pub stale_device_users: Vec<String>,
pub message_secret: Option<[u8; crate::reporting_token::MESSAGE_SECRET_SIZE]>,
pub sender_identity: Jid,
}
#[derive(Debug, thiserror::Error)]
#[error("required sender-key distribution failed")]
pub struct RequiredSenderKeyDistributionError {
#[source]
source: anyhow::Error,
stale_device_users: Vec<String>,
}
impl RequiredSenderKeyDistributionError {
fn new(source: anyhow::Error, stale_device_users: Vec<String>) -> Self {
Self {
source,
stale_device_users,
}
}
pub fn stale_device_users(&self) -> &[String] {
&self.stale_device_users
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum SenderKeyDistributionPolicy {
#[default]
BestEffort,
Required,
}
pub struct GroupStanzaRequest<'a> {
pub group: &'a GroupInfo,
pub own_jid: &'a Jid,
pub own_lid: &'a Jid,
pub account: Option<&'a wa::ADVSignedDeviceIdentity>,
pub to: &'a Jid,
pub message: &'a wa::Message,
pub message_id: &'a str,
pub force_distribution: bool,
pub distribution_targets: Option<Vec<Jid>>,
pub distribution_policy: SenderKeyDistributionPolicy,
pub phash_devices: Option<&'a ResolvedGroupDevices>,
pub edit: Option<&'a crate::types::message::EditAttribute>,
pub extra_nodes: &'a [Node],
pub pre_encoded: Option<&'a [u8]>,
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(name = "wa.send.group_prepare", level = "debug", skip_all, err(Debug))
)]
pub async fn prepare_group_stanza(
runtime: &dyn Runtime,
stores: &mut SignalStores<'_>,
resolver: &dyn SendContextResolver,
request: GroupStanzaRequest<'_>,
) -> Result<PreparedGroupStanza> {
let GroupStanzaRequest {
group: group_info,
own_jid,
own_lid,
account,
to: to_jid,
message,
message_id: request_id,
force_distribution: force_skdm_distribution,
distribution_targets: skdm_target_devices,
distribution_policy,
phash_devices: all_devices_for_phash,
edit,
extra_nodes: extra_stanza_nodes,
pre_encoded,
} = request;
let (own_sending_jid, _) = match group_info.addressing_mode {
crate::types::message::AddressingMode::Lid => (own_lid.clone(), "lid"),
crate::types::message::AddressingMode::Pn => (own_jid.clone(), "pn"),
};
let shared_content = message.message_context_info.is_unset().then(|| {
pre_encoded.map_or_else(
|| std::borrow::Cow::Owned(waproto::codec::message_to_vec(message)),
std::borrow::Cow::Borrowed,
)
});
let existing_secret = crate::reporting_token::extract_message_secret(message);
let reporting_result = match &shared_content {
Some(content) => generate_reporting_token_from_encoded(
message,
content,
request_id,
to_jid,
to_jid,
existing_secret,
),
None => generate_reporting_token(message, request_id, to_jid, to_jid, existing_secret),
};
let reporting_context = reporting_result.as_ref().map(reporting_context_info);
let own_base_jid = own_sending_jid.to_non_ad();
let mut message_children: Vec<Node> = Vec::new();
let mut includes_prekey_message = false;
let mut phash_for_stanza: Option<CompactString> = None;
let mut skdm_encrypted_devices: Vec<Jid> = Vec::new();
let distribution_list: Option<Vec<Jid>> = if let Some(target_devices) = skdm_target_devices {
if target_devices.is_empty() {
None
} else {
log::debug!(
"SKDM distribution to {} specific devices for group {}",
target_devices.len(),
to_jid.observe()
);
Some(target_devices)
}
} else if force_skdm_distribution {
let mut jids_to_resolve: Vec<Jid> = group_info
.participants
.iter()
.map(|jid| {
let base_jid = jid.to_non_ad();
if base_jid.is_lid()
&& let Some(phone_jid) = group_info.phone_jid_for_lid_user(&base_jid.user)
{
log::debug!(
"Using phone number {} for LID {} device query",
phone_jid.observe(),
base_jid.observe()
);
return phone_jid.to_non_ad();
}
base_jid
})
.collect();
let own_pn_mapping = if own_base_jid.is_lid() {
group_info.phone_jid_for_lid_user(&own_base_jid.user)
} else {
None
};
let own_check_user = own_pn_mapping
.map(|pn| pn.user.as_str())
.unwrap_or(own_base_jid.user.as_str());
if !jids_to_resolve.iter().any(|p| p.user == own_check_user) {
jids_to_resolve.push(match own_pn_mapping {
Some(pn) => pn.to_non_ad(),
None => own_base_jid.clone(),
});
}
crate::types::jid::sort_dedup_by_user(&mut jids_to_resolve);
log::debug!(
"Resolving devices for {} participants",
jids_to_resolve.len()
);
let mut resolved_list = resolver.resolve_devices(&jids_to_resolve).await?;
if group_info.addressing_mode == crate::types::message::AddressingMode::Lid {
resolved_list = resolved_list
.into_iter()
.map(|device_jid| group_info.phone_device_jid_into_lid(device_jid))
.collect();
log::debug!(
"Converted {} devices to LID addressing for group {}",
resolved_list.len(),
to_jid.observe()
);
}
crate::types::jid::sort_dedup_by_device(&mut resolved_list);
let own_user = &own_sending_jid.user;
let own_device = own_sending_jid.device;
let before_filter = resolved_list.len();
retain_skdm_distribution_targets(&mut resolved_list, &own_sending_jid);
log::debug!(
"Filtered SKDM devices from {} to {} (excluded sender {}:{} and hosted devices)",
before_filter,
resolved_list.len(),
own_user,
own_device
);
log::debug!(
"SKDM distribution list for {} resolved to {} devices",
to_jid.observe(),
resolved_list.len(),
);
Some(resolved_list)
} else {
None
};
if distribution_policy == SenderKeyDistributionPolicy::Required
&& distribution_list.as_ref().is_none_or(Vec::is_empty)
{
bail!("required sender-key distribution has no targets");
}
if to_jid.is_group() {
if let Some(resolved) = all_devices_for_phash {
phash_for_stanza = resolved.phash(&own_sending_jid);
} else if let Some(src) = distribution_list.as_deref() {
let phash_set = build_group_phash_set(src, &own_sending_jid);
match MessageUtils::participant_list_hash(&phash_set) {
Ok(phash) => phash_for_stanza = Some(CompactString::new(&phash)),
Err(e) => {
log::warn!(
"Failed to compute group phash for {}: {:?}",
to_jid.observe(),
e
)
}
}
}
}
let mut had_unregistered_devices = false;
let mut skdm_rejected_devices: Vec<Jid> = Vec::new();
let sender_key_name = make_sender_key_name(to_jid, &own_sending_jid.to_protocol_address());
let session_guard = resolver
.lock_device_sessions(distribution_list.as_deref().unwrap_or(&[]))
.await;
let session_plan = match distribution_list.as_deref() {
Some(list) => {
let setup_lock = stores
.sender_key_store
.session_setup_lock(&sender_key_name)
.await;
let _setup_guard = setup_lock.lock().await;
match ensure_sessions_for_devices(runtime, stores, resolver, list).await {
Ok(plan) => Some(plan),
Err(error) if distribution_policy == SenderKeyDistributionPolicy::Required => {
return Err(error.context("required sender-key session setup failed"));
}
Err(e) => {
log::warn!(
"SKDM session setup failed for group {}, continuing without distribution: {e}",
to_jid.observe()
);
if is_device_unregistered_error(&e) {
had_unregistered_devices = true;
}
None
}
}
}
None => None,
};
let plaintext = match &shared_content {
Some(content) => {
MessageUtils::pad_with_context_from_encoded(content, reporting_context.as_ref())
}
None => MessageUtils::encode_and_pad_with_context(message, reporting_context.as_ref()),
};
let chain_lock = stores
.sender_key_store
.sender_key_lock(&sender_key_name)
.await;
let chain_guard = chain_lock.lock().await;
if let Some(ref distribution_list) = distribution_list {
let axolotl_skdm_bytes = create_sender_key_distribution_message_for_group(
stores.sender_key_store,
&sender_key_name,
)
.await?;
if let Some(plan) = session_plan {
let skdm_wrapper_msg = wa::Message {
sender_key_distribution_message: buffa::MessageField::some(
wa::message::SenderKeyDistributionMessage {
group_id: Some(to_jid.to_string()),
axolotl_sender_key_distribution_message: Some(axolotl_skdm_bytes),
},
),
..Default::default()
};
let skdm_plaintext_to_encrypt = MessageUtils::encode_and_pad(&skdm_wrapper_msg);
let skdm_hide_decrypt_fail = should_hide_decrypt_fail_for_send(edit, message);
match encrypt_for_devices_with_sessions_detailed(
runtime,
stores,
distribution_list,
&skdm_plaintext_to_encrypt,
skdm_hide_decrypt_fail,
None,
plan,
)
.await
{
Ok(EncryptAttempt {
result,
first_error,
}) => {
let EncryptResult {
participant_nodes,
includes_prekey_message: result_includes_prekey,
encrypted_devices,
had_unregistered_device,
rejected_devices,
} = result;
if distribution_policy == SenderKeyDistributionPolicy::Required
&& (encrypted_devices.len() != distribution_list.len()
|| first_error.is_some())
{
let error = first_error.unwrap_or_else(|| {
anyhow!(
"sender-key distribution encrypted {} of {} required targets",
encrypted_devices.len(),
distribution_list.len()
)
});
let stale_device_users = stale_users_for(
had_unregistered_device,
&rejected_devices,
Some(distribution_list),
&encrypted_devices,
group_info,
);
return Err(RequiredSenderKeyDistributionError::new(
error,
stale_device_users,
)
.into());
}
includes_prekey_message |= result_includes_prekey;
if had_unregistered_device {
had_unregistered_devices = true;
skdm_rejected_devices.extend(rejected_devices);
}
skdm_encrypted_devices = encrypted_devices;
if !participant_nodes.is_empty() {
message_children.push(
NodeBuilder::new("participants")
.children(participant_nodes)
.build(),
);
let device_identity = match distribution_policy {
SenderKeyDistributionPolicy::BestEffort => {
needs_device_identity(includes_prekey_message, account)
.ok()
.flatten()
}
SenderKeyDistributionPolicy::Required => {
needs_device_identity(includes_prekey_message, account)?
}
};
if let Some(device_identity_bytes) = device_identity {
message_children.push(
NodeBuilder::new("device-identity")
.bytes(device_identity_bytes)
.build(),
);
}
}
}
Err(error) if distribution_policy == SenderKeyDistributionPolicy::Required => {
return Err(RequiredSenderKeyDistributionError::new(error, Vec::new()).into());
}
Err(e) => {
log::warn!(
"SKDM distribution failed for group {}, continuing without it: {e}",
to_jid.observe()
);
if is_device_unregistered_error(&e) {
had_unregistered_devices = true;
}
}
}
}
}
drop(session_guard);
let skmsg = encrypt_group_message(
stores.sender_key_store,
&sender_key_name,
&plaintext,
&mut rand::make_rng::<rand::rngs::StdRng>(),
)
.await?;
drop(chain_guard);
let skmsg_ciphertext = skmsg.into_serialized();
let mediatype = media_type_from_message(message);
let hide_decrypt_fail = should_hide_decrypt_fail_for_send(edit, message);
let mut enc_builder = NodeBuilder::new("enc")
.attr("v", stanza::ENC_VERSION)
.attr("type", stanza::ENC_TYPE_SKMSG);
if let Some(mt) = mediatype {
enc_builder = enc_builder.attr("mediatype", mt);
}
enc_builder = enc_builder.bytes(skmsg_ciphertext);
if hide_decrypt_fail {
enc_builder = enc_builder.attr("decrypt-fail", "hide");
}
let content_node = enc_builder.build();
let stanza_type = stanza_type_from_message(message);
let is_status_broadcast = to_jid.is_status_broadcast();
let mut stanza_builder = NodeBuilder::new("message")
.attr("to", to_jid)
.attr("id", request_id)
.attr("type", stanza_type);
if !is_status_broadcast {
stanza_builder =
stanza_builder.attr("addressing_mode", group_info.addressing_mode.as_str());
}
if let Some(edit_attr) = edit
&& *edit_attr != crate::types::message::EditAttribute::Empty
{
stanza_builder = stanza_builder.attr("edit", edit_attr.to_string_val());
}
message_children.push(content_node);
if let Some(ref result) = reporting_result {
message_children.push(build_reporting_node(result));
}
if let Some(phash) = phash_for_stanza {
stanza_builder = stanza_builder.attr("phash", phash);
}
message_children.extend(extra_stanza_nodes.iter().cloned());
let stanza = stanza_builder.children(message_children).build();
let stale_users = stale_users_for(
had_unregistered_devices,
&skdm_rejected_devices,
distribution_list.as_deref(),
&skdm_encrypted_devices,
group_info,
);
Ok(PreparedGroupStanza {
node: stanza,
skdm_devices: distribution_list.unwrap_or_default(),
stale_device_users: stale_users,
message_secret: reporting_result.map(|r| r.message_secret),
sender_identity: own_sending_jid,
})
}
pub(crate) fn build_group_phash_set(devices: &[Jid], own_sending_jid: &Jid) -> Vec<Jid> {
let mut set: Vec<Jid> = devices.iter().filter(|d| !d.is_hosted()).cloned().collect();
if !set
.iter()
.any(|d| d.user == own_sending_jid.user && d.device == own_sending_jid.device)
{
set.push(own_sending_jid.clone());
}
crate::types::jid::sort_dedup_by_device(&mut set);
set
}
pub(crate) fn stale_users_for(
had_unregistered_device: bool,
rejected_devices: &[Jid],
distribution_list: Option<&[Jid]>,
encrypted_devices: &[Jid],
group_info: &GroupInfo,
) -> Vec<String> {
if !had_unregistered_device {
return Vec::new();
}
if rejected_devices.is_empty() {
return collect_stale_device_users(distribution_list, encrypted_devices, group_info);
}
collect_stale_device_users(Some(rejected_devices), &[], group_info)
}
pub(crate) fn collect_stale_device_users(
distribution_list: Option<&[Jid]>,
skdm_encrypted_devices: &[Jid],
group_info: &GroupInfo,
) -> Vec<String> {
let Some(dist) = distribution_list else {
return Vec::new();
};
let is_lid_mode = group_info.addressing_mode == crate::types::message::AddressingMode::Lid;
let encrypted_set: HashSet<&Jid> = skdm_encrypted_devices.iter().collect();
let mut user_set: HashSet<String> = HashSet::new();
for d in dist {
if encrypted_set.contains(d) {
continue;
}
user_set.insert(d.user.to_string());
if is_lid_mode
&& d.is_lid()
&& let Some(pn_jid) = group_info.phone_jid_for_lid_user(&d.user)
&& pn_jid.is_pn()
{
user_set.insert(pn_jid.user.to_string());
}
}
user_set.into_iter().collect()
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(name = "wa.send.skdm_create", level = "debug", skip_all, err(Debug))
)]
pub async fn create_sender_key_distribution_message_for_group(
store: &mut (dyn SenderKeyStore + Send + Sync),
sender_key_name: &SenderKeyName,
) -> Result<Vec<u8>> {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let skdm = crate::libsignal::protocol::create_sender_key_distribution_message(
sender_key_name,
store,
&mut rng,
)
.await?;
Ok(skdm.into_serialized().into_vec())
}
pub fn build_member_label_message(label: String, ts_secs: i64) -> wa::Message {
wa::Message {
protocol_message: buffa::MessageField::some(wa::message::ProtocolMessage {
r#type: Some(wa::message::protocol_message::Type::GroupMemberLabelChange),
member_label: buffa::MessageField::some(wa::MemberLabel {
label: Some(label),
label_timestamp: Some(ts_secs),
}),
..Default::default()
}),
..Default::default()
}
}