diff --git a/components/mbedtls/port/psa_driver/esp_sha/core/psa_crypto_driver_esp_sha1.c b/components/mbedtls/port/psa_driver/esp_sha/core/psa_crypto_driver_esp_sha1.c index 1a508dd0960..3f384ac729d 100644 --- a/components/mbedtls/port/psa_driver/esp_sha/core/psa_crypto_driver_esp_sha1.c +++ b/components/mbedtls/port/psa_driver/esp_sha/core/psa_crypto_driver_esp_sha1.c @@ -24,9 +24,9 @@ static const unsigned char sha1_padding[64] = { 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 }; -psa_status_t esp_sha1_starts(esp_sha1_context *ctx) { +int esp_sha1_starts(esp_sha1_context *ctx) { memset(ctx, 0, sizeof(esp_sha1_context)); - return PSA_SUCCESS; + return ESP_OK; } static void esp_internal_sha1_block_process(esp_sha1_context *ctx, const uint8_t *data) @@ -49,7 +49,7 @@ static void esp_internal_sha_update_state(esp_sha1_context *ctx) } } -static int esp_sha1_update(esp_sha1_context *ctx, const unsigned char *input, size_t ilen) +int esp_sha1_update(esp_sha1_context *ctx, const unsigned char *input, size_t ilen) { size_t fill, left, len; uint32_t local_len = 0; @@ -120,7 +120,7 @@ static int esp_sha1_update(esp_sha1_context *ctx, const unsigned char *input, si return 0; } -static int esp_sha1_finish(esp_sha1_context *ctx, uint8_t *hash) +int esp_sha1_finish(esp_sha1_context *ctx, uint8_t *hash) { int ret = -1; uint32_t last, padn; diff --git a/components/mbedtls/port/psa_driver/esp_sha/include/psa_crypto_driver_esp_sha1.h b/components/mbedtls/port/psa_driver/esp_sha/include/psa_crypto_driver_esp_sha1.h index 9550b5c794f..b59a89b7b6e 100644 --- a/components/mbedtls/port/psa_driver/esp_sha/include/psa_crypto_driver_esp_sha1.h +++ b/components/mbedtls/port/psa_driver/esp_sha/include/psa_crypto_driver_esp_sha1.h @@ -25,7 +25,7 @@ psa_status_t esp_sha1_driver_compute( size_t hash_size, size_t *hash_length); -psa_status_t esp_sha1_starts(esp_sha1_context *ctx); +int esp_sha1_starts(esp_sha1_context *ctx); psa_status_t esp_sha1_driver_update( esp_sha1_context *ctx, diff --git a/components/mbedtls/port/psa_driver/esp_sha/parallel_engine/psa_crypto_driver_esp_sha1.c b/components/mbedtls/port/psa_driver/esp_sha/parallel_engine/psa_crypto_driver_esp_sha1.c index 56e0908c526..7d7da3c9f2c 100644 --- a/components/mbedtls/port/psa_driver/esp_sha/parallel_engine/psa_crypto_driver_esp_sha1.c +++ b/components/mbedtls/port/psa_driver/esp_sha/parallel_engine/psa_crypto_driver_esp_sha1.c @@ -41,7 +41,7 @@ static const unsigned char sha1_padding[64] = { 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 }; -psa_status_t esp_sha1_starts(esp_sha1_context *ctx) +int esp_sha1_starts(esp_sha1_context *ctx) { memset(ctx, 0, sizeof(esp_sha1_context)); ctx->total[0] = 0; @@ -56,10 +56,10 @@ psa_status_t esp_sha1_starts(esp_sha1_context *ctx) ctx->sha_state = ESP_SHA1_STATE_INIT; ctx->first_block = false; ctx->operation_mode = ESP_SHA_MODE_SOFTWARE; - return PSA_SUCCESS; + return ESP_OK; } -static void esp_sha1_software_process( esp_sha1_context *ctx, const unsigned char data[64] ) +void esp_sha1_software_process( esp_sha1_context *ctx, const unsigned char data[64] ) { uint32_t temp, W[16], A, B, C, D, E; @@ -242,7 +242,7 @@ static int esp_internal_sha1_parallel_engine_process( esp_sha1_context *ctx, con return 0; } -static int esp_sha1_update(esp_sha1_context *ctx, const unsigned char *input, size_t ilen) +int esp_sha1_update(esp_sha1_context *ctx, const unsigned char *input, size_t ilen) { int ret = -1; size_t fill; @@ -307,7 +307,7 @@ psa_status_t esp_sha1_driver_update( return PSA_SUCCESS; } -static int esp_sha1_finish(esp_sha1_context *ctx, uint8_t *output) +int esp_sha1_finish(esp_sha1_context *ctx, uint8_t *output) { int ret = -1; uint32_t last, padn; @@ -350,7 +350,6 @@ out: esp_sha_unlock_engine(SHA1); ctx->operation_mode = ESP_SHA_MODE_SOFTWARE; } - memset(ctx, 0, sizeof(esp_sha1_context)); return ret; } diff --git a/components/mbedtls/port/psa_driver/include/psa_crypto_driver_esp_sha.h b/components/mbedtls/port/psa_driver/include/psa_crypto_driver_esp_sha.h index 0f0e8881c8d..06ecba69506 100644 --- a/components/mbedtls/port/psa_driver/include/psa_crypto_driver_esp_sha.h +++ b/components/mbedtls/port/psa_driver/include/psa_crypto_driver_esp_sha.h @@ -64,6 +64,12 @@ psa_status_t esp_sha_hash_abort(esp_sha_hash_operation_t *operation); psa_status_t esp_sha_hash_clone( const esp_sha_hash_operation_t *source_operation, esp_sha_hash_operation_t *target_operation); + +void esp_sha1_software_process( esp_sha1_context *ctx, const unsigned char data[64] ); +int esp_sha1_starts(esp_sha1_context *ctx); +int esp_sha1_update(esp_sha1_context *ctx, const unsigned char *input, size_t ilen); +int esp_sha1_finish(esp_sha1_context *ctx, uint8_t *output); + #endif #ifdef __cplusplus diff --git a/components/wpa_supplicant/esp_supplicant/src/crypto/fastpbkdf2.c b/components/wpa_supplicant/esp_supplicant/src/crypto/fastpbkdf2.c index 352bcd9cdb8..4c470c13ad5 100644 --- a/components/wpa_supplicant/esp_supplicant/src/crypto/fastpbkdf2.c +++ b/components/wpa_supplicant/esp_supplicant/src/crypto/fastpbkdf2.c @@ -29,6 +29,8 @@ #include "mbedtls/esp_config.h" #include "utils/wpa_debug.h" #include "psa/crypto.h" +#define ESP_SHA_DRIVER_ENABLED +#include "psa_crypto_driver_esp_sha.h" /* --- MSVC doesn't support C99 --- */ #ifdef _MSC_VER @@ -41,61 +43,344 @@ #define MIN(a, b) ((a) > (b)) ? (b) : (a) #endif +static inline void write32_be(uint32_t n, uint8_t out[4]) +{ +#if defined(__GNUC__) && __GNUC__ >= 4 && __BYTE_ORDER == __LITTLE_ENDIAN + *(uint32_t *)(out) = __builtin_bswap32(n); +#else + out[0] = (n >> 24) & 0xff; + out[1] = (n >> 16) & 0xff; + out[2] = (n >> 8) & 0xff; + out[3] = n & 0xff; +#endif +} + +/* Prepare block (of blocksz bytes) to contain md padding denoting a msg-size + * message (in bytes). block has a prefix of used bytes. + * + * Message length is expressed in 32 bits (so suitable for sha1, sha256, sha512). */ +static inline void md_pad(uint8_t *block, size_t blocksz, size_t used, size_t msg) +{ + memset(block + used, 0, blocksz - used - 4); + block[used] = 0x80; + block += blocksz - 4; + write32_be((uint32_t)(msg * 8), block); +} + +/* Internal function/type names for hash-specific things. */ +#define HMAC_CTX(_name) HMAC_ ## _name ## _ctx +#define HMAC_INIT(_name) HMAC_ ## _name ## _init +#define HMAC_UPDATE(_name) HMAC_ ## _name ## _update +#define HMAC_FINAL(_name) HMAC_ ## _name ## _final + +#define PBKDF2_F(_name) pbkdf2_f_ ## _name +#define PBKDF2(_name) pbkdf2_ ## _name + +/* This macro expands to decls for the whole implementation for a given + * hash function. Arguments are: + * + * _name like 'sha1', added to symbol names + * _blocksz block size, in bytes + * _hashsz digest output, in bytes + * _ctx hash context type + * _init hash context initialisation function + * args: (_ctx *c) + * _update hash context update function + * args: (_ctx *c, const void *data, size_t ndata) + * _final hash context finish function + * args: (_ctx *c, void *out) + * _xform hash context raw block update function + * args: (_ctx *c, const void *data) + * _xcpy hash context raw copy function (only need copy hash state) + * args: (_ctx * restrict out, const _ctx *restrict in) + * _xtract hash context state extraction + * args: args (_ctx *restrict c, uint8_t *restrict out) + * _xxor hash context xor function (only need xor hash state) + * args: (_ctx *restrict out, const _ctx *restrict in) + * + * The resulting function is named PBKDF2(_name). + */ +#define DECL_PBKDF2(_name, _blocksz, _hashsz, _ctx, \ + _init, _update, _xform, _final, _xcpy, _xtract, _xxor) \ + typedef struct { \ + _ctx inner; \ + _ctx outer; \ + } HMAC_CTX(_name); \ + \ + static inline void HMAC_INIT(_name)(HMAC_CTX(_name) *ctx, \ + const uint8_t *key, size_t nkey) \ + { \ + /* Prepare key: */ \ + uint8_t k[_blocksz] = {0}; \ + \ + /* Shorten long keys. */ \ + if (nkey > _blocksz) \ + { \ + _init(&ctx->inner); \ + _update(&ctx->inner, key, nkey); \ + _final(&ctx->inner, k); \ + \ + key = k; \ + nkey = _hashsz; \ + } \ + \ + /* Standard doesn't cover case where blocksz < hashsz. */ \ + assert(nkey <= _blocksz); \ + \ + /* Right zero-pad short keys. */ \ + if (k != key) \ + memcpy(k, key, nkey); \ + if (_blocksz > nkey) \ + memset(k + nkey, 0, _blocksz - nkey); \ + \ + /* Start inner hash computation */ \ + uint8_t blk_inner[_blocksz]; \ + uint8_t blk_outer[_blocksz]; \ + \ + for (size_t i = 0; i < _blocksz; i++) \ + { \ + blk_inner[i] = 0x36 ^ k[i]; \ + blk_outer[i] = 0x5c ^ k[i]; \ + } \ + \ + _init(&ctx->inner); \ + _update(&ctx->inner, blk_inner, sizeof blk_inner); \ + \ + /* And outer. */ \ + _init(&ctx->outer); \ + _update(&ctx->outer, blk_outer, sizeof blk_outer); \ + } \ + \ + static inline void HMAC_UPDATE(_name)(HMAC_CTX(_name) *ctx, \ + const void *data, size_t ndata) \ + { \ + _update(&ctx->inner, data, ndata); \ + } \ + \ + static inline void HMAC_FINAL(_name)(HMAC_CTX(_name) *ctx, \ + uint8_t out[_hashsz]) \ + { \ + _final(&ctx->inner, out); \ + _update(&ctx->outer, out, _hashsz); \ + _final(&ctx->outer, out); \ + } \ + \ + \ + /* --- PBKDF2 --- */ \ + static inline void PBKDF2_F(_name)(const HMAC_CTX(_name) *startctx, \ + uint32_t counter, \ + const uint8_t *salt, size_t nsalt, \ + uint32_t iterations, \ + uint8_t *out) \ + { \ + uint8_t countbuf[4]; \ + write32_be(counter, countbuf); \ + \ + /* Prepare loop-invariant padding block. */ \ + uint8_t Ublock[_blocksz]; \ + md_pad(Ublock, _blocksz, _hashsz, _blocksz + _hashsz); \ + \ + /* First iteration: \ + * U_1 = PRF(P, S || INT_32_BE(i)) \ + */ \ + HMAC_CTX(_name) ctx = *startctx; \ + HMAC_UPDATE(_name)(&ctx, salt, nsalt); \ + HMAC_UPDATE(_name)(&ctx, countbuf, sizeof countbuf); \ + HMAC_FINAL(_name)(&ctx, Ublock); \ + _ctx result = ctx.outer; \ + \ + /* Subsequent iterations: \ + * U_c = PRF(P, U_{c-1}) \ + */ \ + for (uint32_t i = 1; i < iterations; i++) \ + { \ + /* Complete inner hash with previous U */ \ + _xcpy(&ctx.inner, &startctx->inner); \ + _xform(&ctx.inner, Ublock); \ + _xtract(&ctx.inner, Ublock); \ + /* Complete outer hash with inner output */ \ + _xcpy(&ctx.outer, &startctx->outer); \ + _xform(&ctx.outer, Ublock); \ + _xtract(&ctx.outer, Ublock); \ + _xxor(&result, &ctx.outer); \ + } \ + \ + /* Reform result into output buffer. */ \ + _xtract(&result, out); \ + } \ + \ + static inline void PBKDF2(_name)(const uint8_t *pw, size_t npw, \ + const uint8_t *salt, size_t nsalt, \ + uint32_t iterations, \ + uint8_t *out, size_t nout) \ + { \ + assert(iterations); \ + assert(out && nout); \ + \ + /* Starting point for inner loop. */ \ + HMAC_CTX(_name) ctx; \ + HMAC_INIT(_name)(&ctx, pw, npw); \ + \ + /* How many blocks do we need? */ \ + uint32_t blocks_needed = (uint32_t)(nout + _hashsz - 1) / _hashsz; \ + \ + for (uint32_t counter = 1; counter <= blocks_needed; counter++) \ + { \ + uint8_t block[_hashsz]; \ + PBKDF2_F(_name)(&ctx, counter, salt, nsalt, iterations, block); \ + \ + size_t offset = (counter - 1) * _hashsz; \ + size_t taken = MIN(nout - offset, _hashsz); \ + memcpy(out + offset, block, taken); \ + } \ + } + +static inline void sha1_extract(esp_sha1_context *restrict ctx, uint8_t *restrict out) +{ +#if defined(MBEDTLS_PSA_ACCEL_ALG_SHA_1) +#if CONFIG_IDF_TARGET_ESP32 + /* ESP32 stores internal SHA state in BE format similar to software */ + write32_be(ctx->state[0], out); + write32_be(ctx->state[1], out + 4); + write32_be(ctx->state[2], out + 8); + write32_be(ctx->state[3], out + 12); + write32_be(ctx->state[4], out + 16); +#else + *(uint32_t *)(out) = ctx->state[0]; + *(uint32_t *)(out + 4) = ctx->state[1]; + *(uint32_t *)(out + 8) = ctx->state[2]; + *(uint32_t *)(out + 12) = ctx->state[3]; + *(uint32_t *)(out + 16) = ctx->state[4]; +#endif +#else + write32_be(ctx->MBEDTLS_PRIVATE(state)[0], out); + write32_be(ctx->MBEDTLS_PRIVATE(state)[1], out + 4); + write32_be(ctx->MBEDTLS_PRIVATE(state)[2], out + 8); + write32_be(ctx->MBEDTLS_PRIVATE(state)[3], out + 12); + write32_be(ctx->MBEDTLS_PRIVATE(state)[4], out + 16); +#endif +} + +static inline void sha1_cpy(esp_sha1_context *restrict out, const esp_sha1_context *restrict in) +{ +#if defined(MBEDTLS_PSA_ACCEL_ALG_SHA_1) + out->state[0] = in->state[0]; + out->state[1] = in->state[1]; + out->state[2] = in->state[2]; + out->state[3] = in->state[3]; + out->state[4] = in->state[4]; +#else + out->MBEDTLS_PRIVATE(state)[0] = in->MBEDTLS_PRIVATE(state)[0]; + out->MBEDTLS_PRIVATE(state)[1] = in->MBEDTLS_PRIVATE(state)[1]; + out->MBEDTLS_PRIVATE(state)[2] = in->MBEDTLS_PRIVATE(state)[2]; + out->MBEDTLS_PRIVATE(state)[3] = in->MBEDTLS_PRIVATE(state)[3]; + out->MBEDTLS_PRIVATE(state)[4] = in->MBEDTLS_PRIVATE(state)[4]; +#endif +} + +static inline void sha1_xor(esp_sha1_context *restrict out, const esp_sha1_context *restrict in) +{ +#if defined(MBEDTLS_PSA_ACCEL_ALG_SHA_1) + out->state[0] ^= in->state[0]; + out->state[1] ^= in->state[1]; + out->state[2] ^= in->state[2]; + out->state[3] ^= in->state[3]; + out->state[4] ^= in->state[4]; +#else + out->MBEDTLS_PRIVATE(state)[0] ^= in->MBEDTLS_PRIVATE(state)[0]; + out->MBEDTLS_PRIVATE(state)[1] ^= in->MBEDTLS_PRIVATE(state)[1]; + out->MBEDTLS_PRIVATE(state)[2] ^= in->MBEDTLS_PRIVATE(state)[2]; + out->MBEDTLS_PRIVATE(state)[3] ^= in->MBEDTLS_PRIVATE(state)[3]; + out->MBEDTLS_PRIVATE(state)[4] ^= in->MBEDTLS_PRIVATE(state)[4]; +#endif +} + +static int esp_sha1_init_start(esp_sha1_context *ctx) +{ + esp_sha1_starts(ctx); +#if defined(CONFIG_IDF_TARGET_ESP32) && defined(MBEDTLS_PSA_ACCEL_ALG_SHA_1) + /* Use software mode for esp32 since hardware can't give output more than 20 */ + // esp_mbedtls_set_sha1_mode(ctx, ESP_MBEDTLS_SHA1_SOFTWARE); + ctx->operation_mode = ESP_SHA_MODE_SOFTWARE; + ctx->sha_state = ESP_SHA1_STATE_IN_PROCESS; +#endif + return 0; +} + +#ifndef MBEDTLS_PSA_ACCEL_ALG_SHA_1 +static int sha1_finish(esp_sha1_context *ctx, + unsigned char output[20]) +{ + int ret = -1; + uint32_t used; + uint32_t high, low; + + /* + * Add padding: 0x80 then 0x00 until 8 bytes remain for the length + */ + used = ctx->MBEDTLS_PRIVATE(total)[0] & 0x3F; + + ctx->MBEDTLS_PRIVATE(buffer)[used++] = 0x80; + + if (used <= 56) { + /* Enough room for padding + length in current block */ + memset(ctx->MBEDTLS_PRIVATE(buffer) + used, 0, 56 - used); + } else { + /* We'll need an extra block */ + memset(ctx->MBEDTLS_PRIVATE(buffer) + used, 0, 64 - used); + + esp_sha1_software_process(ctx, ctx->MBEDTLS_PRIVATE(buffer)); + + memset(ctx->MBEDTLS_PRIVATE(buffer), 0, 56); + } + + /* + * Add message length + */ + high = (ctx->MBEDTLS_PRIVATE(total)[0] >> 29) + | (ctx->MBEDTLS_PRIVATE(total)[1] << 3); + low = (ctx->MBEDTLS_PRIVATE(total)[0] << 3); + + write32_be(high, ctx->MBEDTLS_PRIVATE(buffer) + 56); + write32_be(low, ctx->MBEDTLS_PRIVATE(buffer) + 60); + + esp_sha1_software_process(ctx, ctx->MBEDTLS_PRIVATE(buffer)); + + /* + * Output final state + */ + write32_be(ctx->MBEDTLS_PRIVATE(state)[0], output); + write32_be(ctx->MBEDTLS_PRIVATE(state)[1], output + 4); + write32_be(ctx->MBEDTLS_PRIVATE(state)[2], output + 8); + write32_be(ctx->MBEDTLS_PRIVATE(state)[3], output + 12); + write32_be(ctx->MBEDTLS_PRIVATE(state)[4], output + 16); + + ret = 0; + + return ret; +} +#endif + +DECL_PBKDF2(sha1, // _name + 64, // _blocksz + 20, // _hashsz + esp_sha1_context, // _ctx + esp_sha1_init_start, // _init + esp_sha1_update, // _update + esp_sha1_software_process, // _xform +#if defined(MBEDTLS_PSA_ACCEL_ALG_SHA_1) + esp_sha1_finish, // _final +#else + sha1_finish, // _final +#endif + sha1_cpy, // _xcpy + sha1_extract, // _xtract + sha1_xor) // _xxor + void fastpbkdf2_hmac_sha1(const uint8_t *pw, size_t npw, const uint8_t *salt, size_t nsalt, uint32_t iterations, uint8_t *out, size_t nout) { - psa_status_t status; - psa_key_derivation_operation_t operation = PSA_KEY_DERIVATION_OPERATION_INIT; - psa_key_attributes_t attributes = PSA_KEY_ATTRIBUTES_INIT; - psa_key_id_t key_id = 0; - - // Set up key attributes for password - psa_set_key_usage_flags(&attributes, PSA_KEY_USAGE_DERIVE); - psa_set_key_algorithm(&attributes, PSA_ALG_PBKDF2_HMAC(PSA_ALG_SHA_1)); - psa_set_key_type(&attributes, PSA_KEY_TYPE_PASSWORD); - - // Import password as key - status = psa_import_key(&attributes, pw, npw, &key_id); - if (status != PSA_SUCCESS) { - psa_reset_key_attributes(&attributes); - return; - } - - // Set up key derivation - status = psa_key_derivation_setup(&operation, PSA_ALG_PBKDF2_HMAC(PSA_ALG_SHA_1)); - if (status != PSA_SUCCESS) { - goto cleanup; - } - - // Set iteration count - status = psa_key_derivation_input_integer(&operation, PSA_KEY_DERIVATION_INPUT_COST, - iterations); - if (status != PSA_SUCCESS) { - goto cleanup; - } - - // Add salt - status = psa_key_derivation_input_bytes(&operation, PSA_KEY_DERIVATION_INPUT_SALT, - salt, nsalt); - if (status != PSA_SUCCESS) { - goto cleanup; - } - - // Add password - status = psa_key_derivation_input_key(&operation, PSA_KEY_DERIVATION_INPUT_PASSWORD, key_id); - if (status != PSA_SUCCESS) { - goto cleanup; - } - - // Generate output - status = psa_key_derivation_output_bytes(&operation, out, nout); - -cleanup: - psa_key_derivation_abort(&operation); - if (key_id) { - psa_destroy_key(key_id); - } - psa_reset_key_attributes(&attributes); + PBKDF2(sha1)(pw, npw, salt, nsalt, iterations, out, nout); } diff --git a/components/wpa_supplicant/esp_supplicant/src/crypto/fastpsk.c b/components/wpa_supplicant/esp_supplicant/src/crypto/fastpsk.c index dc9c18c5e6f..65f9b3eadbf 100644 --- a/components/wpa_supplicant/esp_supplicant/src/crypto/fastpsk.c +++ b/components/wpa_supplicant/esp_supplicant/src/crypto/fastpsk.c @@ -100,6 +100,26 @@ struct fast_psk_context { uint32_t sum[SHA1_OUTPUT_SZ_WORDS]; /* Intermediate hash result */ }; +/* Acquire SHA1 hardware for exclusive use */ +static inline void sha1_setup(void) +{ +#if SOC_SHA_SUPPORT_PARALLEL_ENG + esp_sha_lock_engine(SHA1); +#else + esp_sha_acquire_hardware(); +#endif +} + +/* Release SHA1 hardware */ +static inline void sha1_teardown(void) +{ +#if SOC_SHA_SUPPORT_PARALLEL_ENG + esp_sha_unlock_engine(SHA1); +#else + esp_sha_release_hardware(); +#endif +} + /* * Pads the given HMAC block context with the appropriate SHA1 padding. * Length is the number of bytes of actual data in the block. @@ -141,49 +161,13 @@ static inline void write32_be(uint32_t n, uint8_t out[4]) void sha1_op(uint32_t blocks[FAST_PSK_SHA1_BLOCKS_BUF_WORDS], uint32_t output[SHA1_OUTPUT_SZ_WORDS]) { - psa_status_t status; - psa_hash_operation_t operation = PSA_HASH_OPERATION_INIT; - - // Initialize output to zero in case of error - memset(output, 0, SHA1_OUTPUT_SZ_WORDS * sizeof(uint32_t)); - - status = psa_hash_setup(&operation, PSA_ALG_SHA_1); - if (status != PSA_SUCCESS) { - ESP_LOGE("fastpsk", "psa_hash_setup failed: %d", status); - return; - } - - // Update with the first block - status = psa_hash_update(&operation, (const uint8_t *)blocks, SHA1_BLOCK_SZ); - if (status != PSA_SUCCESS) { - ESP_LOGE("fastpsk", "psa_hash_update failed: %d", status); - psa_hash_abort(&operation); - return; - } - - // Update with the second block - status = psa_hash_update(&operation, (const uint8_t *)&blocks[SHA1_BLOCK_SZ_WORDS], SHA1_BLOCK_SZ); - if (status != PSA_SUCCESS) { - ESP_LOGE("fastpsk", "psa_hash_update failed: %d", status); - psa_hash_abort(&operation); - return; - } - - // Finish the hash operation - size_t mac_len; - status = psa_hash_finish(&operation, (uint8_t *)output, SHA1_OUTPUT_SZ, &mac_len); - if (status != PSA_SUCCESS) { - ESP_LOGE("fastpsk", "psa_hash_finish failed: %d", status); - memset(output, 0, SHA1_OUTPUT_SZ_WORDS * sizeof(uint32_t)); - return; - } - - // Ensure the output length is correct - if (mac_len != SHA1_OUTPUT_SZ) { - ESP_LOGE("fastpsk", "Unexpected hash length: %zu, expected: %d", mac_len, SHA1_OUTPUT_SZ); - memset(output, 0, SHA1_OUTPUT_SZ_WORDS * sizeof(uint32_t)); - return; - } + esp_sha_set_mode(SHA1); + /* First block */ + esp_sha_block(SHA1, blocks, true); + /* Second block */ + esp_sha_block(SHA1, &blocks[SHA1_BLOCK_SZ_WORDS], false); + /* Read the final digest */ + esp_sha_read_digest_state(SHA1, output); #if CONFIG_IDF_TARGET_ESP32 for (int i = 0; i < SHA1_OUTPUT_SZ_WORDS; i++) { @@ -227,6 +211,8 @@ void fast_psk_f(const char *password, size_t password_len, const uint8_t *ssid, /* Pad the block */ pad_blocks(&ctx->inner, SHA1_BLOCK_SZ + ssid_len + 4); + sha1_setup(); + uint32_t *pi, *po; pi = ctx->inner.whole_words; po = ctx->outer.whole_words; @@ -260,6 +246,8 @@ void fast_psk_f(const char *password, size_t password_len, const uint8_t *ssid, } } + sha1_teardown(); + /* Copy the final result to the output digest */ memcpy(digest, sum, SHA1_OUTPUT_SZ); @@ -269,50 +257,14 @@ void fast_psk_f(const char *password, size_t password_len, const uint8_t *ssid, int esp_fast_psk(const char *password, size_t password_len, const uint8_t *ssid, size_t ssid_len, size_t iterations, uint8_t *output, size_t output_len) { - if (!(ssid_len <= 32 && password_len <= 63 && iterations == 4096 && output_len == 32)) { - return -1; /* Invalid input parameters */ - } + /* Compute the first 16 bytes of the PSK */ + fast_psk_f(password, password_len, ssid, ssid_len, 2, output); - /* Compute the full PSK */ - psa_status_t status; - psa_key_derivation_operation_t operation = PSA_KEY_DERIVATION_OPERATION_INIT; - int ret = -1; /* Track error status */ + /* Replicate the first 16 bytes to form the second half temporarily */ + memcpy(output + SHA1_OUTPUT_SZ, output, 32 - SHA1_OUTPUT_SZ); - // Set up key derivation - status = psa_key_derivation_setup(&operation, PSA_ALG_PBKDF2_HMAC(PSA_ALG_SHA_1)); - if (status != PSA_SUCCESS) { - goto cleanup; - } + /* Compute the second 16 bytes of the PSK */ + fast_psk_f(password, password_len, ssid, ssid_len, 1, output); - // Set iteration count - status = psa_key_derivation_input_integer(&operation, PSA_KEY_DERIVATION_INPUT_COST, iterations); - if (status != PSA_SUCCESS) { - goto cleanup; - } - - // Add salt - status = psa_key_derivation_input_bytes(&operation, PSA_KEY_DERIVATION_INPUT_SALT, - ssid, ssid_len); - if (status != PSA_SUCCESS) { - goto cleanup; - } - - // Add password - status = psa_key_derivation_input_bytes(&operation, PSA_KEY_DERIVATION_INPUT_PASSWORD, - (const uint8_t*)password, password_len); - if (status != PSA_SUCCESS) { - goto cleanup; - } - - // Generate output - status = psa_key_derivation_output_bytes(&operation, output, output_len); - if (status != PSA_SUCCESS) { - goto cleanup; - } - - ret = 0; /* Success */ - -cleanup: - psa_key_derivation_abort(&operation); - return ret; + return 0; /* Success */ }