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 let bytes = STANDARD.decode(encoded)?;
45 Self::from_bytes(&bytes)
46 }
47}
48
49impl AsRef<[u8]> for Key {
50 fn as_ref(&self) -> &[u8] {
51 &self.0
52 }
53}
54
55impl core::fmt::Debug for Key {
56 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
57 f.debug_struct("Key")
58 .field("length", &Self::SIZE)
59 .finish_non_exhaustive()
60 }
61}
62
63pub const X25519_KEY_SIZE: usize = 32;
65
66#[derive(Clone, Zeroize, ZeroizeOnDrop)]
68pub struct X25519KeyPair {
69 pub public_key: [u8; X25519_KEY_SIZE],
71 pub private_key: [u8; X25519_KEY_SIZE],
73}
74
75pub fn generate_key() -> Key {
77 let mut key = [0u8; 32];
78 rand::thread_rng().fill_bytes(&mut key);
79 Key(key)
80}
81
82pub fn derive_key_hkdf(input_key_material: &[u8], salt: Option<&[u8]>, info: &[u8]) -> Result<Key> {
94 let mut okm = derive_key_hkdf_raw(input_key_material, salt, info, Key::SIZE)?;
95 let key = Key::from_bytes(&okm);
96 okm.zeroize();
97 key
98}
99
100pub fn derive_key_hkdf_raw(
109 input_key_material: &[u8],
110 salt: Option<&[u8]>,
111 info: &[u8],
112 length: usize,
113) -> Result<Vec<u8>> {
114 use hkdf::Hkdf;
115 use sha2::Sha256;
116
117 if length == 0 {
118 return Err(Error::KeyDerivationFailed(
119 "HKDF output length must be > 0".to_string(),
120 ));
121 }
122
123 let hk = Hkdf::<Sha256>::new(salt, input_key_material);
124 let mut okm = vec![0u8; length];
125
126 hk.expand(info, &mut okm)
127 .map_err(|e| Error::KeyDerivationFailed(e.to_string()))?;
128
129 Ok(okm)
130}
131
132pub fn derive_key_pbkdf2(password: &[u8], salt: &[u8], iterations: u32) -> Result<Key> {
144 use pbkdf2::pbkdf2_hmac;
145 use sha2::Sha256;
146
147 let mut key = [0u8; 32];
148 pbkdf2_hmac::<Sha256>(password, salt, iterations, &mut key);
149
150 Ok(Key(key))
151}
152
153pub fn generate_x25519_key_pair(seed: Option<&[u8]>) -> Result<X25519KeyPair> {
157 let mut private_key = [0u8; X25519_KEY_SIZE];
158
159 if let Some(seed_bytes) = seed {
160 if seed_bytes.len() != X25519_KEY_SIZE {
161 return Err(Error::InvalidKeyLength {
162 expected: X25519_KEY_SIZE,
163 actual: seed_bytes.len(),
164 });
165 }
166 private_key.copy_from_slice(seed_bytes);
167 } else {
168 rand::thread_rng().fill_bytes(&mut private_key);
169 }
170
171 let public_key = x25519(private_key, X25519_BASEPOINT_BYTES);
172 Ok(X25519KeyPair {
173 public_key,
174 private_key,
175 })
176}
177
178pub fn x25519_shared_secret(
180 our_private_key: &[u8],
181 their_public_key: &[u8],
182) -> Result<[u8; X25519_KEY_SIZE]> {
183 if our_private_key.len() != X25519_KEY_SIZE {
184 return Err(Error::InvalidKeyLength {
185 expected: X25519_KEY_SIZE,
186 actual: our_private_key.len(),
187 });
188 }
189 if their_public_key.len() != X25519_KEY_SIZE {
190 return Err(Error::InvalidKeyLength {
191 expected: X25519_KEY_SIZE,
192 actual: their_public_key.len(),
193 });
194 }
195
196 let mut private_key = [0u8; X25519_KEY_SIZE];
197 private_key.copy_from_slice(our_private_key);
198
199 let mut public_key = [0u8; X25519_KEY_SIZE];
200 public_key.copy_from_slice(their_public_key);
201
202 Ok(x25519(private_key, public_key))
203}
204
205pub fn derive_key_from_shared_secret(shared_secret: &[u8], salt: &str, info: &str) -> Result<Key> {
207 derive_key_hkdf(shared_secret, Some(salt.as_bytes()), info.as_bytes())
208}
209
210#[allow(dead_code)]
212pub fn generate_salt(length: usize) -> Vec<u8> {
213 let mut salt = vec![0u8; length];
214 rand::thread_rng().fill_bytes(&mut salt);
215 salt
216}
217
218#[cfg(test)]
219mod tests {
220 use super::*;
221
222 #[test]
223 fn test_key_generation() {
224 let key1 = generate_key();
225 let key2 = generate_key();
226
227 assert_ne!(key1.as_bytes(), key2.as_bytes());
229
230 assert_eq!(key1.as_bytes().len(), 32);
232 }
233
234 #[test]
235 fn test_key_base64_roundtrip() {
236 let key = generate_key();
237 let encoded = key.to_base64();
238 let decoded = Key::from_base64(&encoded).unwrap();
239
240 assert_eq!(key.as_bytes(), decoded.as_bytes());
241 }
242
243 #[test]
244 fn test_pbkdf2_derivation() {
245 let password = b"test password";
246 let salt = b"random salt here";
247 let iterations = 1000; let key1 = derive_key_pbkdf2(password, salt, iterations).unwrap();
250 let key2 = derive_key_pbkdf2(password, salt, iterations).unwrap();
251
252 assert_eq!(key1.as_bytes(), key2.as_bytes());
254 }
255
256 #[test]
257 fn test_hkdf_derivation() {
258 let ikm = b"input key material";
259 let salt = b"optional salt";
260 let info = b"context info";
261
262 let key1 = derive_key_hkdf(ikm, Some(salt), info).unwrap();
263 let key2 = derive_key_hkdf(ikm, Some(salt), info).unwrap();
264
265 assert_eq!(key1.as_bytes(), key2.as_bytes());
267 }
268
269 #[test]
270 fn test_hkdf_raw_rfc5869_case_1() {
271 let ikm = [0x0b_u8; 22];
272 let salt = [
273 0x00_u8, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c,
274 ];
275 let info = [
276 0xf0_u8, 0xf1, 0xf2, 0xf3, 0xf4, 0xf5, 0xf6, 0xf7, 0xf8, 0xf9,
277 ];
278 let okm = derive_key_hkdf_raw(&ikm, Some(&salt), &info, 42).unwrap();
279
280 let expected = hex::decode(
281 "3cb25f25faacd57a90434f64d0362f2a\
282 2d2d0a90cf1a5a4c5db02d56ecc4c5bf\
283 34007208d5b887185865",
284 )
285 .unwrap();
286
287 assert_eq!(okm, expected);
288 }
289
290 #[test]
291 fn test_x25519_deterministic_generation_from_seed() {
292 let seed = [7_u8; X25519_KEY_SIZE];
293 let a = generate_x25519_key_pair(Some(&seed)).unwrap();
294 let b = generate_x25519_key_pair(Some(&seed)).unwrap();
295
296 assert_eq!(a.private_key, b.private_key);
297 assert_eq!(a.public_key, b.public_key);
298 }
299
300 #[test]
301 fn test_x25519_shared_secret_symmetry() {
302 let alice = generate_x25519_key_pair(None).unwrap();
303 let bob = generate_x25519_key_pair(None).unwrap();
304
305 let s1 = x25519_shared_secret(&alice.private_key, &bob.public_key).unwrap();
306 let s2 = x25519_shared_secret(&bob.private_key, &alice.public_key).unwrap();
307
308 assert_eq!(s1, s2);
309 }
310
311 #[test]
312 fn test_derive_key_from_shared_secret_is_deterministic() {
313 let shared = [0x42_u8; X25519_KEY_SIZE];
314 let key1 =
315 derive_key_from_shared_secret(&shared, "voided-transfer-v1", "key-transfer").unwrap();
316 let key2 =
317 derive_key_from_shared_secret(&shared, "voided-transfer-v1", "key-transfer").unwrap();
318
319 assert_eq!(key1.as_bytes(), key2.as_bytes());
320 }
321}