use crate::constants;
use crate::error::{Error, Result};
use crate::identity::{self, HybridSignature, IdentityPublicKey, IdentitySecretKey};
use crate::primitives::{hkdf, xwing};
use zeroize::{Zeroize, ZeroizeOnDrop, Zeroizing};
pub struct PreKeyBundle {
pub ik_pub: IdentityPublicKey,
pub crypto_version: String,
pub spk_pub: xwing::PublicKey,
pub spk_id: u32,
pub spk_sig: HybridSignature,
pub opk_pub: Option<xwing::PublicKey>,
pub opk_id: Option<u32>,
}
pub struct VerifiedBundle(PreKeyBundle);
impl std::ops::Deref for VerifiedBundle {
type Target = PreKeyBundle;
fn deref(&self) -> &PreKeyBundle {
&self.0
}
}
pub struct SessionInit {
pub crypto_version: String,
pub sender_ik_fingerprint: [u8; 32],
pub recipient_ik_fingerprint: [u8; 32],
pub sender_ek: xwing::PublicKey,
pub ct_ik: xwing::Ciphertext,
pub ct_spk: xwing::Ciphertext,
pub spk_id: u32,
pub ct_opk: Option<xwing::Ciphertext>,
pub opk_id: Option<u32>,
}
#[derive(Zeroize, ZeroizeOnDrop)]
pub struct InitiatedSession {
#[zeroize(skip)]
pub session_init: SessionInit,
root_key: Zeroizing<[u8; 32]>,
initial_chain_key: Zeroizing<[u8; 32]>,
#[zeroize(skip)]
pub ek_pk: xwing::PublicKey,
#[zeroize(skip)]
ek_sk: xwing::SecretKey,
#[zeroize(skip)]
pub sender_sig: HybridSignature,
#[zeroize(skip)]
pub opk_used: bool,
}
impl InitiatedSession {
pub fn take_root_key(&mut self) -> Zeroizing<[u8; 32]> {
std::mem::replace(&mut self.root_key, Zeroizing::new([0u8; 32]))
}
pub fn take_initial_chain_key(&mut self) -> Zeroizing<[u8; 32]> {
std::mem::replace(&mut self.initial_chain_key, Zeroizing::new([0u8; 32]))
}
pub fn ek_sk(&self) -> &xwing::SecretKey {
&self.ek_sk
}
#[cfg(all(feature = "test-utils", debug_assertions))]
#[deprecated(note = "test-utils only — do not call in production code")]
pub fn root_key_ptr(&self) -> *const u8 {
self.root_key.as_ptr()
}
#[cfg(all(feature = "test-utils", debug_assertions))]
#[deprecated(note = "test-utils only — do not call in production code")]
pub fn initial_chain_key_ptr(&self) -> *const u8 {
self.initial_chain_key.as_ptr()
}
}
#[derive(Zeroize, ZeroizeOnDrop)]
pub struct ReceivedSession {
root_key: Zeroizing<[u8; 32]>,
initial_chain_key: Zeroizing<[u8; 32]>,
#[zeroize(skip)]
pub peer_ek: xwing::PublicKey,
}
impl ReceivedSession {
pub fn take_root_key(&mut self) -> Zeroizing<[u8; 32]> {
std::mem::replace(&mut self.root_key, Zeroizing::new([0u8; 32]))
}
pub fn take_initial_chain_key(&mut self) -> Zeroizing<[u8; 32]> {
std::mem::replace(&mut self.initial_chain_key, Zeroizing::new([0u8; 32]))
}
#[cfg(all(feature = "test-utils", debug_assertions))]
#[deprecated(note = "test-utils only — do not call in production code")]
pub fn root_key_ptr(&self) -> *const u8 {
self.root_key.as_ptr()
}
#[cfg(all(feature = "test-utils", debug_assertions))]
#[deprecated(note = "test-utils only — do not call in production code")]
pub fn initial_chain_key_ptr(&self) -> *const u8 {
self.initial_chain_key.as_ptr()
}
}
#[must_use = "bundle verification result must be checked"]
pub fn verify_bundle(bundle: PreKeyBundle, known_ik: &IdentityPublicKey) -> Result<VerifiedBundle> {
if bundle.opk_pub.is_some() != bundle.opk_id.is_some() {
return Err(Error::InvalidData);
}
if bundle.ik_pub != *known_ik {
return Err(Error::BundleVerificationFailed);
}
if bundle.crypto_version != constants::CRYPTO_VERSION {
return Err(Error::BundleVerificationFailed);
}
let mut msg =
Vec::with_capacity(constants::SPK_SIG_LABEL.len() + bundle.spk_pub.as_bytes().len());
msg.extend_from_slice(constants::SPK_SIG_LABEL);
msg.extend_from_slice(bundle.spk_pub.as_bytes());
identity::hybrid_verify(&bundle.ik_pub, &msg, &bundle.spk_sig)
.map_err(|_| Error::BundleVerificationFailed)?;
Ok(VerifiedBundle(bundle))
}
#[must_use = "session establishment result contains key material that must not be discarded"]
pub fn initiate_session(
alice_ik_pk: &IdentityPublicKey,
alice_ik_sk: &IdentitySecretKey,
bundle: &VerifiedBundle,
) -> Result<InitiatedSession> {
let (ek_pk, ek_sk) = xwing::keygen()?;
let bob_xwing_pk = xwing::PublicKey::from_bytes_unchecked(bundle.ik_pub.xwing_pk().to_vec());
let (ct_ik, mut ss_ik) = xwing::encapsulate(&bob_xwing_pk)?;
let (ct_spk, mut ss_spk) = xwing::encapsulate(&bundle.spk_pub)?;
let (ct_opk, mut ss_opk) = if let Some(ref opk_pub) = bundle.opk_pub {
let (ct, ss) = xwing::encapsulate(opk_pub)?;
(Some(ct), Some(ss))
} else {
(None, None)
};
let mut ikm = Zeroizing::new(Vec::with_capacity(96));
ikm.extend_from_slice(ss_ik.as_bytes());
ikm.extend_from_slice(ss_spk.as_bytes());
if let Some(ref ss) = ss_opk {
ikm.extend_from_slice(ss.as_bytes());
}
let info = build_kex_info(alice_ik_pk, &bundle.ik_pub, &ek_pk)?;
let mut session_key = Zeroizing::new([0u8; 64]);
hkdf::hkdf_sha3_256(&constants::HKDF_ZERO_SALT, &ikm, &info, &mut *session_key)?;
ikm.zeroize();
let mut root_key = Zeroizing::new([0u8; 32]);
let mut chain_key = Zeroizing::new([0u8; 32]);
root_key.copy_from_slice(&session_key[..32]);
chain_key.copy_from_slice(&session_key[32..64]);
session_key.zeroize();
ss_ik.0.zeroize();
ss_spk.0.zeroize();
if let Some(ref mut ss) = ss_opk {
ss.0.zeroize();
}
let sender_ik_fingerprint = alice_ik_pk.fingerprint_raw();
let session_init = SessionInit {
crypto_version: constants::CRYPTO_VERSION.to_string(),
sender_ik_fingerprint,
recipient_ik_fingerprint: bundle.ik_pub.fingerprint_raw(),
sender_ek: ek_pk.clone(),
ct_ik,
ct_spk,
spk_id: bundle.spk_id,
ct_opk,
opk_id: bundle.opk_id,
};
let si_encoded = encode_session_init(&session_init)?;
let mut sign_msg = Vec::with_capacity(constants::INITIATOR_SIG_LABEL.len() + si_encoded.len());
sign_msg.extend_from_slice(constants::INITIATOR_SIG_LABEL);
sign_msg.extend_from_slice(&si_encoded);
let sender_sig = identity::hybrid_sign(alice_ik_sk, &sign_msg)?;
let opk_used = session_init.ct_opk.is_some();
Ok(InitiatedSession {
session_init,
root_key,
initial_chain_key: chain_key,
ek_pk,
ek_sk,
sender_sig,
opk_used,
})
}
#[must_use = "session establishment result contains key material that must not be discarded"]
pub fn receive_session(
bob_ik_pk: &IdentityPublicKey,
bob_ik_sk: &IdentitySecretKey,
alice_ik_pk: &IdentityPublicKey,
si: &SessionInit,
sender_sig: &HybridSignature,
spk_sk: &xwing::SecretKey,
opk_sk: Option<&xwing::SecretKey>,
) -> Result<ReceivedSession> {
if si.crypto_version != constants::CRYPTO_VERSION {
return Err(Error::UnsupportedCryptoVersion);
}
use subtle::ConstantTimeEq;
let expected_fp = alice_ik_pk.fingerprint_raw();
if expected_fp.ct_ne(&si.sender_ik_fingerprint).into() {
return Err(Error::InvalidData);
}
let expected_recipient_fp = bob_ik_pk.fingerprint_raw();
if expected_recipient_fp
.ct_ne(&si.recipient_ik_fingerprint)
.into()
{
return Err(Error::InvalidData);
}
let si_encoded = encode_session_init(si)?;
let mut verify_msg =
Vec::with_capacity(constants::INITIATOR_SIG_LABEL.len() + si_encoded.len());
verify_msg.extend_from_slice(constants::INITIATOR_SIG_LABEL);
verify_msg.extend_from_slice(&si_encoded);
identity::hybrid_verify(alice_ik_pk, &verify_msg, sender_sig)?;
if si.ct_opk.is_none() && opk_sk.is_some() {
return Err(Error::InvalidData);
}
if si.ct_opk.is_some() && opk_sk.is_none() {
return Err(Error::InvalidData);
}
let bob_xwing_sk = xwing::SecretKey::from_bytes_unchecked(bob_ik_sk.xwing_sk().to_vec());
let mut ss_ik = xwing::decapsulate(&bob_xwing_sk, &si.ct_ik)?;
let mut ss_spk = xwing::decapsulate(spk_sk, &si.ct_spk)?;
let mut ss_opk = if let Some(ref ct_opk) = si.ct_opk {
Some(xwing::decapsulate(opk_sk.ok_or(Error::Internal)?, ct_opk)?)
} else {
None
};
let mut ikm = Zeroizing::new(Vec::with_capacity(96));
ikm.extend_from_slice(ss_ik.as_bytes());
ikm.extend_from_slice(ss_spk.as_bytes());
if let Some(ref ss) = ss_opk {
ikm.extend_from_slice(ss.as_bytes());
}
let info = build_kex_info(alice_ik_pk, bob_ik_pk, &si.sender_ek)?;
let mut session_key = Zeroizing::new([0u8; 64]);
hkdf::hkdf_sha3_256(&constants::HKDF_ZERO_SALT, &ikm, &info, &mut *session_key)?;
ikm.zeroize();
let mut root_key = Zeroizing::new([0u8; 32]);
let mut chain_key = Zeroizing::new([0u8; 32]);
root_key.copy_from_slice(&session_key[..32]);
chain_key.copy_from_slice(&session_key[32..64]);
session_key.zeroize();
ss_ik.0.zeroize();
ss_spk.0.zeroize();
if let Some(ref mut ss) = ss_opk {
ss.0.zeroize();
}
Ok(ReceivedSession {
root_key,
initial_chain_key: chain_key,
peer_ek: si.sender_ek.clone(),
})
}
fn build_kex_info(
alice_ik: &IdentityPublicKey,
bob_ik: &IdentityPublicKey,
ek: &xwing::PublicKey,
) -> Result<Vec<u8>> {
let cv = constants::CRYPTO_VERSION.as_bytes();
let mut info = Vec::with_capacity(
constants::KEX_HKDF_INFO_PFX.len()
+ 2
+ cv.len()
+ 2
+ constants::LO_PUBLIC_KEY_SIZE
+ 2
+ constants::LO_PUBLIC_KEY_SIZE
+ 2
+ constants::XWING_PUBLIC_KEY_SIZE,
);
info.extend_from_slice(constants::KEX_HKDF_INFO_PFX);
push_u16_prefixed(&mut info, cv)?;
push_u16_prefixed(&mut info, alice_ik.as_bytes())?;
push_u16_prefixed(&mut info, bob_ik.as_bytes())?;
push_u16_prefixed(&mut info, ek.as_bytes())?;
Ok(info)
}
fn require_u16_len(data: &[u8]) -> Result<()> {
if data.len() > u16::MAX as usize {
return Err(Error::InvalidLength {
expected: u16::MAX as usize,
got: data.len(),
});
}
Ok(())
}
fn push_u16_prefixed(buf: &mut Vec<u8>, data: &[u8]) -> Result<()> {
let len = u16::try_from(data.len()).map_err(|_| Error::Internal)?;
buf.extend_from_slice(&len.to_be_bytes());
buf.extend_from_slice(data);
Ok(())
}
pub fn encode_session_init(si: &SessionInit) -> Result<Vec<u8>> {
let mut buf = Vec::with_capacity(5120);
let cv = si.crypto_version.as_bytes();
require_u16_len(cv)?;
push_u16_prefixed(&mut buf, cv)?;
buf.extend_from_slice(&si.sender_ik_fingerprint);
buf.extend_from_slice(&si.recipient_ik_fingerprint);
buf.extend_from_slice(si.sender_ek.as_bytes());
push_u16_prefixed(&mut buf, si.ct_ik.as_bytes())?;
push_u16_prefixed(&mut buf, si.ct_spk.as_bytes())?;
buf.extend_from_slice(&si.spk_id.to_be_bytes());
if si.ct_opk.is_some() != si.opk_id.is_some() {
return Err(Error::InvalidData);
}
if let (Some(ct_opk), Some(opk_id)) = (&si.ct_opk, si.opk_id) {
buf.push(0x01);
push_u16_prefixed(&mut buf, ct_opk.as_bytes())?;
buf.extend_from_slice(&opk_id.to_be_bytes());
} else {
buf.push(0x00);
}
Ok(buf)
}
pub fn decode_session_init(data: &[u8]) -> Result<SessionInit> {
let mut cur = 0usize;
if data.len() < cur + 2 {
return Err(Error::InvalidData);
}
let ver_len = u16::from_be_bytes([data[cur], data[cur + 1]]) as usize;
cur += 2;
if ver_len > 64 {
return Err(Error::InvalidLength {
expected: 64,
got: ver_len,
});
}
if data.len() < cur + ver_len {
return Err(Error::InvalidData);
}
let crypto_version =
String::from_utf8(data[cur..cur + ver_len].to_vec()).map_err(|_| Error::InvalidData)?;
cur += ver_len;
if crypto_version != constants::CRYPTO_VERSION {
return Err(Error::UnsupportedCryptoVersion);
}
if data.len() < cur + 32 {
return Err(Error::InvalidData);
}
let sender_ik_fingerprint: [u8; 32] = data[cur..cur + 32]
.try_into()
.map_err(|_| Error::Internal)?;
cur += 32;
if data.len() < cur + 32 {
return Err(Error::InvalidData);
}
let recipient_ik_fingerprint: [u8; 32] = data[cur..cur + 32]
.try_into()
.map_err(|_| Error::Internal)?;
cur += 32;
if data.len() < cur + constants::XWING_PUBLIC_KEY_SIZE {
return Err(Error::InvalidData);
}
let sender_ek =
xwing::PublicKey::from_bytes(data[cur..cur + constants::XWING_PUBLIC_KEY_SIZE].to_vec())?;
cur += constants::XWING_PUBLIC_KEY_SIZE;
if data.len() < cur + 2 {
return Err(Error::InvalidData);
}
let ct_ik_len = u16::from_be_bytes([data[cur], data[cur + 1]]) as usize;
if ct_ik_len != constants::XWING_CIPHERTEXT_SIZE {
return Err(Error::InvalidData);
}
cur += 2;
if data.len() < cur + ct_ik_len {
return Err(Error::InvalidData);
}
let ct_ik = xwing::Ciphertext::from_bytes(data[cur..cur + ct_ik_len].to_vec())?;
cur += ct_ik_len;
if data.len() < cur + 2 {
return Err(Error::InvalidData);
}
let ct_spk_len = u16::from_be_bytes([data[cur], data[cur + 1]]) as usize;
if ct_spk_len != constants::XWING_CIPHERTEXT_SIZE {
return Err(Error::InvalidData);
}
cur += 2;
if data.len() < cur + ct_spk_len {
return Err(Error::InvalidData);
}
let ct_spk = xwing::Ciphertext::from_bytes(data[cur..cur + ct_spk_len].to_vec())?;
cur += ct_spk_len;
if data.len() < cur + 4 {
return Err(Error::InvalidData);
}
let spk_id = u32::from_be_bytes(data[cur..cur + 4].try_into().map_err(|_| Error::Internal)?);
cur += 4;
if data.len() < cur + 1 {
return Err(Error::InvalidData);
}
let has_opk_byte = data[cur];
cur += 1;
let (ct_opk, opk_id) = match has_opk_byte {
0x00 => (None, None),
0x01 => {
if data.len() < cur + 2 {
return Err(Error::InvalidData);
}
let ct_opk_len = u16::from_be_bytes([data[cur], data[cur + 1]]) as usize;
if ct_opk_len != constants::XWING_CIPHERTEXT_SIZE {
return Err(Error::InvalidData);
}
cur += 2;
if data.len() < cur + ct_opk_len {
return Err(Error::InvalidData);
}
let ct = xwing::Ciphertext::from_bytes(data[cur..cur + ct_opk_len].to_vec())?;
cur += ct_opk_len;
if data.len() < cur + 4 {
return Err(Error::InvalidData);
}
let id =
u32::from_be_bytes(data[cur..cur + 4].try_into().map_err(|_| Error::Internal)?);
cur += 4;
(Some(ct), Some(id))
}
_ => return Err(Error::InvalidData),
};
if cur != data.len() {
return Err(Error::InvalidData);
}
Ok(SessionInit {
crypto_version,
sender_ik_fingerprint,
recipient_ik_fingerprint,
sender_ek,
ct_ik,
ct_spk,
spk_id,
ct_opk,
opk_id,
})
}
pub fn build_first_message_aad(
sender_fingerprint: &[u8; 32],
recipient_fingerprint: &[u8; 32],
si: &SessionInit,
) -> Result<Vec<u8>> {
let si_bytes = encode_session_init(si)?;
build_first_message_aad_from_encoded(sender_fingerprint, recipient_fingerprint, &si_bytes)
}
pub fn build_first_message_aad_from_encoded(
sender_fingerprint: &[u8; 32],
recipient_fingerprint: &[u8; 32],
si_encoded: &[u8],
) -> Result<Vec<u8>> {
if si_encoded.is_empty() {
return Err(Error::InvalidData);
}
let mut aad = Vec::with_capacity(constants::DM_AAD.len() + 32 + 32 + si_encoded.len());
aad.extend_from_slice(constants::DM_AAD);
aad.extend_from_slice(sender_fingerprint);
aad.extend_from_slice(recipient_fingerprint);
aad.extend_from_slice(si_encoded);
Ok(aad)
}
pub fn sign_prekey(
ik_sk: &IdentitySecretKey,
spk_pub: &xwing::PublicKey,
) -> Result<HybridSignature> {
let mut msg = Vec::with_capacity(constants::SPK_SIG_LABEL.len() + spk_pub.as_bytes().len());
msg.extend_from_slice(constants::SPK_SIG_LABEL);
msg.extend_from_slice(spk_pub.as_bytes());
identity::hybrid_sign(ik_sk, &msg)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::Error;
use crate::identity::{GeneratedIdentity, generate_identity};
use crate::ratchet::RatchetState;
use hex_literal::hex;
fn make_bundle(
bob_pk: &IdentityPublicKey,
bob_sk: &IdentitySecretKey,
with_opk: bool,
) -> (PreKeyBundle, xwing::SecretKey, Option<xwing::SecretKey>) {
let (spk_pub, spk_sk) = xwing::keygen().expect("keygen");
let spk_sig = sign_prekey(bob_sk, &spk_pub).expect("sign_prekey");
let (opk_pub, opk_sk, opk_id) = if with_opk {
let (pub_key, sec_key) = xwing::keygen().expect("keygen (opk)");
(Some(pub_key), Some(sec_key), Some(1))
} else {
(None, None, None)
};
let bundle = PreKeyBundle {
ik_pub: IdentityPublicKey(bob_pk.as_bytes().to_vec()),
crypto_version: constants::CRYPTO_VERSION.to_string(),
spk_pub,
spk_id: 42,
spk_sig,
opk_pub,
opk_id,
};
(bundle, spk_sk, opk_sk)
}
#[test]
fn verify_bundle_valid() {
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, _, _) = make_bundle(&bob_pk, &bob_sk, false);
assert!(verify_bundle(bundle, &bob_pk).is_ok());
}
#[test]
fn verify_bundle_ik_mismatch() {
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: other_pk,
..
} = generate_identity().unwrap();
let (bundle, _, _) = make_bundle(&bob_pk, &bob_sk, false);
assert!(matches!(
verify_bundle(bundle, &other_pk),
Err(Error::BundleVerificationFailed)
));
}
#[test]
fn verify_bundle_bad_spk_sig() {
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (mut bundle, _, _) = make_bundle(&bob_pk, &bob_sk, false);
let mut bad_sig = bundle.spk_sig.as_bytes().to_vec();
bad_sig[0] ^= 0xFF;
bundle.spk_sig = HybridSignature::from_bytes(bad_sig).unwrap();
assert!(matches!(
verify_bundle(bundle, &bob_pk),
Err(Error::BundleVerificationFailed)
));
}
#[test]
fn verify_bundle_wrong_crypto_version() {
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (mut bundle, _, _) = make_bundle(&bob_pk, &bob_sk, false);
bundle.crypto_version = "lo-crypto-v999".to_string();
assert!(matches!(
verify_bundle(bundle, &bob_pk),
Err(Error::BundleVerificationFailed)
));
}
#[test]
fn verify_bundle_partial_opk() {
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (mut bundle, _, _) = make_bundle(&bob_pk, &bob_sk, true);
bundle.opk_id = None;
assert!(matches!(
verify_bundle(bundle, &bob_pk),
Err(Error::InvalidData)
));
}
#[test]
fn session_agreement_with_opk() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, opk_sk) = make_bundle(&bob_pk, &bob_sk, true);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let mut initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let mut received = receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&initiated.session_init,
&initiated.sender_sig,
&spk_sk,
opk_sk.as_ref(),
)
.unwrap();
assert_eq!(*initiated.take_root_key(), *received.take_root_key());
assert_eq!(
*initiated.take_initial_chain_key(),
*received.take_initial_chain_key()
);
}
#[test]
fn session_agreement_without_opk() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let mut initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let mut received = receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&initiated.session_init,
&initiated.sender_sig,
&spk_sk,
None,
)
.unwrap();
assert_eq!(*initiated.take_root_key(), *received.take_root_key());
assert_eq!(
*initiated.take_initial_chain_key(),
*received.take_initial_chain_key()
);
}
#[test]
fn kex_derivation_kat() {
let alice_ik =
IdentityPublicKey::from_bytes(vec![0x01u8; constants::LO_PUBLIC_KEY_SIZE]).unwrap();
let bob_ik =
IdentityPublicKey::from_bytes(vec![0x02u8; constants::LO_PUBLIC_KEY_SIZE]).unwrap();
let ek =
xwing::PublicKey::from_bytes(vec![0x03u8; constants::XWING_PUBLIC_KEY_SIZE]).unwrap();
let mut ikm = [0u8; 64];
ikm[..32].fill(0xAA); ikm[32..].fill(0xBB);
let info = build_kex_info(&alice_ik, &bob_ik, &ek).unwrap();
let mut out = [0u8; 64];
hkdf::hkdf_sha3_256(&constants::HKDF_ZERO_SALT, &ikm, &info, &mut out).unwrap();
assert_eq!(
out[..32],
hex!("b90ead32626ae4b0be4864ba8ad2fad4b570c4021d4f20b83c79207dd2a91074"),
"root_key derivation mismatch — info construction or HKDF parameters changed"
);
assert_eq!(
out[32..],
hex!("965624c6263ca5663765bc97f179891ea3e0ee73010b037de182abb3b123d0bf"),
"chain_key derivation mismatch — info construction or HKDF parameters changed"
);
}
#[test]
fn receive_session_fingerprint_mismatch() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: fake_alice_pk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&fake_alice_pk,
&initiated.session_init,
&initiated.sender_sig,
&spk_sk,
None
),
Err(Error::InvalidData)
));
}
#[test]
fn receive_session_wrong_recipient_fingerprint() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: carol_pk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let mut initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
initiated.session_init.recipient_ik_fingerprint = carol_pk.fingerprint_raw();
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&initiated.session_init,
&initiated.sender_sig,
&spk_sk,
None
),
Err(Error::InvalidData)
));
}
#[test]
fn receive_session_wrong_crypto_version() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let mut initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
initiated.session_init.crypto_version = "lo-crypto-v999".to_string();
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&initiated.session_init,
&initiated.sender_sig,
&spk_sk,
None
),
Err(Error::UnsupportedCryptoVersion)
));
}
#[test]
fn receive_session_opk_sk_without_ct_opk() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let (_, extra_sk) = xwing::keygen().unwrap();
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&initiated.session_init,
&initiated.sender_sig,
&spk_sk,
Some(&extra_sk)
),
Err(Error::InvalidData)
));
}
#[test]
fn receive_session_ct_opk_without_opk_sk() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, true);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&initiated.session_init,
&initiated.sender_sig,
&spk_sk,
None
),
Err(Error::InvalidData)
));
}
#[test]
fn receive_session_partial_opk_fields() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, opk_sk) = make_bundle(&bob_pk, &bob_sk, true);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let mut initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
initiated.session_init.opk_id = None;
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&initiated.session_init,
&initiated.sender_sig,
&spk_sk,
opk_sk.as_ref()
),
Err(Error::InvalidData)
));
}
#[test]
fn receive_session_bad_sig() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let mut bad_sig_bytes = initiated.sender_sig.as_bytes().to_vec();
bad_sig_bytes[0] ^= 0xFF;
let bad_sig = HybridSignature::from_bytes(bad_sig_bytes).unwrap();
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&initiated.session_init,
&bad_sig,
&spk_sk,
None
),
Err(Error::VerificationFailed)
));
}
#[test]
fn receive_session_tampered_ct_ik() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let sender_sig = initiated.sender_sig.clone();
let si = &initiated.session_init;
let mut ct_ik_bytes = si.ct_ik.as_bytes().to_vec();
ct_ik_bytes[0] ^= 0xFF;
let tampered = SessionInit {
crypto_version: si.crypto_version.clone(),
sender_ik_fingerprint: si.sender_ik_fingerprint,
recipient_ik_fingerprint: si.recipient_ik_fingerprint,
sender_ek: si.sender_ek.clone(),
ct_ik: xwing::Ciphertext::from_bytes(ct_ik_bytes).unwrap(),
ct_spk: si.ct_spk.clone(),
spk_id: si.spk_id,
ct_opk: si.ct_opk.clone(),
opk_id: si.opk_id,
};
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&tampered,
&sender_sig,
&spk_sk,
None
),
Err(Error::VerificationFailed)
));
}
#[test]
fn receive_session_tampered_sender_ek() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let sender_sig = initiated.sender_sig.clone();
let si = &initiated.session_init;
let mut ek_bytes = si.sender_ek.as_bytes().to_vec();
ek_bytes[0] ^= 0xFF;
let tampered = SessionInit {
crypto_version: si.crypto_version.clone(),
sender_ik_fingerprint: si.sender_ik_fingerprint,
recipient_ik_fingerprint: si.recipient_ik_fingerprint,
sender_ek: xwing::PublicKey::from_bytes(ek_bytes).unwrap(),
ct_ik: si.ct_ik.clone(),
ct_spk: si.ct_spk.clone(),
spk_id: si.spk_id,
ct_opk: si.ct_opk.clone(),
opk_id: si.opk_id,
};
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&tampered,
&sender_sig,
&spk_sk,
None
),
Err(Error::VerificationFailed)
));
}
#[test]
fn receive_session_tampered_ct_spk() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let sender_sig = initiated.sender_sig.clone();
let si = &initiated.session_init;
let mut ct_spk_bytes = si.ct_spk.as_bytes().to_vec();
ct_spk_bytes[0] ^= 0xFF;
let tampered = SessionInit {
crypto_version: si.crypto_version.clone(),
sender_ik_fingerprint: si.sender_ik_fingerprint,
recipient_ik_fingerprint: si.recipient_ik_fingerprint,
sender_ek: si.sender_ek.clone(),
ct_ik: si.ct_ik.clone(),
ct_spk: xwing::Ciphertext::from_bytes(ct_spk_bytes).unwrap(),
spk_id: si.spk_id,
ct_opk: si.ct_opk.clone(),
opk_id: si.opk_id,
};
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&tampered,
&sender_sig,
&spk_sk,
None
),
Err(Error::VerificationFailed)
));
}
#[test]
fn receive_session_tampered_sender_ik_fingerprint() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let sender_sig = initiated.sender_sig.clone();
let si = &initiated.session_init;
let mut bad_fp = si.sender_ik_fingerprint;
bad_fp[0] ^= 0xFF;
let tampered = SessionInit {
crypto_version: si.crypto_version.clone(),
sender_ik_fingerprint: bad_fp,
recipient_ik_fingerprint: si.recipient_ik_fingerprint,
sender_ek: si.sender_ek.clone(),
ct_ik: si.ct_ik.clone(),
ct_spk: si.ct_spk.clone(),
spk_id: si.spk_id,
ct_opk: si.ct_opk.clone(),
opk_id: si.opk_id,
};
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&tampered,
&sender_sig,
&spk_sk,
None
),
Err(Error::InvalidData)
));
}
#[test]
fn receive_session_tampered_crypto_version() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let sender_sig = initiated.sender_sig.clone();
let si = &initiated.session_init;
let tampered = SessionInit {
crypto_version: "lo-crypto-v99".to_string(),
sender_ik_fingerprint: si.sender_ik_fingerprint,
recipient_ik_fingerprint: si.recipient_ik_fingerprint,
sender_ek: si.sender_ek.clone(),
ct_ik: si.ct_ik.clone(),
ct_spk: si.ct_spk.clone(),
spk_id: si.spk_id,
ct_opk: si.ct_opk.clone(),
opk_id: si.opk_id,
};
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&tampered,
&sender_sig,
&spk_sk,
None
),
Err(Error::UnsupportedCryptoVersion)
));
}
#[test]
fn receive_session_wrong_signer() {
let GeneratedIdentity {
public_key: alice_pk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: eve_pk,
secret_key: eve_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb_eve = verify_bundle(bundle, &bob_pk).unwrap();
let eve_initiated = initiate_session(&eve_pk, &eve_sk, &vb_eve).unwrap();
assert!(matches!(
receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&eve_initiated.session_init,
&eve_initiated.sender_sig,
&spk_sk,
None
),
Err(Error::InvalidData)
));
}
#[test]
fn zero_knowledge_impersonation_rejected() {
let GeneratedIdentity {
public_key: alice_pk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
secret_key: eve_sk, ..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let (ek_pk, _ek_sk) = xwing::keygen().unwrap();
let bob_xwing_pk = xwing::PublicKey::from_bytes_unchecked(bob_pk.xwing_pk().to_vec());
let (ct_ik, _ss) = xwing::encapsulate(&bob_xwing_pk).unwrap();
let (ct_spk, _ss) = xwing::encapsulate(&vb.spk_pub).unwrap();
let si = SessionInit {
crypto_version: constants::CRYPTO_VERSION.to_string(),
sender_ik_fingerprint: alice_pk.fingerprint_raw(),
recipient_ik_fingerprint: bob_pk.fingerprint_raw(),
sender_ek: ek_pk,
ct_ik,
ct_spk,
spk_id: vb.spk_id,
ct_opk: None,
opk_id: None,
};
let si_encoded = encode_session_init(&si).unwrap();
let mut sign_msg =
Vec::with_capacity(constants::INITIATOR_SIG_LABEL.len() + si_encoded.len());
sign_msg.extend_from_slice(constants::INITIATOR_SIG_LABEL);
sign_msg.extend_from_slice(&si_encoded);
let eve_sig = identity::hybrid_sign(&eve_sk, &sign_msg).unwrap();
assert!(matches!(
receive_session(&bob_pk, &bob_sk, &alice_pk, &si, &eve_sig, &spk_sk, None),
Err(Error::VerificationFailed)
));
}
#[test]
fn encode_session_init_deterministic() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, _, _) = make_bundle(&bob_pk, &bob_sk, true);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let enc1 = encode_session_init(&initiated.session_init).unwrap();
let enc2 = encode_session_init(&initiated.session_init).unwrap();
assert_eq!(enc1, enc2);
}
#[test]
fn encode_session_init_rejects_oversized_crypto_version() {
let si = SessionInit {
crypto_version: "x".repeat(u16::MAX as usize + 1),
sender_ik_fingerprint: [0u8; 32],
recipient_ik_fingerprint: [0u8; 32],
sender_ek: xwing::PublicKey::from_bytes(vec![0u8; constants::XWING_PUBLIC_KEY_SIZE])
.unwrap(),
ct_ik: xwing::Ciphertext::from_bytes(vec![0u8; constants::XWING_CIPHERTEXT_SIZE])
.unwrap(),
ct_spk: xwing::Ciphertext::from_bytes(vec![0u8; constants::XWING_CIPHERTEXT_SIZE])
.unwrap(),
spk_id: 0,
ct_opk: None,
opk_id: None,
};
assert!(matches!(
encode_session_init(&si),
Err(Error::InvalidLength { .. })
));
}
#[test]
fn build_first_message_aad_structure() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, _, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let fp_a = alice_pk.fingerprint_raw();
let fp_b = bob_pk.fingerprint_raw();
let aad = build_first_message_aad(&fp_a, &fp_b, &initiated.session_init).unwrap();
let si_bytes = encode_session_init(&initiated.session_init).unwrap();
let mut expected = Vec::new();
expected.extend_from_slice(constants::DM_AAD);
expected.extend_from_slice(&fp_a);
expected.extend_from_slice(&fp_b);
expected.extend_from_slice(&si_bytes);
assert_eq!(aad, expected);
}
#[test]
#[allow(clippy::cast_possible_truncation)]
fn kex_info_field_structure() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, _, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let info = build_kex_info(&alice_pk, &bob_pk, &initiated.ek_pk).unwrap();
let cv = b"lo-crypto-v1";
let mut expected = Vec::new();
expected.extend_from_slice(b"lo-kex-v1");
expected.extend_from_slice(&(cv.len() as u16).to_be_bytes());
expected.extend_from_slice(cv);
expected.extend_from_slice(&(alice_pk.as_bytes().len() as u16).to_be_bytes());
expected.extend_from_slice(alice_pk.as_bytes());
expected.extend_from_slice(&(bob_pk.as_bytes().len() as u16).to_be_bytes());
expected.extend_from_slice(bob_pk.as_bytes());
expected.extend_from_slice(&(initiated.ek_pk.as_bytes().len() as u16).to_be_bytes());
expected.extend_from_slice(initiated.ek_pk.as_bytes());
assert_eq!(info, expected);
}
#[test]
fn sign_prekey_verify() {
let GeneratedIdentity {
public_key: pk,
secret_key: sk,
..
} = generate_identity().unwrap();
let (spk_pub, _) = xwing::keygen().unwrap();
let sig = sign_prekey(&sk, &spk_pub).unwrap();
let mut msg = Vec::with_capacity(constants::SPK_SIG_LABEL.len() + spk_pub.as_bytes().len());
msg.extend_from_slice(constants::SPK_SIG_LABEL);
msg.extend_from_slice(spk_pub.as_bytes());
assert!(identity::hybrid_verify(&pk, &msg, &sig).is_ok());
}
#[test]
fn first_message_round_trip() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let mut initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let fp_a = alice_pk.fingerprint_raw();
let fp_b = bob_pk.fingerprint_raw();
let aad = build_first_message_aad(&fp_a, &fp_b, &initiated.session_init).unwrap();
let (payload, ck_alice) = RatchetState::encrypt_first_message(
initiated.take_initial_chain_key(),
b"first message",
&aad,
)
.unwrap();
let mut received = receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&initiated.session_init,
&initiated.sender_sig,
&spk_sk,
None,
)
.unwrap();
let (pt, ck_bob) =
RatchetState::decrypt_first_message(received.take_initial_chain_key(), &payload, &aad)
.unwrap();
assert_eq!(&*pt, b"first message");
assert_eq!(*ck_alice, *ck_bob);
}
fn make_zero_session_init(with_opk: bool) -> SessionInit {
SessionInit {
crypto_version: constants::CRYPTO_VERSION.to_string(),
sender_ik_fingerprint: [0x01u8; 32],
recipient_ik_fingerprint: [0x02u8; 32],
sender_ek: xwing::PublicKey::from_bytes(vec![0x03u8; constants::XWING_PUBLIC_KEY_SIZE])
.unwrap(),
ct_ik: xwing::Ciphertext::from_bytes(vec![0x04u8; constants::XWING_CIPHERTEXT_SIZE])
.unwrap(),
ct_spk: xwing::Ciphertext::from_bytes(vec![0x05u8; constants::XWING_CIPHERTEXT_SIZE])
.unwrap(),
spk_id: 0xDEAD_BEEF,
ct_opk: if with_opk {
Some(
xwing::Ciphertext::from_bytes(vec![0x06u8; constants::XWING_CIPHERTEXT_SIZE])
.unwrap(),
)
} else {
None
},
opk_id: if with_opk { Some(0x1234_5678) } else { None },
}
}
#[test]
fn decode_session_init_roundtrip_without_opk() {
let si = make_zero_session_init(false);
let encoded = encode_session_init(&si).unwrap();
let decoded = decode_session_init(&encoded).unwrap();
assert_eq!(decoded.crypto_version, si.crypto_version);
assert_eq!(decoded.sender_ik_fingerprint, si.sender_ik_fingerprint);
assert_eq!(
decoded.recipient_ik_fingerprint,
si.recipient_ik_fingerprint
);
assert_eq!(decoded.sender_ek.as_bytes(), si.sender_ek.as_bytes());
assert_eq!(decoded.ct_ik.as_bytes(), si.ct_ik.as_bytes());
assert_eq!(decoded.ct_spk.as_bytes(), si.ct_spk.as_bytes());
assert_eq!(decoded.spk_id, si.spk_id);
assert!(decoded.ct_opk.is_none());
assert!(decoded.opk_id.is_none());
let re_encoded = encode_session_init(&decoded).unwrap();
assert_eq!(encoded, re_encoded);
}
#[test]
fn decode_session_init_roundtrip_with_opk() {
let si = make_zero_session_init(true);
let encoded = encode_session_init(&si).unwrap();
let decoded = decode_session_init(&encoded).unwrap();
assert_eq!(decoded.crypto_version, si.crypto_version);
assert_eq!(decoded.sender_ik_fingerprint, si.sender_ik_fingerprint);
assert_eq!(
decoded.recipient_ik_fingerprint,
si.recipient_ik_fingerprint
);
assert_eq!(decoded.sender_ek.as_bytes(), si.sender_ek.as_bytes());
assert_eq!(decoded.ct_ik.as_bytes(), si.ct_ik.as_bytes());
assert_eq!(decoded.ct_spk.as_bytes(), si.ct_spk.as_bytes());
assert_eq!(decoded.spk_id, si.spk_id);
assert_eq!(
decoded.ct_opk.as_ref().unwrap().as_bytes(),
si.ct_opk.as_ref().unwrap().as_bytes()
);
assert_eq!(decoded.opk_id, si.opk_id);
let re_encoded = encode_session_init(&decoded).unwrap();
assert_eq!(encoded, re_encoded);
}
#[test]
fn decode_session_init_empty_returns_invalid_data() {
assert!(matches!(decode_session_init(&[]), Err(Error::InvalidData)));
}
#[test]
fn decode_session_init_truncated_mid_field_returns_invalid_data() {
let encoded = encode_session_init(&make_zero_session_init(false)).unwrap();
for cut in [1, 10, 100, 1500, encoded.len() - 1] {
assert!(
matches!(
decode_session_init(&encoded[..cut]),
Err(Error::InvalidData)
),
"expected InvalidData when truncated to {cut} bytes"
);
}
}
#[test]
fn decode_session_init_trailing_byte_returns_invalid_data() {
let encoded = encode_session_init(&make_zero_session_init(false)).unwrap();
let mut with_trail = encoded;
with_trail.push(0x00);
assert!(matches!(
decode_session_init(&with_trail),
Err(Error::InvalidData)
));
}
#[test]
fn decode_session_init_invalid_utf8_version_returns_invalid_data() {
let mut buf: Vec<u8> = Vec::new();
buf.extend_from_slice(&1u16.to_be_bytes()); buf.push(0xFF); buf.extend_from_slice(&[0x01u8; 32]); buf.extend_from_slice(&[0x02u8; 32]); buf.extend_from_slice(&[0x03u8; constants::XWING_PUBLIC_KEY_SIZE]); buf.extend_from_slice(
&u16::try_from(constants::XWING_CIPHERTEXT_SIZE)
.unwrap()
.to_be_bytes(),
);
buf.extend_from_slice(&[0x04u8; constants::XWING_CIPHERTEXT_SIZE]);
buf.extend_from_slice(
&u16::try_from(constants::XWING_CIPHERTEXT_SIZE)
.unwrap()
.to_be_bytes(),
);
buf.extend_from_slice(&[0x05u8; constants::XWING_CIPHERTEXT_SIZE]);
buf.extend_from_slice(&0u32.to_be_bytes()); buf.push(0x00); assert!(matches!(decode_session_init(&buf), Err(Error::InvalidData)));
}
#[test]
fn decode_session_init_bad_has_opk_byte_returns_invalid_data() {
let encoded = encode_session_init(&make_zero_session_init(false)).unwrap();
let mut bad = encoded;
*bad.last_mut().unwrap() = 0x02;
assert!(matches!(decode_session_init(&bad), Err(Error::InvalidData)));
}
#[test]
fn first_message_wrong_aad() {
let GeneratedIdentity {
public_key: alice_pk,
secret_key: alice_sk,
..
} = generate_identity().unwrap();
let GeneratedIdentity {
public_key: bob_pk,
secret_key: bob_sk,
..
} = generate_identity().unwrap();
let (bundle, spk_sk, _) = make_bundle(&bob_pk, &bob_sk, false);
let vb = verify_bundle(bundle, &bob_pk).unwrap();
let mut initiated = initiate_session(&alice_pk, &alice_sk, &vb).unwrap();
let fp_a = alice_pk.fingerprint_raw();
let fp_b = bob_pk.fingerprint_raw();
let aad = build_first_message_aad(&fp_a, &fp_b, &initiated.session_init).unwrap();
let (payload, _) =
RatchetState::encrypt_first_message(initiated.take_initial_chain_key(), b"first", &aad)
.unwrap();
let mut received = receive_session(
&bob_pk,
&bob_sk,
&alice_pk,
&initiated.session_init,
&initiated.sender_sig,
&spk_sk,
None,
)
.unwrap();
assert!(matches!(
RatchetState::decrypt_first_message(
received.take_initial_chain_key(),
&payload,
b"wrong aad"
),
Err(Error::AeadFailed)
));
}
}