Skip to main content

dryoc/
hkdf.rs

1//! # HKDF key derivation
2//!
3//! [`HkdfSha256`] and [`HkdfSha512`] provide Rustaceous wrappers around
4//! libsodium's HKDF-SHA-256 and HKDF-SHA-512 functions.
5//!
6//! HKDF turns input keying material into one or more independent keys. It has
7//! two steps:
8//!
9//! * extract: mix the input keying material with an optional salt to produce a
10//!   pseudorandom key (PRK)
11//! * expand: derive output bytes from that PRK and a context string
12//!
13//! Use HKDF when you already have keying material, such as a shared secret from
14//! key exchange, and need separate keys for different purposes. The context is
15//! public domain-separation data; changing it changes the derived output.
16//!
17//! # Rustaceous API example
18//!
19//! ```
20//! # #[cfg(feature = "alloc")]
21//! # {
22//! use dryoc::hkdf::{HkdfSha256, HkdfSha256Prk};
23//!
24//! let hkdf: HkdfSha256 =
25//!     HkdfSha256::extract(Some(b"Act IV salt"), b"Now is the winter of our discontent");
26//! let output: Vec<u8> = hkdf
27//!     .expand_to_vec(b"session key", 42)
28//!     .expect("expand failed");
29//! assert_eq!(output.len(), 42);
30//! # }
31//! ```
32//!
33//! # One-shot extract and expand
34//!
35//! ```
36//! # #[cfg(feature = "alloc")]
37//! # {
38//! use dryoc::hkdf::HkdfSha512;
39//!
40//! let output = HkdfSha512::extract_and_expand_to_vec(
41//!     Some(b"optional deployment salt"),
42//!     b"Our remedies oft in ourselves do lie",
43//!     b"application secret",
44//!     64,
45//! )
46//! .expect("expand failed");
47//! assert_eq!(output.len(), 64);
48//! # }
49//! ```
50//!
51//! # Reusing an extracted PRK
52//!
53//! ```
54//! use dryoc::hkdf::{HkdfSha256, HkdfSha256Prk};
55//!
56//! let hkdf = HkdfSha256::extract(Some(b"deployment salt"), b"We know what we are");
57//! let encryption_key: HkdfSha256Prk = hkdf.expand(b"encryption key").expect("expand failed");
58//! let authentication_key: HkdfSha256Prk =
59//!     hkdf.expand(b"authentication key").expect("expand failed");
60//! assert_ne!(encryption_key, authentication_key);
61//! ```
62//!
63//! The concrete expanders are type aliases over [`Hkdf`] and can also be used
64//! through [`HkdfVariant`] in generic code.
65
66#[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
86/// Stack-allocated pseudorandom key for HKDF-SHA-256.
87pub type HkdfSha256Prk = StackByteArray<CRYPTO_KDF_HKDF_SHA256_KEYBYTES>;
88/// Stack-allocated pseudorandom key for HKDF-SHA-512.
89pub type HkdfSha512Prk = StackByteArray<CRYPTO_KDF_HKDF_SHA512_KEYBYTES>;
90/// Stack-allocated HKDF-SHA-256 expander.
91pub type HkdfSha256 = Hkdf<HkdfSha256Variant, HkdfSha256Prk, CRYPTO_KDF_HKDF_SHA256_KEYBYTES>;
92/// Stack-allocated HKDF-SHA-512 expander.
93pub type HkdfSha512 = Hkdf<HkdfSha512Variant, HkdfSha512Prk, CRYPTO_KDF_HKDF_SHA512_KEYBYTES>;
94
95#[derive(Zeroize, Clone, Debug)]
96#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
97/// HKDF expander for a specific [`HkdfVariant`].
98pub 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
107/// HKDF-SHA-256 expander.
108pub type HkdfSha256Expander<Prk> = Hkdf<HkdfSha256Variant, Prk, CRYPTO_KDF_HKDF_SHA256_KEYBYTES>;
109/// HKDF-SHA-512 expander.
110pub type HkdfSha512Expander<Prk> = Hkdf<HkdfSha512Variant, Prk, CRYPTO_KDF_HKDF_SHA512_KEYBYTES>;
111
112/// HKDF-SHA-256 algorithm marker.
113#[derive(Clone, Copy, Debug, Default)]
114pub struct HkdfSha256Variant;
115/// HKDF-SHA-512 algorithm marker.
116#[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    //! # Protected memory type aliases for HKDF
126    //!
127    //! This mod provides protected-memory PRK aliases and locked HKDF aliases.
128    //! Use these aliases when the extracted PRK or expanded output should stay
129    //! in locked memory.
130    //!
131    //! ```
132    //! use dryoc::hkdf::HkdfSha512Expander;
133    //! use dryoc::hkdf::protected::*;
134    //!
135    //! let ikm = HeapBytes::from_slice_into_readonly_locked(b"Truth will come to light.")
136    //!     .expect("ikm failed");
137    //! let hkdf: LockedHkdfSha512 = HkdfSha512Expander::extract(None, &ikm);
138    //! let output: Locked<HeapBytes> = hkdf.expand_to_bytes(b"context", 64).expect("expand failed");
139    //! assert_eq!(output.len(), 64);
140    //! ```
141    use super::*;
142    pub use crate::protected::*;
143
144    /// Heap-allocated, page-aligned pseudorandom key for HKDF-SHA-256.
145    pub type HkdfSha256Prk = HeapByteArray<CRYPTO_KDF_HKDF_SHA256_KEYBYTES>;
146    /// Heap-allocated, page-aligned pseudorandom key for HKDF-SHA-512.
147    pub type HkdfSha512Prk = HeapByteArray<CRYPTO_KDF_HKDF_SHA512_KEYBYTES>;
148
149    /// Locked HKDF-SHA-256 expander.
150    pub type LockedHkdfSha256 = HkdfSha256Expander<Locked<HkdfSha256Prk>>;
151    /// Locked HKDF-SHA-512 expander.
152    pub type LockedHkdfSha512 = HkdfSha512Expander<Locked<HkdfSha512Prk>>;
153}
154
155mod sealed {
156    use crate::error::Error;
157
158    /// The primitive operations behind an [`HkdfVariant`](super::HkdfVariant),
159    /// private to dryoc.
160    pub trait Sealed<const PRK_LENGTH: usize> {
161        /// Minimum output length accepted by this variant.
162        const OUTPUT_BYTES_MIN: usize;
163        /// Maximum output length accepted by this variant.
164        const OUTPUT_BYTES_MAX: usize;
165
166        /// Creates a PRK from input keying material and optional salt.
167        fn extract(prk: &mut [u8; PRK_LENGTH], salt: Option<&[u8]>, ikm: &[u8]);
168        /// Expands a PRK into output keying material, failing if
169        /// `output.len()` is outside the range supported by this variant.
170        fn expand(output: &mut [u8], context: &[u8], prk: &[u8; PRK_LENGTH]) -> Result<(), Error>;
171
172        /// Validates an output length before allocating output storage,
173        /// failing if it is outside
174        /// `OUTPUT_BYTES_MIN..=OUTPUT_BYTES_MAX`.
175        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
190/// HKDF algorithm variant used by [`Hkdf`]: [`HkdfSha256Variant`] or
191/// [`HkdfSha512Variant`].
192///
193/// This trait is sealed so applications cannot plug in custom cryptographic
194/// algorithms; use it to write code that is generic over the provided
195/// variants.
196pub 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    /// Randomly generates a new PRK for HKDF expand.
245    #[must_use]
246    pub fn generate() -> Self {
247        Self {
248            prk: Prk::generate(),
249            _variant: PhantomData,
250        }
251    }
252
253    /// Extracts a PRK from input keying material and optional salt.
254    #[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    /// One-shot HKDF extract-and-expand into a fixed-size output type.
265    ///
266    /// # Errors
267    ///
268    /// Returns an error if `OUTPUT_LENGTH` is outside the range supported by
269    /// the selected HKDF variant.
270    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    /// One-shot HKDF extract-and-expand into a [`Vec`] of `output_len` bytes.
284    ///
285    /// # Errors
286    ///
287    /// Returns an error if `output_len` is outside the range supported by the
288    /// selected HKDF variant.
289    #[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    /// One-shot HKDF extract-and-expand into a runtime-sized byte container
300    /// of `output_len` bytes, such as protected memory.
301    ///
302    /// # Errors
303    ///
304    /// Returns an error if `output_len` is outside the range supported by the
305    /// selected HKDF variant.
306    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    /// Constructs an HKDF expander from a PRK, consuming it.
326    #[must_use]
327    pub fn from_prk(prk: Prk) -> Self {
328        Self {
329            prk,
330            _variant: PhantomData,
331        }
332    }
333
334    /// Moves the PRK out of this expander.
335    #[must_use]
336    pub fn into_prk(self) -> Prk {
337        self.prk
338    }
339
340    /// Expands this PRK into a fixed-size output type.
341    ///
342    /// # Errors
343    ///
344    /// Returns an error if `OUTPUT_LENGTH` is outside the range supported by
345    /// the selected HKDF variant.
346    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    /// Expands this PRK into a [`Vec`] of `output_len` bytes.
364    ///
365    /// # Errors
366    ///
367    /// Returns an error if `output_len` is outside the range supported by the
368    /// selected HKDF variant.
369    #[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    /// Expands this PRK into a runtime-sized byte container of `output_len`
379    /// bytes, such as protected memory.
380    ///
381    /// # Errors
382    ///
383    /// Returns an error if `output_len` is outside the range supported by the
384    /// selected HKDF variant.
385    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    /// RFC 5869 appendix A: `(salt, ikm, info, PRK, OKM)`.
408    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    /// A.1 (basic) and A.3 (no salt, no info) for SHA-256.
417    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    /// A.1 inputs with HKDF-SHA-512. RFC 5869 publishes SHA-256 and SHA-1
443    /// answers only; this OKM is the "OpenSSL-derived" HKDF-SHA512 vector for
444    /// the same inputs in OpenSSL's
445    /// `test/recipes/30-test_evp_data/evpkdf_hkdf.txt`, and the PRK is the
446    /// extract step's value on the way to it. `classic::crypto_kdf` checks the
447    /// same PRK and OKM.
448    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        // The PRK round-trips through `into_prk`/`from_prk`.
514        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        // A missing salt is the all-zero salt of one hash length.
524        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        // Shorter outputs are prefixes; a different context is a different key.
535        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        // SHA-512 allows twice as much output; the SHA-256 maximum is valid
604        // there.
605        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        // Same inputs, different hash: the variants must not collapse.
691        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}