1#[cfg(feature = "alloc")]
67use alloc::vec::Vec;
68use core::marker::PhantomData;
69
70#[cfg(feature = "serde")]
71use serde::{Deserialize, Serialize};
72use zeroize::{Zeroize, ZeroizeOnDrop};
73
74use crate::classic::crypto_kdf::{
75 crypto_kdf_hkdf_sha256_expand, crypto_kdf_hkdf_sha256_extract, crypto_kdf_hkdf_sha512_expand,
76 crypto_kdf_hkdf_sha512_extract,
77};
78use crate::constants::{
79 CRYPTO_KDF_HKDF_SHA256_BYTES_MAX, CRYPTO_KDF_HKDF_SHA256_BYTES_MIN,
80 CRYPTO_KDF_HKDF_SHA256_KEYBYTES, CRYPTO_KDF_HKDF_SHA512_BYTES_MAX,
81 CRYPTO_KDF_HKDF_SHA512_BYTES_MIN, CRYPTO_KDF_HKDF_SHA512_KEYBYTES,
82};
83use crate::error::Error;
84use crate::types::*;
85
86pub type HkdfSha256Prk = StackByteArray<CRYPTO_KDF_HKDF_SHA256_KEYBYTES>;
88pub type HkdfSha512Prk = StackByteArray<CRYPTO_KDF_HKDF_SHA512_KEYBYTES>;
90pub type HkdfSha256 = Hkdf<HkdfSha256Variant, HkdfSha256Prk, CRYPTO_KDF_HKDF_SHA256_KEYBYTES>;
92pub type HkdfSha512 = Hkdf<HkdfSha512Variant, HkdfSha512Prk, CRYPTO_KDF_HKDF_SHA512_KEYBYTES>;
94
95#[derive(Zeroize, Clone, Debug)]
96#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
97pub struct Hkdf<Variant, Prk, const PRK_LENGTH: usize>
99where
100 Variant: HkdfVariant<PRK_LENGTH>,
101 Prk: ByteArray<PRK_LENGTH> + Zeroize + ZeroizeOnDrop,
102{
103 prk: Prk,
104 _variant: PhantomData<Variant>,
105}
106
107pub type HkdfSha256Expander<Prk> = Hkdf<HkdfSha256Variant, Prk, CRYPTO_KDF_HKDF_SHA256_KEYBYTES>;
109pub type HkdfSha512Expander<Prk> = Hkdf<HkdfSha512Variant, Prk, CRYPTO_KDF_HKDF_SHA512_KEYBYTES>;
111
112#[derive(Clone, Copy, Debug, Default)]
114pub struct HkdfSha256Variant;
115#[derive(Clone, Copy, Debug, Default)]
117pub struct HkdfSha512Variant;
118
119#[cfg(any(
120 all(feature = "protected", any(unix, windows)),
121 all(doc, not(doctest), feature = "std")
122))]
123#[cfg_attr(all(feature = "nightly", doc), doc(cfg(feature = "protected")))]
124pub mod protected {
125 use super::*;
142 pub use crate::protected::*;
143
144 pub type HkdfSha256Prk = HeapByteArray<CRYPTO_KDF_HKDF_SHA256_KEYBYTES>;
146 pub type HkdfSha512Prk = HeapByteArray<CRYPTO_KDF_HKDF_SHA512_KEYBYTES>;
148
149 pub type LockedHkdfSha256 = HkdfSha256Expander<Locked<HkdfSha256Prk>>;
151 pub type LockedHkdfSha512 = HkdfSha512Expander<Locked<HkdfSha512Prk>>;
153}
154
155mod sealed {
156 use crate::error::Error;
157
158 pub trait Sealed<const PRK_LENGTH: usize> {
161 const OUTPUT_BYTES_MIN: usize;
163 const OUTPUT_BYTES_MAX: usize;
165
166 fn extract(prk: &mut [u8; PRK_LENGTH], salt: Option<&[u8]>, ikm: &[u8]);
168 fn expand(output: &mut [u8], context: &[u8], prk: &[u8; PRK_LENGTH]) -> Result<(), Error>;
171
172 fn validate_output_len(output_len: usize) -> Result<(), Error> {
176 if output_len < Self::OUTPUT_BYTES_MIN || output_len > Self::OUTPUT_BYTES_MAX {
177 Err(length_error!(
178 crate::ErrorContext::Output,
179 output_len,
180 range Self::OUTPUT_BYTES_MIN,
181 Self::OUTPUT_BYTES_MAX
182 ))
183 } else {
184 Ok(())
185 }
186 }
187 }
188}
189
190pub trait HkdfVariant<const PRK_LENGTH: usize>: sealed::Sealed<PRK_LENGTH> {}
197
198macro_rules! impl_hkdf_variant {
199 ($variant:ty, $prk_len:expr, $bytes_min:expr, $bytes_max:expr, $extract:path, $expand:path) => {
200 impl HkdfVariant<$prk_len> for $variant {}
201
202 impl sealed::Sealed<$prk_len> for $variant {
203 const OUTPUT_BYTES_MAX: usize = $bytes_max;
204 const OUTPUT_BYTES_MIN: usize = $bytes_min;
205
206 fn extract(prk: &mut [u8; $prk_len], salt: Option<&[u8]>, ikm: &[u8]) {
207 $extract(prk, salt, ikm);
208 }
209
210 fn expand(
211 output: &mut [u8],
212 context: &[u8],
213 prk: &[u8; $prk_len],
214 ) -> Result<(), Error> {
215 $expand(output, context, prk)
216 }
217 }
218 };
219}
220
221impl_hkdf_variant!(
222 HkdfSha256Variant,
223 CRYPTO_KDF_HKDF_SHA256_KEYBYTES,
224 CRYPTO_KDF_HKDF_SHA256_BYTES_MIN,
225 CRYPTO_KDF_HKDF_SHA256_BYTES_MAX,
226 crypto_kdf_hkdf_sha256_extract,
227 crypto_kdf_hkdf_sha256_expand
228);
229
230impl_hkdf_variant!(
231 HkdfSha512Variant,
232 CRYPTO_KDF_HKDF_SHA512_KEYBYTES,
233 CRYPTO_KDF_HKDF_SHA512_BYTES_MIN,
234 CRYPTO_KDF_HKDF_SHA512_BYTES_MAX,
235 crypto_kdf_hkdf_sha512_extract,
236 crypto_kdf_hkdf_sha512_expand
237);
238
239impl<Variant, Prk, const PRK_LENGTH: usize> Hkdf<Variant, Prk, PRK_LENGTH>
240where
241 Variant: HkdfVariant<PRK_LENGTH>,
242 Prk: NewByteArray<PRK_LENGTH> + Zeroize + ZeroizeOnDrop,
243{
244 #[must_use]
246 pub fn generate() -> Self {
247 Self {
248 prk: Prk::generate(),
249 _variant: PhantomData,
250 }
251 }
252
253 #[must_use]
255 pub fn extract<Ikm: Bytes + ?Sized>(salt: Option<&[u8]>, ikm: &Ikm) -> Self {
256 let mut prk = Prk::new_byte_array();
257 Variant::extract(prk.as_mut_array(), salt, ikm.as_slice());
258 Self {
259 prk,
260 _variant: PhantomData,
261 }
262 }
263
264 pub fn extract_and_expand<
271 const OUTPUT_LENGTH: usize,
272 Output: NewByteArray<OUTPUT_LENGTH>,
273 Ikm: Bytes + ?Sized,
274 Context: Bytes + ?Sized,
275 >(
276 salt: Option<&[u8]>,
277 ikm: &Ikm,
278 context: &Context,
279 ) -> Result<Output, Error> {
280 Self::extract(salt, ikm).expand(context)
281 }
282
283 #[cfg(feature = "alloc")]
290 pub fn extract_and_expand_to_vec<Ikm: Bytes + ?Sized, Context: Bytes + ?Sized>(
291 salt: Option<&[u8]>,
292 ikm: &Ikm,
293 context: &Context,
294 output_len: usize,
295 ) -> Result<Vec<u8>, Error> {
296 Self::extract(salt, ikm).expand_to_vec(context, output_len)
297 }
298
299 pub fn extract_and_expand_to_bytes<
307 Output: NewBytes + ResizableBytes,
308 Ikm: Bytes + ?Sized,
309 Context: Bytes + ?Sized,
310 >(
311 salt: Option<&[u8]>,
312 ikm: &Ikm,
313 context: &Context,
314 output_len: usize,
315 ) -> Result<Output, Error> {
316 Self::extract(salt, ikm).expand_to_bytes(context, output_len)
317 }
318}
319
320impl<Variant, Prk, const PRK_LENGTH: usize> Hkdf<Variant, Prk, PRK_LENGTH>
321where
322 Variant: HkdfVariant<PRK_LENGTH>,
323 Prk: ByteArray<PRK_LENGTH> + Zeroize + ZeroizeOnDrop,
324{
325 #[must_use]
327 pub fn from_prk(prk: Prk) -> Self {
328 Self {
329 prk,
330 _variant: PhantomData,
331 }
332 }
333
334 #[must_use]
336 pub fn into_prk(self) -> Prk {
337 self.prk
338 }
339
340 pub fn expand<const OUTPUT_LENGTH: usize, Output, Context: Bytes + ?Sized>(
347 &self,
348 context: &Context,
349 ) -> Result<Output, Error>
350 where
351 Output: NewByteArray<OUTPUT_LENGTH>,
352 {
353 Variant::validate_output_len(OUTPUT_LENGTH)?;
354 let mut output = Output::new_byte_array();
355 Variant::expand(
356 output.as_mut_slice(),
357 context.as_slice(),
358 self.prk.as_array(),
359 )?;
360 Ok(output)
361 }
362
363 #[cfg(feature = "alloc")]
370 pub fn expand_to_vec<Context: Bytes + ?Sized>(
371 &self,
372 context: &Context,
373 output_len: usize,
374 ) -> Result<Vec<u8>, Error> {
375 self.expand_to_bytes(context, output_len)
376 }
377
378 pub fn expand_to_bytes<Output: NewBytes + ResizableBytes, Context: Bytes + ?Sized>(
386 &self,
387 context: &Context,
388 output_len: usize,
389 ) -> Result<Output, Error> {
390 Variant::validate_output_len(output_len)?;
391 let mut output = Output::new_bytes();
392 output.resize(output_len, 0);
393 Variant::expand(
394 output.as_mut_slice(),
395 context.as_slice(),
396 self.prk.as_array(),
397 )?;
398 Ok(output)
399 }
400}
401
402#[cfg(all(test, feature = "alloc"))]
403mod tests {
404 use super::*;
405 use crate::utils::test_util::hex as decode;
406
407 struct Case {
409 salt: Option<Vec<u8>>,
410 ikm: Vec<u8>,
411 info: Vec<u8>,
412 prk: Vec<u8>,
413 okm: Vec<u8>,
414 }
415
416 fn sha256_cases() -> [Case; 2] {
418 [
419 Case {
420 salt: Some(decode("000102030405060708090a0b0c")),
421 ikm: vec![0x0b; 22],
422 info: decode("f0f1f2f3f4f5f6f7f8f9"),
423 prk: decode("077709362c2e32df0ddc3f0dc47bba6390b6c73bb50f9c3122ec844ad7c2b3e5"),
424 okm: decode(concat!(
425 "3cb25f25faacd57a90434f64d0362f2a2d2d0a90cf1a5a4c5db02d56ecc4c5bf",
426 "34007208d5b887185865",
427 )),
428 },
429 Case {
430 salt: None,
431 ikm: vec![0x0b; 22],
432 info: Vec::new(),
433 prk: decode("19ef24a32c717b167f33a91d6f648bdf96596776afdb6377ac434c1c293ccb04"),
434 okm: decode(concat!(
435 "8da4e775a563c18f715f802a063c5a31b8a11f5c5ee1879ec3454e5f3c738d2d",
436 "9d201395faa4b61a96c8",
437 )),
438 },
439 ]
440 }
441
442 fn sha512_case() -> Case {
449 Case {
450 salt: Some(decode("000102030405060708090a0b0c")),
451 ikm: vec![0x0b; 22],
452 info: decode("f0f1f2f3f4f5f6f7f8f9"),
453 prk: decode(concat!(
454 "665799823737ded04a88e47e54a5890bb2c3d247c7a4254a8e61350723590a26",
455 "c36238127d8661b88cf80ef802d57e2f7cebcf1e00e083848be19929c61b4237",
456 )),
457 okm: decode(concat!(
458 "832390086cda71fb47625bb5ceb168e4c8e26a1a16ed34d9fc7fe92c14815793",
459 "38da362cb8d9f925d7cb",
460 )),
461 }
462 }
463
464 fn assert_case<Variant, const PRK_LENGTH: usize>(case: &Case)
465 where
466 Variant: HkdfVariant<PRK_LENGTH>,
467 {
468 type H<V, const P: usize> = Hkdf<V, StackByteArray<P>, P>;
469
470 let salt = case.salt.as_deref();
471 let hkdf = H::<Variant, PRK_LENGTH>::extract(salt, case.ikm.as_slice());
472 assert_eq!(hkdf.prk.as_slice(), case.prk.as_slice());
473
474 let okm_len = case.okm.len();
475 assert_eq!(
476 hkdf.expand_to_vec(case.info.as_slice(), okm_len)
477 .expect("expand"),
478 case.okm
479 );
480 let fixed: StackByteArray<42> = hkdf.expand(case.info.as_slice()).expect("expand");
481 assert_eq!(fixed.as_slice(), case.okm.as_slice());
482 let bytes: Vec<u8> = hkdf
483 .expand_to_bytes(case.info.as_slice(), okm_len)
484 .expect("expand");
485 assert_eq!(bytes, case.okm);
486
487 assert_eq!(
488 H::<Variant, PRK_LENGTH>::extract_and_expand_to_vec(
489 salt,
490 case.ikm.as_slice(),
491 case.info.as_slice(),
492 okm_len
493 )
494 .expect("expand"),
495 case.okm
496 );
497 let fixed: StackByteArray<42> = H::<Variant, PRK_LENGTH>::extract_and_expand(
498 salt,
499 case.ikm.as_slice(),
500 case.info.as_slice(),
501 )
502 .expect("expand");
503 assert_eq!(fixed.as_slice(), case.okm.as_slice());
504 let bytes: Vec<u8> = H::<Variant, PRK_LENGTH>::extract_and_expand_to_bytes(
505 salt,
506 case.ikm.as_slice(),
507 case.info.as_slice(),
508 okm_len,
509 )
510 .expect("expand");
511 assert_eq!(bytes, case.okm);
512
513 let prk = hkdf.into_prk();
515 assert_eq!(prk.as_slice(), case.prk.as_slice());
516 assert_eq!(
517 H::<Variant, PRK_LENGTH>::from_prk(prk)
518 .expand_to_vec(case.info.as_slice(), okm_len)
519 .expect("expand"),
520 case.okm
521 );
522
523 if case.salt.is_none() {
525 let empty_salt = H::<Variant, PRK_LENGTH>::extract(Some(&[][..]), case.ikm.as_slice());
526 assert_eq!(empty_salt.prk.as_slice(), case.prk.as_slice());
527 let zero_salt = H::<Variant, PRK_LENGTH>::extract(
528 Some(&[0u8; PRK_LENGTH][..]),
529 case.ikm.as_slice(),
530 );
531 assert_eq!(zero_salt.prk.as_slice(), case.prk.as_slice());
532 }
533
534 let short = H::<Variant, PRK_LENGTH>::extract(salt, case.ikm.as_slice())
536 .expand_to_vec(case.info.as_slice(), okm_len - 1)
537 .expect("expand");
538 assert_eq!(short, &case.okm[..okm_len - 1]);
539 let mut other_info = case.info.clone();
540 other_info.push(0);
541 assert_ne!(
542 H::<Variant, PRK_LENGTH>::extract(salt, case.ikm.as_slice())
543 .expand_to_vec(other_info.as_slice(), okm_len)
544 .expect("expand"),
545 case.okm
546 );
547 }
548
549 #[test]
550 fn rfc5869_sha256_vectors() {
551 for case in &sha256_cases() {
552 assert_case::<HkdfSha256Variant, CRYPTO_KDF_HKDF_SHA256_KEYBYTES>(case);
553 }
554 }
555
556 #[test]
557 fn sha512_a1_inputs_openssl_vector() {
558 assert_case::<HkdfSha512Variant, CRYPTO_KDF_HKDF_SHA512_KEYBYTES>(&sha512_case());
559 }
560
561 #[test]
562 fn output_length_limits_match_the_variant() {
563 let case = &sha256_cases()[0];
564 let hkdf = HkdfSha256::extract(case.salt.as_deref(), case.ikm.as_slice());
565 assert!(
566 hkdf.expand_to_vec(case.info.as_slice(), CRYPTO_KDF_HKDF_SHA256_BYTES_MIN)
567 .expect("min length")
568 .is_empty()
569 );
570 let empty: StackByteArray<0> = hkdf.expand(case.info.as_slice()).expect("min length");
571 assert!(empty.is_empty());
572
573 let max = hkdf
574 .expand_to_vec(case.info.as_slice(), CRYPTO_KDF_HKDF_SHA256_BYTES_MAX)
575 .expect("max length");
576 assert_eq!(max.len(), CRYPTO_KDF_HKDF_SHA256_BYTES_MAX);
577 assert_eq!(&max[..case.okm.len()], case.okm.as_slice());
578 let mut classic = vec![0u8; CRYPTO_KDF_HKDF_SHA256_BYTES_MAX];
579 crypto_kdf_hkdf_sha256_expand(&mut classic, case.info.as_slice(), hkdf.prk.as_array())
580 .expect("classic expand");
581 assert_eq!(max, classic);
582
583 for length in [CRYPTO_KDF_HKDF_SHA256_BYTES_MAX + 1, usize::MAX] {
584 assert!(matches!(
585 hkdf.expand_to_vec(case.info.as_slice(), length),
586 Err(Error::InvalidLength {
587 context: crate::ErrorContext::Output,
588 actual,
589 ..
590 }) if actual == length
591 ));
592 }
593 let too_long: Result<StackByteArray<{ CRYPTO_KDF_HKDF_SHA256_BYTES_MAX + 1 }>, Error> =
594 hkdf.expand(case.info.as_slice());
595 assert!(matches!(
596 too_long,
597 Err(Error::InvalidLength {
598 context: crate::ErrorContext::Output,
599 ..
600 })
601 ));
602
603 let hkdf512 = HkdfSha512::extract(case.salt.as_deref(), case.ikm.as_slice());
606 assert_eq!(
607 hkdf512
608 .expand_to_vec(case.info.as_slice(), CRYPTO_KDF_HKDF_SHA512_BYTES_MAX)
609 .expect("max length")
610 .len(),
611 CRYPTO_KDF_HKDF_SHA512_BYTES_MAX
612 );
613 assert!(
614 hkdf512
615 .expand_to_vec(case.info.as_slice(), CRYPTO_KDF_HKDF_SHA512_BYTES_MAX + 1)
616 .is_err()
617 );
618 }
619
620 #[test]
621 fn matches_classic_extract_and_expand() {
622 use crate::utils::test_util::XorShift64;
623
624 let mut rng = XorShift64::new(0x686b_6466_5f74_6573);
625 for round in 0..6 {
626 let ikm: Vec<u8> = (0..round * 13).map(|_| rng.next_u64() as u8).collect();
627 let salt: Vec<u8> = (0..round * 7).map(|_| rng.next_u64() as u8).collect();
628 let info: Vec<u8> = (0..round * 5).map(|_| rng.next_u64() as u8).collect();
629 let salt = (round % 2 == 0).then_some(salt.as_slice());
630
631 let mut prk256 = [0u8; CRYPTO_KDF_HKDF_SHA256_KEYBYTES];
632 crypto_kdf_hkdf_sha256_extract(&mut prk256, salt, &ikm);
633 let hkdf256 = HkdfSha256::extract(salt, ikm.as_slice());
634 assert_eq!(hkdf256.prk.as_array(), &prk256);
635
636 let mut prk512 = [0u8; CRYPTO_KDF_HKDF_SHA512_KEYBYTES];
637 crypto_kdf_hkdf_sha512_extract(&mut prk512, salt, &ikm);
638 let hkdf512 = HkdfSha512::extract(salt, ikm.as_slice());
639 assert_eq!(hkdf512.prk.as_array(), &prk512);
640
641 for length in [0, 1, 31, 32, 33, 63, 64, 65, 127, 128, 129] {
642 let mut classic = vec![0u8; length];
643 crypto_kdf_hkdf_sha256_expand(&mut classic, &info, &prk256).expect("expand");
644 assert_eq!(
645 hkdf256
646 .expand_to_vec(info.as_slice(), length)
647 .expect("expand"),
648 classic
649 );
650 crypto_kdf_hkdf_sha512_expand(&mut classic, &info, &prk512).expect("expand");
651 assert_eq!(
652 hkdf512
653 .expand_to_vec(info.as_slice(), length)
654 .expect("expand"),
655 classic
656 );
657 }
658 }
659 }
660
661 #[test]
662 fn generic_variant_api_reproduces_rfc5869() {
663 fn extract_and_expand_with_variant<Variant, const PRK_LENGTH: usize>(case: &Case) -> Vec<u8>
664 where
665 Variant: HkdfVariant<PRK_LENGTH>,
666 {
667 Hkdf::<Variant, StackByteArray<PRK_LENGTH>, PRK_LENGTH>::extract_and_expand_to_vec(
668 case.salt.as_deref(),
669 case.ikm.as_slice(),
670 case.info.as_slice(),
671 case.okm.len(),
672 )
673 .expect("expand failed")
674 }
675
676 let case256 = &sha256_cases()[0];
677 let case512 = sha512_case();
678 assert_eq!(
679 extract_and_expand_with_variant::<HkdfSha256Variant, CRYPTO_KDF_HKDF_SHA256_KEYBYTES>(
680 case256
681 ),
682 case256.okm
683 );
684 assert_eq!(
685 extract_and_expand_with_variant::<HkdfSha512Variant, CRYPTO_KDF_HKDF_SHA512_KEYBYTES>(
686 &case512
687 ),
688 case512.okm
689 );
690 assert_ne!(case256.okm, case512.okm);
692 }
693
694 #[cfg(feature = "serde")]
695 #[test]
696 fn serde_round_trip_expands_to_the_rfc5869_output() {
697 let case = &sha256_cases()[0];
698 let hkdf = HkdfSha256::extract(case.salt.as_deref(), case.ikm.as_slice());
699 let json = serde_json::to_string(&hkdf).expect("serialize");
700 let decoded: HkdfSha256 = serde_json::from_str(&json).expect("deserialize");
701 assert_eq!(
702 decoded
703 .expand_to_vec(case.info.as_slice(), case.okm.len())
704 .expect("expand"),
705 case.okm
706 );
707
708 let case = sha512_case();
709 let hkdf = HkdfSha512::extract(case.salt.as_deref(), case.ikm.as_slice());
710 let json = serde_json::to_string(&hkdf).expect("serialize");
711 let decoded: HkdfSha512 = serde_json::from_str(&json).expect("deserialize");
712 assert_eq!(decoded.into_prk().as_slice(), case.prk.as_slice());
713 }
714
715 #[cfg(all(feature = "protected", any(unix, windows)))]
716 #[test]
717 fn locked_expanders_reproduce_rfc5869() {
718 use crate::hkdf::protected::*;
719
720 let case = &sha256_cases()[0];
721 let ikm = HeapBytes::from_slice_into_readonly_locked(&case.ikm).expect("lock ikm");
722 let salt = case
723 .salt
724 .as_ref()
725 .map(|salt| HeapBytes::from_slice_into_readonly_locked(salt).expect("lock salt"));
726 let hkdf: LockedHkdfSha256 =
727 HkdfSha256Expander::extract(salt.as_ref().map(|salt| salt.as_slice()), &ikm);
728 assert_eq!(hkdf.prk.as_slice(), case.prk.as_slice());
729 let okm: Locked<HeapBytes> = hkdf
730 .expand_to_bytes(case.info.as_slice(), case.okm.len())
731 .expect("expand");
732 assert_eq!(okm.as_slice(), case.okm.as_slice());
733 let fixed: Locked<HeapByteArray<42>> = hkdf.expand(case.info.as_slice()).expect("expand");
734 assert_eq!(fixed.as_slice(), case.okm.as_slice());
735
736 let case = sha512_case();
737 let ikm = HeapBytes::from_slice_into_readonly_locked(&case.ikm).expect("lock ikm");
738 let hkdf: LockedHkdfSha512 = HkdfSha512Expander::extract(case.salt.as_deref(), &ikm);
739 assert_eq!(
740 hkdf.expand_to_vec(case.info.as_slice(), case.okm.len())
741 .expect("expand"),
742 case.okm
743 );
744 }
745}