Skip to main content

crypto/
xwing.rs

1//! X-Wing: hybrid post-quantum key encapsulation mechanism (KEM) algorithm (ML-KEM-768 with X25519).
2
3use crate::{
4    curve25519::x25519,
5    mlkem::{self, MlKemError},
6    sha3::{Sha3_256, Shake256},
7};
8
9/// Size of the X-Wing secret key (32 bytes).
10pub const SECRET_KEY_SIZE: usize = 32;
11/// Size of the X-Wing public key in bytes (ML-KEM-768 pk + X25519 pk = 1216).
12pub const PUBLIC_KEY_SIZE: usize = mlkem::PUBLIC_KEY_SIZE_768 + x25519::KEY_SIZE; // 1216
13/// Size of the X-Wing ciphertext in bytes (ML-KEM-768 ct + X25519 shared secret = 1120).
14pub const CIPHERTEXT_SIZE: usize = mlkem::CIPHERTEXT_SIZE_768 + x25519::SHARED_SECRET_SIZE; // 1120
15/// Size of the X-Wing shared secret (32 bytes).
16pub const SHARED_SECRET_SIZE: usize = 32;
17
18const XWING_LABEL: &[u8; 6] = b"\\.//^\\";
19
20/// X-Wing error type.
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
22pub enum XWingError {
23    MlKem(MlKemError),
24}
25
26impl From<MlKemError> for XWingError {
27    fn from(err: MlKemError) -> Self {
28        XWingError::MlKem(err)
29    }
30}
31
32impl core::fmt::Display for XWingError {
33    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
34        match self {
35            XWingError::MlKem(err) => write!(f, "ML-KEM error: {err}"),
36        }
37    }
38}
39
40/// X-Wing hybrid KEM decapsulation (secret) key.
41///
42/// Combines an ML-KEM-768 and an X25519 secret key as specified in the
43/// X-Wing draft. The shared secret is derived via a combiner that hashes
44/// both component secrets together.
45///
46/// # Example
47///
48/// ```ignore
49/// use crypto::xwing::{generate_keypair, SecretKey, PublicKey};
50///
51/// let (secret_key, public_key) = generate_keypair();
52/// let (shared_secret, ciphertext) = public_key.encapsulate();
53/// let decapsulated = secret_key.decapsulate(&ciphertext).unwrap();
54/// assert_eq!(shared_secret, decapsulated);
55/// ```
56#[derive(Clone, Debug, PartialEq, Eq)]
57pub struct SecretKey {
58    bytes: [u8; SECRET_KEY_SIZE],
59    x25519_secret_key: x25519::SecretKey,
60    x25519_public_key_bytes: [u8; x25519::KEY_SIZE],
61    mlkem_secret_key: mlkem::SecretKey768,
62}
63
64impl SecretKey {
65    pub fn to_bytes(&self) -> [u8; SECRET_KEY_SIZE] {
66        self.bytes
67    }
68
69    pub fn decapsulate(&self, ct: &[u8; CIPHERTEXT_SIZE]) -> Result<[u8; SHARED_SECRET_SIZE], XWingError> {
70        let ct_m = &ct[..mlkem::CIPHERTEXT_SIZE_768].try_into().unwrap();
71        let ct_x = x25519::PublicKey::from_bytes(&ct[mlkem::CIPHERTEXT_SIZE_768..].try_into().unwrap());
72
73        let ss_m = self.mlkem_secret_key.decapsulate(&ct_m)?;
74        let ss_x = self.x25519_secret_key.ecdh(&ct_x);
75
76        Ok(combiner(&ss_m, &ss_x, &ct_x.to_bytes(), &self.x25519_public_key_bytes))
77    }
78}
79
80/// X-Wing hybrid KEM encapsulation (public) key.
81///
82/// See [`SecretKey`] for a full usage example.
83#[derive(Clone, Debug, PartialEq, Eq)]
84pub struct PublicKey {
85    mlkem_public_key: mlkem::PublicKey768,
86    x25519_public_key: x25519::PublicKey,
87}
88
89impl PublicKey {
90    pub fn to_bytes(&self) -> [u8; PUBLIC_KEY_SIZE] {
91        let mut bytes = [0u8; PUBLIC_KEY_SIZE];
92        bytes[..mlkem::PUBLIC_KEY_SIZE_768].copy_from_slice(&self.mlkem_public_key.to_bytes());
93        bytes[mlkem::PUBLIC_KEY_SIZE_768..].copy_from_slice(&self.x25519_public_key.to_bytes());
94        bytes
95    }
96
97    #[cfg(feature = "random")]
98    pub fn encapsulate(&self) -> ([u8; SHARED_SECRET_SIZE], [u8; CIPHERTEXT_SIZE]) {
99        let eseed: [u8; 64] = crate::random::random_bytes();
100        self.encapsulate_derand(&eseed)
101    }
102
103    fn encapsulate_derand(&self, eseed: &[u8; 64]) -> ([u8; SHARED_SECRET_SIZE], [u8; CIPHERTEXT_SIZE]) {
104        let ek_x = x25519::SecretKey::from_bytes(&eseed[32..64].try_into().unwrap());
105        let ct_x = ek_x.public_key();
106        let ss_x = ek_x.ecdh(&self.x25519_public_key);
107
108        let m = &eseed[..32].try_into().unwrap();
109        let (ct_m, ss_m) = self.mlkem_public_key.encapsulate_derand(&m);
110
111        let ss = combiner(&ss_m, &ss_x, &ct_x.to_bytes(), &self.x25519_public_key.to_bytes());
112
113        let mut ct = [0u8; CIPHERTEXT_SIZE];
114        ct[..mlkem::CIPHERTEXT_SIZE_768].copy_from_slice(&ct_m);
115        ct[mlkem::CIPHERTEXT_SIZE_768..].copy_from_slice(&ct_x.to_bytes());
116
117        (ss, ct)
118    }
119}
120
121/// Generate an X-Wing keypair.
122///
123/// This is a convenience wrapper around [`SecretKey`] generation.
124///
125/// See [`SecretKey`] for a usage example.
126#[cfg(feature = "random")]
127pub fn generate_keypair() -> (SecretKey, PublicKey) {
128    let seed: [u8; SECRET_KEY_SIZE] = crate::random::random_bytes();
129    generate_keypair_derand(&seed)
130}
131
132/// Generate a deterministic keypair from the given seed (for testing).
133fn generate_keypair_derand(secret_key: &[u8; SECRET_KEY_SIZE]) -> (SecretKey, PublicKey) {
134    let (mlkem_sk, x25519_sk, mlkem_pk, x25519_pk) = expand_decapsulation_key(secret_key);
135
136    let secret_key = SecretKey {
137        bytes: *secret_key,
138        x25519_secret_key: x25519_sk,
139        x25519_public_key_bytes: x25519_pk.to_bytes(),
140        mlkem_secret_key: mlkem_sk,
141    };
142
143    let public_key = PublicKey {
144        mlkem_public_key: mlkem_pk,
145        x25519_public_key: x25519_pk,
146    };
147
148    (secret_key, public_key)
149}
150
151fn expand_decapsulation_key(
152    secret_key: &[u8; 32],
153) -> (mlkem::SecretKey768, x25519::SecretKey, mlkem::PublicKey768, x25519::PublicKey) {
154    let mut expanded_secret_key = [0u8; 96];
155    Shake256::hash(secret_key, &mut expanded_secret_key);
156
157    let (sk_m, pk_m) = derive_mlkeem_keys(&expanded_secret_key);
158
159    let sk_x = x25519::SecretKey::from_bytes(&expanded_secret_key[64..96].try_into().unwrap());
160    let pk_x = sk_x.public_key();
161
162    (sk_m, sk_x, pk_m, pk_x)
163}
164
165fn derive_mlkeem_keys(expnded_secret_key: &[u8; 96]) -> (mlkem::SecretKey768, mlkem::PublicKey768) {
166    mlkem::generate_keypair_768_derand(&expnded_secret_key[..64].try_into().unwrap())
167}
168
169fn combiner(
170    ss_m: &[u8; mlkem::SHARED_SECRET_SIZE],
171    ss_x: &[u8; x25519::KEY_SIZE],
172    ct_x: &[u8; x25519::KEY_SIZE],
173    pk_x: &[u8; x25519::KEY_SIZE],
174) -> [u8; SHARED_SECRET_SIZE] {
175    use crate::Hasher;
176    let mut hasher = Sha3_256::new();
177    hasher.update(ss_m);
178    hasher.update(ss_x);
179    hasher.update(ct_x);
180    hasher.update(pk_x);
181    hasher.update(XWING_LABEL);
182    hasher.sum().as_ref().try_into().unwrap()
183}
184
185#[cfg(test)]
186mod tests {
187    use super::*;
188
189    fn hex_to_array<const N: usize>(hex_str: &str) -> [u8; N] {
190        let bytes = hex::decode(hex_str).unwrap();
191        return bytes.try_into().unwrap();
192    }
193
194    #[test]
195    fn constants() {
196        assert!(PUBLIC_KEY_SIZE == 1216);
197        assert!(CIPHERTEXT_SIZE == 1120);
198    }
199
200    struct TestVector {
201        seed: &'static str,
202        eseed: &'static str,
203        ss: &'static str,
204    }
205
206    const TEST_VECTORS: [TestVector; 3] = [
207        TestVector {
208            seed: "7f9c2ba4e88f827d616045507605853ed73b8093f6efbc88eb1a6eacfa66ef26",
209            eseed: "3cb1eea988004b93103cfb0aeefd2a686e01fa4a58e8a3639ca8a1e3f9ae57e235b8cc873c23dc62b8d260169afa2f75ab916a58d974918835d25e6a435085b2",
210            ss: "d2df0522128f09dd8e2c92b1e905c793d8f57a54c3da25861f10bf4ca613e384",
211        },
212        TestVector {
213            seed: "badfd6dfaac359a5efbb7bcc4b59d538df9a04302e10c8bc1cbf1a0b3a5120ea",
214            eseed: "17cda7cfad765f5623474d368ccca8af0007cd9f5e4c849f167a580b14aabdefaee7eef47cb0fca9767be1fda69419dfb927e9df07348b196691abaeb580b32d",
215            ss: "f2e86241c64d60f6649fbc6c5b7d17180b780a3f34355e64a85749949c45f150",
216        },
217        TestVector {
218            seed: "ef58538b8d23f87732ea63b02b4fa0f4873360e2841928cd60dd4cee8cc0d4c9",
219            eseed: "22a96188d032675c8ac850933c7aff1533b94c834adbb69c6115bad4692d8619f90b0cdf8a7b9c264029ac185b70b83f2801f2f4b3f70c593ea3aeeb613a7f1b",
220            ss: "953f7f4e8c5b5049bdc771d1dffada0dd961477d1a2ae0988baa7ea6898d893f",
221        },
222    ];
223
224    #[test]
225    fn test_vectors_from_draft() {
226        for (i, tv) in TEST_VECTORS.iter().enumerate() {
227            let seed: [u8; 32] = hex_to_array(tv.seed);
228            let eseed: [u8; 64] = hex_to_array(tv.eseed);
229            let expected_ss: [u8; 32] = hex_to_array(tv.ss);
230
231            let (secret_key, pk) = generate_keypair_derand(&seed);
232            assert_eq!(secret_key.to_bytes(), seed, "vector {i}: sk mismatch");
233
234            let (ss, ct) = pk.encapsulate_derand(&eseed);
235            assert_eq!(ss, expected_ss, "vector {i}: encaps ss mismatch");
236
237            let decapsulated_ss = secret_key.decapsulate(&ct).unwrap();
238            assert_eq!(decapsulated_ss, expected_ss, "vector {i}: decaps ss mismatch");
239        }
240    }
241
242    #[test]
243    fn round_trip() {
244        let (secret_key, public_key) = generate_keypair();
245        let (ss, ct) = public_key.encapsulate();
246        let decapsulated = secret_key.decapsulate(&ct).unwrap();
247        assert_eq!(ss, decapsulated);
248    }
249
250    #[test]
251    fn round_trip_many() {
252        for _ in 0..10 {
253            let (secret_key, public_key) = generate_keypair();
254            let (ss, ct) = public_key.encapsulate();
255            let decapsulated = secret_key.decapsulate(&ct).unwrap();
256            assert_eq!(ss, decapsulated);
257        }
258    }
259
260    #[test]
261    fn decapsulation_with_wrong_key_produces_different_secret() {
262        let (_, pk_a) = generate_keypair();
263        let (sk_b, _) = generate_keypair();
264
265        let (ss_a, ct) = pk_a.encapsulate();
266        let ss_b = sk_b.decapsulate(&ct).unwrap();
267        assert_ne!(ss_a, ss_b);
268    }
269
270    #[test]
271    fn tampered_ciphertext_produces_different_secret() {
272        let (secret_key, public_key) = generate_keypair();
273        let (ss, mut ct) = public_key.encapsulate();
274
275        ct[0] ^= 0x80;
276
277        let tampered_ss = secret_key.decapsulate(&ct).unwrap();
278        assert_ne!(ss, tampered_ss);
279    }
280
281    #[test]
282    fn derandomized_keygen_is_deterministic() {
283        let seed: [u8; 32] = hex_to_array("7f9c2ba4e88f827d616045507605853ed73b8093f6efbc88eb1a6eacfa66ef26");
284        let (sk1, pk1) = generate_keypair_derand(&seed);
285        let (sk2, pk2) = generate_keypair_derand(&seed);
286        assert_eq!(sk1.to_bytes(), sk2.to_bytes());
287        assert_eq!(pk1.to_bytes(), pk2.to_bytes());
288    }
289
290    #[test]
291    fn derandomized_encaps_is_deterministic() {
292        let seed: [u8; 32] = hex_to_array("7f9c2ba4e88f827d616045507605853ed73b8093f6efbc88eb1a6eacfa66ef26");
293        let eseed: [u8; 64] = hex_to_array(
294            "3cb1eea988004b93103cfb0aeefd2a686e01fa4a58e8a3639ca8a1e3f9ae57e235b8cc873c23dc62b8d260169afa2f75ab916a58d974918835d25e6a435085b2",
295        );
296        let (_, pk) = generate_keypair_derand(&seed);
297
298        let (ss1, ct1) = pk.encapsulate_derand(&eseed);
299        let (ss2, ct2) = pk.encapsulate_derand(&eseed);
300        assert_eq!(ct1, ct2);
301        assert_eq!(ss1, ss2);
302    }
303
304    #[test]
305    fn xwing_label_is_correct() {
306        assert_eq!(XWING_LABEL.len(), 6);
307        assert_eq!(hex::encode(XWING_LABEL), "5c2e2f2f5e5c");
308    }
309
310    #[test]
311    fn expand_decapsulation_key_is_deterministic() {
312        let seed: [u8; 32] = hex_to_array("7f9c2ba4e88f827d616045507605853ed73b8093f6efbc88eb1a6eacfa66ef26");
313
314        let (sk_m1, sk_x1, pk_m1, pk_x1) = expand_decapsulation_key(&seed);
315        let (sk_m2, sk_x2, pk_m2, pk_x2) = expand_decapsulation_key(&seed);
316        assert_eq!(sk_m1, sk_m2);
317        assert_eq!(sk_x1, sk_x2);
318        assert_eq!(pk_m1, pk_m2);
319        assert_eq!(pk_x1, pk_x2);
320    }
321
322    #[test]
323    fn combiner_is_deterministic() {
324        let ss_m = [0x01u8; 32];
325        let ss_x = [0x02u8; 32];
326        let ct_x = [0x03u8; 32];
327        let pk_x = [0x04u8; 32];
328
329        let result1 = combiner(&ss_m, &ss_x, &ct_x, &pk_x);
330        let result2 = combiner(&ss_m, &ss_x, &ct_x, &pk_x);
331        assert_eq!(result1, result2);
332    }
333}