use crate::iq::node::{extract_content_bytes, extract_content_uint, required_child};
use crate::iq::spec::IqSpec;
use crate::prekeys::PreKeyUtils;
use crate::protocol::ProtocolNode;
use crate::request::InfoQuery;
use anyhow::anyhow;
use wacore_binary::builder::NodeBuilder;
use wacore_binary::encoder::{ByteWriter, EncodeNode, Encoder};
use wacore_binary::{Jid, Server};
use wacore_binary::{Node, NodeContent, NodeContentRef, NodeRef};
pub use crate::libsignal::protocol::{PreKeyBundle, PublicKey};
#[derive(Debug, Clone)]
pub struct PreKeyCountResponse {
pub count: usize,
}
#[derive(Debug, Clone, Default)]
pub struct PreKeyCountSpec;
impl PreKeyCountSpec {
pub fn new() -> Self {
Self
}
}
impl IqSpec for PreKeyCountSpec {
type Response = PreKeyCountResponse;
fn build_iq(&self) -> InfoQuery<'static> {
let count_node = NodeBuilder::new("count").build();
InfoQuery::get(
"encrypt",
Jid::new("", Server::Pn),
Some(NodeContent::Nodes(vec![count_node])),
)
}
fn parse_response(&self, response: &NodeRef<'_>) -> Result<Self::Response, anyhow::Error> {
let count_node = response
.get_optional_child("count")
.ok_or_else(|| anyhow!("Missing <count> node in response"))?;
let count_str = count_node.attrs().optional_string("value");
let count_str = count_str.as_deref().unwrap_or("0");
let count = count_str.parse::<usize>().unwrap_or(0);
Ok(PreKeyCountResponse { count })
}
}
#[derive(Debug, Clone, PartialEq, Eq, crate::WireEnum)]
pub enum PreKeyFetchReason {
#[wire = "identity"]
Identity,
#[wire = "retry"]
Retry,
#[wire_fallback]
Other(String),
}
#[derive(Debug, Clone)]
pub struct PreKeyFetchSpec {
pub jids: Vec<Jid>,
pub reason: Option<PreKeyFetchReason>,
}
impl PreKeyFetchSpec {
pub fn new(jids: Vec<Jid>) -> Self {
Self { jids, reason: None }
}
pub fn with_reason(jids: Vec<Jid>, reason: PreKeyFetchReason) -> Self {
Self {
jids,
reason: Some(reason),
}
}
}
impl IqSpec for PreKeyFetchSpec {
type Response = std::collections::HashMap<Jid, PreKeyBundle>;
fn build_iq(&self) -> InfoQuery<'static> {
let content = PreKeyUtils::build_fetch_prekeys_request(
&self.jids,
self.reason.as_ref().map(|r| r.as_str()),
);
InfoQuery::get(
"encrypt",
Jid::new("", Server::Pn),
Some(NodeContent::Nodes(vec![content])),
)
}
fn parse_response(&self, response: &NodeRef<'_>) -> Result<Self::Response, anyhow::Error> {
PreKeyUtils::parse_prekeys_response(response)
}
}
#[derive(Debug, Clone, Default)]
pub struct DigestKeyBundleSpec;
impl DigestKeyBundleSpec {
pub fn new() -> Self {
Self
}
}
#[derive(Debug, Clone)]
pub struct DigestKeyBundleResponse {
pub reg_id: u32,
pub identity: Vec<u8>,
pub skey_id: u32,
pub skey_pubkey: Vec<u8>,
pub skey_signature: Vec<u8>,
pub prekey_ids: Vec<u32>,
pub hash: Vec<u8>,
}
impl IqSpec for DigestKeyBundleSpec {
type Response = DigestKeyBundleResponse;
fn build_iq(&self) -> InfoQuery<'static> {
let digest_node = NodeBuilder::new("digest").build();
InfoQuery::get(
"encrypt",
Jid::new("", Server::Pn),
Some(NodeContent::Nodes(vec![digest_node])),
)
}
fn parse_response(&self, response: &NodeRef<'_>) -> Result<Self::Response, anyhow::Error> {
let digest_node = response
.get_optional_child("digest")
.ok_or_else(|| anyhow::anyhow!("missing <digest> child in response"))?;
let reg_node = required_child(digest_node, "registration")?;
let reg_id = extract_content_uint(Some(reg_node));
let identity_node = required_child(digest_node, "identity")?;
let identity = match identity_node.content.as_deref() {
Some(NodeContentRef::Bytes(b)) if !b.is_empty() => b.to_vec(),
_ => return Err(anyhow!("missing or empty bytes in <identity>")),
};
let skey_node = digest_node.get_optional_child("skey");
let (skey_id, skey_pubkey, skey_signature) = if let Some(skey) = skey_node {
(
extract_content_uint(skey.get_optional_child("id")),
extract_content_bytes(skey.get_optional_child("value")),
extract_content_bytes(skey.get_optional_child("signature")),
)
} else {
(0, Vec::new(), Vec::new())
};
let prekey_ids = digest_node
.get_optional_child("list")
.and_then(|list| list.children())
.map(|children| {
children
.iter()
.map(|child| extract_content_uint(Some(child)))
.collect()
})
.unwrap_or_default();
let hash_node = required_child(digest_node, "hash")?;
let hash = match hash_node.content.as_deref() {
Some(NodeContentRef::Bytes(b)) if !b.is_empty() => b.to_vec(),
_ => return Err(anyhow!("missing or empty bytes in <hash>")),
};
Ok(DigestKeyBundleResponse {
reg_id,
identity,
skey_id,
skey_pubkey,
skey_signature,
prekey_ids,
hash,
})
}
}
#[derive(Debug, Clone)]
pub struct PreKeyUploadSpec {
pub registration_id: u32,
pub identity_key: PublicKey,
pub signed_pre_key_id: u32,
pub signed_pre_key_public: PublicKey,
pub signed_pre_key_signature: Vec<u8>,
pub pre_keys: Vec<(u32, PublicKey)>,
}
impl PreKeyUploadSpec {
pub fn new(
registration_id: u32,
identity_key: PublicKey,
signed_pre_key_id: u32,
signed_pre_key_public: PublicKey,
signed_pre_key_signature: Vec<u8>,
pre_keys: Vec<(u32, PublicKey)>,
) -> Self {
Self {
registration_id,
identity_key,
signed_pre_key_id,
signed_pre_key_public,
signed_pre_key_signature,
pre_keys,
}
}
}
struct PreKeyUploadIqNode<'a> {
request_id: &'a str,
spec: &'a PreKeyUploadSpec,
}
impl EncodeNode for PreKeyUploadIqNode<'_> {
fn tag(&self) -> &str {
"iq"
}
fn attrs_len(&self) -> usize {
4 }
fn has_content(&self) -> bool {
true
}
fn encode_attrs<'a, W: ByteWriter>(
&self,
encoder: &mut Encoder<'a, W>,
) -> wacore_binary::Result<()> {
encoder.write_string("id")?;
encoder.write_string(self.request_id)?;
encoder.write_string("xmlns")?;
encoder.write_string("encrypt")?;
encoder.write_string("type")?;
encoder.write_string("set")?;
encoder.write_string("to")?;
encoder.write_jid_owned(&Jid::new("", Server::Pn))?;
Ok(())
}
fn encode_content<'a, W: ByteWriter>(
&self,
encoder: &mut Encoder<'a, W>,
) -> wacore_binary::Result<()> {
let spec = self.spec;
encoder.write_list_start(5)?;
encoder.write_list_start(2)?; encoder.write_string("registration")?;
encoder.write_bytes_with_len(&spec.registration_id.to_be_bytes())?;
encoder.write_list_start(2)?;
encoder.write_string("type")?;
encoder.write_bytes_with_len(&[5u8])?;
encoder.write_list_start(2)?;
encoder.write_string("identity")?;
encoder.write_bytes_with_len(spec.identity_key.public_key_bytes())?;
encoder.write_list_start(2)?; encoder.write_string("list")?;
encoder.write_list_start(spec.pre_keys.len())?;
for &(id, ref pk) in &spec.pre_keys {
encoder.write_list_start(2)?; encoder.write_string("key")?;
encoder.write_list_start(2)?; encoder.write_list_start(2)?; encoder.write_string("id")?;
encoder.write_bytes_with_len(&id.to_be_bytes()[1..])?;
encoder.write_list_start(2)?; encoder.write_string("value")?;
encoder.write_bytes_with_len(pk.public_key_bytes())?;
}
encoder.write_list_start(2)?; encoder.write_string("skey")?;
encoder.write_list_start(3)?; encoder.write_list_start(2)?;
encoder.write_string("id")?;
encoder.write_bytes_with_len(&spec.signed_pre_key_id.to_be_bytes()[1..])?;
encoder.write_list_start(2)?;
encoder.write_string("value")?;
encoder.write_bytes_with_len(spec.signed_pre_key_public.public_key_bytes())?;
encoder.write_list_start(2)?;
encoder.write_string("signature")?;
encoder.write_bytes_with_len(&spec.signed_pre_key_signature)?;
Ok(())
}
}
impl IqSpec for PreKeyUploadSpec {
type Response = ();
fn build_iq(&self) -> InfoQuery<'static> {
let content = PreKeyUtils::build_upload_prekeys_request(
self.registration_id,
self.identity_key.public_key_bytes(),
self.signed_pre_key_id,
self.signed_pre_key_public.public_key_bytes(),
&self.signed_pre_key_signature,
self.pre_keys
.iter()
.map(|(id, pk)| (*id, pk.public_key_bytes())),
);
InfoQuery::set(
"encrypt",
Jid::new("", Server::Pn),
Some(NodeContent::Nodes(content)),
)
}
fn encode_iq_direct(&self, request_id: &str, out: &mut Vec<u8>) -> Result<bool, anyhow::Error> {
out.reserve(self.pre_keys.len() * 40 + 256);
let node = PreKeyUploadIqNode {
request_id,
spec: self,
};
let mut encoder = Encoder::new_vec(out)?;
encoder.write_node(&node)?;
Ok(true)
}
fn parse_response(&self, _response: &NodeRef<'_>) -> Result<Self::Response, anyhow::Error> {
Ok(())
}
}
fn truncate_to_3bytes(id: u32) -> Vec<u8> {
debug_assert!(id <= 0x00FF_FFFF, "prekey id exceeds 3-byte range: {id}");
id.to_be_bytes()[1..].to_vec()
}
fn expand_from_3bytes(bytes: &[u8]) -> Result<u32, anyhow::Error> {
if bytes.len() != 3 {
return Err(anyhow!("Expected 3 bytes for ID, got {}", bytes.len()));
}
Ok(u32::from_be_bytes([0, bytes[0], bytes[1], bytes[2]]))
}
#[derive(Debug, Clone)]
pub struct SignedPreKeyNode {
pub id: u32,
pub public_bytes: Vec<u8>,
pub signature: Vec<u8>,
}
impl SignedPreKeyNode {
pub fn new(id: u32, public_bytes: Vec<u8>, signature: Vec<u8>) -> Self {
Self {
id,
public_bytes,
signature,
}
}
}
impl ProtocolNode for SignedPreKeyNode {
fn tag(&self) -> &'static str {
"skey"
}
fn into_node(self) -> Node {
NodeBuilder::new("skey")
.children([
NodeBuilder::new("id")
.bytes(truncate_to_3bytes(self.id))
.build(),
NodeBuilder::new("value").bytes(self.public_bytes).build(),
NodeBuilder::new("signature").bytes(self.signature).build(),
])
.build()
}
fn try_from_node_ref(node: &NodeRef<'_>) -> Result<Self, anyhow::Error> {
if node.tag != "skey" {
return Err(anyhow!("expected <skey>, got <{}>", node.tag));
}
let id_node = required_child(node, "id")?;
let id_bytes = match id_node.content.as_deref() {
Some(NodeContentRef::Bytes(b)) => b,
_ => return Err(anyhow!("missing bytes in <id>")),
};
let id = expand_from_3bytes(id_bytes)?;
let value_node = required_child(node, "value")?;
let public_bytes = match value_node.content.as_deref() {
Some(NodeContentRef::Bytes(b)) => b.to_vec(),
_ => return Err(anyhow!("missing bytes in <value>")),
};
if public_bytes.len() != 32 {
return Err(anyhow!("signed prekey public key must be 32 bytes"));
}
let sig_node = required_child(node, "signature")?;
let signature = match sig_node.content.as_deref() {
Some(NodeContentRef::Bytes(b)) => b.to_vec(),
_ => return Err(anyhow!("missing bytes in <signature>")),
};
if signature.len() != 64 {
return Err(anyhow!("signed prekey signature must be 64 bytes"));
}
Ok(Self {
id,
public_bytes,
signature,
})
}
}
#[derive(Debug, Clone)]
pub struct OneTimePreKeyNode {
pub id: u32,
pub public_bytes: Vec<u8>,
}
impl OneTimePreKeyNode {
pub fn new(id: u32, public_bytes: Vec<u8>) -> Self {
Self { id, public_bytes }
}
}
impl ProtocolNode for OneTimePreKeyNode {
fn tag(&self) -> &'static str {
"key"
}
fn into_node(self) -> Node {
NodeBuilder::new("key")
.children([
NodeBuilder::new("id")
.bytes(truncate_to_3bytes(self.id))
.build(),
NodeBuilder::new("value").bytes(self.public_bytes).build(),
])
.build()
}
fn try_from_node_ref(node: &NodeRef<'_>) -> Result<Self, anyhow::Error> {
if node.tag != "key" {
return Err(anyhow!("expected <key>, got <{}>", node.tag));
}
let id_node = required_child(node, "id")?;
let id_bytes = match id_node.content.as_deref() {
Some(NodeContentRef::Bytes(b)) => b,
_ => return Err(anyhow!("missing bytes in <id>")),
};
let id = expand_from_3bytes(id_bytes)?;
let value_node = required_child(node, "value")?;
let public_bytes = match value_node.content.as_deref() {
Some(NodeContentRef::Bytes(b)) => b.to_vec(),
_ => return Err(anyhow!("missing bytes in <value>")),
};
if public_bytes.len() != 32 {
return Err(anyhow!("one-time prekey public key must be 32 bytes"));
}
Ok(Self { id, public_bytes })
}
}
#[derive(Debug, Clone)]
pub struct PreKeyBundleUserNode {
pub jid: Jid,
pub registration_id: u32,
pub identity_key: Vec<u8>, pub signed_pre_key: SignedPreKeyNode,
pub one_time_pre_key: Option<OneTimePreKeyNode>,
pub device_identity: Option<Vec<u8>>, }
impl PreKeyBundleUserNode {
pub fn from_bundle(
jid: Jid,
bundle: &PreKeyBundle,
device_identity: Option<Vec<u8>>,
) -> Result<Self, anyhow::Error> {
let registration_id = bundle.registration_id()?;
let identity_key = bundle
.identity_key()?
.public_key()
.public_key_bytes()
.to_vec();
let signed_pre_key_id: u32 = bundle.signed_pre_key_id()?.into();
let signed_pre_key_public = bundle.signed_pre_key_public()?.public_key_bytes().to_vec();
let signed_pre_key_signature = bundle.signed_pre_key_signature()?.to_vec();
let signed_pre_key = SignedPreKeyNode::new(
signed_pre_key_id,
signed_pre_key_public,
signed_pre_key_signature,
);
let one_time_pre_key = match (bundle.pre_key_id()?, bundle.pre_key_public()?) {
(Some(id), Some(pk)) => {
let pre_key_id: u32 = id.into();
let public_bytes = pk.public_key_bytes().to_vec();
Some(OneTimePreKeyNode::new(pre_key_id, public_bytes))
}
_ => None,
};
Ok(Self {
jid,
registration_id,
identity_key,
signed_pre_key,
one_time_pre_key,
device_identity,
})
}
}
impl ProtocolNode for PreKeyBundleUserNode {
fn tag(&self) -> &'static str {
"user"
}
fn into_node(self) -> Node {
let mut children = vec![
NodeBuilder::new("registration")
.bytes(self.registration_id.to_be_bytes().to_vec())
.build(),
NodeBuilder::new("type").bytes(vec![5]).build(),
NodeBuilder::new("identity")
.bytes(self.identity_key)
.build(),
self.signed_pre_key.into_node(),
];
if let Some(otpk) = self.one_time_pre_key {
children.push(otpk.into_node());
}
if let Some(dev_id) = self.device_identity {
children.push(NodeBuilder::new("device-identity").bytes(dev_id).build());
}
NodeBuilder::new("user")
.attr("jid", self.jid)
.attr("type", "result")
.children(children)
.build()
}
fn try_from_node_ref(node: &NodeRef<'_>) -> Result<Self, anyhow::Error> {
if node.tag != "user" {
return Err(anyhow!("expected <user>, got <{}>", node.tag));
}
let jid = node
.attrs()
.optional_jid("jid")
.ok_or_else(|| anyhow!("missing required attribute jid"))?;
let reg_node = required_child(node, "registration")?;
let reg_bytes = match reg_node.content.as_deref() {
Some(NodeContentRef::Bytes(b)) => b,
_ => return Err(anyhow!("missing bytes in <registration>")),
};
if reg_bytes.len() != 4 {
return Err(anyhow!("registration ID must be 4 bytes"));
}
let registration_id =
u32::from_be_bytes([reg_bytes[0], reg_bytes[1], reg_bytes[2], reg_bytes[3]]);
let identity_node = required_child(node, "identity")?;
let identity_key = match identity_node.content.as_deref() {
Some(NodeContentRef::Bytes(b)) => b.to_vec(),
_ => return Err(anyhow!("missing bytes in <identity>")),
};
if identity_key.len() != 32 {
return Err(anyhow!("identity key must be 32 bytes"));
}
let skey_node = required_child(node, "skey")?;
let signed_pre_key = SignedPreKeyNode::try_from_node_ref(skey_node)?;
let one_time_pre_key = match node.get_optional_child("key") {
Some(n) => Some(OneTimePreKeyNode::try_from_node_ref(n)?),
None => None,
};
let device_identity = match node.get_optional_child("device-identity") {
Some(n) => match n.content.as_deref() {
Some(NodeContentRef::Bytes(b)) => Some(b.to_vec()),
_ => return Err(anyhow!("device-identity must be bytes")),
},
None => None,
};
Ok(Self {
jid,
registration_id,
identity_key,
signed_pre_key,
one_time_pre_key,
device_identity,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_prekey_count_spec_build_iq() {
let spec = PreKeyCountSpec::new();
let iq = spec.build_iq();
assert_eq!(iq.namespace, "encrypt");
assert_eq!(iq.query_type, crate::request::InfoQueryType::Get);
if let Some(NodeContent::Nodes(nodes)) = &iq.content {
assert_eq!(nodes.len(), 1);
assert_eq!(nodes[0].tag, "count");
} else {
panic!("Expected NodeContent::Nodes");
}
}
#[test]
fn test_prekey_count_spec_parse_response() {
let spec = PreKeyCountSpec::new();
let response = NodeBuilder::new("iq")
.attr("type", "result")
.children([NodeBuilder::new("count").attr("value", "42").build()])
.build();
let result = spec.parse_response(&response.as_node_ref()).unwrap();
assert_eq!(result.count, 42);
}
#[test]
fn test_prekey_count_spec_parse_response_missing_value() {
let spec = PreKeyCountSpec::new();
let response = NodeBuilder::new("iq")
.attr("type", "result")
.children([NodeBuilder::new("count").build()])
.build();
let result = spec.parse_response(&response.as_node_ref()).unwrap();
assert_eq!(result.count, 0); }
#[test]
fn test_prekey_fetch_spec_build_iq() {
let jids = vec![
"1234567890:0@s.whatsapp.net".parse().unwrap(),
"0987654321:0@s.whatsapp.net".parse().unwrap(),
];
let spec = PreKeyFetchSpec::new(jids);
let iq = spec.build_iq();
assert_eq!(iq.namespace, "encrypt");
assert_eq!(iq.query_type, crate::request::InfoQueryType::Get);
if let Some(NodeContent::Nodes(nodes)) = &iq.content {
assert_eq!(nodes.len(), 1);
assert_eq!(nodes[0].tag, "key");
} else {
panic!("Expected NodeContent::Nodes");
}
}
#[test]
fn test_prekey_fetch_spec_with_reason() {
let jids = vec!["1234567890:0@s.whatsapp.net".parse().unwrap()];
let spec = PreKeyFetchSpec::with_reason(jids, PreKeyFetchReason::Retry);
assert_eq!(spec.reason, Some(PreKeyFetchReason::Retry));
}
#[test]
fn test_digest_key_bundle_spec_build_iq() {
let spec = DigestKeyBundleSpec::new();
let iq = spec.build_iq();
assert_eq!(iq.namespace, "encrypt");
assert_eq!(iq.query_type, crate::request::InfoQueryType::Get);
if let Some(NodeContent::Nodes(nodes)) = &iq.content {
assert_eq!(nodes.len(), 1);
assert_eq!(nodes[0].tag, "digest");
} else {
panic!("Expected NodeContent::Nodes");
}
}
#[test]
fn test_digest_key_bundle_spec_parse_response() {
let spec = DigestKeyBundleSpec::new();
let hash_bytes = vec![0xAA; 20];
let response = NodeBuilder::new("iq")
.attr("type", "result")
.children([NodeBuilder::new("digest")
.children([
NodeBuilder::new("registration")
.bytes(12345u32.to_be_bytes().to_vec())
.build(),
NodeBuilder::new("type").bytes(vec![5]).build(),
NodeBuilder::new("identity").bytes(vec![0x01; 32]).build(),
NodeBuilder::new("skey")
.children([
NodeBuilder::new("id").bytes(vec![0x00, 0x00, 0x01]).build(),
NodeBuilder::new("value").bytes(vec![0x02; 32]).build(),
NodeBuilder::new("signature").bytes(vec![0x03; 64]).build(),
])
.build(),
NodeBuilder::new("list")
.children([
NodeBuilder::new("id").bytes(vec![0x00, 0x00, 0x0A]).build(),
NodeBuilder::new("id").bytes(vec![0x00, 0x00, 0x0B]).build(),
])
.build(),
NodeBuilder::new("hash").bytes(hash_bytes.clone()).build(),
])
.build()])
.build();
let result = spec.parse_response(&response.as_node_ref()).unwrap();
assert_eq!(result.reg_id, 12345);
assert_eq!(result.identity, vec![0x01; 32]);
assert_eq!(result.skey_id, 1);
assert_eq!(result.skey_pubkey, vec![0x02; 32]);
assert_eq!(result.skey_signature, vec![0x03; 64]);
assert_eq!(result.prekey_ids, vec![10, 11]);
assert_eq!(result.hash, hash_bytes);
}
#[test]
fn test_digest_key_bundle_spec_parse_response_empty() {
let spec = DigestKeyBundleSpec::new();
let response = NodeBuilder::new("iq").attr("type", "result").build();
assert!(spec.parse_response(&response.as_node_ref()).is_err());
}
#[test]
fn test_prekey_upload_spec_build_iq() {
let pk = |b| PublicKey::from_djb_public_key_bytes(&[b; 32]).unwrap();
let spec = PreKeyUploadSpec::new(
12345, pk(1), 1, pk(2), vec![3u8; 64], vec![(100, pk(4)), (101, pk(5))], );
let iq = spec.build_iq();
assert_eq!(iq.namespace, "encrypt");
assert_eq!(iq.query_type, crate::request::InfoQueryType::Set);
if let Some(NodeContent::Nodes(nodes)) = &iq.content {
assert_eq!(nodes.len(), 5);
assert_eq!(nodes[0].tag, "registration");
assert_eq!(nodes[1].tag, "type");
assert_eq!(nodes[2].tag, "identity");
assert_eq!(nodes[3].tag, "list");
assert_eq!(nodes[4].tag, "skey");
if let Some(list_children) = nodes[3].children() {
assert_eq!(list_children.len(), 2);
assert_eq!(list_children[0].tag, "key");
assert_eq!(list_children[1].tag, "key");
} else {
panic!("Expected list to have children");
}
} else {
panic!("Expected NodeContent::Nodes");
}
}
#[test]
fn test_prekey_upload_spec_parse_response() {
let pk = |b| PublicKey::from_djb_public_key_bytes(&[b; 32]).unwrap();
let spec = PreKeyUploadSpec::new(12345, pk(1), 1, pk(2), vec![3u8; 64], vec![(100, pk(4))]);
let response = NodeBuilder::new("iq").attr("type", "result").build();
let result = spec.parse_response(&response.as_node_ref());
assert!(result.is_ok());
}
#[test]
fn test_truncate_to_3bytes() {
assert_eq!(truncate_to_3bytes(0x00345678), vec![0x34, 0x56, 0x78]);
assert_eq!(truncate_to_3bytes(0x00000001), vec![0x00, 0x00, 0x01]);
assert_eq!(truncate_to_3bytes(0x00ABCDEF), vec![0xAB, 0xCD, 0xEF]);
}
#[test]
fn test_expand_from_3bytes() {
assert_eq!(expand_from_3bytes(&[0x34, 0x56, 0x78]).unwrap(), 0x00345678);
assert_eq!(expand_from_3bytes(&[0x00, 0x00, 0x01]).unwrap(), 0x00000001);
assert_eq!(expand_from_3bytes(&[0xAB, 0xCD, 0xEF]).unwrap(), 0x00ABCDEF);
}
#[test]
fn test_signed_prekey_node_round_trip() {
let original = SignedPreKeyNode::new(0x00345678, vec![1; 32], vec![2; 64]);
let node = original.clone().into_node();
assert_eq!(node.tag, "skey");
let parsed = SignedPreKeyNode::try_from_node(&node).unwrap();
assert_eq!(parsed.id, original.id);
assert_eq!(parsed.public_bytes, original.public_bytes);
assert_eq!(parsed.signature, original.signature);
}
#[test]
fn test_onetime_prekey_node_round_trip() {
let original = OneTimePreKeyNode::new(0xABCDEF, vec![3; 32]);
let node = original.clone().into_node();
assert_eq!(node.tag, "key");
let parsed = OneTimePreKeyNode::try_from_node(&node).unwrap();
assert_eq!(parsed.id, original.id);
assert_eq!(parsed.public_bytes, original.public_bytes);
}
#[test]
fn test_prekey_bundle_user_node_structure() {
let jid: Jid = "1234567890:33@s.whatsapp.net".parse().unwrap();
let user_node = PreKeyBundleUserNode {
jid: jid.clone(),
registration_id: 12345,
identity_key: vec![1; 32],
signed_pre_key: SignedPreKeyNode::new(100, vec![2; 32], vec![3; 64]),
one_time_pre_key: Some(OneTimePreKeyNode::new(200, vec![4; 32])),
device_identity: Some(vec![5; 128]),
};
let node = user_node.into_node();
assert_eq!(node.tag, "user");
assert_eq!(node.attrs().optional_jid("jid"), Some(jid));
assert_eq!(
node.attrs().optional_string("type").as_deref(),
Some("result")
);
if let Some(children) = node.children() {
assert_eq!(children.len(), 6);
assert_eq!(children[0].tag, "registration");
assert_eq!(children[1].tag, "type");
assert_eq!(children[2].tag, "identity");
assert_eq!(children[3].tag, "skey");
assert_eq!(children[4].tag, "key");
assert_eq!(children[5].tag, "device-identity");
} else {
panic!("Expected user node to have children");
}
}
#[test]
fn test_prekey_bundle_user_node_round_trip() {
let jid: Jid = "1234567890:33@s.whatsapp.net".parse().unwrap();
let original = PreKeyBundleUserNode {
jid: jid.clone(),
registration_id: 12345,
identity_key: vec![1; 32],
signed_pre_key: SignedPreKeyNode::new(100, vec![2; 32], vec![3; 64]),
one_time_pre_key: Some(OneTimePreKeyNode::new(200, vec![4; 32])),
device_identity: Some(vec![5; 128]),
};
let node = original.clone().into_node();
let parsed = PreKeyBundleUserNode::try_from_node(&node).unwrap();
assert_eq!(parsed.jid, original.jid);
assert_eq!(parsed.registration_id, original.registration_id);
assert_eq!(parsed.identity_key, original.identity_key);
assert_eq!(parsed.signed_pre_key.id, original.signed_pre_key.id);
assert_eq!(
parsed.signed_pre_key.public_bytes,
original.signed_pre_key.public_bytes
);
assert_eq!(
parsed.signed_pre_key.signature,
original.signed_pre_key.signature
);
assert!(parsed.one_time_pre_key.is_some());
assert_eq!(
parsed.one_time_pre_key.as_ref().unwrap().id,
original.one_time_pre_key.as_ref().unwrap().id
);
assert_eq!(parsed.device_identity, original.device_identity);
}
#[test]
fn test_prekey_bundle_user_node_without_optional_fields() {
let jid: Jid = "1234567890:0@s.whatsapp.net".parse().unwrap();
let original = PreKeyBundleUserNode {
jid: jid.clone(),
registration_id: 54321,
identity_key: vec![9; 32],
signed_pre_key: SignedPreKeyNode::new(500, vec![8; 32], vec![7; 64]),
one_time_pre_key: None,
device_identity: None,
};
let node = original.clone().into_node();
if let Some(children) = node.children() {
assert_eq!(children.len(), 4);
assert_eq!(children[0].tag, "registration");
assert_eq!(children[1].tag, "type");
assert_eq!(children[2].tag, "identity");
assert_eq!(children[3].tag, "skey");
} else {
panic!("Expected user node to have children");
}
let parsed = PreKeyBundleUserNode::try_from_node(&node).unwrap();
assert_eq!(parsed.jid, original.jid);
assert_eq!(parsed.registration_id, original.registration_id);
assert!(parsed.one_time_pre_key.is_none());
assert!(parsed.device_identity.is_none());
}
fn assert_direct_encode_matches_marshal(num_prekeys: u32) {
use crate::iq::spec::IqSpec;
use crate::libsignal::protocol::KeyPair;
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let identity = KeyPair::generate(&mut rng);
let signed_prekey = KeyPair::generate(&mut rng);
let sig = identity
.private_key
.calculate_signature(&signed_prekey.public_key.serialize(), &mut rng)
.unwrap();
let pre_keys: Vec<(u32, crate::libsignal::protocol::PublicKey)> = (1..=num_prekeys)
.map(|id| {
let kp = KeyPair::generate(&mut rng);
(id, kp.public_key)
})
.collect();
let spec = PreKeyUploadSpec::new(
12345,
identity.public_key,
1,
signed_prekey.public_key,
sig.to_vec(),
pre_keys,
);
let request_id = "test-req-id-123";
let mut direct_buf = Vec::new();
let used = spec
.encode_iq_direct(request_id, &mut direct_buf)
.expect("encode_iq_direct should succeed");
assert!(used);
let iq = spec.build_iq();
let iq_node = wacore_binary::builder::NodeBuilder::new("iq")
.attr("id", request_id)
.attr("xmlns", iq.namespace)
.attr("type", iq.query_type.as_str())
.attr("to", iq.to)
.apply_content(iq.content)
.build();
let marshal_buf =
wacore_binary::marshal_auto(&iq_node).expect("marshal_auto should succeed");
assert_eq!(
direct_buf, marshal_buf,
"encode_iq_direct must match build_iq + marshal for {num_prekeys} prekeys"
);
}
#[test]
fn test_encode_iq_direct_matches_marshal_5_keys() {
assert_direct_encode_matches_marshal(5);
}
#[test]
fn test_encode_iq_direct_matches_marshal_0_keys() {
assert_direct_encode_matches_marshal(0);
}
#[test]
fn test_encode_iq_direct_matches_marshal_1_key() {
assert_direct_encode_matches_marshal(1);
}
}