@@ -86,6 +86,126 @@ namespace hmac_cpp {
8686 }
8787 }
8888
89+ void HmacContext::init (const void * key_ptr, size_t key_len) {
90+ if (key_len > 0 && key_ptr == nullptr )
91+ throw std::invalid_argument (" Null key with non-zero length" );
92+
93+ switch (type_) {
94+ case TypeHash::SHA1 :
95+ block_size_ = hmac_hash::SHA1 ::BLOCK_SIZE ;
96+ digest_size_ = hmac_hash::SHA1 ::DIGEST_SIZE ;
97+ break ;
98+ case TypeHash::SHA256 :
99+ block_size_ = hmac_hash::SHA256 ::SHA224_256_BLOCK_SIZE ;
100+ digest_size_ = hmac_hash::SHA256 ::DIGEST_SIZE ;
101+ break ;
102+ case TypeHash::SHA512 :
103+ block_size_ = hmac_hash::SHA512 ::SHA384_512_BLOCK_SIZE ;
104+ digest_size_ = hmac_hash::SHA512 ::DIGEST_SIZE ;
105+ break ;
106+ default :
107+ throw std::invalid_argument (" Unsupported hash type" );
108+ }
109+
110+ secure_buffer<uint8_t > key (block_size_);
111+ if (key_len > block_size_) {
112+ auto hashed = get_hash (key_ptr, key_len, type_);
113+ std::copy (hashed.begin (), hashed.end (), key.begin ());
114+ if (hashed.size () < block_size_)
115+ std::fill (key.begin () + hashed.size (), key.end (), 0 );
116+ secure_zero (hashed.data (), hashed.size ());
117+ } else {
118+ if (key_len > 0 )
119+ std::memcpy (key.data (), key_ptr, key_len);
120+ if (key_len < block_size_)
121+ std::fill (key.begin () + key_len, key.end (), 0 );
122+ }
123+
124+ okeypad_ = secure_buffer<uint8_t >(block_size_);
125+ secure_buffer<uint8_t > ipad (block_size_);
126+ for (size_t i = 0 ; i < block_size_; ++i) {
127+ const uint8_t k = key[i];
128+ ipad[i] = k ^ 0x36 ;
129+ okeypad_[i] = k ^ 0x5c ;
130+ }
131+
132+ switch (type_) {
133+ case TypeHash::SHA1 :
134+ sha1_.init ();
135+ sha1_.update (ipad.data (), block_size_);
136+ break ;
137+ case TypeHash::SHA256 :
138+ sha256_.init ();
139+ sha256_.update (ipad.data (), block_size_);
140+ break ;
141+ case TypeHash::SHA512 :
142+ sha512_.init ();
143+ sha512_.update (ipad.data (), block_size_);
144+ break ;
145+ default :
146+ throw std::invalid_argument (" Unsupported hash type" );
147+ }
148+
149+ secure_zero (key.data (), key.size ());
150+ secure_zero (ipad.data (), ipad.size ());
151+ }
152+
153+ void HmacContext::update (const void * data_ptr, size_t data_len) {
154+ if (data_len > 0 && data_ptr == nullptr )
155+ throw std::invalid_argument (" Null data pointer with non-zero length" );
156+ const uint8_t * p = static_cast <const uint8_t *>(data_ptr);
157+ switch (type_) {
158+ case TypeHash::SHA1 :
159+ sha1_.update (p, data_len);
160+ break ;
161+ case TypeHash::SHA256 :
162+ sha256_.update (p, data_len);
163+ break ;
164+ case TypeHash::SHA512 :
165+ sha512_.update (p, data_len);
166+ break ;
167+ default :
168+ throw std::invalid_argument (" Unsupported hash type" );
169+ }
170+ }
171+
172+ void HmacContext::final (uint8_t * out_ptr, size_t out_len) {
173+ if (out_ptr == nullptr )
174+ throw std::invalid_argument (" Null output pointer" );
175+ if (out_len < digest_size_)
176+ throw std::invalid_argument (" Output buffer too small" );
177+
178+ secure_buffer<uint8_t > inner (digest_size_);
179+
180+ switch (type_) {
181+ case TypeHash::SHA1 :
182+ sha1_.finish (inner.data ());
183+ sha1_.init ();
184+ sha1_.update (okeypad_.data (), block_size_);
185+ sha1_.update (inner.data (), digest_size_);
186+ sha1_.finish (out_ptr);
187+ break ;
188+ case TypeHash::SHA256 :
189+ sha256_.finish (inner.data ());
190+ sha256_.init ();
191+ sha256_.update (okeypad_.data (), block_size_);
192+ sha256_.update (inner.data (), digest_size_);
193+ sha256_.finish (out_ptr);
194+ break ;
195+ case TypeHash::SHA512 :
196+ sha512_.finish (inner.data ());
197+ sha512_.init ();
198+ sha512_.update (okeypad_.data (), block_size_);
199+ sha512_.update (inner.data (), digest_size_);
200+ sha512_.finish (out_ptr);
201+ break ;
202+ default :
203+ throw std::invalid_argument (" Unsupported hash type" );
204+ }
205+
206+ secure_zero (inner.data (), inner.size ());
207+ }
208+
89209 std::vector<uint8_t > get_hmac (const void * key_ptr, size_t key_len, const void * msg_ptr, size_t msg_len, TypeHash type) {
90210 if ((key_len > 0 && key_ptr == nullptr ) || (msg_len > 0 && msg_ptr == nullptr ))
91211 throw std::invalid_argument (" Null pointer with non-zero length" );
0 commit comments