@@ -241,7 +241,7 @@ namespace hmac_cpp {
241241 return dk;
242242 }
243243
244- std::vector <uint8_t > hkdf_extract_sha256 (
244+ secure_buffer <uint8_t , true > hkdf_extract_sha256_secure (
245245 const void * ikm_ptr, size_t ikm_len,
246246 const void * salt_ptr, size_t salt_len) {
247247 std::vector<uint8_t > salt_buf;
@@ -250,11 +250,19 @@ namespace hmac_cpp {
250250 salt_ptr = salt_buf.data ();
251251 salt_len = salt_buf.size ();
252252 }
253- auto prk = get_hmac (salt_ptr, salt_len, ikm_ptr, ikm_len, TypeHash::SHA256 );
253+ secure_buffer<uint8_t , true > prk (
254+ get_hmac (salt_ptr, salt_len, ikm_ptr, ikm_len, TypeHash::SHA256 ));
254255 return prk;
255256 }
256257
257- std::vector<uint8_t > hkdf_expand_sha256 (
258+ std::vector<uint8_t > hkdf_extract_sha256 (
259+ const void * ikm_ptr, size_t ikm_len,
260+ const void * salt_ptr, size_t salt_len) {
261+ auto prk = hkdf_extract_sha256_secure (ikm_ptr, ikm_len, salt_ptr, salt_len);
262+ return std::vector<uint8_t >(prk.begin (), prk.end ());
263+ }
264+
265+ secure_buffer<uint8_t , true > hkdf_expand_sha256_secure (
258266 const void * prk_ptr, size_t prk_len,
259267 const void * info_ptr, size_t info_len,
260268 size_t L) {
@@ -264,34 +272,46 @@ namespace hmac_cpp {
264272 if (L > 255 * HashLen)
265273 throw std::invalid_argument (" HKDF: L too large" );
266274
267- std::vector<uint8_t > okm;
268- okm.reserve (L);
269- std::vector<uint8_t > previous;
275+ secure_buffer<uint8_t , true > okm (L);
276+ secure_buffer<uint8_t , true > previous;
270277 size_t n = (L + HashLen - 1 ) / HashLen;
278+ size_t offset = 0 ;
271279 for (size_t i = 1 ; i <= n; ++i) {
272- std::vector<uint8_t > input (previous.begin (), previous.end ());
273- if (info_ptr && info_len)
274- input.insert (input.end (),
275- reinterpret_cast <const uint8_t *>(info_ptr),
276- reinterpret_cast <const uint8_t *>(info_ptr) + info_len);
277- input.push_back (static_cast <uint8_t >(i));
278- auto t = get_hmac (prk_ptr, prk_len, input.data (), input.size (), TypeHash::SHA256 );
279- size_t take = (i == n) ? (L - okm.size ()) : t.size ();
280- okm.insert (okm.end (), t.begin (), t.begin () + take);
281- previous.assign (t.begin (), t.end ());
280+ size_t info_bytes = (info_ptr && info_len) ? info_len : 0 ;
281+ size_t input_len = previous.size () + info_bytes + 1 ;
282+ secure_buffer<uint8_t , true > input (input_len);
283+ if (previous.size ())
284+ std::memcpy (input.data (), previous.data (), previous.size ());
285+ if (info_bytes)
286+ std::memcpy (input.data () + previous.size (), info_ptr, info_len);
287+ input[input_len - 1 ] = static_cast <uint8_t >(i);
288+ secure_buffer<uint8_t , true > t (
289+ get_hmac (prk_ptr, prk_len, input.data (), input.size (), TypeHash::SHA256 ));
290+ size_t take = (i == n) ? (L - offset) : t.size ();
291+ std::memcpy (okm.data () + offset, t.data (), take);
292+ offset += take;
293+ previous = t;
282294 secure_zero (t.data (), t.size ());
283295 secure_zero (input.data (), input.size ());
284296 }
285297 secure_zero (previous.data (), previous.size ());
286298 return okm;
287299 }
288300
301+ std::vector<uint8_t > hkdf_expand_sha256 (
302+ const void * prk_ptr, size_t prk_len,
303+ const void * info_ptr, size_t info_len,
304+ size_t L) {
305+ auto okm = hkdf_expand_sha256_secure (prk_ptr, prk_len, info_ptr, info_len, L);
306+ return std::vector<uint8_t >(okm.begin (), okm.end ());
307+ }
308+
289309 KeyIv hkdf_key_iv_256 (const void * ikm_ptr, size_t ikm_len,
290310 const void * salt_ptr, size_t salt_len,
291311 const std::string& context) {
292- auto prk = hkdf_extract_sha256 (ikm_ptr, ikm_len, salt_ptr, salt_len);
293- auto okm = hkdf_expand_sha256 (prk.data (), prk.size (),
294- context.data (), context.size (), 44 );
312+ auto prk = hkdf_extract_sha256_secure (ikm_ptr, ikm_len, salt_ptr, salt_len);
313+ auto okm = hkdf_expand_sha256_secure (prk.data (), prk.size (),
314+ context.data (), context.size (), 44 );
295315 KeyIv out{};
296316 std::copy (okm.begin (), okm.begin () + 32 , out.key .begin ());
297317 std::copy (okm.begin () + 32 , okm.begin () + 44 , out.iv .begin ());
0 commit comments