use crate::libsignal::protocol::PublicKey;
use crate::store::traits::DeviceInfo;
const ADV_PREFIX_ACCOUNT_SIGNATURE: &[u8] = &[6, 0];
const ADV_PREFIX_DEVICE_SIGNATURE: &[u8] = &[6, 1];
const ADV_HOSTED_PREFIX_ACCOUNT_SIGNATURE: &[u8] = &[6, 5];
const ADV_HOSTED_PREFIX_DEVICE_SIGNATURE: &[u8] = &[6, 6];
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AdvValidation {
Valid,
Invalid,
NoAccountKey,
}
pub fn validate_adv_with_identity_key(
device_identity_bytes: &[u8],
fetched_identity_key: &[u8; 32],
account_identity_fallback: Option<&[u8; 32]>,
) -> AdvValidation {
let Ok(signed) = waproto::codec::adv_signed_device_identity_decode(device_identity_bytes)
else {
return AdvValidation::Invalid;
};
let (Some(details), Some(account_sig), Some(device_sig)) = (
signed.details.as_deref(),
signed.account_signature.as_deref(),
signed.device_signature.as_deref(),
) else {
return AdvValidation::Invalid;
};
let account_key: &[u8] = match signed.account_signature_key.as_deref() {
Some(k) if !k.is_empty() => k,
_ => match account_identity_fallback {
Some(f) => f.as_slice(),
None => return AdvValidation::NoAccountKey,
},
};
let (Ok(account_pub), Ok(device_pub)) = (
PublicKey::from_djb_public_key_bytes(account_key),
PublicKey::from_djb_public_key_bytes(fetched_identity_key),
) else {
return AdvValidation::Invalid;
};
let verified = [
(ADV_PREFIX_ACCOUNT_SIGNATURE, ADV_PREFIX_DEVICE_SIGNATURE),
(
ADV_HOSTED_PREFIX_ACCOUNT_SIGNATURE,
ADV_HOSTED_PREFIX_DEVICE_SIGNATURE,
),
]
.into_iter()
.any(|(account_prefix, device_prefix)| {
let account_msg = [account_prefix, details, fetched_identity_key].concat();
let device_msg = [device_prefix, details, fetched_identity_key, account_key].concat();
account_pub.verify_signature(&account_msg, account_sig)
&& device_pub.verify_signature(&device_msg, device_sig)
});
if verified {
AdvValidation::Valid
} else {
AdvValidation::Invalid
}
}
#[derive(Debug, Clone)]
pub struct DecodedKeyIndex {
pub raw_id: u32,
pub timestamp: u64,
pub current_index: u32,
pub valid_indexes: Vec<u32>,
}
pub fn decode_key_index_list(signed_bytes: &[u8]) -> Option<DecodedKeyIndex> {
let signed = waproto::codec::adv_signed_key_index_list_decode(signed_bytes).ok()?;
let details_bytes = signed.details.as_ref()?;
let key_index = waproto::codec::adv_key_index_list_decode(details_bytes.as_slice()).ok()?;
let raw_id = key_index.raw_id?;
let timestamp = key_index.timestamp?;
let current_index = key_index.current_index.unwrap_or(0);
Some(DecodedKeyIndex {
raw_id,
timestamp,
current_index,
valid_indexes: key_index.valid_indexes,
})
}
pub fn filter_devices_by_key_index(
devices: &[DeviceInfo],
decoded: &DecodedKeyIndex,
) -> Vec<DeviceInfo> {
devices
.iter()
.filter(|device| should_retain_device(device, decoded))
.cloned()
.collect()
}
pub fn retain_devices_by_key_index(devices: &mut Vec<DeviceInfo>, decoded: &DecodedKeyIndex) {
devices.retain(|device| should_retain_device(device, decoded));
}
fn should_retain_device(device: &DeviceInfo, decoded: &DecodedKeyIndex) -> bool {
device.device_id == 0 || is_key_index_valid(device.key_index, decoded)
}
pub fn is_key_index_valid(key_index: Option<u32>, decoded: &DecodedKeyIndex) -> bool {
match key_index {
Some(ki) => decoded.valid_indexes.contains(&ki) || ki > decoded.current_index,
None => false,
}
}
#[cfg(test)]
#[allow(clippy::disallowed_methods)]
mod tests {
use super::*;
use buffa::Message;
fn dev(id: u32, key_index: Option<u32>) -> DeviceInfo {
DeviceInfo::new(id, key_index)
}
#[test]
fn primary_device_always_kept() {
let devices = vec![dev(0, None), dev(5, Some(3))];
let decoded = DecodedKeyIndex {
raw_id: 1,
timestamp: 100,
current_index: 10,
valid_indexes: vec![], };
let result = filter_devices_by_key_index(&devices, &decoded);
assert_eq!(result.len(), 1);
assert_eq!(result[0].device_id, 0);
}
#[test]
fn valid_index_kept_invalid_removed() {
let devices = vec![dev(0, None), dev(11, Some(5)), dev(12, Some(7))];
let decoded = DecodedKeyIndex {
raw_id: 1,
timestamp: 100,
current_index: 10,
valid_indexes: vec![7], };
let result = filter_devices_by_key_index(&devices, &decoded);
assert_eq!(result.len(), 2); assert!(result.iter().any(|d| d.device_id == 0));
assert!(result.iter().any(|d| d.device_id == 12));
assert!(!result.iter().any(|d| d.device_id == 11));
}
#[test]
fn device_newer_than_current_index_kept() {
let devices = vec![dev(0, None), dev(15, Some(20))];
let decoded = DecodedKeyIndex {
raw_id: 1,
timestamp: 100,
current_index: 10,
valid_indexes: vec![7],
};
let result = filter_devices_by_key_index(&devices, &decoded);
assert_eq!(result.len(), 2); }
#[test]
fn device_without_key_index_removed() {
let devices = vec![dev(0, None), dev(5, None)];
let decoded = DecodedKeyIndex {
raw_id: 1,
timestamp: 100,
current_index: 10,
valid_indexes: vec![7],
};
let result = filter_devices_by_key_index(&devices, &decoded);
assert_eq!(result.len(), 1); assert_eq!(result[0].device_id, 0);
}
#[test]
fn is_key_index_valid_in_valid_set() {
let decoded = DecodedKeyIndex {
raw_id: 1,
timestamp: 100,
current_index: 5,
valid_indexes: vec![3, 7],
};
assert!(is_key_index_valid(Some(3), &decoded));
assert!(is_key_index_valid(Some(7), &decoded));
}
#[test]
fn is_key_index_valid_not_in_valid_set() {
let decoded = DecodedKeyIndex {
raw_id: 1,
timestamp: 100,
current_index: 5,
valid_indexes: vec![3, 7],
};
assert!(!is_key_index_valid(Some(4), &decoded));
}
#[test]
fn is_key_index_valid_newer_than_current() {
let decoded = DecodedKeyIndex {
raw_id: 1,
timestamp: 100,
current_index: 5,
valid_indexes: vec![3],
};
assert!(is_key_index_valid(Some(10), &decoded));
}
#[test]
fn is_key_index_valid_none_rejected() {
let decoded = DecodedKeyIndex {
raw_id: 1,
timestamp: 100,
current_index: 5,
valid_indexes: vec![3, 7],
};
assert!(!is_key_index_valid(None, &decoded));
}
#[test]
fn decode_roundtrip() {
use buffa::Message;
let key_index = waproto::whatsapp::ADVKeyIndexList {
raw_id: Some(42),
timestamp: Some(1000),
current_index: Some(5),
valid_indexes: vec![3, 5, 7],
..Default::default()
};
let details = key_index.encode_to_vec();
let signed = waproto::whatsapp::ADVSignedKeyIndexList {
details: Some(details),
..Default::default()
};
let bytes = signed.encode_to_vec();
let decoded = decode_key_index_list(&bytes).unwrap();
assert_eq!(decoded.raw_id, 42);
assert_eq!(decoded.timestamp, 1000);
assert_eq!(decoded.current_index, 5);
assert_eq!(decoded.valid_indexes, vec![3, 5, 7]);
}
use crate::libsignal::protocol::KeyPair;
fn signed_identity_opts(
account: &KeyPair,
device: &KeyPair,
details: &[u8],
hosted: bool,
include_account_key: bool,
) -> Vec<u8> {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let identity = device.public_key.public_key_bytes();
let account_key = account.public_key.public_key_bytes();
let (acct_prefix, dev_prefix): (&[u8], &[u8]) = if hosted {
(
ADV_HOSTED_PREFIX_ACCOUNT_SIGNATURE,
ADV_HOSTED_PREFIX_DEVICE_SIGNATURE,
)
} else {
(ADV_PREFIX_ACCOUNT_SIGNATURE, ADV_PREFIX_DEVICE_SIGNATURE)
};
let account_sig = account
.private_key
.calculate_signature(&[acct_prefix, details, identity].concat(), &mut rng)
.unwrap()
.to_vec();
let device_sig = device
.private_key
.calculate_signature(
&[dev_prefix, details, identity, account_key].concat(),
&mut rng,
)
.unwrap()
.to_vec();
waproto::whatsapp::ADVSignedDeviceIdentity {
details: Some(details.to_vec()),
account_signature_key: include_account_key.then(|| account_key.to_vec()),
account_signature: Some(account_sig),
device_signature: Some(device_sig),
}
.encode_to_vec()
}
fn signed_identity(
account: &KeyPair,
device: &KeyPair,
details: &[u8],
hosted: bool,
) -> Vec<u8> {
signed_identity_opts(account, device, details, hosted, true)
}
fn id32(kp: &KeyPair) -> [u8; 32] {
kp.public_key.public_key_bytes().try_into().unwrap()
}
#[test]
fn adv_chain_valid_accepted() {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let account = KeyPair::generate(&mut rng);
let device = KeyPair::generate(&mut rng);
let bytes = signed_identity(&account, &device, b"details", false);
assert_eq!(
validate_adv_with_identity_key(&bytes, &id32(&device), None),
AdvValidation::Valid
);
}
#[test]
fn adv_chain_hosted_prefix_accepted() {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let account = KeyPair::generate(&mut rng);
let device = KeyPair::generate(&mut rng);
let bytes = signed_identity(&account, &device, b"hosted-details", true);
assert_eq!(
validate_adv_with_identity_key(&bytes, &id32(&device), None),
AdvValidation::Valid
);
}
#[test]
fn adv_chain_rejects_substituted_identity() {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let account = KeyPair::generate(&mut rng);
let device = KeyPair::generate(&mut rng);
let attacker = KeyPair::generate(&mut rng);
let bytes = signed_identity(&account, &device, b"details", false);
assert_eq!(
validate_adv_with_identity_key(&bytes, &id32(&attacker), None),
AdvValidation::Invalid
);
}
#[test]
fn adv_chain_rejects_missing_device_signature() {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let account = KeyPair::generate(&mut rng);
let device = KeyPair::generate(&mut rng);
let no_dev_sig = waproto::whatsapp::ADVSignedDeviceIdentity {
details: Some(b"details".to_vec()),
account_signature_key: Some(account.public_key.public_key_bytes().to_vec()),
account_signature: Some(vec![0u8; 64]),
device_signature: None,
}
.encode_to_vec();
assert_eq!(
validate_adv_with_identity_key(&no_dev_sig, &id32(&device), None),
AdvValidation::Invalid
);
}
#[test]
fn adv_chain_rejects_garbage() {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let device = KeyPair::generate(&mut rng);
assert_eq!(
validate_adv_with_identity_key(&[1, 2, 3, 4], &id32(&device), None),
AdvValidation::Invalid
);
}
#[test]
fn adv_chain_missing_account_key_verifies_with_fallback() {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let account = KeyPair::generate(&mut rng);
let device = KeyPair::generate(&mut rng);
let bytes = signed_identity_opts(&account, &device, b"details", false, false);
assert_eq!(
validate_adv_with_identity_key(&bytes, &id32(&device), Some(&id32(&account))),
AdvValidation::Valid
);
}
#[test]
fn adv_chain_missing_account_key_no_fallback_is_no_account_key() {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let account = KeyPair::generate(&mut rng);
let device = KeyPair::generate(&mut rng);
let bytes = signed_identity_opts(&account, &device, b"details", false, false);
assert_eq!(
validate_adv_with_identity_key(&bytes, &id32(&device), None),
AdvValidation::NoAccountKey
);
}
#[test]
fn adv_chain_missing_account_key_wrong_fallback_is_invalid() {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let account = KeyPair::generate(&mut rng);
let device = KeyPair::generate(&mut rng);
let attacker = KeyPair::generate(&mut rng);
let bytes = signed_identity_opts(&account, &device, b"details", false, false);
assert_eq!(
validate_adv_with_identity_key(&bytes, &id32(&device), Some(&id32(&attacker))),
AdvValidation::Invalid
);
}
#[test]
fn adv_chain_in_blob_key_takes_precedence_over_fallback() {
let mut rng = rand::make_rng::<rand::rngs::StdRng>();
let account = KeyPair::generate(&mut rng);
let device = KeyPair::generate(&mut rng);
let unrelated = KeyPair::generate(&mut rng);
let bytes = signed_identity(&account, &device, b"details", false);
assert_eq!(
validate_adv_with_identity_key(&bytes, &id32(&device), Some(&id32(&unrelated))),
AdvValidation::Valid
);
}
}