1use crate::{HmacAlgorithm, HmacProvider};
5
6#[derive(Debug, PartialEq, Eq)]
7pub enum HkdfError {
8 OutputTooLong,
10}
11
12pub trait HkdfProvider: HmacProvider {
26 fn hkdf_extract(
31 &mut self,
32 alg: <Self as HmacProvider>::Algorithm,
33 salt: Option<&[u8]>,
34 ikm: &[u8],
35 ) -> Result<impl AsRef<[u8]> + use<Self>, HkdfError>;
36
37 fn hkdf_expand(
39 &mut self,
40 alg: <Self as HmacProvider>::Algorithm,
41 prk: &[u8],
42 info: &[u8],
43 okm: &mut [u8],
44 ) -> Result<(), HkdfError>;
45
46 fn hkdf(
48 &mut self,
49 alg: <Self as HmacProvider>::Algorithm,
50 salt: Option<&[u8]>,
51 ikm: &[u8],
52 info: &[u8],
53 okm: &mut [u8],
54 ) -> Result<(), HkdfError> {
55 let prk = self.hkdf_extract(alg.clone(), salt, ikm)?;
56 self.hkdf_expand(alg, prk.as_ref(), info, okm)
57 }
58}
59
60impl<H: HmacProvider> HkdfProvider for H {
61 fn hkdf_extract(
62 &mut self,
63 alg: <Self as HmacProvider>::Algorithm,
64 salt: Option<&[u8]>,
65 ikm: &[u8],
66 ) -> Result<impl AsRef<[u8]> + use<H>, HkdfError> {
67 let mut zero_salt = <<H as HmacProvider>::Algorithm as HmacAlgorithm>::MaxLenBuf::default();
71 let zero_salt = zero_salt.as_mut();
72 let hash_len = alg.len();
73 debug_assert!(
74 hash_len <= zero_salt.len(),
75 "algorithm length is longer than type's announced maximum HMAC length"
76 );
77 let salt_bytes = salt.unwrap_or(&zero_salt[..hash_len]);
78 Ok(self.hmac_with_keydata(alg, salt_bytes, ikm))
80 }
81
82 fn hkdf_expand(
83 &mut self,
84 alg: <Self as HmacProvider>::Algorithm,
85 prk: &[u8],
86 info: &[u8],
87 okm: &mut [u8],
88 ) -> Result<(), HkdfError> {
89 let hash_len = alg.len();
90 if okm.len() > 255 * hash_len {
91 return Err(HkdfError::OutputTooLong);
92 }
93 let mut t = <<H as HmacProvider>::Algorithm as HmacAlgorithm>::MaxLenBuf::default();
94 let t = t.as_mut();
95 debug_assert!(
96 hash_len <= t.len(),
97 "algorithm length is longer than type's announced maximum HMAC length"
98 );
99 let mut t_len = 0usize;
100 let mut pos = 0usize;
101
102 while pos < okm.len() {
103 let counter = (pos / hash_len + 1) as u8;
105 let mut state = self.init_with_keydata(alg.clone(), prk);
107 if t_len > 0 {
108 HmacProvider::update(self, &mut state, &t[..t_len]);
109 }
110 HmacProvider::update(self, &mut state, info);
111 HmacProvider::update(self, &mut state, &[counter]);
112 let result = HmacProvider::finalize(self, state);
113 let result_bytes = result.as_ref();
114 debug_assert_eq!(
115 result_bytes.len(),
116 hash_len,
117 "algorithm did not produce its announced fixed length as output"
118 );
119 t[..hash_len].copy_from_slice(result_bytes);
120 t_len = hash_len;
121
122 let take = (okm.len() - pos).min(hash_len);
123 okm[pos..pos + take].copy_from_slice(&t[..take]);
124 pos += take;
125 }
126 Ok(())
127 }
128}