1use crate::{Error, Result};
4use alloc::vec::Vec;
5use rand::RngCore;
6use x25519_dalek::{x25519, X25519_BASEPOINT_BYTES};
7use zeroize::{Zeroize, ZeroizeOnDrop};
8
9#[derive(Clone, Zeroize, ZeroizeOnDrop)]
11pub struct Key([u8; 32]);
12
13impl Key {
14 pub const SIZE: usize = 32;
16
17 pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
19 if bytes.len() != Self::SIZE {
20 return Err(Error::InvalidKeyLength {
21 expected: Self::SIZE,
22 actual: bytes.len(),
23 });
24 }
25 let mut key = [0u8; 32];
26 key.copy_from_slice(bytes);
27 Ok(Key(key))
28 }
29
30 pub fn as_bytes(&self) -> &[u8; 32] {
32 &self.0
33 }
34
35 pub fn to_base64(&self) -> String {
37 use base64::{engine::general_purpose::STANDARD, Engine};
38 STANDARD.encode(&self.0)
39 }
40
41 pub fn from_base64(encoded: &str) -> Result<Self> {
43 use base64::{engine::general_purpose::STANDARD, Engine};
44 if encoded.len() != 44 || !encoded.ends_with('=') {
48 return Err(Error::InvalidKeyFormat(
49 "AES-256 key must use canonical padded Base64".to_string(),
50 ));
51 }
52 let mut bytes = STANDARD.decode(encoded)?;
53 let result = if STANDARD.encode(&bytes) != encoded {
54 Err(Error::InvalidKeyFormat(
55 "AES-256 key must use canonical padded Base64".to_string(),
56 ))
57 } else {
58 Self::from_bytes(&bytes)
59 };
60 bytes.zeroize();
61 result
62 }
63}
64
65impl AsRef<[u8]> for Key {
66 fn as_ref(&self) -> &[u8] {
67 &self.0
68 }
69}
70
71impl core::fmt::Debug for Key {
72 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
73 f.debug_struct("Key")
74 .field("length", &Self::SIZE)
75 .finish_non_exhaustive()
76 }
77}
78
79pub const X25519_KEY_SIZE: usize = 32;
81
82pub const HKDF_SHA256_MAX_OUTPUT: usize = 255 * 32;
84
85pub const PBKDF2_MIN_ITERATIONS: u32 = 100_000;
87
88pub const PBKDF2_MAX_ITERATIONS: u32 = 1_000_000;
93
94pub const PBKDF2_MIN_SALT_SIZE: usize = 16;
96
97pub const PBKDF2_MAX_SALT_SIZE: usize = 1024;
99
100#[derive(Clone, Zeroize, ZeroizeOnDrop)]
102pub struct X25519KeyPair {
103 pub public_key: [u8; X25519_KEY_SIZE],
105 pub private_key: [u8; X25519_KEY_SIZE],
107}
108
109pub fn generate_key() -> Key {
111 let mut key = [0u8; 32];
112 rand::thread_rng().fill_bytes(&mut key);
113 Key(key)
114}
115
116pub fn derive_key_hkdf(input_key_material: &[u8], salt: Option<&[u8]>, info: &[u8]) -> Result<Key> {
128 let mut okm = derive_key_hkdf_raw(input_key_material, salt, info, Key::SIZE)?;
129 let key = Key::from_bytes(&okm);
130 okm.zeroize();
131 key
132}
133
134pub fn derive_key_hkdf_raw(
143 input_key_material: &[u8],
144 salt: Option<&[u8]>,
145 info: &[u8],
146 length: usize,
147) -> Result<Vec<u8>> {
148 use hkdf::Hkdf;
149 use sha2::Sha256;
150
151 if length == 0 {
152 return Err(Error::KeyDerivationFailed(
153 "HKDF output length must be > 0".to_string(),
154 ));
155 }
156 if length > HKDF_SHA256_MAX_OUTPUT {
157 return Err(Error::KeyDerivationFailed(format!(
158 "HKDF-SHA256 output length must be at most {HKDF_SHA256_MAX_OUTPUT} bytes"
159 )));
160 }
161
162 let hk = Hkdf::<Sha256>::new(salt, input_key_material);
163 let mut okm = vec![0u8; length];
164
165 hk.expand(info, &mut okm)
166 .map_err(|e| Error::KeyDerivationFailed(e.to_string()))?;
167
168 Ok(okm)
169}
170
171pub fn derive_key_pbkdf2(password: &[u8], salt: &[u8], iterations: u32) -> Result<Key> {
183 use pbkdf2::pbkdf2_hmac;
184 use sha2::Sha256;
185
186 validate_pbkdf2_parameters(salt, iterations)?;
187
188 let mut key = [0u8; 32];
189 pbkdf2_hmac::<Sha256>(password, salt, iterations, &mut key);
190
191 Ok(Key(key))
192}
193
194pub fn validate_pbkdf2_parameters(salt: &[u8], iterations: u32) -> Result<()> {
196 if !(PBKDF2_MIN_SALT_SIZE..=PBKDF2_MAX_SALT_SIZE).contains(&salt.len()) {
197 return Err(Error::InvalidConfiguration(format!(
198 "PBKDF2 salt must be between {PBKDF2_MIN_SALT_SIZE} and {PBKDF2_MAX_SALT_SIZE} bytes"
199 )));
200 }
201 if !(PBKDF2_MIN_ITERATIONS..=PBKDF2_MAX_ITERATIONS).contains(&iterations) {
202 return Err(Error::InvalidConfiguration(format!(
203 "PBKDF2 iterations must be between {PBKDF2_MIN_ITERATIONS} and {PBKDF2_MAX_ITERATIONS}"
204 )));
205 }
206 Ok(())
207}
208
209pub fn generate_x25519_key_pair(seed: Option<&[u8]>) -> Result<X25519KeyPair> {
213 let mut private_key = [0u8; X25519_KEY_SIZE];
214
215 if let Some(seed_bytes) = seed {
216 if seed_bytes.len() != X25519_KEY_SIZE {
217 return Err(Error::InvalidKeyLength {
218 expected: X25519_KEY_SIZE,
219 actual: seed_bytes.len(),
220 });
221 }
222 private_key.copy_from_slice(seed_bytes);
223 } else {
224 rand::thread_rng().fill_bytes(&mut private_key);
225 }
226
227 let public_key = x25519(private_key, X25519_BASEPOINT_BYTES);
228 Ok(X25519KeyPair {
229 public_key,
230 private_key,
231 })
232}
233
234pub fn x25519_shared_secret(
236 our_private_key: &[u8],
237 their_public_key: &[u8],
238) -> Result<[u8; X25519_KEY_SIZE]> {
239 if our_private_key.len() != X25519_KEY_SIZE {
240 return Err(Error::InvalidKeyLength {
241 expected: X25519_KEY_SIZE,
242 actual: our_private_key.len(),
243 });
244 }
245 if their_public_key.len() != X25519_KEY_SIZE {
246 return Err(Error::InvalidKeyLength {
247 expected: X25519_KEY_SIZE,
248 actual: their_public_key.len(),
249 });
250 }
251
252 let mut private_key = [0u8; X25519_KEY_SIZE];
253 private_key.copy_from_slice(our_private_key);
254
255 let mut public_key = [0u8; X25519_KEY_SIZE];
256 public_key.copy_from_slice(their_public_key);
257
258 let mut shared_secret = x25519(private_key, public_key);
259 private_key.zeroize();
260
261 let shared_secret_or = shared_secret.iter().fold(0u8, |acc, byte| acc | byte);
262 if shared_secret_or == 0 {
263 shared_secret.zeroize();
264 return Err(Error::KeyDerivationFailed(
265 "X25519 agreement rejected an all-zero shared secret from a low-order or invalid public key"
266 .to_string(),
267 ));
268 }
269
270 Ok(shared_secret)
271}
272
273pub fn derive_key_from_shared_secret(shared_secret: &[u8], salt: &str, info: &str) -> Result<Key> {
275 if shared_secret.len() != X25519_KEY_SIZE {
276 return Err(Error::InvalidKeyLength {
277 expected: X25519_KEY_SIZE,
278 actual: shared_secret.len(),
279 });
280 }
281 if shared_secret.iter().fold(0u8, |acc, byte| acc | byte) == 0 {
282 return Err(Error::KeyDerivationFailed(
283 "shared secret must not be all zero".to_string(),
284 ));
285 }
286 derive_key_hkdf(shared_secret, Some(salt.as_bytes()), info.as_bytes())
287}
288
289#[allow(dead_code)]
291pub fn generate_salt(length: usize) -> Vec<u8> {
292 let mut salt = vec![0u8; length];
293 rand::thread_rng().fill_bytes(&mut salt);
294 salt
295}
296
297#[cfg(test)]
298mod tests {
299 use super::*;
300
301 #[test]
302 fn test_key_generation() {
303 let key1 = generate_key();
304 let key2 = generate_key();
305
306 assert_ne!(key1.as_bytes(), key2.as_bytes());
308
309 assert_eq!(key1.as_bytes().len(), 32);
311 }
312
313 #[test]
314 fn test_key_base64_roundtrip() {
315 let key = generate_key();
316 let encoded = key.to_base64();
317 let decoded = Key::from_base64(&encoded).unwrap();
318
319 assert_eq!(key.as_bytes(), decoded.as_bytes());
320 }
321
322 #[test]
323 fn test_key_base64_rejects_noncanonical_and_wrong_length_encodings() {
324 assert!(Key::from_base64("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAB=").is_err());
325 assert!(Key::from_base64("AA==").is_err());
326 assert!(Key::from_base64("not base64").is_err());
327 assert!(Key::from_base64(&"A".repeat(1024 * 1024)).is_err());
328 }
329
330 #[test]
331 fn test_pbkdf2_derivation() {
332 let password = b"test password";
333 let salt = b"random salt here";
334 let iterations = PBKDF2_MIN_ITERATIONS;
335
336 let key1 = derive_key_pbkdf2(password, salt, iterations).unwrap();
337 let key2 = derive_key_pbkdf2(password, salt, iterations).unwrap();
338
339 assert_eq!(key1.as_bytes(), key2.as_bytes());
341 }
342
343 #[test]
344 fn test_hkdf_derivation() {
345 let ikm = b"input key material";
346 let salt = b"optional salt";
347 let info = b"context info";
348
349 let key1 = derive_key_hkdf(ikm, Some(salt), info).unwrap();
350 let key2 = derive_key_hkdf(ikm, Some(salt), info).unwrap();
351
352 assert_eq!(key1.as_bytes(), key2.as_bytes());
354 }
355
356 #[test]
357 fn test_hkdf_raw_rfc5869_case_1() {
358 let ikm = [0x0b_u8; 22];
359 let salt = [
360 0x00_u8, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c,
361 ];
362 let info = [
363 0xf0_u8, 0xf1, 0xf2, 0xf3, 0xf4, 0xf5, 0xf6, 0xf7, 0xf8, 0xf9,
364 ];
365 let okm = derive_key_hkdf_raw(&ikm, Some(&salt), &info, 42).unwrap();
366
367 let expected = hex::decode(
368 "3cb25f25faacd57a90434f64d0362f2a\
369 2d2d0a90cf1a5a4c5db02d56ecc4c5bf\
370 34007208d5b887185865",
371 )
372 .unwrap();
373
374 assert_eq!(okm, expected);
375 }
376
377 #[test]
378 fn test_hkdf_rejects_oversized_output_before_allocation() {
379 let error =
380 derive_key_hkdf_raw(b"ikm", None, b"info", HKDF_SHA256_MAX_OUTPUT + 1).unwrap_err();
381 assert!(matches!(error, Error::KeyDerivationFailed(_)));
382 }
383
384 #[test]
385 fn test_pbkdf2_rejects_unsafe_parameters() {
386 assert!(derive_key_pbkdf2(b"password", &[0u8; 16], 0).is_err());
387 assert!(derive_key_pbkdf2(b"password", &[0u8; 16], PBKDF2_MAX_ITERATIONS + 1).is_err());
388 assert!(derive_key_pbkdf2(
389 b"password",
390 &[0u8; PBKDF2_MIN_SALT_SIZE - 1],
391 PBKDF2_MIN_ITERATIONS
392 )
393 .is_err());
394 }
395
396 #[test]
397 fn test_x25519_deterministic_generation_from_seed() {
398 let seed = [7_u8; X25519_KEY_SIZE];
399 let a = generate_x25519_key_pair(Some(&seed)).unwrap();
400 let b = generate_x25519_key_pair(Some(&seed)).unwrap();
401
402 assert_eq!(a.private_key, b.private_key);
403 assert_eq!(a.public_key, b.public_key);
404 }
405
406 #[test]
407 fn test_x25519_shared_secret_symmetry() {
408 let alice = generate_x25519_key_pair(None).unwrap();
409 let bob = generate_x25519_key_pair(None).unwrap();
410
411 let s1 = x25519_shared_secret(&alice.private_key, &bob.public_key).unwrap();
412 let s2 = x25519_shared_secret(&bob.private_key, &alice.public_key).unwrap();
413
414 assert_eq!(s1, s2);
415 }
416
417 #[test]
418 fn test_x25519_rejects_low_order_public_keys() {
419 let private_key = [7_u8; X25519_KEY_SIZE];
420 let mut one = [0_u8; X25519_KEY_SIZE];
421 one[0] = 1;
422
423 for public_key in [[0_u8; X25519_KEY_SIZE], one] {
424 let error = x25519_shared_secret(&private_key, &public_key).unwrap_err();
425 assert_eq!(
426 error,
427 Error::KeyDerivationFailed(
428 "X25519 agreement rejected an all-zero shared secret from a low-order or invalid public key"
429 .to_string()
430 )
431 );
432 }
433 }
434
435 #[test]
436 fn test_derive_key_from_shared_secret_is_deterministic() {
437 let shared = [0x42_u8; X25519_KEY_SIZE];
438 let key1 =
439 derive_key_from_shared_secret(&shared, "voided-transfer-v1", "key-transfer").unwrap();
440 let key2 =
441 derive_key_from_shared_secret(&shared, "voided-transfer-v1", "key-transfer").unwrap();
442
443 assert_eq!(key1.as_bytes(), key2.as_bytes());
444 }
445
446 #[test]
447 fn test_derive_key_from_shared_secret_rejects_invalid_material() {
448 assert!(derive_key_from_shared_secret(&[0u8; 32], "salt", "info").is_err());
449 assert!(derive_key_from_shared_secret(&[1u8; 31], "salt", "info").is_err());
450 assert!(derive_key_from_shared_secret(&[1u8; 33], "salt", "info").is_err());
451 }
452}