Skip to main content

crypto/
hkdf.rs

1//! HMAC (Hash-based Message Authentication Code)
2
3use crate::{Hash, Hasher, HkdfError, hmac::Hmac};
4
5const DEFAULT_SALT: [u8; 64] = [0u8; 64];
6
7/// HKDF extract step: `PRK = HMAC-Hash(salt, IKM)`.
8///
9/// If `salt` is `None`, a string of `H::OUTPUT_SIZE` zero bytes is used.
10///
11/// # Example
12///
13/// ```ignore
14/// use crypto::hkdf;
15/// use crypto::sha2::Sha256;
16///
17/// let prk = hkdf::extract::<Sha256>(Some(b"salt"), b"input key material");
18/// ```
19pub fn extract<H: Hasher>(salt: Option<&[u8]>, ikm: &[u8]) -> Hash {
20    let salt = salt.unwrap_or(&DEFAULT_SALT[..H::OUTPUT_SIZE]);
21    let mut mac = Hmac::<H>::new(salt);
22    mac.update(ikm);
23    return mac.finalize();
24}
25
26/// HKDF expand step: `OKM = T(1) || T(2) || ...`, where
27/// `T(i) = HMAC-Hash(PRK, T(i-1) || info || i)`.
28///
29/// The output is written into `okm`. The length of the output is determined by
30/// `okm.len()`.
31///
32/// # Example
33///
34/// ```ignore
35/// use crypto::hkdf;
36/// use crypto::sha2::Sha256;
37///
38/// let prk = hkdf::extract::<Sha256>(Some(b"salt"), b"input key material");
39/// let mut okm = [0u8; 32];
40/// hkdf::expand::<Sha256>(&mut okm, &prk, b"context info").unwrap();
41/// ```
42///
43/// # Error
44///
45/// Returns an error if `okm.len() > 255 * H::OUTPUT_SIZE` or if `prk.len() < H::OUTPUT_SIZE`.
46pub fn expand<H: Hasher>(okm: &mut [u8], prk: &[u8], info: &[u8]) -> Result<(), HkdfError> {
47    let n = okm.len();
48
49    if prk.len() < H::OUTPUT_SIZE {
50        return Err(HkdfError::PrkIsTooShort(H::OUTPUT_SIZE));
51    }
52
53    if n > 255 * H::OUTPUT_SIZE {
54        return Err(HkdfError::OutputIsTooLong);
55    }
56
57    if n == 0 {
58        return Ok(());
59    }
60
61    let mut t = [0u8; 64];
62    let mut t_len = 0usize;
63    let mut offset = 0usize;
64    let mut counter = 1u8;
65
66    while offset < n {
67        let mut mac = Hmac::<H>::new(&prk[..H::OUTPUT_SIZE]);
68        mac.update(&t[..t_len]);
69        mac.update(info);
70        mac.update(&[counter]);
71        let block = mac.finalize();
72        let block_bytes = block.as_ref();
73        let chunk_len = (n - offset).min(H::OUTPUT_SIZE);
74        okm[offset..offset + chunk_len].copy_from_slice(&block_bytes[..chunk_len]);
75        t[..H::OUTPUT_SIZE].copy_from_slice(block_bytes);
76        t_len = H::OUTPUT_SIZE;
77        offset += chunk_len;
78        counter = counter.wrapping_add(1);
79    }
80
81    Ok(())
82}
83
84/// One-shot HKDF: extract-then-expand in a single call.
85///
86/// # Example
87///
88/// ```ignore
89/// use crypto::hkdf;
90/// use crypto::sha2::Sha256;
91///
92/// let okm: [u8; 32] = hkdf::derive_key::<Sha256, 32>(
93///     b"input key material",
94///     b"context info",
95///     Some(b"salt"),
96/// ).unwrap();
97/// ```
98///
99/// # Error
100///
101/// Returns an error if `N > 255 * H::OUTPUT_SIZE`.
102pub fn derive_key<H: Hasher, const N: usize>(
103    ikm: &[u8],
104    info: &[u8],
105    salt: Option<&[u8]>,
106) -> Result<[u8; N], HkdfError> {
107    let prk = extract::<H>(salt, ikm);
108    let mut okm = [0u8; N];
109    expand::<H>(&mut okm, prk.as_ref(), info)?;
110    Ok(okm)
111}
112
113#[cfg(test)]
114mod tests {
115    use super::*;
116    use crate::sha2::{Sha256, Sha384, Sha512};
117
118    struct TestVector {
119        ikm: &'static str,
120        salt: Option<&'static str>,
121        info: &'static str,
122        expected_prk: &'static str,
123        expected_okm: &'static str,
124    }
125
126    fn decode_hex(input: &str) -> Vec<u8> {
127        let input = input.replace(|c: char| c.is_whitespace(), "");
128        (0..input.len())
129            .step_by(2)
130            .map(|i| u8::from_str_radix(&input[i..i + 2], 16).unwrap())
131            .collect()
132    }
133
134    const SHA256_VECTORS: [TestVector; 4] = [
135        TestVector {
136            ikm: "0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b",
137            salt: Some("000102030405060708090a0b0c"),
138            info: "f0f1f2f3f4f5f6f7f8f9",
139            expected_prk: "077709362c2e32df0ddc3f0dc47bba6390b6c73bb50f9c3122ec844ad7c2b3e5",
140            expected_okm: "3cb25f25faacd57a90434f64d0362f2a2d2d0a90cf1a5a4c5db02d56ecc4c5bf34007208d5b887185865",
141        },
142        TestVector {
143            ikm: "000102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f\
144                  202122232425262728292a2b2c2d2e2f303132333435363738393a3b3c3d3e3f\
145                  404142434445464748494a4b4c4d4e4f",
146            salt: Some(
147                "606162636465666768696a6b6c6d6e6f707172737475767778797a7b7c7d7e7f\
148                 808182838485868788898a8b8c8d8e8f909192939495969798999a9b9c9d9e9f\
149                 a0a1a2a3a4a5a6a7a8a9aaabacadaeaf",
150            ),
151            info: "b0b1b2b3b4b5b6b7b8b9babbbcbdbebfc0c1c2c3c4c5c6c7c8c9cacbcccdcecf\
152                  d0d1d2d3d4d5d6d7d8d9dadbdcdddedfe0e1e2e3e4e5e6e7e8e9eaebecedeeef\
153                  f0f1f2f3f4f5f6f7f8f9fafbfcfdfeff",
154            expected_prk: "06a6b88c5853361a06104c9ceb35b45cef760014904671014a193f40c15fc244",
155            expected_okm: "b11e398dc80327a1c8e7f78c596a49344f012eda2d4efad8a050cc4c19afa97c59045a99cac7827271cb41c65e590e09da3275600c2f09b8367793a9aca3db71cc30c58179ec3e87c14c01d5c1f3434f1d87",
156        },
157        TestVector {
158            ikm: "0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b",
159            salt: Some(""),
160            info: "",
161            expected_prk: "19ef24a32c717b167f33a91d6f648bdf96596776afdb6377ac434c1c293ccb04",
162            expected_okm: "8da4e775a563c18f715f802a063c5a31b8a11f5c5ee1879ec3454e5f3c738d2d9d201395faa4b61a96c8",
163        },
164        TestVector {
165            ikm: "0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b",
166            salt: None,
167            info: "",
168            expected_prk: "19ef24a32c717b167f33a91d6f648bdf96596776afdb6377ac434c1c293ccb04",
169            expected_okm: "8da4e775a563c18f715f802a063c5a31b8a11f5c5ee1879ec3454e5f3c738d2d9d201395faa4b61a96c8",
170        },
171    ];
172
173    const SHA512_VECTORS: [TestVector; 4] = [
174        TestVector {
175            ikm: "0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b",
176            salt: Some("000102030405060708090a0b0c"),
177            info: "f0f1f2f3f4f5f6f7f8f9",
178            expected_prk: "665799823737ded04a88e47e54a5890bb2c3d247c7a4254a8e61350723590a26c36238127d8661b88cf80ef802d57e2f7cebcf1e00e083848be19929c61b4237",
179            expected_okm: "832390086cda71fb47625bb5ceb168e4c8e26a1a16ed34d9fc7fe92c1481579338da362cb8d9f925d7cb",
180        },
181        TestVector {
182            ikm: "000102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f\
183                  202122232425262728292a2b2c2d2e2f303132333435363738393a3b3c3d3e3f\
184                  404142434445464748494a4b4c4d4e4f",
185            salt: Some(
186                "606162636465666768696a6b6c6d6e6f707172737475767778797a7b7c7d7e7f\
187                 808182838485868788898a8b8c8d8e8f909192939495969798999a9b9c9d9e9f\
188                 a0a1a2a3a4a5a6a7a8a9aaabacadaeaf",
189            ),
190            info: "b0b1b2b3b4b5b6b7b8b9babbbcbdbebfc0c1c2c3c4c5c6c7c8c9cacbcccdcecf\
191                  d0d1d2d3d4d5d6d7d8d9dadbdcdddedfe0e1e2e3e4e5e6e7e8e9eaebecedeeef\
192                  f0f1f2f3f4f5f6f7f8f9fafbfcfdfeff",
193            expected_prk: "35672542907d4e142c00e84499e74e1de08be86535f924e022804ad775dde27ec86cd1e5b7d178c74489bdbeb30712beb82d4f97416c5a94ea81ebdf3e629e4a",
194            expected_okm: "ce6c97192805b346e6161e821ed165673b84f400a2b514b2fe23d84cd189ddf1b695b48cbd1c8388441137b3ce28f16aa64ba33ba466b24df6cfcb021ecff235f6a2056ce3af1de44d572097a8505d9e7a93",
195        },
196        TestVector {
197            ikm: "0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b",
198            salt: Some(""),
199            info: "",
200            expected_prk: "fd200c4987ac491313bd4a2a13287121247239e11c9ef82802044b66ef357e5b194498d0682611382348572a7b1611de54764094286320578a863f36562b0df6",
201            expected_okm: "f5fa02b18298a72a8c23898a8703472c6eb179dc204c03425c970e3b164bf90fff22d04836d0e2343bac",
202        },
203        TestVector {
204            ikm: "0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b0b",
205            salt: None,
206            info: "",
207            expected_prk: "fd200c4987ac491313bd4a2a13287121247239e11c9ef82802044b66ef357e5b194498d0682611382348572a7b1611de54764094286320578a863f36562b0df6",
208            expected_okm: "f5fa02b18298a72a8c23898a8703472c6eb179dc204c03425c970e3b164bf90fff22d04836d0e2343bac",
209        },
210    ];
211
212    #[test]
213    fn hkdf_sha256_vectors() {
214        for (i, vector) in SHA256_VECTORS.iter().enumerate() {
215            let ikm = decode_hex(vector.ikm);
216            let salt = vector.salt.map(decode_hex);
217            let info = decode_hex(vector.info);
218            let expected_prk = decode_hex(vector.expected_prk);
219            let expected_okm = decode_hex(vector.expected_okm);
220
221            let prk = extract::<Sha256>(salt.as_deref(), &ikm);
222            assert_eq!(prk.as_ref(), expected_prk.as_slice(), "vector {} PRK", i);
223
224            let mut buf = [0u8; 82];
225            match expected_okm.len() {
226                42 => expand::<Sha256>(&mut buf[..42], prk.as_ref(), &info).unwrap(),
227                82 => expand::<Sha256>(&mut buf[..82], prk.as_ref(), &info).unwrap(),
228                _ => unreachable!(),
229            };
230            assert_eq!(&buf[..expected_okm.len()], expected_okm, "vector {} OKM", i);
231
232            let derived = match expected_okm.len() {
233                42 => derive_key::<Sha256, 42>(&ikm, &info, salt.as_deref()).unwrap().to_vec(),
234                82 => derive_key::<Sha256, 82>(&ikm, &info, salt.as_deref()).unwrap().to_vec(),
235                _ => unreachable!(),
236            };
237            assert_eq!(derived, expected_okm, "vector {} derive_key OKM", i);
238        }
239    }
240
241    #[test]
242    fn hkdf_sha512_vectors() {
243        for (i, vector) in SHA512_VECTORS.iter().enumerate() {
244            let ikm = decode_hex(vector.ikm);
245            let salt = vector.salt.map(decode_hex);
246            let info = decode_hex(vector.info);
247            let expected_prk = decode_hex(vector.expected_prk);
248            let expected_okm = decode_hex(vector.expected_okm);
249
250            let prk = extract::<Sha512>(salt.as_deref(), &ikm);
251            assert_eq!(prk.as_ref(), expected_prk.as_slice(), "vector {} PRK", i);
252
253            let mut buf = [0u8; 82];
254            match expected_okm.len() {
255                42 => expand::<Sha512>(&mut buf[..42], prk.as_ref(), &info).unwrap(),
256                82 => expand::<Sha512>(&mut buf[..82], prk.as_ref(), &info).unwrap(),
257                _ => unreachable!(),
258            };
259            assert_eq!(&buf[..expected_okm.len()], expected_okm, "vector {} OKM", i);
260
261            let derived = match expected_okm.len() {
262                42 => derive_key::<Sha512, 42>(&ikm, &info, salt.as_deref()).unwrap().to_vec(),
263                82 => derive_key::<Sha512, 82>(&ikm, &info, salt.as_deref()).unwrap().to_vec(),
264                _ => unreachable!(),
265            };
266            assert_eq!(derived, expected_okm, "vector {} derive_key OKM", i);
267        }
268    }
269
270    #[test]
271    fn hkdf_zero_length_output() {
272        let prk = [0u8; 32];
273        let mut buf = [0u8; 0];
274        assert_eq!(expand::<Sha256>(&mut buf, &prk, b""), Ok(()));
275        assert_eq!(derive_key::<Sha256, 0>(b"ikm", b"info", None).unwrap(), [] as [u8; 0]);
276    }
277
278    #[test]
279    fn hkdf_expand_panics_when_output_is_too_large() {
280        let prk = [0u8; 32];
281        let mut buf = vec![0u8; Sha256::BLOCK_SIZE * 300];
282        assert_eq!(expand::<Sha256>(&mut buf, &prk, b""), Err(HkdfError::OutputIsTooLong));
283    }
284
285    #[test]
286    fn hkdf_expand_panics_when_prk_is_too_short() {
287        let mut buf = [0u8; 32];
288        assert_eq!(
289            expand::<Sha256>(&mut buf, &[0u8; 31], b""),
290            Err(HkdfError::PrkIsTooShort(Sha256::OUTPUT_SIZE))
291        );
292    }
293
294    // --- Wycheproof test vectors ---
295
296    #[test]
297    fn hkdf_sha256_wycheproof() {
298        // Maximum valid HKDF-SHA-256 output: 255 * 32 = 8160 bytes.
299        const MAX_OKM: usize = 8160;
300        const SIZE_TOO_LARGE: usize = 8161;
301
302        let data: serde_json::Value =
303            serde_json::from_str(include_str!("../testdata/wycheproof/testvectors_v1/hkdf_sha256_test.json")).unwrap();
304        let mut valid_tested = 0u64;
305        let mut invalid_tested = 0u64;
306        for group in data["testGroups"].as_array().unwrap() {
307            for test in group["tests"].as_array().unwrap() {
308                let ikm_hex = test["ikm"].as_str().unwrap();
309                let salt_hex = test["salt"].as_str().unwrap();
310                let info_hex = test["info"].as_str().unwrap();
311                let size = test["size"].as_u64().unwrap() as usize;
312                let expected_okm_hex = test["okm"].as_str().unwrap();
313                let result = test["result"].as_str().unwrap();
314
315                let ikm = hex::decode(ikm_hex).unwrap();
316                let info = hex::decode(info_hex).unwrap();
317                let salt: Option<Vec<u8>> = if salt_hex.is_empty() {
318                    None
319                } else {
320                    Some(hex::decode(salt_hex).unwrap())
321                };
322
323                if result == "valid" {
324                    let okm = derive_key::<Sha256, MAX_OKM>(&ikm, &info, salt.as_deref()).unwrap();
325                    let okm_hex = hex::encode(&okm[..size]);
326                    assert_eq!(
327                        okm_hex, expected_okm_hex,
328                        "wycheproof HKDF-SHA-256 tcId={} size={}",
329                        test["tcId"], size
330                    );
331                    valid_tested += 1;
332                } else {
333                    assert_eq!(
334                        derive_key::<Sha256, SIZE_TOO_LARGE>(&ikm, &info, salt.as_deref()),
335                        Err(HkdfError::OutputIsTooLong),
336                        "wycheproof HKDF-SHA-256 tcId={} size={} should reject",
337                        test["tcId"],
338                        size
339                    );
340                    invalid_tested += 1;
341                }
342            }
343        }
344        assert!(valid_tested > 0, "no valid HKDF-SHA-256 wycheproof tests were run");
345        assert!(invalid_tested > 0, "no invalid HKDF-SHA-256 wycheproof tests were run");
346    }
347
348    #[test]
349    fn hkdf_sha512_wycheproof() {
350        // Maximum valid HKDF-SHA-512 output: 255 * 64 = 16320 bytes.
351        const MAX_OKM: usize = 16320;
352        const SIZE_TOO_LARGE: usize = 16321;
353
354        let data: serde_json::Value =
355            serde_json::from_str(include_str!("../testdata/wycheproof/testvectors_v1/hkdf_sha512_test.json")).unwrap();
356        let mut valid_tested = 0u64;
357        let mut invalid_tested = 0u64;
358        for group in data["testGroups"].as_array().unwrap() {
359            for test in group["tests"].as_array().unwrap() {
360                let ikm_hex = test["ikm"].as_str().unwrap();
361                let salt_hex = test["salt"].as_str().unwrap();
362                let info_hex = test["info"].as_str().unwrap();
363                let size = test["size"].as_u64().unwrap() as usize;
364                let expected_okm_hex = test["okm"].as_str().unwrap();
365                let result = test["result"].as_str().unwrap();
366
367                let ikm = hex::decode(ikm_hex).unwrap();
368                let info = hex::decode(info_hex).unwrap();
369                let salt: Option<Vec<u8>> = if salt_hex.is_empty() {
370                    None
371                } else {
372                    Some(hex::decode(salt_hex).unwrap())
373                };
374
375                if result == "valid" {
376                    let okm = derive_key::<Sha512, MAX_OKM>(&ikm, &info, salt.as_deref()).unwrap();
377                    let okm_hex = hex::encode(&okm[..size]);
378                    assert_eq!(
379                        okm_hex, expected_okm_hex,
380                        "wycheproof HKDF-SHA-512 tcId={} size={}",
381                        test["tcId"], size
382                    );
383                    valid_tested += 1;
384                } else {
385                    assert_eq!(
386                        derive_key::<Sha512, SIZE_TOO_LARGE>(&ikm, &info, salt.as_deref()),
387                        Err(HkdfError::OutputIsTooLong),
388                        "wycheproof HKDF-SHA-512 tcId={} size={} should reject",
389                        test["tcId"],
390                        size
391                    );
392                    invalid_tested += 1;
393                }
394            }
395        }
396        assert!(valid_tested > 0, "no valid HKDF-SHA-512 wycheproof tests were run");
397        assert!(invalid_tested > 0, "no invalid HKDF-SHA-512 wycheproof tests were run");
398    }
399
400    #[test]
401    fn hkdf_sha384_wycheproof() {
402        // Maximum valid HKDF-SHA-384 output: 255 * 48 = 12240 bytes.
403        const MAX_OKM: usize = 12240;
404        const SIZE_TOO_LARGE: usize = 12241;
405
406        let data: serde_json::Value =
407            serde_json::from_str(include_str!("../testdata/wycheproof/testvectors_v1/hkdf_sha384_test.json")).unwrap();
408        let mut valid_tested = 0u64;
409        let mut invalid_tested = 0u64;
410        for group in data["testGroups"].as_array().unwrap() {
411            for test in group["tests"].as_array().unwrap() {
412                let ikm_hex = test["ikm"].as_str().unwrap();
413                let salt_hex = test["salt"].as_str().unwrap();
414                let info_hex = test["info"].as_str().unwrap();
415                let size = test["size"].as_u64().unwrap() as usize;
416                let expected_okm_hex = test["okm"].as_str().unwrap();
417                let result = test["result"].as_str().unwrap();
418
419                let ikm = hex::decode(ikm_hex).unwrap();
420                let info = hex::decode(info_hex).unwrap();
421                let salt: Option<Vec<u8>> = if salt_hex.is_empty() {
422                    None
423                } else {
424                    Some(hex::decode(salt_hex).unwrap())
425                };
426
427                if result == "valid" {
428                    let okm = derive_key::<Sha384, MAX_OKM>(&ikm, &info, salt.as_deref()).unwrap();
429                    let okm_hex = hex::encode(&okm[..size]);
430                    assert_eq!(
431                        okm_hex, expected_okm_hex,
432                        "wycheproof HKDF-SHA-384 tcId={} size={}",
433                        test["tcId"], size
434                    );
435                    valid_tested += 1;
436                } else {
437                    assert_eq!(
438                        derive_key::<Sha384, SIZE_TOO_LARGE>(&ikm, &info, salt.as_deref()),
439                        Err(HkdfError::OutputIsTooLong),
440                        "wycheproof HKDF-SHA-384 tcId={} size={} should reject",
441                        test["tcId"],
442                        size
443                    );
444                    invalid_tested += 1;
445                }
446            }
447        }
448        assert!(valid_tested > 0, "no valid HKDF-SHA-384 wycheproof tests were run");
449        assert!(invalid_tested > 0, "no invalid HKDF-SHA-384 wycheproof tests were run");
450    }
451}