use anyhow::{Result, anyhow, ensure};
use sha2::{Digest, Sha256};
use crate::secret_enc_addon::{
AddonContext, MESSAGE_SECRET_SIZE, ModificationType, build_aad, decrypt_addon, encrypt_addon,
};
const GCM_IV_SIZE: usize = 12;
const GCM_TAG_SIZE: usize = 16;
fn poll_vote_addon_ctx<'a>(
stanza_id: &'a str,
poll_creator_jid: &'a str,
voter_jid: &'a str,
) -> AddonContext<'a> {
AddonContext {
stanza_id,
parent_msg_original_sender: poll_creator_jid,
modification_sender: voter_jid,
modification_type: ModificationType::PollVote,
}
}
pub fn compute_option_hash(option_name: &str) -> [u8; 32] {
let mut hasher = Sha256::new();
hasher.update(option_name.as_bytes());
hasher.finalize().into()
}
pub fn derive_vote_encryption_key(
message_secret: &[u8],
stanza_id: &str,
poll_creator_jid: &str,
voter_jid: &str,
) -> Result<[u8; 32]> {
crate::secret_enc_addon::derive_use_case_secret(
message_secret,
&poll_vote_addon_ctx(stanza_id, poll_creator_jid, voter_jid),
)
}
pub fn encrypt_poll_vote_with_secret(
selected_option_hashes: &[Vec<u8>],
message_secret: &[u8],
stanza_id: &str,
poll_creator_jid: &str,
voter_jid: &str,
) -> Result<(Vec<u8>, [u8; GCM_IV_SIZE])> {
let plaintext = encode_selected_options(selected_option_hashes);
encrypt_addon(
&plaintext,
message_secret,
&poll_vote_addon_ctx(stanza_id, poll_creator_jid, voter_jid),
)
}
fn encode_selected_options(selected_option_hashes: &[Vec<u8>]) -> Vec<u8> {
use buffa::encoding::{Tag, WireType};
let encoded_len = selected_option_hashes
.iter()
.map(|hash| 1 + buffa::types::bytes_encoded_len(hash))
.sum();
let mut plaintext = Vec::with_capacity(encoded_len);
for hash in selected_option_hashes {
Tag::new(1, WireType::LengthDelimited).encode(&mut plaintext);
buffa::types::encode_bytes(hash, &mut plaintext);
}
plaintext
}
pub fn decrypt_poll_vote(
enc_payload: &[u8],
iv: &[u8],
encryption_key: &[u8; 32],
stanza_id: &str,
voter_jid: &str,
) -> Result<Vec<Vec<u8>>> {
use crate::libsignal::crypto::aes_256_gcm_decrypt;
let nonce: &[u8; GCM_IV_SIZE] = iv
.try_into()
.map_err(|_| anyhow!("Invalid IV size: expected {GCM_IV_SIZE}, got {}", iv.len()))?;
if enc_payload.len() < GCM_TAG_SIZE {
return Err(anyhow!(
"Encrypted payload too short: need at least {GCM_TAG_SIZE} bytes for tag, got {}",
enc_payload.len()
));
}
let aad = build_aad(&poll_vote_addon_ctx(stanza_id, "", voter_jid));
let mut plaintext = Vec::with_capacity(enc_payload.len().saturating_sub(GCM_TAG_SIZE));
aes_256_gcm_decrypt(encryption_key, nonce, &aad, enc_payload, &mut plaintext)
.map_err(|_| anyhow!("Poll vote GCM tag verification failed"))?;
decode_selected_options(&plaintext)
}
#[derive(Debug, Clone, Copy)]
pub struct PollVoteAddressing<'a> {
pub poll_creator_jid: &'a str,
pub voter_jid: &'a str,
}
#[derive(Debug, Clone, Copy)]
pub struct PollVoteCiphertext<'a> {
pub enc_payload: &'a [u8],
pub enc_iv: &'a [u8],
}
pub fn decrypt_poll_vote_with_fallback(
ciphertext: PollVoteCiphertext<'_>,
message_secret: &[u8],
stanza_id: &str,
primary: PollVoteAddressing<'_>,
fallback: Option<PollVoteAddressing<'_>>,
) -> Result<Vec<Vec<u8>>> {
match decrypt_poll_vote_with_secret(
ciphertext,
message_secret,
stanza_id,
primary.poll_creator_jid,
primary.voter_jid,
) {
Ok(v) => Ok(v),
Err(primary_err) => match fallback {
Some(fb) => decrypt_poll_vote_with_secret(
ciphertext,
message_secret,
stanza_id,
fb.poll_creator_jid,
fb.voter_jid,
)
.map_err(|fb_err| {
anyhow!("poll vote decrypt failed: primary={primary_err}; fallback={fb_err}")
}),
None => Err(primary_err),
},
}
}
pub fn visit_decrypted_poll_vote_with_fallback<F>(
enc_payload: &[u8],
iv: &[u8],
message_secret: &[u8],
stanza_id: &str,
primary: PollVoteAddressing<'_>,
fallback: Option<PollVoteAddressing<'_>>,
mut visit: F,
) -> Result<()>
where
F: FnMut(&[u8]),
{
match visit_poll_vote_with_secret(
enc_payload,
iv,
message_secret,
stanza_id,
primary.poll_creator_jid,
primary.voter_jid,
&mut visit,
) {
Ok(()) => Ok(()),
Err(primary_err) => match fallback {
Some(fb) => visit_poll_vote_with_secret(
enc_payload,
iv,
message_secret,
stanza_id,
fb.poll_creator_jid,
fb.voter_jid,
&mut visit,
)
.map_err(|fb_err| {
anyhow!("poll vote decrypt failed: primary={primary_err}; fallback={fb_err}")
}),
None => Err(primary_err),
},
}
}
pub fn decrypt_poll_vote_with_secret(
ciphertext: PollVoteCiphertext<'_>,
message_secret: &[u8],
stanza_id: &str,
poll_creator_jid: &str,
voter_jid: &str,
) -> Result<Vec<Vec<u8>>> {
let plaintext = decrypt_poll_vote_payload_with_secret(
ciphertext,
message_secret,
stanza_id,
poll_creator_jid,
voter_jid,
)?;
decode_selected_options(&plaintext)
}
pub fn decrypt_poll_vote_payload_with_secret(
ciphertext: PollVoteCiphertext<'_>,
message_secret: &[u8],
stanza_id: &str,
poll_creator_jid: &str,
voter_jid: &str,
) -> Result<Vec<u8>> {
ensure!(
message_secret.len() == MESSAGE_SECRET_SIZE,
"message_secret must be {MESSAGE_SECRET_SIZE} bytes, got {}",
message_secret.len()
);
decrypt_addon(
ciphertext.enc_payload,
ciphertext.enc_iv,
message_secret,
&poll_vote_addon_ctx(stanza_id, poll_creator_jid, voter_jid),
)
}
fn visit_poll_vote_with_secret<F>(
enc_payload: &[u8],
iv: &[u8],
message_secret: &[u8],
stanza_id: &str,
poll_creator_jid: &str,
voter_jid: &str,
visit: &mut F,
) -> Result<()>
where
F: FnMut(&[u8]),
{
let plaintext = decrypt_poll_vote_payload_with_secret(
PollVoteCiphertext {
enc_payload,
enc_iv: iv,
},
message_secret,
stanza_id,
poll_creator_jid,
voter_jid,
)?;
scan_selected_options(&plaintext, |_| {})?;
scan_selected_options(&plaintext, visit)
}
fn decode_selected_options(plaintext: &[u8]) -> Result<Vec<Vec<u8>>> {
let mut selected_options = Vec::new();
scan_selected_options(plaintext, |selected| {
selected_options.push(selected.to_vec());
})?;
Ok(selected_options)
}
fn scan_selected_options<'a, F>(plaintext: &'a [u8], mut visit: F) -> Result<()>
where
F: FnMut(&'a [u8]),
{
use buffa::encoding::{Tag, WireType, skip_field_depth};
let mut cur = plaintext;
while !cur.is_empty() {
let tag = Tag::decode(&mut cur)?;
match tag.field_number() {
1 => {
if tag.wire_type() != WireType::LengthDelimited {
return Err(buffa::DecodeError::WireTypeMismatch {
field_number: 1,
expected: WireType::LengthDelimited as u8,
actual: tag.wire_type() as u8,
}
.into());
}
visit(buffa::types::borrow_bytes(&mut cur)?);
}
_ => skip_field_depth(tag, &mut cur, buffa::RECURSION_LIMIT)?,
}
}
Ok(())
}
#[cfg(test)]
#[allow(clippy::disallowed_methods)]
mod tests {
use super::*;
#[test]
fn option_hash_deterministic() {
let h1 = compute_option_hash("Option A");
let h2 = compute_option_hash("Option A");
let h3 = compute_option_hash("Option B");
assert_eq!(h1, h2);
assert_ne!(h1, h3);
assert_eq!(h1.len(), 32);
}
#[test]
fn vote_encrypt_decrypt_roundtrip() {
use buffa::Message;
let secret = [0xCDu8; 32];
let stanza_id = "3EB0ABCD1234";
let creator = "creator@s.whatsapp.net";
let voter = "voter@s.whatsapp.net";
let hashes = vec![
compute_option_hash("Yes").to_vec(),
compute_option_hash("No").to_vec(),
];
let (enc, iv) =
encrypt_poll_vote_with_secret(&hashes, &secret, stanza_id, creator, voter).unwrap();
let out = decrypt_poll_vote_with_secret(
PollVoteCiphertext {
enc_payload: &enc,
enc_iv: &iv,
},
&secret,
stanza_id,
creator,
voter,
)
.unwrap();
assert_eq!(out, hashes);
let plaintext = decrypt_poll_vote_payload_with_secret(
PollVoteCiphertext {
enc_payload: &enc,
enc_iv: &iv,
},
&secret,
stanza_id,
creator,
voter,
)
.unwrap();
let vote_message = waproto::whatsapp::message::PollVoteMessage {
selected_options: hashes,
};
assert_eq!(plaintext, vote_message.encode_to_vec());
}
#[test]
fn payload_decrypt_rejects_invalid_message_secret_before_ciphertext() {
let error = decrypt_poll_vote_payload_with_secret(
PollVoteCiphertext {
enc_payload: &[],
enc_iv: &[],
},
&[0u8; MESSAGE_SECRET_SIZE - 1],
"id",
"creator@s.whatsapp.net",
"voter@s.whatsapp.net",
)
.unwrap_err();
assert!(
error
.to_string()
.contains("message_secret must be 32 bytes")
);
}
#[test]
fn selected_option_encoding_matches_message_encoding() {
use buffa::Message;
let hashes = vec![
compute_option_hash("Yes").to_vec(),
compute_option_hash("No").to_vec(),
];
let vote_msg = waproto::whatsapp::message::PollVoteMessage {
selected_options: hashes.clone(),
};
assert_eq!(encode_selected_options(&hashes), vote_msg.encode_to_vec());
}
#[test]
fn legacy_decrypt_path_still_works() {
let secret = [0xCDu8; 32];
let stanza_id = "3EB0ABCD1234";
let creator = "creator@s.whatsapp.net";
let voter = "voter@s.whatsapp.net";
let hashes = vec![compute_option_hash("Yes").to_vec()];
let (enc, iv) =
encrypt_poll_vote_with_secret(&hashes, &secret, stanza_id, creator, voter).unwrap();
let key = derive_vote_encryption_key(&secret, stanza_id, creator, voter).unwrap();
let out = decrypt_poll_vote(&enc, &iv, &key, stanza_id, voter).unwrap();
assert_eq!(out, hashes);
}
#[test]
fn wrong_voter_fails() {
let secret = [0xEFu8; 32];
let (enc, iv) = encrypt_poll_vote_with_secret(
&[compute_option_hash("Yes").to_vec()],
&secret,
"id",
"c@s.whatsapp.net",
"v@s.whatsapp.net",
)
.unwrap();
assert!(
decrypt_poll_vote_with_secret(
PollVoteCiphertext {
enc_payload: &enc,
enc_iv: &iv,
},
&secret,
"id",
"c@s.whatsapp.net",
"wrong@s.whatsapp.net"
)
.is_err()
);
}
#[test]
fn fallback_recovers_across_addressing() {
let secret = [0x33u8; 32];
let stanza_id = "3EB0FALLBACK";
let creator_pn = "5511999999999@s.whatsapp.net";
let voter_pn = "5511888888888@s.whatsapp.net";
let (enc, iv) = encrypt_poll_vote_with_secret(
&[compute_option_hash("Yes").to_vec()],
&secret,
stanza_id,
creator_pn,
voter_pn,
)
.unwrap();
let creator_lid = "111111111111111@lid";
let voter_lid = "222222222222222@lid";
let out = decrypt_poll_vote_with_fallback(
PollVoteCiphertext {
enc_payload: &enc,
enc_iv: &iv,
},
&secret,
stanza_id,
PollVoteAddressing {
poll_creator_jid: creator_lid,
voter_jid: voter_lid,
},
Some(PollVoteAddressing {
poll_creator_jid: creator_pn,
voter_jid: voter_pn,
}),
)
.unwrap();
assert_eq!(out, vec![compute_option_hash("Yes").to_vec()]);
}
#[test]
fn visit_fallback_yields_selected_hashes() {
let secret = [0x34u8; 32];
let stanza_id = "3EB0VISIT";
let creator_pn = "5511999999999@s.whatsapp.net";
let voter_pn = "5511888888888@s.whatsapp.net";
let expected = compute_option_hash("Yes");
let (enc, iv) = encrypt_poll_vote_with_secret(
&[expected.to_vec()],
&secret,
stanza_id,
creator_pn,
voter_pn,
)
.unwrap();
let mut visited = Vec::new();
visit_decrypted_poll_vote_with_fallback(
&enc,
&iv,
&secret,
stanza_id,
PollVoteAddressing {
poll_creator_jid: "111111111111111@lid",
voter_jid: "222222222222222@lid",
},
Some(PollVoteAddressing {
poll_creator_jid: creator_pn,
voter_jid: voter_pn,
}),
|hash| visited.push(<[u8; 32]>::try_from(hash).unwrap()),
)
.unwrap();
assert_eq!(visited, vec![expected]);
}
#[test]
fn fallback_primary_succeeds_without_using_fallback() {
let secret = [0x44u8; 32];
let stanza_id = "id";
let creator = "c@s.whatsapp.net";
let voter = "v@s.whatsapp.net";
let (enc, iv) = encrypt_poll_vote_with_secret(
&[compute_option_hash("A").to_vec()],
&secret,
stanza_id,
creator,
voter,
)
.unwrap();
let out = decrypt_poll_vote_with_fallback(
PollVoteCiphertext {
enc_payload: &enc,
enc_iv: &iv,
},
&secret,
stanza_id,
PollVoteAddressing {
poll_creator_jid: creator,
voter_jid: voter,
},
Some(PollVoteAddressing {
poll_creator_jid: "wrong@lid",
voter_jid: "wrong@lid",
}),
)
.unwrap();
assert_eq!(out, vec![compute_option_hash("A").to_vec()]);
}
#[test]
fn fallback_combined_error_when_both_fail() {
let secret = [0x55u8; 32];
let stanza_id = "id";
let (enc, iv) = encrypt_poll_vote_with_secret(
&[compute_option_hash("A").to_vec()],
&secret,
stanza_id,
"c@s.whatsapp.net",
"v@s.whatsapp.net",
)
.unwrap();
let err = decrypt_poll_vote_with_fallback(
PollVoteCiphertext {
enc_payload: &enc,
enc_iv: &iv,
},
&secret,
stanza_id,
PollVoteAddressing {
poll_creator_jid: "x@lid",
voter_jid: "y@lid",
},
Some(PollVoteAddressing {
poll_creator_jid: "x@s.whatsapp.net",
voter_jid: "y@s.whatsapp.net",
}),
)
.unwrap_err();
let s = err.to_string();
assert!(s.contains("primary="), "got: {s}");
assert!(s.contains("fallback="), "got: {s}");
}
#[test]
fn fallback_none_propagates_primary_error() {
let secret = [0x66u8; 32];
let (enc, iv) = encrypt_poll_vote_with_secret(
&[compute_option_hash("A").to_vec()],
&secret,
"id",
"c@s.whatsapp.net",
"v@s.whatsapp.net",
)
.unwrap();
assert!(
decrypt_poll_vote_with_fallback(
PollVoteCiphertext {
enc_payload: &enc,
enc_iv: &iv,
},
&secret,
"id",
PollVoteAddressing {
poll_creator_jid: "wrong@lid",
voter_jid: "wrong@lid",
},
None,
)
.is_err()
);
}
#[test]
fn empty_vote_roundtrip() {
let secret = [0xEFu8; 32];
let (enc, iv) = encrypt_poll_vote_with_secret(
&[],
&secret,
"id",
"c@s.whatsapp.net",
"v@s.whatsapp.net",
)
.unwrap();
let out = decrypt_poll_vote_with_secret(
PollVoteCiphertext {
enc_payload: &enc,
enc_iv: &iv,
},
&secret,
"id",
"c@s.whatsapp.net",
"v@s.whatsapp.net",
)
.unwrap();
assert!(out.is_empty());
}
}