diff --git a/src/wh_client_crypto.c b/src/wh_client_crypto.c index 93ef3eb9c..ddc5218ef 100644 --- a/src/wh_client_crypto.c +++ b/src/wh_client_crypto.c @@ -11410,6 +11410,7 @@ static int _MlKemMakeKeyDma(whClientContext* ctx, int level, uint8_t buffer[WC_ML_KEM_MAX_PRIVATE_KEY_SIZE]; uint16_t buffer_len = sizeof(buffer); + uint16_t res_len = 0; if ((ctx == NULL) || (key == NULL)) { return WH_ERROR_BADARGS; @@ -11465,7 +11466,7 @@ static int _MlKemMakeKeyDma(whClientContext* ctx, int level, } if (ret == WH_ERROR_OK) { do { - ret = wh_Client_RecvResponse(ctx, &group, &action, &req_len, + ret = wh_Client_RecvResponse(ctx, &group, &action, &res_len, WOLFHSM_CFG_COMM_DATA_LEN, (uint8_t*)dataPtr); } while (ret == WH_ERROR_NOTREADY); @@ -11479,28 +11480,33 @@ static int _MlKemMakeKeyDma(whClientContext* ctx, int level, ret = _getCryptoResponse(dataPtr, WC_PK_TYPE_PQC_KEM_KEYGEN, (uint8_t**)&res); if (ret >= 0) { - key_id = (whKeyId)res->keyId; - if (inout_key_id != NULL) { - *inout_key_id = key_id; + const uint32_t hdr_sz = + sizeof(whMessageCrypto_GenericResponseHeader) + sizeof(*res); + if (res_len < hdr_sz) { + ret = WH_ERROR_ABORTED; } - if (key != NULL) { - wh_Client_MlKemSetKeyId(key, key_id); - /* buffer holds the exported key (EPHEMERAL) or the public key - * (cached keygen); keySize bounds the DMA write. An empty result - * means the caller requested key material the server did not - * return. */ - if (res->keySize == 0) { - ret = WH_ERROR_ABORTED; - } - else if (res->keySize > buffer_len) { - ret = WH_ERROR_BADARGS; + else { + key_id = (whKeyId)res->keyId; + if (inout_key_id != NULL) { + *inout_key_id = key_id; } - else { - ret = wh_Crypto_MlKemDeserializeKey( - buffer, (uint16_t)res->keySize, key); + if (key != NULL) { + wh_Client_MlKemSetKeyId(key, key_id); + /* buffer holds the exported key (EPHEMERAL) or the public + * key (cached keygen); keySize bounds the DMA write. An + * empty result means the server returned no key material. */ + if (res->keySize == 0) { + ret = WH_ERROR_ABORTED; + } + else if (res->keySize > buffer_len) { + ret = WH_ERROR_BADARGS; + } + else { + ret = wh_Crypto_MlKemDeserializeKey( + buffer, (uint16_t)res->keySize, key); + } } } - } } @@ -11886,6 +11892,7 @@ int wh_Client_LmsMakeKeyDma(whClientContext* ctx, LmsKey* key, uint16_t action = WC_ALGO_TYPE_PK; uint16_t req_len = sizeof(whMessageCrypto_GenericRequestHeader) + sizeof(*req); + uint16_t res_len = 0; if (req_len > WOLFHSM_CFG_COMM_DATA_LEN) { return WH_ERROR_BADARGS; @@ -11921,7 +11928,7 @@ int wh_Client_LmsMakeKeyDma(whClientContext* ctx, LmsKey* key, } if (ret == WH_ERROR_OK) { do { - ret = wh_Client_RecvResponse(ctx, &group, &action, &req_len, + ret = wh_Client_RecvResponse(ctx, &group, &action, &res_len, WOLFHSM_CFG_COMM_DATA_LEN, (uint8_t*)dataPtr); } while (ret == WH_ERROR_NOTREADY); @@ -11936,11 +11943,19 @@ int wh_Client_LmsMakeKeyDma(whClientContext* ctx, LmsKey* key, WC_PK_TYPE_PQC_STATEFUL_SIG_KEYGEN, (uint8_t**)&res); if (ret >= 0) { - key_id = (whKeyId)res->keyId; - if (inout_key_id != NULL) { - *inout_key_id = key_id; + const uint32_t hdr_sz = + sizeof(whMessageCrypto_GenericResponseHeader) + + sizeof(*res); + if (res_len < hdr_sz) { + ret = WH_ERROR_ABORTED; + } + else { + key_id = (whKeyId)res->keyId; + if (inout_key_id != NULL) { + *inout_key_id = key_id; + } + wh_Client_LmsSetKeyId(key, key_id); } - wh_Client_LmsSetKeyId(key, key_id); } } @@ -11997,6 +12012,7 @@ int wh_Client_LmsSignDma(whClientContext* ctx, const byte* msg, word32 msgSz, uint16_t action = WC_ALGO_TYPE_PK; uint16_t req_len = sizeof(whMessageCrypto_GenericRequestHeader) + sizeof(*req); + uint16_t res_len = 0; if (req_len > WOLFHSM_CFG_COMM_DATA_LEN) { return WH_ERROR_BADARGS; @@ -12024,7 +12040,7 @@ int wh_Client_LmsSignDma(whClientContext* ctx, const byte* msg, word32 msgSz, } if (ret == WH_ERROR_OK) { do { - ret = wh_Client_RecvResponse(ctx, &group, &action, &req_len, + ret = wh_Client_RecvResponse(ctx, &group, &action, &res_len, WOLFHSM_CFG_COMM_DATA_LEN, (uint8_t*)dataPtr); } while (ret == WH_ERROR_NOTREADY); @@ -12041,7 +12057,13 @@ int wh_Client_LmsSignDma(whClientContext* ctx, const byte* msg, word32 msgSz, ret = _getCryptoResponse(dataPtr, WC_PK_TYPE_PQC_STATEFUL_SIG_SIGN, (uint8_t**)&res); if (ret >= 0) { - if (res->sigLen > sigCap) { + const uint32_t hdr_sz = + sizeof(whMessageCrypto_GenericResponseHeader) + + sizeof(*res); + if (res_len < hdr_sz) { + ret = WH_ERROR_ABORTED; + } + else if (res->sigLen > sigCap) { ret = WH_ERROR_BADARGS; } else { @@ -12098,6 +12120,7 @@ int wh_Client_LmsVerifyDma(whClientContext* ctx, const byte* sig, word32 sigSz, uint16_t action = WC_ALGO_TYPE_PK; uint16_t req_len = sizeof(whMessageCrypto_GenericRequestHeader) + sizeof(*req); + uint16_t res_len = 0; if (req_len > WOLFHSM_CFG_COMM_DATA_LEN) { return WH_ERROR_BADARGS; @@ -12124,7 +12147,7 @@ int wh_Client_LmsVerifyDma(whClientContext* ctx, const byte* sig, word32 sigSz, } if (ret == WH_ERROR_OK) { do { - ret = wh_Client_RecvResponse(ctx, &group, &action, &req_len, + ret = wh_Client_RecvResponse(ctx, &group, &action, &res_len, WOLFHSM_CFG_COMM_DATA_LEN, (uint8_t*)dataPtr); } while (ret == WH_ERROR_NOTREADY); @@ -12142,8 +12165,16 @@ int wh_Client_LmsVerifyDma(whClientContext* ctx, const byte* sig, word32 sigSz, WC_PK_TYPE_PQC_STATEFUL_SIG_VERIFY, (uint8_t**)&resp); if (ret >= 0) { - *res = (int)resp->res; - ret = WH_ERROR_OK; + const uint32_t hdr_sz = + sizeof(whMessageCrypto_GenericResponseHeader) + + sizeof(*resp); + if (res_len < hdr_sz) { + ret = WH_ERROR_ABORTED; + } + else { + *res = (int)resp->res; + ret = WH_ERROR_OK; + } } } } @@ -12183,6 +12214,7 @@ int wh_Client_LmsSigsLeftDma(whClientContext* ctx, LmsKey* key) uint16_t action = WC_ALGO_TYPE_PK; uint16_t req_len = sizeof(whMessageCrypto_GenericRequestHeader) + sizeof(*req); + uint16_t res_len = 0; if (req_len > WOLFHSM_CFG_COMM_DATA_LEN) { return WH_ERROR_BADARGS; @@ -12195,7 +12227,7 @@ int wh_Client_LmsSigsLeftDma(whClientContext* ctx, LmsKey* key) (uint8_t*)dataPtr); if (ret == WH_ERROR_OK) { do { - ret = wh_Client_RecvResponse(ctx, &group, &action, &req_len, + ret = wh_Client_RecvResponse(ctx, &group, &action, &res_len, WOLFHSM_CFG_COMM_DATA_LEN, (uint8_t*)dataPtr); } while (ret == WH_ERROR_NOTREADY); @@ -12205,9 +12237,17 @@ int wh_Client_LmsSigsLeftDma(whClientContext* ctx, LmsKey* key) WC_PK_TYPE_PQC_STATEFUL_SIG_SIGS_LEFT, (uint8_t**)&res); if (ret >= 0) { - /* The server mirrors wc_LmsKey_SigsLeft(), which is a - * boolean. Normalize so the only nonzero return is 1. */ - ret = (res->sigsLeft != 0) ? 1 : 0; + const uint32_t hdr_sz = + sizeof(whMessageCrypto_GenericResponseHeader) + + sizeof(*res); + if (res_len < hdr_sz) { + ret = WH_ERROR_ABORTED; + } + else { + /* The server mirrors wc_LmsKey_SigsLeft(), which is a + * boolean. Normalize so the only nonzero return is 1. */ + ret = (res->sigsLeft != 0) ? 1 : 0; + } } } } @@ -12323,6 +12363,7 @@ int wh_Client_XmssMakeKeyDma(whClientContext* ctx, XmssKey* key, uint16_t action = WC_ALGO_TYPE_PK; uint16_t req_len = sizeof(whMessageCrypto_GenericRequestHeader) + sizeof(*req); + uint16_t res_len = 0; if (req_len > WOLFHSM_CFG_COMM_DATA_LEN) { return WH_ERROR_BADARGS; @@ -12368,7 +12409,7 @@ int wh_Client_XmssMakeKeyDma(whClientContext* ctx, XmssKey* key, } if (ret == WH_ERROR_OK) { do { - ret = wh_Client_RecvResponse(ctx, &group, &action, &req_len, + ret = wh_Client_RecvResponse(ctx, &group, &action, &res_len, WOLFHSM_CFG_COMM_DATA_LEN, (uint8_t*)dataPtr); } while (ret == WH_ERROR_NOTREADY); @@ -12383,11 +12424,19 @@ int wh_Client_XmssMakeKeyDma(whClientContext* ctx, XmssKey* key, WC_PK_TYPE_PQC_STATEFUL_SIG_KEYGEN, (uint8_t**)&res); if (ret >= 0) { - key_id = (whKeyId)res->keyId; - if (inout_key_id != NULL) { - *inout_key_id = key_id; + const uint32_t hdr_sz = + sizeof(whMessageCrypto_GenericResponseHeader) + + sizeof(*res); + if (res_len < hdr_sz) { + ret = WH_ERROR_ABORTED; + } + else { + key_id = (whKeyId)res->keyId; + if (inout_key_id != NULL) { + *inout_key_id = key_id; + } + wh_Client_XmssSetKeyId(key, key_id); } - wh_Client_XmssSetKeyId(key, key_id); } } @@ -12444,6 +12493,7 @@ int wh_Client_XmssSignDma(whClientContext* ctx, const byte* msg, word32 msgSz, uint16_t action = WC_ALGO_TYPE_PK; uint16_t req_len = sizeof(whMessageCrypto_GenericRequestHeader) + sizeof(*req); + uint16_t res_len = 0; if (req_len > WOLFHSM_CFG_COMM_DATA_LEN) { return WH_ERROR_BADARGS; @@ -12471,7 +12521,7 @@ int wh_Client_XmssSignDma(whClientContext* ctx, const byte* msg, word32 msgSz, } if (ret == WH_ERROR_OK) { do { - ret = wh_Client_RecvResponse(ctx, &group, &action, &req_len, + ret = wh_Client_RecvResponse(ctx, &group, &action, &res_len, WOLFHSM_CFG_COMM_DATA_LEN, (uint8_t*)dataPtr); } while (ret == WH_ERROR_NOTREADY); @@ -12488,7 +12538,13 @@ int wh_Client_XmssSignDma(whClientContext* ctx, const byte* msg, word32 msgSz, ret = _getCryptoResponse(dataPtr, WC_PK_TYPE_PQC_STATEFUL_SIG_SIGN, (uint8_t**)&res); if (ret >= 0) { - if (res->sigLen > sigCap) { + const uint32_t hdr_sz = + sizeof(whMessageCrypto_GenericResponseHeader) + + sizeof(*res); + if (res_len < hdr_sz) { + ret = WH_ERROR_ABORTED; + } + else if (res->sigLen > sigCap) { ret = WH_ERROR_BADARGS; } else { @@ -12546,6 +12602,7 @@ int wh_Client_XmssVerifyDma(whClientContext* ctx, const byte* sig, uint16_t action = WC_ALGO_TYPE_PK; uint16_t req_len = sizeof(whMessageCrypto_GenericRequestHeader) + sizeof(*req); + uint16_t res_len = 0; if (req_len > WOLFHSM_CFG_COMM_DATA_LEN) { return WH_ERROR_BADARGS; @@ -12572,7 +12629,7 @@ int wh_Client_XmssVerifyDma(whClientContext* ctx, const byte* sig, } if (ret == WH_ERROR_OK) { do { - ret = wh_Client_RecvResponse(ctx, &group, &action, &req_len, + ret = wh_Client_RecvResponse(ctx, &group, &action, &res_len, WOLFHSM_CFG_COMM_DATA_LEN, (uint8_t*)dataPtr); } while (ret == WH_ERROR_NOTREADY); @@ -12590,8 +12647,16 @@ int wh_Client_XmssVerifyDma(whClientContext* ctx, const byte* sig, WC_PK_TYPE_PQC_STATEFUL_SIG_VERIFY, (uint8_t**)&resp); if (ret >= 0) { - *res = (int)resp->res; - ret = WH_ERROR_OK; + const uint32_t hdr_sz = + sizeof(whMessageCrypto_GenericResponseHeader) + + sizeof(*resp); + if (res_len < hdr_sz) { + ret = WH_ERROR_ABORTED; + } + else { + *res = (int)resp->res; + ret = WH_ERROR_OK; + } } } } @@ -12631,6 +12696,7 @@ int wh_Client_XmssSigsLeftDma(whClientContext* ctx, XmssKey* key) uint16_t action = WC_ALGO_TYPE_PK; uint16_t req_len = sizeof(whMessageCrypto_GenericRequestHeader) + sizeof(*req); + uint16_t res_len = 0; if (req_len > WOLFHSM_CFG_COMM_DATA_LEN) { return WH_ERROR_BADARGS; @@ -12643,7 +12709,7 @@ int wh_Client_XmssSigsLeftDma(whClientContext* ctx, XmssKey* key) (uint8_t*)dataPtr); if (ret == WH_ERROR_OK) { do { - ret = wh_Client_RecvResponse(ctx, &group, &action, &req_len, + ret = wh_Client_RecvResponse(ctx, &group, &action, &res_len, WOLFHSM_CFG_COMM_DATA_LEN, (uint8_t*)dataPtr); } while (ret == WH_ERROR_NOTREADY); @@ -12653,9 +12719,17 @@ int wh_Client_XmssSigsLeftDma(whClientContext* ctx, XmssKey* key) WC_PK_TYPE_PQC_STATEFUL_SIG_SIGS_LEFT, (uint8_t**)&res); if (ret >= 0) { - /* The server mirrors wc_XmssKey_SigsLeft(), which is a - * boolean. Normalize so the only nonzero return is 1. */ - ret = (res->sigsLeft != 0) ? 1 : 0; + const uint32_t hdr_sz = + sizeof(whMessageCrypto_GenericResponseHeader) + + sizeof(*res); + if (res_len < hdr_sz) { + ret = WH_ERROR_ABORTED; + } + else { + /* The server mirrors wc_XmssKey_SigsLeft(), which is a + * boolean. Normalize so the only nonzero return is 1. */ + ret = (res->sigsLeft != 0) ? 1 : 0; + } } } }