Skip to main content

pq_msg/exchange/
pair.rs

1use ml_kem::{
2    DecapsulationKey1024, EncapsulationKey1024, KeyExport, MlKem1024,
3    kem::{Decapsulate, Encapsulate, Kem},
4};
5use zeroize::Zeroizing;
6
7use crate::errors::CryptoError;
8
9/// Size of an ML-KEM-1024 public (encapsulation) key in bytes
10pub const PUBLIC_KEY_BYTES: usize = 1568;
11/// Size of an ML-KEM-1024 secret key in bytes (the 64-byte FIPS 203 seed)
12pub const SECRET_KEY_BYTES: usize = 64;
13/// Size of an ML-KEM-1024 ciphertext in bytes
14pub const CIPHERTEXT_BYTES: usize = 1568;
15/// Size of an ML-KEM shared secret in bytes
16pub const SHARED_SECRET_BYTES: usize = 32;
17
18/// A shared secret that is wiped from memory when dropped
19pub type SharedSecret = Zeroizing<[u8; SHARED_SECRET_BYTES]>;
20
21/// A Key Encapsulation Mechanism (KEM) pair using ML-KEM (formerly Kyber)
22///
23/// This struct represents a post-quantum cryptography key pair used for
24/// key encapsulation and decapsulation operations. It utilizes ML-KEM-1024,
25/// which provides 256-bit equivalent security strength.
26/// The secret key is wiped from memory when the pair is dropped.
27pub struct KEMPair {
28    sec_key: DecapsulationKey1024,
29}
30
31impl KEMPair {
32    /// Creates a new random KEM pair
33    ///
34    /// # Returns
35    /// A new KEMPair with generated public and secret keys
36    pub fn create() -> Self {
37        let (sec_key, _) = MlKem1024::generate_keypair();
38        Self { sec_key }
39    }
40
41    /// Creates a KEM pair from separate public and secret key bytes
42    ///
43    /// # Arguments
44    /// * `pub_key` - The public key bytes
45    /// * `sec_key` - The secret key bytes (64-byte seed)
46    ///
47    /// # Returns
48    /// - `Result<KEMPair, CryptoError>`: The constructed KEMPair, or an error if
49    ///   the lengths are wrong or the public key does not belong to the secret key
50    pub fn from_bytes(pub_key: &[u8], sec_key: &[u8]) -> Result<Self, CryptoError> {
51        let seed: Zeroizing<[u8; SECRET_KEY_BYTES]> = Zeroizing::new(
52            sec_key
53                .try_into()
54                .map_err(|_| CryptoError::IncongruentLength(SECRET_KEY_BYTES, sec_key.len()))?,
55        );
56        let pair = Self {
57            sec_key: DecapsulationKey1024::from_seed((*seed).into()),
58        };
59        if pair.pub_key_bytes()[..] != *pub_key {
60            return Err(CryptoError::InvalidKey);
61        }
62        Ok(pair)
63    }
64
65    /// Returns the public key bytes, which are shared with peers
66    pub fn pub_key_bytes(&self) -> [u8; PUBLIC_KEY_BYTES] {
67        self.sec_key.encapsulation_key().to_bytes().into()
68    }
69
70    /// Converts the key pair to raw byte arrays
71    ///
72    /// # Returns
73    /// A tuple containing the public key and the secret key (wiped when dropped)
74    pub fn to_bytes(&self) -> ([u8; PUBLIC_KEY_BYTES], Zeroizing<[u8; SECRET_KEY_BYTES]>) {
75        (
76            self.pub_key_bytes(),
77            Zeroizing::new(self.sec_key.to_bytes().into()),
78        )
79    }
80
81    /// Converts the key pair to a single byte vector with public key followed by secret key
82    ///
83    /// # Returns
84    /// A vector containing the concatenated public and secret key bytes (wiped when dropped)
85    pub fn to_bytes_uniform(&self) -> Zeroizing<Vec<u8>> {
86        let (pub_key, sec_key) = self.to_bytes();
87        Zeroizing::new([&pub_key[..], &sec_key[..]].concat())
88    }
89
90    /// Creates a KEM pair from a single byte slice containing both public and secret keys
91    ///
92    /// # Arguments
93    /// * `bytes` - The concatenated public and secret key bytes
94    ///
95    /// # Returns
96    /// - `Result<KEMPair, CryptoError>`: The constructed KEMPair or an error
97    pub fn from_bytes_uniform(bytes: &[u8]) -> Result<Self, CryptoError> {
98        if bytes.len() != PUBLIC_KEY_BYTES + SECRET_KEY_BYTES {
99            return Err(CryptoError::IncongruentLength(
100                PUBLIC_KEY_BYTES + SECRET_KEY_BYTES,
101                bytes.len(),
102            ));
103        }
104        let (pub_key, sec_key) = bytes.split_at(PUBLIC_KEY_BYTES);
105        Self::from_bytes(pub_key, sec_key)
106    }
107
108    /// Encapsulates a fresh shared secret for the holder of `receiver_pubkey`
109    ///
110    /// # Arguments
111    /// * `receiver_pubkey` - The receiver's public key
112    ///
113    /// # Returns
114    /// - `Result<(SharedSecret, [u8; CIPHERTEXT_BYTES]), CryptoError>`: The shared secret
115    ///   and the ciphertext to send to the receiver, or an error if the key is invalid
116    pub fn encapsulate(
117        receiver_pubkey: &[u8; PUBLIC_KEY_BYTES],
118    ) -> Result<(SharedSecret, [u8; CIPHERTEXT_BYTES]), CryptoError> {
119        let ek = EncapsulationKey1024::new(receiver_pubkey.into())
120            .map_err(|_| CryptoError::InvalidKey)?;
121        let (ciphertext, shared_secret) = ek.encapsulate();
122        Ok((Zeroizing::new(shared_secret.into()), ciphertext.into()))
123    }
124
125    /// Decapsulates a shared secret from the provided ciphertext using this pair's secret key
126    ///
127    /// ML-KEM uses implicit rejection: a tampered ciphertext yields an unrelated
128    /// secret instead of an error, so a mismatch only shows up once decryption fails.
129    ///
130    /// # Arguments
131    /// * `ciphertext` - The ciphertext received from the sender
132    ///
133    /// # Returns
134    /// The decapsulated shared secret
135    pub fn decapsulate(&self, ciphertext: &[u8; CIPHERTEXT_BYTES]) -> SharedSecret {
136        Zeroizing::new(self.sec_key.decapsulate(ciphertext.into()).into())
137    }
138}
139
140#[cfg(test)]
141mod tests {
142    use super::*;
143
144    #[test]
145    fn test_keypair() {
146        let keypair = KEMPair::create();
147        let (pub_key, sec_key) = keypair.to_bytes();
148        let new_keypair = KEMPair::from_bytes(&pub_key, &sec_key[..]).unwrap();
149        assert_eq!(keypair.to_bytes_uniform(), new_keypair.to_bytes_uniform());
150
151        let uniform = KEMPair::from_bytes_uniform(&keypair.to_bytes_uniform()).unwrap();
152        assert_eq!(keypair.to_bytes_uniform(), uniform.to_bytes_uniform());
153    }
154
155    #[test]
156    fn test_from_bytes_rejects_foreign_pubkey() {
157        let keypair = KEMPair::create();
158        let other = KEMPair::create();
159        let result = KEMPair::from_bytes(&other.pub_key_bytes(), &keypair.to_bytes().1[..]);
160        assert!(matches!(result, Err(CryptoError::InvalidKey)));
161        assert!(matches!(
162            KEMPair::from_bytes(&keypair.pub_key_bytes(), &[0u8; 10]),
163            Err(CryptoError::IncongruentLength(SECRET_KEY_BYTES, 10))
164        ));
165    }
166
167    #[test]
168    fn test_invalid_inputs() {
169        let keypair = KEMPair::create();
170
171        assert!(matches!(
172            KEMPair::from_bytes_uniform(&[0u8; 10]),
173            Err(CryptoError::IncongruentLength(n, 10)) if n == PUBLIC_KEY_BYTES + SECRET_KEY_BYTES
174        ));
175        // A truncated public key can't match the secret key
176        assert!(matches!(
177            KEMPair::from_bytes(&keypair.pub_key_bytes()[1..], &keypair.to_bytes().1[..]),
178            Err(CryptoError::InvalidKey)
179        ));
180        // Not a valid ML-KEM public key: coefficients out of range
181        assert!(matches!(
182            KEMPair::encapsulate(&[0xFF; PUBLIC_KEY_BYTES]),
183            Err(CryptoError::InvalidKey)
184        ));
185    }
186
187    #[test]
188    fn test_encapsulate_decapsulate() {
189        let receiver = KEMPair::create();
190
191        let (shared_secret, ciphertext) = KEMPair::encapsulate(&receiver.pub_key_bytes()).unwrap();
192        let dec_shared_secret = receiver.decapsulate(&ciphertext);
193
194        assert_eq!(
195            shared_secret, dec_shared_secret,
196            "Difference in shared secrets!"
197        );
198
199        // A different receiver gets an unrelated secret (implicit rejection)
200        assert_ne!(shared_secret, KEMPair::create().decapsulate(&ciphertext));
201    }
202}