use sha3::{Digest, Sha3_256};
use super::{mlkem, random, x25519};
use crate::constants;
use crate::error::{Error, Result};
use zeroize::{Zeroize, ZeroizeOnDrop, Zeroizing};
#[derive(Clone, Eq)]
pub struct PublicKey(pub(crate) Vec<u8>);
impl PartialEq for PublicKey {
fn eq(&self, other: &Self) -> bool {
use subtle::ConstantTimeEq;
if self.0.len() != other.0.len() {
return false;
}
self.0.ct_eq(&other.0).into()
}
}
#[derive(Zeroize, ZeroizeOnDrop)]
pub struct SecretKey(pub(crate) Vec<u8>);
#[derive(Clone, PartialEq, Eq)]
pub struct Ciphertext(pub(crate) Vec<u8>);
#[derive(Zeroize, ZeroizeOnDrop)]
pub struct SharedSecret(pub(crate) [u8; 32]);
impl PublicKey {
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self> {
if bytes.len() != constants::XWING_PUBLIC_KEY_SIZE {
return Err(Error::InvalidLength {
expected: constants::XWING_PUBLIC_KEY_SIZE,
got: bytes.len(),
});
}
Ok(Self(bytes))
}
pub(crate) fn from_bytes_unchecked(bytes: Vec<u8>) -> Self {
assert_eq!(
bytes.len(),
constants::XWING_PUBLIC_KEY_SIZE,
"from_bytes_unchecked called with wrong size"
);
Self(bytes)
}
pub fn x25519_pk(&self) -> &[u8] {
&self.0[..32]
}
pub fn mlkem_pk(&self) -> &[u8] {
&self.0[32..]
}
}
impl SecretKey {
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self> {
let mut bytes = Zeroizing::new(bytes);
let expected = 32 + mlkem::sk_len();
if bytes.len() != expected {
return Err(Error::InvalidLength {
expected,
got: bytes.len(),
});
}
Ok(Self(std::mem::take(&mut *bytes)))
}
pub(crate) fn from_bytes_unchecked(bytes: Vec<u8>) -> Self {
assert_eq!(
bytes.len(),
32 + mlkem::sk_len(),
"from_bytes_unchecked called with wrong size"
);
Self(bytes)
}
pub(crate) fn x25519_sk(&self) -> &[u8] {
&self.0[..32]
}
pub(crate) fn mlkem_sk(&self) -> &[u8] {
&self.0[32..]
}
}
impl Ciphertext {
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self> {
if bytes.len() != constants::XWING_CIPHERTEXT_SIZE {
return Err(Error::InvalidLength {
expected: constants::XWING_CIPHERTEXT_SIZE,
got: bytes.len(),
});
}
Ok(Self(bytes))
}
}
impl SharedSecret {
pub fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
}
const XWING_LABEL: &[u8; 6] = b"\\.//^\\";
const _: () = assert!(
XWING_LABEL[0] == 0x5c
&& XWING_LABEL[1] == 0x2e
&& XWING_LABEL[2] == 0x2f
&& XWING_LABEL[3] == 0x2f
&& XWING_LABEL[4] == 0x5e
&& XWING_LABEL[5] == 0x5c,
"XWING_LABEL must be 5c 2e 2f 2f 5e 5c per draft-09 §5.3"
);
fn combiner(ss_m: &[u8], ss_x: &[u8], ct_x: &[u8], pk_x: &[u8]) -> [u8; 32] {
let mut hasher = Sha3_256::new();
hasher.update(ss_m);
hasher.update(ss_x);
hasher.update(ct_x);
hasher.update(pk_x);
hasher.update(XWING_LABEL);
hasher.finalize().into()
}
#[must_use = "key material must not be discarded"]
pub fn keygen() -> Result<(PublicKey, SecretKey)> {
let (x_pk, x_sk) = x25519::keygen();
let (m_pk, m_sk) = mlkem::keygen()?;
assert_eq!(
32 + mlkem::pk_len(),
constants::XWING_PUBLIC_KEY_SIZE,
"X-Wing public key size mismatch — update XWING_PUBLIC_KEY_SIZE"
);
assert_eq!(
32 + mlkem::sk_len(),
constants::XWING_SECRET_KEY_SIZE,
"X-Wing secret key size mismatch — update XWING_SECRET_KEY_SIZE"
);
assert_eq!(
32 + mlkem::ct_len(),
constants::XWING_CIPHERTEXT_SIZE,
"X-Wing ciphertext size mismatch — update XWING_CIPHERTEXT_SIZE"
);
let mut pk = Vec::with_capacity(32 + m_pk.as_bytes().len());
pk.extend_from_slice(x_pk.as_bytes());
pk.extend_from_slice(m_pk.as_bytes());
let mut sk = Zeroizing::new(Vec::with_capacity(32 + m_sk.as_bytes().len()));
sk.extend_from_slice(x_sk.as_bytes());
sk.extend_from_slice(m_sk.as_bytes());
Ok((PublicKey(pk), SecretKey(std::mem::take(&mut *sk))))
}
#[must_use = "shared secret and ciphertext must not be discarded"]
pub fn encapsulate(pk: &PublicKey) -> Result<(Ciphertext, SharedSecret)> {
if pk.0.len() != constants::XWING_PUBLIC_KEY_SIZE {
return Err(Error::InvalidLength {
expected: constants::XWING_PUBLIC_KEY_SIZE,
got: pk.0.len(),
});
}
let mut ek_sk_bytes = [0u8; 32];
random::random_bytes(&mut ek_sk_bytes);
let ek_sk = x25519::SecretKey::from_bytes(ek_sk_bytes);
ek_sk_bytes.zeroize();
let ek_pk = x25519::public_from_secret(&ek_sk);
let recipient_x_pk = x25519::PublicKey::from_bytes({
let mut buf = [0u8; 32];
buf.copy_from_slice(pk.x25519_pk());
buf
});
let mut raw_ss_x = x25519::dh(&ek_sk, &recipient_x_pk).unwrap_or([0u8; 32]);
let ss_x = Zeroizing::new(raw_ss_x);
raw_ss_x.zeroize();
let mlkem_pk = mlkem::PublicKey::from_bytes_unchecked(pk.mlkem_pk().to_vec());
let (mlkem_ct, mlkem_ss) = mlkem::encapsulate(&mlkem_pk)?;
let mut ct = Vec::with_capacity(32 + mlkem_ct.as_bytes().len());
ct.extend_from_slice(ek_pk.as_bytes());
ct.extend_from_slice(mlkem_ct.as_bytes());
let mut ss = combiner(
mlkem_ss.as_bytes(),
&*ss_x,
ek_pk.as_bytes(),
pk.x25519_pk(),
);
let shared = SharedSecret(ss);
ss.zeroize();
Ok((Ciphertext(ct), shared))
}
#[must_use = "shared secret must not be discarded"]
pub fn decapsulate(sk: &SecretKey, ct: &Ciphertext) -> Result<SharedSecret> {
debug_assert_eq!(
sk.0.len(),
constants::XWING_SECRET_KEY_SIZE,
"SecretKey constructed with wrong length"
);
if ct.0.len() != constants::XWING_CIPHERTEXT_SIZE {
return Err(Error::InvalidLength {
expected: constants::XWING_CIPHERTEXT_SIZE,
got: ct.0.len(),
});
}
let ct_x = &ct.0[..32]; let ct_m = &ct.0[32..];
let mut x_sk_bytes = [0u8; 32];
x_sk_bytes.copy_from_slice(sk.x25519_sk());
let our_x_sk = x25519::SecretKey::from_bytes(x_sk_bytes);
x_sk_bytes.zeroize();
let peer_ek = x25519::PublicKey::from_bytes({
let mut buf = [0u8; 32];
buf.copy_from_slice(ct_x);
buf
});
let mut raw_ss_x = x25519::dh(&our_x_sk, &peer_ek).unwrap_or([0u8; 32]);
let ss_x = Zeroizing::new(raw_ss_x);
raw_ss_x.zeroize();
let mlkem_sk = mlkem::SecretKey::from_bytes_unchecked(sk.mlkem_sk().to_vec());
let mlkem_ct = mlkem::Ciphertext::from_bytes_unchecked(ct_m.to_vec());
let mlkem_ss = mlkem::decapsulate(&mlkem_sk, &mlkem_ct)?;
let pk_x = x25519::public_from_secret(&our_x_sk);
let mut ss = combiner(mlkem_ss.as_bytes(), &*ss_x, ct_x, pk_x.as_bytes());
let shared = SharedSecret(ss);
ss.zeroize();
Ok(shared)
}
#[cfg(test)]
mod tests {
use super::*;
use hex_literal::hex;
use sha3::{Digest, Sha3_256};
#[test]
fn keygen_sizes() {
let (pk, sk) = keygen().unwrap();
assert_eq!(pk.as_bytes().len(), 1216);
assert_eq!(sk.as_bytes().len(), 2432);
assert!(pk.x25519_pk().iter().any(|&b| b != 0));
assert_eq!(pk.mlkem_pk().len(), 1184);
}
#[test]
fn round_trip() {
let (pk, sk) = keygen().unwrap();
let (ct, ss_enc) = encapsulate(&pk).unwrap();
let ss_dec = decapsulate(&sk, &ct).unwrap();
assert_eq!(ss_enc.as_bytes(), ss_dec.as_bytes());
}
#[test]
fn combiner_kat() {
let ss_m = [0x01u8; 32];
let ss_x = [0x02u8; 32];
let ct_x = [0x03u8; 32];
let pk_x = [0x04u8; 32];
let result = combiner(&ss_m, &ss_x, &ct_x, &pk_x);
let expected: [u8; 32] = Sha3_256::new()
.chain_update(ss_m)
.chain_update(ss_x)
.chain_update(ct_x)
.chain_update(pk_x)
.chain_update(XWING_LABEL)
.finalize()
.into();
assert_eq!(result, expected);
}
#[test]
fn label_hex_value() {
assert_eq!(XWING_LABEL, &[0x5c, 0x2e, 0x2f, 0x2f, 0x5e, 0x5c]);
}
#[test]
fn label_is_six_bytes() {
assert_eq!(XWING_LABEL.len(), 6);
}
#[test]
fn combiner_order_matters_ss() {
let a = [0x01u8; 32];
let b = [0x02u8; 32];
let c = [0x03u8; 32];
let d = [0x04u8; 32];
let out1 = combiner(&a, &b, &c, &d);
let out2 = combiner(&b, &a, &c, &d);
assert_ne!(
out1, out2,
"swapping ss_M and ss_X must produce different output"
);
}
#[test]
fn combiner_order_matters_ct_pk() {
let a = [0x01u8; 32];
let b = [0x02u8; 32];
let c = [0x03u8; 32];
let d = [0x04u8; 32];
let out1 = combiner(&a, &b, &c, &d);
let out2 = combiner(&a, &b, &d, &c);
assert_ne!(
out1, out2,
"swapping ct_X and pk_X must produce different output"
);
}
#[test]
fn combiner_label_is_last() {
let ss_m = [0x01u8; 32];
let ss_x = [0x02u8; 32];
let ct_x = [0x03u8; 32];
let pk_x = [0x04u8; 32];
let actual = combiner(&ss_m, &ss_x, &ct_x, &pk_x);
let label_first: [u8; 32] = Sha3_256::new()
.chain_update(XWING_LABEL)
.chain_update(ss_m)
.chain_update(ss_x)
.chain_update(ct_x)
.chain_update(pk_x)
.finalize()
.into();
assert_ne!(
actual, label_first,
"label-first ordering must differ from combiner"
);
let label_last: [u8; 32] = Sha3_256::new()
.chain_update(ss_m)
.chain_update(ss_x)
.chain_update(ct_x)
.chain_update(pk_x)
.chain_update(XWING_LABEL)
.finalize()
.into();
assert_eq!(
actual, label_last,
"label-last ordering must match combiner"
);
}
#[test]
fn encapsulate_wrong_pk_size() {
assert!(matches!(
PublicKey::from_bytes(vec![0u8; 100]),
Err(crate::error::Error::InvalidLength {
expected: 1216,
got: 100
})
));
}
#[test]
fn decapsulate_wrong_ct_size() {
assert!(matches!(
Ciphertext::from_bytes(vec![0u8; 100]),
Err(crate::error::Error::InvalidLength {
expected: 1120,
got: 100
})
));
}
#[test]
fn sk_from_bytes_wrong_size() {
assert!(matches!(
SecretKey::from_bytes(vec![0u8; 100]),
Err(crate::error::Error::InvalidLength {
expected: 2432,
got: 100
})
));
}
#[test]
fn shared_secret_is_32_bytes() {
let (pk, _sk) = keygen().unwrap();
let (_ct, ss) = encapsulate(&pk).unwrap();
assert!(ss.as_bytes().iter().any(|&b| b != 0));
}
#[test]
fn independent_encapsulations_differ() {
let (pk1, _sk1) = keygen().unwrap();
let (pk2, _sk2) = keygen().unwrap();
let (_ct1, ss1) = encapsulate(&pk1).unwrap();
let (_ct2, ss2) = encapsulate(&pk2).unwrap();
assert_ne!(
ss1.as_bytes(),
ss2.as_bytes(),
"independent encapsulations must produce different shared secrets"
);
}
#[test]
fn low_order_x25519_does_not_error() {
let (pk, sk) = keygen().unwrap();
let mlkem_part = pk.mlkem_pk().to_vec();
let small_order_points: Vec<[u8; 32]> = vec![
[0u8; 32], {
let mut p = [0u8; 32]; p[0] = 1;
p
},
{
let mut p = [0xffu8; 32]; p[31] = 0x7f;
p[0] = 0xec;
p
},
{
let mut p = [0xffu8; 32]; p[31] = 0x7f;
p[0] = 0xed;
p
},
];
let mut shared_secrets = Vec::new();
for point in &small_order_points {
let mut bad_pk_bytes = point.to_vec();
bad_pk_bytes.extend_from_slice(&mlkem_part);
let bad_pk = PublicKey::from_bytes(bad_pk_bytes).unwrap();
let (ct, ss_enc) = encapsulate(&bad_pk).unwrap();
assert_ne!(
ss_enc.as_bytes(),
&[0u8; 32],
"combiner must not produce all-zero SS for small-order point {:02x?}",
&point[..4]
);
let ss_dec = decapsulate(&sk, &ct).unwrap();
assert_ne!(
ss_enc.as_bytes(),
ss_dec.as_bytes(),
"low-order x25519 must cause combiner divergence between enc and dec"
);
shared_secrets.push(ss_enc);
}
for i in 0..shared_secrets.len() {
for j in (i + 1)..shared_secrets.len() {
assert_ne!(
shared_secrets[i].as_bytes(),
shared_secrets[j].as_bytes(),
"shared secrets for points {} and {} should differ",
i,
j
);
}
}
}
#[test]
fn round_trip_repeated() {
for _ in 0..1000 {
let (pk, sk) = keygen().unwrap();
let (ct, ss_enc) = encapsulate(&pk).unwrap();
let ss_dec = decapsulate(&sk, &ct).unwrap();
assert_eq!(ss_enc.as_bytes(), ss_dec.as_bytes());
}
}
fn expand_draft09_seed(seed: &[u8; 32]) -> SecretKey {
use ml_kem::array::Array;
use ml_kem::{B32, EncodedSizeUser, KemCore, MlKem768};
use sha3::Shake256;
use sha3::digest::{ExtendableOutput, Update, XofReader};
use zeroize::Zeroize;
let mut expanded = [0u8; 96];
let mut hasher = Shake256::default();
hasher.update(seed.as_ref());
hasher.finalize_xof().read(&mut expanded);
let d: B32 = Array::from(<[u8; 32]>::try_from(&expanded[0..32]).unwrap());
let z: B32 = Array::from(<[u8; 32]>::try_from(&expanded[32..64]).unwrap());
let skx: [u8; 32] = expanded[64..96].try_into().unwrap();
expanded.zeroize();
let (dk, _ek) = MlKem768::generate_deterministic(&d, &z);
let mlkem_bytes = dk.as_bytes().to_vec();
let mut sk_bytes = Vec::with_capacity(32 + mlkem_bytes.len());
sk_bytes.extend_from_slice(&skx);
sk_bytes.extend_from_slice(&mlkem_bytes);
SecretKey(sk_bytes)
}
#[test]
fn xwing_draft09_decap_kat() {
let seed: [u8; 32] =
hex!("7f9c2ba4e88f827d616045507605853ed73b8093f6efbc88eb1a6eacfa66ef26");
let ct_bytes: [u8; 1120] = hex!(
"b83aa828d4d62b9a83ceffe1d3d3bb1ef31264643c070c5798927e41fb07914a273f8f96"
"e7826cd5375a283d7da885304c5de0516a0f0654243dc5b97f8bfeb831f68251219aabdd"
"723bc6512041acbaef8af44265524942b902e68ffd23221cda70b1b55d776a92d1143ea3"
"a0c475f63ee6890157c7116dae3f62bf72f60acd2bb8cc31ce2ba0de364f52b8ed38c79d"
"719715963a5dd3842d8e8b43ab704e4759b5327bf027c63c8fa857c4908d5a8a7b88ac7f"
"2be394d93c3706ddd4e698cc6ce370101f4d0213254238b4a2e8821b6e414a1cf20f6c12"
"44b699046f5a01caa0a1a55516300b40d2048c77cc73afba79afeea9d2c0118bdf2adb88"
"70dc328c5516cc45b1a2058141039e2c90a110a9e16b318dfb53bd49a126d6b73f215787"
"517b8917cc01cabd107d06859854ee8b4f9861c226d3764c87339ab16c3667d2f49384e5"
"5456dd40414b70a6af841585f4c90c68725d57704ee8ee7ce6e2f9be582dbee985e038ff"
"c346ebfb4e22158b6c84374a9ab4a44e1f91de5aac5197f89bc5e5442f51f9a5937b102b"
"a3beaebf6e1c58380a4a5fedce4a4e5026f88f528f59ffd2db41752b3a3d90efabe46389"
"9b7d40870c530c8841e8712b733668ed033adbfafb2d49d37a44d4064e5863eb0af0a08d"
"47b3cc888373bc05f7a33b841bc2587c57eb69554e8a3767b7506917b6b70498727f16ea"
"c1a36ec8d8cfaf751549f2277db277e8a55a9a5106b23a0206b4721fa9b3048552c5bd5b"
"594d6e247f38c18c591aea7f56249c72ce7b117afcc3a8621582f9cf71787e183dee0936"
"7976e98409ad9217a497df888042384d7707a6b78f5f7fb8409e3b535175373461b77600"
"2d799cbad62860be70573ecbe13b246e0da7e93a52168e0fb6a9756b895ef7f0147a0dc8"
"1bfa644b088a9228160c0f9acf1379a2941cd28c06ebc80e44e17aa2f8177010afd78a97"
"ce0868d1629ebb294c5151812c583daeb88685220f4da9118112e07041fcc24d5564a99f"
"dbde28869fe0722387d7a9a4d16e1cc8555917e09944aa5ebaaaec2cf62693afad42a3f5"
"18fce67d273cc6c9fb5472b380e8573ec7de06a3ba2fd5f931d725b493026cb0acbd3fe6"
"2d00e4c790d965d7a03a3c0b4222ba8c2a9a16e2ac658f572ae0e746eafc4feba023576f"
"08942278a041fb82a70a595d5bacbf297ce2029898a71e5c3b0d1c6228b485b1ade509b3"
"5fbca7eca97b2132e7cb6bc465375146b7dceac969308ac0c2ac89e7863eb8943015b243"
"14cafb9c7c0e85fe543d56658c213632599efabfc1ec49dd8c88547bb2cc40c9d38cbd30"
"99b4547840560531d0188cd1e9c23a0ebee0a03d5577d66b1d2bcb4baaf21cc7fef1e038"
"06ca96299df0dfbc56e1b2b43e4fc20c37f834c4af62127e7dae86c3c25a2f696ac8b589"
"dec71d595bfbe94b5ed4bc07d800b330796fda89edb77be0294136139354eb8cd3759157"
"8f9c600dd9be8ec6219fdd507adf3397ed4d68707b8d13b24ce4cd8fb22851bfe9d63240"
"7f31ed6f7cb1600de56f17576740ce2a32fc5145030145cfb97e63e0e41d354274a079d3"
"e6fb2e15"
);
let expected_ss: [u8; 32] =
hex!("d2df0522128f09dd8e2c92b1e905c793d8f57a54c3da25861f10bf4ca613e384");
let mut lo_ct = Vec::with_capacity(1120);
lo_ct.extend_from_slice(&ct_bytes[1088..]); lo_ct.extend_from_slice(&ct_bytes[..1088]);
let sk = expand_draft09_seed(&seed);
let ct = Ciphertext::from_bytes(lo_ct).unwrap();
let ss = decapsulate(&sk, &ct).unwrap();
assert_eq!(ss.as_bytes(), &expected_ss);
}
}