1use ml_kem::{
2 DecapsulationKey1024, EncapsulationKey1024, KeyExport, MlKem1024,
3 kem::{Decapsulate, Encapsulate, Kem},
4};
5use zeroize::Zeroizing;
6
7use crate::errors::CryptoError;
8
9pub const PUBLIC_KEY_BYTES: usize = 1568;
11pub const SECRET_KEY_BYTES: usize = 64;
13pub const CIPHERTEXT_BYTES: usize = 1568;
15pub const SHARED_SECRET_BYTES: usize = 32;
17
18pub type SharedSecret = Zeroizing<[u8; SHARED_SECRET_BYTES]>;
20
21pub struct KEMPair {
28 sec_key: DecapsulationKey1024,
29}
30
31impl KEMPair {
32 pub fn create() -> Self {
37 let (sec_key, _) = MlKem1024::generate_keypair();
38 Self { sec_key }
39 }
40
41 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 pub fn pub_key_bytes(&self) -> [u8; PUBLIC_KEY_BYTES] {
67 self.sec_key.encapsulation_key().to_bytes().into()
68 }
69
70 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 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 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 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 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 assert!(matches!(
177 KEMPair::from_bytes(&keypair.pub_key_bytes()[1..], &keypair.to_bytes().1[..]),
178 Err(CryptoError::InvalidKey)
179 ));
180 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 assert_ne!(shared_secret, KEMPair::create().decapsulate(&ciphertext));
201 }
202}