use crate::iq::prekeys::{OneTimePreKeyNode, SignedPreKeyNode};
use crate::libsignal::protocol::PublicKey;
use crate::protocol::ProtocolNode;
use wacore_binary::builder::NodeBuilder;
use wacore_binary::{Node, NodeContent, NodeRef};
pub const MAX_RETRY_COUNT: u8 = 5;
pub const MIN_RETRY_COUNT_FOR_KEYS: u8 = 2;
pub const MIN_RETRY_FOR_BASE_KEY_CHECK: u8 = 2;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
#[allow(dead_code)] pub enum RetryReason {
UnknownError = 0,
NoSession = 1,
InvalidKey = 2,
InvalidKeyId = 3,
InvalidMessage = 4,
InvalidSignature = 5,
FutureMessage = 6,
BadMac = 7,
InvalidSession = 8,
InvalidMsgKey = 9,
BadBroadcastEphemeralSetting = 10,
UnknownCompanionNoPrekey = 11,
AdvFailure = 12,
StatusRevokeDelay = 13,
}
impl RetryReason {
pub fn as_str(&self) -> &'static str {
match self {
Self::UnknownError => "unknown",
Self::NoSession => "no_session",
Self::InvalidKey => "invalid_key",
Self::InvalidKeyId => "invalid_key_id",
Self::InvalidMessage => "invalid_message",
Self::InvalidSignature => "invalid_signature",
Self::FutureMessage => "future_message",
Self::BadMac => "bad_mac",
Self::InvalidSession => "invalid_session",
Self::InvalidMsgKey => "invalid_msg_key",
Self::BadBroadcastEphemeralSetting => "bad_broadcast_ephemeral",
Self::UnknownCompanionNoPrekey => "unknown_companion",
Self::AdvFailure => "adv_failure",
Self::StatusRevokeDelay => "status_revoke_delay",
}
}
}
pub fn get_bytes_content(node: &Node) -> Option<&[u8]> {
match &node.content {
Some(NodeContent::Bytes(b)) => Some(b.as_slice()),
_ => None,
}
}
fn parse_registration_id(bytes: &[u8]) -> Option<u32> {
match bytes.len() {
4 => Some(u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]])),
n if n > 4 => None,
0 => None,
_ => {
let mut arr = [0u8; 4];
let start = 4 - bytes.len();
arr[start..].copy_from_slice(bytes);
Some(u32::from_be_bytes(arr))
}
}
}
pub fn extract_registration_id_from_node(node: &Node) -> Option<u32> {
let registration_node = node.get_optional_child("registration")?;
parse_registration_id(get_bytes_content(registration_node)?)
}
pub fn extract_registration_id_from_node_ref(node: &NodeRef<'_>) -> Option<u32> {
let registration_node = node.get_optional_child("registration")?;
parse_registration_id(registration_node.content_bytes()?)
}
pub fn should_include_keys(retry_count: u8, _reason: RetryReason) -> bool {
should_include_keys_with_policy(retry_count, false, false)
}
pub fn should_include_keys_with_policy(
retry_count: u8,
force_include_keys: bool,
is_stateless: bool,
) -> bool {
force_include_keys || is_stateless || retry_count >= MIN_RETRY_COUNT_FOR_KEYS
}
pub fn should_drop_unknown_device_retry(keys_present: bool, device_known: bool) -> bool {
!keys_present && !device_known
}
pub fn build_retry_keys_node(
identity_pub: &PublicKey,
prekey_id: u32,
prekey_pub: &PublicKey,
signed_prekey_id: u32,
signed_prekey_pub: &PublicKey,
signed_prekey_signature: Vec<u8>,
device_identity: Vec<u8>,
) -> Node {
NodeBuilder::new("keys")
.children([
NodeBuilder::new("type").bytes(vec![5u8]).build(),
NodeBuilder::new("identity")
.bytes(identity_pub.public_key_bytes().to_vec())
.build(),
OneTimePreKeyNode::new(prekey_id, prekey_pub.public_key_bytes().to_vec()).into_node(),
SignedPreKeyNode::new(
signed_prekey_id,
signed_prekey_pub.public_key_bytes().to_vec(),
signed_prekey_signature,
)
.into_node(),
NodeBuilder::new("device-identity")
.bytes(device_identity)
.build(),
])
.build()
}
#[cfg(test)]
mod tests {
use super::*;
use std::borrow::Cow;
use wacore_binary::Attrs;
#[test]
fn get_bytes_content_extracts_bytes() {
let node = Node {
tag: Cow::Borrowed("test"),
attrs: Attrs::new(),
content: Some(NodeContent::Bytes(vec![1, 2, 3, 4])),
};
assert_eq!(get_bytes_content(&node), Some(&[1, 2, 3, 4][..]));
}
#[test]
fn get_bytes_content_returns_none_for_string() {
let node = Node {
tag: Cow::Borrowed("test"),
attrs: Attrs::new(),
content: Some(NodeContent::String("hello".into())),
};
assert_eq!(get_bytes_content(&node), None);
}
#[test]
fn get_bytes_content_returns_none_for_empty() {
let node = Node {
tag: Cow::Borrowed("test"),
attrs: Attrs::new(),
content: None,
};
assert_eq!(get_bytes_content(&node), None);
}
#[test]
fn extract_registration_id_4_bytes() {
let reg_node = Node {
tag: Cow::Borrowed("registration"),
attrs: Attrs::new(),
content: Some(NodeContent::Bytes(vec![0x00, 0x01, 0x02, 0x03])),
};
let parent = Node {
tag: Cow::Borrowed("receipt"),
attrs: Attrs::new(),
content: Some(NodeContent::Nodes(vec![reg_node])),
};
assert_eq!(extract_registration_id_from_node(&parent), Some(0x00010203));
}
#[test]
fn extract_registration_id_3_bytes() {
let reg_node = Node {
tag: Cow::Borrowed("registration"),
attrs: Attrs::new(),
content: Some(NodeContent::Bytes(vec![0x01, 0x02, 0x03])),
};
let parent = Node {
tag: Cow::Borrowed("receipt"),
attrs: Attrs::new(),
content: Some(NodeContent::Nodes(vec![reg_node])),
};
assert_eq!(extract_registration_id_from_node(&parent), Some(0x00010203));
}
#[test]
fn extract_registration_id_missing() {
let parent = Node {
tag: Cow::Borrowed("receipt"),
attrs: Attrs::new(),
content: Some(NodeContent::Nodes(vec![])),
};
assert_eq!(extract_registration_id_from_node(&parent), None);
}
#[test]
fn extract_registration_id_empty_bytes() {
let reg_node = Node {
tag: Cow::Borrowed("registration"),
attrs: Attrs::new(),
content: Some(NodeContent::Bytes(vec![])),
};
let parent = Node {
tag: Cow::Borrowed("receipt"),
attrs: Attrs::new(),
content: Some(NodeContent::Nodes(vec![reg_node])),
};
assert_eq!(extract_registration_id_from_node(&parent), None);
}
#[test]
fn extract_registration_id_rejects_oversized() {
let reg_node = Node {
tag: Cow::Borrowed("registration"),
attrs: Attrs::new(),
content: Some(NodeContent::Bytes(vec![0x01, 0x02, 0x03, 0x04, 0x05])),
};
let parent = Node {
tag: Cow::Borrowed("receipt"),
attrs: Attrs::new(),
content: Some(NodeContent::Nodes(vec![reg_node])),
};
assert_eq!(extract_registration_id_from_node(&parent), None);
assert_eq!(
extract_registration_id_from_node_ref(&parent.as_node_ref()),
None
);
}
#[test]
fn extract_registration_id_from_node_ref_matches_owned() {
let reg_node = Node {
tag: Cow::Borrowed("registration"),
attrs: Attrs::new(),
content: Some(NodeContent::Bytes(vec![0x00, 0x01, 0x02, 0x03])),
};
let parent = Node {
tag: Cow::Borrowed("receipt"),
attrs: Attrs::new(),
content: Some(NodeContent::Nodes(vec![reg_node])),
};
assert_eq!(
extract_registration_id_from_node_ref(&parent.as_node_ref()),
Some(0x00010203)
);
let reg_node_short = Node {
tag: Cow::Borrowed("registration"),
attrs: Attrs::new(),
content: Some(NodeContent::Bytes(vec![0x01, 0x02, 0x03])),
};
let parent_short = Node {
tag: Cow::Borrowed("receipt"),
attrs: Attrs::new(),
content: Some(NodeContent::Nodes(vec![reg_node_short])),
};
assert_eq!(
extract_registration_id_from_node_ref(&parent_short.as_node_ref()),
Some(0x00010203)
);
}
#[test]
fn should_not_include_keys_on_first_normal_retry() {
assert!(!should_include_keys(1, RetryReason::NoSession));
assert!(!should_include_keys(1, RetryReason::InvalidMessage));
}
#[test]
fn should_include_keys_when_forced() {
assert!(should_include_keys_with_policy(1, true, false));
}
#[test]
fn should_include_keys_for_stateless_recipient() {
assert!(should_include_keys_with_policy(1, false, true));
}
#[test]
fn should_include_keys_at_retry_threshold() {
assert!(should_include_keys(2, RetryReason::UnknownError));
}
#[test]
fn should_include_keys_after_retry_threshold() {
assert!(should_include_keys(3, RetryReason::BadMac));
}
#[test]
fn drop_unknown_device_retry_only_without_bundle() {
assert!(!should_drop_unknown_device_retry(true, false));
assert!(should_drop_unknown_device_retry(false, false));
assert!(!should_drop_unknown_device_retry(true, true));
assert!(!should_drop_unknown_device_retry(false, true));
}
#[test]
fn constants_match_wa_web() {
assert_eq!(MAX_RETRY_COUNT, 5);
assert_eq!(MIN_RETRY_COUNT_FOR_KEYS, 2);
assert_eq!(MIN_RETRY_FOR_BASE_KEY_CHECK, 2);
}
#[test]
fn retry_keys_bundle_emits_raw_32_byte_curve_values() {
use crate::libsignal::protocol::KeyPair;
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let identity = KeyPair::generate(&mut rng);
let prekey = KeyPair::generate(&mut rng);
let signed_prekey = KeyPair::generate(&mut rng);
let keys = build_retry_keys_node(
&identity.public_key,
201,
&prekey.public_key,
100,
&signed_prekey.public_key,
vec![7u8; 64],
vec![9u8; 16],
);
let key_node = keys.get_optional_child("key").expect("<key> child present");
let parsed_prekey =
OneTimePreKeyNode::try_from_node(key_node).expect("one-time prekey should parse");
assert_eq!(
parsed_prekey.public_bytes.as_slice(),
prekey.public_key.public_key_bytes()
);
let skey_node = keys
.get_optional_child("skey")
.expect("<skey> child present");
let parsed_skey =
SignedPreKeyNode::try_from_node(skey_node).expect("signed prekey should parse");
assert_eq!(
parsed_skey.public_bytes.as_slice(),
signed_prekey.public_key.public_key_bytes()
);
}
#[test]
fn retry_reason_as_str_is_stable() {
let cases = [
(RetryReason::UnknownError, "unknown"),
(RetryReason::NoSession, "no_session"),
(RetryReason::InvalidKey, "invalid_key"),
(RetryReason::InvalidKeyId, "invalid_key_id"),
(RetryReason::InvalidMessage, "invalid_message"),
(RetryReason::InvalidSignature, "invalid_signature"),
(RetryReason::FutureMessage, "future_message"),
(RetryReason::BadMac, "bad_mac"),
(RetryReason::InvalidSession, "invalid_session"),
(RetryReason::InvalidMsgKey, "invalid_msg_key"),
(
RetryReason::BadBroadcastEphemeralSetting,
"bad_broadcast_ephemeral",
),
(RetryReason::UnknownCompanionNoPrekey, "unknown_companion"),
(RetryReason::AdvFailure, "adv_failure"),
(RetryReason::StatusRevokeDelay, "status_revoke_delay"),
];
for (reason, expected) in cases {
assert_eq!(reason.as_str(), expected);
}
}
}